Checkpointing and logging¶
IO handler¶
- class stix.training.training_io_handler.TrainingIOHandler(config=None, data_upload_fn=None)¶
Bases:
objectAn 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.Checkpointerinstances for training state and potentially dataset state saving and restoration.- Parameters:
config (TrainingIOHandlerConfig | None)
data_upload_fn (Callable[[str | PathLike], Future | None] | None)
- 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.
- 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
graphdefand (when a restore flag is off) itsopt_state/rng_keyare 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 totraining_state’s.restore_rng_key (bool) – Continue the checkpointed PRNG stream. When
False, reset totraining_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:
- 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:
BaseModelPydantic config holding all settings relevant for the training IO handler.
- Parameters:
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, theclear_previous_checkpointsguard is skipped. Defaults toNone.- Type:
str | os.PathLike | None
- restore_dir¶
Directory to restore checkpoints from. If
None, will default tocheckpoint_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_keepto save. If True we keep the best byeval_loss(plus the most recent!), else keep the most recentmax_to_keep. The default isTrue.- Type:
- restore_checkpoint_if_exists¶
Whether to restore a previous checkpoint if it exists. By default, this is
False.- Type:
- 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
- clear_previous_checkpoints¶
Whether to clear the previous checkpoints if any exist. Note that this setting can not be set to
Trueif one selects to restore a checkpoint. The default isFalse.- Type:
- enable_async_checkpointing¶
Whether Orbax should write checkpoints asynchronously. Defaults to
False.- Type:
- 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:
EnumEnum 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.
Checkpointer¶
- class stix.training.checkpointer.Checkpointer(config)¶
Bases:
objectPersist and restore a
TrainingStatevia Orbax.- Parameters:
config (CheckpointerConfig)
- Config¶
alias of
CheckpointerConfig
- save(training_state, step, eval_loss, dataset_state=None)¶
Persist
training_stateunderstepas independently-loadable items.The state is split across four Orbax items:
params,opt_state,ema_params(the raw, non-de-biased EMA params) andextra(the small scalars), so each can be restored on its own. In particular, sampling can easily recover the de-biased params viaema_params+extra(seerestore_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
TrainingStateandDatasetState, 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
graphdefand (when a restore flag is off) itsopt_state/rng_keyare 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 toNone.restore_optimizer_state (bool) – Keep the checkpointed optimizer state (Adam momentum). When
False, reset totraining_state’s.restore_rng_key (bool) – Continue the checkpointed PRNG stream. When
False, reset totraining_state’s key.
- Returns:
The restored training state and dataset state, or
Noneif 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_stepsdisagrees with thenum_stepsrecorded in the checkpoint’s custom metadata.- Return type:
- restore_ema(template, step=None)¶
Restore the debiased EMA parameters for offline sampling/evaluation.
Reads only the
ema_paramsandextraitems, neverparamsoropt_state, then applies the same bias correction the live loop uses for evaluation (ema / (1 - decay ** num_steps)).- Parameters:
- Returns:
The debiased EMA parameters, ready to merge with a
graphdef.- Raises:
ValueError – If no checkpoint exists.
- Return type:
- 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:
BaseModelConfiguration for
Checkpointer.- Parameters:
- checkpoint_dir¶
Directory to write to.
- Type:
- max_to_keep¶
Maximum number of checkpoints to retain.
Nonedisables 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_keepto save. If True we keep the best byeval_loss(plus the most recent!), else keep the most recentmax_to_keep.- Type:
- enable_async_checkpointing¶
Whether Orbax writes asynchronously. Defaults to
False(synchronous) in case of crashes.- Type:
- 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 OrbaxCheckpointManager. Generic by design:metadatais any JSON-serialisable dict and the caller decides the contents (e.g. a serialised run config plus a human label).
- stix.training.checkpointer.read_run_metadata(run_dir)¶
Read metadata written by
save_run_metadata(), orNoneif absent.Manager-free on purpose: this only reads a JSON sibling file, so it never constructs a
CheckpointManagerand so cannot race a live run’s in-flight checkpoint save (the failure mode that constructing one has).
EMA¶
- stix.training.ema.get_debiased_ema(ema_params, decay, num_updates)¶
Bias-corrected EMA (Adam convention).
Assumes
ema_paramswas zero-initialised and has been updatednum_updatestimes. The correctionema / (1 - decay^num_updates)makes the result unbiased under stationary-parameter expectations.Must not be called with
num_updates == 0(division by zero).- Parameters:
- Returns:
The bias-corrected EMA parameters, same structure as
ema_params.- Return type:
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_tormse_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.