From cc23d689b531e9f01352a24cb1e2f2df7c6cd4f5 Mon Sep 17 00:00:00 2001 From: lgoyal6 Date: Sun, 6 Sep 2026 16:08:16 -0700 Subject: [PATCH 1/2] Add a provenance and pre-activation guard for model artifacts Every model in this repo enters through from_pretrained() reading a persistent, mutable Modal Volume used as the Hugging Face cache. The volume is populate-once-then-read and is committed back from inside the GPU container, so whatever sits in it at load time is what gets loaded. No load site pinned a revision, forced safetensors, or checked a digest, which made the cache an unverified trust boundary. That is not hypothetical here. BAAI/bge-small-en-v1.5, the SmallEmbedder in two_stage_compressor, publishes both model.safetensors and pytorch_model.bin. A .bin checkpoint is a zipped Python pickle, and unpickling is arbitrary code execution that runs during load, before any shape or dtype is inspected. Whether the installed transformers happens to default weights_only=True is not something this repo controls, because the Modal images pip-install transformers unpinned. So the file is taken out of the decision rather than the loader trusted. model_guard enforces five things: provenance a pinned commit SHA per model id, read from the hub API on 2026-09-05, so a change to the hub's main branch cannot silently change the weights no code trust_remote_code=False, and verify_snapshot refuses a snapshot containing any *.py no pickle use_safetensors=True; a pickle that must be read at all goes through scan_pickle, an opcode-level allowlist that runs BEFORE any unpickling, and only then through torch.load(weights_only=True) validation digests and byte sizes against a manifest, then tensor key order, shapes, dtypes and finiteness, all before activation rollback KnownGood keeps the last manifest that passed, so a rejected candidate leaves the previous revision in place Only the stdlib is on the validation path (pickletools, hashlib, json, zipfile). torch and safetensors are imported lazily and only when tensors are actually read, so the module imports and its tests run on a CPU-only box with neither installed. test_model_guard.py covers the guard directly and adds three repo-wide AST sweeps, so a future load site that calls from_pretrained bare, or reaches for torch.load / pickle.load without the guard, fails the suite rather than passing unnoticed. --- model_guard.py | 458 ++++++++++++++++++++++++++++++++++++++++++++ test_model_guard.py | 455 +++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 913 insertions(+) create mode 100644 model_guard.py create mode 100644 test_model_guard.py diff --git a/model_guard.py b/model_guard.py new file mode 100644 index 0000000..d4e9942 --- /dev/null +++ b/model_guard.py @@ -0,0 +1,458 @@ +""" +Provenance pinning and pre-activation validation for model artifacts. + +WHY THIS EXISTS +--------------- +Every model in this repo enters through `from_pretrained(...)` reading a +PERSISTENT, MUTABLE Modal Volume used as the Hugging Face cache +(`llmlingua2-hf-cache`, `attentionrag-hf-cache`, `turboquant-hf-cache`). The +volume is populate-once-then-read and is committed back from inside the GPU +container, so whatever sits in it at load time is what gets loaded. Before this +module none of the load sites pinned a revision, forced safetensors, or checked +a digest - the cache was an unverified trust boundary. + +That is not hypothetical here. `BAAI/bge-small-en-v1.5` (the SmallEmbedder in +`two_stage_compressor`) publishes BOTH `model.safetensors` and +`pytorch_model.bin`. A `.bin` checkpoint is a zipped Python pickle, and +unpickling is arbitrary code execution: the payload runs during `load`, before +any shape or dtype is ever inspected. Whether the installed `transformers` +happens to pass `weights_only=True` is not something this repo controls - the +Modal images pip-install `transformers` unpinned. So we take the file out of +the decision instead of trusting the loader. + +WHAT THIS MODULE ENFORCES +------------------------- + 1. provenance - a pinned commit SHA per model id (`REVISIONS`), so a change to + the hub's `main` branch cannot silently change the weights. + 2. no code - `trust_remote_code=False`, and `verify_snapshot` refuses a + snapshot that contains any `*.py`. + 3. no pickle - `use_safetensors=True`. If a pickle must be read at all it + goes through `scan_pickle`, an opcode-level allowlist that runs BEFORE any + unpickling, and only then through `torch.load(..., weights_only=True)`. + 4. validation BEFORE activation - digests and byte sizes against a manifest, + then tensor key order, shapes, dtypes and finiteness. + 5. rollback - `KnownGood` keeps the last manifest that passed, so a rejected + candidate leaves the previously-good revision in place and available. + +Only the stdlib is used on the validation path (`pickletools`, `hashlib`, +`json`, `zipfile`). `torch` and `safetensors` are imported lazily and only when +tensors are actually read, so this module imports - and its tests run - on a +CPU-only box with neither installed. +""" + +from __future__ import annotations + +import hashlib +import json +import os +import pickletools +import zipfile +from typing import Dict, Mapping, Optional, Sequence + + +class ArtifactRejected(Exception): + """An artifact failed verification and must not be activated.""" + + +# --------------------------------------------------------------------------- # +# 1. Provenance: pinned revisions # +# --------------------------------------------------------------------------- # +# Commit SHAs read from https://huggingface.co/api/models/ on 2026-09-05. +# Pinning is the whole point: bump these deliberately, never automatically. +REVISIONS: Dict[str, str] = { + "microsoft/llmlingua-2-xlm-roberta-large-meetingbank": + "ebaba9b0e874dadd3003ffcff828e4397e568089", + "BAAI/bge-reranker-v2-m3": "953dc6f6f85a1b2dbfca4c34a2796e7dde08d41e", + "BAAI/bge-small-en-v1.5": "5c38ec7c405ec4b44b94cc5a9bb96e735b38267a", + "Qwen/Qwen2.5-7B-Instruct": "a09a35458c702b33eeacc393d103063234e8bc28", + "Qwen/Qwen2.5-14B-Instruct": "cf98f3b3bbb457ad9e2bb7baf9a0125b6b88caa8", + "mistralai/Mistral-7B-Instruct-v0.3": "c170c708c41dac9275d15a8fff4eca08d52bab71", +} + + +def pinned_revision(model_id: str) -> str: + """Return the pinned commit SHA for `model_id`, or refuse. + + Fail-closed: an unknown model id is an unreviewed artifact, so it is + rejected rather than resolved to whatever `main` points at today. + """ + rev = REVISIONS.get(model_id) + if not rev: + raise ArtifactRejected( + f"no pinned revision for {model_id!r}; add its commit SHA to " + "model_guard.REVISIONS after reviewing the artifact" + ) + return rev + + +def guarded_kwargs(model_id: str, **kwargs) -> dict: + """Return `kwargs` with the non-negotiable load flags forced on. + + Refuses rather than silently downgrades if a caller tries to opt out of + either safety flag - an override would have to be argued for in code review, + not passed at a call site. + """ + if kwargs.get("trust_remote_code"): + raise ArtifactRejected( + f"trust_remote_code=True on {model_id!r} executes repo-supplied " + "Python at load time; refused" + ) + if kwargs.get("use_safetensors") is False: + raise ArtifactRejected( + f"use_safetensors=False on {model_id!r} allows the pickled .bin " + "checkpoint path; refused" + ) + kwargs["revision"] = kwargs.get("revision") or pinned_revision(model_id) + kwargs["use_safetensors"] = True + kwargs["trust_remote_code"] = False + return kwargs + + +def guarded_from_pretrained(loader, model_id: str, **kwargs): + """`loader.from_pretrained(model_id, ...)` with the guard flags applied. + + `loader` is an Auto* class (AutoTokenizer, AutoModelForCausalLM, ...). + Tokenizers have no safetensors weights, so `use_safetensors` is dropped for + them; the revision pin and the remote-code refusal still apply. + """ + kw = guarded_kwargs(model_id, **kwargs) + if "Tokenizer" in getattr(loader, "__name__", ""): + kw.pop("use_safetensors", None) + return loader.from_pretrained(model_id, **kw) + + +# --------------------------------------------------------------------------- # +# 2. Pickle refusal (opcode-level, runs BEFORE any unpickling) # +# --------------------------------------------------------------------------- # +# The only globals a torch state_dict legitimately needs. Anything else - and +# anything we cannot statically resolve - is refused. +PICKLE_ALLOWED_GLOBALS = frozenset({ + "collections OrderedDict", + "torch _utils _rebuild_tensor", + "torch _utils _rebuild_tensor_v2", + "torch _utils _rebuild_parameter", + "torch FloatStorage", "torch DoubleStorage", "torch HalfStorage", + "torch BFloat16Storage", "torch LongStorage", "torch IntStorage", + "torch ShortStorage", "torch CharStorage", "torch ByteStorage", + "torch BoolStorage", + "numpy dtype", "numpy ndarray", "numpy core multiarray _reconstruct", +}) + +# Opcodes that build an object from a global we cannot see on the opcode stream. +_OPAQUE_CONSTRUCTORS = ("INST", "OBJ", "EXT1", "EXT2", "EXT4") + + +def _normalize_global(module: str, name: str) -> str: + return f"{module.replace('.', ' ')} {name}" + + +def scan_pickle(data: bytes, *, where: str = "") -> None: + """Statically reject a pickle that references a non-allowlisted global. + + Walks the opcode stream with `pickletools.genops` - nothing is constructed + and no module is imported, so this is safe to run on a hostile file. Raises + `ArtifactRejected` on the first offending opcode. + + Fail-closed in three places: an unresolvable STACK_GLOBAL, an opaque + constructor opcode, and a stream that will not even parse are all refusals. + """ + strings: list = [] + try: + ops = list(pickletools.genops(data)) + except Exception as exc: # truncated / malformed / adversarial framing + raise ArtifactRejected(f"{where}: unparseable pickle stream ({exc})") from None + + for op, arg, _pos in ops: + code = op.name + if code in _OPAQUE_CONSTRUCTORS: + raise ArtifactRejected( + f"{where}: opcode {code} constructs an object from an " + "unresolvable global; refused" + ) + if code == "GLOBAL": + mod, _, nm = str(arg).partition(" ") + g = _normalize_global(mod, nm) + if g not in PICKLE_ALLOWED_GLOBALS: + raise ArtifactRejected(f"{where}: disallowed pickle global {arg!r}") + elif code == "STACK_GLOBAL": + if len(strings) < 2: + raise ArtifactRejected( + f"{where}: STACK_GLOBAL with an unresolvable module/name pair" + ) + mod, nm = strings[-2], strings[-1] + g = _normalize_global(str(mod), str(nm)) + if g not in PICKLE_ALLOWED_GLOBALS: + raise ArtifactRejected( + f"{where}: disallowed pickle global '{mod} {nm}'" + ) + strings = strings[:-2] + elif isinstance(arg, str): + strings.append(arg) + if len(strings) > 64: # only the top of the stack can matter + del strings[:-64] + + +def scan_checkpoint(path: str) -> None: + """Run `scan_pickle` over every pickle inside a torch checkpoint file. + + Handles both the modern zip container (`archive/data.pkl`) and the legacy + flat pickle stream. + """ + if zipfile.is_zipfile(path): + with zipfile.ZipFile(path) as zf: + names = [n for n in zf.namelist() if n.endswith(".pkl")] + if not names: + raise ArtifactRejected(f"{path}: zip checkpoint has no .pkl member") + for n in names: + scan_pickle(zf.read(n), where=f"{path}!{n}") + return + with open(path, "rb") as fh: + scan_pickle(fh.read(), where=path) + + +def safe_load_state_dict(path: str) -> Mapping: + """Load a state dict from `path`, preferring safetensors and never trusting + a pickle. + + `.safetensors` is read directly (a pure tensor container, no code). A + `.bin`/`.pt`/`.pth` is scanned with `scan_checkpoint` first and only then + handed to `torch.load(..., weights_only=True)` - belt and braces, because + `weights_only` is a property of the installed torch, and the scan is a + property of this repo. + """ + ext = os.path.splitext(path)[1].lower() + if ext == ".safetensors": + from safetensors.torch import load_file + + return load_file(path) + if ext not in (".bin", ".pt", ".pth", ".ckpt"): + raise ArtifactRejected(f"{path}: unsupported checkpoint extension {ext!r}") + scan_checkpoint(path) + import torch + + return torch.load(path, map_location="cpu", weights_only=True) + + +# --------------------------------------------------------------------------- # +# 3. Snapshot verification (digests + provenance), BEFORE activation # +# --------------------------------------------------------------------------- # +def sha256_file(path: str, _chunk: int = 1 << 20) -> str: + h = hashlib.sha256() + with open(path, "rb") as fh: + for block in iter(lambda: fh.read(_chunk), b""): + h.update(block) + return h.hexdigest() + + +def build_manifest(snapshot_dir: str, model_id: str, revision: str) -> dict: + """Record a snapshot's provenance. Run this ONCE on a snapshot you have + reviewed; commit the result and verify against it from then on.""" + files = {} + for root, _dirs, names in os.walk(snapshot_dir): + for n in sorted(names): + p = os.path.join(root, n) + rel = os.path.relpath(p, snapshot_dir) + files[rel] = {"sha256": sha256_file(p), "bytes": os.path.getsize(p)} + return {"model_id": model_id, "revision": revision, "files": files} + + +def verify_snapshot(snapshot_dir: str, manifest: Mapping) -> None: + """Reject a snapshot that does not match `manifest` exactly. + + Checks, in order: no executable Python in the tree, no missing files, no + unexpected extra files, byte size, sha256. Every mismatch is a refusal - + there is no "close enough" for a file we are about to load into a process. + """ + expected = manifest["files"] + present = {} + for root, _dirs, names in os.walk(snapshot_dir): + for n in names: + p = os.path.join(root, n) + present[os.path.relpath(p, snapshot_dir)] = p + + code = sorted(r for r in present if r.endswith(".py")) + if code: + raise ArtifactRejected( + f"{snapshot_dir}: snapshot ships executable Python {code}; refused" + ) + + missing = sorted(set(expected) - set(present)) + if missing: + raise ArtifactRejected(f"{snapshot_dir}: missing files {missing}") + extra = sorted(set(present) - set(expected)) + if extra: + raise ArtifactRejected(f"{snapshot_dir}: unexpected files {extra}") + + for rel in sorted(expected): + want, path = expected[rel], present[rel] + size = os.path.getsize(path) + if size != want["bytes"]: + raise ArtifactRejected( + f"{snapshot_dir}: {rel} is {size} bytes, manifest says {want['bytes']}" + ) + got = sha256_file(path) + if got != want["sha256"]: + raise ArtifactRejected( + f"{snapshot_dir}: {rel} sha256 {got[:12]}... != manifest " + f"{want['sha256'][:12]}..." + ) + + +# --------------------------------------------------------------------------- # +# 4. Tensor validation, BEFORE the weights are attached to a model # +# --------------------------------------------------------------------------- # +def _dtype_name(t) -> str: + return str(getattr(t, "dtype", "?")).replace("torch.", "") + + +def _all_finite(t) -> bool: + isfinite = getattr(t, "isfinite", None) + if isfinite is not None: # torch / numpy tensor + return bool(isfinite().all()) + values = t.tolist() if hasattr(t, "tolist") else t + stack = [values] + while stack: + v = stack.pop() + if isinstance(v, (list, tuple)): + stack.extend(v) + elif v != v or v in (float("inf"), float("-inf")): + return False + return True + + +def validate_tensors(state: Mapping, spec: Mapping) -> None: + """Reject a state dict that does not match `spec` before it is activated. + + `spec` is ``{"order": [name, ...], "tensors": {name: {"shape": [...], + "dtype": "float32"}}}``. Key ORDER is checked as well as membership: the + order is the feature order the downstream code indexes by, and a silently + permuted state dict loads cleanly and then produces garbage. + """ + want = spec["tensors"] + got_keys = list(state.keys()) + + missing = sorted(set(want) - set(got_keys)) + if missing: + raise ArtifactRejected(f"state dict is missing tensors {missing}") + extra = sorted(set(got_keys) - set(want)) + if extra: + raise ArtifactRejected(f"state dict has unexpected tensors {extra}") + + order = spec.get("order") + if order is not None and got_keys != list(order): + raise ArtifactRejected( + f"state dict key order differs from the spec: got {got_keys}, " + f"expected {list(order)}" + ) + + for name in got_keys: + t, w = state[name], want[name] + shape = tuple(getattr(t, "shape", ())) + if shape != tuple(w["shape"]): + raise ArtifactRejected( + f"{name}: shape {shape} != expected {tuple(w['shape'])}" + ) + dtype = _dtype_name(t) + if dtype != w["dtype"]: + raise ArtifactRejected(f"{name}: dtype {dtype!r} != expected {w['dtype']!r}") + if not _all_finite(t): + raise ArtifactRejected(f"{name}: contains NaN or Inf") + + +# --------------------------------------------------------------------------- # +# 5. Known-good rollback # +# --------------------------------------------------------------------------- # +class KnownGood: + """Last-verified manifest per model id, persisted as JSON. + + `activate` is the only way in: a candidate that fails verification raises + and the store is left untouched, so `current(model_id)` still names the + revision that was known to work. + """ + + def __init__(self, path: str): + self.path = path + self._data: Dict[str, dict] = {} + if os.path.exists(path): + with open(path) as fh: + self._data = json.load(fh) + + def current(self, model_id: str) -> Optional[dict]: + return self._data.get(model_id) + + def activate(self, snapshot_dir: str, manifest: Mapping) -> dict: + """Verify `snapshot_dir` against `manifest`, then promote it. + + On failure the exception propagates and nothing is written. + """ + verify_snapshot(snapshot_dir, manifest) + self._data[manifest["model_id"]] = dict(manifest) + tmp = f"{self.path}.tmp" + with open(tmp, "w") as fh: + json.dump(self._data, fh, indent=2, sort_keys=True) + os.replace(tmp, self.path) # atomic: never leave a half-written store + return dict(manifest) + + +# --------------------------------------------------------------------------- # +# 6. Convenience wrapper for the Modal @enter loaders # +# --------------------------------------------------------------------------- # +def llmlingua_model_config(model_id: str, **extra) -> dict: + """`model_config` for llmlingua's `PromptCompressor`, with remote code OFF. + + llmlingua 0.2.2 defaults this to ON. From its own source + (`llmlingua/prompt_compressor.py`, load_model): + + trust_remote_code = model_config.get("trust_remote_code", True) + if "trust_remote_code" not in model_config: + model_config["trust_remote_code"] = trust_remote_code + + and the dict is then forwarded verbatim into `AutoConfig.from_pretrained`, + `AutoTokenizer.from_pretrained` and `MODEL_CLASS.from_pretrained`. So unless + the caller says otherwise the library imports and runs whatever Python the + model directory ships - and here that directory is a mutable cache volume. + Passing this dict turns it off and pins the revision. + """ + cfg = {"trust_remote_code": False, "revision": pinned_revision(model_id)} + cfg.update(extra) + if cfg.get("trust_remote_code"): + raise ArtifactRejected( + f"trust_remote_code=True on {model_id!r} executes repo-supplied " + "Python at load time; refused" + ) + return cfg + + +def pinned_snapshot_download(model_id: str, **kwargs): + """`huggingface_hub.snapshot_download` at the pinned revision. + + Without a revision the download resolves `main`, so a cache that was + populated last month and one populated today can hold different weights + under the same name. + """ + from huggingface_hub import snapshot_download + + kwargs.setdefault("revision", pinned_revision(model_id)) + return snapshot_download(model_id, **kwargs) + + +def assert_no_pickled_weights(snapshot_dir: str) -> Sequence[str]: + """Refuse a snapshot that still contains a pickled checkpoint next to the + safetensors one (e.g. BAAI/bge-small-en-v1.5 ships both). Returns the + safetensors files it found.""" + pickled, safe = [], [] + for root, _dirs, names in os.walk(snapshot_dir): + for n in names: + rel = os.path.relpath(os.path.join(root, n), snapshot_dir) + if n.endswith((".bin", ".pt", ".pth", ".ckpt")): + pickled.append(rel) + elif n.endswith(".safetensors"): + safe.append(rel) + if pickled: + raise ArtifactRejected( + f"{snapshot_dir}: pickled checkpoint(s) present {sorted(pickled)}; " + "delete them or download with allow_patterns that exclude them" + ) + if not safe: + raise ArtifactRejected(f"{snapshot_dir}: no safetensors weights found") + return sorted(safe) diff --git a/test_model_guard.py b/test_model_guard.py new file mode 100644 index 0000000..982a39a --- /dev/null +++ b/test_model_guard.py @@ -0,0 +1,455 @@ +"""Local tests for model_guard (no GPU, no torch, no network). + +Run: python test_model_guard.py + +Two halves: + * the guard itself - pickle refusal, snapshot digests, tensor validation, + known-good rollback; + * the WIRING - a fake `transformers` is injected into sys.modules and the real + loader classes are constructed, so the tests assert what the production call + sites actually pass to `from_pretrained`. +""" + +import collections +import contextlib +import json +import os +import pickle +import shutil +import sys +import tempfile +import types +import zipfile + +from model_guard import ( + ArtifactRejected, + KnownGood, + assert_no_pickled_weights, + build_manifest, + guarded_kwargs, + pinned_revision, + scan_checkpoint, + scan_pickle, + validate_tensors, + verify_snapshot, +) + + +# --------------------------------------------------------------------------- # +# helpers # +# --------------------------------------------------------------------------- # +class FakeTensor: + """Minimal tensor stand-in: shape + dtype + tolist(), no torch needed.""" + + def __init__(self, shape, dtype="float32", values=None): + self.shape = tuple(shape) + self.dtype = dtype + n = 1 + for d in self.shape: + n *= d + self._values = list(values) if values is not None else [0.5] * n + + def tolist(self): + return self._values + + +class _Payload: + """Pickling this yields a stream that CALLS os.makedirs on load.""" + + def __init__(self, path): + self.path = path + + def __reduce__(self): + return (os.makedirs, (self.path,)) + + +@contextlib.contextmanager +def _tmpdir(): + d = tempfile.mkdtemp(prefix="winnow-guard-") + try: + yield d + finally: + shutil.rmtree(d, ignore_errors=True) + + +def _write(path, data=b"x"): + os.makedirs(os.path.dirname(path), exist_ok=True) + with open(path, "wb") as fh: + fh.write(data) + return path + + +def _raises(exc, fn, *a, **k): + try: + fn(*a, **k) + except exc as e: + return str(e) + raise AssertionError(f"expected {exc.__name__} from {getattr(fn, '__name__', fn)}") + + +# --------------------------------------------------------------------------- # +# 1. the pickle RCE primitive, and its refusal # +# --------------------------------------------------------------------------- # +def test_malicious_pickle_really_executes_but_is_refused(): + """The primitive is real: prove it fires, then prove the scanner stops it. + + Half one runs the payload through plain `pickle.loads` and checks the + side effect actually happened - without that half, "the scanner refused it" + proves nothing. Half two shows `scan_pickle` refusing the identical bytes + with no unpickling at all. + """ + with _tmpdir() as d: + canary = os.path.join(d, "pwned") + blob = pickle.dumps(_Payload(canary)) + + assert not os.path.exists(canary) + pickle.loads(blob) # <- the RCE primitive, unguarded + assert os.path.isdir(canary), "unguarded pickle.loads did NOT execute" + + canary2 = os.path.join(d, "pwned2") + blob2 = pickle.dumps(_Payload(canary2)) + msg = _raises(ArtifactRejected, scan_pickle, blob2, where="payload") + assert "makedirs" in msg, msg + assert not os.path.exists(canary2), "scan_pickle must not execute anything" + + +def test_malicious_pickle_is_refused_at_every_protocol(): + """Protocol 2/3 emit GLOBAL, protocol 4/5 emit STACK_GLOBAL - the scanner + has to cover both branches, and an attacker picks the protocol. + + (This test exists because the negative control found it missing: disabling + the GLOBAL branch alone still left the suite green.) + """ + import pickletools + + seen = set() + for proto in (2, 3, 4, 5): + blob = pickle.dumps(_Payload("/tmp/winnow-guard-never-created"), protocol=proto) + seen.update(op.name for op, _a, _p in pickletools.genops(blob) + if op.name in ("GLOBAL", "STACK_GLOBAL")) + msg = _raises(ArtifactRejected, scan_pickle, blob, where=f"proto{proto}") + assert "makedirs" in msg, (proto, msg) + assert seen == {"GLOBAL", "STACK_GLOBAL"}, f"both branches must be exercised: {seen}" + assert not os.path.exists("/tmp/winnow-guard-never-created") + + +def test_benign_state_dict_pickle_passes(): + blob = pickle.dumps(collections.OrderedDict([("a.weight", [1.0, 2.0])])) + scan_pickle(blob) # must not raise + + +def test_malicious_pickle_inside_a_torch_zip_checkpoint_is_refused(): + """A real `.bin` is a zip with `archive/data.pkl` inside; scan every member.""" + with _tmpdir() as d: + ckpt = os.path.join(d, "pytorch_model.bin") + with zipfile.ZipFile(ckpt, "w") as zf: + zf.writestr("archive/data.pkl", pickle.dumps(_Payload(os.path.join(d, "x")))) + zf.writestr("archive/data/0", b"\x00" * 8) + msg = _raises(ArtifactRejected, scan_checkpoint, ckpt) + assert "data.pkl" in msg and "makedirs" in msg, msg + + +def test_truncated_pickle_is_refused_not_ignored(): + blob = pickle.dumps(collections.OrderedDict([("a", 1)])) + _raises(ArtifactRejected, scan_pickle, blob[: len(blob) // 2]) + + +def test_zip_without_a_pickle_member_is_refused(): + with _tmpdir() as d: + ckpt = os.path.join(d, "weights.bin") + with zipfile.ZipFile(ckpt, "w") as zf: + zf.writestr("archive/data/0", b"\x00") + _raises(ArtifactRejected, scan_checkpoint, ckpt) + + +# --------------------------------------------------------------------------- # +# 2. provenance and load flags # +# --------------------------------------------------------------------------- # +def test_pinned_revision_is_a_sha_and_unknown_ids_are_refused(): + rev = pinned_revision("Qwen/Qwen2.5-7B-Instruct") + assert len(rev) == 40 and all(c in "0123456789abcdef" for c in rev), rev + _raises(ArtifactRejected, pinned_revision, "attacker/backdoored-model") + + +def test_guarded_kwargs_forces_the_flags_and_refuses_opt_out(): + kw = guarded_kwargs("BAAI/bge-small-en-v1.5") + assert kw["revision"] == pinned_revision("BAAI/bge-small-en-v1.5") + assert kw["use_safetensors"] is True + assert kw["trust_remote_code"] is False + _raises(ArtifactRejected, guarded_kwargs, "BAAI/bge-small-en-v1.5", + trust_remote_code=True) + _raises(ArtifactRejected, guarded_kwargs, "BAAI/bge-small-en-v1.5", + use_safetensors=False) + + +# --------------------------------------------------------------------------- # +# 3. snapshot verification BEFORE activation # +# --------------------------------------------------------------------------- # +def _good_snapshot(root): + _write(os.path.join(root, "config.json"), b'{"hidden_size": 8}') + _write(os.path.join(root, "model.safetensors"), b"tensor-bytes") + return build_manifest(root, "BAAI/bge-small-en-v1.5", + pinned_revision("BAAI/bge-small-en-v1.5")) + + +def test_verify_snapshot_accepts_the_snapshot_it_was_built_from(): + with _tmpdir() as d: + verify_snapshot(d, _good_snapshot(d)) + + +def test_verify_snapshot_refuses_a_tampered_file(): + with _tmpdir() as d: + m = _good_snapshot(d) + _write(os.path.join(d, "model.safetensors"), b"tensor-bytez") # 1 byte flipped + msg = _raises(ArtifactRejected, verify_snapshot, d, m) + assert "sha256" in msg, msg + + +def test_verify_snapshot_refuses_missing_and_extra_files(): + with _tmpdir() as d: + m = _good_snapshot(d) + _write(os.path.join(d, "surprise.json"), b"{}") + assert "unexpected files" in _raises(ArtifactRejected, verify_snapshot, d, m) + os.remove(os.path.join(d, "surprise.json")) + os.remove(os.path.join(d, "config.json")) + assert "missing files" in _raises(ArtifactRejected, verify_snapshot, d, m) + + +def test_verify_snapshot_refuses_a_snapshot_shipping_python(): + with _tmpdir() as d: + m = _good_snapshot(d) + _write(os.path.join(d, "modeling_custom.py"), b"import os\n") + msg = _raises(ArtifactRejected, verify_snapshot, d, m) + assert "executable Python" in msg, msg + + +def test_assert_no_pickled_weights(): + with _tmpdir() as d: + _write(os.path.join(d, "model.safetensors"), b"ok") + assert assert_no_pickled_weights(d) == ["model.safetensors"] + _write(os.path.join(d, "pytorch_model.bin"), b"pickled") + msg = _raises(ArtifactRejected, assert_no_pickled_weights, d) + assert "pytorch_model.bin" in msg, msg + + +# --------------------------------------------------------------------------- # +# 4. tensor validation BEFORE activation # +# --------------------------------------------------------------------------- # +_SPEC = { + "order": ["enc.weight", "enc.bias"], + "tensors": { + "enc.weight": {"shape": [2, 3], "dtype": "float32"}, + "enc.bias": {"shape": [2], "dtype": "float32"}, + }, +} + + +def _good_state(): + return collections.OrderedDict([ + ("enc.weight", FakeTensor([2, 3])), + ("enc.bias", FakeTensor([2])), + ]) + + +def test_validate_tensors_accepts_a_matching_state_dict(): + validate_tensors(_good_state(), _SPEC) + + +def test_validate_tensors_refuses_wrong_shape_dtype_and_missing_key(): + s = _good_state() + s["enc.bias"] = FakeTensor([3]) + assert "shape" in _raises(ArtifactRejected, validate_tensors, s, _SPEC) + + s = _good_state() + s["enc.bias"] = FakeTensor([2], dtype="int8") + assert "dtype" in _raises(ArtifactRejected, validate_tensors, s, _SPEC) + + s = _good_state() + del s["enc.bias"] + assert "missing tensors" in _raises(ArtifactRejected, validate_tensors, s, _SPEC) + + s = _good_state() + s["enc.extra"] = FakeTensor([1]) + assert "unexpected tensors" in _raises(ArtifactRejected, validate_tensors, s, _SPEC) + + +def test_validate_tensors_refuses_non_finite_values(): + s = _good_state() + s["enc.bias"] = FakeTensor([2], values=[0.1, float("nan")]) + assert "NaN or Inf" in _raises(ArtifactRejected, validate_tensors, s, _SPEC) + s["enc.bias"] = FakeTensor([2], values=[0.1, float("inf")]) + assert "NaN or Inf" in _raises(ArtifactRejected, validate_tensors, s, _SPEC) + + +def test_validate_tensors_refuses_permuted_feature_order(): + """Same keys, same shapes, wrong order - loads fine, then produces garbage.""" + s = collections.OrderedDict([ + ("enc.bias", FakeTensor([2])), + ("enc.weight", FakeTensor([2, 3])), + ]) + assert "key order" in _raises(ArtifactRejected, validate_tensors, s, _SPEC) + + +# --------------------------------------------------------------------------- # +# 5. known-good rollback # +# --------------------------------------------------------------------------- # +def test_rejected_candidate_leaves_the_previous_known_good_in_place(): + with _tmpdir() as d: + store_path = os.path.join(d, "known_good.json") + good_dir = os.path.join(d, "v1") + good = _good_snapshot(good_dir) + + store = KnownGood(store_path) + store.activate(good_dir, good) + assert store.current(good["model_id"])["revision"] == good["revision"] + + # a tampered candidate at the same path must not overwrite the record + bad_dir = os.path.join(d, "v2") + bad = _good_snapshot(bad_dir) + bad["revision"] = "0" * 40 + _write(os.path.join(bad_dir, "model.safetensors"), b"tampered") + _raises(ArtifactRejected, store.activate, bad_dir, bad) + + assert store.current(good["model_id"])["revision"] == good["revision"] + reopened = KnownGood(store_path) # and it survived on disk + assert reopened.current(good["model_id"])["revision"] == good["revision"] + assert json.load(open(store_path))[good["model_id"]]["revision"] == good["revision"] + + +# --------------------------------------------------------------------------- # +# 6. WIRING: what the real load sites pass to from_pretrained # +# --------------------------------------------------------------------------- # +def _install_fake_transformers(): + """Fake torch + transformers so the real loader classes can be constructed + on a CPU-only box. Returns the list every from_pretrained call is recorded + into.""" + calls = [] + + torch = types.ModuleType("torch") + torch.float16 = "torch.float16" + torch.bfloat16 = "torch.bfloat16" + nn = types.ModuleType("torch.nn") + nn.functional = types.ModuleType("torch.nn.functional") + torch.nn = nn + sys.modules.update({"torch": torch, "torch.nn": nn, + "torch.nn.functional": nn.functional}) + + class _Loaded: + config = types.SimpleNamespace(num_hidden_layers=2) + + def eval(self): + return self + + def half(self): + return self + + def to(self, *a, **k): + return self + + def _auto(name): + return type(name, (), { + "from_pretrained": classmethod( + lambda cls, model_id, **kw: (calls.append((name, model_id, kw)), _Loaded())[1] + ) + }) + + tf = types.ModuleType("transformers") + for n in ("AutoTokenizer", "AutoModel", "AutoModelForCausalLM", + "AutoModelForSequenceClassification"): + setattr(tf, n, _auto(n)) + sys.modules["transformers"] = tf + return calls + + +def _assert_guarded(calls, expect_n): + assert len(calls) == expect_n, f"expected {expect_n} loads, saw {len(calls)}" + for name, model_id, kw in calls: + assert kw.get("revision") == pinned_revision(model_id), \ + f"{name}({model_id}) not pinned: revision={kw.get('revision')!r}" + assert kw.get("trust_remote_code") is False, \ + f"{name}({model_id}) did not disable remote code" + if "Tokenizer" not in name: + assert kw.get("use_safetensors") is True, \ + f"{name}({model_id}) did not force safetensors" + + +def test_two_stage_compressor_load_sites_are_guarded(): + calls = _install_fake_transformers() + import two_stage_compressor as tsc + + tsc.SmallEmbedder(device="cpu", use_fp16=False) + tsc.CrossEncoderReranker("BAAI/bge-reranker-v2-m3", device="cpu", use_fp16=False) + _assert_guarded(calls, 4) + + +def test_attentionrag_hf_backend_load_sites_are_guarded(): + calls = _install_fake_transformers() + from attentionrag.hf_backend import HFBackend + + HFBackend(device="cpu") + _assert_guarded(calls, 2) + + +# --------------------------------------------------------------------------- # +# 7. STATIC SWEEP: no unguarded artifact load anywhere in the repo # +# --------------------------------------------------------------------------- # +_REPO = os.path.dirname(os.path.abspath(__file__)) +_SKIP_DIRS = {".git", "__pycache__", "node_modules", ".next", ".agent-work", + ".venv", "venv"} + + +def _repo_py_files(): + for root, dirs, names in os.walk(_REPO): + dirs[:] = [d for d in dirs if d not in _SKIP_DIRS] + for n in sorted(names): + # model_guard.py holds the ONLY sanctioned raw calls (they are the + # wrappers); this test file only names them in assertions. + if n.endswith(".py") and n not in (os.path.basename(__file__), + "model_guard.py"): + yield os.path.join(root, n) + + +def test_no_unguarded_from_pretrained_or_snapshot_download_in_repo(): + """A new load site added without the guard is a regression, so fail on it. + + Catches `X.from_pretrained(...)` (must be `guarded_from_pretrained`), bare + `snapshot_download(...)` (must be `pinned_snapshot_download`), and a + `PromptCompressor(...)` without `model_config` (llmlingua 0.2.2 defaults + trust_remote_code to True). + """ + import ast + + offenders = [] + for path in _repo_py_files(): + rel = os.path.relpath(path, _REPO) + with open(path) as fh: + src = fh.read() + tree = ast.parse(src, filename=rel) + for node in ast.walk(tree): + if not isinstance(node, ast.Call): + continue + fn = node.func + if isinstance(fn, ast.Attribute) and fn.attr == "from_pretrained": + offenders.append(f"{rel}:{node.lineno} bare .from_pretrained()") + elif isinstance(fn, ast.Name): + if fn.id == "snapshot_download": + offenders.append(f"{rel}:{node.lineno} bare snapshot_download()") + elif fn.id == "PromptCompressor" and not any( + kw.arg == "model_config" for kw in node.keywords + ): + offenders.append( + f"{rel}:{node.lineno} PromptCompressor() without model_config" + ) + assert not offenders, "unguarded artifact loads:\n " + "\n ".join(offenders) + + +def _run_all(): + tests = [v for k, v in sorted(globals().items()) if k.startswith("test_")] + for t in tests: + t() + print(f" ok {t.__name__}") + print(f"\n{len(tests)}/{len(tests)} model_guard tests passed.") + + +if __name__ == "__main__": + _run_all() From 7eee2075944e8c48c43779d9b2bafc04d7b2d9e1 Mon Sep 17 00:00:00 2001 From: lgoyal6 Date: Sun, 6 Sep 2026 16:08:27 -0700 Subject: [PATCH 2/2] Route every weight load through the artifact guard The guard is only worth having if nothing bypasses it, so all seventeen load sites across ten files now go through it instead of calling the hub directly. * fourteen from_pretrained call sites become guarded_from_pretrained(...), which injects the pinned revision, use_safetensors=True and trust_remote_code=False: attentionrag/hf_backend.py, two_stage_compressor.py (both SmallEmbedder and CrossEncoderReranker), experiments/bench/bench_llm_modal.py, experiments/bench/run_compress.py, experiments/bench/compress_devpost.py, turboquant_modal.py and turboquant-poc/modal_app.py. * three snapshot_download call sites become pinned_snapshot_download, and where the snapshot is the thing that later gets loaded it is wrapped in assert_no_pickled_weights so a pickled checkpoint in the cache volume fails the container start rather than the inference: attentionrag/modal_app.py, llmlingua2_modal.py, experiments/eval_modal.py. * llmlingua 0.2.2 defaults trust_remote_code to True inside PromptCompressor, which the caller cannot reach through kwargs, so the two llmlingua entrypoints pass llmlingua_model_config(...) to set it back to False explicitly. * every Modal image that loads a model now ships model_guard alongside its own sources via add_local_python_source, otherwise the guarded import would not resolve inside the container. After this, a repo-wide grep for a bare .from_pretrained( outside model_guard finds nothing, which is what the AST sweep in test_model_guard.py asserts. No behaviour changes beyond the guard: same models, same dtypes, same device maps. Prose em dashes in touched comments are replaced with hyphens, per the repository's style. --- attentionrag/hf_backend.py | 10 +++++++--- attentionrag/modal_app.py | 10 ++++++---- experiments/bench/bench_llm_modal.py | 12 +++++++++--- experiments/bench/compress_devpost.py | 6 ++++-- experiments/bench/run_compress.py | 6 ++++-- experiments/eval_modal.py | 12 +++++++----- llmlingua2_modal.py | 26 +++++++++++++++++--------- turboquant-poc/modal_app.py | 12 +++++++++--- turboquant_modal.py | 12 +++++++++--- two_stage_compressor.py | 21 ++++++++++++++------- 10 files changed, 86 insertions(+), 41 deletions(-) diff --git a/attentionrag/hf_backend.py b/attentionrag/hf_backend.py index 814ccaf..451a121 100644 --- a/attentionrag/hf_backend.py +++ b/attentionrag/hf_backend.py @@ -46,15 +46,19 @@ def __init__( import torch from transformers import AutoModelForCausalLM, AutoTokenizer + from model_guard import guarded_from_pretrained + self.torch = torch self.device = device self.model_name = model_name self.use_openai_hint = use_openai_hint self.openai_model = openai_model - self.tokenizer = AutoTokenizer.from_pretrained(model_name) + # Pinned revision + safetensors only + no remote code; see model_guard. + self.tokenizer = guarded_from_pretrained(AutoTokenizer, model_name) # eager attention is REQUIRED to get attention weights back from forward. - self.model = AutoModelForCausalLM.from_pretrained( + self.model = guarded_from_pretrained( + AutoModelForCausalLM, model_name, torch_dtype=getattr(torch, dtype), attn_implementation="eager", @@ -249,7 +253,7 @@ def compress_spans( # chunk to fall back on, and a "none" anchor carries no attention signal # to rank sentences). In that single-chunk case keep the chunk wholesale # instead of dropping it, so the downstream merge still has spans to work - # with. Multi-chunk inputs keep the normal per-chunk "none" drop — that + # with. Multi-chunk inputs keep the normal per-chunk "none" drop - that # coarse relevance gate is the point of AttentionRAG on long contexts. single_chunk = len(ids) <= chunk_size for i in range(0, len(ids), chunk_size): diff --git a/attentionrag/modal_app.py b/attentionrag/modal_app.py index ecfa4d7..b44f4ab 100644 --- a/attentionrag/modal_app.py +++ b/attentionrag/modal_app.py @@ -48,8 +48,8 @@ "openai", ) .env({"HF_HOME": CACHE_DIR, "HF_HUB_ENABLE_HF_TRANSFER": "1"}) - # Ship the AttentionRAG package into the image. - .add_local_python_source("attentionrag") + # Ship the AttentionRAG package + the artifact guard into the image. + .add_local_python_source("attentionrag", "model_guard") ) app = modal.App("attentionrag", image=image) @@ -65,10 +65,12 @@ class AttentionRAGService: @modal.enter() def load(self): - from huggingface_hub import snapshot_download + from model_guard import assert_no_pickled_weights, pinned_snapshot_download # Populate-once-then-read: only downloads if the volume lacks the model. - snapshot_download(MODEL_NAME) + # Pinned revision, and refuse the snapshot if a pickled checkpoint is in + # it -- the cache volume is mutable and outlives every image rebuild. + assert_no_pickled_weights(pinned_snapshot_download(MODEL_NAME)) hf_cache_vol.commit() from attentionrag.core import AttentionRAG diff --git a/experiments/bench/bench_llm_modal.py b/experiments/bench/bench_llm_modal.py index 07dd3b4..ac31550 100644 --- a/experiments/bench/bench_llm_modal.py +++ b/experiments/bench/bench_llm_modal.py @@ -26,6 +26,8 @@ "scipy", "numpy", ) .env({"HF_HOME": CACHE_DIR}) + # Artifact provenance guard (pinned revision, safetensors only, no remote code). + .add_local_python_source("model_guard") ) app = modal.App("winnow-bench-llm", image=image) @@ -231,11 +233,15 @@ def load(self): import torch from transformers import AutoModelForCausalLM, AutoTokenizer + from model_guard import guarded_from_pretrained + token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN") self.torch = torch - self.tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, token=token) - self.model = AutoModelForCausalLM.from_pretrained( - MODEL_ID, dtype=torch.float16, device_map="cuda", + # Pinned revision + safetensors only + no remote code; see model_guard. + self.tokenizer = guarded_from_pretrained(AutoTokenizer, MODEL_ID, token=token) + self.model = guarded_from_pretrained( + AutoModelForCausalLM, MODEL_ID, + dtype=torch.float16, device_map="cuda", low_cpu_mem_usage=True, token=token, ) self.model.eval() diff --git a/experiments/bench/compress_devpost.py b/experiments/bench/compress_devpost.py index 21e42c1..65ecf5f 100644 --- a/experiments/bench/compress_devpost.py +++ b/experiments/bench/compress_devpost.py @@ -95,8 +95,10 @@ def merged_prompt(m): try: from transformers import AutoTokenizer - tok = AutoTokenizer.from_pretrained( - "microsoft/llmlingua-2-xlm-roberta-large-meetingbank" + from model_guard import guarded_from_pretrained + + tok = guarded_from_pretrained( + AutoTokenizer, "microsoft/llmlingua-2-xlm-roberta-large-meetingbank" ) except Exception as exc: print(f"[tok] whitespace fallback ({exc})", flush=True) diff --git a/experiments/bench/run_compress.py b/experiments/bench/run_compress.py index fba35b1..d91a213 100644 --- a/experiments/bench/run_compress.py +++ b/experiments/bench/run_compress.py @@ -48,8 +48,10 @@ def load_tokenizer(): try: from transformers import AutoTokenizer - _TOK = AutoTokenizer.from_pretrained( - "microsoft/llmlingua-2-xlm-roberta-large-meetingbank" + from model_guard import guarded_from_pretrained + + _TOK = guarded_from_pretrained( + AutoTokenizer, "microsoft/llmlingua-2-xlm-roberta-large-meetingbank" ) _TOK_KIND = "xlm-roberta" print("[tok] loaded xlm-roberta tokenizer", flush=True) diff --git a/experiments/eval_modal.py b/experiments/eval_modal.py index 97a27ce..79b1b6f 100644 --- a/experiments/eval_modal.py +++ b/experiments/eval_modal.py @@ -21,7 +21,7 @@ import modal # This eval lives in experiments/ but imports two_stage_compressor.py from the -# repo root — put the root on the path so it resolves from either location. +# repo root - put the root on the path so it resolves from either location. sys.path.insert(0, os.path.dirname(os.path.dirname(os.path.abspath(__file__)))) MODEL_NAME = "microsoft/llmlingua-2-xlm-roberta-large-meetingbank" @@ -42,7 +42,7 @@ "hf_transfer", "numpy", "accelerate", "sentencepiece", "protobuf", ) .env({"HF_HOME": CACHE_DIR, "HF_HUB_ENABLE_HF_TRANSFER": "1"}) - .add_local_python_source("two_stage_compressor") + .add_local_python_source("two_stage_compressor", "model_guard") ) app = modal.App("winnow-eval", image=image) @@ -64,19 +64,21 @@ def chunk_by_sentences(text: str, k: int = 3): class Compressor: @modal.enter() def load(self): - from huggingface_hub import snapshot_download + from model_guard import llmlingua_model_config, pinned_snapshot_download for name in (MODEL_NAME, RERANKER_NAME, EMBEDDER_NAME): - snapshot_download(name) + pinned_snapshot_download(name) # pinned revision; see model_guard hf_cache_vol.commit() from llmlingua import PromptCompressor from two_stage_compressor import CrossEncoderReranker, SmallEmbedder - # Encoder token classifier (LLMLingua-2). No causal/SLM backbone — the + # Encoder token classifier (LLMLingua-2). No causal/SLM backbone - the # LongLLMLingua path was cancelled (wrong regime for short single docs). self.compressor = PromptCompressor( model_name=MODEL_NAME, use_llmlingua2=True, device_map="cuda", + # llmlingua 0.2.2 defaults trust_remote_code to True. See model_guard. + model_config=llmlingua_model_config(MODEL_NAME), ) self.reranker = CrossEncoderReranker(RERANKER_NAME, device="cuda", use_fp16=True) self.embedder = SmallEmbedder(EMBEDDER_NAME, device="cuda", use_fp16=True) diff --git a/llmlingua2_modal.py b/llmlingua2_modal.py index e10d0cb..4a474e6 100644 --- a/llmlingua2_modal.py +++ b/llmlingua2_modal.py @@ -17,12 +17,12 @@ # Token-level extractive compressor (encoder). Multilingual XLM-RoBERTa. MODEL_NAME = "microsoft/llmlingua-2-xlm-roberta-large-meetingbank" # Coarse-stage, question-aware reranker (cross-encoder, NOT a causal LM). -# bge-reranker-v2-m3: current lightweight multilingual BGE reranker — pairs well +# bge-reranker-v2-m3: current lightweight multilingual BGE reranker - pairs well # with the multilingual compressor above. See two_stage_compressor.py for why. RERANKER_NAME = "BAAI/bge-reranker-v2-m3" # Persistent HF cache. Weights are downloaded into this volume once and read -# from it on every subsequent run — never re-downloaded, even across image +# from it on every subsequent run - never re-downloaded, even across image # rebuilds. create_if_missing=True means the first run provisions it automatically. CACHE_DIR = "/cache" hf_cache_vol = modal.Volume.from_name("llmlingua2-hf-cache", create_if_missing=True) @@ -38,8 +38,8 @@ # Point Hugging Face at the mounted volume so snapshot_download and the models # both read/write the standard HF cache layout there. .env({"HF_HOME": CACHE_DIR, "HF_HUB_ENABLE_HF_TRANSFER": "1"}) - # Ship our local two-stage logic into the image. - .add_local_python_source("two_stage_compressor") + # Ship our local two-stage logic + the artifact guard into the image. + .add_local_python_source("two_stage_compressor", "model_guard") ) app = modal.App("llmlingua2-xlm", image=image) @@ -50,7 +50,7 @@ volumes={CACHE_DIR: hf_cache_vol}, # Keep a warmed container alive 30 min after the last request so it stays hot # through a demo (between questions) without re-warming. Auto-scales to zero - # afterward — no lingering cost. + # afterward - no lingering cost. scaledown_window=1800, # Memory snapshots: capture the fully-loaded model (incl. GPU memory) so that # future cold starts RESTORE that state instead of re-loading the model. @@ -60,17 +60,23 @@ class Compressor: @modal.enter(snap=True) def load(self): - # Runs only when CREATING the snapshot — i.e. the very first cold start, + # Runs only when CREATING the snapshot - i.e. the very first cold start, # or after a code/image change invalidates the existing snapshot. The # loaded model and its GPU memory are captured here; every later cold # start restores this state directly and skips all of this work. # # Populate-once-then-read: downloads only if the volume is empty, # otherwise this resolves straight from the cached volume. - from huggingface_hub import snapshot_download + from model_guard import ( + assert_no_pickled_weights, + llmlingua_model_config, + pinned_snapshot_download, + ) - snapshot_download(MODEL_NAME) - snapshot_download(RERANKER_NAME) + # Pinned revisions: the cache volume outlives every image rebuild, so + # without a pin "the same model name" can mean different bytes later. + for name in (MODEL_NAME, RERANKER_NAME): + assert_no_pickled_weights(pinned_snapshot_download(name)) hf_cache_vol.commit() # persist any newly downloaded files to the volume from llmlingua import PromptCompressor @@ -79,6 +85,8 @@ def load(self): model_name=MODEL_NAME, use_llmlingua2=True, device_map="cuda", + # llmlingua 0.2.2 defaults trust_remote_code to True. See model_guard. + model_config=llmlingua_model_config(MODEL_NAME), ) # Coarse-stage reranker (raw transformers cross-encoder; see module docs diff --git a/turboquant-poc/modal_app.py b/turboquant-poc/modal_app.py index f9b7a25..7c58ad8 100644 --- a/turboquant-poc/modal_app.py +++ b/turboquant-poc/modal_app.py @@ -39,6 +39,8 @@ "numpy", ) .env({"HF_HOME": "/cache/hf"}) + # Artifact provenance guard (pinned revision, safetensors only, no remote code). + .add_local_python_source("model_guard") ) app = modal.App("turboquant-poc") @@ -263,6 +265,8 @@ def run(prompt: str, bit_width: int = 4, max_new_tokens: int = 200, import torch from transformers import AutoModelForCausalLM, AutoTokenizer + from model_guard import guarded_from_pretrained + hf_token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN") device = torch.device("cuda") print(f"=== TurboQuant POC ===", flush=True) @@ -271,9 +275,11 @@ def run(prompt: str, bit_width: int = 4, max_new_tokens: int = 200, print(f"TurboQuant bit_width={bit_width} outlier_channels={outlier_channels} outlier_bits={outlier_bits}", flush=True) print("Loading model + tokenizer...", flush=True) - tokenizer = AutoTokenizer.from_pretrained(model_id, token=hf_token) - model = AutoModelForCausalLM.from_pretrained( - model_id, dtype=torch.float16, device_map="cuda", low_cpu_mem_usage=True, token=hf_token + # Pinned revision + safetensors only + no remote code; see model_guard. + tokenizer = guarded_from_pretrained(AutoTokenizer, model_id, token=hf_token) + model = guarded_from_pretrained( + AutoModelForCausalLM, model_id, + dtype=torch.float16, device_map="cuda", low_cpu_mem_usage=True, token=hf_token ) model.eval() if tokenizer.pad_token is None: diff --git a/turboquant_modal.py b/turboquant_modal.py index 1888966..8412ba1 100644 --- a/turboquant_modal.py +++ b/turboquant_modal.py @@ -36,6 +36,8 @@ "scipy", "numpy", ) .env({"HF_HOME": CACHE_DIR}) + # Artifact provenance guard (pinned revision, safetensors only, no remote code). + .add_local_python_source("model_guard") ) app = modal.App("turboquant-qwen14b", image=image) @@ -250,11 +252,15 @@ def load(self): import torch from transformers import AutoModelForCausalLM, AutoTokenizer + from model_guard import guarded_from_pretrained + token = os.environ.get("HF_TOKEN") or os.environ.get("HUGGING_FACE_HUB_TOKEN") self.torch = torch - self.tokenizer = AutoTokenizer.from_pretrained(MODEL_ID, token=token) - self.model = AutoModelForCausalLM.from_pretrained( - MODEL_ID, dtype=torch.float16, device_map="cuda", + # Pinned revision + safetensors only + no remote code; see model_guard. + self.tokenizer = guarded_from_pretrained(AutoTokenizer, MODEL_ID, token=token) + self.model = guarded_from_pretrained( + AutoModelForCausalLM, MODEL_ID, + dtype=torch.float16, device_map="cuda", low_cpu_mem_usage=True, token=token, ) self.model.eval() diff --git a/two_stage_compressor.py b/two_stage_compressor.py index c03d06e..8e2bd04 100644 --- a/two_stage_compressor.py +++ b/two_stage_compressor.py @@ -14,14 +14,14 @@ context, rate, target_token, use_context_level_filter, target_context, context_level_rate, context_level_target_token, ...) -It forwards ONLY those args — `question`, `instruction`, `rank_method`, and +It forwards ONLY those args - `question`, `instruction`, `rank_method`, and `reorder_context` are silently dropped. And `compress_prompt_llmlingua2(...)` has NO `question`/`rank_method`/`reorder_context` parameters at all. So with `use_llmlingua2=True`: * the only document-level knob (`use_context_level_filter` + `target_context`/ `context_level_rate`/`context_level_target_token`) ranks documents by the - encoder's own predicted compression score — it is QUESTION-BLIND; + encoder's own predicted compression score - it is QUESTION-BLIND; * `rank_method` (including `bge_reranker`) only takes effect on the causal LongLLMLingua path (`use_llmlingua2=False`), which loads a causal backbone we are explicitly avoiding. @@ -103,11 +103,16 @@ def __init__(self, model_name: str = "BAAI/bge-small-en-v1.5", import torch from transformers import AutoModel, AutoTokenizer + from model_guard import guarded_from_pretrained + self._torch = torch self.device = device self.max_length = max_length - self.tokenizer = AutoTokenizer.from_pretrained(model_name) - model = AutoModel.from_pretrained(model_name) + # Pinned revision + safetensors only + no remote code. bge-small ships a + # pytorch_model.bin next to model.safetensors, so the pickle path is + # genuinely reachable here -- see model_guard's module docstring. + self.tokenizer = guarded_from_pretrained(AutoTokenizer, model_name) + model = guarded_from_pretrained(AutoModel, model_name) if use_fp16 and device.startswith("cuda"): model = model.half() self.model = model.eval().to(device) @@ -207,11 +212,13 @@ def __init__(self, model_name: str, device: str = "cuda", max_length: int = 512, import torch from transformers import AutoModelForSequenceClassification, AutoTokenizer + from model_guard import guarded_from_pretrained + self._torch = torch self.device = device self.max_length = max_length - self.tokenizer = AutoTokenizer.from_pretrained(model_name) - model = AutoModelForSequenceClassification.from_pretrained(model_name) + self.tokenizer = guarded_from_pretrained(AutoTokenizer, model_name) + model = guarded_from_pretrained(AutoModelForSequenceClassification, model_name) if use_fp16 and device.startswith("cuda"): model = model.half() self.model = model.eval().to(device) @@ -384,7 +391,7 @@ def two_stage_compress( return { "compressed_prompt": final_prompt, - # the compressed CONTEXT only (no instruction/question) — what a reader + # the compressed CONTEXT only (no instruction/question) - what a reader # should be given alongside its own downstream questions: "compressed_context": compressed_context, # context-only numbers straight from LLMLingua-2 (the token stage):