"""GPT-2 W_U row-norm frequency check. At-scale test of the jspace-nanogpt geometric finding: at char scale, the unembedding row norms anti-correlate with token frequency (r(||W_U[k]||, log10 f) = -0.693, Spearman -0.808) and this geometry is LEARNED (r ~ 0 at init). Does it generalize to a real LM with V = 50,257? Cheap by design: load lm_head/wte rows, take row norms, correlate against unigram token counts from a corpus. No Jacobian, no training, no GPU. Usage: .venv/bin/python src/wu_row_norm_check.py [--models gpt2,gpt2-medium] Corpus: wikitext-103-raw-v1 (train) from HF. Counts are token-level unigram frequencies after GPT-2 byte-level BPE encoding. """ import argparse import time import numpy as np import safetensors.torch import torch from huggingface_hub import hf_hub_download from transformers import GPT2TokenizerFast WIKITEXT_REPO = "Salesforce/wikitext" WIKITEXT_PARQUETS = [ "wikitext-103-raw-v1/train-00000-of-00002.parquet", "wikitext-103-raw-v1/train-00001-of-00002.parquet", ] def pearson(x, y): x = x - x.mean() y = y - y.mean() denom = np.sqrt((x * x).sum() * (y * y).sum()) return float((x * y).sum() / denom) if denom > 0 else float("nan") def spearman(x, y): def rankdata(a): order = np.argsort(a, kind="mergesort") ranks = np.empty_like(order, dtype=float) ranks[order] = np.arange(1, a.size + 1) return ranks return pearson(rankdata(x), rankdata(y)) def load_wte(model_id: str) -> np.ndarray: path = hf_hub_download(model_id, "model.safetensors") st = safetensors.torch.load_file(path) candidates = ["transformer.wte.weight", "wte.weight", "model.embed_tokens.weight"] key = next((k for k in candidates if k in st), None) if key is None: raise KeyError( f"no wte/embedding weight found in {model_id}; " f"have: {sorted(st.keys())[:10]} ..." ) wte = st[key] print(f"[{model_id}] {key} {tuple(wte.shape)} dtype={wte.dtype}") # GPT-2 ties lm_head to wte, so row norms of wte == row norms of W_U. return torch.norm(wte.float(), dim=1).numpy() def load_corpus() -> list[str]: import pyarrow.parquet as pq texts: list[str] = [] for name in WIKITEXT_PARQUETS: path = hf_hub_download(WIKITEXT_REPO, name, repo_type="dataset") print(f"loading corpus shard from {path}") table = pq.read_table(path, columns=["text"]) texts.extend(table.column("text").to_pylist()) print(f"corpus: {len(texts)} rows") return texts def count_tokens(tokenizer, texts: list[str], batch_size: int = 1000) -> np.ndarray: counts = np.zeros(tokenizer.vocab_size, dtype=np.int64) n_tokens = 0 t0 = time.time() for start in range(0, len(texts), batch_size): chunk = texts[start : start + batch_size] enc = tokenizer(chunk, add_special_tokens=False) for ids in enc["input_ids"]: np.add.at(counts, ids, 1) n_tokens += len(ids) if (start // batch_size) % 20 == 0: print( f" rows {start}/{len(texts)} tokens {n_tokens:,} " f"elapsed {time.time() - t0:.0f}s" ) print(f"done: {n_tokens:,} tokens in {time.time() - t0:.0f}s") return counts def report(model_id: str, norms: np.ndarray, counts: np.ndarray, tokenizer): freqs = counts.astype(np.float64) seen = freqs > 0 logf = np.log10(freqs + 1.0) # smoothed; +1 keeps unseen tokens finite def r_raw(mask): return pearson(freqs[mask], norms[mask]) def r_log(mask): return pearson(logf[mask], norms[mask]) def r_sp(mask): return spearman(freqs[mask], norms[mask]) n = norms.size n_seen = int(seen.sum()) print(f"\n===== {model_id} ===== V={n} tokens seen={n_seen} " f"({100 * n_seen / n:.1f}%)") print(f" r(raw freq, norm) all V: {r_raw(np.ones(n, bool)):+.3f}") print(f" r(raw freq, norm) seen only: {r_raw(seen):+.3f}") print(f" r(log10 freq, norm) all V: {r_log(np.ones(n, bool)):+.3f}") print(f" r(log10 freq, norm) seen only: {r_log(seen):+.3f}") for thr in (5, 100, 1000): m = freqs >= thr print(f" r(log10 freq, norm) freq>={thr:<6} " f"(n={int(m.sum()):>5}): {r_log(m):+.3f}") print(f" Spearman(raw, norm) seen only: {r_sp(seen):+.3f}") # Decile table: quantiles of log-frequency vs mean norm. if n_seen > 10: qs = np.quantile(logf[seen], np.linspace(0, 1, 11)) idx = np.digitize(logf, qs[1:-1]) print(" log10-freq decile -> mean norm:") for d in range(10): m = (idx == d) & seen if m.sum() == 0: continue lo, hi = qs[d], qs[d + 1] print(f" [{lo:5.2f},{hi:5.2f}] n={int(m.sum()):>5} " f"mean norm={norms[m].mean():.3f}") # Top tokens by norm (narrative color). order = np.argsort(norms)[::-1][:15] print(" top-15 tokens by row norm (token | norm | count | rank-by-freq):") for i in order: tok = tokenizer.decode([int(i)]) rank = int((freqs > freqs[i]).sum()) + 1 print(f" {tok!r:>14} {norms[i]:.3f} {int(freqs[i]):>8,} #{rank:,}") return { "model": model_id, "r_raw_all": r_raw(np.ones(n, bool)), "r_log_all": r_log(np.ones(n, bool)), "r_log_seen": r_log(seen), "spearman_seen": r_sp(seen), "n_seen": n_seen, } def main(): ap = argparse.ArgumentParser() ap.add_argument("--models", default="gpt2", help="comma-separated HF ids") ap.add_argument("--no-corpus", action="store_true", help="skip tokenization (use saved counts if present)") args = ap.parse_args() tokenizer = GPT2TokenizerFast.from_pretrained("openai-community/gpt2") if args.no_corpus: counts = np.load("/tmp/wu_counts.npy") else: texts = load_corpus() counts = count_tokens(tokenizer, texts) np.save("/tmp/wu_counts.npy", counts) results = [] for mid in args.models.split(","): mid = mid.strip() norms = load_wte(mid) results.append(report(mid, norms, counts, tokenizer)) print("\n===== SUMMARY =====") for r_ in results: print( f"{r_['model']}: r(raw)={r_['r_raw_all']:+.3f} " f"r(log10, all V)={r_['r_log_all']:+.3f} " f"r(log10, seen)={r_['r_log_seen']:+.3f} " f"Spearman(seen)={r_['spearman_seen']:+.3f}" ) if __name__ == "__main__": main()