"""Re-implement TORQUE's two-stage outlier retention (arXiv 2609.36032v1) and measure it. What it does ------------ TORQUE (Ben Basat, Mitzenmacher, Vargaftik, 28 Sep 2026) wraps a rotation-based quantizer: keep the k largest input coordinates exactly, rotate the rest, keep the rotated values above a threshold c exactly, and quantize the remaining "inliers" with s bits, choosing (k, c, s) per vector under one expected bit budget. This script rebuilds that around the simplest inlier quantizer there is, a Lloyd-Max (MSE-optimal, biased) scalar codebook for a truncated standard Gaussian behind a randomized Hadamard rotation, and runs five checks: 1. Gaussian model: exact error integrals, no sampling. Post-rotation retention against the same quantizer without it, at budgets from 2 to 8 bits in quarter-bit steps, plus a baseline that packs floor(2^b) levels two values per 2b-bit code at half-bit budgets. 2. A Monte Carlo spot check of (1) on real rotations at d = 256. 3. Synthetic skewed and heavy-tailed inputs (signed lognormal, Student-t) at 8 bits: which stage does the work. 4. Qwen2.5-0.5B on the WikiText-2 test split, the paper's own setting: activation vectors (four per decoder block, 256-coordinate chunks) and the key and value caches (one 64-coordinate vector per head and token), at budgets 2 to 8 in half-bit steps, with the per-vector side information free (as the baseline pays none) and charged. 5. WikiText-2 perplexity with the key and value caches quantized inside the model, which the paper does not report, with a plain per-vector absmax integer quantizer (no rotation) as a third reference at whole-bit budgets. Design choices that are ours, not the paper's: the inlier quantizer (Lloyd-Max, biased); fractional inlier bits implemented as a mix of the two neighbouring integer codebooks; retained values stored as float16 with 8-bit indices (the paper's 16 + 8 bits); the eight token positions per activation segment (31, 63, ..., 255); NMSE averaged per recorded vector; the side-information encoding in the charged variant (a 1-bit on/off flag, 4 bits for c, and the two retention counts at ceil(log2(d+1)) bits each; s is implied by the rest). Inputs (downloaded from Hugging Face on first run, pinned by revision): Qwen/Qwen2.5-0.5B revision 060db6499f32faf8b98477b0a26969ef7d8b9987 Salesforce/wikitext revision b08601e04326c79dfdd32d625aee71d232d685c3 file wikitext-2-raw-v1/test-00000-of-00001.parquet Outputs: datasets/torque-outlier-retention.csv (every number the article quotes), content/images/torque-outlier-retention/torque.png, and a printed report. Run: python code/torque-outlier-retention.py (about an hour on 32 CPU threads) python code/torque-outlier-retention.py --ppl-segments 8 --no-chart (quicker, noisier) python code/torque-outlier-retention.py --chart-only (redraw the figure from the CSV) The CSV is rewritten after every stage, so an interrupted run keeps what it finished. Requires: Python 3.11+, numpy 2.3.5, scipy 1.17.0, torch 2.12.0 (CPU build is enough), transformers 4.56.2, huggingface_hub 0.36.2, pandas 3.0.2 with pyarrow 24.0.0 (parquet), matplotlib 3.10.9 (chart only). No GPU. """ from __future__ import annotations import argparse import csv import math import os import sys import time import warnings import numpy as np from scipy.linalg import hadamard from scipy.special import erfc, ndtri HERE = os.path.dirname(os.path.abspath(__file__)) ROOT = os.path.dirname(HERE) SLUG = "torque-outlier-retention" CSV_PATH = os.path.join(ROOT, "datasets", SLUG + ".csv") IMG_PATH = os.path.join(ROOT, "content", "images", SLUG, "torque.png") MODEL, MODEL_REV = "Qwen/Qwen2.5-0.5B", "060db6499f32faf8b98477b0a26969ef7d8b9987" WIKITEXT, WIKITEXT_REV = "Salesforce/wikitext", "b08601e04326c79dfdd32d625aee71d232d685c3" WIKITEXT_FILE = "wikitext-2-raw-v1/test-00000-of-00001.parquet" # The paper's candidate thresholds (Appendix D) and bit grid S = {1 + j/64 : j = 0..448}. C_SET = [1.75, 2, 2.25, 2.5, 2.75, 3, 3.5, 4, 4.5, 5, 6, math.inf] GRID, S_MIN, S_MAX = 64, 1.0, 8.0 VALUE_BITS, INDEX_BITS = 16, 8 # the paper's v = 16 and p = 8 L_RET = VALUE_BITS + INDEX_BITS # 24 bits per retained coordinate, either stage POW2 = [2 ** s for s in range(1, 9)] PACKED = {2.5: 5, 3.5: 11, 4.5: 22, 5.5: 45, 6.5: 90, 7.5: 181} # floor(2^b) levels, packed in pairs ACT_POSITIONS = list(range(31, 256, 32)) # eight tokens per 256-token segment (our choice) CHECKS = 0 ROWS: list[dict] = [] def check(cond: bool, msg: str) -> None: global CHECKS CHECKS += 1 if not cond: raise AssertionError(msg) def row(experiment, data, budget, method, metric, value, note=""): ROWS.append({"experiment": experiment, "data": data, "budget_bits": budget, "method": method, "metric": metric, "value": value, "note": note}) # ---------------------------------------------------------------- Gaussian codebooks SQ2, INV_SQ2PI = math.sqrt(2.0), 1.0 / math.sqrt(2.0 * math.pi) def phi(x): x = np.asarray(x, dtype=float) with np.errstate(over="ignore", invalid="ignore"): out = INV_SQ2PI * np.exp(-0.5 * x * x) return np.where(np.isfinite(x), out, 0.0) def tail(x): """P(G > x) for G ~ N(0, 1), accurate in the tail.""" return 0.5 * erfc(np.asarray(x, dtype=float) / SQ2) def xphi(x): x = np.asarray(x, dtype=float) with np.errstate(over="ignore", invalid="ignore"): out = x * INV_SQ2PI * np.exp(-0.5 * x * x) return np.where(np.isfinite(x), out, 0.0) def lloyd_max(c: float, L: int, iters: int = 200000, tol: float = 1e-12): """L-level MSE-optimal quantizer for N(0,1) restricted to [-c, c] (c may be inf). Returns (cuts, levels, eps) with eps = E[(Q(G) - G)^2 ; |G| <= c], exact integrals.""" h, odd = L // 2, L % 2 == 1 top = 1.0 if not math.isfinite(c) else 1.0 - float(tail(c / math.sqrt(3.0))) # start from the high-resolution optimum (point density ~ phi^(1/3), i.e. N(0, 3) quantiles) u = (0.5 + (top - 0.5) * (np.arange(1, h + 2) - 0.5) / (h + 0.5)) if odd else \ (0.5 + (top - 0.5) * np.arange(0, h + 1) / h) t = math.sqrt(3.0) * ndtri(np.clip(u, 0.5, 1.0)) if not odd: t[0] = 0.0 t[-1] = c for _ in range(iters): a, b = t[:-1], t[1:] y = (phi(a) - phi(b)) / (tail(a) - tail(b)) # centroid of each positive cell tn = t.copy() if odd: tn[0] = 0.5 * y[0] tn[1:-1] = 0.5 * (y[:-1] + y[1:]) # nearest-neighbour cuts delta = float(np.max(np.abs(tn[:-1] - t[:-1]))) t = tn if delta < tol: break a, b = t[:-1], t[1:] mass = tail(a) - tail(b) y = (phi(a) - phi(b)) / mass second, first = mass + xphi(a) - xphi(b), phi(a) - phi(b) # int g^2 phi, int g phi eps = 2.0 * float(np.sum(second - 2 * y * first + y * y * mass)) if odd: eps += 2.0 * float((0.5 - tail(t[0])) - xphi(t[0])) # centre cell reconstructs to 0 levels = np.concatenate([-y[::-1], [0.0], y]) cuts = np.concatenate([-t[:-1][::-1], t[:-1]]) else: levels = np.concatenate([-y[::-1], y]) cuts = np.concatenate([-t[1:-1][::-1], [0.0], t[1:-1]]) return cuts, levels, eps CB: dict = {} def build_codebooks() -> None: for c in C_SET: for L in POW2: CB[(c, L)] = lloyd_max(c, L) for L in PACKED.values(): CB[(math.inf, L)] = lloyd_max(math.inf, L) def p_out(c: float) -> float: return 0.0 if not math.isfinite(c) else 2.0 * float(tail(c)) def eps_s(c: float, s: float) -> float: """Gaussian-model error at s bits per inlier; fractional s mixes the two integer codebooks.""" s0 = int(math.floor(s + 1e-12)) f = s - s0 e0 = CB[(c, 2 ** s0)][2] return e0 if f < 1e-12 else (1 - f) * e0 + f * CB[(c, 2 ** (s0 + 1))][2] def best_s(avail: float, p: float): """Largest grid s with p*L_RET + (1-p)*s <= avail bits per coordinate, or None.""" s = (avail - p * L_RET) / (1 - p) s = min(S_MAX, math.floor((s - S_MIN) * GRID + 1e-9) / GRID + S_MIN) return s if s >= S_MIN - 1e-12 else None # ---------------------------------------------------------------- the quantizer def side_bits_on(d: int) -> int: return 1 + 4 + 2 * math.ceil(math.log2(d + 1)) # flag, c index, k, post-rotation count def candidates(d: int, b: float, stage: str): """(c, k, s, side_bits) plans. The paper's rule: for each (c, s) the largest feasible k wins, equivalently for each (c, k) the largest feasible s; we enumerate (c, k).""" if stage == "none": return [(math.inf, 0, b, 0)] charged = stage == "charged" side = side_bits_on(d) if charged else 0 out = [] if charged: # the 'off' plan: plain quantizer, one flag bit out.append((math.inf, 0, best_s((d * b - 1) / d, 0.0), 1)) for c in C_SET: if stage == "pre" and math.isfinite(c): continue p = p_out(c) kmax = 0 if stage == "post" else d for k in range(0, kmax + 1): s = best_s((d * b - side - L_RET * k) / d, p) if s is None: break if charged and k == 0 and not math.isfinite(c): continue # same as the off plan but dearer out.append((c, k, s, side)) return out HCACHE: dict = {} def hmat(d: int): if d not in HCACHE: HCACHE[d] = hadamard(d).astype(np.float64) / math.sqrt(d) return HCACHE[d] def encode_decode(X: np.ndarray, b: float, stage: str, signs: np.ndarray, packed: bool = False): """Quantize each row of X (n x d) at an expected b bits per coordinate. stage: none | pre | post | both | charged. Returns (xhat, realized bits per row, k per row).""" n, d = X.shape Hd = hmat(d) sq = X * X norm2 = sq.sum(1) order = np.argsort(-sq, axis=1, kind="stable") cums = np.concatenate([np.zeros((n, 1)), np.cumsum(np.take_along_axis(sq, order, axis=1), axis=1)], axis=1) cand = candidates(d, b, stage) best_obj, best_j = np.full(n, np.inf), np.zeros(n, int) for j, (c, k, s, side) in enumerate(cand): e = CB[(math.inf, PACKED[b])][2] if packed else eps_s(c, s) obj = np.maximum(1.0 - cums[:, k] / norm2, 0.0) * e # the bound rho_k * eps(c, s) better = obj < best_obj best_obj[better], best_j[better] = obj[better], j ks = np.array([cand[j][1] for j in best_j]) pre = np.zeros((n, d), dtype=bool) for kk in np.unique(ks): if kk: sel = np.where(ks == kk)[0] pre[sel[:, None], order[sel, :kk]] = True Xr = np.where(pre, 0.0, X) r = np.sqrt((Xr * Xr).sum(1)) ok = r > 0 Z = np.zeros_like(X) Z[ok] = Xr[ok] * (math.sqrt(d) / r[ok])[:, None] U = (Z * signs) @ Hd # randomized Hadamard rotation Uhat = np.zeros_like(U) bits = np.zeros(n) for j in np.unique(best_j): c, k, s, side = cand[j] sel = np.where((best_j == j) & ok)[0] bits[best_j == j] += side + L_RET * k if sel.size == 0: continue Us = U[sel] post = np.abs(Us) > c inl = ~post n_in = inl.sum(1) if packed: cuts, lev, _ = CB[(math.inf, PACKED[b])] Uq = lev[np.searchsorted(cuts, Us)] bits[sel] += np.ceil(n_in / 2) * 2 * b else: s0 = int(math.floor(s + 1e-12)) n_hi = np.floor((s - s0) * n_in + 0.5).astype(int) hi = inl & (np.cumsum(inl, axis=1) <= n_hi[:, None]) lo = inl & ~hi Uq = np.where(post, Us.astype(np.float16).astype(np.float64), 0.0) for mask, L in ((hi, 2 ** (s0 + 1)), (lo, 2 ** s0)): if mask.any(): cuts, lev, _ = CB[(c, L)] Uq[mask] = lev[np.searchsorted(cuts, Us[mask])] bits[sel] += n_hi * (s0 + 1) + (n_in - n_hi) * s0 + L_RET * post.sum(1) Uhat[sel] = Uq Xt = ((Uhat @ Hd) * signs) * (r / math.sqrt(d))[:, None] Xt[pre] = X[pre].astype(np.float16).astype(np.float64) return Xt, bits, ks def absmax_int(X: np.ndarray, b: float) -> np.ndarray: """Plain per-vector symmetric absmax rounding to 2^(b-1) - 1 steps a side, no rotation. Like the rotated quantizer's norm, the one scale per vector is kept outside the budget.""" q = 2 ** (int(b) - 1) - 1 m = np.abs(X).max(1, keepdims=True) m = np.where(m > 0, m, 1.0) return np.round(X / m * q) / q * m def nmse(X, Xh): return ((Xh - X) ** 2).sum(1) / (X * X).sum(1) def random_signs(rng, shape): return rng.choice(np.array([-1.0, 1.0]), size=shape) # ---------------------------------------------------------------- 1-3: Gaussian model, MC, synthetic QUARTER = [x / 4 for x in range(8, 33)] HALF = [x / 2 for x in range(4, 17)] def gaussian_model() -> dict: out = {} for b in QUARTER: base = eps_s(math.inf, b) best = (base, math.inf, b) for c in C_SET: p = p_out(c) for j in range(449): s = 1 + j / GRID if p * L_RET + (1 - p) * s <= b + 1e-12 and eps_s(c, s) < best[0]: best = (eps_s(c, s), c, s) red = 100 * (1 - best[0] / base) out[b] = (base, best, red) row("gaussian_model", "N(0,1)", b, "mixed_baseline", "eps", base) row("gaussian_model", "N(0,1)", b, "torque_post", "eps", best[0], "c=%s s=%.6f" % (best[1], best[2])) row("gaussian_model", "N(0,1)", b, "torque_post", "reduction_pct", red) if b in PACKED: row("gaussian_model", "N(0,1)", b, "packed_baseline", "eps", CB[(math.inf, PACKED[b])][2], "levels=%d" % PACKED[b]) return out def monte_carlo(gm: dict) -> None: rng = np.random.default_rng(7) d, n = 256, 4000 X = rng.standard_normal((n, d)) sg = random_signs(rng, (n, d)) for b in (4.0, 4.25, 8.0): e_base = nmse(X, encode_decode(X, b, "none", sg)[0]).mean() e_post = nmse(X, encode_decode(X, b, "post", sg)[0]).mean() row("monte_carlo_d256", "N(0,1) rotated", b, "mixed_baseline", "nmse", e_base) row("monte_carlo_d256", "N(0,1) rotated", b, "torque_post", "nmse", e_post) check(abs(e_base / gm[b][0] - 1) < 0.03, "MC baseline within 3%% of the model at b=%s" % b) check(abs(e_post / gm[b][1][0] - 1) < 0.03, "MC TORQUE within 3%% of the model at b=%s" % b) def synthetic() -> dict: """Signed lognormal (log-variance theta^2) and Student-t (nu) at 8 bits, d = 256, 256 vectors x 5 seeds.""" out = {} for fam, params in (("lognormal", (0.1, 1.0, 10.0)), ("student_t", (30, 5, 3))): for prm in params: res = {} for stage in ("none", "pre", "post", "both"): vals = [] for seed in range(5): rng = np.random.default_rng([seed, 11]) if fam == "lognormal": X = rng.choice([-1.0, 1.0], (256, 256)) * np.exp(rng.normal(0, math.sqrt(prm), (256, 256))) else: X = rng.standard_t(prm, (256, 256)) X /= np.linalg.norm(X, axis=1, keepdims=True) sg = random_signs(rng, X.shape) vals.append(nmse(X, encode_decode(X, 8.0, stage, sg)[0]).mean()) res[stage] = float(np.mean(vals)) for stage in ("pre", "post", "both"): red = 100 * (1 - res[stage] / res["none"]) row("synthetic_8bit", "%s=%s" % (fam, prm), 8.0, "torque_" + stage, "reduction_pct", red) out[(fam, prm)] = res return out # ---------------------------------------------------------------- 4-5: Qwen2.5-0.5B def load_model_and_text(): import pandas as pd import torch from huggingface_hub import hf_hub_download from transformers import AutoModelForCausalLM, AutoTokenizer tok = AutoTokenizer.from_pretrained(MODEL, revision=MODEL_REV) model = AutoModelForCausalLM.from_pretrained(MODEL, revision=MODEL_REV, dtype=torch.float32, attn_implementation="eager").eval() path = hf_hub_download(WIKITEXT, WIKITEXT_FILE, repo_type="dataset", revision=WIKITEXT_REV) ids = tok("\n\n".join(pd.read_parquet(path)["text"].tolist()), return_tensors="pt").input_ids[0] segs = ids[: ids.numel() // 256 * 256].view(-1, 256) cfg = model.config check((cfg.hidden_size, cfg.intermediate_size, cfg.num_hidden_layers) == (896, 4864, 24), "Qwen2.5-0.5B shape") check((cfg.num_key_value_heads, cfg.hidden_size // cfg.num_attention_heads) == (2, 64), "2 KV heads of 64") row("setup", "wikitext-2-raw-v1 test", "", "tokens", "count", int(ids.numel())) return model, segs def capture(model, segs): import torch acts = {"qkv_in": [], "o_in": [], "gateup_in": [], "down_in": []} def grab(name): def hook(_mod, args): acts[name].append(args[0][0, ACT_POSITIONS, :].detach().double().numpy()) return hook hooks = [] for layer in model.model.layers: hooks += [layer.self_attn.q_proj.register_forward_pre_hook(grab("qkv_in")), layer.self_attn.o_proj.register_forward_pre_hook(grab("o_in")), layer.mlp.gate_proj.register_forward_pre_hook(grab("gateup_in")), layer.mlp.down_proj.register_forward_pre_hook(grab("down_in"))] with torch.no_grad(): for i in range(8): # 8 segments x 8 tokens = 64 tokens model(segs[i:i + 1], use_cache=False) for h in hooks: h.remove() acts = {k: np.concatenate(v, 0) for k, v in acts.items()} keys, vals = [], [] with torch.no_grad(): for i in range(16): # 16 segments, first 255 tokens cached cache = model(segs[i:i + 1, :255], use_cache=True).past_key_values for layer in cache.layers: keys.append(layer.keys[0].double().numpy().reshape(-1, 64)) vals.append(layer.values[0].double().numpy().reshape(-1, 64)) K, V = np.concatenate(keys), np.concatenate(vals) check(all(a.shape[0] == 64 * 24 for a in acts.values()), "64 tokens x 24 blocks per activation type") check(K.shape == (16 * 24 * 2 * 255, 64) and V.shape == K.shape, "195,840 key and value vectors") for name, X in list(acts.items()) + [("keys", K), ("values", V)]: check(np.abs(X).max() < 65504, "%s fits float16 retention" % name) sq = X * X srt = -np.sort(-sq, axis=1) if name in acts: srt = -np.sort(-(X[:, :256] ** 2), axis=1) sq = X[:, :256] ** 2 row("stats", name, "", "top1_energy_share", "median", float(np.median(srt[:, 0] / sq.sum(1)))) row("stats", name, "", "top4_energy_share", "median", float(np.median(srt[:, :4].sum(1) / sq.sum(1)))) row("stats", name, "", "max_abs", "value", float(np.abs(X).max())) return acts, K, V def chunks(D: int): out, start = [], 0 while start < D: w = 256 if D - start >= 256 else D - start out.append((start, w)) start += w return out def real_nmse(acts, K, V) -> dict: check(chunks(896) == [(0, 256), (256, 256), (512, 256), (768, 128)], "896 = 3 x 256 + 128") check(len(chunks(4864)) == 19 and chunks(4864)[-1] == (4608, 256), "4864 = 19 x 256") rng = np.random.default_rng(20261006) sets = {"activations": [acts["qkv_in"], acts["o_in"], acts["gateup_in"], acts["down_in"]], "keys": [K], "values": [V]} signs = {name: [{cw: random_signs(rng, (X.shape[0], cw[1])) for cw in chunks(X.shape[1])} for X in arrs] for name, arrs in sets.items()} out = {} for name, arrs in sets.items(): for b in HALF: methods = [("none", False, "baseline"), ("pre", False, "torque_pre"), ("post", False, "torque_post"), ("both", False, "torque_both"), ("charged", False, "torque_charged")] if b in PACKED: methods.append(("none", True, "packed_baseline")) for stage, packed, label in methods: errs, bits, kept = [], [], [] for X, sg in zip(arrs, signs[name]): err, bt, kk = np.zeros(X.shape[0]), np.zeros(X.shape[0]), np.zeros(X.shape[0]) for (st, w) in chunks(X.shape[1]): Xc = X[:, st:st + w] Xh, bc, kc = encode_decode(Xc, b, stage, sg[(st, w)], packed) err += ((Xh - Xc) ** 2).sum(1) bt += bc kk += kc errs.append(err / (X * X).sum(1)) bits.append(bt / X.shape[1]) kept.append(kk) e, bt, kk = np.concatenate(errs), np.concatenate(bits), np.concatenate(kept) out[(name, b, label)] = (float(e.mean()), float(bt.mean()), float(bt.max()), float(kk.mean())) row("qwen_nmse", name, b, label, "nmse", float(e.mean())) row("qwen_nmse", name, b, label, "bits_per_coord_mean", float(bt.mean())) row("qwen_nmse", name, b, label, "pre_kept_per_vector_mean", float(kk.mean())) if name != "activations" and b == int(b): X = arrs[0] e = nmse(X, absmax_int(X, b)) out[(name, b, "absmax_int")] = (float(e.mean()), b, b, 0.0) row("qwen_nmse", name, b, "absmax_int", "nmse", float(e.mean())) for label in ("torque_pre", "torque_post", "torque_both", "torque_charged", "packed_baseline", "absmax_int"): if (name, b, label) in out: red = 100 * (1 - out[(name, b, label)][0] / out[(name, b, "baseline")][0]) row("qwen_nmse", name, b, label, "reduction_pct", red) print(" %-11s b=%.1f base %.4g both %.4g (%.1f%%) charged %.4g" % ( name, b, out[(name, b, "baseline")][0], out[(name, b, "torque_both")][0], 100 * (1 - out[(name, b, "torque_both")][0] / out[(name, b, "baseline")][0]), out[(name, b, "torque_charged")][0]), flush=True) return out PPL_BUDGETS = (4.0, 4.5, 5.0, 5.5, 6.0, 6.5, 7.0, 7.5, 8.0) def perplexity(model, segs, n_segments: int, budgets=PPL_BUDGETS) -> dict: """Quantize every key (after RoPE) and value vector inside attention; mean NLL per segment.""" import torch import torch.nn.functional as F from transformers.models.qwen2 import modeling_qwen2 as mq original = mq.eager_attention_forward state = {"cfg": None, "seg": 0} def qdq(t, layer, kind): b, method = state["cfg"] X = t[0].double().numpy().reshape(-1, 64) if method == "absmax_int": Xh = absmax_int(X, b) else: rng = np.random.default_rng([state["seg"], layer, kind]) Xh = encode_decode(X, b, "both" if method == "torque_both" else "none", random_signs(rng, X.shape))[0] return torch.from_numpy(Xh.reshape(t.shape)).to(t.dtype) def patched(module, query, key, value, attention_mask, scaling, dropout=0.0, **kw): if state["cfg"] is not None: key, value = qdq(key, module.layer_idx, 0), qdq(value, module.layer_idx, 1) return original(module, query, key, value, attention_mask, scaling=scaling, dropout=dropout, **kw) mq.eager_attention_forward = patched configs = [("fp32", None)] for b in budgets: configs += [("baseline", (b, "baseline")), ("torque_both", (b, "torque_both"))] if b == int(b): configs.append(("absmax_int", (b, "absmax_int"))) out = {} try: for label, cfg in configs: state["cfg"] = cfg nll = [] t0 = time.time() with torch.no_grad(): for i in range(n_segments): state["seg"] = i logits = model(segs[i:i + 1], use_cache=False).logits[0, :-1].float() nll.append(F.cross_entropy(logits, segs[i, 1:], reduction="mean").item()) b = "" if cfg is None else cfg[0] out[(label, b)] = np.array(nll) for i, v in enumerate(nll): row("qwen_ppl", "segment_%02d" % i, b, label, "mean_nll", v) row("qwen_ppl", "all", b, label, "perplexity", float(math.exp(np.mean(nll)))) write_csv() # keep finished configurations if the run stops print(" ppl %-16s b=%-4s %.4f (%.0f s)" % (label, b, math.exp(np.mean(nll)), time.time() - t0), flush=True) finally: mq.eager_attention_forward = original rng = np.random.default_rng(99) for (label, b), v in out.items(): if label in ("torque_both", "absmax_int"): diff = v - out[("baseline", b)] boots = [diff[rng.integers(0, len(diff), len(diff))].mean() for _ in range(10000)] lo, hi = np.percentile(boots, [2.5, 97.5]) row("qwen_ppl", "all", b, label, "nll_diff_vs_baseline_mean", float(diff.mean())) row("qwen_ppl", "all", b, label, "nll_diff_vs_baseline_ci_lo", float(lo)) row("qwen_ppl", "all", b, label, "nll_diff_vs_baseline_ci_hi", float(hi)) return out # ---------------------------------------------------------------- report and chart def worked_example() -> dict: """d = 256 at 4.5 bits: the best post-only plan against keeping four inputs exactly at c = 3, s = 4.""" d, b = 256, 4.5 p3 = p_out(3.0) T = d * (p3 * L_RET + (1 - p3) * 4.0) k = math.floor((d * b - T) / L_RET) post_only = min(((eps_s(c, s), c, s) for c, kk, s, _ in candidates(d, b, "post")), key=lambda z: z[0]) breakeven = 1 - post_only[0] / eps_s(3.0, 4.0) out = {"budget": d * b, "p3": p3, "T": T, "k": k, "post_only": post_only, "eps34": eps_s(3.0, 4.0), "breakeven_top4_share": breakeven} for key, v in (("budget_bits", d * b), ("p_out_c3", p3), ("plan_bits_before_retention", T), ("k_max", k), ("post_only_eps", post_only[0]), ("post_only_c", post_only[1]), ("post_only_s", post_only[2]), ("eps_c3_s4", eps_s(3.0, 4.0)), ("breakeven_top4_share", breakeven)): row("worked_example_d256_b4.5", "", 4.5, "", key, v) return out def draw_chart() -> None: """Draw the article's figure from the shipped CSV alone (also: --chart-only).""" import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt with open(CSV_PATH, encoding="utf-8") as fh: data = list(csv.DictReader(fh)) def series(experiment, name, method, metric): pts = sorted((float(r["budget_bits"]), float(r["value"])) for r in data if r["experiment"] == experiment and r["data"] == name and r["method"] == method and r["metric"] == metric) return [p[0] for p in pts], [p[1] for p in pts] gx, gy = series("gaussian_model", "N(0,1)", "torque_post", "reduction_pct") teps = dict(zip(*series("gaussian_model", "N(0,1)", "torque_post", "eps"))) px, peps = series("gaussian_model", "N(0,1)", "packed_baseline", "eps") surface, ink, ink2, grid, base = "#fcfcfb", "#0b0b0b", "#52514e", "#e1e0d9", "#c3c2b7" blue, orange, aqua = "#2a78d6", "#eb6834", "#1baf7a" # categorical slots 1-3, fixed order mk = dict(ms=6, mec=surface, mew=1.2, lw=1.6, solid_capstyle="round", solid_joinstyle="round") fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(10.4, 4.1), dpi=110, facecolor=surface) ax1.plot(gx, gy, color=blue, marker="o", label="vs mixing two codebooks", **mk) ax1.plot(px, [100 * (1 - teps[b] / e) for b, e in zip(px, peps)], color=orange, ls="none", marker="D", label="vs a packed code (half bits)", **mk) ax1.set_title("Gaussian model: keeping rotated values", color=ink, fontsize=10.5, loc="left") for name, col in (("activations", blue), ("keys", orange), ("values", aqua)): ax2.plot(*series("qwen_nmse", name, "torque_both", "reduction_pct"), color=col, marker="o", label=name, **mk) ax2.set_title("Qwen2.5-0.5B on WikiText-2, both stages", color=ink, fontsize=10.5, loc="left") for ax in (ax1, ax2): ax.set_facecolor(surface) ax.axhline(0, color=base, lw=1.0, zorder=1) ax.grid(True, color=grid, lw=0.6) ax.set_axisbelow(True) ax.set_xticks(range(2, 9)) ax.set_xlabel("expected bits per coordinate", color=ink2) ax.set_ylabel("NMSE reduction (%)", color=ink2) for sp in ("top", "right"): ax.spines[sp].set_visible(False) for sp in ("left", "bottom"): ax.spines[sp].set_color(base) ax.tick_params(colors=ink2, labelsize=8.5) ax.legend(frameon=False, fontsize=8.2, loc="upper left", labelcolor=ink) fig.tight_layout() os.makedirs(os.path.dirname(IMG_PATH), exist_ok=True) fig.savefig(IMG_PATH, dpi=110, facecolor=surface) plt.close(fig) def write_csv() -> None: os.makedirs(os.path.dirname(CSV_PATH), exist_ok=True) with open(CSV_PATH, "w", newline="", encoding="utf-8") as fh: w = csv.DictWriter(fh, fieldnames=["experiment", "data", "budget_bits", "method", "metric", "value", "note"], lineterminator="\n") w.writeheader() for r in ROWS: v = r["value"] w.writerow({**r, "value": ("%.10g" % v) if isinstance(v, float) else v}) def main() -> None: ap = argparse.ArgumentParser(description=__doc__, formatter_class=argparse.RawDescriptionHelpFormatter) ap.add_argument("--ppl-segments", type=int, default=64, help="256-token segments for perplexity (default 64)") ap.add_argument("--no-chart", action="store_true") ap.add_argument("--chart-only", action="store_true", help="redraw the figure from the existing CSV and stop") a = ap.parse_args() sys.stdout.reconfigure(encoding="utf-8") warnings.filterwarnings("ignore") if a.chart_only: draw_chart() return t0 = time.time() build_codebooks() check(abs(CB[(math.inf, 2)][2] - (1 - 2 / math.pi)) < 1e-9, "1-bit Lloyd-Max error is 1 - 2/pi") check(abs(CB[(math.inf, 2)][1][1] - math.sqrt(2 / math.pi)) < 1e-9, "1-bit level is sqrt(2/pi)") for L in POW2[1:]: check(CB[(math.inf, L)][2] < CB[(math.inf, L // 2)][2] / 3, "each extra bit cuts error by > 3x (L=%d)" % L) for c in C_SET[:-1]: check(CB[(c, 16)][2] < CB[(math.inf, 16)][2], "truncating at c=%s lowers the inlier error" % c) rng = np.random.default_rng(3) g = rng.standard_normal(2_000_000) for c, L in ((math.inf, 16), (2.5, 8), (math.inf, 22)): cuts, lev, e = CB[(c, L)] inl = np.abs(g) <= c mc = float(np.sum((lev[np.searchsorted(cuts, g[inl])] - g[inl]) ** 2) / g.size) check(abs(mc / e - 1) < 0.01, "exact error matches sampling within 1%% (c=%s, L=%d)" % (c, L)) print("codebooks: %d built in %.0f s" % (len(CB), time.time() - t0), flush=True) gm = gaussian_model() check(gm[2.0][2] == 0 and gm[3.0][2] == 0, "no Gaussian-model gain at 2 and 3 bits") check(all(gm[b][2] >= 0 for b in QUARTER), "retention never hurts under the model it optimises") check(all(gm[b + 0.25][2] > gm[b][2] for b in (3.0, 4.0, 5.0, 6.0, 7.0)), "sawtooth: gain jumps just above integers") monte_carlo(gm) ex = worked_example() check(ex["k"] == 4, "four inputs fit at c=3, s=4") syn = synthetic() check(syn[("lognormal", 10.0)]["pre"] < 0.05 * syn[("lognormal", 10.0)]["none"], "skewed: pre-rotation stage dominant") check(syn[("lognormal", 0.1)]["pre"] > 0.97 * syn[("lognormal", 0.1)]["none"], "near-Gaussian: pre stage idle") print("model + synthetic done in %.0f s" % (time.time() - t0), flush=True) write_csv() model, segs = load_model_and_text() acts, K, V = capture(model, segs) print("captured in %.0f s" % (time.time() - t0), flush=True) real = real_nmse(acts, K, V) for name in ("activations", "keys", "values"): for b in HALF: check(real[(name, b, "torque_both")][1] <= b + 1e-9, "%s b=%s: mean realized bits within budget" % (name, b)) check(real[(name, b, "torque_charged")][1] <= b + 1e-9, "%s b=%s: charged within budget" % (name, b)) check(abs(real[(name, b, "baseline")][1] - b) < 1e-9, "%s b=%s: baseline spends exactly b" % (name, b)) write_csv() if not a.no_chart: draw_chart() print("NMSE stage done in %.0f s" % (time.time() - t0), flush=True) perplexity(model, segs, a.ppl_segments) write_csv() print("rows %d, checks %d passed, %.0f s" % (len(ROWS), CHECKS, time.time() - t0)) if __name__ == "__main__": main()