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, orNonefor 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.