Training loop

class stix.training.training_loop.TrainingLoop(train_data, val_data, loss_pipeline, gen_model, optimizer_tx, config, *, num_batches_per_eval_step=5, eval_fn=None, io_handler=None)

Bases: object

Step-based training loop for stochastic interpolant models.

The loop runs for config.num_steps optimizer updates. Evaluation is triggered every config.eval_every_n_steps (if set) and always at the end of training. There is no concept of an epoch: train_data and val_data must be iterators (objects supporting next(...)) that yield enough batches for the full run; the caller is responsible for any cycling/shuffling.

Parameters:
training_state

The current training state.

best_eval_loss

Best eval loss since this process started; reset on resume, not restored (the durable best lives in the checkpoint store).

best_eval_step

The step at which best_eval_loss was found (same resume caveat).

Config

alias of TrainingLoopConfig

run()

Run the full training loop. Final state is on self.training_state.

The PRNG stream is carried on training_state.rng_key rather than reseeded, so a resumed run continues the same noise/time stream. The key is re-stamped into training_state at each evaluation boundary (right before _log_window, which is where the IO handler’s save_checkpoint fires) so a restored checkpoint resumes from the exact key the loop would next split.

Notes

  • Checkpoints are written via self.io_handler.save_checkpoint inside _log_window, i.e. at the evaluation cadence. Box death between evaluations therefore loses at most one evaluation window of steps.

  • End-to-end bit-identity across a resume also requires a replayable train_data iterator. When the iterator exposes get_state/set_state (e.g. a grain dataset iterator), its cursor is checkpointed at each evaluation boundary and restored into the live iterator on resume, so the data stream continues where it left off. Iterators without those methods restart their stream on resume.

  • best_eval_loss/best_eval_step/best_params are not restored on resume: they track the best since this process started. The durable best checkpoint is tracked separately by the checkpoint store (best_step()).

Return type:

None

property ema_model: GenerativeModel

Reconstruct a live generative model with debiased EMA parameters.

property best_params: State | None

The parameters from the best evaluation step (EMA-corrected if enabled).

restore_model(params=None)

Merge params with the graphdef and return a live model.

Parameters:

params (State | None) – Parameters to restore. Defaults to best_params.

Returns:

A reconstructed GenerativeModel with the given parameters.

Raises:

ValueError – If params is None and no best parameters have been recorded (no evaluation has run).

Return type:

GenerativeModel

class stix.training.training_loop_config.TrainingLoopConfig(*, num_steps, eval_every_n_steps=None, num_gradient_accumulation_steps=1, random_seed=42, ema_decay=0.99, use_ema_params_for_eval=True, run_eval_at_start=True)

Bases: BaseModel

Pydantic config holding all settings related to the TrainingLoop class.

Parameters:
  • num_steps (Annotated[int, Gt(gt=0)])

  • eval_every_n_steps (Annotated[int, FieldInfo(annotation=NoneType, required=True, metadata=[Gt(gt=0)])] | None)

  • num_gradient_accumulation_steps (Annotated[int, Gt(gt=0)])

  • random_seed (int)

  • ema_decay (Annotated[float, Gt(gt=0.0), Le(le=1.0)])

  • use_ema_params_for_eval (bool)

  • run_eval_at_start (bool)

num_steps

Total number of optimizer updates to run.

Type:

Annotated[int, FieldInfo(annotation=NoneType, required=True, metadata=[Gt(gt=0)])]

eval_every_n_steps

Run evaluation every N optimizer steps. If None, evaluation only runs at the end of training (and optionally at the start, see run_eval_at_start).

Type:

Annotated[int, FieldInfo(annotation=NoneType, required=True, metadata=[Gt(gt=0)])] | None

num_gradient_accumulation_steps

Number of sub-batches to accumulate gradients over before applying an optimizer update. Default is 1.

Type:

Annotated[int, FieldInfo(annotation=NoneType, required=True, metadata=[Gt(gt=0)])]

random_seed

A random seed.

Type:

int

ema_decay

The EMA decay rate.

Type:

Annotated[float, FieldInfo(annotation=NoneType, required=True, metadata=[Gt(gt=0.0), Le(le=1.0)])]

use_ema_params_for_eval

Whether to use the EMA parameters for evaluation, set to True by default.

Type:

bool

run_eval_at_start

Whether to run an evaluation on the validation set before the first training step. True by default.

Type:

bool

model_config = {}

Configuration for the model, should be a dictionary conforming to [ConfigDict][pydantic.config.ConfigDict].

class stix.training.training_state.TrainingState(params, opt_state, ema_params, num_steps, ema_decay, rng_key, graphdef)

Bases: PyTreeNode

Holds mutable training state across steps.

Parameters:
params

Current trainable parameters (nnx.Param only).

Type:

flax.nnx.statelib.State

opt_state

The optax optimizer state.

Type:

jax.Array | numpy.ndarray | numpy.bool | numpy.number | bool | int | float | complex | Iterable[ArrayTree] | Mapping[Any, ArrayTree]

ema_params

Exponentially averaged parameters for evaluation.

Type:

flax.nnx.statelib.State

num_steps

Total number of optimizer steps taken (jax.Array so it lives on-device and can be used inside JIT’d code).

Type:

jax.Array

ema_decay

The EMA decay rate used to update ema_params. Stored on the state (set from TrainingLoopConfig.ema_decay at init) so it is checkpointed alongside ema_params and num_steps; this lets offline restore debias the EMA without the training config.

Type:

jax.Array

rng_key

The PRNG key carrying the noise/time sampling stream, so a resumed run continues the same stream rather than reseeding. Stored as a PRNGKeyArray (uint32[2]) so it is checkpointed alongside the other leaves; inert inside the JIT’d train step (the step takes its key as an explicit argument).

Type:

jax.Array

graphdef

The nnx graph definition (static structure of the model).

Type:

flax.nnx.graphlib.GraphDef

replace(**updates)

Returns a new object replacing the specified fields with new values.