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

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
10 changes: 7 additions & 3 deletions attentionrag/hf_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down Expand Up @@ -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):
Expand Down
10 changes: 6 additions & 4 deletions attentionrag/modal_app.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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
Expand Down
12 changes: 9 additions & 3 deletions experiments/bench/bench_llm_modal.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down Expand Up @@ -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()
Expand Down
6 changes: 4 additions & 2 deletions experiments/bench/compress_devpost.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
6 changes: 4 additions & 2 deletions experiments/bench/run_compress.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand Down
12 changes: 7 additions & 5 deletions experiments/eval_modal.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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)
Expand All @@ -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)
Expand Down
26 changes: 17 additions & 9 deletions llmlingua2_modal.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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)
Expand All @@ -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.
Expand All @@ -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
Expand All @@ -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
Expand Down
Loading
Loading