BatchManager#
- class liesel.optim.BatchManager(batches, epoch_size='strict', *, sampling_weights=None)[source]#
Bases:
objectCoordinates multiple
Batchesobjects as one batching interface.A
BatchManageris useful when a model contains observed branches with different observation sizes. Each containedBatchesobject 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 ofBatchesobjects. Theirposition_keysmust 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’ssample_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
batchesis empty, if anyposition_keysare claimed by more than one child, ifepoch_sizeis invalid, or if strict sizing is used with unequal childBatches.n_full_batches.
Notes
The properties
axis_size,batch_size, andbatch_sample_scalesreturn tuples in child-batch order. The scalar aliases are available only when all children have the same likelihood scale. With unequal scales, usescaled_log_lik()so each branch is scaled by its own sample-size ratio.Use manual
BatchManager([Batches(...)])construction when child groups need custom per-branchsample_sizeorbatch_sample_sizevalues.Like
Batches,start_epoch()mutates and returnsself.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
BatchManagerfrom 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
Number of observations for each contained batch object.
Batch index matrices selected for the joint epoch.
Common likelihood scaling factor.
Likelihood scaling factors for each contained batch object.
Batch sample sizes for each contained batch object.
Batch size for each contained batch object.
"strict","min","max", or a positive integer.Whether every child represents one full-data batch.
Number of joint batch steps in one epoch.
Position keys claimed by all contained batch objects.
Full-data sample sizes for each contained batch object.
a vector for a single child, or a mapping from one child position key per group to its vector.
Non-empty sequence of
Batchesobjects.