Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
14 changes: 10 additions & 4 deletions docs/loading_weights.md
Original file line number Diff line number Diff line change
Expand Up @@ -194,10 +194,16 @@ model = Qwen3TextGenerate.from_weights(
It works for both conversion paths. Way 1 has nothing to cache beyond the downloaded file.

The cache key includes the source identity, the backend and dtype, and the quantization
recipe, so it cannot hand back a stale or differently configured model. For an `hf:` id the
source identity is the resolved **commit SHA**, so a repo that moves invalidates the entry.
A miss falls back to the normal path silently. On an ephemeral machine (Colab, CI) point
`ZEROMODELS_HOME` at persistent storage or the cache buys you nothing.
recipe to separate differently configured models. For an `hf:` id the source identity is
the resolved **commit SHA**, so a repo that moves invalidates the entry. A metadata
fingerprint detects stale or accidentally damaged entries; it is stored in the cache and is
not a signature or tamper-resistance mechanism.

The cache is trusted local input: its Keras metadata names classes that are resolved during
deserialization. Do not place `ZEROMODELS_HOME` somewhere writable by untrusted users. New
cache directories use owner-only permissions where the filesystem supports them. A miss
falls back to the normal path silently. On an ephemeral machine (Colab, CI), point
`ZEROMODELS_HOME` at trusted persistent storage or the cache buys you nothing.

## Loading big checkpoints

Expand Down
6 changes: 4 additions & 2 deletions zeromodels/base/base_mixin.py
Original file line number Diff line number Diff line change
Expand Up @@ -304,8 +304,10 @@ def from_weights(
Applies to models loaded as float (a built functional graph can't
be re-quantized from a serialized skeleton, so quantized loads are
not cached); a cache miss / failure silently falls back to the
source path. Best on a persistent disk: set ``ZEROMODELS_HOME``
on ephemeral boxes.
source path. Cache metadata is trusted local input and may name
classes for Keras to deserialize; use only a location that is not
writable by untrusted users. Best on a persistent disk: set
``ZEROMODELS_HOME`` on ephemeral boxes.
**kwargs: Forwarded to the model constructor (or to
``from_hf`` when applicable).

Expand Down
42 changes: 36 additions & 6 deletions zeromodels/conversion/converted_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,8 @@ def cache_root():
``$ZEROMODELS_HOME/converted`` (else ``~/.cache/zeromodels/converted``),
self-managed like the HF cache. On an ephemeral box (Colab), point
``ZEROMODELS_HOME`` at a persistent mount (Drive) to keep the benefit
across sessions.
across sessions. Cache metadata is trusted local input, not authenticated
content, so the location must not be writable by untrusted users.
"""
home = os.environ.get(
"ZEROMODELS_HOME",
Expand Down Expand Up @@ -142,7 +143,26 @@ def save_converted(model, directory, quantization, load_dtype=None):
"""
from safetensors.numpy import save_file

os.makedirs(directory, exist_ok=True)
root = os.path.abspath(cache_root())
target = os.path.abspath(directory)
directories = [target]
try:
if os.path.commonpath((root, target)) == root:
directories.insert(0, root)
except ValueError:
# Different drives on Windows cannot share a common path. ``directory``
# is then an explicit external target rather than the configured cache.
pass
for cache_directory in directories:
os.makedirs(cache_directory, mode=0o700, exist_ok=True)
if os.name == "posix":
try:
os.chmod(cache_directory, 0o700)
except OSError:
# Some mounted filesystems (for example cloud-drive FUSE mounts)
# do not implement chmod. The documented trusted-cache requirement
# still applies there.
pass
weights = list(model.weights)

keys = [f"{i:06d}" for i in range(len(weights))]
Expand Down Expand Up @@ -176,7 +196,7 @@ def save_converted(model, directory, quantization, load_dtype=None):
"backend": keras.backend.backend(),
"load_dtype": load_dtype,
"config": config,
"architecture_hash": _json_hash(config),
"architecture_fingerprint": _json_hash(config),
"quantization": quant_id(quantization),
"keying": "index",
"keys": keys,
Expand All @@ -193,7 +213,10 @@ def load_converted(directory, quantization, load_dtype):

Deserializes the config to the model, then streams each cached tensor onto
its weight by position. Raises on any count / shape / keying mismatch so the
caller can fall back to the source.
caller can fall back to the source. The architecture fingerprint detects
stale or accidental corruption only. Because Keras deserialization resolves
the classes named in ``meta.json``, cache directories must be trusted and not
writable by untrusted users.
"""
from zeromodels.base.base_mixin import build_dtype_scope

Expand All @@ -214,8 +237,15 @@ def load_converted(directory, quantization, load_dtype):
raise ValueError(
f"Converted cache {key}={meta.get(key)!r} does not match {value!r}."
)
if meta.get("architecture_hash") != _json_hash(meta.get("config")):
raise ValueError("Converted cache architecture config is corrupt.")
fingerprint = meta.get("architecture_fingerprint")
if fingerprint is None:
# Compatibility with caches written before the field was accurately
# named. This legacy value has the same staleness-only semantics.
fingerprint = meta.get("architecture_hash")
if fingerprint != _json_hash(meta.get("config")):
raise ValueError(
"Converted cache architecture fingerprint is stale or damaged."
)

with build_dtype_scope(load_dtype):
model = keras.saving.deserialize_keras_object(meta["config"])
Expand Down
Loading