diff options
Diffstat (limited to 'src/wu_row_norm_check.py')
| -rw-r--r-- | src/wu_row_norm_check.py | 189 |
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() |
