Fahim_s_444

Ai Context DNA

Aug 22nd, 2026
127
0
Never
Not a member of Pastebin yet? Sign Up, it unlocks many cool features!
text 23.56 KB | None | 0 0
  1. """
  2. Context-DNA compression: a load-bearing evaluation, not a shape-check.
  3.  
  4. What changed vs. the earlier scripts, and why:
  5.  
  6. 1. TRAINED, not random-init. Untrained encoder/decoder numbers are meaningless
  7. (dominated by init scale, not by what the architecture can learn to preserve).
  8. This script actually runs an optimization loop and reports loss curves.
  9.  
  10. 2. Perceiver-style attention bottleneck instead of (a) mean-pooling, which
  11. collapses token identity before compression even starts, or (b) a single
  12. flatten-linear layer, whose parameter count explodes (chunk_size * d_model
  13. * k_dna * 2) and which is just linear PCA in a trenchcoat. The bottleneck
  14. here uses learned latent queries that cross-attend into the chunk, and
  15. learned position queries that cross-attend back out. Parameter count is
  16. independent of chunk_size, so this scales to much larger chunks/contexts.
  17.  
  18. 3. A real baseline: PCA (best possible LINEAR compression at the same latent
  19. budget, closed-form via SVD, zero training needed). If the learned model
  20. can't beat PCA, the extra machinery isn't earning its complexity.
  21.  
  22. 4. Diverse text with an INJECTED, UNIQUE FACT per document (a random 5-digit
  23. code in a sentence), instead of one repeated sentence. Repetition makes
  24. compression artificially easy; this doesn't.
  25.  
  26. 5. The metric that actually matters: after compress -> decompress, can the
  27. frozen LM's own output head still predict the correct token at the fact's
  28. position? Whole-chunk MSE can look good while the 2-3 dimensions that
  29. actually encode "48213" get washed out. Top-1/top-5 next-token accuracy
  30. at the fact position is a much sharper, more honest test.
  31.  
  32. 6. Averaged over many (document, chunk) pairs, not one query on one sample.
  33.  
  34. Usage:
  35. python bscm_eval.py --mode real --model_id Qwen/Qwen2.5-0.5B --n_docs 40 --steps 800
  36. python bscm_eval.py --mode synthetic # no internet/model needed, sanity-checks the pipeline
  37.  
  38. Requires: torch, and (for --mode real) transformers + internet access to download the model.
  39. """
  40.  
  41. import argparse
  42. import random
  43. import string
  44. import torch
  45. import torch.nn as nn
  46. import torch.nn.functional as F
  47.  
  48.  
  49. # --------------------------------------------------------------------------
  50. # Architecture
  51. # --------------------------------------------------------------------------
  52.  
  53. class LatentBottleneck(nn.Module):
  54. """Perceiver-style compressive autoencoder for a chunk of hidden states.
  55.  
  56. Encoder: n_latents learned queries cross-attend into the S-token chunk
  57. -> DNA code [n_latents, d_model].
  58. Decoder: S learned position queries cross-attend into the DNA code
  59. -> reconstructed chunk [S, d_model].
  60.  
  61. Param count depends on n_latents and d_model, NOT on chunk_size, which is
  62. the key fix over a flatten-linear encoder/decoder.
  63. """
  64.  
  65. def __init__(self, d_model, chunk_size, n_latents=8, n_heads=4):
  66. super().__init__()
  67. self.chunk_size = chunk_size
  68. self.n_latents = n_latents
  69.  
  70. self.enc_latents = nn.Parameter(torch.randn(n_latents, d_model) * 0.02)
  71. self.enc_attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True)
  72. self.enc_ff = nn.Sequential(nn.Linear(d_model, d_model), nn.GELU(), nn.Linear(d_model, d_model))
  73. self.enc_ln1 = nn.LayerNorm(d_model)
  74. self.enc_ln2 = nn.LayerNorm(d_model)
  75.  
  76. self.dec_queries = nn.Parameter(torch.randn(chunk_size, d_model) * 0.02)
  77. self.dec_attn = nn.MultiheadAttention(d_model, n_heads, batch_first=True)
  78. self.dec_ff = nn.Sequential(nn.Linear(d_model, d_model), nn.GELU(), nn.Linear(d_model, d_model))
  79. self.dec_ln1 = nn.LayerNorm(d_model)
  80. self.dec_ln2 = nn.LayerNorm(d_model)
  81.  
  82. def encode(self, chunk):
  83. B = chunk.shape[0]
  84. latents = self.enc_latents.unsqueeze(0).expand(B, -1, -1)
  85. attended, _ = self.enc_attn(latents, chunk, chunk)
  86. latents = self.enc_ln1(latents + attended)
  87. latents = self.enc_ln2(latents + self.enc_ff(latents))
  88. return latents
  89.  
  90. def decode(self, dna):
  91. B = dna.shape[0]
  92. queries = self.dec_queries.unsqueeze(0).expand(B, -1, -1)
  93. attended, _ = self.dec_attn(queries, dna, dna)
  94. out = self.dec_ln1(queries + attended)
  95. out = self.dec_ln2(out + self.dec_ff(out))
  96. return out
  97.  
  98. def forward(self, chunk):
  99. dna = self.encode(chunk)
  100. return self.decode(dna), dna
  101.  
  102. def num_params(self):
  103. return sum(p.numel() for p in self.parameters())
  104.  
  105.  
  106. class PCABaseline:
  107. """Best possible LINEAR compressor at a given latent budget. Fit once via
  108. SVD on training chunks (no gradient descent needed -- this is the
  109. closed-form optimum for linear reconstruction, so it's a fair floor for
  110. "is the learned nonlinear model actually buying us anything")."""
  111.  
  112. def __init__(self, n_components):
  113. self.n_components = n_components
  114. self.mean = None
  115. self.components = None # [d_model, n_components]
  116.  
  117. def fit(self, flat_chunks): # flat_chunks: [N, d_model]
  118. self.mean = flat_chunks.mean(dim=0, keepdim=True)
  119. centered = flat_chunks - self.mean
  120. U, S, V = torch.pca_lowrank(centered, q=self.n_components)
  121. self.components = V[:, : self.n_components] # [d_model, n_components]
  122.  
  123. def reconstruct(self, flat_chunks):
  124. centered = flat_chunks - self.mean
  125. coeffs = centered @ self.components # [N, n_components]
  126. recon = coeffs @ self.components.T + self.mean
  127. return recon
  128.  
  129.  
  130. class MeanPoolBaseline(nn.Module):
  131. """The original naive approach, kept as a baseline so the improvement is
  132. quantified rather than assumed."""
  133.  
  134. def __init__(self, d_model, chunk_size, k_dna):
  135. super().__init__()
  136. self.enc = nn.Sequential(nn.Linear(d_model, k_dna), nn.Tanh())
  137. self.dec = nn.Linear(k_dna, chunk_size * d_model)
  138. self.chunk_size = chunk_size
  139. self.d_model = d_model
  140.  
  141. def forward(self, chunk): # chunk: [B, S, D]
  142. pooled = chunk.mean(dim=1)
  143. dna = self.enc(pooled)
  144. recon = self.dec(dna).view(-1, self.chunk_size, self.d_model)
  145. return recon, dna
  146.  
  147. def num_params(self):
  148. return sum(p.numel() for p in self.parameters())
  149.  
  150.  
  151. # --------------------------------------------------------------------------
  152. # Fact localization (verified separately against HF's offset_mapping convention)
  153. # --------------------------------------------------------------------------
  154.  
  155. def find_fact_token_indices(offset_mapping, fact_char_start, fact_char_end):
  156. idxs = []
  157. for i, (s, e) in enumerate(offset_mapping):
  158. if s == e:
  159. continue
  160. if s < fact_char_end and e > fact_char_start:
  161. idxs.append(i)
  162. return idxs
  163.  
  164.  
  165. # --------------------------------------------------------------------------
  166. # Data: real LM hidden states with an injected, unique fact per document
  167. # --------------------------------------------------------------------------
  168.  
  169. FILLER_SENTENCES = [
  170. "The quarterly report covers regional sales performance in detail.",
  171. "Weather patterns this year have been unusually volatile across the coast.",
  172. "The committee reviewed several proposals before reaching a decision.",
  173. "Engineers debated the merits of the new cooling system for hours.",
  174. "Historical records suggest the settlement was founded in the early period.",
  175. "The recipe calls for a slow simmer over low heat for best results.",
  176. "Traffic congestion has worsened following the recent construction.",
  177. "The novel explores themes of memory and displacement across generations.",
  178. "Researchers collected samples from twelve different locations.",
  179. "The orchestra rehearsed the symphony's second movement extensively.",
  180. ]
  181.  
  182.  
  183. def make_document(chunk_size, n_chunks, fact_chunk_idx, rng):
  184. """Build a document of n_chunks * chunk_size (approx, in words) with a
  185. unique injected fact sentence placed in one designated chunk."""
  186. words_per_chunk_target = chunk_size # rough proxy; real tokenization will differ
  187. fact_code = "".join(rng.choice(string.digits) for _ in range(5))
  188. fact_sentence = f" The secret access code is {fact_code}."
  189.  
  190. chunks_text = []
  191. for c in range(n_chunks):
  192. filler = " ".join(rng.choice(FILLER_SENTENCES) for _ in range(max(3, words_per_chunk_target // 12)))
  193. if c == fact_chunk_idx:
  194. filler = filler + fact_sentence
  195. chunks_text.append(filler)
  196.  
  197. full_text = " ".join(chunks_text)
  198. return full_text, fact_code
  199.  
  200.  
  201. def build_real_dataset(model_id, n_docs, chunk_size, n_chunks_per_doc, device, seed=0):
  202. from transformers import AutoTokenizer, AutoModelForCausalLM
  203.  
  204. rng = random.Random(seed)
  205. print(f"Loading {model_id} ...")
  206. tokenizer = AutoTokenizer.from_pretrained(model_id)
  207. model = AutoModelForCausalLM.from_pretrained(model_id)
  208. model.to(device).eval()
  209.  
  210. docs = []
  211. for d in range(n_docs):
  212. fact_chunk_idx = rng.randrange(n_chunks_per_doc)
  213. text, fact_code = make_document(chunk_size, n_chunks_per_doc, fact_chunk_idx, rng)
  214.  
  215. enc = tokenizer(text, return_tensors="pt", return_offsets_mapping=True, truncation=True,
  216. max_length=chunk_size * n_chunks_per_doc)
  217. input_ids = enc.input_ids.to(device)
  218. offsets = enc.offset_mapping[0].tolist()
  219.  
  220. with torch.no_grad():
  221. out = model(input_ids, output_hidden_states=True)
  222. hidden = out.hidden_states[-1][0] # [seq_len, d_model], pre-final-norm
  223.  
  224. seq_len = hidden.shape[0]
  225. n_full_chunks = seq_len // chunk_size
  226. if n_full_chunks < 2:
  227. continue
  228. hidden = hidden[: n_full_chunks * chunk_size]
  229.  
  230. fact_char_start = text.find(fact_code)
  231. fact_char_end = fact_char_start + len(fact_code)
  232. fact_token_idxs = find_fact_token_indices(offsets, fact_char_start, fact_char_end)
  233. fact_token_idxs = [i for i in fact_token_idxs if i < n_full_chunks * chunk_size]
  234. if not fact_token_idxs:
  235. continue
  236. fact_chunk = fact_token_idxs[0] // chunk_size
  237. fact_local_idxs = [i % chunk_size for i in fact_token_idxs if i // chunk_size == fact_chunk]
  238. fact_target_ids = input_ids[0, fact_token_idxs].tolist()
  239.  
  240. docs.append({
  241. "hidden": hidden.detach(),
  242. "n_chunks": n_full_chunks,
  243. "fact_chunk": fact_chunk,
  244. "fact_local_idxs": fact_local_idxs,
  245. "fact_target_ids": fact_target_ids,
  246. })
  247.  
  248. return docs, model, tokenizer
  249.  
  250.  
  251. def build_synthetic_dataset(d_model, n_docs, chunk_size, n_chunks_per_doc, device, seed=0):
  252. """No internet needed. Hidden states are an AR(1) process per feature
  253. (so nearby tokens are correlated, like real embeddings, unlike iid noise).
  254. The 'fact' is a distinctive fixed vector pattern injected at one token
  255. position; retrieval is judged by nearest-neighbor match against a small
  256. codebook of decoy vectors, standing in for the LM vocabulary/head."""
  257. g = torch.Generator().manual_seed(seed)
  258. codebook_size = 50
  259. codebook = torch.randn(codebook_size, d_model, generator=g) * 2.0
  260.  
  261. docs = []
  262. for d in range(n_docs):
  263. seq_len = chunk_size * n_chunks_per_doc
  264. noise = torch.randn(seq_len, d_model, generator=g) * 0.3
  265. hidden = torch.cumsum(noise, dim=0)
  266. hidden = hidden - hidden.mean(dim=0, keepdim=True)
  267.  
  268. fact_chunk = torch.randint(0, n_chunks_per_doc, (1,), generator=g).item()
  269. fact_local_idx = torch.randint(0, chunk_size, (1,), generator=g).item()
  270. fact_code_idx = torch.randint(0, codebook_size, (1,), generator=g).item()
  271. global_idx = fact_chunk * chunk_size + fact_local_idx
  272. hidden[global_idx] = codebook[fact_code_idx] + torch.randn(d_model, generator=g) * 0.05
  273.  
  274. docs.append({
  275. "hidden": hidden.to(device),
  276. "n_chunks": n_chunks_per_doc,
  277. "fact_chunk": fact_chunk,
  278. "fact_local_idxs": [fact_local_idx],
  279. "fact_target_ids": [fact_code_idx],
  280. })
  281. return docs, codebook.to(device)
  282.  
  283.  
  284. # --------------------------------------------------------------------------
  285. # Training + evaluation
  286. # --------------------------------------------------------------------------
  287.  
  288. def train_bottleneck(model, train_chunks, steps, lr, device, batch_size=16, log_every=100):
  289. model.to(device).train()
  290. opt = torch.optim.Adam(model.parameters(), lr=lr)
  291. n = train_chunks.shape[0]
  292. losses = []
  293. for step in range(steps):
  294. idx = torch.randint(0, n, (min(batch_size, n),))
  295. batch = train_chunks[idx].to(device)
  296. recon, _ = model(batch)
  297. loss = F.mse_loss(recon, batch)
  298. opt.zero_grad()
  299. loss.backward()
  300. opt.step()
  301. losses.append(loss.item())
  302. if step % log_every == 0 or step == steps - 1:
  303. print(f" step {step:4d} mse {loss.item():.5f}")
  304. model.eval()
  305. return losses
  306.  
  307.  
  308. def evaluate_fact_retrieval_real(model, lm_model, eval_docs, chunk_size, device):
  309. """Returns dict of aggregate metrics using the LM's own head as the ground
  310. truth for whether the fact is recoverable, not just low-level MSE."""
  311. top1_hits, top5_hits, n_facts = 0, 0, 0
  312. mse_all, mse_fact_pos = [], []
  313.  
  314. final_norm = getattr(lm_model.model, "norm", None)
  315. lm_head = lm_model.get_output_embeddings()
  316.  
  317. with torch.no_grad():
  318. for doc in eval_docs:
  319. hidden = doc["hidden"]
  320. c = doc["fact_chunk"]
  321. chunk = hidden[c * chunk_size:(c + 1) * chunk_size].unsqueeze(0).to(device)
  322. recon, _ = model(chunk)
  323.  
  324. mse_all.append(F.mse_loss(recon, chunk).item())
  325. local_idxs = doc["fact_local_idxs"]
  326. target_ids = doc["fact_target_ids"]
  327.  
  328. recon_fact = recon[0, local_idxs] # [n_fact_toks, d_model]
  329. orig_fact = chunk[0, local_idxs]
  330. mse_fact_pos.append(F.mse_loss(recon_fact, orig_fact).item())
  331.  
  332. normed = final_norm(recon_fact) if final_norm is not None else recon_fact
  333. logits = lm_head(normed) # [n_fact_toks, vocab]
  334. top5 = logits.topk(5, dim=-1).indices
  335.  
  336. for i, tgt in enumerate(target_ids):
  337. n_facts += 1
  338. if top5[i, 0].item() == tgt:
  339. top1_hits += 1
  340. if tgt in top5[i].tolist():
  341. top5_hits += 1
  342.  
  343. return {
  344. "chunk_mse": sum(mse_all) / len(mse_all),
  345. "fact_position_mse": sum(mse_fact_pos) / len(mse_fact_pos),
  346. "fact_top1_acc": top1_hits / max(n_facts, 1),
  347. "fact_top5_acc": top5_hits / max(n_facts, 1),
  348. "n_facts_evaluated": n_facts,
  349. }
  350.  
  351.  
  352. def evaluate_fact_retrieval_synthetic(model, eval_docs, codebook, chunk_size, device):
  353. top1_hits, n_facts = 0, 0
  354. mse_all, mse_fact_pos = [], []
  355.  
  356. with torch.no_grad():
  357. for doc in eval_docs:
  358. hidden = doc["hidden"]
  359. c = doc["fact_chunk"]
  360. chunk = hidden[c * chunk_size:(c + 1) * chunk_size].unsqueeze(0).to(device)
  361. recon, _ = model(chunk)
  362.  
  363. mse_all.append(F.mse_loss(recon, chunk).item())
  364. local_idxs = doc["fact_local_idxs"]
  365. target_ids = doc["fact_target_ids"]
  366.  
  367. recon_fact = recon[0, local_idxs]
  368. orig_fact = chunk[0, local_idxs]
  369. mse_fact_pos.append(F.mse_loss(recon_fact, orig_fact).item())
  370.  
  371. sims = F.cosine_similarity(recon_fact.unsqueeze(1), codebook.unsqueeze(0), dim=-1)
  372. pred = sims.argmax(dim=-1)
  373. for i, tgt in enumerate(target_ids):
  374. n_facts += 1
  375. if pred[i].item() == tgt:
  376. top1_hits += 1
  377.  
  378. return {
  379. "chunk_mse": sum(mse_all) / len(mse_all),
  380. "fact_position_mse": sum(mse_fact_pos) / len(mse_fact_pos),
  381. "fact_top1_acc": top1_hits / max(n_facts, 1),
  382. "n_facts_evaluated": n_facts,
  383. }
  384.  
  385.  
  386. def evaluate_pca(pca, eval_docs, chunk_size, device, real_probe=None, lm_model=None, codebook=None):
  387. mse_all, mse_fact_pos = [], []
  388. top1_hits, n_facts = 0, 0
  389. with torch.no_grad():
  390. for doc in eval_docs:
  391. hidden = doc["hidden"]
  392. c = doc["fact_chunk"]
  393. chunk = hidden[c * chunk_size:(c + 1) * chunk_size].to(device)
  394. recon = pca.reconstruct(chunk)
  395. mse_all.append(F.mse_loss(recon, chunk).item())
  396.  
  397. local_idxs = doc["fact_local_idxs"]
  398. target_ids = doc["fact_target_ids"]
  399. recon_fact = recon[local_idxs]
  400. orig_fact = chunk[local_idxs]
  401. mse_fact_pos.append(F.mse_loss(recon_fact, orig_fact).item())
  402.  
  403. if real_probe and lm_model is not None:
  404. final_norm = getattr(lm_model.model, "norm", None)
  405. lm_head = lm_model.get_output_embeddings()
  406. normed = final_norm(recon_fact) if final_norm is not None else recon_fact
  407. logits = lm_head(normed)
  408. pred = logits.argmax(dim=-1)
  409. for i, tgt in enumerate(target_ids):
  410. n_facts += 1
  411. if pred[i].item() == tgt:
  412. top1_hits += 1
  413. elif codebook is not None:
  414. sims = F.cosine_similarity(recon_fact.unsqueeze(1), codebook.unsqueeze(0), dim=-1)
  415. pred = sims.argmax(dim=-1)
  416. for i, tgt in enumerate(target_ids):
  417. n_facts += 1
  418. if pred[i].item() == tgt:
  419. top1_hits += 1
  420.  
  421. result = {
  422. "chunk_mse": sum(mse_all) / len(mse_all),
  423. "fact_position_mse": sum(mse_fact_pos) / len(mse_fact_pos),
  424. "fact_top1_acc": top1_hits / max(n_facts, 1),
  425. "n_facts_evaluated": n_facts,
  426. }
  427. return result
  428.  
  429.  
  430. # --------------------------------------------------------------------------
  431. # Main
  432. # --------------------------------------------------------------------------
  433.  
  434. def main():
  435. ap = argparse.ArgumentParser()
  436. ap.add_argument("--mode", choices=["real", "synthetic"], default="synthetic")
  437. ap.add_argument("--model_id", default="Qwen/Qwen2.5-0.5B")
  438. ap.add_argument("--n_docs", type=int, default=40)
  439. ap.add_argument("--chunk_size", type=int, default=64)
  440. ap.add_argument("--n_chunks_per_doc", type=int, default=6)
  441. ap.add_argument("--n_latents", type=int, default=8)
  442. ap.add_argument("--steps", type=int, default=800)
  443. ap.add_argument("--lr", type=float, default=1e-3)
  444. ap.add_argument("--seed", type=int, default=0)
  445. args = ap.parse_args()
  446.  
  447. device = "cuda" if torch.cuda.is_available() else "cpu"
  448. torch.manual_seed(args.seed)
  449. print(f"device: {device}")
  450.  
  451. lm_model = None
  452. codebook = None
  453.  
  454. if args.mode == "real":
  455. try:
  456. docs, lm_model, tokenizer = build_real_dataset(
  457. args.model_id, args.n_docs, args.chunk_size, args.n_chunks_per_doc, device, args.seed)
  458. if len(docs) < 8:
  459. raise RuntimeError("too few usable documents (fact localization failed too often)")
  460. d_model = docs[0]["hidden"].shape[-1]
  461. except Exception as e:
  462. print(f"[real mode failed: {e}] -- falling back to synthetic mode. "
  463. f"Run this on a machine with internet access to huggingface.co for real results.")
  464. args.mode = "synthetic"
  465.  
  466. if args.mode == "synthetic":
  467. d_model = 896 # matches Qwen2.5-0.5B for apples-to-apples config
  468. docs, codebook = build_synthetic_dataset(
  469. d_model, args.n_docs, args.chunk_size, args.n_chunks_per_doc, device, args.seed)
  470.  
  471. n_eval = max(1, len(docs) // 5)
  472. eval_docs, train_docs = docs[:n_eval], docs[n_eval:]
  473. print(f"docs: {len(docs)} total -> {len(train_docs)} train / {len(eval_docs)} eval")
  474.  
  475. all_chunks = []
  476. for doc in train_docs:
  477. h = doc["hidden"]
  478. n_c = doc["n_chunks"]
  479. for c in range(n_c):
  480. all_chunks.append(h[c * args.chunk_size:(c + 1) * args.chunk_size])
  481. train_chunks = torch.stack(all_chunks) # [N, S, D]
  482. print(f"training chunks: {train_chunks.shape}")
  483.  
  484. # --- Model 1: Latent Bottleneck (the proposed architecture) ---
  485. print("\n=== Training LatentBottleneck ===")
  486. lb = LatentBottleneck(d_model=d_model, chunk_size=args.chunk_size, n_latents=args.n_latents)
  487. train_bottleneck(lb, train_chunks, args.steps, args.lr, device)
  488.  
  489. # --- Model 2: Mean-pool baseline (the original naive idea) ---
  490. # NOTE: we do NOT match the LatentBottleneck's total latent scalar budget
  491. # (n_latents * d_model) here on purpose. A single Linear(k_dna, chunk_size
  492. # * d_model) decoder has chunk_size * d_model * k_dna parameters, so
  493. # matching the budget exactly would require ~400M+ params for this decoder
  494. # alone (confirmed: it OOMs). That blowup is itself the finding -- it's
  495. # why a flatten/mean-pool decoder doesn't scale, independent of whether
  496. # the reconstruction quality is any good. We give it k_dna = n_latents
  497. # scalars instead (a smaller, tractable budget) and still expect the
  498. # LatentBottleneck to win on fact retrieval despite MeanPool having a much
  499. # larger raw parameter count in its remaining layers.
  500. print("\n=== Training MeanPoolBaseline ===")
  501. mp = MeanPoolBaseline(d_model=d_model, chunk_size=args.chunk_size, k_dna=args.n_latents)
  502. train_bottleneck(mp, train_chunks, args.steps, args.lr, device)
  503.  
  504. # --- Model 3: PCA baseline (closed-form optimal linear compressor) ---
  505. print("\n=== Fitting PCA baseline ===")
  506. flat_train = train_chunks.reshape(-1, d_model).to(device)
  507. pca = PCABaseline(n_components=args.n_latents)
  508. pca.fit(flat_train)
  509.  
  510. # --- Evaluation ---
  511. print("\n=== Evaluation (held-out documents) ===")
  512. if args.mode == "real":
  513. res_lb = evaluate_fact_retrieval_real(lb, lm_model, eval_docs, args.chunk_size, device)
  514. res_mp = evaluate_fact_retrieval_real(mp, lm_model, eval_docs, args.chunk_size, device)
  515. res_pca = evaluate_pca(pca, eval_docs, args.chunk_size, device, real_probe=True, lm_model=lm_model)
  516. else:
  517. res_lb = evaluate_fact_retrieval_synthetic(lb, eval_docs, codebook, args.chunk_size, device)
  518. res_mp = evaluate_fact_retrieval_synthetic(mp, eval_docs, codebook, args.chunk_size, device)
  519. res_pca = evaluate_pca(pca, eval_docs, args.chunk_size, device, codebook=codebook)
  520.  
  521. def fmt(r):
  522. s = f"chunk_mse={r['chunk_mse']:.5f} fact_pos_mse={r['fact_position_mse']:.5f} " \
  523. f"fact_top1={r['fact_top1_acc']:.2%}"
  524. if "fact_top5_acc" in r:
  525. s += f" fact_top5={r['fact_top5_acc']:.2%}"
  526. s += f" (n_facts={r['n_facts_evaluated']})"
  527. return s
  528.  
  529. print("\n--- RESULTS ---")
  530. print(f"LatentBottleneck params={lb.num_params():>10,} {fmt(res_lb)}")
  531. print(f"MeanPoolBaseline params={mp.num_params():>10,} {fmt(res_mp)}")
  532. print(f"PCA (closed-form) params={'0 (SVD)':>10} {fmt(res_pca)}")
  533. print(f"\ncompression ratio (elements): {args.chunk_size * d_model / (args.n_latents * d_model):.1f}x")
  534. print("\nInterpretation: fact_top1_acc is the number that matters. chunk_mse can look")
  535. print("good while the specific fact detail is unrecoverable -- if fact_top1_acc for")
  536. print("LatentBottleneck isn't clearly ahead of PCA at the same latent budget, the")
  537. print("attention machinery isn't earning its complexity yet, and that's the honest")
  538. print("result to report, not a reason to hide the comparison.")
  539.  
  540.  
  541. if __name__ == "__main__":
  542. main()
Add Comment
Please, Sign In to add comment