summaryrefslogtreecommitdiff
path: root/src/wu_row_norm_check.py
diff options
context:
space:
mode:
Diffstat (limited to 'src/wu_row_norm_check.py')
-rw-r--r--src/wu_row_norm_check.py189
1 files changed, 189 insertions, 0 deletions
diff --git a/src/wu_row_norm_check.py b/src/wu_row_norm_check.py
new file mode 100644
index 0000000..e7e7473
--- /dev/null
+++ b/src/wu_row_norm_check.py
@@ -0,0 +1,189 @@
+"""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()