"""Reproduce the mandatory-operator-bottleneck claim on a released cl33-opLM checkpoint. Zeroing the emitted operators (Bs=0 everywhere: identity rotors in the scan, zeroed readout features) multiplies perplexity by orders of magnitude — the model has no other path to output. Paper: "One Object" §1/§2/§7; 314x measured on the chat checkpoint (prose validation); expect the same order of magnitude on WikiText-103. Usage: python repro_bottleneck.py --ckpt cl33_oplm_chat_236m.pt [--iters 20] Deps: torch, transformers, datasets (model_v2.py + so33.py from this bundle) """ import argparse, math, sys from pathlib import Path import torch, torch.nn.functional as F sys.path.insert(0, str(Path(__file__).resolve().parent)) from model_v2 import OpEmitV2Config, OpEmitLMv2 ap = argparse.ArgumentParser() ap.add_argument("--ckpt", default="cl33_oplm_chat_236m.pt") ap.add_argument("--iters", type=int, default=20) ap.add_argument("--seq", type=int, default=1024) ap.add_argument("--slice", default=None, help="path to eval_slice_prose_val.npy for the EXACT in-domain number (270x)") a = ap.parse_args() dev = "cuda" if torch.cuda.is_available() else "cpu" d = torch.load(a.ckpt, map_location="cpu", weights_only=False) cfg = OpEmitV2Config(**{k: v for k, v in d["config"].items() if k in OpEmitV2Config.__dataclass_fields__}) m = OpEmitLMv2(cfg); m.load_state_dict(d["model"]); m.eval().to(dev) for p in m.parameters(): p.requires_grad_(False) print(f"loaded {a.ckpt} | {sum(p.numel() for p in m.parameters())/1e6:.1f}M params | step {d.get('step')}") if a.slice: import numpy as np rows = np.load(a.slice) # (24,1025) frozen prose-val slice ids = None print(f"frozen slice: {rows.shape[0]} x {rows.shape[1]} tokens ({a.slice})") else: from transformers import GPT2TokenizerFast from datasets import load_dataset tok = GPT2TokenizerFast.from_pretrained("gpt2") text = "\n\n".join(load_dataset("Salesforce/wikitext", "wikitext-103-raw-v1", split="test")["text"]) ids = tok(text, add_special_tokens=False)["input_ids"] print(f"wikitext-103 test: {len(ids)/1e6:.2f}M tokens") def ce_pass(zero_ops: bool): tot, n = 0.0, 0 if a.slice: batches = [(torch.tensor(r[:-1][None].astype("int64"), device=dev), torch.tensor(r[1:][None].astype("int64"), device=dev)) for r in rows] else: batches = [(torch.tensor([ids[i*a.seq:i*a.seq+a.seq]], device=dev), torch.tensor([ids[i*a.seq+1:i*a.seq+a.seq+1]], device=dev)) for i in range(a.iters)] for x, y in batches: with torch.no_grad(): Bs, Bq, Bk = m.emit(x) if zero_ops: Bs = torch.zeros_like(Bs) logits = m.assemble(Bs, Bq, Bk, scan_only=getattr(cfg, "scan_only", False)) if isinstance(logits, tuple): logits = logits[0] tot += float(F.cross_entropy(logits.reshape(-1, logits.shape[-1]).float(), y.reshape(-1), reduction="sum")) n += y.numel() return tot / n ce_nat = ce_pass(False); ce_off = ce_pass(True) print(f"\n native : CE {ce_nat:.3f} PPL {math.exp(ce_nat):9.1f}") print(f" ops off: CE {ce_off:.3f} PPL {math.exp(ce_off):9.1f}") print(f" BOTTLENECK RATIO: {math.exp(ce_off)/math.exp(ce_nat):.0f}x " f"(expected: 270x on the shipped frozen slice; ~106-112x on wikitext; " f"the paper's original draw measured 314x — order of magnitude is the claim)")