Not a member of Pastebin yet?
Sign Up,
it unlocks many cool features!
- """
- Context-DNA compression: a load-bearing evaluation, not a shape-check.
- What changed vs. the earlier scripts, and why:
- 1. TRAINED, not random-init. Untrained encoder/decoder numbers are meaningless
- (dominated by init scale, not by what the architecture can learn to preserve).
- This script actually runs an optimization loop and reports loss curves.
- 2. Perceiver-style attention bottleneck instead of (a) mean-pooling, which
- collapses token identity before compression even starts, or (b) a single
- flatten-linear layer, whose parameter count explodes (chunk_size * d_model
- * k_dna * 2) and which is just linear PCA in a trenchcoat. The bottleneck
- here uses learned latent queries that cross-attend into the chunk, and
- learned position queries that cross-attend back out. Parameter count is
- independent of chunk_size, so this scales to much larger chunks/contexts.
- 3. A real baseline: PCA (best possible LINEAR compression at the same latent
- budget, closed-form via SVD, zero training needed). If the learned model
- can't beat PCA, the extra machinery isn't earning its complexity.
- 4. Diverse text with an INJECTED, UNIQUE FACT per document (a random 5-digit
- code in a sentence), instead of one repeated sentence. Repetition makes
- compression artificially easy; this doesn't.
- 5. The metric that actually matters: after compress -> decompress, can the
- frozen LM's own output head still predict the correct token at the fact's
- position? Whole-chunk MSE can look good while the 2-3 dimensions that
- actually encode "48213" get washed out. Top-1/top-5 next-token accuracy
- at the fact position is a much sharper, more honest test.
- 6. Averaged over many (document, chunk) pairs, not one query on one sample.
- Usage:
- python bscm_eval.py --mode real --model_id Qwen/Qwen2.5-0.5B --n_docs 40 --steps 800
- python bscm_eval.py --mode synthetic # no internet/model needed, sanity-checks the pipeline
- Requires: torch, and (for --mode real) transformers + internet access to download the model.
- """
- import argparse
- import random
- import string
- import torch
- import torch.nn as nn
- import torch.nn.functional as F
- # --------------------------------------------------------------------------
- # Architecture
- # --------------------------------------------------------------------------
- class LatentBottleneck(nn.Module):
- """Perceiver-style compressive autoencoder for a chunk of hidden states.
- Encoder: n_latents learned queries cross-attend into the S-token chunk
- -> DNA code [n_latents, d_model].
- Decoder: S learned position queries cross-attend into the DNA code
- -> reconstructed chunk [S, d_model].
- Param count depends on n_latents and d_model, NOT on chunk_size, which is
- the key fix over a flatten-linear encoder/decoder.
- """
- def __init__(self, d_model, chunk_size, n_latents=8, n_heads=4):
- super().__init__()
- self.chunk_size = chunk_size
- self.n_latents = n_latents
- self.enc_latents = nn.Parameter(torch.randn(n_latents, d_model) * 0.02)
- self.enc_attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True)
- self.enc_ff = nn.Sequential(nn.Linear(d_model, d_model), nn.GELU(), nn.Linear(d_model, d_model))
- self.enc_ln1 = nn.LayerNorm(d_model)
- self.enc_ln2 = nn.LayerNorm(d_model)
- self.dec_queries = nn.Parameter(torch.randn(chunk_size, d_model) * 0.02)
- self.dec_attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True)
- self.dec_ff = nn.Sequential(nn.Linear(d_model, d_model), nn.GELU(), nn.Linear(d_model, d_model))
- self.dec_ln1 = nn.LayerNorm(d_model)
- self.dec_ln2 = nn.LayerNorm(d_model)
- def encode(self, chunk):
- B = chunk.shape[0]
- latents = self.enc_latents.unsqueeze(0).expand(B, -1, -1)
- attended, _ = self.enc_attn(latents, chunk, chunk)
- latents = self.enc_ln1(latents + attended)
- latents = self.enc_ln2(latents + self.enc_ff(latents))
- return latents
- def decode(self, dna):
- B = dna.shape[0]
- queries = self.dec_queries.unsqueeze(0).expand(B, -1, -1)
- attended, _ = self.dec_attn(queries, dna, dna)
- out = self.dec_ln1(queries + attended)
- out = self.dec_ln2(out + self.dec_ff(out))
- return out
- def forward(self, chunk):
- dna = self.encode(chunk)
- return self.decode(dna), dna
- def num_params(self):
- return sum(p.numel() for p in self.parameters())
- class PCABaseline:
- """Best possible LINEAR compressor at a given latent budget. Fit once via
- SVD on training chunks (no gradient descent needed -- this is the
- closed-form optimum for linear reconstruction, so it's a fair floor for
- "is the learned nonlinear model actually buying us anything")."""
- def __init__(self, n_components):
- self.n_components = n_components
- self.mean = None
- self.components = None # [d_model, n_components]
- def fit(self, flat_chunks): # flat_chunks: [N, d_model]
- self.mean = flat_chunks.mean(dim=0, keepdim=True)
- centered = flat_chunks - self.mean
- U, S, V = torch.pca_lowrank(centered, q=self.n_components)
- self.components = V[:, : self.n_components] # [d_model, n_components]
- def reconstruct(self, flat_chunks):
- centered = flat_chunks - self.mean
- coeffs = centered @ self.components # [N, n_components]
- recon = coeffs @ self.components.T + self.mean
- return recon
- class MeanPoolBaseline(nn.Module):
- """The original naive approach, kept as a baseline so the improvement is
- quantified rather than assumed."""
- def __init__(self, d_model, chunk_size, k_dna):
- super().__init__()
- self.enc = nn.Sequential(nn.Linear(d_model, k_dna), nn.Tanh())
- self.dec = nn.Linear(k_dna, chunk_size * d_model)
- self.chunk_size = chunk_size
- self.d_model = d_model
- def forward(self, chunk): # chunk: [B, S, D]
- pooled = chunk.mean(dim=1)
- dna = self.enc(pooled)
- recon = self.dec(dna).view(-1, self.chunk_size, self.d_model)
- return recon, dna
- def num_params(self):
- return sum(p.numel() for p in self.parameters())
- # --------------------------------------------------------------------------
- # Fact localization (verified separately against HF's offset_mapping convention)
- # --------------------------------------------------------------------------
- def find_fact_token_indices(offset_mapping, fact_char_start, fact_char_end):
- idxs = []
- for i, (s, e) in enumerate(offset_mapping):
- if s == e:
- continue
- if s < fact_char_end and e > fact_char_start:
- idxs.append(i)
- return idxs
- # --------------------------------------------------------------------------
- # Data: real LM hidden states with an injected, unique fact per document
- # --------------------------------------------------------------------------
- FILLER_SENTENCES = [
- "The quarterly report covers regional sales performance in detail.",
- "Weather patterns this year have been unusually volatile across the coast.",
- "The committee reviewed several proposals before reaching a decision.",
- "Engineers debated the merits of the new cooling system for hours.",
- "Historical records suggest the settlement was founded in the early period.",
- "The recipe calls for a slow simmer over low heat for best results.",
- "Traffic congestion has worsened following the recent construction.",
- "The novel explores themes of memory and displacement across generations.",
- "Researchers collected samples from twelve different locations.",
- "The orchestra rehearsed the symphony's second movement extensively.",
- ]
- def make_document(chunk_size, n_chunks, fact_chunk_idx, rng):
- """Build a document of n_chunks * chunk_size (approx, in words) with a
- unique injected fact sentence placed in one designated chunk."""
- words_per_chunk_target = chunk_size # rough proxy; real tokenization will differ
- fact_code = "".join(rng.choice(string.digits) for _ in range(5))
- fact_sentence = f" The secret access code is {fact_code}."
- chunks_text = []
- for c in range(n_chunks):
- filler = " ".join(rng.choice(FILLER_SENTENCES) for _ in range(max(3, words_per_chunk_target // 12)))
- if c == fact_chunk_idx:
- filler = filler + fact_sentence
- chunks_text.append(filler)
- full_text = " ".join(chunks_text)
- return full_text, fact_code
- def build_real_dataset(model_id, n_docs, chunk_size, n_chunks_per_doc, device, seed=0):
- from transformers import AutoTokenizer, AutoModelForCausalLM
- rng = random.Random(seed)
- print(f"Loading {model_id} ...")
- tokenizer = AutoTokenizer.from_pretrained(model_id)
- model = AutoModelForCausalLM.from_pretrained(model_id)
- model.to(device).eval()
- docs = []
- for d in range(n_docs):
- fact_chunk_idx = rng.randrange(n_chunks_per_doc)
- text, fact_code = make_document(chunk_size, n_chunks_per_doc, fact_chunk_idx, rng)
- enc = tokenizer(text, return_tensors="pt", return_offsets_mapping=True, truncation=True,
- max_length=chunk_size * n_chunks_per_doc)
- input_ids = enc.input_ids.to(device)
- offsets = enc.offset_mapping[0].tolist()
- with torch.no_grad():
- out = model(input_ids, output_hidden_states=True)
- hidden = out.hidden_states[-1][0] # [seq_len, d_model], pre-final-norm
- seq_len = hidden.shape[0]
- n_full_chunks = seq_len // chunk_size
- if n_full_chunks < 2:
- continue
- hidden = hidden[: n_full_chunks * chunk_size]
- fact_char_start = text.find(fact_code)
- fact_char_end = fact_char_start + len(fact_code)
- fact_token_idxs = find_fact_token_indices(offsets, fact_char_start, fact_char_end)
- fact_token_idxs = [i for i in fact_token_idxs if i < n_full_chunks * chunk_size]
- if not fact_token_idxs:
- continue
- fact_chunk = fact_token_idxs[0] // chunk_size
- fact_local_idxs = [i % chunk_size for i in fact_token_idxs if i // chunk_size == fact_chunk]
- fact_target_ids = input_ids[0, fact_token_idxs].tolist()
- docs.append({
- "hidden": hidden.detach(),
- "n_chunks": n_full_chunks,
- "fact_chunk": fact_chunk,
- "fact_local_idxs": fact_local_idxs,
- "fact_target_ids": fact_target_ids,
- })
- return docs, model, tokenizer
- def build_synthetic_dataset(d_model, n_docs, chunk_size, n_chunks_per_doc, device, seed=0):
- """No internet needed. Hidden states are an AR(1) process per feature
- (so nearby tokens are correlated, like real embeddings, unlike iid noise).
- The 'fact' is a distinctive fixed vector pattern injected at one token
- position; retrieval is judged by nearest-neighbor match against a small
- codebook of decoy vectors, standing in for the LM vocabulary/head."""
- g = torch.Generator().manual_seed(seed)
- codebook_size = 50
- codebook = torch.randn(codebook_size, d_model, generator=g) * 2.0
- docs = []
- for d in range(n_docs):
- seq_len = chunk_size * n_chunks_per_doc
- noise = torch.randn(seq_len, d_model, generator=g) * 0.3
- hidden = torch.cumsum(noise, dim=0)
- hidden = hidden - hidden.mean(dim=0, keepdim=True)
- fact_chunk = torch.randint(0, n_chunks_per_doc, (1,), generator=g).item()
- fact_local_idx = torch.randint(0, chunk_size, (1,), generator=g).item()
- fact_code_idx = torch.randint(0, codebook_size, (1,), generator=g).item()
- global_idx = fact_chunk * chunk_size + fact_local_idx
- hidden[global_idx] = codebook[fact_code_idx] + torch.randn(d_model, generator=g) * 0.05
- docs.append({
- "hidden": hidden.to(device),
- "n_chunks": n_chunks_per_doc,
- "fact_chunk": fact_chunk,
- "fact_local_idxs": [fact_local_idx],
- "fact_target_ids": [fact_code_idx],
- })
- return docs, codebook.to(device)
- # --------------------------------------------------------------------------
- # Training + evaluation
- # --------------------------------------------------------------------------
- def train_bottleneck(model, train_chunks, steps, lr, device, batch_size=16, log_every=100):
- model.to(device).train()
- opt = torch.optim.Adam(model.parameters(), lr=lr)
- n = train_chunks.shape[0]
- losses = []
- for step in range(steps):
- idx = torch.randint(0, n, (min(batch_size, n),))
- batch = train_chunks[idx].to(device)
- recon, _ = model(batch)
- loss = F.mse_loss(recon, batch)
- opt.zero_grad()
- loss.backward()
- opt.step()
- losses.append(loss.item())
- if step % log_every == 0 or step == steps - 1:
- print(f" step {step:4d} mse {loss.item():.5f}")
- model.eval()
- return losses
- def evaluate_fact_retrieval_real(model, lm_model, eval_docs, chunk_size, device):
- """Returns dict of aggregate metrics using the LM's own head as the ground
- truth for whether the fact is recoverable, not just low-level MSE."""
- top1_hits, top5_hits, n_facts = 0, 0, 0
- mse_all, mse_fact_pos = [], []
- final_norm = getattr(lm_model.model, "norm", None)
- lm_head = lm_model.get_output_embeddings()
- with torch.no_grad():
- for doc in eval_docs:
- hidden = doc["hidden"]
- c = doc["fact_chunk"]
- chunk = hidden[c * chunk_size:(c + 1) * chunk_size].unsqueeze(0).to(device)
- recon, _ = model(chunk)
- mse_all.append(F.mse_loss(recon, chunk).item())
- local_idxs = doc["fact_local_idxs"]
- target_ids = doc["fact_target_ids"]
- recon_fact = recon[0, local_idxs] # [n_fact_toks, d_model]
- orig_fact = chunk[0, local_idxs]
- mse_fact_pos.append(F.mse_loss(recon_fact, orig_fact).item())
- normed = final_norm(recon_fact) if final_norm is not None else recon_fact
- logits = lm_head(normed) # [n_fact_toks, vocab]
- top5 = logits.topk(5, dim=-1).indices
- for i, tgt in enumerate(target_ids):
- n_facts += 1
- if top5[i, 0].item() == tgt:
- top1_hits += 1
- if tgt in top5[i].tolist():
- top5_hits += 1
- return {
- "chunk_mse": sum(mse_all) / len(mse_all),
- "fact_position_mse": sum(mse_fact_pos) / len(mse_fact_pos),
- "fact_top1_acc": top1_hits / max(n_facts, 1),
- "fact_top5_acc": top5_hits / max(n_facts, 1),
- "n_facts_evaluated": n_facts,
- }
- def evaluate_fact_retrieval_synthetic(model, eval_docs, codebook, chunk_size, device):
- top1_hits, n_facts = 0, 0
- mse_all, mse_fact_pos = [], []
- with torch.no_grad():
- for doc in eval_docs:
- hidden = doc["hidden"]
- c = doc["fact_chunk"]
- chunk = hidden[c * chunk_size:(c + 1) * chunk_size].unsqueeze(0).to(device)
- recon, _ = model(chunk)
- mse_all.append(F.mse_loss(recon, chunk).item())
- local_idxs = doc["fact_local_idxs"]
- target_ids = doc["fact_target_ids"]
- recon_fact = recon[0, local_idxs]
- orig_fact = chunk[0, local_idxs]
- mse_fact_pos.append(F.mse_loss(recon_fact, orig_fact).item())
- sims = F.cosine_similarity(recon_fact.unsqueeze(1), codebook.unsqueeze(0), dim=-1)
- pred = sims.argmax(dim=-1)
- for i, tgt in enumerate(target_ids):
- n_facts += 1
- if pred[i].item() == tgt:
- top1_hits += 1
- return {
- "chunk_mse": sum(mse_all) / len(mse_all),
- "fact_position_mse": sum(mse_fact_pos) / len(mse_fact_pos),
- "fact_top1_acc": top1_hits / max(n_facts, 1),
- "n_facts_evaluated": n_facts,
- }
- def evaluate_pca(pca, eval_docs, chunk_size, device, real_probe=None, lm_model=None, codebook=None):
- mse_all, mse_fact_pos = [], []
- top1_hits, n_facts = 0, 0
- with torch.no_grad():
- for doc in eval_docs:
- hidden = doc["hidden"]
- c = doc["fact_chunk"]
- chunk = hidden[c * chunk_size:(c + 1) * chunk_size].to(device)
- recon = pca.reconstruct(chunk)
- mse_all.append(F.mse_loss(recon, chunk).item())
- local_idxs = doc["fact_local_idxs"]
- target_ids = doc["fact_target_ids"]
- recon_fact = recon[local_idxs]
- orig_fact = chunk[local_idxs]
- mse_fact_pos.append(F.mse_loss(recon_fact, orig_fact).item())
- if real_probe and lm_model is not None:
- final_norm = getattr(lm_model.model, "norm", None)
- lm_head = lm_model.get_output_embeddings()
- normed = final_norm(recon_fact) if final_norm is not None else recon_fact
- logits = lm_head(normed)
- pred = logits.argmax(dim=-1)
- for i, tgt in enumerate(target_ids):
- n_facts += 1
- if pred[i].item() == tgt:
- top1_hits += 1
- elif codebook is not None:
- sims = F.cosine_similarity(recon_fact.unsqueeze(1), codebook.unsqueeze(0), dim=-1)
- pred = sims.argmax(dim=-1)
- for i, tgt in enumerate(target_ids):
- n_facts += 1
- if pred[i].item() == tgt:
- top1_hits += 1
- result = {
- "chunk_mse": sum(mse_all) / len(mse_all),
- "fact_position_mse": sum(mse_fact_pos) / len(mse_fact_pos),
- "fact_top1_acc": top1_hits / max(n_facts, 1),
- "n_facts_evaluated": n_facts,
- }
- return result
- # --------------------------------------------------------------------------
- # Main
- # --------------------------------------------------------------------------
- def main():
- ap = argparse.ArgumentParser()
- ap.add_argument("--mode", choices=["real", "synthetic"], default="synthetic")
- ap.add_argument("--model_id", default="Qwen/Qwen2.5-0.5B")
- ap.add_argument("--n_docs", type=int, default=40)
- ap.add_argument("--chunk_size", type=int, default=64)
- ap.add_argument("--n_chunks_per_doc", type=int, default=6)
- ap.add_argument("--n_latents", type=int, default=8)
- ap.add_argument("--steps", type=int, default=800)
- ap.add_argument("--lr", type=float, default=1e-3)
- ap.add_argument("--seed", type=int, default=0)
- args = ap.parse_args()
- device = "cuda" if torch.cuda.is_available() else "cpu"
- torch.manual_seed(args.seed)
- print(f"device: {device}")
- lm_model = None
- codebook = None
- if args.mode == "real":
- try:
- docs, lm_model, tokenizer = build_real_dataset(
- args.model_id, args.n_docs, args.chunk_size, args.n_chunks_per_doc, device, args.seed)
- if len(docs) < 8:
- raise RuntimeError("too few usable documents (fact localization failed too often)")
- d_model = docs[0]["hidden"].shape[-1]
- except Exception as e:
- print(f"[real mode failed: {e}] -- falling back to synthetic mode. "
- f"Run this on a machine with internet access to huggingface.co for real results.")
- args.mode = "synthetic"
- if args.mode == "synthetic":
- d_model = 896 # matches Qwen2.5-0.5B for apples-to-apples config
- docs, codebook = build_synthetic_dataset(
- d_model, args.n_docs, args.chunk_size, args.n_chunks_per_doc, device, args.seed)
- n_eval = max(1, len(docs) // 5)
- eval_docs, train_docs = docs[:n_eval], docs[n_eval:]
- print(f"docs: {len(docs)} total -> {len(train_docs)} train / {len(eval_docs)} eval")
- all_chunks = []
- for doc in train_docs:
- h = doc["hidden"]
- n_c = doc["n_chunks"]
- for c in range(n_c):
- all_chunks.append(h[c * args.chunk_size:(c + 1) * args.chunk_size])
- train_chunks = torch.stack(all_chunks) # [N, S, D]
- print(f"training chunks: {train_chunks.shape}")
- # --- Model 1: Latent Bottleneck (the proposed architecture) ---
- print("\n=== Training LatentBottleneck ===")
- lb = LatentBottleneck(d_model=d_model, chunk_size=args.chunk_size, n_latents=args.n_latents)
- train_bottleneck(lb, train_chunks, args.steps, args.lr, device)
- # --- Model 2: Mean-pool baseline (the original naive idea) ---
- # NOTE: we do NOT match the LatentBottleneck's total latent scalar budget
- # (n_latents * d_model) here on purpose. A single Linear(k_dna, chunk_size
- # * d_model) decoder has chunk_size * d_model * k_dna parameters, so
- # matching the budget exactly would require ~400M+ params for this decoder
- # alone (confirmed: it OOMs). That blowup is itself the finding -- it's
- # why a flatten/mean-pool decoder doesn't scale, independent of whether
- # the reconstruction quality is any good. We give it k_dna = n_latents
- # scalars instead (a smaller, tractable budget) and still expect the
- # LatentBottleneck to win on fact retrieval despite MeanPool having a much
- # larger raw parameter count in its remaining layers.
- print("\n=== Training MeanPoolBaseline ===")
- mp = MeanPoolBaseline(d_model=d_model, chunk_size=args.chunk_size, k_dna=args.n_latents)
- train_bottleneck(mp, train_chunks, args.steps, args.lr, device)
- # --- Model 3: PCA baseline (closed-form optimal linear compressor) ---
- print("\n=== Fitting PCA baseline ===")
- flat_train = train_chunks.reshape(-1, d_model).to(device)
- pca = PCABaseline(n_components=args.n_latents)
- pca.fit(flat_train)
- # --- Evaluation ---
- print("\n=== Evaluation (held-out documents) ===")
- if args.mode == "real":
- res_lb = evaluate_fact_retrieval_real(lb, lm_model, eval_docs, args.chunk_size, device)
- res_mp = evaluate_fact_retrieval_real(mp, lm_model, eval_docs, args.chunk_size, device)
- res_pca = evaluate_pca(pca, eval_docs, args.chunk_size, device, real_probe=True, lm_model=lm_model)
- else:
- res_lb = evaluate_fact_retrieval_synthetic(lb, eval_docs, codebook, args.chunk_size, device)
- res_mp = evaluate_fact_retrieval_synthetic(mp, eval_docs, codebook, args.chunk_size, device)
- res_pca = evaluate_pca(pca, eval_docs, args.chunk_size, device, codebook=codebook)
- def fmt(r):
- s = f"chunk_mse={r['chunk_mse']:.5f} fact_pos_mse={r['fact_position_mse']:.5f} " \
- f"fact_top1={r['fact_top1_acc']:.2%}"
- if "fact_top5_acc" in r:
- s += f" fact_top5={r['fact_top5_acc']:.2%}"
- s += f" (n_facts={r['n_facts_evaluated']})"
- return s
- print("\n--- RESULTS ---")
- print(f"LatentBottleneck params={lb.num_params():>10,} {fmt(res_lb)}")
- print(f"MeanPoolBaseline params={mp.num_params():>10,} {fmt(res_mp)}")
- print(f"PCA (closed-form) params={'0 (SVD)':>10} {fmt(res_pca)}")
- print(f"\ncompression ratio (elements): {args.chunk_size * d_model / (args.n_latents * d_model):.1f}x")
- print("\nInterpretation: fact_top1_acc is the number that matters. chunk_mse can look")
- print("good while the specific fact detail is unrecoverable -- if fact_top1_acc for")
- print("LatentBottleneck isn't clearly ahead of PCA at the same latent budget, the")
- print("attention machinery isn't earning its complexity yet, and that's the honest")
- print("result to report, not a reason to hide the comparison.")
- if __name__ == "__main__":
- main()
Add Comment
Please, Sign In to add comment