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/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() 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):