Skip to main content

model_memo

Process-lifetime reuse of model code and weights downloaded from the Hub.

A waved background run re-enters the whole DAG once per wave, so a model inference step loads its model once per wave. Nothing about that model changes between waves: the class is the same code and the weights are the same bytes. This module downloads each one once and hands every later wave the class plus a path to the weights already on local disk.

The scope is the process rather than the run. The pod keeps one process alive across scheduled reruns, growth reruns and wave continuations, so all of them reuse what the first one fetched. Nothing here survives a restart: the weights live in a temporary directory removed when the process exits.

Weights never sit on the volume in the clear. They are sealed with AES-GCM under a key this process generates at random and keeps only in memory, and read_weights decrypts them back on the way to deserialize. The key is never written anywhere, so a weights file that outlives its process - the orphans reported at start-up - is unreadable bytes rather than a model.

A failure to use the disk is not a failure to load a model. Every path through this module either returns something the caller can deserialize or leaves the memo empty so the next wave downloads again - it never raises on the caller's behalf.

Module​

Functions​

get_model_memo​

def get_model_memo() ‑> ModelMemo:

Return the process-wide model memo, creating it on first use.

Classes​

MemoisedModel​

class MemoisedModel(    model_cls: type[ModelProtocol],    weights_path: Path | None = None,    weights_bytes: bytes | None = None,):

A model class and the weights to load into an instance of it.

Exactly one of weights_path and weights_bytes is set when the model has weights at all; both are None when it has none.

Arguments

  • model_cls: The model class downloaded from the Hub.
  • weights_path: Path to the weights on local disk.
  • weights_bytes: The weights themselves, set only when they could not be written to disk. Such a result is not memoised.

Variables​

  • static model_cls : type[ModelProtocol]
  • static weights_bytes : bytes | None
  • static weights_path : Path | None

ModelMemo​

class ModelMemo(root: Path | None = None):

Remembers model classes and weights for the life of the process.

Weights are encrypted under a key generated here and held only in memory, so the bytes on the volume are readable by this process alone and by nothing at all once it exits.

Arguments

  • root: Directory to create the weights directory inside. Defaults to the configured cache directory, which is a real disk in the deployed pod
  • unlike the system temporary directory, which may be a tmpfs and would put the weights back in memory.

Methods​


get_or_load​

def get_or_load(    self, key: ModelMemoKey, load: Callable[[], tuple[type[ModelProtocol], bytes | None]],) ‑> MemoisedModel:

Return the memoised entry for key, calling load on a miss.

load runs while this key's lock is held, so concurrent waves wanting the same model download it once. The lock is per key: the inference steps of one DAG run in parallel against different models and must not wait on each other.

Arguments

  • key: Identifies the model.
  • load: Fetches the model class and its weights from the Hub. Called at most once per key per process, unless the result carried no weights, could not be written to disk, or was invalidated.

Returns The model class and whatever weights accompany it.

invalidate​

def invalidate(self, key: ModelMemoKey) ‑> None:

Forget key, deleting its weights file if it has one.

For a caller whose weights would not load: the next get_or_load downloads them again.

The delete happens under the key's lock, with the entry, because a weights file is named after its key alone. Unlinking outside the lock would let a concurrent get_or_load install a fresh file at that same path first, and this delete would then remove the new file rather than the stale one - leaving a memoised entry pointing at nothing.

Arguments

  • key: Identifies the model to forget.

read_weights​

def read_weights(self, entry: MemoisedModel) ‑> bytes | None:

Return entry's weights, decrypting them if they are on disk.

Arguments

  • entry: The memoised model whose weights to read.

Returns The weights, or None if the model has none.

Raises

  • DecryptError: If the file was altered or was written by another process, whose key this one does not have.
  • OSError: If the file could not be read.

ModelMemoKey​

class ModelMemoKey(    hub_host: str,    username: str,    model_name: str,    model_version: int,    project_id: str | None,):

Identifies one model as the Hub resolves it.

project_id is part of the identity, not just of the request. The Hub gates model access by project, and one pod process serves several projects, so a memo shared between them would serve one project bytes that were authorised for another without the Hub ever being asked. Two projects entitled to the same model therefore download it once each, which costs one extra download per process and keeps the access check where it belongs.

Arguments

  • hub_host: Host of the Hub the model came from. Distinguishes the same model name on staging from production.
  • username: The model's owner on the Hub.
  • model_name: The model's name on the Hub.
  • model_version: The model's version. Pinned by the task, so new weights normally arrive as a new version and therefore a new key.
  • project_id: The project the model was fetched for, or None for a publicly accessible model.

Variables​

  • static hub_host : str
  • static model_name : str
  • static model_version : int
  • static project_id : str | None
  • static username : str
  • digest : str - A file-name-safe digest of this key.

    Used to name the weights file, so that a model name taken from customer-editable task YAML cannot shape a path.