Checkpointing and logging

IO handler

class stix.training.training_io_handler.TrainingIOHandler(config=None, data_upload_fn=None)

Bases: object

An IO handler class for the training loop.

This handles checkpointing as well as specialised logging, e.g., to some external logger that a user can provide. Checkpointing is delegated to stix.training.checkpointer.Checkpointer instances for training state and potentially dataset state saving and restoration.

Parameters:
Config

alias of TrainingIOHandlerConfig

attach_logger(logger)

Attach one training loop logging function to the IO handler.

The logging function must take in three parameters and should not return anything. The three parameters are a logging category which describes what type of data is logged (it is an enum), the data dictionary to log, and the current step number.

Parameters:

logger (Callable[[LoggingCategory, dict[str, Any], int], None]) – The logging function to add.

Return type:

None

log(category, to_log, step_number)

Log data via the logging functions stored in this class.

Parameters:
  • category (LoggingCategory) – A logging category which describes what type of data is logged (it is an enum).

  • to_log (dict[str, Any]) – A data dictionary to log (typically, metrics).

  • step_number (int) – The current step number.

Return type:

None

save_checkpoint(training_state, step_number, eval_loss, dataset_state=None)

Save a model checkpoint and upload it if an upload function is configured.

Parameters:
  • training_state (TrainingState) – The training state to save.

  • step_number (int) – The current step number.

  • eval_loss (float) – Scalar evaluation loss used for best-checkpoint ranking.

  • dataset_state (Any | Mapping[str, Any] | None) – The dataset state to save. If None, the dataset state will not be saved. Defaults to None.

Return type:

None

restore_checkpoint(training_state, dataset_state=None, restore_optimizer_state=True, restore_rng_key=True)

Restore a TrainingState, defaulting to the latest checkpoint.

Parameters:
  • training_state (TrainingState) – A freshly-initialised training state of the right structure; its static graphdef and (when a restore flag is off) its opt_state / rng_key are reused.

  • dataset_state (Any | Mapping[str, Any] | None) – A freshly-initialised dataset state of the right structure. If None or if the checkpoint does not contain a dataset state, the returned dataset state will be None. Defaults to None.

  • restore_optimizer_state (bool) – Keep the checkpointed optimizer state (Adam momentum). When False, reset to training_state’s.

  • restore_rng_key (bool) – Continue the checkpointed PRNG stream. When False, reset to training_state’s key.

Returns:

The restored training state and dataset state, or the inputs unchanged if no checkpoint exists.

Raises:

CheckpointRestorationError – If restoration is requested but checkpointing is disabled.

Return type:

tuple[TrainingState, Any | Mapping[str, Any] | None]

wait_until_finished()

Wait until local checkpoints and uploads are finished due to their asynchronous nature.

To be called at the end of a training run.

Return type:

None

class stix.training.training_io_handler.TrainingIOHandlerConfig(*, checkpoint_dir=None, restore_dir=None, max_to_keep=5, save_interval_steps=None, keep_best=True, enable_async_checkpointing=False, restore_checkpoint_if_exists=False, step_to_restore=None, restore_optimizer_state=True, clear_previous_checkpoints=False)

Bases: BaseModel

Pydantic config holding all settings relevant for the training IO handler.

Parameters:
  • checkpoint_dir (str | PathLike | None)

  • restore_dir (str | PathLike | None)

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

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

  • keep_best (bool)

  • enable_async_checkpointing (bool)

  • restore_checkpoint_if_exists (bool)

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

  • restore_optimizer_state (bool)

  • clear_previous_checkpoints (bool)

checkpoint_dir

Root checkpoint directory, e.g. a local path or any URI scheme Orbax supports. When None, the clear_previous_checkpoints guard is skipped. Defaults to None.

Type:

str | os.PathLike | None

restore_dir

Directory to restore checkpoints from. If None, will default to checkpoint_dir.

Type:

str | os.PathLike | None

max_to_keep

Maximum number of checkpoints to retain. The default is 5.

Type:

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

save_interval_steps

Save a checkpoint every N training steps. The default is None, which disables checkpointing.

Type:

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

keep_best

Which max_to_keep to save. If True we keep the best by eval_loss (plus the most recent!), else keep the most recent max_to_keep. The default is True.

Type:

bool

restore_checkpoint_if_exists

Whether to restore a previous checkpoint if it exists. By default, this is False.

Type:

bool

step_to_restore

The step number to restore. The default is None, which means the latest step will be restored.

Type:

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

restore_optimizer_state

Whether to also restore the optimizer state. Default is True.

Type:

bool

clear_previous_checkpoints

Whether to clear the previous checkpoints if any exist. Note that this setting can not be set to True if one selects to restore a checkpoint. The default is False.

Type:

bool

enable_async_checkpointing

Whether Orbax should write checkpoints asynchronously. Defaults to False.

Type:

bool

model_config = {'extra': 'forbid'}

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

class stix.training.training_io_handler.LoggingCategory(*values)

Bases: Enum

Enum class for logging categories.

These values provide a signal to a logging function what type of data is being logged.

BEST_MODEL

Information about the current best model is logged.

TRAIN_METRICS

Metrics for the training set are logged.

EVAL_METRICS

Metrics for the validation set are logged.

TEST_METRICS

Metrics for the test set are logged.

SYSTEM_METRICS

Per-process system metrics (runtime, throughput) are logged.

CLEANUP_AFTER_CKPT_RESTORATION

Allows the logger to clean itself up after a checkpoint has been restored.

exception stix.training.training_io_handler.CheckpointRestorationError

Bases: Exception

Exception to be raised if issues occur during checkpoint restoration.

Checkpointer

class stix.training.checkpointer.Checkpointer(config)

Bases: object

Persist and restore a TrainingState via Orbax.

Parameters:

config (CheckpointerConfig)

Config

alias of CheckpointerConfig

save(training_state, step, eval_loss, dataset_state=None)

Persist training_state under step as independently-loadable items.

The state is split across four Orbax items: params, opt_state, ema_params (the raw, non-de-biased EMA params) and extra (the small scalars), so each can be restored on its own. In particular, sampling can easily recover the de-biased params via ema_params + extra (see restore_ema()).

Parameters:
  • training_state (TrainingState) – The training state to persist (pulled to host first).

  • step (int) – The step number used as the checkpoint key.

  • eval_loss (float) – Validation loss at this step, used for best-checkpoint ranking and stored in the checkpoint metrics.

  • dataset_state (Any | Mapping[str, Any] | None) – The dataset state to persist. If None, the dataset state will not be saved. Defaults to None.

Return type:

None

restore(training_state, dataset_state=None, step=None, restore_optimizer_state=True, restore_rng_key=True)

Restore a TrainingState and DatasetState, defaulting to the latest checkpoint.

If the checkpoint does not contain a dataset state, the returned dataset state will be None.

Parameters:
  • training_state (TrainingState) – A freshly-initialised training state of the right structure; its static graphdef and (when a restore flag is off) its opt_state / rng_key are reused.

  • dataset_state (Any | Mapping[str, Any] | None) – A freshly-initialised dataset state of the right structure. If None or if the checkpoint does not contain a dataset state, the returned dataset state will be None. Defaults to None.

  • step (int | None) – Step to restore. If None, defaults to latest_step(). Defaults to None.

  • restore_optimizer_state (bool) – Keep the checkpointed optimizer state (Adam momentum). When False, reset to training_state’s.

  • restore_rng_key (bool) – Continue the checkpointed PRNG stream. When False, reset to training_state’s key.

Returns:

The restored training state and dataset state, or None if no checkpoint exists. If no dataset state is provided or if the checkpoint does not contain a dataset state, the returned dataset state will be None.

Raises:

ValueError – If the restored num_steps disagrees with the num_steps recorded in the checkpoint’s custom metadata.

Return type:

tuple[TrainingState, Any | Mapping[str, Any] | None] | None

restore_ema(template, step=None)

Restore the debiased EMA parameters for offline sampling/evaluation.

Reads only the ema_params and extra items, never params or opt_state, then applies the same bias correction the live loop uses for evaluation (ema / (1 - decay ** num_steps)).

Parameters:
  • template (State) – A params State of the right structure, e.g. nnx.split(model, nnx.Param)[1]. No optimizer or full training state is needed.

  • step (int | None) – Step to restore. Defaults to latest_step().

Returns:

The debiased EMA parameters, ready to merge with a graphdef.

Raises:

ValueError – If no checkpoint exists.

Return type:

State

latest_step()

Return the most recent checkpoint step, or None.

Return type:

int | None

best_step()

Return the best-eval_loss checkpoint step, or None.

Return type:

int | None

wait_until_finished()

Block until any pending asynchronous saves complete.

Return type:

None

class stix.training.checkpointer.CheckpointerConfig(*, checkpoint_dir, max_to_keep=5, save_interval_steps=100, keep_best=True, enable_async_checkpointing=False)

Bases: BaseModel

Configuration for Checkpointer.

Parameters:
checkpoint_dir

Directory to write to.

Type:

str | os.PathLike

max_to_keep

Maximum number of checkpoints to retain. None disables pruning entirely, which is the right setting for read-only managers pointed at a foreign directory. Default is 5.

Type:

Annotated[int, annotated_types.Gt(gt=0)] | None

save_interval_steps

Save a checkpoint every N training steps. The default is 100 training steps.

Type:

Annotated[int, annotated_types.Gt(gt=0)]

keep_best

Which max_to_keep to save. If True we keep the best by eval_loss (plus the most recent!), else keep the most recent max_to_keep.

Type:

bool

enable_async_checkpointing

Whether Orbax writes asynchronously. Defaults to False (synchronous) in case of crashes.

Type:

bool

model_config = {'arbitrary_types_allowed': True}

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

stix.training.checkpointer.save_run_metadata(run_dir, metadata)

Persist run-level metadata as JSON at the run root.

This is run-level (not checkpoint-level) state, so it lives beside checkpoints/ and is read/written without an Orbax CheckpointManager. Generic by design: metadata is any JSON-serialisable dict and the caller decides the contents (e.g. a serialised run config plus a human label).

Parameters:
  • run_dir (str) – The run root (parent of checkpoints/).

  • metadata (dict[str, Any]) – A JSON-serialisable dict describing the run.

Return type:

None

stix.training.checkpointer.read_run_metadata(run_dir)

Read metadata written by save_run_metadata(), or None if absent.

Manager-free on purpose: this only reads a JSON sibling file, so it never constructs a CheckpointManager and so cannot race a live run’s in-flight checkpoint save (the failure mode that constructing one has).

Parameters:

run_dir (str) – The run root (parent of checkpoints/).

Returns:

The metadata dict, or None if the metadata file does not exist.

Return type:

dict[str, Any] | None

EMA

stix.training.ema.get_debiased_ema(ema_params, decay, num_updates)

Bias-corrected EMA (Adam convention).

Assumes ema_params was zero-initialised and has been updated num_updates times. The correction ema / (1 - decay^num_updates) makes the result unbiased under stationary-parameter expectations.

Must not be called with num_updates == 0 (division by zero).

Parameters:
  • ema_params (State) – The (biased) EMA parameter tree, zero-initialised at step 0.

  • decay (float) – The EMA decay rate used during accumulation.

  • num_updates (int | Array) – Number of EMA updates applied so far.

Returns:

The bias-corrected EMA parameters, same structure as ema_params.

Return type:

State

Loggers

stix.training.training_loggers.log_metrics_to_line(category, to_log, step_number)

Logging function for the training loop which logs the metrics to a single line.

This function also converts MSE metrics to RMSE before logging them.

Parameters:
  • category (LoggingCategory) – The logging category describing what type of data is currently logged.

  • to_log (dict[str, Any]) – The data to log (typically, the metrics).

  • step_number (int) – The current step number.

Return type:

None

stix.training.training_loggers.log_metrics_to_table(category, to_log, step_number)

Logging function for the training loop which logs the metrics to a nice table.

The table will be printed to the command line.

This function also converts MSE metrics to RMSE before logging them.

Parameters:
  • category (LoggingCategory) – The logging category describing what type of data is currently logged.

  • to_log (dict[str, Any]) – The data to log (typically, the metrics).

  • step_number (int) – The current step number.

Return type:

None

stix.training.training_loggers.convert_mse_to_rmse_in_logs(to_log)

Convert metrics whose key contains mse_ to rmse_ by taking the square root.

Applied at log time rather than during accumulation, because the square root must be taken after averaging rather than before.

Parameters:

to_log (dict[str, Any]) – The metrics dictionary.

Returns:

The metrics dictionary with any MSE entries converted to RMSE.

Return type:

dict[str, Any]