spacr.torch_artifacts

Versioned, crash-safe PyTorch model artifacts used by deep_spacr.

Older spaCR releases wrote complete nn.Module objects with torch.save. Those files remain readable here, but new artifacts store state dictionaries plus the information needed to reconstruct the model and resume training.

Functions

atomic_torch_save(→ str)

Write payload beside path and atomically replace the target.

build_model_from_configuration(→ torch.nn.Module)

Reconstruct a TorchModel without pretrained weights.

capture_rng_state(→ dict[str, Any])

Capture random-generator state needed for deterministic continuation.

dependency_versions(→ dict[str, str])

Return the runtime versions that materially affect a model artifact.

load_model_artifact(→ tuple[torch.nn.Module, dict[str, ...)

Load current artifacts and legacy full-module/state-dict checkpoints.

make_model_artifact(→ dict[str, Any])

Build the canonical serializable spaCR PyTorch artifact.

model_configuration(→ dict[str, Any])

Return the constructor information required to rebuild model.

restore_rng_state(→ None)

Restore a state returned by capture_rng_state().

restore_training_state(→ dict[str, Any])

Restore optimizer/scheduler/RNG state and return training metadata.

save_model_artifact(→ str)

Build and atomically save a canonical spaCR model artifact.

Module Contents

spacr.torch_artifacts.atomic_torch_save(payload: Any, path: str) → str[source]

Write payload beside path and atomically replace the target.

Parameters:
  • payload – object to serialize with torch.save().

  • path – destination artifact path to replace atomically.

spacr.torch_artifacts.build_model_from_configuration(config: collections.abc.Mapping[str, Any]) → torch.nn.Module[source]

Reconstruct a TorchModel without pretrained weights.

Parameters:

config – recorded model-constructor settings.

spacr.torch_artifacts.capture_rng_state() → dict[str, Any][source]

Capture random-generator state needed for deterministic continuation.

spacr.torch_artifacts.dependency_versions() → dict[str, str][source]

Return the runtime versions that materially affect a model artifact.

spacr.torch_artifacts.load_model_artifact(path: str, *, map_location: Any = 'cpu', model: torch.nn.Module | None = None, strict: bool = True) → tuple[torch.nn.Module, dict[str, Any]][source]

Load current artifacts and legacy full-module/state-dict checkpoints.

The returned metadata dict always contains legacy. Current artifacts retain their optimizer/scheduler/RNG state so callers can resume training.

Parameters:
  • path – checkpoint to read. Unpickled with weights_only=False, so only files you trust are safe to pass.

  • map_location – forwarded to torch.load(). The cpu default lets a GPU-trained checkpoint load on a machine without a GPU; an unrecognised device string raises RuntimeError from torch.

  • model – None rebuilds the architecture from the recorded config – a legacy bare state dict records none, so maxvit_t is assumed silently. A module passed here is loaded IN PLACE and returned as the same object; for legacy full-module files it is ignored and the file’s own module comes back.

  • strict – forwarded to load_state_dict; True raises RuntimeError on any key mismatch, False tolerates missing and unexpected keys so a mismatched architecture loads quietly with parts still randomly initialised. Also ignored for legacy full-module files.

Raises:

ValueError – the file is neither a module nor a checkpoint mapping, its artifact_version is not the supported one, no state dictionary was found, or the config names no architecture and no model was supplied.

Returns:

(model, metadata).

spacr.torch_artifacts.make_model_artifact(model: torch.nn.Module, *, optimizer=None, scheduler=None, epoch: int | None = None, metrics: collections.abc.Mapping[str, Any] | None = None, best_metric: float | None = None, epochs_without_improvement: int = 0, preprocessing: collections.abc.Mapping[str, Any] | None = None, classes: list[str] | None = None, channels: list[str] | None = None, include_rng: bool = True, artifact_role: str = 'model') → dict[str, Any][source]

Build the canonical serializable spaCR PyTorch artifact.

Everything needed to REBUILD the model, not only to load its weights: model_configuration() records the architecture so build_model_from_configuration() can reconstruct it without the caller remembering what it was.

Parameters:
  • model – the module to serialise. Its state_dict and its configuration are both captured.

  • optimizer – optimiser whose state to store, so training can RESUME rather than restart. Silently stored as None if it has no state_dict, which is what makes a plain object safe to pass.

  • scheduler – learning-rate scheduler, same contract as optimizer.

  • epoch – epoch this artifact was written at. None records 0.

  • metrics – whatever the caller measured, stored verbatim. Not interpreted, so the keys are the caller’s own.

  • best_metric – the best value seen so far, for checkpoint selection on resume. None means “no best recorded”, which is not the same as zero.

  • epochs_without_improvement – early-stopping counter, carried so a resumed run does not forget how close it was to stopping.

  • preprocessing – the transform the inputs were prepared with. Without it a loaded model can be fed differently-normalised images and can produce incorrect results without failing.

  • classes – class names in OUTPUT-COLUMN order. The order is the contract – a reordered list silently relabels every prediction.

  • channels – input channel names, in channel order, same contract.

  • include_rng – capture Python/NumPy/torch RNG state, so a resumed run continues the same stream. Turn it off for a smaller artifact when exact resumption does not matter.

  • artifact_role – what this file IS – model for a trained model, another role for a companion artifact – recorded so a loader can tell them apart.

Returns:

the artifact dict, ready for atomic_torch_save().

spacr.torch_artifacts.model_configuration(model: torch.nn.Module) → dict[str, Any][source]

Return the constructor information required to rebuild model.

Parameters:

model – module whose reconstruction settings should be captured.

spacr.torch_artifacts.restore_rng_state(state: collections.abc.Mapping[str, Any] | None) → None[source]

Restore a state returned by capture_rng_state().

Parameters:

state – captured generator states, or None for a no-op.

spacr.torch_artifacts.restore_training_state(payload: collections.abc.Mapping[str, Any], *, optimizer=None, scheduler=None, restore_random_generators: bool = True) → dict[str, Any][source]

Restore optimizer/scheduler/RNG state and return training metadata.

Parameters:
  • payload – the metadata dict returned by load_model_artifact(). Must be a mapping – None raises AttributeError.

  • optimizer – applied only when the payload also carries an optimizer_state_dict; otherwise silently skipped. Restoring it overwrites the live optimiser settings, learning rate included, with the checkpoint’s.

  • scheduler – same contract, restoring last_epoch and with it the position in the schedule.

  • restore_random_generators – reseeds the PROCESS-WIDE Python, NumPy and torch generators, not just this model’s. A no-op for artifacts written with include_rng=False.

Returns:

a copy of the payload’s training_state; {} when it is absent, as it is for legacy full-module checkpoints.

spacr.torch_artifacts.save_model_artifact(model: torch.nn.Module, path: str, **kwargs) → str[source]

Build and atomically save a canonical spaCR model artifact.

Parameters:
  • model – module whose configuration and state should be saved.

  • path – destination artifact path.