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:
objectStep-based training loop for stochastic interpolant models.
The loop runs for
config.num_stepsoptimizer updates. Evaluation is triggered everyconfig.eval_every_n_steps(if set) and always at the end of training. There is no concept of an epoch:train_dataandval_datamust be iterators (objects supportingnext(...)) that yield enough batches for the full run; the caller is responsible for any cycling/shuffling.- Parameters:
loss_pipeline (LossPipeline)
gen_model (GenerativeModel)
optimizer_tx (GradientTransformation)
config (TrainingLoopConfig)
num_batches_per_eval_step (int | None)
eval_fn (Callable[[GenerativeModel], float] | None)
io_handler (TrainingIOHandler | None)
- 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_losswas 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_keyrather than reseeded, so a resumed run continues the same noise/time stream. The key is re-stamped intotraining_stateat each evaluation boundary (right before_log_window, which is where the IO handler’ssave_checkpointfires) so a restored checkpoint resumes from the exact key the loop would next split.Notes
Checkpoints are written via
self.io_handler.save_checkpointinside_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_dataiterator. When the iterator exposesget_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_paramsare 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
GenerativeModelwith the given parameters.- Raises:
ValueError – If
paramsisNoneand no best parameters have been recorded (no evaluation has run).- Return type:
- 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:
BaseModelPydantic config holding all settings related to the
TrainingLoopclass.- Parameters:
- 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, seerun_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)])]
- 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
Trueby default.- Type:
- run_eval_at_start¶
Whether to run an evaluation on the validation set before the first training step.
Trueby default.- Type:
- 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:
PyTreeNodeHolds mutable training state across steps.
- Parameters:
- params¶
Current trainable parameters (
nnx.Paramonly).- Type:
- 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:
- num_steps¶
Total number of optimizer steps taken (
jax.Arrayso it lives on-device and can be used inside JIT’d code).- Type:
- ema_decay¶
The EMA decay rate used to update
ema_params. Stored on the state (set fromTrainingLoopConfig.ema_decayat init) so it is checkpointed alongsideema_paramsandnum_steps; this lets offline restore debias the EMA without the training config.- Type:
- 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:
- graphdef¶
The nnx graph definition (static structure of the model).
- replace(**updates)¶
Returns a new object replacing the specified fields with new values.