BatchManager

Contents

BatchManager#

class liesel.optim.BatchManager(batches, epoch_size='strict', *, sampling_weights=None)[source]#

Bases: object

Coordinates multiple Batches objects as one batching interface.

A BatchManager is useful when a model contains observed branches with different observation sizes. Each contained Batches object owns the slicing rules for one branch. The manager combines them into one joint batched position for every optimizer step.

Parameters:
  • batches (Sequence[Batches]) – Non-empty sequence of Batches objects. Their position_keys must not overlap.

  • epoch_size (Literal['strict', 'min', 'max'] | int, default: 'strict') – Epoch length policy: "strict", "min", "max", or a positive integer.

  • sampling_weights (InitVar, default: None) – Optional keyword-only weights: a vector for a single child, or a mapping from one child position key per group to its vector. Each vector applies to the entire group and requires that child’s sample_with_replacement=True. Unknown keys and multiple entries for one group are rejected. Supplied weights override existing child weights on a copy; omitted groups retain their existing sampling configuration.

Raises:

ValueError – If batches is empty, if any position_keys are claimed by more than one child, if epoch_size is invalid, or if strict sizing is used with unequal child Batches.n_full_batches.

Notes

The properties axis_size, batch_size, and batch_sample_scales return tuples in child-batch order. The scalar aliases are available only when all children have the same likelihood scale. With unequal scales, use scaled_log_lik() so each branch is scaled by its own sample-size ratio.

Use manual BatchManager([Batches(...)]) construction when child groups need custom per-branch sample_size or batch_sample_size values.

Like Batches, start_epoch() mutates and returns self.

Examples

Combine two equally long batch sequences:

>>> import jax.numpy as jnp
>>> from liesel.optim import BatchManager, Batches
>>> manager = BatchManager(
...     [
...         Batches(["x"], axis_size=6, batch_size=2, shuffle=False),
...         Batches(["y"], axis_size=9, batch_size=3, shuffle=False),
...     ]
... )
>>> manager.n_full_batches
3
>>> position = {"x": jnp.arange(6), "y": jnp.arange(9)}
>>> batched = manager.get_batched_position(position, 1)
>>> batched["x"].tolist(), batched["y"].tolist()
([2, 3], [3, 4, 5])

With epoch_size="max", shorter branches assemble additional shuffled passes:

>>> import jax
>>> manager = BatchManager(
...     [
...         Batches(["x"], axis_size=6, batch_size=2, shuffle=True),
...         Batches(["y"], axis_size=8, batch_size=4, shuffle=True),
...     ],
...     epoch_size="max",
... ).start_epoch(jax.random.key(0))
>>> manager.n_full_batches
3

Per-branch scaling agrees with a manual scaled log-likelihood calculation:

>>> import liesel.model as lsl
>>> import tensorflow_probability.substrates.jax.distributions as tfd
>>> y1 = lsl.Var.new_obs(
...     jnp.arange(6.0),
...     lsl.Dist(tfd.Normal, loc=0.0, scale=1.0),
...     name="y1",
... )
>>> y2 = lsl.Var.new_obs(
...     jnp.arange(8.0),
...     lsl.Dist(tfd.Normal, loc=0.0, scale=1.0),
...     name="y2",
... )
>>> model = lsl.Model([y1, y2])
>>> manager = BatchManager(
...     [
...         Batches(["y1"], axis_size=6, batch_size=2, shuffle=True),
...         Batches(["y2"], axis_size=8, batch_size=4, shuffle=True),
...     ],
...     epoch_size="max",
... )
>>> batch = manager.get_batched_position(model.extract_position(["y1", "y2"]), 0)
>>> state = model.update_state(batch, model.state)
>>> manual = (
...     3.0 * state["y1_log_prob"].value.sum()
...     + 2.0 * state["y2_log_prob"].value.sum()
... )
>>> bool(jnp.allclose(manager.scaled_log_lik(model, state), manual))
True

Methods

correction_factors(batch_index)

Return extra per-index correction factors for each child, in order.

extract_batched_position(interface, ...)

Extracts observed data from a model state and returns one joint batch.

from_model(model, batch_size[, ...])

Builds a BatchManager from inferred or explicit groups.

from_split(split, batch_size[, shuffle, ...])

Build a manager from training data, including a single split.

get_batched_position(position, batch_index)

Returns the joint batched position for one optimizer step.

permute_indices(key)

Returns fresh epoch indices for every contained batch object.

scaled_log_lik(model, model_state, *[, ...])

Returns a log likelihood with per-child batch scaling.

start_epoch(key)

Starts a new joint epoch.

Attributes

axis_size

Number of observations for each contained batch object.

batch_indices

Batch index matrices selected for the joint epoch.

batch_sample_scale

Common likelihood scaling factor.

batch_sample_scales

Likelihood scaling factors for each contained batch object.

batch_sample_sizes

Batch sample sizes for each contained batch object.

batch_size

Batch size for each contained batch object.

epoch_size

"strict", "min", "max", or a positive integer.

is_full_data

Whether every child represents one full-data batch.

n_full_batches

Number of joint batch steps in one epoch.

position_keys

Position keys claimed by all contained batch objects.

sample_sizes

Full-data sample sizes for each contained batch object.

sampling_weights

a vector for a single child, or a mapping from one child position key per group to its vector.

batches

Non-empty sequence of Batches objects.