From bc7ddeaf4b8d47d423c7ac8cc5b7d3711b849076 Mon Sep 17 00:00:00 2001 From: helpmatteo Date: Mon, 31 Aug 2026 17:44:45 +0200 Subject: [PATCH] refactor: consolidate the ported experiment code onto shared helpers (0.2.1) Behaviour-preserving cleanup of the code shipped in 0.2.0, verified against the paper's own run artifacts where they exist locally (cross-lens summaries and worked-example tables, error tails, paired matched-KL records, DE bank: byte-identical apart from added provenance). Structure - research/cross_lens.py, research/wsd.py, research/row_geometry.py hold the helpers that were copied across scripts (jlens import, model/lens loading, dump parsing, Wilson CIs, agreement rule, bundle IO, cluster bootstrap, spherical k-means, bare-first token resolver); scripts route through qwen_readout.load_sae / encode_topk, factorizers.topk_mask, data.center_normalize_rows, utils.load_causal_lm / write_json, run_io.run_provenance; one script loader in tests/conftest.py. - Unused AmbiStory / anchors paths removed from the WSD scripts. Paper-level choices made explicit (defaults reproduce the paper) - --centering {live,trained} on the cross-lens, causal, stability and WSD scripts (data.centering_mean); on Qwen3.5-2B the two means differ by ~0.04% of a centred row norm. - --agreement-rule {half,strict} and --null-population {all,cross} on the cross-lens aggregators; --knn-exclude-self on loo_core_recovery. Fixed - feature_group_matching: null pool no longer depends on PYTHONHASHSEED; below-null tail uses the unrounded Jaccard. - run_readout_baseline_comparisons: per-row-support methods aligned before A-B subtraction (coverage/compactness columns only); k-means fitted once per (n_clusters, seed); determinism claim corrected. - fit_jlens: OOM fallback reaches dim_batch 1; shard/merge resume guarded by metadata sidecars. fit_ridge_lens: --holdout >= 1, holdout-only diagnostic, guarded resume. - run_wsd_feature_alignment: left truncation in cloze mode; --analyze-only no longer overwrites run_config.json. run_causal_contribution_validation: self-test's realized-change check compares against the dense LM head. - nearest_rows_baseline: empty tables no longer crash. - analyze_error_tails bootstrap vectorised (identical draws, 3x faster). Docs: THIRD_PARTY.md lists jlens (Apache-2.0, git source), pandas, scikit-learn, CoarseWSD-20 and the C4 lens corpora; config headers and REPRODUCE rows corrected; version 0.2.1. Co-Authored-By: Claude Fable 5 --- CITATION.cff | 2 +- README.md | 13 +- configs/sweeps/README.md | 16 +- configs/sweeps/lowk_qwen35_0p8b_16x_base.yaml | 5 + configs/sweeps/lowk_qwen35_2b_16x_base.yaml | 5 + .../qwen35_0p8b_paper_topk_32x_k256_s1.yaml | 22 +- .../qwen35_0p8b_paper_topk_32x_k256_s2.yaml | 20 +- configs/sweeps/qwen35_2b_seedvar_base.yaml | 5 + data/cross_lens/README.md | 8 +- data/wsd/wsd_sense_anchors.json | 101 -- docs/DATA.md | 6 +- docs/REPRODUCE.md | 40 +- docs/THIRD_PARTY.md | 49 +- pyproject.toml | 4 +- scripts/README.md | 2 + .../analyze/analyze_wsd_classifier_framing.py | 415 +------ scripts/analyze/analyze_wsd_sense_groups.py | 217 ++-- scripts/analyze/cross_lens_antonym_layers.py | 172 ++- .../analyze/cross_lens_three_lens_prompt.py | 36 +- scripts/analyze/nearest_rows_baseline.py | 58 +- scripts/data/build_cross_lens_de_bank.py | 44 +- scripts/data/extract_model_readout.py | 9 +- scripts/eval/aggregate_cross_lens_en_de.py | 227 ++-- scripts/eval/aggregate_cross_lens_en_zh.py | 183 +-- scripts/eval/analyze_error_tails.py | 75 +- scripts/eval/cross_seed_stability.py | 63 +- scripts/eval/feature_group_matching.py | 106 +- scripts/eval/loo_core_recovery.py | 107 +- scripts/eval/paired_matched_kl_bootstrap.py | 16 +- .../run_causal_contribution_validation.py | 107 +- .../compute_cross_lens_shared_feature.py | 22 +- scripts/run/fit_jlens.py | 188 ++- scripts/run/fit_ridge_lens.py | 137 ++- scripts/run/run_cross_lens_readouts.py | 177 ++- .../run_qwen_profanity_suppression_eval.py | 3 +- .../run/run_readout_baseline_comparisons.py | 293 ++--- scripts/run/run_wsd_feature_alignment.py | 804 +++---------- src/sparse_readout_prism/data.py | 35 + src/sparse_readout_prism/research/__init__.py | 6 +- .../research/cross_lens.py | 582 +++++++++ .../research/row_geometry.py | 80 ++ .../research/seed_stability.py | 170 ++- src/sparse_readout_prism/research/wsd.py | 220 ++++ tests/conftest.py | 27 + tests/test_baseline_and_dla.py | 148 ++- tests/test_causal_contribution_validation.py | 75 +- tests/test_cross_lens_scripts.py | 627 +++++++++- tests/test_error_tails_and_nearest_rows.py | 161 ++- tests/test_lexical_control_directions.py | 19 +- tests/test_paired_matched_kl_bootstrap.py | 43 +- tests/test_seed_stability_scripts.py | 461 ++++++- tests/test_wsd_sense_groups.py | 1070 +++++++++++++++-- uv.lock | 2 +- 53 files changed, 4915 insertions(+), 2568 deletions(-) delete mode 100644 data/wsd/wsd_sense_anchors.json create mode 100644 src/sparse_readout_prism/research/cross_lens.py create mode 100644 src/sparse_readout_prism/research/row_geometry.py create mode 100644 src/sparse_readout_prism/research/wsd.py create mode 100644 tests/conftest.py diff --git a/CITATION.cff b/CITATION.cff index 4a37c68..e047683 100644 --- a/CITATION.cff +++ b/CITATION.cff @@ -23,7 +23,7 @@ authors: family-names: Lane affiliation: University of Cambridge license: MIT -version: "0.2.0" +version: "0.2.1" date-released: "2026-08-31" repository-code: "https://github.com/hematteo/sparse-readout-prism" preferred-citation: diff --git a/README.md b/README.md index 723a87a..6b98ac3 100644 --- a/README.md +++ b/README.md @@ -66,6 +66,15 @@ script and command line in [`docs/REPRODUCE.md`](docs/REPRODUCE.md) §2): with the seed-variation and low-k training configs; - the sense-labelled evaluation on CoarseWSD-20. +0.2.1 consolidates that code onto the repo's shared helpers (`research/cross_lens.py`, +`research/wsd.py`, `research/row_geometry.py`; the library's checkpoint reader, +top-k kernel and centering helper), fixes the small bugs found in review, records +provenance in every result file, and exposes the choices the paper runs made +implicitly as flags whose defaults reproduce the paper: `--centering {live,trained}` +(full-vocabulary vs. training row mean), `--agreement-rule {half,strict}` and +`--null-population {all,cross}` (cross-lens aggregators), and +`--knn-exclude-self` (leave-one-out core recovery). + ## Install Requires Python 3.11 or 3.12. @@ -211,7 +220,8 @@ src/sparse_readout_prism/ core library research/ flat helpers shared by several scripts/ entry points (Qwen readout + query-decomposition toolkits, prompt banks, registry, run IO, - seed-stability contrast pipeline); + seed-stability pipeline, row-geometry helpers, + cross-lens toolkit, CoarseWSD-20 helpers); logic used by a single script stays inline in that script @@ -241,7 +251,6 @@ configs/ data/ query_banks/ curated prompt banks (paper inputs) cross_lens/ EN-ZH / EN-DE cross-lens prompt banks + exclusion report - wsd/ optional sense-gloss anchors for the CoarseWSD-20 runs appendix/ audited literals behind the Appendix K tables audit/ feature-label audit annotations + counts (Appendix L) diff --git a/configs/sweeps/README.md b/configs/sweeps/README.md index 43fa52d..76c8e27 100644 --- a/configs/sweeps/README.md +++ b/configs/sweeps/README.md @@ -38,7 +38,7 @@ Two kinds of file live here: | `deepseek_r1_distill_llama_8b_paper.yaml`, `…_16x_k128.yaml` | run config | R1-Distill-Llama-8B finalists. | | `ministral3_8b_paper.yaml`, `…_16x_k128.yaml` | run config | Ministral-3-8B-Base finalists (32×/k256 and strict-budget 16×/k128). | | `qwen35_2b_seedvar_base.yaml` | run config | Qwen3.5-2B seed-variation base (32×/k256, seed 0). The stability appendix's three-seed families and the multi-seed rows of the low-k table are this file with `--set` overrides of `run.seed` / `run.init_seed` / `data.data_seed`, `factorizer.d_features` / `factorizer.k` / `evaluation.k`, and `run.name` / `run.output_dir` (cells listed in the file header). | -| `qwen35_0p8b_paper_topk_32x_k256_s1.yaml`, `…_s2.yaml` | run config | Qwen3.5-0.8B 32×/k256 seed-window cells (seeds 1 and 2; seed 0 is the archived finalist). | +| `qwen35_0p8b_paper_topk_32x_k256_s1.yaml`, `…_s2.yaml` | run config | Qwen3.5-0.8B 32×/k256 seed-window cells (seeds 1 and 2; seed 0 is the archived finalist, `qwen0p8b_k256` in the exp2 registry / `qwen0p8b_k256_32x` in the fidelity registry). Recipe matched on the keys this trainer reads (see the file headers). | | `lowk_qwen35_2b_16x_base.yaml`, `lowk_qwen35_0p8b_16x_base.yaml` | run config | Low activation-budget cells at 16× width (`k` in {32, 64} via `--set factorizer.k=… --set evaluation.k=…`), behind the sparsity-budget table. | The provisional Gemma operating points in the registries were trained with the @@ -56,10 +56,16 @@ by the run-config guard, before any data is loaded.) The run files behind the paper's seed-variation and low-k cells carried a few extra keys (`data.row_preprocessing`, `training.init_mode`, `lr_schedule`, `lambda_mode`, `final_lr_frac`, `auxk_*`, `artifacts.make_figures`) that no -code path in the shipped trainer reads, and did not read in the trainer that -produced those cells either; they are omitted here rather than shipped as if -they were live. Every value the trainer does read is identical to the paper -files. +code path in the shipped trainer reads. The Qwen3.5-2B seed-variation and +low-k cells were trained with this trainer, so for them the keys were inert; +the two Qwen3.5-0.8B seed cells were launched from an earlier trainer tree +that is not available to check, and the released paper dictionaries record +the same keys in their checkpoint configs (see the `config` block of any +Hub checkpoint), so treat "identical recipe" for those as covering the keys +this trainer reads. The keys are omitted here rather than shipped as if they +were live. The five seed-variation / low-k run configs carry +`training.log_every: 250` and `training.checkpoint_every: 5000` (logging and +checkpoint cadence only) for parity with `paper_phase2_base.yaml`. The cross-model finalist configs (R1-Distill-Qwen-7B, R1-Distill-Llama-8B, Ministral) are launched with the generic trainer — diff --git a/configs/sweeps/lowk_qwen35_0p8b_16x_base.yaml b/configs/sweeps/lowk_qwen35_0p8b_16x_base.yaml index 71e1657..82fe73f 100644 --- a/configs/sweeps/lowk_qwen35_0p8b_16x_base.yaml +++ b/configs/sweeps/lowk_qwen35_0p8b_16x_base.yaml @@ -23,6 +23,9 @@ # data.path is the {W_U_orig, h_LN} extraction of Qwen3.5-0.8B # (scripts/data/extract_model_readout.py), the artifact # configs/registries/result1_query_fidelity_cluster.yaml resolves for this model. +# training.log_every / training.checkpoint_every (250 / 5000) are logging and +# checkpoint cadence only, added for parity with paper_phase2_base.yaml; they do +# not change the trained dictionary. run: name: lowk_qwen35_0p8b_16x @@ -54,6 +57,8 @@ training: prism_top_n: 1 lambda_prism: 0.001 row_sampling: hybrid_50freq_50uniform + log_every: 250 # cadence only: parity with paper_phase2_base.yaml + checkpoint_every: 5000 # cadence only: parity with paper_phase2_base.yaml evaluation: k: 64 diff --git a/configs/sweeps/lowk_qwen35_2b_16x_base.yaml b/configs/sweeps/lowk_qwen35_2b_16x_base.yaml index ccbb782..06967ea 100644 --- a/configs/sweeps/lowk_qwen35_2b_16x_base.yaml +++ b/configs/sweeps/lowk_qwen35_2b_16x_base.yaml @@ -26,6 +26,9 @@ # data.path is the {W_U_orig, h_LN} extraction of Qwen3.5-2B # (scripts/data/extract_model_readout.py), the artifact # configs/registries/result1_query_fidelity_cluster.yaml resolves for this model. +# training.log_every / training.checkpoint_every (250 / 5000) are logging and +# checkpoint cadence only, added for parity with paper_phase2_base.yaml; they do +# not change the trained dictionary. run: name: lowk_qwen35_2b_16x @@ -57,6 +60,8 @@ training: prism_top_n: 1 lambda_prism: 0.001 row_sampling: hybrid_50freq_50uniform + log_every: 250 # cadence only: parity with paper_phase2_base.yaml + checkpoint_every: 5000 # cadence only: parity with paper_phase2_base.yaml evaluation: k: 64 diff --git a/configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s1.yaml b/configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s1.yaml index 7346d69..1bd49d6 100644 --- a/configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s1.yaml +++ b/configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s1.yaml @@ -1,8 +1,17 @@ # Qwen3.5-0.8B readout SAE, 32x/k=256 recipe, seed 1 (tab:app-k-qwen08b-seed-window). -# Repeats the archived seed-0 dictionary setting (qwen0p8b_k256_32x in -# configs/registries/exp2_selected_sae_checkpoints.yaml) with only -# run.seed / run.init_seed / data.data_seed changed. Sister file: -# qwen35_0p8b_paper_topk_32x_k256_s2.yaml (seed 2). +# Matches the archived seed-0 dictionary setting (qwen0p8b_k256 in +# configs/registries/exp2_selected_sae_checkpoints.yaml; the same checkpoint is +# qwen0p8b_k256_32x in configs/registries/result1_query_fidelity_cluster.yaml) +# on every key this trainer reads, with only run.seed / run.init_seed / +# data.data_seed changed. Sister file: qwen35_0p8b_paper_topk_32x_k256_s2.yaml +# (seed 2). +# +# The archived run file also carried keys this trainer does not read +# (data.row_preprocessing, training.init_mode / lr_schedule / lambda_mode / +# final_lr_frac / auxk_*, artifacts.make_figures). The seed-0 cell was launched +# from an earlier trainer tree that is not available to check, so whether those +# keys were live there is not verified; they are omitted here, and "same recipe" +# means the keys this trainer reads. # # Launched as-is, no --set overrides: # uv run python scripts/train/train_readout_sae_from_config.py \ @@ -11,6 +20,9 @@ # d_features 32768 = 32 x d_model (1024). data.path is the {W_U_orig, h_LN} # extraction of Qwen3.5-0.8B (scripts/data/extract_model_readout.py), the artifact # configs/registries/result1_query_fidelity_cluster.yaml resolves for this model. +# training.log_every / training.checkpoint_every (250 / 5000) are logging and +# checkpoint cadence only, added for parity with paper_phase2_base.yaml; they do +# not change the trained dictionary. run: name: qwen35_0p8b_paper_topk_32x_k256_s1 @@ -42,6 +54,8 @@ training: prism_top_n: 1 lambda_prism: 0.001 row_sampling: hybrid_50freq_50uniform + log_every: 250 # cadence only: parity with paper_phase2_base.yaml + checkpoint_every: 5000 # cadence only: parity with paper_phase2_base.yaml evaluation: k: 256 diff --git a/configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s2.yaml b/configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s2.yaml index 444ac20..727380f 100644 --- a/configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s2.yaml +++ b/configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s2.yaml @@ -1,8 +1,17 @@ # Qwen3.5-0.8B readout SAE, 32x/k=256 recipe, seed 2 (tab:app-k-qwen08b-seed-window). # Sister file of qwen35_0p8b_paper_topk_32x_k256_s1.yaml: the archived seed-0 -# dictionary setting (qwen0p8b_k256_32x in -# configs/registries/exp2_selected_sae_checkpoints.yaml) with only -# run.seed / run.init_seed / data.data_seed changed. +# dictionary setting (qwen0p8b_k256 in +# configs/registries/exp2_selected_sae_checkpoints.yaml; the same checkpoint is +# qwen0p8b_k256_32x in configs/registries/result1_query_fidelity_cluster.yaml) +# on every key this trainer reads, with only run.seed / run.init_seed / +# data.data_seed changed. +# +# The archived run file also carried keys this trainer does not read +# (data.row_preprocessing, training.init_mode / lr_schedule / lambda_mode / +# final_lr_frac / auxk_*, artifacts.make_figures). The seed-0 cell was launched +# from an earlier trainer tree that is not available to check, so whether those +# keys were live there is not verified; they are omitted here, and "same recipe" +# means the keys this trainer reads. # # Launched as-is, no --set overrides: # uv run python scripts/train/train_readout_sae_from_config.py \ @@ -11,6 +20,9 @@ # d_features 32768 = 32 x d_model (1024). data.path is the {W_U_orig, h_LN} # extraction of Qwen3.5-0.8B (scripts/data/extract_model_readout.py), the artifact # configs/registries/result1_query_fidelity_cluster.yaml resolves for this model. +# training.log_every / training.checkpoint_every (250 / 5000) are logging and +# checkpoint cadence only, added for parity with paper_phase2_base.yaml; they do +# not change the trained dictionary. run: name: qwen35_0p8b_paper_topk_32x_k256_s2 @@ -42,6 +54,8 @@ training: prism_top_n: 1 lambda_prism: 0.001 row_sampling: hybrid_50freq_50uniform + log_every: 250 # cadence only: parity with paper_phase2_base.yaml + checkpoint_every: 5000 # cadence only: parity with paper_phase2_base.yaml evaluation: k: 256 diff --git a/configs/sweeps/qwen35_2b_seedvar_base.yaml b/configs/sweeps/qwen35_2b_seedvar_base.yaml index 2e52a3a..0750723 100644 --- a/configs/sweeps/qwen35_2b_seedvar_base.yaml +++ b/configs/sweeps/qwen35_2b_seedvar_base.yaml @@ -28,6 +28,9 @@ # d_features 65536 = 32 x d_model (2048). data.path is the {W_U_orig, h_LN} # extraction of Qwen3.5-2B (scripts/data/extract_model_readout.py), the artifact # configs/registries/result1_query_fidelity_cluster.yaml resolves for this model. +# training.log_every / training.checkpoint_every (250 / 5000) are logging and +# checkpoint cadence only, added for parity with paper_phase2_base.yaml; they do +# not change the trained dictionary. run: name: qwen35_2b_seedvar_base @@ -59,6 +62,8 @@ training: prism_top_n: 1 lambda_prism: 0.001 row_sampling: hybrid_50freq_50uniform + log_every: 250 # cadence only: parity with paper_phase2_base.yaml + checkpoint_every: 5000 # cadence only: parity with paper_phase2_base.yaml evaluation: k: 256 diff --git a/data/cross_lens/README.md b/data/cross_lens/README.md index c5b8196..a841eb1 100644 --- a/data/cross_lens/README.md +++ b/data/cross_lens/README.md @@ -18,10 +18,10 @@ Every record has these keys (`lang_a` / `lang_b` only in the EN-DE bank): | `concept` | str | The concept the prompt elicits; used to exclude same-concept pairs from the shuffled-pairing null. | | `prompt` | str | The full prompt fed to the model; the final position is scored. | | `form_a` | str | The in-context answer form (the language or script the prompt asks for). | -| `form_b` | str | The same concept's canonical form in the other language. | -| `targets` | list[str] | Every surface the runner decomposes: `form_a`, `form_b`, any extra forms, then the nulls. | -| `null_targets` | list[str] | Unrelated tokens (script- or language-matched) for the unrelated-token null floor. | -| `answer` | str | The expected answer (`form_a` without its leading space). | +| `form_b` | str | The same concept's canonical form in the other language. Controls in the EN-DE bank repeat `form_a` here (`7`, ` Paris`), since their surface does not depend on language. | +| `targets` | list[str] | Every surface the runner decomposes, `null_targets` last. Opens with `form_a`; `form_b` follows, but extra forms may sit between them (`cle_03`: ` large`, ` big`, `大`; the three-lens bank lists the answer's own variants before the other-language forms) or after `form_b` (EN-ZH controls: ` 7`, `7`). EN-DE controls therefore list `form_a` twice (`['7', '7', ' Schuh', ' shoe']`). | +| `null_targets` | list[str] | Unrelated tokens (script- or language-matched) for the unrelated-token null floor. Empty for the EN-ZH controls. | +| `answer` | str | The expected continuation. Usually `form_a` without its leading space; for three EN-ZH items the answer is a multi-character word whose single-token first character is `form_a` (`clz_09` 睡觉 / 睡, `clz_16` 月亮 / 月, `exm_03` 红色 / 红), so the probed token is the answer's prefix. | | `lang_a`, `lang_b` | str | EN-DE bank only: the language of `form_a` / `form_b` (`de` or `en`). | Each target string is probed by the id of its first token, and the aggregators diff --git a/data/wsd/wsd_sense_anchors.json b/data/wsd/wsd_sense_anchors.json deleted file mode 100644 index 43bf4a5..0000000 --- a/data/wsd/wsd_sense_anchors.json +++ /dev/null @@ -1,101 +0,0 @@ -{ - "_meta": { - "purpose": "Frozen sense-anchor token sets for weights-only feature typing (CoarseWSD-20). Typing rule: beta_f(sense) = mean anchor code (one-vs-rest contrast); a feature is sense-typed when its max one-vs-rest beta exceeds tau from the random-anchor null.", - "rule": "Per sense: the class-map sense name plus curated common-word associates; the ambiguous target itself is never an anchor; no anchor is reused across senses of the same word; anchors are filtered per tokenizer to leading-space single tokens at runtime and the surviving counts are reported.", - "frozen": "2026-08-01, before any occurrence-mode run was executed", - "sense_ids": "match CoarseWSD-20 class_map.txt per word" - }, - "apple": { - "0": ["iPhone", "Mac", "software", "computer", "technology", "company"], - "1": ["fruit", "orchard", "pear", "cherry", "juice", "ripe"] - }, - "arm": { - "0": ["processor", "chip", "architecture", "Intel", "silicon", "computing"], - "1": ["shoulder", "elbow", "hand", "leg", "wrist", "limb"] - }, - "bank": { - "0": ["money", "loan", "deposit", "financial", "credit", "account"], - "1": ["river", "shore", "stream", "flood", "water", "erosion"] - }, - "bass": { - "0": ["guitar", "riff", "band", "rock", "amplifier"], - "1": ["singer", "tenor", "choir", "opera", "voice"], - "2": ["cello", "orchestra", "violin", "symphony", "concerto"] - }, - "bow": { - "0": ["ship", "stern", "hull", "vessel", "sailors"], - "1": ["arrow", "archer", "quiver", "shooting", "hunting"], - "2": ["violin", "strings", "cello", "rosin", "instrument"] - }, - "chair": { - "0": ["committee", "president", "board", "meeting", "elected"], - "1": ["table", "furniture", "seat", "cushion", "wooden"] - }, - "club": { - "0": ["society", "members", "organization", "association", "founded"], - "1": ["nightclub", "dancing", "bar", "DJ", "party"], - "2": ["weapon", "stick", "bat", "blunt", "wielded"] - }, - "crane": { - "0": ["machinery", "lifting", "construction", "tower", "load", "hoist"], - "1": ["bird", "heron", "stork", "wetland", "feathers", "migratory"] - }, - "deck": { - "0": ["ship", "sailors", "hull", "aboard", "vessel"], - "1": ["porch", "patio", "balcony", "wooden", "backyard"] - }, - "digit": { - "0": ["number", "numeral", "decimal", "integer", "arithmetic"], - "1": ["finger", "toe", "thumb", "hand", "limb"] - }, - "hood": { - "0": ["villain", "comics", "Marvel", "superhero", "character"], - "1": ["car", "engine", "bonnet", "vehicle", "windshield"], - "2": ["cloak", "jacket", "head", "garment", "sweatshirt"] - }, - "java": { - "0": ["Indonesia", "island", "coffee", "Jakarta", "volcano"], - "1": ["programming", "software", "code", "compiler", "language"] - }, - "mole": { - "0": ["burrow", "rodent", "digging", "fur", "underground"], - "1": ["spy", "agent", "intelligence", "undercover", "espionage"], - "2": ["molar", "chemistry", "gram", "concentration", "atoms"], - "3": ["sauce", "Mexican", "chili", "chocolate", "cuisine"], - "4": ["pier", "harbor", "breakwater", "jetty", "seawall"] - }, - "pitcher": { - "0": ["baseball", "innings", "strikeout", "mound", "batter"], - "1": ["jug", "water", "pour", "ceramic", "vessel"] - }, - "pound": { - "0": ["kilogram", "weight", "ounce", "mass", "grams"], - "1": ["sterling", "currency", "pence", "British", "dollar"] - }, - "seal": { - "0": ["marine", "whiskers", "blubber", "colony", "swimming"], - "1": ["singer", "album", "song", "Grammy", "vocalist"], - "2": ["emblem", "stamp", "official", "insignia", "wax"], - "3": ["gasket", "valve", "leak", "pressure", "rubber"] - }, - "spring": { - "0": ["water", "geyser", "mineral", "well", "flowing"], - "1": ["summer", "winter", "autumn", "season", "blossom"], - "2": ["coil", "mechanical", "tension", "metal", "compression"] - }, - "square": { - "0": ["rectangle", "triangle", "geometry", "shape", "sides"], - "1": ["company", "corporation", "payments", "games"], - "2": ["plaza", "town", "market", "city", "fountain"], - "3": ["squared", "multiplication", "integer", "root", "arithmetic"] - }, - "trunk": { - "0": ["tree", "bark", "branches", "roots", "wood"], - "1": ["car", "luggage", "rear", "vehicle", "storage"], - "2": ["torso", "body", "abdomen", "muscles"] - }, - "yard": { - "0": ["meter", "feet", "inches", "garden", "lawn"], - "1": ["mast", "sail", "rigging", "spar", "ship"] - } -} diff --git a/docs/DATA.md b/docs/DATA.md index 0dc78b7..28527aa 100644 --- a/docs/DATA.md +++ b/docs/DATA.md @@ -15,14 +15,12 @@ A record is one **case** — a prompt, a position to score, two targets (`A`, `B`), and an expected side. Every script that consumes a bank reads the same schema; the banks are model-agnostic. -Two further paper-input sets live beside them. [`data/cross_lens/`](../data/cross_lens/) +One further paper-input set lives beside them. [`data/cross_lens/`](../data/cross_lens/) holds the EN–ZH and EN–DE prompt banks of the cross-lens study (one record per prompt: two surface forms, an unrelated null target, family and control tags) together with the EN–DE cognate-exclusion report and the one-prompt bank of the three-lens worked example; schema and provenance are in -[`data/cross_lens/README.md`](../data/cross_lens/README.md). [`data/wsd/`](../data/wsd/) -holds the optional sense-gloss anchors that `scripts/run/run_wsd_feature_alignment.py` -accepts through `--anchors`; the paper's CoarseWSD-20 runs did not use them. +[`data/cross_lens/README.md`](../data/cross_lens/README.md). ## 2. Extracted readouts (per-model artifacts) diff --git a/docs/REPRODUCE.md b/docs/REPRODUCE.md index 873304c..817107f 100644 --- a/docs/REPRODUCE.md +++ b/docs/REPRODUCE.md @@ -148,7 +148,7 @@ check. | Table / metrics behind: benchmark-derived task targets | `scripts/run/run_benchmark_derived_query_suite.py`, `scripts/figures/compute_benchmark_task_group_examples.py` | `--out-dir results/benchmark_derived_query_suite_task_group_v2_qwen_20260523` (the suite uses task-group contrasts — the paper's `task_group_v2` design — and does not derive the dir from it; pass `--out-dir` explicitly) | | Metrics behind: selected-score / family / selection queries (`fig:qwen-general-readout-scores-basic`, `fig:app-qwen-general-readout-scores-extra`) | `scripts/figures/compute_general_readout_queries.py` | query specs are built in (the paper's jury/guilty panels plus the abstention-family contrast); `--checkpoint` points at the Qwen3.5-2B 32x/k256 selected SAE in `configs/registries/result1_query_fidelity_cluster.yaml` | | Metrics behind: literature-prompt all-layer trace | `scripts/figures/compute_all_layer_literature_prompt.py` | requires `--checkpoint` (the Qwen3.5-2B 32x/k256 selected SAE); prompt/model defaults reproduce the paper trace | -| `tab:app-direct-geometry-grid`; SRP / best-alternative / lead columns of `tab:readout-score-fidelity-summary` | `scripts/run/run_readout_baseline_comparisons.py` | one run per softcap-free readout, `--model` in Qwen3.5-0.8B, Qwen3.5-2B, Qwen3.5-9B, Ministral-3-8B-Base, R1-Distill-Qwen-7B, R1-Distill-Llama-8B: `--registry configs/registries/result1_query_fidelity_cluster.yaml --bank-dir data/query_banks --banks curated_ab,case_candidates,model_native --model-native-file qwen_gemma_result1_model_native_prompts_c4.jsonl --operating-point fidelity --methods sparse_rp,nearest_row_ridge_top128,knn_basis_top128,row_cluster_d16384_k256,row_cluster_d65536_k256,row_cluster_hard_d65536,pca_256 --max-native 500 --max-curated 320 --max-cases 40 --max-len 64 --seed 0 --out-root results/direct_geometry_runs/`; coverage = `accepted_rate` with `accepted_ci_lo/hi` in `baseline_by_method.csv` (unfloored rho_0, 400 base-case resamples). The C4 model-native slice is regenerated with `scripts/data/build_query_banks.py --native-source c4` (see §0) | +| `tab:app-direct-geometry-grid`; SRP / best-alternative / lead columns of `tab:readout-score-fidelity-summary` | `scripts/run/run_readout_baseline_comparisons.py` | one run per softcap-free readout, `--model` in Qwen3.5-0.8B, Qwen3.5-2B, Qwen3.5-9B, Ministral-3-8B-Base, R1-Distill-Qwen-7B, R1-Distill-Llama-8B: `--registry configs/registries/result1_query_fidelity_cluster.yaml --bank-dir data/query_banks --banks curated_ab,case_candidates,model_native --model-native-file qwen_gemma_result1_model_native_prompts_c4.jsonl --operating-point fidelity --methods sparse_rp,nearest_row_ridge_top128,knn_basis_top128,row_cluster_d16384_k256,row_cluster_d65536_k256,row_cluster_hard_d65536,pca_256 --max-native 500 --max-curated 320 --max-cases 40 --max-len 64 --seed 0 --out-root results/direct_geometry_runs/`; coverage = `accepted_rate` with `accepted_ci_lo/hi` in `baseline_by_method.csv` (unfloored rho_0, 400 base-case resamples). The C4 model-native slice is regenerated with `scripts/data/build_query_banks.py --native-source c4` (see §0) Since 0.2.1 the per-feature coverage/compactness columns for `nearest_row_ridge_top*`, `knn_basis_top*` and `row_cluster_*` are computed on aligned supports (the scalar coverage/sign metrics the tables report are unchanged). | | Result-1 five-model query-fidelity tables / §4.5 distributional readout metrics; per-model reconstruction with cluster-bootstrap intervals (`tab:app-fidelity-cis`) | `scripts/run/run_query_fidelity_bank.py` | `configs/registries/result1_query_fidelity_cluster.yaml` + the curated/model-native banks under `data/query_banks/`. `tab:app-fidelity-cis` reads, per model, `sign_agreement` (+`_ci_lo/_ci_hi`), `median_rho5` (+CI) and `pass_rho5_lt_0.50` (+CI) from the group metrics: `rho5` is the floored relative error \|m_exact − m_sparse\| / (\|m_exact\| + 0.5) and every interval comes from the same 400 base-case resamples | | Appendix J lexical-edit / readout-side control stress test (`tab:lexical-control-stress-test`, `tab:lexical-control-primary-methods`, `tab:lexical-control-cross-model-results`) | `scripts/run/run_qwen_profanity_suppression_eval.py` | `--checkpoint --model-id --out-dir --device cuda --dtype bfloat16 --batch-size 8 --prompt-limit 16 --pair-limit 9 --max-open-prompts 6 --num-samples-per-prompt 2 --max-new-tokens 24` (the original 16-prompt × 9-pair grid; `baseline_comparison.csv` / `candidate_constrained_summary.csv` at scale 16 give the primary-methods table). Cross-model rows: Qwen3.5-2B and DeepSeek-R1-Distill-Qwen-7B at 32x/k256, Ministral-3-8B-Base at 16x/k128 (`configs/registries/result1_query_fidelity_cluster.yaml`). Note: with the extended scale grid, `choose_operating_point` can select a scale above 16; read the scale-16 rows from `candidate_constrained_summary.csv` | | Appendix K sweep / selection tables (`tab:app-k-model-finalists`, `tab:app-model-suite-current`, and the other `tab:app-k-*` sweep ledgers) | — (no figure ships; numbers hand-transcribed from the audited literals in [`data/appendix/appendix_k_sweep_tables.json`](../data/appendix/appendix_k_sweep_tables.json)) | `configs/sweeps/arch_frontier_160m.yaml`, `configs/sweeps/paper_phase2_frontier.yaml` | @@ -161,30 +161,30 @@ check. | Appendix K: dense / negative-control diagnostic (`omp_diagnostic` in [`data/appendix/appendix_k_sweep_tables.json`](../data/appendix/appendix_k_sweep_tables.json)) | `scripts/eval/dense_control_diagnostic.py` | `--w-u .pt --checkpoint --setting "Qwen-9B 32x, k=256" --n-rows 20000 --bootstrap 1000` (encoder vs LS/NNLS-on-support vs signed/nonneg OMP vs dense rank-k, scored by rowEV) | | Appendix L: feature-label audit (`tab:app-main-case-study-feature-audit` data; `tab:app-qualitative-feature-label-audit` counts) | `scripts/analyze/audit_feature_labels.py` | `substrate --feature-ids 36,4095,… --model-id Qwen/Qwen3.5-2B` emits per-feature top rows; `aggregate --annotations data/audit/feature_label_audit_annotations.csv --validate-against data/audit/feature_label_audit.json` tallies the counts (classification is human; see note below) | | Matched-KL frontier metrics (`fig:lexical-matched-kl-frontier`) | `scripts/run/run_qwen_profanity_suppression_eval.py` | once per model (Qwen/Qwen3.5-2B, Qwen/Qwen3.5-0.8B, Qwen/Qwen3.5-9B, deepseek-ai/DeepSeek-R1-Distill-Qwen-7B; 32x/k256 checkpoints): `--checkpoint --model-id --out-dir --device cuda --dtype bfloat16 --batch-size 8 --max-open-prompts 0` (full 20-prompt × 17-pair × 14-scale × 8-method grid); the frontier plots the `split == heldout` rows of `candidate_constrained_summary.csv` (`median_kl_bits` vs `mean_bad_prob_reduction` / `candidate_flip_rate` per method and scale) | -| Paired differences at matched KL quoted with `fig:lexical-matched-kl-frontier` (SRP − mean-row / PCA rank-1 / PCA rank-4; term- and prompt-clustered 95% CIs) | `scripts/eval/paired_matched_kl_bootstrap.py` | `--input qwen2b=/candidate_constrained_rows.csv --input qwen0p8b=... --input qwen9b=... --input r1qwen7b=... --out paired_matched_kl_results.json --n-boot 10000 --seed 0` (default `--target-kl 0.02 0.05 0.1 0.2`); CPU, ~17 s for four models. Reproduces the paper's paired results exactly from the paper CSVs | +| Paired differences at matched KL quoted with `fig:lexical-matched-kl-frontier` (SRP − mean-row / PCA rank-1 / PCA rank-4; term- and prompt-clustered 95% CIs) | `scripts/eval/paired_matched_kl_bootstrap.py` | `--input qwen2b=/candidate_constrained_rows.csv --input qwen0p8b=... --input qwen9b=... --input r1qwen7b=... --out paired_matched_kl_results.json --n-boot 10000 --seed 0` (default `--target-kl 0.02 0.05 0.1 0.2`); CPU, ~17 s for four models. Output is a JSON object whose `comparisons` list holds the 96 paper records (order follows the `--input` order) beside a `provenance` block; reproduces the paper's paired results exactly from the paper CSVs | | `tab:app-cross-model-nulls` | `scripts/run/run_readout_baseline_comparisons.py` | one run per model Qwen3.5-0.8B/2B/9B, Gemma-4-E2B, Gemma-4-E4B with `--methods sparse_rp,shuffled_row_code,random_support_same_magnitudes` (all other flags as in the row above); coverage = `accepted_rate`, sign = `sign_agreement` in `baseline_by_method.csv` | | `fig:app-robustness-baseline-comparisons` (Qwen3.5-2B reference panel; the 848-cluster CI paragraph of the fidelity-results subsection) | `scripts/run/run_readout_baseline_comparisons.py` | `--model Qwen3.5-2B --methods sparse_rp,shuffled_row_code,random_support_same_magnitudes,pca_64,pca_256,pca_1024,nearest_row_ridge_top128` (other flags as above); `baseline_by_method.csv`, `baseline_by_method_margin_bin.csv`, `baseline_feature_compactness.csv` | | `tab:app-error-tails`, `tab:app-error-tail-margins` (and the per-model mean/p95/max ranges quoted in that subsection) | `scripts/eval/analyze_error_tails.py` | `--input-root results/direct_geometry_runs --input-prefix --model-tags qwen0p8b,qwen2b,qwen9b,ministral8b,r1qwen7b,r1llama8b --primary-model-tag qwen2b --out-dir results/error_tails --n-boot 10000 --seed 20260711` over the six direct-geometry run dirs; `summary_pooled_by_method.csv` (8,009 contrasts per method), the `sparse_rp` block of `summary_pooled_by_method_margin_bin.csv`, `summary_by_model_method.csv`. Verified to reproduce both tables from the paper's raw run CSVs | | `tab:app-nearest-rows` | `scripts/analyze/nearest_rows_baseline.py` | `--model-id Qwen/Qwen3.5-2B --revision 15852e8c16360a2fea060d615a32b45270f8a8fc --contrast bug,insect --contrast bug,error --top-n 12 --out-dir results/nearest_rows_qwen2b` (full-vocabulary centering mean, contrast tokens excluded); the `centered_*` columns of `nearest_rows_table.csv` are the table, `nearest_rows.json` holds the raw-cosine ranking too | -| Table 2 predicted-vs-realized readout-side changes, one row per model (`tab:causal-validation-summary`; per-model rows and protocol in `app:causal-validation`) | `scripts/eval/run_causal_contribution_validation.py` | once per model, fidelity operating point (32x, k=256): `--model-id --checkpoint /k256_32x/checkpoint.pt")> --bank-dir data/query_banks --out-dir results/causal_ --device cuda --dtype bfloat16 --seed 0`, with defaults `--banks curated_ab,case_candidates,model_native --max-native 300 --top-features 10 --random-per-case 10 --rho-gate 0.5 --n-boot 2000 --max-len 64`. Model → Hub dir: `Qwen/Qwen3.5-0.8B`→`qwen3.5-0.8b`, `Qwen/Qwen3.5-2B`→`qwen3.5-2b`, `Qwen/Qwen3.5-9B`→`qwen3.5-9b`, `mistralai/Ministral-3-8B-Base-2512`→`ministral-3-8b`, `deepseek-ai/DeepSeek-R1-Distill-Qwen-7B`→`r1-distill-qwen-7b`, `deepseek-ai/DeepSeek-R1-Distill-Llama-8B`→`r1-distill-llama-8b`. Table columns come from `summary.json`: r²/CI/slope/n ← `gated_predicted`, Random ← `gated_random_control.r2`; `ungated_predicted` is the all-pairs comparison. The `model_native` bank carries no A/B targets and contributes zero contrasts (a missing file is logged and skipped) | -| Appendix harness self-test line (`app:causal-validation`: r²=1.000, slope 1.04, zero-scoring controls) | `scripts/eval/run_causal_contribution_validation.py` | `--self-test` (CPU, no model/checkpoint; prints `r2=1.000 slope=1.044`, three PASS lines) | -| Appendix I: cross-seed stability of contrast explanations (`tab:app-cross-seed-stability`) | `scripts/eval/cross_seed_stability.py` | three same-recipe checkpoints trained from `configs/sweeps/qwen35_2b_seedvar_base.yaml`; `--dictionaries results/qwen35_2b_seedvar/qwen2b_d65536_k256_s0/checkpoint.pt …_s1/checkpoint.pt …_s2/checkpoint.pt --w-u "$SRP_ARCHIVE_ROOT/data/qwen35-2b/qwen35_2b.pt" --bank data/query_banks/qwen_gemma_result1_curated_ab.jsonl --width-tag 32x --out results/seed_stability/cross_seed_stability_32x.json` (defaults `--max-contrasts 150 --seed 0 --top-m 8 --top-r 12 --n-sample 4096 --n-hidden 512`; 16x column: the `qwen2b_d32768_k128_s{0,1,2}` checkpoints, `--width-tag 16x`) | -| Appendix I: cross-seed feature-group matching incl. the below-null direction check (`tab:app-feature-group-matching`) | `scripts/eval/feature_group_matching.py` | same dictionaries / `--w-u` / `--bank` as above; `--width-tag 32x --direction-check --out results/seed_stability/feature_group_matching_32x.json` (defaults `--n-query 100 --n-null 500 --greedy-pool 200 --greedy-k 3 --min-group-size 4 --direction-null 500 --direction-seed 0`; 16x: the k128 checkpoints, `--width-tag 16x`) | -| Appendix I: held-out recovery of the stable core (`tab:app-loo-core-recovery`, plus the 0.444 cross-recipe figure) | `scripts/eval/loo_core_recovery.py` | same dictionaries / `--w-u` / `--bank`; `--width-tag 32x --reference-dict --out results/seed_stability/loo_core_recovery_32x.json` (defaults `--side-n 96 --n-clusters 16384 --kmeans-iters 12 --kmeans-seed 0 --n-boot 2000 --bootstrap-seed 1`; 16x: k128 checkpoints, `--width-tag 16x`, no `--reference-dict`; k-means wants a GPU/MPS) | +| Table 2 predicted-vs-realized readout-side changes, one row per model (`tab:causal-validation-summary`; per-model rows and protocol in `app:causal-validation`) | `scripts/eval/run_causal_contribution_validation.py` | once per model, fidelity operating point (32x, k=256): `--model-id --checkpoint /k256_32x/checkpoint.pt")> --bank-dir data/query_banks --out-dir results/causal_ --device cuda --dtype bfloat16 --seed 0`, with defaults `--banks curated_ab,case_candidates,model_native --max-native 300 --top-features 10 --random-per-case 10 --rho-gate 0.5 --n-boot 2000 --max-len 64 --centering live` (`--centering trained` centres on the checkpoint's stored `row_mean`, else the tokenizer's text-token mean; the paper runs used the live full-vocabulary mean, which differs by ~0.04% of a centred row norm on Qwen3.5-2B). Model → Hub dir: `Qwen/Qwen3.5-0.8B`→`qwen3.5-0.8b`, `Qwen/Qwen3.5-2B`→`qwen3.5-2b`, `Qwen/Qwen3.5-9B`→`qwen3.5-9b`, `mistralai/Ministral-3-8B-Base-2512`→`ministral-3-8b`, `deepseek-ai/DeepSeek-R1-Distill-Qwen-7B`→`r1-distill-qwen-7b`, `deepseek-ai/DeepSeek-R1-Distill-Llama-8B`→`r1-distill-llama-8b`. Table columns come from `summary.json`: r²/CI/slope/n ← `gated_predicted`, Random ← `gated_random_control.r2`; `ungated_predicted` is the all-pairs comparison. The `model_native` bank carries no A/B targets and contributes zero contrasts (a missing file is logged and skipped) | +| Appendix harness self-test line (`app:causal-validation`: r²=1.000, slope 1.04, zero-scoring controls) | `scripts/eval/run_causal_contribution_validation.py` | `--self-test` (CPU, no model/checkpoint; prints `r2=1.000 slope=1.044` and three PASS lines, the third comparing each stored realized change against the dense LM-head margin change) | +| Appendix I: cross-seed stability of contrast explanations (`tab:app-cross-seed-stability`) | `scripts/eval/cross_seed_stability.py` | three same-recipe checkpoints trained from `configs/sweeps/qwen35_2b_seedvar_base.yaml`; `--dictionaries results/qwen35_2b_seedvar/qwen2b_d65536_k256_s0/checkpoint.pt …_s1/checkpoint.pt …_s2/checkpoint.pt --w-u "$SRP_ARCHIVE_ROOT/data/qwen35-2b/qwen35_2b.pt" --bank data/query_banks/qwen_gemma_result1_curated_ab.jsonl --width-tag 32x --out results/seed_stability/cross_seed_stability_32x.json` (defaults `--max-contrasts 150 --seed 0 --top-m 8 --top-r 12 --n-sample 4096 --n-hidden 512`; 16x column: the `qwen2b_d32768_k128_s{0,1,2}` checkpoints, `--width-tag 16x`) `--centering live` is the default (paper); `--centering trained` centres on the dictionaries' stored `row_mean`, else the payload's token-mask mean. Output JSON carries `provenance`. | +| Appendix I: cross-seed feature-group matching incl. the below-null direction check (`tab:app-feature-group-matching`) | `scripts/eval/feature_group_matching.py` | same dictionaries / `--w-u` / `--bank` as above; `--width-tag 32x --direction-check --out results/seed_stability/feature_group_matching_32x.json` (defaults `--n-query 100 --n-null 500 --greedy-pool 200 --greedy-k 3 --min-group-size 4 --direction-null 500 --direction-seed 0`; 16x: the k128 checkpoints, `--width-tag 16x`) `--centering` as above. Since 0.2.1 the frequency-matched null pool is sorted, so null draws no longer depend on `PYTHONHASHSEED`; they are a fresh reproducible sample rather than the paper run's specific draws (the paper's above-null counts were computed before this fix). | +| Appendix I: held-out recovery of the stable core (`tab:app-loo-core-recovery`, plus the 0.444 cross-recipe figure) | `scripts/eval/loo_core_recovery.py` | same dictionaries / `--w-u` / `--bank`; `--width-tag 32x --reference-dict --out results/seed_stability/loo_core_recovery_32x.json` (defaults `--side-n 96 --n-clusters 16384 --kmeans-iters 12 --kmeans-seed 0 --n-boot 2000 --bootstrap-seed 1`; 16x: k128 checkpoints, `--width-tag 16x`, no `--reference-dict`; k-means wants a GPU/MPS) `--centering` as above; `--knn-exclude-self` (default, paper) drops the query row from the kNN side set while the SRP and cluster sides keep it — `--no-knn-exclude-self` removes that asymmetry. | | Appendix K: sparsity budget vs dictionary usage (`tab:app-low-k-dead-features`) | `scripts/train/train_readout_sae_from_config.py` | multi-seed rows: `--config configs/sweeps/qwen35_2b_seedvar_base.yaml` with `--set factorizer.d_features= --set factorizer.k= --set evaluation.k= --set run.seed= --set run.init_seed= --set data.data_seed= --set run.output_dir=results/qwen35_2b_seedvar/qwen2b_d_k_s --set run.name=qwen2b_d_k_s` over cells (32768,256,{0,1,2}) (65536,128,{0,1,2}) (32768,64,{1,2}) (plus the Appendix I family (65536,256,{0,1,2}) (32768,128,{0,1,2})); single-seed rows: `--config configs/sweeps/lowk_qwen35_2b_16x_base.yaml` and `--config configs/sweeps/lowk_qwen35_0p8b_16x_base.yaml` with `--set factorizer.k=<32 or 64> --set evaluation.k=<32 or 64>` and `run.name`/`run.output_dir`; dead / rare / top-1 / KL read from each cell's `metrics.json` | -| Appendix K: Qwen3.5-0.8B 32x/k256 seed window (`tab:app-k-qwen08b-seed-window`) | `scripts/train/train_readout_sae_from_config.py` | `--config configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s1.yaml` and `--config configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s2.yaml`, no `--set`; seed 0 row is the archived `qwen0p8b_k256_32x` dictionary in `configs/registries/exp2_selected_sae_checkpoints.yaml` | -| Appendix "Sense-Labelled Evaluation": frozen CoarseWSD-20 readout bundles and coverage statistics (`app:sense-labelled-evaluation` data paragraph: contexts scored, single-token coverage, row-gate coverage, median target rank) | `scripts/run/run_wsd_feature_alignment.py` | `--dataset coarsewsd20 --data-root --model-id --checkpoint --out-dir results/wsd_core/coarsewsd20_ --splits train,test --batch-size 8 (2B) / 4 (9B, R1-Llama) --n-boot 5000 --seed 0`; writes `representations.pt` (input to the two rows below), `audit.json`, `scoring_summary.json`, `metrics.json`; `--analyze-only --out-dir --n-boot 5000 --seed 0` recomputes `metrics.json` without inference | -| Sense alignment on CoarseWSD-20 (`tab:app-sense-alignment`; group-size 1/2/4/8 sweep in the appendix prose) | `scripts/analyze/analyze_wsd_sense_groups.py` | `--bundle results/wsd_core/coarsewsd20_/representations.pt --out results/wsd_sense_groups/coarsewsd20___srp.json --group-sizes 1,2,4,8 --selector mean_diff --n-null 200 --n-boot 2000 --seed 0` (defaults: `--basis srp`, train-standardized); the table reads `_g8.json` fields `word_mean_balanced`, `word_mean_null_balanced`, `word_mean_majority_balanced`, `words_beating_null_p95_balanced`, `words_beating_majority_balanced`, `word_mean_full_account_balanced`, `word_mean_hidden_balanced`; CPU only. Verified against the paper run outputs cell for cell | -| Appendix "Sense-Labelled Evaluation", classifier-framing paragraph (projections vs signed contributions, row-coefficient shuffles, random features) | `scripts/analyze/analyze_wsd_classifier_framing.py` | `--bundle results/wsd_core/coarsewsd20_/representations.pt --out results/wsd_classifier_framing/coarsewsd20_.json --ks 5,10,20,50 --primary-k 10 --primary-encoding weighted --n-boot 5000 --n-null-seeds 20 --seed 0` (all defaults); the paragraph quotes `primary.methods..accuracy`; CPU only | -| Fitted lenses for the cross-lens study (inputs to all `fig:cross-lens-*` / `tab:app-cross-lens-*`) | `scripts/run/fit_jlens.py` | `prompts --prompts-json /c4_prompts_en_seed0.json --n-prompts 1000 --seed 0` (default `--min-chars 800 --c4-config en`); `prompts … c4_prompts_zh_seed0.json --n-prompts 1000 --seed 0 --min-chars 300 --c4-config zh`; same with `--c4-config de`; `smoke --model-id Qwen/Qwen3.5-0.8B --prompts-json --out /smoke_0p8b.lens.pt --dim-batch 16`; then per corpus `fit --model-id Qwen/Qwen3.5-9B --prompts-json --n-prompts 100 --shard $i --num-shards 4 --n-layers 12 --dim-batch 8 --ckpt-dir --out-dir ` for i in 0..3 and `merge --shard-dir --num-shards 4 --out /qwen35_9b_jlens_{en,zh,de}_seed0_n100.pt`; `--n-prompts 300` for the `*_n300.pt` refits. Needs `uv sync --extra lens` | -| Ridge translators (`tab:app-cross-lens-extension` rows 3–4, `fig:cross-lens-extension`) | `scripts/run/fit_ridge_lens.py` | `--model-id Qwen/Qwen3.5-9B --prompts-json /c4_prompts_en_seed0.json --n-prompts 100 --holdout 10 --layers-from /qwen35_9b_jlens_en_seed0_n100.pt --ckpt-dir --tag ridge_en --out /qwen35_9b_ridgelens_en_seed0_n100.pt`; repeat with `c4_prompts_zh_seed0.json --tag ridge_zh --out …ridgelens_zh_seed0_n100.pt` (layers still from the EN Jacobian lens). Holdout R² and deep top-1 agreement land in `.report.json` | -| Readout dumps behind every cross-lens artifact | `scripts/run/run_cross_lens_readouts.py` | once per (lens, bank): `--model-id Qwen/Qwen3.5-9B --lens --sae --k 128 --prompts data/cross_lens/cross_lens_prompts_en_zh.json --n-positions 1 --decompose-top1 --out /en_zh__.json` for the EN/ZH Jacobian n=100 and n=300 lenses and the EN/ZH ridge translators; `--prompts data/cross_lens/cross_lens_prompts_en_de.json` for the EN/DE Jacobian n=100 and n=300 lenses; three-lens example: `--prompts data/cross_lens/cross_lens_three_lens_prompts.json --n-positions 1 --decompose-top1 --top-feats 10` under the EN, ZH, DE n=100 lenses (optionally `--lens identity --layers-from `) | -| `tab:app-cross-lens-families`; `tab:app-cross-lens-extension` rows 1, 3, 4, 5; `fig:cross-lens-extension` (EN–ZH cells) | `scripts/eval/aggregate_cross_lens_en_zh.py` | `--prompts data/cross_lens/cross_lens_prompts_en_zh.json --seed 0 --out /_summary.json [--rows-csv …]` with `--lens-a/--lens-b` = (jlens_en_n100, jlens_zh_n100) main study; (ridge_en, ridge_zh) ridge row; (jlens_en_n100, ridge_en) Jacobian-vs-ridge row; (jlens_en_n300, jlens_zh_n300) n=300 row. Per-family agreement/CI and the "Lens-only split" column are `per_group` and `lens_only_split`; floors are `null_within_lens` (/159) and `null_shuffle_cross_lens` (/240). Verified to reproduce the paper run's summary key for key | -| `tab:app-cross-lens-extension` rows 2, 6; `fig:cross-lens-extension` (EN–DE cells); second-language-pair per-family counts, 73/90 and 76/90 divergence | `scripts/eval/aggregate_cross_lens_en_de.py` | `--lens-a /en_de__jlens_en_n100.json --lens-b /en_de__jlens_de_n100.json --prompts data/cross_lens/cross_lens_prompts_en_de.json --seed 0 --out /en_de_jlens_n100_summary.json`; same with the `_n300` dumps. Divergence is `divergence_rate`; floors /204 and /270 | -| EN–DE bank + cognate-exclusion report (inputs to the EN–DE cells) | `scripts/data/build_cross_lens_de_bank.py` | `--tokenizer Qwen/Qwen3.5-9B --out data/cross_lens/cross_lens_prompts_en_de.json --report data/cross_lens/cross_lens_prompts_en_de_exclusion_report.json` (shipped outputs are in `data/cross_lens/`) | -| `tab:app-cross-lens-antonym-layers` (and the f112 contributions quoted beside it) | `scripts/analyze/cross_lens_antonym_layers.py` | `run --model-id Qwen/Qwen3.5-9B --lens /qwen35_9b_jlens_en_seed0_n100.pt --sae --k 128 --out /antonym_layers_en.json` (default `--prompt '"小"的反义词是"'`); same with the ZH lens → `antonym_layers_zh.json`; then `table --dump EN=/antonym_layers_en.json --dump ZH=/antonym_layers_zh.json --layers 24,26,29,final --out-csv /antonym_layers_table.csv --out-features-csv /antonym_layers_features.csv` | -| `tab:app-cross-lens-de-antonym-layers` (and the groß/large/big feature ids and 15.16/15.60/22.19 logits in the prose) | `scripts/analyze/cross_lens_three_lens_prompt.py` | `--dump EN=/three_lens__jlens_en.json --dump ZH=/three_lens__jlens_zh.json --dump DE=/three_lens__jlens_de.json --prompt-id antonym_de_01 --layers 24,26,29 --out-csv /three_lens_top1.csv --out-targets-csv /three_lens_targets.csv --out-features-csv /three_lens_features.csv` (table = `top1_token`, `top1_lens_logit`; quoted groß logits = `original_logit` column) | -| `fig:cross-lens-butterfly` (metrics only) | `scripts/figures/compute_cross_lens_shared_feature.py` | `--dump EN=/en_zh__jlens_en_n100.json --dump ZH=/en_zh__jlens_zh_n100.json --prompt-id fac_03 --layer 26 --out-csv /cross_lens_shared_feature.csv --out-json /cross_lens_shared_feature.json` (top-1 logits 38.5 / 39.2, f23180 shares 0.752 / 0.629, largest other +1.69 / +1.55) | +| Appendix K: Qwen3.5-0.8B 32x/k256 seed window (`tab:app-k-qwen08b-seed-window`) | `scripts/train/train_readout_sae_from_config.py` | `--config configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s1.yaml` and `--config configs/sweeps/qwen35_0p8b_paper_topk_32x_k256_s2.yaml`, no `--set`; seed 0 row is the archived dictionary registered as `qwen0p8b_k256` in `configs/registries/exp2_selected_sae_checkpoints.yaml` (the same checkpoint is `qwen0p8b_k256_32x` in `result1_query_fidelity_cluster.yaml`) | +| Appendix "Sense-Labelled Evaluation": frozen CoarseWSD-20 readout bundles and coverage statistics (`app:sense-labelled-evaluation` data paragraph: contexts scored, single-token coverage, row-gate coverage, median target rank) | `scripts/run/run_wsd_feature_alignment.py` | `--dataset coarsewsd20 --data-root --model-id --checkpoint --out-dir results/wsd_core/coarsewsd20_ --splits train,test --batch-size 8 (2B) / 4 (9B, R1-Llama) --n-boot 5000 --seed 0`; writes `representations.pt` (input to the two rows below), `audit.json`, `scoring_summary.json`, `metrics.json`; `--analyze-only --out-dir --n-boot 5000 --seed 0` recomputes `metrics.json` without inference Flags since 0.2.1: `--centering {live,trained}` (default `live`, paper), `--k` (default: the checkpoint's k); `--analyze-only` writes `analysis_config.json` and leaves `run_config.json` untouched; `scoring_summary.json` records `truncation` (cloze prompts are truncated from the left, a 0.2.1 fix); `metrics.json` carries `provenance`. Only CoarseWSD-20 is supported. | +| Sense alignment on CoarseWSD-20 (`tab:app-sense-alignment`; group-size 1/2/4/8 sweep in the appendix prose) | `scripts/analyze/analyze_wsd_sense_groups.py` | `--bundle results/wsd_core/coarsewsd20_/representations.pt --out results/wsd_sense_groups/coarsewsd20___srp.json --group-sizes 1,2,4,8 --selector mean_diff --n-null 200 --n-boot 2000 --seed 0` (defaults: `--basis srp`, train-standardized); the table reads `_g8.json` fields `word_mean_balanced`, `word_mean_null_balanced`, `word_mean_majority_balanced`, `words_beating_null_p95_balanced`, `words_beating_majority_balanced`, `word_mean_full_account_balanced`, `word_mean_hidden_balanced`; CPU only. Verified against the paper run outputs cell for cell `--revision` (default: the cached `main` snapshot) and, for the `--basis` geometry controls, `--centering` (default `live`) with an optional `--checkpoint` supplying the stored `row_mean`; output carries `provenance` with sorted keys. Verified identical to the paper run outputs after the 0.2.1 refactor. | +| Appendix "Sense-Labelled Evaluation", classifier-framing paragraph (projections vs signed contributions, row-coefficient shuffles, random features) | `scripts/analyze/analyze_wsd_classifier_framing.py` | `--bundle results/wsd_core/coarsewsd20_/representations.pt --out results/wsd_classifier_framing/coarsewsd20_.json --ks 5,10,20,50 --primary-k 10 --primary-encoding weighted --n-boot 5000 --n-null-seeds 20 --seed 0` (all defaults); the paragraph quotes `primary.methods..accuracy`; CPU only Output carries `provenance` with sorted keys; verified identical to the paper run outputs after the 0.2.1 refactor. | +| Fitted lenses for the cross-lens study (inputs to all `fig:cross-lens-*` / `tab:app-cross-lens-*`) | `scripts/run/fit_jlens.py` | `prompts --prompts-json /c4_prompts_en_seed0.json --n-prompts 1000 --seed 0` (default `--min-chars 800 --c4-config en`); `prompts … c4_prompts_zh_seed0.json --n-prompts 1000 --seed 0 --min-chars 300 --c4-config zh`; same with `--c4-config de`; `smoke --model-id Qwen/Qwen3.5-0.8B --prompts-json --out /smoke_0p8b.lens.pt --dim-batch 16`; then per corpus `fit --model-id Qwen/Qwen3.5-9B --prompts-json --n-prompts 100 --shard $i --num-shards 4 --n-layers 12 --dim-batch 8 --ckpt-dir --out-dir ` for i in 0..3 and `merge --shard-dir --num-shards 4 --out /qwen35_9b_jlens_{en,zh,de}_seed0_n100.pt`; `--n-prompts 300` for the `*_n300.pt` refits. Needs `uv sync --extra lens` `fit`/`smoke` take `--device` (default `cuda`); `fit` writes `shard{i}.meta.json` beside each shard and refuses to resume a shard whose sidecar differs; `merge` requires the sidecars and writes `.meta.json`; `--dim-batch` halves on OOM down to 1. | +| Ridge translators (`tab:app-cross-lens-extension` rows 3–4, `fig:cross-lens-extension`) | `scripts/run/fit_ridge_lens.py` | `--model-id Qwen/Qwen3.5-9B --prompts-json /c4_prompts_en_seed0.json --n-prompts 100 --holdout 10 --layers-from /qwen35_9b_jlens_en_seed0_n100.pt --ckpt-dir --tag ridge_en --out /qwen35_9b_ridgelens_en_seed0_n100.pt`; repeat with `c4_prompts_zh_seed0.json --tag ridge_zh --out …ridgelens_zh_seed0_n100.pt` (layers still from the EN Jacobian lens). Holdout R² and deep top-1 agreement land in `.report.json` `--holdout` must be >= 1; `--device` (default `cuda`); accumulator resume is guarded by `/.meta.json`; the report carries `provenance`. The lambda grid is scaled per token but applied to the N-token Gram sum, so both paper fits selected the grid's top value at every layer — kept as the paper ran it. | +| Readout dumps behind every cross-lens artifact | `scripts/run/run_cross_lens_readouts.py` | once per (lens, bank): `--model-id Qwen/Qwen3.5-9B --lens --checkpoint --k 128 --prompts data/cross_lens/cross_lens_prompts_en_zh.json --n-positions 1 --decompose-top1 --out /en_zh__.json` for the EN/ZH Jacobian n=100 and n=300 lenses and the EN/ZH ridge translators; `--prompts data/cross_lens/cross_lens_prompts_en_de.json` for the EN/DE Jacobian n=100 and n=300 lenses; three-lens example: `--prompts data/cross_lens/cross_lens_three_lens_prompts.json --n-positions 1 --decompose-top1 --top-feats 10` under the EN, ZH, DE n=100 lenses (optionally `--lens identity --layers-from `) `--k` defaults to the checkpoint's k; `--centering live` (default, paper) or `trained`; `--device` (default `cuda`); `.manifest.json` records k, centering and provenance. | +| `tab:app-cross-lens-families`; `tab:app-cross-lens-extension` rows 1, 3, 4, 5; `fig:cross-lens-extension` (EN–ZH cells) | `scripts/eval/aggregate_cross_lens_en_zh.py` | `--prompts data/cross_lens/cross_lens_prompts_en_zh.json --seed 0 --out /_summary.json [--rows-csv …]` with `--lens-a/--lens-b` = (jlens_en_n100, jlens_zh_n100) main study; (ridge_en, ridge_zh) ridge row; (jlens_en_n100, ridge_en) Jacobian-vs-ridge row; (jlens_en_n300, jlens_zh_n300) n=300 row. Per-family agreement/CI and the "Lens-only split" column are `per_group` and `lens_only_split`; floors are `null_within_lens` (/159) and `null_shuffle_cross_lens` (/240). Verified to reproduce the paper run's summary key for key `--agreement-rule half` (default, paper: at least half of the mid-band layers; `strict` = more than half, which gives 72/80 on the main cell). Summary carries `provenance`. | +| `tab:app-cross-lens-extension` rows 2, 6; `fig:cross-lens-extension` (EN–DE cells); second-language-pair per-family counts, 73/90 and 76/90 divergence | `scripts/eval/aggregate_cross_lens_en_de.py` | `--lens-a /en_de__jlens_en_n100.json --lens-b /en_de__jlens_de_n100.json --prompts data/cross_lens/cross_lens_prompts_en_de.json --seed 0 --out /en_de_jlens_n100_summary.json`; same with the `_n300` dumps. Divergence is `divergence_rate`; floors /204 and /270 `--agreement-rule half` (default) and `--null-population all` (default, /204; `cross` drops the 12 controls, /180, the EN–ZH population). Summary carries `provenance`. | +| EN–DE bank + cognate-exclusion report (inputs to the EN–DE cells) | `scripts/data/build_cross_lens_de_bank.py` | `--tokenizer Qwen/Qwen3.5-9B --out data/cross_lens/cross_lens_prompts_en_de.json --report data/cross_lens/cross_lens_prompts_en_de_exclusion_report.json` (shipped outputs are in `data/cross_lens/`) The exclusion report carries `provenance`; the bank is byte-identical to the shipped file. | +| `tab:app-cross-lens-antonym-layers` (and the f112 contributions quoted beside it) | `scripts/analyze/cross_lens_antonym_layers.py` | `run --model-id Qwen/Qwen3.5-9B --lens /qwen35_9b_jlens_en_seed0_n100.pt --checkpoint --k 128 --out /antonym_layers_en.json` (default `--prompt '"小"的反义词是"'`); same with the ZH lens → `antonym_layers_zh.json`; then `table --dump EN=/antonym_layers_en.json --dump ZH=/antonym_layers_zh.json --layers 24,26,29,final --out-csv /antonym_layers_table.csv --out-features-csv /antonym_layers_features.csv` `run` also accepts `--centering` / `--device` and writes `.manifest.json`; `table` writes `.manifest.json`. | +| `tab:app-cross-lens-de-antonym-layers` (and the groß/large/big feature ids and 15.16/15.60/22.19 logits in the prose) | `scripts/analyze/cross_lens_three_lens_prompt.py` | `--dump EN=/three_lens__jlens_en.json --dump ZH=/three_lens__jlens_zh.json --dump DE=/three_lens__jlens_de.json --prompt-id antonym_de_01 --layers 24,26,29 --out-csv /three_lens_top1.csv --out-targets-csv /three_lens_targets.csv --out-features-csv /three_lens_features.csv` (table = `top1_token`, `top1_lens_logit`; quoted groß logits = `original_logit` column) Writes `.manifest.json`. | +| `fig:cross-lens-butterfly` (metrics only) | `scripts/figures/compute_cross_lens_shared_feature.py` | `--dump EN=/en_zh__jlens_en_n100.json --dump ZH=/en_zh__jlens_zh_n100.json --prompt-id fac_03 --layer 26 --out-csv /cross_lens_shared_feature.csv --out-json /cross_lens_shared_feature.json` (top-1 logits 38.5 / 39.2, f23180 shares 0.752 / 0.629, largest other +1.69 / +1.55) Writes `.manifest.json`, or `provenance` inside `--out-json` when given. | ### Notes and paper-only artifacts diff --git a/docs/THIRD_PARTY.md b/docs/THIRD_PARTY.md index 213dc5d..63bc65e 100644 --- a/docs/THIRD_PARTY.md +++ b/docs/THIRD_PARTY.md @@ -7,9 +7,11 @@ relevant upstream licenses and citation requirements. ## Code -No third-party source code is vendored in this repository. All Python -dependencies are standard packages installed from PyPI. The main runtime -dependencies (see `pyproject.toml`) are: +No third-party source code is vendored in this repository. The Python +dependencies are installed from PyPI, with one exception: the Jacobian-lens +reference implementation (`jlens`, below) is installed from its git repository +through an optional extra. The main runtime dependencies (see `pyproject.toml`) +are: - `torch` - `transformers` @@ -18,12 +20,25 @@ dependencies (see `pyproject.toml`) are: - `scipy` - `datasets` - `pyyaml` +- `pandas` (BSD-3-Clause; `scripts/eval/analyze_error_tails.py`) +- `scikit-learn` (BSD-3-Clause; `scripts/run/run_wsd_feature_alignment.py`, + `scripts/analyze/analyze_wsd_classifier_framing.py`) `torch` installs from the default PyPI index on every platform (CPU/MPS wheels on macOS, CUDA wheels on Linux). No alternative wheel index is configured; if your GPU needs a specific CUDA build, add a local `[tool.uv.sources]` pin and re-lock (see the comment in `pyproject.toml`). +### Jacobian lens (`jlens`) + +The cross-lens study (`scripts/run/fit_jlens.py` and the cross-lens scripts +that load its fitted lenses; they import it lazily) uses the reference +Jacobian-lens implementation from , +licensed under Apache-2.0. It is not published on PyPI: the optional `lens` +extra (`uv sync --extra lens`) installs it from git, and `uv.lock` pins it to +commit `581d398613e5602a5af361e1c34d3a92ea82ba8e`. The rest of the repository +installs and tests without it, and none of its code is vendored here. + ## Data ### C4 (`allenai/c4`) @@ -65,6 +80,13 @@ uv run python scripts/data/build_query_banks.py \ so no raw C4 text is redistributed in this repository. +The cross-lens study fits its lenses on seeded prompt dumps streamed from the +C4 `en`, `zh` and `de` configs (`scripts/run/fit_jlens.py prompts --c4-config +{en,zh,de}`; `train` split, streaming, `shuffle(seed=0)`, 1,000 prompts per +language). The dumps are raw C4 text and are not shipped; they carry the same +license and citation as above (ODC-BY 1.0, subject to the Common Crawl terms of +use). + ### WikiText-2 (`wikitext`, config `wikitext-2-raw-v1`) WikiText-2 is the hidden-state extraction corpus referenced by the model @@ -78,6 +100,27 @@ License: CC-BY-SA 3.0. Citation: Merity et al. (2016), "Pointer Sentinel Mixture Models" (WikiText). +### CoarseWSD-20 + +The sense-labelled evaluation (`scripts/run/run_wsd_feature_alignment.py +--dataset coarsewsd20 --data-root `, followed by +`scripts/analyze/analyze_wsd_sense_groups.py` and +`scripts/analyze/analyze_wsd_classifier_framing.py`) reads CoarseWSD-20 from a +local clone of +(`data/CoarseWSD-20`). The dataset is not redistributed here; the shipped +outputs are per-model feature bundles and aggregate metrics, not its text. + +License: the upstream repository declares no license (no `LICENSE` file, no +license statement in its README, and an empty license field on the GitHub +repository record at the time of this release), so no license can be quoted for +the dataset as given upstream. CoarseWSD-20 is built from English Wikipedia +sentences, whose text is under CC BY-SA. Check the upstream repository before +redistributing the data. + +Citation: Loureiro, Rezaee, Pilehvar and Camacho-Collados (2021), "Analysis and +Evaluation of Language Models for Word Sense Disambiguation", *Computational +Linguistics* 47(2), 387-443. + ### Benchmark-derived query suites `scripts/run/run_benchmark_derived_query_suite.py` builds readout-contrast diff --git a/pyproject.toml b/pyproject.toml index 5bf6fc3..55599b8 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "sparse-readout-prism" -version = "0.2.0" +version = "0.2.1" description = "Sparse Readout Prism: a sparse LM-head basis that factorizes language-model unembedding rows into reusable readout features and decomposes selected logits into signed feature contributions plus an explicit residual." readme = "README.md" license = { file = "LICENSE" } @@ -39,7 +39,7 @@ dependencies = [ "scipy>=1.15.3,<2", # scripts/eval/dense_control_diagnostic.py (scipy.optimize.nnls) "datasets>=4.8.4,<5", # scripts/data/extract_model_readout.py + build_query_banks.py "pyyaml>=6,<7", # config loading in utils.py + research/registry.py (was only transitive) - "pandas>=2.2,<3", # scripts/eval/analyze_error_tails.py + analyze_task_grounded_fidelity.py (CSV aggregation) + "pandas>=2.2,<3", # scripts/eval/analyze_error_tails.py (CSV aggregation) "scikit-learn>=1.5,<2", # scripts/run/run_wsd_feature_alignment.py + scripts/analyze/analyze_wsd_*.py ] diff --git a/scripts/README.md b/scripts/README.md index 64eccc3..2c0208a 100644 --- a/scripts/README.md +++ b/scripts/README.md @@ -38,6 +38,8 @@ The cross-lens scripts (`run/fit_jlens.py`, `run/fit_ridge_lens.py`, `run/run_cross_lens_readouts.py`, `analyze/cross_lens_antonym_layers.py run`) import the Jacobian-lens reference implementation lazily; install it with `uv sync --extra lens`. Every other script runs from the base environment. +The four GPU cross-lens scripts take `--device` (default `cuda`) and `--centering` +(default `live`, the paper's centering); see `docs/REPRODUCE.md`. The earlier exploratory families (causal / attention-head circuits, CoT-faithfulness and agentic-injection probes, the 160M transcoder/SAE lineage, diff --git a/scripts/analyze/analyze_wsd_classifier_framing.py b/scripts/analyze/analyze_wsd_classifier_framing.py index ac9392d..c750a54 100644 --- a/scripts/analyze/analyze_wsd_classifier_framing.py +++ b/scripts/analyze/analyze_wsd_classifier_framing.py @@ -23,9 +23,6 @@ the local score gate: sign preservation and |residual| / (|centered exact score| + 0.5) < 0.5. -AmbiStory bundles go through the same encoding path (story display against -the two gloss-anchor displays); the paper uses CoarseWSD-20 only. - Reads the ``representations.pt`` bundles written by ``scripts/run/run_wsd_feature_alignment.py``; CPU-only, deterministic under ``--seed``. Paper run, once per model bundle (Qwen3.5-2B, Qwen3.5-9B, @@ -36,13 +33,15 @@ --out results/wsd_classifier_framing/coarsewsd20_qwen2b.json \ --ks 5,10,20,50 --primary-k 10 --primary-encoding weighted \ --n-boot 5000 --n-null-seeds 20 --seed 0 + +Fixed in 0.2.1: the AmbiStory bundle path, which the paper does not use, was +removed; the bundle reader, shuffle null and bootstrap helpers moved to +``sparse_readout_prism.research.wsd``. """ from __future__ import annotations import argparse -import hashlib -import json from collections import defaultdict from pathlib import Path from typing import Any @@ -50,19 +49,22 @@ import numpy as np import torch +from sparse_readout_prism.research.run_io import run_provenance +from sparse_readout_prism.research.wsd import ( + bootstrap_mean, + l2_normalize, + load_bundle, + percentile_ci, + shuffled_srp, + stable_seed, + word_splits, +) +from sparse_readout_prism.utils import write_json + METHODS = ("srp", "projection_only", "shuffled_srp", "random_srp_features") ENCODINGS = ("signed", "weighted") - - -def stable_seed(text: str, seed: int) -> int: - digest = hashlib.sha1(text.encode("utf-8")).hexdigest()[:8] - return (int(digest, 16) + seed) % (2**32) - - -def write_json(path: Path, value: Any) -> None: - path.parent.mkdir(parents=True, exist_ok=True) - path.write_text(json.dumps(value, indent=2, ensure_ascii=False) + "\n") +COARSE_METRICS = ("accuracy", "balanced_accuracy", "macro_f1") def score_gate(bundle: dict[str, Any]) -> tuple[np.ndarray, dict[str, float]]: @@ -85,40 +87,16 @@ def score_gate(bundle: dict[str, Any]) -> tuple[np.ndarray, dict[str, float]]: } -def validate_feature_alignment(bundle: dict[str, Any]) -> None: - metadata = bundle["metadata"] - feature_ids = bundle["feature_ids"].numpy() - first_by_target: dict[str, np.ndarray] = {} - for row, ids in zip(metadata, feature_ids): - target = str(row["target"]) - if target in first_by_target: - if not np.array_equal(first_by_target[target], ids): - raise ValueError(f"feature coordinates change within target {target!r}") - else: - first_by_target[target] = ids.copy() - - def source_matrix(bundle: dict[str, Any], method: str, seed: int) -> np.ndarray: projection = bundle["projection"].float().numpy() contribution = bundle["contribution"].float().numpy() - beta = bundle["beta"].float().numpy() - metadata = bundle["metadata"] if method == "srp" or method == "random_srp_features": return contribution if method == "projection_only": return projection if method != "shuffled_srp": raise KeyError(method) - shuffled = np.empty_like(contribution) - by_target: dict[str, list[int]] = defaultdict(list) - for i, row in enumerate(metadata): - by_target[str(row["target"])].append(i) - for target, indices in by_target.items(): - rng = np.random.default_rng(stable_seed(target, seed)) - permutation = rng.permutation(beta.shape[1]) - idx = np.asarray(indices) - shuffled[idx] = projection[idx] * beta[idx][:, permutation] - return shuffled + return shuffled_srp(projection, bundle["beta"].float().numpy(), bundle["metadata"], seed) def encode_top_features( @@ -152,179 +130,7 @@ def encode_top_features( elif encoding != "weighted": raise KeyError(encoding) encoded[row_index, selected] = values - norms = np.linalg.norm(encoded, axis=1, keepdims=True) - return encoded / np.maximum(norms, 1e-12) - - -def safe_spearman(x: list[float], y: list[float]) -> float: - from scipy.stats import spearmanr - - if len(x) < 3 or np.ptp(x) < 1e-8 or np.ptp(y) < 1e-8: - return float("nan") - return float(spearmanr(x, y).statistic) - - -def cluster_bootstrap_delta( - rows: list[dict[str, Any]], - cluster_key: str, - statistic, - n_boot: int, - seed: int, -) -> list[float]: - grouped: dict[str, list[dict[str, Any]]] = defaultdict(list) - for row in rows: - grouped[str(row[cluster_key])].append(row) - keys = sorted(grouped) - rng = np.random.default_rng(seed) - values: list[float] = [] - for _ in range(n_boot): - sample: list[dict[str, Any]] = [] - for selected in rng.choice(keys, size=len(keys), replace=True): - sample.extend(grouped[str(selected)]) - value = float(statistic(sample)) - if np.isfinite(value): - values.append(value) - return values - - -def confidence_interval(values: list[float]) -> list[float]: - if not values: - return [float("nan"), float("nan")] - return [float(np.percentile(values, 2.5)), float(np.percentile(values, 97.5))] - - -def ambistory_rows(bundle: dict[str, Any], encoded: np.ndarray, gate: np.ndarray) -> list[dict[str, Any]]: - metadata = bundle["metadata"] - anchor_index = { - (str(row["target"]), str(row["anchor_key"])): i for i, row in enumerate(metadata) if row["kind"] == "anchor" - } - rows: list[dict[str, Any]] = [] - for i, row in enumerate(metadata): - if row["kind"] != "story" or not gate[i]: - continue - anchor_ids = [anchor_index.get((str(row["target"]), str(key))) for key in row["anchor_keys"]] - if any(index is None or not gate[int(index)] for index in anchor_ids): - continue - similarities = [float(encoded[i] @ encoded[int(index)]) for index in anchor_ids] - human_margin = float(row["human_margin"]) - predicted_margin = similarities[0] - similarities[1] - rows.append( - { - "item_id": str(row["item_id"]), - "setup_id": str(row["setup_id"]), - "target": str(row["target"]), - "human_margin": human_margin, - "predicted_margin": predicted_margin, - "eligible_gap1": int(abs(human_margin) >= 1.0), - "predicted_tie": int(abs(predicted_margin) < 1e-12), - "correct": int(human_margin != 0 and np.sign(human_margin) == np.sign(predicted_margin)), - "tie_half_score": ( - 0.5 - if abs(predicted_margin) < 1e-12 - else float(human_margin != 0 and np.sign(human_margin) == np.sign(predicted_margin)) - ), - } - ) - return rows - - -def summarize_ambistory_rows(rows: list[dict[str, Any]]) -> dict[str, Any]: - eligible = [row for row in rows if row["eligible_gap1"]] - eligible_non_ties = [row for row in eligible if not row["predicted_tie"]] - non_ties = [row for row in rows if row["human_margin"] != 0] - return { - "n_contexts": len(rows), - "n_setups": len({row["setup_id"] for row in rows}), - "n_targets": len({row["target"] for row in rows}), - "preference_spearman": safe_spearman( - [row["human_margin"] for row in rows], - [row["predicted_margin"] for row in rows], - ), - "preferred_sense_accuracy_gap1": ( - float(np.mean([row["tie_half_score"] for row in eligible])) if eligible else float("nan") - ), - "preferred_sense_accuracy_gap1_non_abstain": ( - float(np.mean([row["correct"] for row in eligible_non_ties])) if eligible_non_ties else float("nan") - ), - "preferred_sense_coverage_gap1": (len(eligible_non_ties) / len(eligible) if eligible else float("nan")), - "n_gap1": len(eligible), - "preferred_sense_accuracy_non_tie": ( - float(np.mean([row["correct"] for row in non_ties])) if non_ties else float("nan") - ), - } - - -def paired_ambistory_comparison( - srp_rows: list[dict[str, Any]], - baseline_rows: list[dict[str, Any]], - n_boot: int, - seed: int, -) -> dict[str, Any]: - baseline_by_id = {row["item_id"]: row for row in baseline_rows} - paired = [] - for srp in srp_rows: - baseline = baseline_by_id.get(srp["item_id"]) - if baseline is None: - continue - paired.append( - { - "setup_id": srp["setup_id"], - "human_margin": srp["human_margin"], - "eligible_gap1": srp["eligible_gap1"], - "srp_margin": srp["predicted_margin"], - "baseline_margin": baseline["predicted_margin"], - "srp_correct": srp["correct"], - "baseline_correct": baseline["correct"], - "srp_tie_half_score": srp["tie_half_score"], - "baseline_tie_half_score": baseline["tie_half_score"], - } - ) - - def rho_delta(sample): - human = [row["human_margin"] for row in sample] - return safe_spearman(human, [row["srp_margin"] for row in sample]) - safe_spearman( - human, [row["baseline_margin"] for row in sample] - ) - - eligible = [row for row in paired if row["eligible_gap1"]] - - def accuracy_delta(sample): - return float(np.mean([row["srp_tie_half_score"] for row in sample])) - float( - np.mean([row["baseline_tie_half_score"] for row in sample]) - ) - - return { - "n_contexts": len(paired), - "preference_spearman_difference": float(rho_delta(paired)), - "preference_spearman_difference_ci": confidence_interval( - cluster_bootstrap_delta(paired, "setup_id", rho_delta, n_boot, seed) - ), - "preferred_sense_accuracy_gap1_difference": float(accuracy_delta(eligible)), - "preferred_sense_accuracy_gap1_difference_ci": confidence_interval( - cluster_bootstrap_delta(eligible, "setup_id", accuracy_delta, n_boot, seed + 1) - ), - } - - -def bootstrap_ambistory_method(rows: list[dict[str, Any]], n_boot: int, seed: int) -> dict[str, Any]: - eligible = [row for row in rows if row["eligible_gap1"]] - - def rho(sample): - return safe_spearman( - [row["human_margin"] for row in sample], - [row["predicted_margin"] for row in sample], - ) - - def accuracy(sample): - return float(np.mean([row["tie_half_score"] for row in sample])) - - return { - **summarize_ambistory_rows(rows), - "preference_spearman_ci": confidence_interval(cluster_bootstrap_delta(rows, "setup_id", rho, n_boot, seed)), - "preferred_sense_accuracy_gap1_ci": confidence_interval( - cluster_bootstrap_delta(eligible, "setup_id", accuracy, n_boot, seed + 1) - ), - } + return l2_normalize(encoded) def summarize_null_values(values: list[float], observed: float) -> dict[str, Any]: @@ -341,123 +147,25 @@ def summarize_null_values(values: list[float], observed: float) -> dict[str, Any } -def ambistory_null_seed_sensitivity( - bundle: dict[str, Any], - gate: np.ndarray, - primary_k: int, - primary_encoding: str, - n_null_seeds: int, - seed: int, - observed: dict[str, Any], -) -> dict[str, Any]: - result: dict[str, Any] = {} - for method in ("shuffled_srp", "random_srp_features"): - rho_values: list[float] = [] - accuracy_values: list[float] = [] - for offset in range(n_null_seeds): - encoded = encode_top_features(bundle, method, primary_k, primary_encoding, seed + offset) - summary = summarize_ambistory_rows(ambistory_rows(bundle, encoded, gate)) - rho_values.append(float(summary["preference_spearman"])) - accuracy_values.append(float(summary["preferred_sense_accuracy_gap1"])) - result[method] = { - "preference_spearman": summarize_null_values(rho_values, float(observed["preference_spearman"])), - "preferred_sense_accuracy_gap1": summarize_null_values( - accuracy_values, float(observed["preferred_sense_accuracy_gap1"]) - ), - } - return result - - -def analyze_ambistory( - bundle: dict[str, Any], - ks: list[int], - primary_k: int, - primary_encoding: str, - n_boot: int, - n_null_seeds: int, - seed: int, -) -> dict[str, Any]: - gate, gate_summary = score_gate(bundle) - metrics: dict[str, Any] = { - "dataset": "ambistory", - "question": "Do top contributing features identify the human-preferred sense?", - "fidelity_gate": gate_summary, - "results": {}, - "primary": { - "encoding": primary_encoding, - "k": primary_k, - "methods": {}, - "comparisons": {}, - }, - } - primary_rows: dict[str, list[dict[str, Any]]] = {} - for encoding in ENCODINGS: - metrics["results"][encoding] = {} - for k in ks: - metrics["results"][encoding][str(k)] = {} - for method in METHODS: - encoded = encode_top_features(bundle, method, k, encoding, seed) - rows = ambistory_rows(bundle, encoded, gate) - metrics["results"][encoding][str(k)][method] = summarize_ambistory_rows(rows) - if encoding == primary_encoding and k == primary_k: - primary_rows[method] = rows - for method_index, method in enumerate(METHODS): - metrics["primary"]["methods"][method] = bootstrap_ambistory_method( - primary_rows[method], n_boot, seed + 50 + method_index * 2 - ) - for comparison_index, baseline in enumerate(METHODS[1:]): - metrics["primary"]["comparisons"][f"srp_minus_{baseline}"] = paired_ambistory_comparison( - primary_rows["srp"], - primary_rows[baseline], - n_boot, - seed + 100 + comparison_index * 10, - ) - metrics["primary"]["null_seed_sensitivity"] = ambistory_null_seed_sensitivity( - bundle, - gate, - primary_k, - primary_encoding, - n_null_seeds, - seed, - metrics["primary"]["methods"]["srp"], - ) - return metrics - - def coarse_per_word(bundle: dict[str, Any], encoded: np.ndarray, gate: np.ndarray) -> dict[str, dict[str, float]]: from sklearn.metrics import accuracy_score, balanced_accuracy_score, f1_score - metadata = bundle["metadata"] - words = sorted({str(row["target"]) for row in metadata}) per_word: dict[str, dict[str, float]] = {} - for word in words: - train_idx = [ - i for i, row in enumerate(metadata) if str(row["target"]) == word and row["split"] == "train" and gate[i] - ] - test_idx = [ - i for i, row in enumerate(metadata) if str(row["target"]) == word and row["split"] == "test" and gate[i] - ] - if not train_idx or not test_idx: - continue - train_y = np.asarray([str(metadata[i]["sense"]) for i in train_idx]) - test_y = np.asarray([str(metadata[i]["sense"]) for i in test_idx]) - senses = sorted(set(train_y)) - if len(senses) < 2 or not set(test_y).issubset(set(senses)): - continue + for split in word_splits(bundle["metadata"], keep=gate): centroids = [] - for sense in senses: - centroid = encoded[np.asarray(train_idx)[train_y == sense]].mean(axis=0) + for sense in split.senses: + centroid = encoded[split.train_idx[split.train_y == sense]].mean(axis=0) centroid /= max(np.linalg.norm(centroid), 1e-12) centroids.append(centroid) - scores = encoded[test_idx] @ np.stack(centroids).T - predictions = np.asarray([senses[index] for index in scores.argmax(axis=1)]) - per_word[word] = { - "n_train": len(train_idx), - "n_test": len(test_idx), - "n_senses": len(senses), - "accuracy": float(accuracy_score(test_y, predictions)), - "balanced_accuracy": float(balanced_accuracy_score(test_y, predictions)), - "macro_f1": float(f1_score(test_y, predictions, average="macro")), + scores = encoded[split.test_idx] @ np.stack(centroids).T + predictions = np.asarray([split.senses[index] for index in scores.argmax(axis=1)]) + per_word[split.word] = { + "n_train": len(split.train_idx), + "n_test": len(split.test_idx), + "n_senses": len(split.senses), + "accuracy": float(accuracy_score(split.test_y, predictions)), + "balanced_accuracy": float(balanced_accuracy_score(split.test_y, predictions)), + "macro_f1": float(f1_score(split.test_y, predictions, average="macro")), } return per_word @@ -467,10 +175,7 @@ def summarize_coarse(per_word: dict[str, dict[str, float]]) -> dict[str, Any]: "n_words": len(per_word), "n_train": int(sum(row["n_train"] for row in per_word.values())), "n_test": int(sum(row["n_test"] for row in per_word.values())), - **{ - metric: float(np.mean([row[metric] for row in per_word.values()])) - for metric in ("accuracy", "balanced_accuracy", "macro_f1") - }, + **{metric: float(np.mean([row[metric] for row in per_word.values()])) for metric in COARSE_METRICS}, } @@ -483,11 +188,10 @@ def paired_word_comparison( words = sorted(set(srp) & set(baseline)) rng = np.random.default_rng(seed) result: dict[str, Any] = {"n_words": len(words)} - for metric in ("accuracy", "balanced_accuracy", "macro_f1"): + for metric in COARSE_METRICS: differences = np.asarray([srp[word][metric] - baseline[word][metric] for word in words], dtype=float) - boot = [float(rng.choice(differences, size=len(differences), replace=True).mean()) for _ in range(n_boot)] result[f"{metric}_difference"] = float(differences.mean()) - result[f"{metric}_difference_ci"] = confidence_interval(boot) + result[f"{metric}_difference_ci"] = percentile_ci(bootstrap_mean(differences, rng, n_boot)) result[f"{metric}_positive_words"] = int((differences > 0).sum()) return result @@ -496,8 +200,7 @@ def bootstrap_coarse_method(per_word: dict[str, dict[str, float]], n_boot: int, result = summarize_coarse(per_word) values = np.asarray([row["macro_f1"] for row in per_word.values()], dtype=float) rng = np.random.default_rng(seed) - boot = [float(rng.choice(values, size=len(values), replace=True).mean()) for _ in range(n_boot)] - result["macro_f1_ci"] = confidence_interval(boot) + result["macro_f1_ci"] = percentile_ci(bootstrap_mean(values, rng, n_boot)) return result @@ -595,12 +298,11 @@ def self_test() -> int: assert np.isclose(np.linalg.norm(signed[0]), 1.0) weighted = encode_top_features(bundle, "srp", 2, "weighted", 0) assert abs(weighted[0, 1]) > abs(weighted[0, 2]) > 0 - validate_feature_alignment(bundle) print("self-test passed") return 0 -def parse_args() -> argparse.Namespace: +def parse_args(argv: list[str] | None = None) -> argparse.Namespace: parser = argparse.ArgumentParser() parser.add_argument("--bundle", type=Path, help="representations.pt from run_wsd_feature_alignment.py") parser.add_argument("--out", type=Path, help="metrics JSON") @@ -611,11 +313,11 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--n-null-seeds", type=int, default=20) parser.add_argument("--seed", type=int, default=0) parser.add_argument("--self-test", action="store_true") - return parser.parse_args() + return parser.parse_args(argv) -def main() -> int: - args = parse_args() +def main(argv: list[str] | None = None) -> int: + args = parse_args(argv) if args.self_test: return self_test() if args.bundle is None or args.out is None: @@ -623,31 +325,17 @@ def main() -> int: ks = [int(value.strip()) for value in args.ks.split(",") if value.strip()] if args.primary_k not in ks: raise SystemExit("--primary-k must be included in --ks") - bundle = torch.load(args.bundle, map_location="cpu", weights_only=False) - validate_feature_alignment(bundle) - dataset = str(bundle["run"]["dataset"]) - if dataset == "ambistory": - metrics = analyze_ambistory( - bundle, - ks, - args.primary_k, - args.primary_encoding, - args.n_boot, - args.n_null_seeds, - args.seed, - ) - elif dataset == "coarsewsd20": - metrics = analyze_coarse( - bundle, - ks, - args.primary_k, - args.primary_encoding, - args.n_boot, - args.n_null_seeds, - args.seed, - ) - else: - raise ValueError(f"unsupported dataset {dataset!r}") + provenance = run_provenance(args) + bundle = load_bundle(args.bundle) + metrics = analyze_coarse( + bundle, + ks, + args.primary_k, + args.primary_encoding, + args.n_boot, + args.n_null_seeds, + args.seed, + ) metrics["run"] = bundle["run"] metrics["analysis"] = { "bundle": str(args.bundle), @@ -660,7 +348,8 @@ def main() -> int: "methods": list(METHODS), "encodings": list(ENCODINGS), } - write_json(args.out, metrics) + metrics["provenance"] = provenance + write_json(metrics, args.out) print(f"wrote {args.out}") return 0 diff --git a/scripts/analyze/analyze_wsd_sense_groups.py b/scripts/analyze/analyze_wsd_sense_groups.py index da23919..1442db6 100644 --- a/scripts/analyze/analyze_wsd_sense_groups.py +++ b/scripts/analyze/analyze_wsd_sense_groups.py @@ -36,7 +36,11 @@ model forward pass. These need the unembedding matrix, read from ``--w-u`` (an extraction payload with ``W_U_orig``, as written by ``scripts/data/extract_model_readout.py``) or, failing that, from the local -Hugging Face cache for ``--model-id``. The paper's table reports ``srp`` only. +Hugging Face cache of ``--model-id`` at ``--revision`` (no download). Their +rows are centred under ``--centering`` (``live``, the default: the +full-vocabulary mean of that matrix; ``trained``: the ``--checkpoint``'s stored +``row_mean``, else the mean over the payload's ``token_mask`` or the +tokenizer's text-token rows). The paper's table reports ``srp`` only. Reads the ``representations.pt`` bundles written by ``scripts/run/run_wsd_feature_alignment.py``; CPU-only, deterministic under @@ -50,57 +54,63 @@ The table reads the ``_g8`` file; the group-size sweep in the appendix prose reads all four. + +Fixed in 0.2.1: the cached Hugging Face weights are resolved through +``snapshot_download(local_files_only=True)`` at ``--revision`` (default: the +cached ``main`` ref) instead of the lexicographically last snapshot directory; +the direct-geometry centering mean goes through ``data.centering_mean`` under +``--centering``; the bundle reader and split rule moved to +``sparse_readout_prism.research.wsd``. """ from __future__ import annotations import argparse -import json -import os from collections import Counter from pathlib import Path +from typing import Any import numpy as np import torch +from sparse_readout_prism.data import centering_mean, token_mask_from_tokenizer +from sparse_readout_prism.research.qwen_readout import load_sae +from sparse_readout_prism.research.run_io import run_provenance +from sparse_readout_prism.research.wsd import load_bundle, percentile_ci, word_splits +from sparse_readout_prism.utils import write_json + -def load_wu_from_artifact(path: Path) -> torch.Tensor: - """Load the unembedding matrix (V, d_model) fp32 from an extraction payload (``W_U_orig`` / ``W_U``).""" +def load_wu_from_artifact(path: Path) -> tuple[torch.Tensor, torch.Tensor | None]: + """Unembedding matrix (V, d_model) fp32 and the ``token_mask`` (if stored) from an extraction payload.""" payload = torch.load(path, map_location="cpu", weights_only=True) w = payload.get("W_U_orig", payload.get("W_U")) if w is None: raise KeyError(f"{path}: no 'W_U_orig'/'W_U' key (keys: {list(payload)})") - return w.float() + token_mask = payload.get("token_mask") + return w.float(), (token_mask.bool() if token_mask is not None else None) -def load_wu(model_id: str) -> torch.Tensor: - """Load the LM-head weight (V, d_model) fp32 from the local Hugging Face cache. +def load_wu(model_id: str, revision: str | None = None) -> torch.Tensor: + """LM-head weight (V, d_model) fp32 from the local Hugging Face cache, no download. - Reads only the tensor from safetensors; falls back to the tied embedding - when no lm_head.weight exists (e.g. Qwen3.5-2B, tie_word_embeddings). The - cache root follows ``HF_HUB_CACHE``, then ``HF_HOME/hub``, then - ``~/.cache/huggingface/hub``. + ``huggingface_hub.snapshot_download(..., local_files_only=True)`` resolves + ``revision`` (default: the cached ``main`` ref) to the snapshot directory; + only the tensor is read from the safetensors shards, falling back to the + tied embedding when no ``lm_head.weight`` exists (e.g. Qwen3.5-2B, + ``tie_word_embeddings``). """ - from glob import glob - + from huggingface_hub import snapshot_download from safetensors import safe_open - hub = os.environ.get("HF_HUB_CACHE") - if not hub: - hub = str(Path(os.environ.get("HF_HOME") or (Path.home() / ".cache" / "huggingface")) / "hub") - cache = Path(hub) / f"models--{model_id.replace('/', '--')}" - snaps = sorted(cache.glob("snapshots/*")) - if not snaps: - raise FileNotFoundError(f"no local snapshot for {model_id}") - snap = snaps[-1] - shards = sorted(glob(str(snap / "*.safetensors"))) + snap = Path(snapshot_download(model_id, revision=revision, local_files_only=True)) + shards = sorted(snap.glob("*.safetensors")) for suffix in ("lm_head.weight", "embed_tokens.weight"): for shard in shards: - with safe_open(shard, framework="pt") as f: + with safe_open(str(shard), framework="pt") as f: for key in f.keys(): if key.endswith(suffix): w = f.get_tensor(key).float() - print(f"[wu] {model_id}: {key} {tuple(w.shape)} from {Path(shard).name}") + print(f"[wu] {model_id}: {key} {tuple(w.shape)} from {shard.name}") return w raise KeyError(f"no lm_head/embed_tokens weight in {snap}") @@ -108,6 +118,7 @@ def load_wu(model_id: str) -> torch.Tensor: def geometry_contributions( basis: str, W: torch.Tensor, # (V, d_model) fp32 + row_mean: torch.Tensor, # (d_model,) centering mean, from data.centering_mean token_ids: list[int], hidden_by_word: dict[int, torch.Tensor], # token_id -> (n_word, d_model) fp32 pca_rank: int, @@ -117,12 +128,18 @@ def geometry_contributions( ) -> dict[int, dict[str, np.ndarray]]: """Per-word contribution matrices for a direct-geometry basis. - Math is line-for-line the vectorized form of readout_baselines.py: - rows are mean-centered and unit-normalized, the target row is rebuilt in - the basis, and contribution_j(h) = row_norm * coeff_j * (h . basis_j). - Returns {token_id: {contribution (n, F), exact (n,), recon (n,)}}. + Vectorized form of the baseline runner's methods in + ``scripts/run/run_readout_baseline_comparisons.py`` (script classes, so + they cannot be imported here): ``pca256`` mirrors ``DensePCAMethod`` + (``svd_lowrank`` on the centred, unit-normalised rows, q = rank + 16, + niter = 4), ``ridge128`` mirrors ``NearestRowRidgeMethod`` (signed ridge fit + on the top-k cosine neighbours, self excluded) and ``knn128`` mirrors + ``KNNBasisMethod`` (the cosine-weighted neighbourhood mean rescaled by the + target's projection onto it). Rows are centred on ``row_mean`` and + unit-normalised, the target row is rebuilt in the basis, and + contribution_j(h) = row_norm * coeff_j * (h . basis_j). Returns + {token_id: {contribution (n, F), exact (n,), recon (n,)}}. """ - row_mean = W.mean(dim=0) # (d_model,) centered = W - row_mean row_norm_all = centered.norm(dim=1).clamp_min(1e-8) row_normalized_all = centered / row_norm_all[:, None] @@ -215,7 +232,7 @@ def select_anchors( def word_accuracy( - contribution: np.ndarray, anchors: dict[str, int], senses: list[str], gold: np.ndarray + contribution: np.ndarray, anchors: dict[str, list[int]], senses: list[str], gold: np.ndarray ) -> tuple[float, np.ndarray]: anchor_matrix = np.stack([contribution[:, anchors[s]].sum(axis=1) for s in senses], axis=1) # (n, n_senses) predictions = np.array([senses[i] for i in anchor_matrix.argmax(axis=1)]) @@ -228,7 +245,51 @@ def balanced_accuracy(predictions: np.ndarray, gold: np.ndarray) -> float: return float(np.mean(recalls)) -def main() -> int: +def ncm_balanced(train_x: np.ndarray, test_x: np.ndarray, train_y: np.ndarray, test_y: np.ndarray) -> float: + """Nearest class-centroid on train-standardized features: the + no-selection reference (how much sense information the representation + carries when the whole account is used).""" + mu = train_x.mean(axis=0) + sd = train_x.std(axis=0) + 1e-6 + ztr = (train_x - mu) / sd + zte = (test_x - mu) / sd + senses = sorted(set(train_y)) + cents = np.stack([ztr[train_y == s].mean(axis=0) for s in senses]) + dist = ((zte[:, None, :] - cents[None]) ** 2).sum(axis=-1) # (n_test, n_senses) + pred = np.array([senses[i] for i in dist.argmin(axis=1)]) + return balanced_accuracy(pred, test_y) + + +def geometry_row_mean( + W: torch.Tensor, + mode: str, + *, + token_mask: torch.Tensor | None, + checkpoint: Path | None, + model_id: str, + revision: str | None, +) -> torch.Tensor: + """``--centering`` for the direct-geometry controls, via ``data.centering_mean``. + + ``live``: the full-vocabulary mean of ``W``. ``trained``: the checkpoint's + stored ``row_mean`` when ``--checkpoint`` is given, else the mean over + ``token_mask`` (the ``--w-u`` payload's, else built from ``model_id``'s + tokenizer in the local cache). + """ + ckpt = None + if mode == "trained": + if checkpoint is not None: + *_rest, ckpt_row_mean = load_sae(checkpoint) + ckpt = {"row_mean": ckpt_row_mean} if ckpt_row_mean is not None else None + if token_mask is None: + from transformers import AutoTokenizer + + tok = AutoTokenizer.from_pretrained(model_id, revision=revision, local_files_only=True) + token_mask = token_mask_from_tokenizer(tok, W.shape[0]) + return centering_mean(W, mode=mode, token_mask=token_mask, ckpt=ckpt) + + +def parse_args(argv: list[str] | None = None) -> argparse.Namespace: parser = argparse.ArgumentParser(description=__doc__.split("\n")[0]) parser.add_argument( "--bundle", type=Path, required=True, help="representations.pt from run_wsd_feature_alignment.py" @@ -254,12 +315,23 @@ def main() -> int: help="srp: bundle contributions; others: direct-W_U geometry controls", ) parser.add_argument("--model-id", default=None, help="override bundle run.model_id for --basis controls") + parser.add_argument("--revision", default=None, help="HF revision of --model-id in the local cache (default: main)") parser.add_argument( "--w-u", type=Path, default=None, help="extraction payload (.pt with W_U_orig) for --basis controls; default: local HF cache of --model-id", ) + parser.add_argument( + "--centering", + choices=("live", "trained"), + default="live", + help="row centering for --basis controls: live = full-vocabulary mean of W_U (the paper's runs); " + "trained = --checkpoint's stored row_mean, else the token_mask mean", + ) + parser.add_argument( + "--checkpoint", type=Path, default=None, help="dictionary checkpoint whose row_mean --centering trained uses" + ) parser.add_argument("--pca-rank", type=int, default=256) parser.add_argument("--neighbor-k", type=int, default=128) parser.add_argument("--ridge-lam", type=float, default=1e-3) @@ -268,19 +340,34 @@ def main() -> int: action="store_true", help="skip train-statistic standardization of contributions (default: standardize)", ) - args = parser.parse_args() + return parser.parse_args(argv) + + +def main(argv: list[str] | None = None) -> int: + args = parse_args(argv) + provenance = run_provenance(args) - bundle = torch.load(args.bundle, map_location="cpu", weights_only=False) + bundle = load_bundle(args.bundle) metadata = bundle["metadata"] contribution = bundle["contribution"].float().numpy() # (n_items, k_support) exact = bundle["exact_logit"].numpy() recon = bundle["reconstructed_logit"].numpy() if args.basis != "srp": + model_id = args.model_id or bundle["run"]["model_id"] + token_mask = None if args.w_u is not None: - wu = load_wu_from_artifact(args.w_u) + wu, token_mask = load_wu_from_artifact(args.w_u) else: - wu = load_wu(args.model_id or bundle["run"]["model_id"]) + wu = load_wu(model_id, args.revision) + row_mean = geometry_row_mean( + wu, + args.centering, + token_mask=token_mask, + checkpoint=args.checkpoint, + model_id=model_id, + revision=args.revision, + ) hidden = bundle["hidden"].float() # (n_items, d_model) word_tid: dict[str, int] = {} word_idx: dict[str, list[int]] = {} @@ -291,6 +378,7 @@ def main() -> int: geo = geometry_contributions( args.basis, wu, + row_mean, [word_tid[w] for w in sorted(word_idx)], hidden_by_word, args.pca_rank, @@ -317,34 +405,11 @@ def main() -> int: group_sizes = [int(x) for x in args.group_sizes.split(",")] if args.group_sizes else [args.group_size] - def ncm_balanced(train_x: np.ndarray, test_x: np.ndarray, train_y: np.ndarray, test_y: np.ndarray) -> float: - """Nearest class-centroid on train-standardized features: the - no-selection reference (how much sense information the representation - carries when the whole account is used).""" - mu = train_x.mean(axis=0) - sd = train_x.std(axis=0) + 1e-6 - ztr = (train_x - mu) / sd - zte = (test_x - mu) / sd - senses = sorted(set(train_y)) - cents = np.stack([ztr[train_y == s].mean(axis=0) for s in senses]) - dist = ((zte[:, None, :] - cents[None]) ** 2).sum(axis=-1) # (n_test, n_senses) - pred = np.array([senses[i] for i in dist.argmin(axis=1)]) - return balanced_accuracy(pred, test_y) - hidden_all = bundle["hidden"].float().numpy() # (n_items, d_model) - words = sorted({row["target"] for row in metadata}) word_cache: dict[str, dict] = {} - for word in words: - train_idx = np.array([i for i, r in enumerate(metadata) if r["target"] == word and r["split"] == "train"]) - test_idx = np.array([i for i, r in enumerate(metadata) if r["target"] == word and r["split"] == "test"]) - if len(train_idx) == 0 or len(test_idx) == 0: - continue - train_y = np.array([metadata[i]["sense"] for i in train_idx]) - test_y = np.array([metadata[i]["sense"] for i in test_idx]) - senses = sorted(set(train_y)) - if len(senses) < 2 or not set(test_y).issubset(set(senses)): - continue - + for split in word_splits(metadata): + train_idx, test_idx = split.train_idx, split.test_idx + train_y, test_y = split.train_y, split.test_y train_c = contribution[train_idx] test_c = contribution[test_idx] if not args.raw_scale: @@ -355,12 +420,12 @@ def ncm_balanced(train_x: np.ndarray, test_x: np.ndarray, train_y: np.ndarray, t sd = train_c.std(axis=0) + 1e-6 train_c = (train_c - mu) / sd test_c = (test_c - mu) / sd - word_cache[word] = { + word_cache[split.word] = { "train_idx": train_idx, "test_idx": test_idx, "train_y": train_y, "test_y": test_y, - "senses": senses, + "senses": split.senses, "train_c": train_c, "test_c": test_c, # basis-independent references, computed once per word @@ -369,11 +434,11 @@ def ncm_balanced(train_x: np.ndarray, test_x: np.ndarray, train_y: np.ndarray, t } for group_size in group_sizes: - run_group(args, group_size, word_cache, gate) + run_group(args, group_size, word_cache, gate, provenance) return 0 -def run_group(args, group_size: int, word_cache: dict, gate: np.ndarray) -> None: +def run_group(args, group_size: int, word_cache: dict, gate: np.ndarray, provenance: dict[str, Any]) -> None: rng = np.random.default_rng(args.seed) per_word: dict[str, dict] = {} pooled_correct: list[int] = [] @@ -402,8 +467,8 @@ def run_group(args, group_size: int, word_cache: dict, gate: np.ndarray) -> None boot[b] = correct[sample].mean() strat = np.concatenate([pool[rng.integers(0, len(pool), len(pool))] for pool in sense_pools.values()]) boot_bal[b] = balanced_accuracy(predictions[strat], test_y[strat]) - ci = [float(np.percentile(boot, 2.5)), float(np.percentile(boot, 97.5))] - bal_ci = [float(np.percentile(boot_bal, 2.5)), float(np.percentile(boot_bal, 97.5))] + ci = percentile_ci(boot) + bal_ci = percentile_ci(boot_bal) # label-shuffle null: re-select anchors under permuted train labels null_accs = np.empty(args.n_null) @@ -468,16 +533,10 @@ def run_group(args, group_size: int, word_cache: dict, gate: np.ndarray) -> None "seed": args.seed, "n_words": len(names), "word_mean_accuracy": float(word_acc.mean()), - "word_mean_accuracy_ci": [ - float(np.percentile(boot_mean, 2.5)), - float(np.percentile(boot_mean, 97.5)), - ], + "word_mean_accuracy_ci": percentile_ci(boot_mean), "standardized": not args.raw_scale, "word_mean_balanced": float(word_bal.mean()), - "word_mean_balanced_ci": [ - float(np.percentile(boot_mean_bal, 2.5)), - float(np.percentile(boot_mean_bal, 97.5)), - ], + "word_mean_balanced_ci": percentile_ci(boot_mean_bal), "word_mean_majority": float(word_majority.mean()), "word_mean_majority_balanced": float(word_majority_bal.mean()), "word_mean_chance_balanced": float(word_chance_bal.mean()), @@ -494,13 +553,11 @@ def run_group(args, group_size: int, word_cache: dict, gate: np.ndarray) -> None "pooled_gated_accuracy": float(np.mean(pooled_correct_gated)), "gate_fraction_overall": float(len(pooled_correct_gated) / max(len(pooled_correct), 1)), "per_word": per_word, + "provenance": provenance, } out = args.out if args.group_sizes is None else args.out.with_name(f"{args.out.stem}_g{group_size}.json") - out.parent.mkdir(parents=True, exist_ok=True) - tmp = out.with_suffix(".tmp") - tmp.write_text(json.dumps(summary, indent=2)) - os.replace(tmp, out) + write_json(summary, out, atomic=True) print( f"{args.bundle.parent.name} [{args.basis} g={group_size}]: words={summary['n_words']} " f"bal acc={summary['word_mean_balanced']:.3f} " diff --git a/scripts/analyze/cross_lens_antonym_layers.py b/scripts/analyze/cross_lens_antonym_layers.py index 76eb705..5ae1e40 100644 --- a/scripts/analyze/cross_lens_antonym_layers.py +++ b/scripts/analyze/cross_lens_antonym_layers.py @@ -11,18 +11,25 @@ the layer x lens table of top-1 token and softmax share (and the dominant feature of each probed form under each lens) from the per-lens JSON files. -Decomposition basis: the k=128 seed-0 readout dictionary for Qwen3.5-9B (8x, -D=32768), Hugging Face ``hematteo/sparse-readout-prism`` file -``qwen3.5-9b/k128_8x/checkpoint.pt``. ``jlens`` is imported lazily -(``uv sync --extra lens``). +Decomposition basis: ``--checkpoint``, the k=128 seed-0 readout dictionary for +Qwen3.5-9B (8x, D=32768), Hugging Face ``hematteo/sparse-readout-prism`` file +``qwen3.5-9b/k128_8x/checkpoint.pt``, read with ``research.qwen_readout.load_sae``; +``--k`` defaults to the checkpoint's trained k (a different explicit value +warns) and ``--centering`` selects the centering mean (``live`` = full-vocabulary +mean of the live head, the paper; ``trained`` = the checkpoint's ``row_mean`` or +the tokenizer's text-token mean). ``jlens`` is imported lazily +(``uv sync --extra lens``); the model is loaded through +``research.cross_lens.load_lens_model`` on ``--device`` (default ``cuda``). +``run`` writes ``.manifest.json`` (resolved k, centering, provenance) +next to its JSON; ``table`` writes ``.manifest.json`` next to its CSVs. Paper runs:: cross_lens_antonym_layers.py run --model-id Qwen/Qwen3.5-9B \ - --lens /qwen35_9b_jlens_en_seed0_n100.pt --sae --k 128 \ + --lens /qwen35_9b_jlens_en_seed0_n100.pt --checkpoint --k 128 \ --out /antonym_layers_en.json cross_lens_antonym_layers.py run --model-id Qwen/Qwen3.5-9B \ - --lens /qwen35_9b_jlens_zh_seed0_n100.pt --sae --k 128 \ + --lens /qwen35_9b_jlens_zh_seed0_n100.pt --checkpoint --k 128 \ --out /antonym_layers_zh.json cross_lens_antonym_layers.py table --dump EN=/antonym_layers_en.json \ --dump ZH=/antonym_layers_zh.json --layers 24,26,29,final \ @@ -34,65 +41,56 @@ import argparse import json import os +from functools import partial from pathlib import Path import torch -import torch.nn.functional as F -from sparse_readout_prism.utils import write_csv +from sparse_readout_prism.research.cross_lens import ( + CENTERING_MODES, + centering_row_mean, + decompose_token, + feature_top_tokens, + import_jlens, + load_lens_model, + parse_dump_args, + resolve_k, + write_manifest, +) +from sparse_readout_prism.research.qwen_readout import load_sae +from sparse_readout_prism.utils import resolve_device, write_csv def log(msg): print(f"[antonym] {msg}", flush=True) -def _import_jlens(): - try: - import jlens - except ImportError as e: - raise ImportError( - "jlens (the Jacobian-lens reference implementation, Apache-2.0, " - "github.com/anthropics/jacobian-lens) is not installed; install it with `uv sync --extra lens`" - ) from e - return jlens - - -def load_sae(checkpoint): - """Raw TopK dictionary tensors from a training checkpoint (decoder rows unit-normalised).""" - ckpt = torch.load(checkpoint, map_location="cpu", weights_only=True) - state = ckpt["model_state_dict"] - decoder = state["decoder"].float().contiguous() # (d_features, d_model) - encoder_w = state["encoder.weight"].float().contiguous() # (d_features, d_model) - encoder_b = state["encoder.bias"].float().contiguous() # (d_features,) - decoder = decoder / decoder.norm(dim=1, keepdim=True).clamp_min(1e-8) - return decoder, encoder_w, encoder_b - - -def encode_topk(x, encoder_w, encoder_b, k): - acts = F.relu(x @ encoder_w.T + encoder_b) - values, indices = torch.topk(acts, k=min(k, acts.shape[-1]), dim=-1) - code = torch.zeros_like(acts) - code.scatter_(dim=-1, index=indices, src=values) - return code - - @torch.no_grad() def cmd_run(args) -> None: - import transformers + jlens = import_jlens() - jlens = _import_jlens() - - hf = transformers.AutoModelForCausalLM.from_pretrained(args.model_id, dtype=torch.bfloat16).cuda().eval() - tok = transformers.AutoTokenizer.from_pretrained(args.model_id) - model = jlens.from_hf(hf, tok) + device = resolve_device(args.device) + loaded = load_lens_model(args.model_id, device=device) + tok, model = loaded.tok, loaded.lens_model lens = jlens.JacobianLens.load(args.lens) layers = sorted(lens.jacobians) - J = {l: lens.jacobians[l].float().cuda() for l in layers} - W = hf.get_output_embeddings().weight.detach().float().cuda() # (vocab, d_model) - row_mean = W.mean(dim=0) - final_norm = hf.model.norm - decoder, encoder_w, encoder_b = load_sae(Path(args.sae)) - decoder, encoder_w, encoder_b = decoder.cuda(), encoder_w.cuda(), encoder_b.cuda() + J = {l: lens.jacobians[l].float().to(device) for l in layers} + W = loaded.lm_head.weight.detach().float().to(device) # (vocab, d_model) + final_norm = loaded.final_norm + decoder, encoder_w, encoder_b, config, ckpt_row_mean = load_sae(args.checkpoint) + k = resolve_k(args.k, config) + row_mean = centering_row_mean(W, args.centering, tok=tok, ckpt_row_mean=ckpt_row_mean) # (d_model,) + decoder, encoder_w, encoder_b = decoder.to(device), encoder_w.to(device), encoder_b.to(device) + decompose = partial( + decompose_token, + W=W, + row_mean=row_mean, + decoder=decoder, + encoder_w=encoder_w, + encoder_b=encoder_b, + k=k, + top_feats=16, + ) captured = {} @@ -112,7 +110,7 @@ def norm_hook(module, inputs, output): per_layer, final_logits, _ = lens.apply(model, args.prompt, positions=[-1]) h_ref = captured["final_norm_out"][0, -1].float().clone() - final_logits = final_logits.float().squeeze().cuda() + final_logits = final_logits.float().squeeze().to(device) greedy = tok.decode([int(final_logits.argmax())]) log(f"prompt={args.prompt!r} greedy={greedy!r}") @@ -132,22 +130,6 @@ def probe_ids(): probes = probe_ids() log(f"single-token probes: {list(probes)}") - def decompose(h_state, token_id): - W_row = W[token_id] - centered = W_row - row_mean - norm = centered.norm().clamp_min(1e-8) - code = encode_topk((centered / norm)[None, :], encoder_w, encoder_b, args.k)[0] - contributions = norm * code * (h_state @ decoder.T) - active = torch.nonzero(contributions != 0).flatten() - top = active[contributions[active].abs().argsort(descending=True)[:16]] - return { - "original_logit": float(h_state @ W_row), - "base": float(h_state @ row_mean), - "feature_sum": float(contributions.sum()), - "residual": float(h_state @ W_row) - float(h_state @ row_mean) - float(contributions.sum()), - "top_features": [{"id": int(f), "contribution": float(contributions[f])} for f in top], - } - rows, seen = {}, set() for l in layers: h = captured[("out", l)][0, -1].float() @@ -171,30 +153,14 @@ def decompose(h_state, token_id): for h in handles: h.remove() - fids = torch.tensor(sorted(seen), dtype=torch.long) - enc, bias = encoder_w[fids], encoder_b[fids] - best_s = torch.full((len(fids), 10), -float("inf"), device=W.device) - best_i = torch.zeros((len(fids), 10), dtype=torch.long, device=W.device) - for s in range(0, W.shape[0], 8192): - rws = W[s : s + 8192] - c = rws - row_mean - x = c / c.norm(dim=1, keepdim=True).clamp_min(1e-8) - sc = F.relu(x @ enc.T + bias).T - m = torch.cat([best_s, sc], 1) - mi = torch.cat([best_i, torch.arange(s, s + rws.shape[0], device=W.device).expand(len(fids), -1)], 1) - best_s, keep = torch.topk(m, k=10, dim=1) - best_i = torch.gather(mi, 1, keep) - labels = { - int(f): [tok.decode([t]) for t, sc in zip(best_i[j].tolist(), best_s[j].tolist()) if sc > 0] - for j, f in enumerate(fids.tolist()) - } + labels = feature_top_tokens(W, row_mean, seen, encoder_w, encoder_b, tok, top_tokens=10) out = Path(args.out) out.parent.mkdir(parents=True, exist_ok=True) payload = { "model": args.model_id, "lens": args.lens, - "sae": args.sae, + "sae": args.checkpoint, "prompt": args.prompt, "greedy": greedy, "layers": layers, @@ -204,6 +170,7 @@ def decompose(h_state, token_id): tmp = out.with_suffix(".json.tmp") tmp.write_text(json.dumps(payload, ensure_ascii=False)) os.replace(tmp, out) + write_manifest(out, args, prompt=args.prompt, greedy=greedy, layers=layers, k=k, centering=args.centering) for L in [str(l) for l in layers] + ["final"]: jl = " · ".join(f"{t}({p}%)" for t, p in rows[L]["jlens_top10"][:3]) @@ -212,17 +179,6 @@ def decompose(h_state, token_id): log(f"wrote -> {out}") -def parse_dump_args(specs: list[str]) -> list[tuple[str, Path]]: - """Parse repeated ``LABEL=path`` arguments, preserving order.""" - out = [] - for spec in specs: - label, sep, path = spec.partition("=") - if not sep or not label or not path: - raise ValueError(f"--dump expects LABEL=path, got {spec!r}") - out.append((label, Path(path))) - return out - - def antonym_layer_table(dumps: dict[str, dict], layers: list[str]) -> tuple[list[dict], list[dict]]: """Layer x lens rows of top-1 token / softmax share, and per-probe dominant features. @@ -273,12 +229,17 @@ def cmd_table(args) -> None: for L in layers: cells = [t for t in table if t["layer"] == L] print(f"{L:>5s} " + "".join(f"{c['top1_token']!r} {c['top1_softmax_pct']}%".rjust(28) for c in cells)) + written = [] if args.out_csv: write_csv(args.out_csv, table) + written.append(args.out_csv) print(f"wrote {args.out_csv}") if args.out_features_csv: write_csv(args.out_features_csv, features) + written.append(args.out_features_csv) print(f"wrote {args.out_features_csv}") + if written: + write_manifest(written[0], args, prompt=first["prompt"], layers=layers, lenses=labels) def main(argv: list[str] | None = None) -> int: @@ -288,9 +249,24 @@ def main(argv: list[str] | None = None) -> int: sp = sub.add_parser("run", help="read the prompt through one fitted lens (GPU)") sp.add_argument("--model-id", default="Qwen/Qwen3.5-9B") sp.add_argument("--lens", required=True, help="fitted lens .pt (JacobianLens container)") - sp.add_argument("--sae", required=True, help="readout SAE checkpoint.pt (paper: Qwen3.5-9B k=128 seed 0, 8x)") - sp.add_argument("--k", type=int, default=128) + sp.add_argument( + "--checkpoint", required=True, help="readout dictionary checkpoint.pt (paper: Qwen3.5-9B k=128 seed 0, 8x)" + ) + sp.add_argument( + "--k", + type=int, + default=None, + help="active features per row code (default: the checkpoint's trained k; a different value warns)", + ) + sp.add_argument( + "--centering", + choices=CENTERING_MODES, + default="live", + help="centering mean: live = full-vocabulary mean of the live head (paper), " + "trained = the checkpoint's row_mean (else the tokenizer's text-token mean)", + ) sp.add_argument("--prompt", default='"小"的反义词是"') + sp.add_argument("--device", default="cuda", help="torch device for the model and the dictionary (paper: cuda)") sp.add_argument("--out", required=True, help="output JSON") sp.set_defaults(fn=cmd_run) diff --git a/scripts/analyze/cross_lens_three_lens_prompt.py b/scripts/analyze/cross_lens_three_lens_prompt.py index d38040d..827abe8 100644 --- a/scripts/analyze/cross_lens_three_lens_prompt.py +++ b/scripts/analyze/cross_lens_three_lens_prompt.py @@ -8,7 +8,8 @@ logit (the table), the full top-5 reading, the lens logit, rank and decomposed (fp32) logit of every probed target (e.g. ' groß', ' large', ' big', '大'), and the dominant readout feature of each target's decomposition and of the lens's -own top-1 token. +own top-1 token. When any CSV is written, ``.manifest.json`` records +the provenance next to it. Paper run (the German antonym prompt ``antonym_de_01``, answer ``groß``, read by the English-, Chinese- and German-fitted 100-prompt Jacobian lenses; dumps produced @@ -28,23 +29,10 @@ import argparse import json -from pathlib import Path +from sparse_readout_prism.research.cross_lens import POS, parse_dump_args, write_manifest from sparse_readout_prism.utils import write_csv -POS = "-1" - - -def parse_dump_args(specs: list[str]) -> list[tuple[str, Path]]: - """Parse repeated ``LABEL=path`` arguments, preserving order.""" - out = [] - for spec in specs: - label, sep, path = spec.partition("=") - if not sep or not label or not path: - raise ValueError(f"--dump expects LABEL=path, got {spec!r}") - out.append((label, Path(path))) - return out - def select_record(dump: dict, prompt_id: str | None) -> dict: records = dump["records"] @@ -168,15 +156,15 @@ def main(argv: list[str] | None = None) -> int: f"{r['dominant_contribution']:+.2f} (sum {r['feature_sum']:.1f}, resid {r['residual']:.2f})" ) - if args.out_csv: - write_csv(args.out_csv, top1_rows) - print(f"wrote {args.out_csv}") - if args.out_targets_csv: - write_csv(args.out_targets_csv, target_rows) - print(f"wrote {args.out_targets_csv}") - if args.out_features_csv: - write_csv(args.out_features_csv, feature_rows) - print(f"wrote {args.out_features_csv}") + written = [] + outputs = ((args.out_csv, top1_rows), (args.out_targets_csv, target_rows), (args.out_features_csv, feature_rows)) + for path, rows in outputs: + if path: + write_csv(path, rows) + written.append(path) + print(f"wrote {path}") + if written: + write_manifest(written[0], args, prompt_id=first["id"], layers=layers, lenses=labels) return 0 diff --git a/scripts/analyze/nearest_rows_baseline.py b/scripts/analyze/nearest_rows_baseline.py index 898f834..ac6008d 100644 --- a/scripts/analyze/nearest_rows_baseline.py +++ b/scripts/analyze/nearest_rows_baseline.py @@ -30,12 +30,15 @@ rank, token_id, token, cosine) and ``nearest_rows_table.csv`` (the table layout: per rank the centered and raw token with cosines rounded to two decimals and glosses for the non-Latin tokens that occur in the paper's lists). + +Changed in 0.2.1: both CSVs are written through ``utils.write_csv`` with fixed +column lists (same columns, same bytes); an empty result (``--top-n 0``) now +gives 0-byte CSVs instead of crashing the table writer. """ from __future__ import annotations import argparse -import csv import json from pathlib import Path @@ -44,7 +47,7 @@ from sparse_readout_prism.data import resolve_row_mean from sparse_readout_prism.research.qwen_readout import load_qwen_model from sparse_readout_prism.research.run_io import run_provenance -from sparse_readout_prism.utils import find_lm_head_with_path +from sparse_readout_prism.utils import find_lm_head_with_path, write_csv # Glosses for the non-Latin tokens that enter the paper's top-12 lists. GLOSS = { @@ -58,6 +61,18 @@ " ошибки": "errors (ru)", } +LONG_FIELDS = ["contrast", "variant", "rank", "token_id", "token", "cosine"] +TABLE_FIELDS = [ + "contrast", + "rank", + "centered_cosine", + "centered_token", + "centered_gloss", + "raw_cosine", + "raw_token", + "raw_gloss", +] + def single_token_id(tokenizer, text: str) -> int | None: ids = tokenizer.encode(text, add_special_tokens=False) @@ -166,6 +181,31 @@ def table_rows(contrast_results: dict[str, dict]) -> list[dict]: return rows +def long_rows(contrast_results: dict[str, dict]) -> list[dict]: + """Long-form listing (contrast, variant, rank, token_id, token, cosine); the token column is ``repr``'d.""" + return [ + { + "contrast": name, + "variant": variant, + "rank": row["rank"], + "token_id": row["token_id"], + "token": repr(row["token"]), + "cosine": row["cosine"], + } + for name, c in contrast_results.items() + for variant in ("top_centered", "top_raw") + for row in c[variant] + ] + + +def write_tables(out_dir: Path, contrast_results: dict[str, dict]) -> None: + """``nearest_rows.csv`` (long form) and ``nearest_rows_table.csv`` (table layout); empty results give 0-byte files.""" + write_csv(out_dir / "nearest_rows.csv", long_rows(contrast_results), fieldnames=LONG_FIELDS, write_empty=True) + write_csv( + out_dir / "nearest_rows_table.csv", table_rows(contrast_results), fieldnames=TABLE_FIELDS, write_empty=True + ) + + def parse_contrast(spec: str) -> tuple[str, str]: parts = [p.strip() for p in spec.split(",")] if len(parts) != 2 or not all(parts): @@ -241,19 +281,7 @@ def main() -> int: } (args.out_dir / "nearest_rows.json").write_text(json.dumps(results, indent=2, ensure_ascii=False, default=str)) - with (args.out_dir / "nearest_rows.csv").open("w", newline="", encoding="utf-8") as f: - w = csv.writer(f) - w.writerow(["contrast", "variant", "rank", "token_id", "token", "cosine"]) - for name, c in contrast_results.items(): - for variant in ("top_centered", "top_raw"): - for row in c[variant]: - w.writerow([name, variant, row["rank"], row["token_id"], repr(row["token"]), row["cosine"]]) - - table = table_rows(contrast_results) - with (args.out_dir / "nearest_rows_table.csv").open("w", newline="", encoding="utf-8") as f: - w = csv.DictWriter(f, fieldnames=list(table[0].keys())) - w.writeheader() - w.writerows(table) + write_tables(args.out_dir, contrast_results) for name, c in contrast_results.items(): print(f"\n=== {name} ({c['definition']}) ===") diff --git a/scripts/data/build_cross_lens_de_bank.py b/scripts/data/build_cross_lens_de_bank.py index 2990c2d..a31e811 100644 --- a/scripts/data/build_cross_lens_de_bank.py +++ b/scripts/data/build_cross_lens_de_bank.py @@ -15,12 +15,16 @@ first (no leading space inside a quoted frame, a leading space otherwise), then the alternative, and drops the item if neither is a single token. - Cognate exclusion: after NFKD diacritic stripping, ss/eszett folding and - casefolding, items with Levenshtein(form_a, form_b) <= 2 or identical forms - are excluded. Exclusions are reported, not hidden (30 of 132 candidates in - the shipped bank). + casefolding (``research.cross_lens.fold_text``, shared with the EN-DE + aggregator's lexical call), items with Levenshtein(form_a, form_b) <= 2 or + identical forms are excluded. Exclusions are reported, not hidden (30 of + 132 candidates in the shipped bank). - null_targets are language-matched unrelated nouns [de_null, en_null], assigned round-robin, skipping any null equal to either form. +The exclusion report carries a ``provenance`` block (args, git commit, versions); +the bank JSON itself is unchanged by it. + Paper run:: build_cross_lens_de_bank.py --tokenizer Qwen/Qwen3.5-9B \ @@ -32,9 +36,11 @@ import argparse import json -import unicodedata from pathlib import Path +from sparse_readout_prism.research.cross_lens import fold_text +from sparse_readout_prism.research.run_io import run_provenance + # (group, concept, prompt, de_surface, en_surface, answer_lang) # answer_lang: which side is the in-context gold continuation (form_a). CANDIDATES = [ @@ -582,12 +588,6 @@ ) -def norm(s: str) -> str: - s = s.strip().casefold().replace("ß", "ss") - s = unicodedata.normalize("NFKD", s) - return "".join(c for c in s if not unicodedata.combining(c)) - - def levenshtein(a: str, b: str) -> int: if len(a) < len(b): a, b = b, a @@ -619,16 +619,10 @@ def build_bank(tok) -> tuple[list[dict], list[dict]]: pid = f"{group}_{counters[group]:02d}" is_ctrl = group.startswith("ctrl_") - if not is_ctrl and (norm(de) == norm(en) or levenshtein(norm(de), norm(en)) <= 2): - excluded.append( - { - "id": pid, - "concept": concept, - "de": de, - "en": en, - "reason": f"cognate (lev={levenshtein(norm(de), norm(en))})", - } - ) + de_f, en_f = fold_text(de), fold_text(en) + lev = levenshtein(de_f, en_f) + if not is_ctrl and (de_f == en_f or lev <= 2): + excluded.append({"id": pid, "concept": concept, "de": de, "en": en, "reason": f"cognate (lev={lev})"}) continue # In-context continuation: quoted frames take no leading space, plain @@ -652,14 +646,15 @@ def build_bank(tok) -> tuple[list[dict], list[dict]]: form_a, form_b = (de_tok, en_tok) if answer_lang == "de" else (en_tok, de_tok) - # Language-matched nulls, skipping collisions with either form. + # Language-matched nulls (space-prefixed when single-token), skipping + # collisions with either form. nulls = [] - for pool, prefer in ((DE_NULLS, True), (EN_NULLS, True)): + for pool in (DE_NULLS, EN_NULLS): for step in range(len(pool)): cand = pool[(null_i + step) % len(pool)] - if norm(cand) in (norm(de), norm(en)): + if fold_text(cand) in (de_f, en_f): continue - cand_tok = single_token(tok, cand, prefer_space=prefer) + cand_tok = single_token(tok, cand, prefer_space=True) if cand_tok is not None: nulls.append(cand_tok) break @@ -710,6 +705,7 @@ def main(argv: list[str] | None = None) -> int: "n_ctrl": len(prompts) - n_cross, "per_group": {g: sum(1 for p in prompts if p["group"] == g) for g in sorted({p["group"] for p in prompts})}, "excluded": excluded, + "provenance": run_provenance(args), } report_path = Path(args.report) report_path.parent.mkdir(parents=True, exist_ok=True) diff --git a/scripts/data/extract_model_readout.py b/scripts/data/extract_model_readout.py index b0c3406..203dc80 100644 --- a/scripts/data/extract_model_readout.py +++ b/scripts/data/extract_model_readout.py @@ -23,6 +23,8 @@ import torch +from sparse_readout_prism.data import token_mask_from_tokenizer + def log(msg: str) -> None: print(f"[extract {time.strftime('%H:%M:%S')}] {msg}", flush=True) @@ -290,11 +292,8 @@ def _remove_hooks(): # tokenizer vocab (padded/unused rows on models whose embedding matrix is # larger than the tokenizer) and (b) special tokens — including the # additional specials multimodal tokenizers register for vision/image ids. - token_mask = torch.zeros(vocab, dtype=torch.bool) - token_mask[: min(vocab, len(tok))] = True - special_ids = [i for i in (getattr(tok, "all_special_ids", None) or []) if 0 <= i < vocab] - token_mask[special_ids] = False - log(f"token_mask: {int(token_mask.sum())}/{vocab} rows kept ({len(special_ids)} specials dropped)") + token_mask = token_mask_from_tokenizer(tok, vocab) + log(f"token_mask: {int(token_mask.sum())}/{vocab} rows kept") payload = { "W_U_orig": W_U.float().cpu().contiguous(), # (vocab, d_model) fp32 diff --git a/scripts/eval/aggregate_cross_lens_en_de.py b/scripts/eval/aggregate_cross_lens_en_de.py index e731387..9a8eaf0 100644 --- a/scripts/eval/aggregate_cross_lens_en_de.py +++ b/scripts/eval/aggregate_cross_lens_en_de.py @@ -15,10 +15,23 @@ token is lexical rather than by script. In order: (1) membership in the bank item's own gold forms (the form in the item's language); (2) an umlaut / eszett cue means German; (3) membership in small embedded EN / DE common-word lists -(casefolded, diacritic-normalised); else OTHER. The per-prompt call is the -mid-band plurality vote (a 2-2 tie resolves to the alphabetically first label, -DE < EN < OTHER). The headline agreement metric never uses the language call; -it is a descriptive diagnostic. +(casefolded, diacritic-normalised through ``research.cross_lens.fold_text``); +else OTHER. The per-prompt call is the mid-band plurality vote (a 2-2 tie +resolves to the alphabetically first label, DE < EN < OTHER). The headline +agreement metric never uses the language call; it is a descriptive diagnostic. + +Options that change the population, not the arithmetic: + - ``--agreement-rule half`` (default, the paper) passes a vote on at least + ``ceil(n/2)`` scored mid-band layers (2 of 4); ``strict`` needs more than + half (3 of 4). Applies to every vote, divergence included. + - ``--null-population all`` (default, the paper) pools the unrelated-token + null over every bank item, controls included (204 comparisons on the + shipped bank); ``cross`` drops the 12 controls (180 comparisons), which + matches the EN-ZH population, whose controls carry no nulls. + +The metric bodies live in ``research.cross_lens`` (``aggregate_pair``); this +entry point supplies the EN-DE families, the lexical call and the divergence +column. The summary JSON carries a ``provenance`` block. Paper runs (``--prompts data/cross_lens/cross_lens_prompts_en_de.json --seed 0``):: @@ -33,16 +46,21 @@ from __future__ import annotations import argparse -import json -import math -import random -import unicodedata -from pathlib import Path +from sparse_readout_prism.research.cross_lens import ( + AGREEMENT_RULES, + NULL_POPULATIONS, + PairSpec, + aggregate_pair, + fold_text, + load_bank_items, + load_dump_records, + print_summary, + write_summary, +) +from sparse_readout_prism.research.run_io import run_provenance from sparse_readout_prism.utils import write_csv -MIDBAND = ["21", "24", "26", "29"] -POS = "-1" CROSS_GROUPS = ( "antonym_de", "trans_de2en", @@ -78,34 +96,18 @@ ) -def wilson(k, n, z=1.96): - if n == 0: - return (0.0, 0.0, 1.0) - p = k / n - d = 1 + z * z / n - c = (p + z * z / (2 * n)) / d - h = z * math.sqrt(p * (1 - p) / n + z * z / (4 * n * n)) / d - return (p, max(0.0, c - h), min(1.0, c + h)) - - -def norm(s): - s = s.strip().casefold().replace("ß", "ss") - s = unicodedata.normalize("NFKD", s) - return "".join(c for c in s if not unicodedata.combining(c)) - - -def lang_of(tok, item): +def lang_of(tok: str, item: dict) -> str: t = tok.strip() if not t: return "OTHER" - n = norm(t) + n = fold_text(t) forms = { "de": item["form_a"] if item["lang_a"] == "de" else item["form_b"], "en": item["form_a"] if item["lang_a"] == "en" else item["form_b"], } - if n == norm(forms["de"]): + if n == fold_text(forms["de"]): return "DE" - if n == norm(forms["en"]): + if n == fold_text(forms["en"]): return "EN" if any(c in t for c in "äöüÄÖÜß"): return "DE" @@ -117,24 +119,16 @@ def lang_of(tok, item): return "OTHER" -def dom_feat(rec, layer, target): - t = rec["layers"].get(layer, {}).get(POS, {}).get("targets", {}).get(target) - if not t or not t.get("top_features"): - return None - return t["top_features"][0]["id"] - - -def majority_same(rec_a, rec_b, target_a, target_b): - same = total = 0 - for L in MIDBAND: - fa, fb = dom_feat(rec_a, L, target_a), dom_feat(rec_b, L, target_b) - if fa is None or fb is None: - continue - total += 1 - same += int(fa == fb) - if total == 0: - return None - return same >= (total + 1) // 2 +SPEC = PairSpec( + tag_b="de", + cross_groups=CROSS_GROUPS, + surface_column="lang", + surface_call=lang_of, + split_labels=("EN", "DE"), + split_caption="EN-lens=EN & DE-lens=DE", + surface_caption="top-1 lexical call", + divergence=True, +) def main(argv: list[str] | None = None) -> int: @@ -143,120 +137,39 @@ def main(argv: list[str] | None = None) -> int: p.add_argument("--lens-b", required=True, help="dump under the German-fitted lens (output keys *_de)") p.add_argument("--prompts", required=True, help="EN-DE prompt bank JSON") p.add_argument("--seed", type=int, default=0, help="seed for the shuffled-pairing null") + p.add_argument( + "--agreement-rule", + choices=AGREEMENT_RULES, + default="half", + help="mid-band majority: half = at least ceil(n/2) layers (paper), strict = more than half", + ) + p.add_argument( + "--null-population", + choices=NULL_POPULATIONS, + default="all", + help="rows pooled into the unrelated-token floor: all = every bank item incl. controls (paper, /204), " + "cross = cross families only (/180, the EN-ZH population)", + ) p.add_argument("--out", required=True, help="summary JSON (per-group rates, floors, per-prompt rows)") p.add_argument("--rows-csv", default=None, help="optional per-prompt rows as CSV") args = p.parse_args(argv) - ar = {r["id"]: r for r in json.loads(Path(args.lens_a).read_text())["records"]} - br = {r["id"]: r for r in json.loads(Path(args.lens_b).read_text())["records"]} - bank_list = json.loads(Path(args.prompts).read_text())["prompts"] - assert len({v["id"] for v in bank_list}) == len(bank_list), "duplicate prompt ids in bank" - bank = {v["id"]: v for v in bank_list} - ids = sorted(set(ar) & set(br) & set(bank)) - print(f"[agg] {len(ids)} prompts present in both dumps") - - rows = [] - for rid in ids: - v, a, b = bank[rid], ar[rid], br[rid] - fa, fb = v["form_a"], v["form_b"] - row = {"id": rid, "group": v["group"], "concept": v.get("concept", "")} - row["cross_lens_pass"] = majority_same(a, b, fa, fa) - row["cross_form_en"] = majority_same(a, a, fa, fb) - row["cross_form_de"] = majority_same(b, b, fa, fb) - nulls = [majority_same(a, a, fa, nt) for nt in v.get("null_targets", [])] - nulls = [x for x in nulls if x is not None] - row["null_hits"] = sum(nulls) - row["null_total"] = len(nulls) - # top-1 language under each lens (mid-band vote), lexical call - for tag, rec in (("en", a), ("de", b)): - langs = [] - for L in MIDBAND: - top5 = rec["layers"].get(L, {}).get(POS, {}).get("top5") - if top5: - langs.append(lang_of(top5[0][0], v)) - # Plurality vote; a 2-2 tie resolves to the alphabetically first label - # (DE < EN < OTHER), fixed so the vote is deterministic. - row[f"top1_lang_{tag}"] = max(sorted(set(langs)), key=langs.count) if langs else "NA" - # divergence: do the two lenses' top-1 token strings differ (mid-band vote)? - diff = tot = 0 - for L in MIDBAND: - ta = a["layers"].get(L, {}).get(POS, {}).get("top5") - tb = b["layers"].get(L, {}).get(POS, {}).get("top5") - if ta and tb: - tot += 1 - diff += int(ta[0][0] != tb[0][0]) - row["token_diverges"] = (diff >= (tot + 1) // 2) if tot else None - rows.append(row) - - rng = random.Random(args.seed) - cross_ids = [r["id"] for r in rows if r["group"] in CROSS_GROUPS] - shuffle_hits = shuffle_total = 0 - for rid in cross_ids: - others = [x for x in cross_ids if bank[x]["concept"] != bank[rid]["concept"]] - for oid in rng.sample(others, min(3, len(others))): - res = majority_same(ar[rid], br[oid], bank[rid]["form_a"], bank[oid]["form_a"]) - if res is not None: - shuffle_total += 1 - shuffle_hits += int(res) - - def rate(sel): - vals = [r for r in rows if sel(r) and r["cross_lens_pass"] is not None] - k = sum(r["cross_lens_pass"] for r in vals) - return k, len(vals), wilson(k, len(vals)) - - summary: dict = {"per_group": {}, "rows": rows} - print("\n=== CROSS-LENS DOMINANT-FEATURE AGREEMENT (majority of mid-band) ===") - for grp in sorted({r["group"] for r in rows}): - k, n, (pt, lo, hi) = rate(lambda r, g=grp: r["group"] == g) - summary["per_group"][grp] = {"pass": k, "n": n, "rate": pt, "ci": [lo, hi]} - print(f" {grp:14s} {k:3d}/{n:<3d} {pt:.2f} [{lo:.2f}, {hi:.2f}]") - k, n, (pt, lo, hi) = rate(lambda r: r["group"] in CROSS_GROUPS) - summary["headline"] = {"pass": k, "n": n, "rate": pt, "ci": [lo, hi]} - print(f" {'ALL CROSS':14s} {k:3d}/{n:<3d} {pt:.2f} [{lo:.2f}, {hi:.2f}]") - - cf_en = [r["cross_form_en"] for r in rows if r["group"] in CROSS_GROUPS and r["cross_form_en"] is not None] - cf_de = [r["cross_form_de"] for r in rows if r["group"] in CROSS_GROUPS and r["cross_form_de"] is not None] - nk = sum(r["null_hits"] for r in rows) - nn = sum(r["null_total"] for r in rows) - summary["cross_form_en"] = {"pass": sum(cf_en), "n": len(cf_en)} - summary["cross_form_de"] = {"pass": sum(cf_de), "n": len(cf_de)} - summary["null_within_lens"] = {"pass": nk, "n": nn, "rate": wilson(nk, nn)[0]} - summary["null_shuffle_cross_lens"] = { - "pass": shuffle_hits, - "n": shuffle_total, - "rate": wilson(shuffle_hits, shuffle_total)[0], - } - div = [r["token_diverges"] for r in rows if r["group"] in CROSS_GROUPS and r["token_diverges"] is not None] - summary["divergence_rate"] = {"diverging": sum(div), "n": len(div)} - print("\n=== ONE FEATURE CARRIES BOTH FORMS (within-lens) ===") - print(f" EN lens: {sum(cf_en)}/{len(cf_en)} DE lens: {sum(cf_de)}/{len(cf_de)}") - print("\n=== NULL FLOORS ===") - print(f" (a) form vs unrelated token, within-lens: {nk}/{nn} ({wilson(nk, nn)[0]:.2f})") - print( - f" (b) cross-lens shuffled prompts: {shuffle_hits}/{shuffle_total} " - f"({wilson(shuffle_hits, shuffle_total)[0]:.2f})" + summary = aggregate_pair( + load_dump_records(args.lens_a), + load_dump_records(args.lens_b), + load_bank_items(args.prompts), + SPEC, + seed=args.seed, + rule=args.agreement_rule, + null_population=args.null_population, ) - print("\n=== TOKEN DIVERGENCE (descriptive) ===") - print(f" top-1 differs between lenses on {sum(div)}/{len(div)} cross prompts") - print("\n=== LANGUAGE-FOLLOWS-LENS (top-1 lexical call, mid-band vote) ===") - summary["lens_only_split"] = {} - for grp in sorted({r["group"] for r in rows}): - sel = [r for r in rows if r["group"] == grp] - flip = sum(1 for r in sel if r["top1_lang_en"] == "EN" and r["top1_lang_de"] == "DE") - summary["lens_only_split"][grp] = {"pass": flip, "n": len(sel)} - print(f" {grp:14s} EN-lens=EN & DE-lens=DE on {flip}/{len(sel)}") - cross_sel = [r for r in rows if r["group"] in CROSS_GROUPS] - flip = sum(1 for r in cross_sel if r["top1_lang_en"] == "EN" and r["top1_lang_de"] == "DE") - summary["lens_only_split"]["all_cross"] = {"pass": flip, "n": len(cross_sel)} - print(f" {'ALL CROSS':14s} EN-lens=EN & DE-lens=DE on {flip}/{len(cross_sel)}") - - out = Path(args.out) - out.parent.mkdir(parents=True, exist_ok=True) - with out.open("w", encoding="utf-8") as f: - json.dump(summary, f, ensure_ascii=False, indent=1) - print(f"\n[agg] wrote {out}") + print(f"[agg] {len(summary['rows'])} prompts present in both dumps") + print_summary(summary, SPEC) + summary["provenance"] = run_provenance(args) + write_summary(summary, args.out) + print(f"\n[agg] wrote {args.out}") if args.rows_csv: - write_csv(args.rows_csv, rows) + write_csv(args.rows_csv, summary["rows"]) print(f"[agg] wrote {args.rows_csv}") return 0 diff --git a/scripts/eval/aggregate_cross_lens_en_zh.py b/scripts/eval/aggregate_cross_lens_en_zh.py index 21ca8f7..fca4a7d 100644 --- a/scripts/eval/aggregate_cross_lens_en_zh.py +++ b/scripts/eval/aggregate_cross_lens_en_zh.py @@ -23,10 +23,18 @@ - control families (``ctrl_*``) reported separately, never pooled into the headline. +Majority rule. ``--agreement-rule half`` (default, the paper) passes a vote on at +least ``ceil(n/2)`` of the scored mid-band layers (2 of 4); ``strict`` needs more +than half (3 of 4). The rule applies to every vote above (headline, within-lens +cross-form, both null floors). The metric bodies live in +``research.cross_lens`` (``aggregate_pair``); this entry point supplies the +EN-ZH families and the script-based surface call. + Slot semantics: ``--lens-a`` is the English-fitted lens and ``--lens-b`` the Chinese-fitted lens (output keys ``*_en`` / ``*_zh``). For the construction comparison the two slots hold the English-fitted Jacobian lens and the -English-fitted ridge translator; the key names are unchanged. +English-fitted ridge translator; the key names are unchanged. The summary JSON +carries a ``provenance`` block. Paper runs (all with ``--prompts data/cross_lens/cross_lens_prompts_en_zh.json --seed 0``):: @@ -47,16 +55,20 @@ from __future__ import annotations import argparse -import json -import math -import random import unicodedata -from pathlib import Path +from sparse_readout_prism.research.cross_lens import ( + AGREEMENT_RULES, + PairSpec, + aggregate_pair, + load_bank_items, + load_dump_records, + print_summary, + write_summary, +) +from sparse_readout_prism.research.run_io import run_provenance from sparse_readout_prism.utils import write_csv -MIDBAND = ["21", "24", "26", "29"] -POS = "-1" CROSS_GROUPS = ( "antonym_zh", "trans_zh2en", @@ -68,17 +80,7 @@ ) -def wilson(k, n, z=1.96): - if n == 0: - return (0.0, 0.0, 1.0) - p = k / n - d = 1 + z * z / n - c = (p + z * z / (2 * n)) / d - h = z * math.sqrt(p * (1 - p) / n + z * z / (4 * n * n)) / d - return (p, max(0.0, c - h), min(1.0, c + h)) - - -def script_of(tok): +def script_of(tok: str) -> str: for ch in tok: if ch.strip() == "" or not ch.isalnum(): continue @@ -90,25 +92,15 @@ def script_of(tok): return "OTHER" -def dom_feat(rec, layer, target): - t = rec["layers"].get(layer, {}).get(POS, {}).get("targets", {}).get(target) - if not t or not t.get("top_features"): - return None - return t["top_features"][0]["id"] - - -def majority_same(rec_a, rec_b, target_a, target_b): - """Majority-of-midband agreement between dom(target_a in rec_a) and dom(target_b in rec_b).""" - same = total = 0 - for L in MIDBAND: - fa, fb = dom_feat(rec_a, L, target_a), dom_feat(rec_b, L, target_b) - if fa is None or fb is None: - continue - total += 1 - same += int(fa == fb) - if total == 0: - return None - return same >= (total + 1) // 2 +SPEC = PairSpec( + tag_b="zh", + cross_groups=CROSS_GROUPS, + surface_column="script", + surface_call=lambda tok, _item: script_of(tok), + split_labels=("LATIN", "CJK"), + split_caption="EN=Latin & ZH=CJK", + surface_caption="top-1 script", +) def main(argv: list[str] | None = None) -> int: @@ -117,110 +109,31 @@ def main(argv: list[str] | None = None) -> int: p.add_argument("--lens-b", required=True, help="dump under the Chinese-fitted lens (output keys *_zh)") p.add_argument("--prompts", required=True, help="EN-ZH prompt bank JSON") p.add_argument("--seed", type=int, default=0, help="seed for the shuffled-pairing null") + p.add_argument( + "--agreement-rule", + choices=AGREEMENT_RULES, + default="half", + help="mid-band majority: half = at least ceil(n/2) layers (paper), strict = more than half", + ) p.add_argument("--out", required=True, help="summary JSON (per-group rates, floors, per-prompt rows)") p.add_argument("--rows-csv", default=None, help="optional per-prompt rows as CSV") args = p.parse_args(argv) - enr = {r["id"]: r for r in json.loads(Path(args.lens_a).read_text())["records"]} - zhr = {r["id"]: r for r in json.loads(Path(args.lens_b).read_text())["records"]} - bank = {v["id"]: v for v in json.loads(Path(args.prompts).read_text())["prompts"]} - ids = sorted(set(enr) & set(zhr) & set(bank)) - print(f"[agg] {len(ids)} prompts present in both dumps") - - rows = [] - for rid in ids: - v, e, z = bank[rid], enr[rid], zhr[rid] - fa, fb = v["form_a"], v["form_b"] - row = {"id": rid, "group": v["group"], "concept": v.get("concept", "")} - # headline: same dominant feature for the SAME concept token, lens A vs lens B - row["cross_lens_pass"] = majority_same(e, z, fa, fa) - # one feature carries both surface forms, within each lens - row["cross_form_en"] = majority_same(e, e, fa, fb) - row["cross_form_zh"] = majority_same(z, z, fa, fb) - # null floor (a): form_a vs unrelated token, within lens A - nulls = [majority_same(e, e, fa, nt) for nt in v.get("null_targets", [])] - nulls = [x for x in nulls if x is not None] - row["null_hits"] = sum(nulls) - row["null_total"] = len(nulls) - # surface script of top-1 under each lens (mid-band vote) - for tag, rec in (("en", e), ("zh", z)): - scripts = [] - for L in MIDBAND: - top5 = rec["layers"].get(L, {}).get(POS, {}).get("top5") - if top5: - scripts.append(script_of(top5[0][0])) - # Plurality vote; a 2-2 tie resolves to the alphabetically first label - # (CJK < LATIN < OTHER), fixed so the vote is deterministic. - row[f"top1_script_{tag}"] = max(sorted(set(scripts)), key=scripts.count) if scripts else "NA" - rows.append(row) - - # null floor (b): shuffled pairing, form_a of prompt i under lens A vs form_a of j under lens B - rng = random.Random(args.seed) - cross_ids = [r["id"] for r in rows if r["group"] in CROSS_GROUPS] - shuffle_hits = shuffle_total = 0 - for rid in cross_ids: - others = [x for x in cross_ids if bank[x]["concept"] != bank[rid]["concept"]] - for oid in rng.sample(others, min(3, len(others))): - res = majority_same(enr[rid], zhr[oid], bank[rid]["form_a"], bank[oid]["form_a"]) - if res is not None: - shuffle_total += 1 - shuffle_hits += int(res) - - def rate(sel): - vals = [r for r in rows if sel(r) and r["cross_lens_pass"] is not None] - k = sum(r["cross_lens_pass"] for r in vals) - return k, len(vals), wilson(k, len(vals)) - - summary: dict = {"per_group": {}, "rows": rows} - print("\n=== CROSS-LENS DOMINANT-FEATURE AGREEMENT (majority of mid-band) ===") - for grp in sorted({r["group"] for r in rows}): - k, n, (pt, lo, hi) = rate(lambda r, g=grp: r["group"] == g) - summary["per_group"][grp] = {"pass": k, "n": n, "rate": pt, "ci": [lo, hi]} - print(f" {grp:14s} {k:3d}/{n:<3d} {pt:.2f} [{lo:.2f}, {hi:.2f}]") - k, n, (pt, lo, hi) = rate(lambda r: r["group"] in CROSS_GROUPS) - summary["headline"] = {"pass": k, "n": n, "rate": pt, "ci": [lo, hi]} - print(f" {'ALL CROSS':14s} {k:3d}/{n:<3d} {pt:.2f} [{lo:.2f}, {hi:.2f}]") - - cf_en = [r["cross_form_en"] for r in rows if r["group"] in CROSS_GROUPS and r["cross_form_en"] is not None] - cf_zh = [r["cross_form_zh"] for r in rows if r["group"] in CROSS_GROUPS and r["cross_form_zh"] is not None] - nk = sum(r["null_hits"] for r in rows) - nn = sum(r["null_total"] for r in rows) - summary["cross_form_en"] = {"pass": sum(cf_en), "n": len(cf_en)} - summary["cross_form_zh"] = {"pass": sum(cf_zh), "n": len(cf_zh)} - summary["null_within_lens"] = {"pass": nk, "n": nn, "rate": wilson(nk, nn)[0]} - summary["null_shuffle_cross_lens"] = { - "pass": shuffle_hits, - "n": shuffle_total, - "rate": wilson(shuffle_hits, shuffle_total)[0], - } - print("\n=== ONE FEATURE CARRIES BOTH FORMS (within-lens) ===") - print(f" EN lens: {sum(cf_en)}/{len(cf_en)} ZH lens: {sum(cf_zh)}/{len(cf_zh)}") - print("\n=== NULL FLOORS ===") - print(f" (a) form vs unrelated token, within-lens: {nk}/{nn} ({wilson(nk, nn)[0]:.2f})") - print( - f" (b) cross-lens shuffled prompts: {shuffle_hits}/{shuffle_total} " - f"({wilson(shuffle_hits, shuffle_total)[0]:.2f})" + summary = aggregate_pair( + load_dump_records(args.lens_a), + load_dump_records(args.lens_b), + load_bank_items(args.prompts), + SPEC, + seed=args.seed, + rule=args.agreement_rule, ) - - print("\n=== LANGUAGE-FOLLOWS-LENS (top-1 script, mid-band vote) ===") - summary["lens_only_split"] = {} - for grp in sorted({r["group"] for r in rows}): - sel = [r for r in rows if r["group"] == grp] - flip = sum(1 for r in sel if r["top1_script_en"] == "LATIN" and r["top1_script_zh"] == "CJK") - summary["lens_only_split"][grp] = {"pass": flip, "n": len(sel)} - print(f" {grp:14s} EN=Latin & ZH=CJK on {flip}/{len(sel)}") - cross_sel = [r for r in rows if r["group"] in CROSS_GROUPS] - flip = sum(1 for r in cross_sel if r["top1_script_en"] == "LATIN" and r["top1_script_zh"] == "CJK") - summary["lens_only_split"]["all_cross"] = {"pass": flip, "n": len(cross_sel)} - print(f" {'ALL CROSS':14s} EN=Latin & ZH=CJK on {flip}/{len(cross_sel)}") - - out = Path(args.out) - out.parent.mkdir(parents=True, exist_ok=True) - with out.open("w", encoding="utf-8") as f: - json.dump(summary, f, ensure_ascii=False, indent=1) - print(f"\n[agg] wrote {out}") + print(f"[agg] {len(summary['rows'])} prompts present in both dumps") + print_summary(summary, SPEC) + summary["provenance"] = run_provenance(args) + write_summary(summary, args.out) + print(f"\n[agg] wrote {args.out}") if args.rows_csv: - write_csv(args.rows_csv, rows) + write_csv(args.rows_csv, summary["rows"]) print(f"[agg] wrote {args.rows_csv}") return 0 diff --git a/scripts/eval/analyze_error_tails.py b/scripts/eval/analyze_error_tails.py index d314d70..084b7c1 100644 --- a/scripts/eval/analyze_error_tails.py +++ b/scripts/eval/analyze_error_tails.py @@ -53,6 +53,12 @@ uv run python scripts/eval/analyze_error_tails.py \\ --input-root results/direct_geometry_runs \\ --out-dir results/error_tails --n-boot 10000 --seed 20260711 + +Changed in 0.2.1: the two cluster bootstraps gather each resample through flat +per-cluster offsets (``_flat_clusters`` / ``_resample_index``) instead of a +per-cluster ``np.concatenate`` list; the random draws, the element order of +every resample and hence every output are unchanged (byte-identical on the six +paper run dirs apart from ``summary.json`` provenance). """ from __future__ import annotations @@ -160,19 +166,43 @@ def point_metrics(df: pd.DataFrame) -> dict[str, float | int]: return out +def _flat_clusters(clusters: dict[str, np.ndarray]) -> tuple[np.ndarray, np.ndarray, np.ndarray]: + """Flatten ``{cluster: row positions}`` (visited in sorted-key order) into one + position array plus per-cluster ``(start, length)`` offsets, so a cluster + resample becomes a single repeat/arange gather (``_resample_index``).""" + parts = [np.asarray(clusters[k], dtype=np.intp) for k in sorted(clusters)] + lengths = np.array([p.size for p in parts], dtype=np.intp) + starts = np.cumsum(lengths) - lengths + order = np.concatenate(parts) if parts else np.empty(0, dtype=np.intp) + return order, starts, lengths + + +def _resample_index(sampled: np.ndarray, starts: np.ndarray, lengths: np.ndarray) -> np.ndarray: + """Positions (into the flattened order) of the clusters in ``sampled``: each + cluster's rows contiguous and in order, clusters in sample order -- the + sequence ``np.concatenate([clusters[keys[i]] for i in sampled])`` yields, + without the per-cluster Python loop.""" + lens = lengths[sampled] + total = int(lens.sum()) + if total == 0: + return np.empty(0, dtype=np.intp) + return np.repeat(starts[sampled] - (np.cumsum(lens) - lens), lens) + np.arange(total, dtype=np.intp) + + def bootstrap_ci(df: pd.DataFrame, n_boot: int, seed: int) -> dict[str, float]: clusters = {str(k): idx.to_numpy() for k, idx in df.groupby("base_case_id", sort=True).groups.items()} keys = sorted(clusters) if len(keys) < 2 or n_boot <= 0: return {} rng = np.random.default_rng(seed) - rho = df["rho"].to_numpy(float) - sign = df["sign_match_strict"].to_numpy(float) - accepted = df["accepted"].to_numpy(float) + order, starts, lengths = _flat_clusters(clusters) + rho = df["rho"].to_numpy(float)[order] + sign = df["sign_match_strict"].to_numpy(float)[order] + accepted = df["accepted"].to_numpy(float)[order] stats = np.empty((n_boot, 5), dtype=float) for b in range(n_boot): sampled = rng.integers(0, len(keys), size=len(keys)) - idx = np.concatenate([clusters[keys[i]] for i in sampled]) + idx = _resample_index(sampled, starts, lengths) rb = rho[idx] stats[b] = ( rb.mean(), @@ -210,9 +240,25 @@ def summarize( return pd.DataFrame(records) +PAIRED_COLS = ("rho", "sign_match_strict", "accepted") + + +def _paired_stats(v: dict[str, np.ndarray], idx: np.ndarray) -> np.ndarray: + """SRP-minus-baseline mean rho, p95 rho, sign agreement and accepted rate over rows ``idx``.""" + rho_srp, rho_base = v["rho_srp"][idx], v["rho_base"][idx] + return np.array( + [ + rho_srp.mean() - rho_base.mean(), + np.percentile(rho_srp, 95) - np.percentile(rho_base, 95), + v["sign_match_strict_srp"][idx].mean() - v["sign_match_strict_base"][idx].mean(), + v["accepted_srp"][idx].mean() - v["accepted_base"][idx].mean(), + ] + ) + + def paired_differences(qdf: pd.DataFrame, n_boot: int, seed: int) -> pd.DataFrame: keys = ["bank", "base_case_id", "case_id", "query"] - cols = keys + ["rho", "sign_match_strict", "accepted"] + cols = keys + list(PAIRED_COLS) a = qdf[qdf.method == ANCHOR][cols].copy() results: list[dict] = [] for method in sorted(set(qdf.method) - {ANCHOR}): @@ -223,23 +269,14 @@ def paired_differences(qdf: pd.DataFrame, n_boot: int, seed: int) -> pd.DataFram clusters = {str(k): idx.to_numpy() for k, idx in m.groupby("base_case_id", sort=True).groups.items()} ckeys = sorted(clusters) rng = np.random.default_rng(seed + int(hashlib.sha1(method.encode()).hexdigest()[:7], 16)) - - def diffs(frame: pd.DataFrame) -> np.ndarray: - return np.array( - [ - frame.rho_srp.mean() - frame.rho_base.mean(), - np.percentile(frame.rho_srp, 95) - np.percentile(frame.rho_base, 95), - frame.sign_match_strict_srp.mean() - frame.sign_match_strict_base.mean(), - frame.accepted_srp.mean() - frame.accepted_base.mean(), - ] - ) - - point = diffs(m) + values = {f"{c}_{side}": m[f"{c}_{side}"].to_numpy(float) for c in PAIRED_COLS for side in ("srp", "base")} + point = _paired_stats(values, np.arange(len(m))) + order, starts, lengths = _flat_clusters(clusters) + flat = {name: arr[order] for name, arr in values.items()} boots = np.empty((n_boot, 4), float) for i in range(n_boot): sample = rng.integers(0, len(ckeys), size=len(ckeys)) - idx = np.concatenate([clusters[ckeys[j]] for j in sample]) - boots[i] = diffs(m.iloc[idx]) + boots[i] = _paired_stats(flat, _resample_index(sample, starts, lengths)) rec: dict[str, float | str | int] = {"baseline": method, "n_pairs": len(m)} for j, name in enumerate(("mean_rho", "p95_rho", "sign_agreement", "accepted_rate")): rec[f"delta_srp_minus_baseline_{name}"] = float(point[j]) diff --git a/scripts/eval/cross_seed_stability.py b/scripts/eval/cross_seed_stability.py index 2612916..372bafe 100644 --- a/scripts/eval/cross_seed_stability.py +++ b/scripts/eval/cross_seed_stability.py @@ -24,6 +24,11 @@ cross-contrast null's 90th percentile. Everything is local to the extraction payload and the checkpoints; no model forward pass. +``--centering {live,trained}`` selects the centering mean: ``live`` (default) +is the full-vocabulary mean of the payload's ``W_U``, how the paper's runs were +computed; ``trained`` is the dictionaries' stored training mean. The output +JSON carries a ``provenance`` block (command line, args, git hash, versions). + Dictionary-family discipline: pass dictionaries from ONE recipe (the seed-variation family trained from ``configs/sweeps/qwen35_2b_seedvar_base.yaml``). The script refuses dictionaries that differ in width or ``k``. Cross-recipe @@ -54,6 +59,7 @@ import numpy as np import torch +from sparse_readout_prism.research.run_io import run_provenance from sparse_readout_prism.research.seed_stability import ( center_rows, contrast_features, @@ -62,19 +68,13 @@ load_contrast_pairs, load_dictionary, load_readout, + resolve_centering, + summary_stats, + unit_rows, ) from sparse_readout_prism.utils import to_jsonable -def _stats(v) -> dict: - v = np.asarray(v, dtype=float) - return dict(mean=float(v.mean()), median=float(np.median(v)), p10=float(np.percentile(v, 10)), n=int(len(v))) - - -def _unit(d: torch.Tensor) -> torch.Tensor: - return d / d.norm(dim=1, keepdim=True).clamp_min(1e-8) - - def run( dicts: list, labels: list[int], @@ -84,13 +84,14 @@ def run( pairs: list[tuple], rng: np.random.Generator, *, + row_mean: torch.Tensor, top_m: int, top_r: int, n_sample: int, n_hidden: int, width_tag: str | None, ) -> dict: - W_c, rn, W_n = center_rows(W) + W_c, rn, W_n = center_rows(W, row_mean) n_seeds = len(dicts) cache: dict[int, str] = {} @@ -98,7 +99,7 @@ def run( side_sets: list[list[tuple[set, set]]] = [[] for _ in dicts] # [seed][ci] = (T_pos, T_neg) used_feats: list[set[int]] = [set() for _ in dicts] for si, d in enumerate(dicts): - dec = d[0] + dec = d.decoder for _a, _b, ia, ib in pairs: top_pos, top_neg = contrast_features(W_n, rn, d, ia, ib, top_m) used_feats[si].update(top_pos + top_neg) @@ -133,14 +134,14 @@ def run( basis = {} for s1 in range(n_seeds): for s2 in range(s1 + 1, n_seeds): - d1, d2 = dicts[s1][0], dicts[s2][0] - d1n, d2n = _unit(d1), _unit(d2) + d1, d2 = dicts[s1].decoder, dicts[s2].decoder + d1n, d2n = unit_rows(d1), unit_rows(d2) samp = torch.tensor(rng.choice(d1.shape[0], min(n_sample, d1.shape[0]), replace=False)) nn_rand = (d1n[samp] @ d2n.T).max(dim=1).values used = torch.tensor(sorted(used_feats[s1])) nn_used = (d1n[used] @ d2n.T).max(dim=1).values basis[f"s{labels[s1]}-s{labels[s2]}"] = dict( - nn_cos_random=_stats(nn_rand.tolist()), nn_cos_used=_stats(nn_used.tolist()) + nn_cos_random=summary_stats(nn_rand.tolist(), (10,)), nn_cos_used=summary_stats(nn_used.tolist(), (10,)) ) # projection correlation on held-out hidden states @@ -150,9 +151,9 @@ def run( H = H[torch.tensor(rng.choice(H.shape[0], min(n_hidden, H.shape[0]), replace=False))] for s1 in range(n_seeds): for s2 in range(s1 + 1, n_seeds): - d1, d2 = dicts[s1][0], dicts[s2][0] + d1, d2 = dicts[s1].decoder, dicts[s2].decoder used = torch.tensor(sorted(used_feats[s1])) - match = (_unit(d1[used]) @ _unit(d2).T).argmax(dim=1) + match = (unit_rows(d1[used]) @ unit_rows(d2).T).argmax(dim=1) p1 = H @ d1[used].T p2 = H @ d2[match].T r = torch.corrcoef(torch.stack([p1.flatten(), p2.flatten()]))[0, 1] @@ -160,13 +161,13 @@ def run( return dict( width=width_tag, - k=int(dicts[0][3]), + k=int(dicts[0].k), n_contrasts=len(pairs), top_m=top_m, top_r=top_r, - same_side=_stats(same_side), - cross_side=_stats(cross_side), - cross_contrast_null=_stats(cross_contrast), + same_side=summary_stats(same_side, (10,)), + cross_side=summary_stats(cross_side, (10,)), + cross_contrast_null=summary_stats(cross_contrast, (10,)), null_p90=null90, frac_contrasts_above_null_p90=frac_above_null, basis_nn_cosine=basis, @@ -184,9 +185,16 @@ def main() -> int: help="checkpoint.pt of each seed, one recipe (paper: seeds 0, 1, 2 of one width)", ) ap.add_argument("--seed-labels", nargs="+", type=int, default=None, help="labels for output keys (default 0..n-1)") - ap.add_argument("--w-u", type=Path, required=True, help="extraction payload {W_U_orig, h_LN, ...}") + ap.add_argument("--w-u", type=Path, required=True, help="extraction payload {W_U_orig, h_LN, token_mask, ...}") ap.add_argument("--bank", type=Path, required=True, help="curated A/B JSONL bank (target_a / target_b)") ap.add_argument("--tokenizer", default="Qwen/Qwen3.5-2B", help="HF tokenizer id or local path") + ap.add_argument( + "--centering", + choices=("live", "trained"), + default="live", + help="centering mean: full-vocabulary mean of W_U (live, the paper's runs) or the dictionaries' stored " + "training mean (trained)", + ) ap.add_argument("--width-tag", default=None, help='label stored in the output, e.g. "32x"') ap.add_argument("--max-contrasts", type=int, default=150) ap.add_argument("--seed", type=int, default=0) @@ -204,16 +212,17 @@ def main() -> int: ap.error("--seed-labels must have one entry per dictionary") rng = np.random.default_rng(args.seed) - W, h_LN = load_readout(args.w_u) + W, h_LN, token_mask = load_readout(args.w_u) from transformers import AutoTokenizer tok = AutoTokenizer.from_pretrained(args.tokenizer) dicts = [load_dictionary(p) for p in args.dictionaries] for lab, d in zip(labels, dicts): - print(f"seed {lab}: D={d[0].shape[0]} k={d[3]}") - if len({(d[0].shape[0], d[3]) for d in dicts}) != 1: + print(f"seed {lab}: D={d.decoder.shape[0]} k={d.k}") + if len({(d.decoder.shape[0], d.k) for d in dicts}) != 1: raise SystemExit("dictionaries differ in width or k; pass checkpoints from one dictionary family") + row_mean = resolve_centering(W, dicts, args.centering, token_mask, tok=tok) pairs = load_contrast_pairs(args.bank, tok, rng, args.max_contrasts) print(f"contrasts: {len(pairs)} unique single-token pairs") @@ -226,19 +235,21 @@ def main() -> int: tok, pairs, rng, + row_mean=row_mean, top_m=args.top_m, top_r=args.top_r, n_sample=args.n_sample, n_hidden=args.n_hidden, width_tag=args.width_tag, ) - args.out.parent.mkdir(parents=True, exist_ok=True) - args.out.write_text(json.dumps(to_jsonable(out), indent=2) + "\n") print(json.dumps({k: v for k, v in out.items() if k != "basis_nn_cosine"}, indent=2)) print( "basis NN-cosine (used features):", {k: round(v["nn_cos_used"]["median"], 3) for k, v in out["basis_nn_cosine"].items()}, ) + out["provenance"] = run_provenance(args) + args.out.parent.mkdir(parents=True, exist_ok=True) + args.out.write_text(json.dumps(to_jsonable(out), indent=2) + "\n") print(f"wrote {args.out}") return 0 diff --git a/scripts/eval/feature_group_matching.py b/scripts/eval/feature_group_matching.py index fb7363e..a335742 100644 --- a/scripts/eval/feature_group_matching.py +++ b/scripts/eval/feature_group_matching.py @@ -36,6 +36,11 @@ Seed pairs are all ordered pairs (query, candidate) with query index below candidate index, so three dictionaries give s0->s1, s0->s2, s1->s2. +``--centering {live,trained}`` selects the centering mean: ``live`` (default) +is the full-vocabulary mean of the payload's ``W_U``, how the paper's runs were +computed; ``trained`` is the dictionaries' stored training mean. The output +JSON carries a ``provenance`` block (command line, args, git hash, versions). + Dictionary-family discipline: pass dictionaries from ONE recipe (the seed-variation family trained from ``configs/sweeps/qwen35_2b_seedvar_base.yaml``); the script refuses dictionaries that differ in width or ``k``. @@ -52,6 +57,25 @@ --out results/seed_stability/feature_group_matching_32x.json (16x: the ``qwen2b_d32768_k128_s{0,1,2}`` checkpoints, ``--width-tag 16x``.) + +Fixed in 0.2.1: + +* The frequency-matched null pool is enumerated in sorted token order + (``null_token_pool``). It was enumerated in ``Counter`` insertion order, + which follows ``frozenset`` iteration and therefore ``PYTHONHASHSEED``: the + pseudo-groups drawn for the null changed from process to process, and, + because ``rng.choice(..., replace=False, p=...)`` consumes a data-dependent + amount of generator state, so did the query features sampled for every seed + pair after the first. The paper's runs were produced under an unrecorded + hash seed, so a rerun reproduces the protocol and the first seed pair's + query population exactly, while the null draws (hence ``null99``, + ``frac_above_null_p99``, ``null_best_jaccard``) and the later pairs' query + samples are a fresh, now reproducible, sample of the same distributions. +* The below-null tail of ``--direction-check`` is selected with the unrounded + best Jaccard, the value ``frac_above_null_p99`` is computed from. The + per-query record's 4-decimal ``best_jaccard`` was compared before, so a + query within 5e-5 of ``null99`` could be counted on different sides by the + two summaries. On the paper's 16x/32x outputs the two counts already agree. """ from __future__ import annotations @@ -65,13 +89,17 @@ import numpy as np import torch +from sparse_readout_prism.research.run_io import run_provenance from sparse_readout_prism.research.seed_stability import ( center_rows, contrast_features, load_contrast_pairs, load_dictionary, load_readout, + resolve_centering, + summary_stats, token_strings, + unit_rows, ) from sparse_readout_prism.utils import resolve_device, to_jsonable @@ -89,6 +117,19 @@ def all_feature_groups(dec: torch.Tensor, W_c: torch.Tensor, tok, device, cache: return [frozenset(token_strings(tok, top_rows[f].tolist(), cache)) for f in range(D)] +def null_token_pool(cand_groups: list[frozenset]) -> tuple[list[str], np.ndarray]: + """Frequency-matched null pool: tokens in sorted order and their group-frequency weights. + + Sorted enumeration is what makes the null draws independent of the hash + seed (``Counter`` insertion order follows ``frozenset`` iteration). + """ + tok_freq = Counter(t for g in cand_groups for t in g) + pool_toks = sorted(tok_freq) + pool_p = np.array([tok_freq[t] for t in pool_toks], dtype=float) + pool_p /= pool_p.sum() + return pool_toks, pool_p + + def best_matches(query: frozenset, inv_index: dict, groups: list[frozenset], greedy_pool: int, greedy_k: int): """``(best_jaccard, best_feature, greedy recall@1..greedy_k)`` over all candidate groups.""" counts: Counter = Counter() @@ -118,21 +159,6 @@ def best_matches(query: frozenset, inv_index: dict, groups: list[frozenset], gre return best_j, best_f, recalls -def _stats(v) -> dict: - v = np.asarray(v, dtype=float) - return dict( - mean=float(v.mean()), - median=float(np.median(v)), - p90=float(np.percentile(v, 90)), - p99=float(np.percentile(v, 99)), - n=int(len(v)), - ) - - -def _unit(d: torch.Tensor) -> torch.Tensor: - return d / d.norm(dim=1, keepdim=True).clamp_min(1e-8) - - def run( dicts: list, labels: list[int], @@ -142,6 +168,7 @@ def run( rng: np.random.Generator, device, *, + row_mean: torch.Tensor, top_m: int, top_r: int, n_query: int, @@ -154,11 +181,11 @@ def run( direction_null: int, direction_seed: int, ) -> dict: - W_c, rn, W_n = center_rows(W) + W_c, rn, W_n = center_rows(W, row_mean) cache: dict[int, str] = {} # token groups of every feature, every seed - groups = [all_feature_groups(d[0], W_c, tok, device, cache, top_r) for d in dicts] + groups = [all_feature_groups(d.decoder, W_c, tok, device, cache, top_r) for d in dicts] print("feature groups computed", flush=True) # used features per seed (same construction as cross_seed_stability.py) @@ -171,8 +198,9 @@ def run( print("used features:", [len(u) for u in used]) seed_pairs = list(itertools.combinations(range(len(dicts)), 2)) # (query seed, candidate seed) - unit_decs = [_unit(d[0]) for d in dicts] + unit_decs = [unit_rows(d.decoder) for d in dicts] results: dict[str, dict] = {} + unrounded_best: dict[str, list[float]] = {} # per pair, the best Jaccards before 4-dp rounding for qs, cs in seed_pairs: cand_groups = groups[cs] inv: dict[str, list[int]] = {} @@ -180,11 +208,7 @@ def run( for t in g: inv.setdefault(t, []).append(f) - # frequency-matched null token pool - tok_freq = Counter(t for g in cand_groups for t in g) - pool_toks = list(tok_freq) - pool_p = np.array([tok_freq[t] for t in pool_toks], dtype=float) - pool_p /= pool_p.sum() + pool_toks, pool_p = null_token_pool(cand_groups) # queries: sampled used features of the query seed with large-enough groups cands = [f for f in sorted(used[qs]) if len(groups[qs][f]) >= min_group_size] @@ -222,21 +246,22 @@ def run( n_rec3 = [r[-1] for r in n_rec] res = dict( n_query=int(len(qs_feats)), - best_jaccard=_stats(q_best), - null_best_jaccard=_stats(n_best), + best_jaccard=summary_stats(q_best, (90, 99)), + null_best_jaccard=summary_stats(n_best, (90, 99)), frac_above_null_p99=float((np.array(q_best) > null99).mean()), frac_jaccard_ge_050=float((np.array(q_best) >= 0.5).mean()), frac_jaccard_ge_025=float((np.array(q_best) >= 0.25).mean()), - recall_at_1=_stats([r[0] for r in q_rec]), - recall_at_3=_stats(rec3), - null_recall_at_3=_stats(n_rec3), + recall_at_1=summary_stats([r[0] for r in q_rec], (90, 99)), + recall_at_3=summary_stats(rec3, (90, 99)), + null_recall_at_3=summary_stats(n_rec3, (90, 99)), frac_recall3_above_null_p99=float((np.array(rec3) > np.percentile(n_rec3, 99)).mean()), - matched_decoder_cosine=_stats(q_cos), + matched_decoder_cosine=summary_stats(q_cos, (90, 99)), null99=null99, per_query=q_records, ) key = f"s{labels[qs]}->s{labels[cs]}" results[key] = res + unrounded_best[key] = q_best print( f"{key}: best-J med {res['best_jaccard']['median']:.3f} " f"(null p99 {null99:.3f}), >null {res['frac_above_null_p99']:.2f}, " @@ -247,11 +272,12 @@ def run( if direction_check: drng = np.random.default_rng(direction_seed) - d_model = dicts[0][0].shape[1] + d_model = dicts[0].decoder.shape[1] for qs, cs in seed_pairs: key = f"s{labels[qs]}->s{labels[cs]}" r = results[key] - fails = [q for q in r["per_query"] if q["best_jaccard"] <= r["null99"]] + # same comparison as frac_above_null_p99 (unrounded best Jaccard), so the two summaries agree + fails = [q for q, bj in zip(r["per_query"], unrounded_best[key]) if bj <= r["null99"]] dq, dc = unit_decs[qs], unit_decs[cs] g = torch.tensor(drng.standard_normal((direction_null, d_model)), dtype=torch.float32) g = g / g.norm(dim=1, keepdim=True) @@ -294,9 +320,16 @@ def main() -> int: help="checkpoint.pt of each seed, one recipe (paper: seeds 0, 1, 2 of one width)", ) ap.add_argument("--seed-labels", nargs="+", type=int, default=None, help="labels for output keys (default 0..n-1)") - ap.add_argument("--w-u", type=Path, required=True, help="extraction payload {W_U_orig, ...}") + ap.add_argument("--w-u", type=Path, required=True, help="extraction payload {W_U_orig, token_mask, ...}") ap.add_argument("--bank", type=Path, required=True, help="curated A/B JSONL bank (target_a / target_b)") ap.add_argument("--tokenizer", default="Qwen/Qwen3.5-2B", help="HF tokenizer id or local path") + ap.add_argument( + "--centering", + choices=("live", "trained"), + default="live", + help="centering mean: full-vocabulary mean of W_U (live, the paper's runs) or the dictionaries' stored " + "training mean (trained)", + ) ap.add_argument("--width-tag", default=None, help='label stored in the output, e.g. "32x"') ap.add_argument("--max-contrasts", type=int, default=150) ap.add_argument("--seed", type=int, default=0) @@ -323,16 +356,17 @@ def main() -> int: rng = np.random.default_rng(args.seed) device = resolve_device(args.device) print(f"device={device}") - W, _ = load_readout(args.w_u) + W, _, token_mask = load_readout(args.w_u) from transformers import AutoTokenizer tok = AutoTokenizer.from_pretrained(args.tokenizer) dicts = [load_dictionary(p) for p in args.dictionaries] for lab, d in zip(labels, dicts): - print(f"seed {lab}: D={d[0].shape[0]} k={d[3]}") - if len({(d[0].shape[0], d[3]) for d in dicts}) != 1: + print(f"seed {lab}: D={d.decoder.shape[0]} k={d.k}") + if len({(d.decoder.shape[0], d.k) for d in dicts}) != 1: raise SystemExit("dictionaries differ in width or k; pass checkpoints from one dictionary family") + row_mean = resolve_centering(W, dicts, args.centering, token_mask, tok=tok) pairs = load_contrast_pairs(args.bank, tok, rng, args.max_contrasts) print(f"contrasts: {len(pairs)} unique single-token pairs") @@ -345,6 +379,7 @@ def main() -> int: pairs, rng, device, + row_mean=row_mean, top_m=args.top_m, top_r=args.top_r, n_query=args.n_query, @@ -357,6 +392,7 @@ def main() -> int: direction_null=args.direction_null, direction_seed=args.direction_seed, ) + out["provenance"] = run_provenance(args) args.out.parent.mkdir(parents=True, exist_ok=True) args.out.write_text(json.dumps(to_jsonable(out), indent=2) + "\n") print(f"wrote {args.out}") diff --git a/scripts/eval/loo_core_recovery.py b/scripts/eval/loo_core_recovery.py index 4ce294f..8a639ec 100644 --- a/scripts/eval/loo_core_recovery.py +++ b/scripts/eval/loo_core_recovery.py @@ -20,6 +20,25 @@ contrast-clustered bootstrap (resample contrasts, keep all their cells, ``--n-boot`` replicates) for each method's mean recall and for the SRP/kNN ratio. +Two conventions of the paper's run are kept and made explicit: + +* The kNN side set excludes the target row itself (``--knn-exclude-self``, + the default). The SRP side sets (top-R rows of the contrast's features) and + the cluster side set (the target's cluster members) do not exclude it, so + the target token can sit in a core and be recoverable by SRP and by the + cluster control but never by kNN. ``--no-knn-exclude-self`` admits the + target's own row (cosine 1, so it always takes one of the ``--side-n`` + slots). +* Cluster membership is the assignment that produced the last k-means update + (``spherical_kmeans_unit(..., final_assignment=False)``), one Lloyd step + behind the centroids the members are ranked against. The ``row_cluster_*`` + baselines of the direct-geometry grid assign against the final centroids. + +``--centering {live,trained}`` selects the centering mean: ``live`` (default) +is the full-vocabulary mean of the payload's ``W_U``, how the paper's runs were +computed; ``trained`` is the dictionaries' stored training mean. The output +JSON carries a ``provenance`` block (command line, args, git hash, versions). + Dictionary-family discipline: ``--dictionaries`` must come from ONE recipe (the seed-variation family trained from ``configs/sweeps/qwen35_2b_seedvar_base.yaml``); the script refuses dictionaries that differ in width or ``k``. A dictionary @@ -55,6 +74,8 @@ import numpy as np import torch +from sparse_readout_prism.research.row_geometry import spherical_kmeans_unit +from sparse_readout_prism.research.run_io import run_provenance from sparse_readout_prism.research.seed_stability import ( center_rows, contrast_features, @@ -62,33 +83,26 @@ load_contrast_pairs, load_dictionary, load_readout, + resolve_centering, token_strings, ) from sparse_readout_prism.utils import resolve_device, to_jsonable -def kmeans_unit(X: torch.Tensor, n: int, iters: int, seed: int, device) -> tuple[torch.Tensor, torch.Tensor]: - """Spherical k-means on unit rows: ``(centroids (n, d), assignment (N,))`` on CPU.""" - g = torch.Generator().manual_seed(seed) - Xd = X.to(device) - C = Xd[torch.randperm(X.shape[0], generator=g)[:n].to(device)].clone() - ones = torch.ones(X.shape[0], device=device) - assign = torch.empty(X.shape[0], dtype=torch.long, device=device) - for it in range(iters): - for s in range(0, X.shape[0], 4096): - assign[s : s + 4096] = (Xd[s : s + 4096] @ C.T).argmax(dim=1) - C_new = torch.zeros_like(C) - cnt = torch.zeros(n, device=device) - C_new.index_add_(0, assign, Xd) - cnt.index_add_(0, assign, ones) - dead = cnt == 0 - C = C_new / cnt.clamp_min(1.0)[:, None] - nd = int(dead.sum()) - if nd: - C[dead] = Xd[torch.randperm(X.shape[0], generator=g)[:nd].to(device)] - C = C / C.norm(dim=1, keepdim=True).clamp_min(1e-8) - print(f" kmeans iter {it + 1}/{iters} (dead={nd})", flush=True) - return C.cpu(), assign.cpu() +def knn_side_sets( + W_n: torch.Tensor, pairs: list[tuple], side_n: int, tok, cache: dict[int, str], *, exclude_self: bool +) -> list[tuple[set[str], set[str]]]: + """Per contrast, the ``(A, B)`` token sets of each target's top ``side_n`` cosine neighbours.""" + out = [] + for _a, _b, ia, ib in pairs: + sides = [] + for i in (ia, ib): + sims = W_n @ W_n[i] + if exclude_self: + sims[i] = -1 + sides.append(token_strings(tok, torch.topk(sims, side_n).indices.tolist(), cache)) + out.append(tuple(sides)) + return out def run( @@ -98,6 +112,7 @@ def run( pairs: list[tuple], device, *, + row_mean: torch.Tensor, top_m: int, top_r: int, side_n: int, @@ -108,8 +123,9 @@ def run( bootstrap_seed: int, reference_dict=None, width_tag: str | None, + knn_exclude_self: bool = True, ) -> dict: - W_c, rn, W_n = center_rows(W) + W_c, rn, W_n = center_rows(W, row_mean) cache: dict[int, str] = {} n_seeds = len(dicts) @@ -119,8 +135,8 @@ def dict_side_sets(d) -> list[tuple[set, set]]: top_pos, top_neg = contrast_features(W_n, rn, d, ia, ib, top_m) out.append( ( - feature_token_set(top_pos, d[0], W_c, tok, top_r, cache), - feature_token_set(top_neg, d[0], W_c, tok, top_r, cache), + feature_token_set(top_pos, d.decoder, W_c, tok, top_r, cache), + feature_token_set(top_neg, d.decoder, W_c, tok, top_r, cache), ) ) return out @@ -130,14 +146,7 @@ def dict_side_sets(d) -> list[tuple[set, set]]: side_sets.append(dict_side_sets(d)) print(f"seed {si} side sets done", flush=True) - knn_sets = [] - for _a, _b, ia, ib in pairs: - sides = [] - for i in (ia, ib): - sims = W_n @ W_n[i] - sims[i] = -1 - sides.append(token_strings(tok, torch.topk(sims, side_n).indices.tolist(), cache)) - knn_sets.append(tuple(sides)) + knn_sets = knn_side_sets(W_n, pairs, side_n, tok, cache, exclude_self=knn_exclude_self) ref_sets = None if reference_dict is not None: @@ -145,7 +154,10 @@ def dict_side_sets(d) -> list[tuple[set, set]]: print("reference dictionary side sets done", flush=True) print(f"fitting k-means n_clusters={n_clusters} on {device} ...", flush=True) - C, assign = kmeans_unit(W_n, n_clusters, kmeans_iters, kmeans_seed, device) + C, assign = spherical_kmeans_unit( + W_n.to(device), n_clusters, kmeans_seed, iters=kmeans_iters, return_assignments=True, final_assignment=False + ) + C, assign = C.cpu(), assign.cpu() def cluster_side(i: int) -> set[str]: members = (assign == assign[i]).nonzero().flatten() @@ -226,9 +238,16 @@ def main() -> int: required=True, help="checkpoint.pt of each seed, one recipe, at least three (paper: seeds 0, 1, 2 of one width)", ) - ap.add_argument("--w-u", type=Path, required=True, help="extraction payload {W_U_orig, ...}") + ap.add_argument("--w-u", type=Path, required=True, help="extraction payload {W_U_orig, token_mask, ...}") ap.add_argument("--bank", type=Path, required=True, help="curated A/B JSONL bank (target_a / target_b)") ap.add_argument("--tokenizer", default="Qwen/Qwen3.5-2B", help="HF tokenizer id or local path") + ap.add_argument( + "--centering", + choices=("live", "trained"), + default="live", + help="centering mean: full-vocabulary mean of W_U (live, the paper's runs) or the dictionaries' stored " + "training mean (trained)", + ) ap.add_argument("--width-tag", default=None, help='label stored in the output, e.g. "32x"') ap.add_argument( "--reference-dict", @@ -241,6 +260,12 @@ def main() -> int: ap.add_argument("--top-m", type=int, default=8, help="features per side") ap.add_argument("--top-r", type=int, default=12, help="rows per feature for the token summary") ap.add_argument("--side-n", type=int, default=96, help="tokens per kNN / cluster side set") + ap.add_argument( + "--knn-exclude-self", + action=argparse.BooleanOptionalAction, + default=True, + help="drop the target row from its own kNN side set (the paper's run); --no-knn-exclude-self admits it", + ) ap.add_argument("--n-clusters", type=int, default=16384) ap.add_argument("--kmeans-iters", type=int, default=12) ap.add_argument("--kmeans-seed", type=int, default=0) @@ -255,17 +280,18 @@ def main() -> int: rng = np.random.default_rng(args.seed) device = resolve_device(args.device) - W, _ = load_readout(args.w_u) + W, _, token_mask = load_readout(args.w_u) from transformers import AutoTokenizer tok = AutoTokenizer.from_pretrained(args.tokenizer) dicts = [load_dictionary(p) for p in args.dictionaries] for si, d in enumerate(dicts): - print(f"seed {si}: D={d[0].shape[0]} k={d[3]}") - if len({(d[0].shape[0], d[3]) for d in dicts}) != 1: + print(f"seed {si}: D={d.decoder.shape[0]} k={d.k}") + if len({(d.decoder.shape[0], d.k) for d in dicts}) != 1: raise SystemExit("dictionaries differ in width or k; pass checkpoints from one dictionary family") reference = load_dictionary(args.reference_dict) if args.reference_dict is not None else None + row_mean = resolve_centering(W, dicts, args.centering, token_mask, tok=tok) pairs = load_contrast_pairs(args.bank, tok, rng, args.max_contrasts) print(f"contrasts: {len(pairs)} unique single-token pairs") @@ -276,6 +302,7 @@ def main() -> int: tok, pairs, device, + row_mean=row_mean, top_m=args.top_m, top_r=args.top_r, side_n=args.side_n, @@ -286,10 +313,12 @@ def main() -> int: bootstrap_seed=args.bootstrap_seed, reference_dict=reference, width_tag=args.width_tag, + knn_exclude_self=args.knn_exclude_self, ) + print(json.dumps(to_jsonable(out), indent=2)) + out["provenance"] = run_provenance(args) args.out.parent.mkdir(parents=True, exist_ok=True) args.out.write_text(json.dumps(to_jsonable(out), indent=2) + "\n") - print(json.dumps(to_jsonable(out), indent=2)) print(f"wrote {args.out}") return 0 diff --git a/scripts/eval/paired_matched_kl_bootstrap.py b/scripts/eval/paired_matched_kl_bootstrap.py index cfe3cc3..84b9825 100644 --- a/scripts/eval/paired_matched_kl_bootstrap.py +++ b/scripts/eval/paired_matched_kl_bootstrap.py @@ -21,9 +21,15 @@ ``bad_prob_reduction`` (primary) and ``flip`` (secondary). Comparisons with fewer than 20 aligned candidates are skipped. -Output: a JSON list with one record per (model, target KL, baseline, outcome): -the matched scales and their median KLs, ``mean_diff``, ``ci_term`` and -``ci_prompt`` (95% percentile intervals). +Output: a JSON object with ``comparisons``, a list with one record per (model, +target KL, baseline, outcome) holding the matched scales and their median KLs, +``mean_diff``, ``ci_term`` and ``ci_prompt`` (95% percentile intervals), and +``provenance`` (command line, args, git hash, versions, timestamp). + +Changed in 0.2.1: the output gained the ``provenance`` key, so the top level is +an object rather than the bare list the paper file used; the ``comparisons`` +records are unchanged (the paper's 96 records reproduce exactly from the paper +CSVs). Paper run (10,000 resamples, seed 0, target KL 0.02 / 0.05 / 0.10 / 0.20):: @@ -44,6 +50,8 @@ import numpy as np +from sparse_readout_prism.research.run_io import run_provenance + BASELINES = ["mean_row_direction", "pca_group_rank4", "pca_group_direction"] SRP = "feature_suppression" DEFAULT_TARGET_KLS = [0.02, 0.05, 0.1, 0.2] @@ -188,7 +196,7 @@ def main(argv: list[str] | None = None) -> int: paired_comparisons(model, rows, target_kls=list(args.target_kl), n_boot=args.n_boot, seed=args.seed) ) args.out.parent.mkdir(parents=True, exist_ok=True) - args.out.write_text(json.dumps(out_rows, indent=2)) + args.out.write_text(json.dumps({"comparisons": out_rows, "provenance": run_provenance(args)}, indent=2)) print(f"\nwrote {args.out} ({len(out_rows)} comparisons)") return 0 diff --git a/scripts/eval/run_causal_contribution_validation.py b/scripts/eval/run_causal_contribution_validation.py index fbecd58..d10f09c 100644 --- a/scripts/eval/run_causal_contribution_validation.py +++ b/scripts/eval/run_causal_contribution_validation.py @@ -36,6 +36,25 @@ is realized-on-predicted. Slope > 1 means the realized change exceeds the prediction; slope < 1 means it falls short. +Centering (row-mean convention). The codes z are TopK codes of rows centered +against a row mean and per-row normalized (``data.center_normalize_rows``). +``--centering live`` (the default, and how the paper's six runs were computed) +centers against the full-vocabulary mean of the live float32 ``lm_head`` +weight, ``W_U.mean(0)`` over all V rows including special and padded rows. +``--centering trained`` centers the way the dictionary was trained +(``data.centering_mean(mode="trained")``): the checkpoint's stored ``row_mean`` +when it has one, else the mean over the tokenizer's text-token rows +(``data.token_mask_from_tokenizer``). The max-abs gap between the checkpoint +mean and the live mean is logged either way. On the released Qwen3.5-2B +dictionary the two means differ by ~0.04% of a centered row norm. + +Token resolution uses ``research.row_geometry.resolve_single_token_bare_first`` +(the bare form first, then the space-prefixed form), as the paper run did. The +fidelity runners use the space-first, special-rejecting +``research.registry.resolve_single_token_strict``; the two orders can pick +different rows for terms where both forms are single tokens, so the bare-first +order is kept here to reproduce the paper's contrast set rather than unified. + ``summary.json`` subsets: ``gated_predicted`` (covered contrasts, rho below ``--rho-gate`` with sign agreement; the table's r^2, CI, slope and n), ``ungated_predicted`` (all contrasts) and ``gated_random_control`` (the table's @@ -54,7 +73,8 @@ single token. Outputs: ``causal_rows.csv`` (one row per prediction-realization pair), -``summary.json`` and ``manifest.json`` (run provenance). +``summary.json`` and ``manifest.json`` (run provenance, including +``--centering``). Paper run: six models, each at the fidelity operating point (32x width, k=256; Hub path ``/k256_32x/checkpoint.pt``), every other flag at its @@ -75,10 +95,18 @@ Defaults the paper run relied on: ``--banks curated_ab,case_candidates,model_native --max-native 300 --top-features 10 --random-per-case 10 --rho-gate 0.5 ---n-boot 2000 --max-len 64``. +--n-boot 2000 --max-len 64 --centering live``. ``--self-test`` runs the synthetic residual-free check of the math core on CPU -(no model, no checkpoint) and exits. +(no model, no checkpoint) and exits. Its third check builds a dense two-row LM +head W = [w_A; w_B] with w_A - w_B = q and re-reads the margin from W before +and after ablating h, so the stored realized change is compared against an +independent measurement rather than against the formula that produced it. + +Changed in 0.2.1: ``--centering {live,trained}`` (default ``live``, the paper's +behaviour); the file-local resolver was replaced by the shared bare-first one +(same variant order, same ids); the self-test's third check is the dense-head +comparison described above (it used to re-evaluate (h . d_i)(q . d_i)). """ from __future__ import annotations @@ -91,9 +119,10 @@ import numpy as np import torch -from sparse_readout_prism.data import center_normalize_rows +from sparse_readout_prism.data import center_normalize_rows, centering_mean, token_mask_from_tokenizer from sparse_readout_prism.factorizers import TopKSAE, load_factorizer from sparse_readout_prism.research.qwen_readout import encode_topk +from sparse_readout_prism.research.row_geometry import resolve_single_token_bare_first from sparse_readout_prism.research.run_io import load_bank, run_provenance from sparse_readout_prism.utils import find_lm_head, load_causal_lm, resolve_device, set_seed, write_csv, write_json @@ -110,7 +139,7 @@ def log(msg: str) -> None: # --------------------------------------------------------------------------- # -# Banks and tokens +# Banks, tokens and centering # --------------------------------------------------------------------------- # @@ -131,13 +160,19 @@ def load_banks(bank_dir: Path, banks: list[str], caps: dict[str, int]) -> list[d return rows -def single_token_id(tok, term: str): - """Token id of ``term`` if it (bare first, then space-prefixed) is a single token, else None.""" - for variant in (term, " " + term): - ids = tok.encode(variant, add_special_tokens=False) - if len(ids) == 1: - return ids[0] - return None +def select_row_mean(mode: str, W_U: torch.Tensor, *, ckpt: dict, tok) -> tuple[torch.Tensor, str]: + """Centering mean ``(d,)`` under ``--centering`` plus a label for the log line. + + ``live`` is the full-vocabulary mean of ``W_U`` exactly as handed in (the + live float32 lm_head weight, reduced on its own device); ``trained`` is the + checkpoint's stored ``row_mean``, else the mean over the tokenizer's + text-token rows. Both go through ``data.centering_mean``. + """ + if mode == "live": + return centering_mean(W_U, mode="live").to(W_U.device), "the live full-vocabulary mean" + token_mask = token_mask_from_tokenizer(tok, W_U.shape[0]).to(W_U.device) + source = "the checkpoint row_mean" if ckpt.get("row_mean") is not None else "the text-token mean of the live W_U" + return centering_mean(W_U, mode="trained", token_mask=token_mask, ckpt=ckpt).to(W_U.device), source # --------------------------------------------------------------------------- # @@ -208,6 +243,12 @@ def cluster_bootstrap_r2(rows: list[dict], n_boot: int, seed: int) -> tuple[floa return float(np.percentile(stats, 2.5)), float(np.percentile(stats, 97.5)) +def dense_head_margin(W: torch.Tensor, state: torch.Tensor) -> float: + """Margin (W state)_0 - (W state)_1 read off a dense two-row LM head ``W`` ``(2, d)``.""" + logits = W.double() @ state.double() # (2,) + return float(logits[0] - logits[1]) + + def self_test_metrics(seed: int = 0) -> dict: """Synthetic residual-free check of the math core (no model). @@ -236,27 +277,34 @@ def self_test_metrics(seed: int = 0) -> dict: ctrl_pred = np.array([abs(p[1]) for p in pairs if p[3] == 1]) ctrl_max_abs_pred = float(ctrl_pred.max(initial=0.0)) ok2 = ctrl_max_abs_pred < 1e-6 # random features carry ~0 prediction - # Identity check: realized values equal the algebraic form - i0 = pairs[0][0] - lhs = pairs[0][2] - rhs = float((W_dec[i0] @ h) * (W_dec[i0] @ q)) - identity_abs_err = abs(lhs - rhs) - ok3 = identity_abs_err < 1e-4 + # Dense-head check: a two-row LM head W = [w_A; w_B] with w_A - w_B = q. The + # margin is read off W before and after ablating h along d_i; the stored + # delta_real must equal -(m(h') - m(h)) for every tested feature. The margin + # never sees the (h . d_i)(q . d_i) formula that produced delta_real. + w_b = torch.randn(d) + W_head = torch.stack([w_b + q, w_b]) # (2, d) dense LM head, rows A and B + m_base = dense_head_margin(W_head, h) + dense_head_abs_err = max( + abs((dense_head_margin(W_head, h - (h @ W_dec[fid]) * W_dec[fid]) - m_base) + delta_real) + for fid, _c_pred, delta_real, _is_random in pairs + ) + ok3 = dense_head_abs_err < 1e-4 return { "r2": r2, "slope": slope, "ctrl_max_abs_pred": ctrl_max_abs_pred, - "identity_abs_err": identity_abs_err, + "dense_head_abs_err": dense_head_abs_err, "checks": [ ("r2/slope on span-constructed q", ok1), ("random controls predict ~0", ok2), - ("realized identity", ok3), + ("realized change matches the dense LM head", ok3), ], "pairs": pairs, "h": h, "q": q, "beta": beta, "W_dec": W_dec, + "W_head": W_head, } @@ -283,6 +331,13 @@ def main() -> int: ap.add_argument("--out-dir", type=Path) ap.add_argument("--device", default="cuda") ap.add_argument("--dtype", choices=["bfloat16", "float32"], default="bfloat16") + ap.add_argument( + "--centering", + choices=["live", "trained"], + default="live", + help="row-centering mean: live = full-vocabulary mean of the live lm_head (paper); " + "trained = the checkpoint's row_mean, else the tokenizer text-token mean", + ) ap.add_argument("--top-features", type=int, default=10) ap.add_argument("--random-per-case", type=int, default=10) ap.add_argument("--max-len", type=int, default=64) @@ -312,7 +367,6 @@ def main() -> int: lm_head = find_lm_head(model) W_U = lm_head.weight.detach().float().to(device) # (V, d) V, d = W_U.shape - row_mean = W_U.mean(dim=0) # (d,) centering mean of the live readout log(f"W_U ({V}, {d}) from live lm_head") log(f"loading dictionary {args.checkpoint}") @@ -326,10 +380,15 @@ def main() -> int: enc_b = sae.encoder.bias.detach().float() # (D,) dec_norm = W_dec.norm(dim=1) W_dec_unit = W_dec / dec_norm[:, None].clamp_min(1e-8) # (D, d) unit rows + row_mean, mean_source = select_row_mean(args.centering, W_U, ckpt=ckpt, tok=tok) # (d,) ckpt_row_mean = ckpt.get("row_mean") if ckpt_row_mean is not None: - gap = float((ckpt_row_mean.float().to(device) - row_mean).abs().max()) - log(f"checkpoint row_mean vs live W_U mean: max abs gap {gap:.3e} (the live mean is used)") + live_mean = row_mean if args.centering == "live" else centering_mean(W_U, mode="live").to(device) + gap = float((ckpt_row_mean.float().to(device) - live_mean).abs().max()) + log( + f"checkpoint row_mean vs live W_U mean: max abs gap {gap:.3e} " + f"(--centering {args.centering}: {mean_source} is used)" + ) log(f"dictionary D={W_dec.shape[0]} k={k} (decoder norms {dec_norm.min():.3f}-{dec_norm.max():.3f})") banks = [b.strip() for b in args.banks.split(",") if b.strip()] @@ -373,7 +432,7 @@ def raw_beta(a_id: int, b_id: int) -> torch.Tensor: if not (a and b and prompt): n_skip += 1 continue - a_id, b_id = single_token_id(tok, a), single_token_id(tok, b) + a_id, b_id = resolve_single_token_bare_first(tok, a), resolve_single_token_bare_first(tok, b) if a_id is None or b_id is None or a_id == b_id: n_skip += 1 continue diff --git a/scripts/figures/compute_cross_lens_shared_feature.py b/scripts/figures/compute_cross_lens_shared_feature.py index 172270a..3aaa9a6 100644 --- a/scripts/figures/compute_cross_lens_shared_feature.py +++ b/scripts/figures/compute_cross_lens_shared_feature.py @@ -10,7 +10,8 @@ reports whether the dominant feature is shared across the lenses, each lens's share carried by the shared feature, and the feature's top unembedding rows from the dump's ``feature_top_tokens`` labels. No rendering; the figure is drawn in -the paper source from these numbers. +the paper source from these numbers. The metrics JSON carries a ``provenance`` +block; a CSV-only run writes ``.manifest.json`` instead. Paper run (factual-recall prompt ``fac_03``, "The capital of China is Beijing. The capital of the UK is", layer 26, English- and Chinese-fitted 100-prompt Jacobian @@ -25,23 +26,11 @@ import argparse import json -from pathlib import Path +from sparse_readout_prism.research.cross_lens import POS, parse_dump_args, write_manifest +from sparse_readout_prism.research.run_io import run_provenance from sparse_readout_prism.utils import write_csv, write_json -POS = "-1" - - -def parse_dump_args(specs: list[str]) -> list[tuple[str, Path]]: - """Parse repeated ``LABEL=path`` arguments, preserving order.""" - out = [] - for spec in specs: - label, sep, path = spec.partition("=") - if not sep or not label or not path: - raise ValueError(f"--dump expects LABEL=path, got {spec!r}") - out.append((label, Path(path))) - return out - def lens_top1_summary(record: dict, layer: str) -> dict: """Top-1 token, its logit, and the dominant-feature share of its decomposition at one layer.""" @@ -134,8 +123,11 @@ def main(argv: list[str] | None = None) -> int: write_csv(args.out_csv, rows) print(f"wrote {args.out_csv}") if args.out_json: + metrics["provenance"] = run_provenance(args) write_json(metrics, args.out_json) print(f"wrote {args.out_json}") + elif args.out_csv: + write_manifest(args.out_csv, args, prompt_id=args.prompt_id, layer=args.layer) return 0 diff --git a/scripts/run/fit_jlens.py b/scripts/run/fit_jlens.py index 692bdd7..f91f450 100644 --- a/scripts/run/fit_jlens.py +++ b/scripts/run/fit_jlens.py @@ -21,12 +21,22 @@ Apache-2.0, https://github.com/anthropics/jacobian-lens); this script only orchestrates sharding and storage. Jacobians accumulate in fp32 and are saved in fp16. Install the dependency with ``uv sync --extra lens``; the import is lazy so -the rest of the repository installs and tests without it. +the rest of the repository installs and tests without it. The model is loaded +through ``research.cross_lens.load_lens_model`` (``utils.load_causal_lm`` on +``--device``, default ``cuda``). Source layers are ``pick_source_layers(n_layers, --n-layers)``: evenly spaced over 5-95% depth, excluding layer 0 and the final (target) layer. For Qwen3.5-9B and ``--n-layers 12`` this gives {2, 5, 7, 10, 12, 14, 17, 19, 21, 24, 26, 29}. +Resume safety. ``fit`` writes ``shard{i}.meta.json`` next to the shard lens +(model id, prompt manifest, sha1 of the prompt slice ``[start, start+n_prompts)``, +``n_prompts``, ``start``, shard index and count, ``n_layers``) and refuses to +reuse an existing ``shard{i}.lens.pt`` or ``shard{i}.ckpt.pt`` whose sidecar +differs. ``merge`` requires the sidecars, checks that every shard agrees on +them, and records the shared values (plus per-shard prompt counts and +provenance) in ``.meta.json``. + Paper runs (Qwen3.5-9B, four shards on four GPUs, ``--dim-batch 8``):: # seeded prompt dumps, 1000 prompts each (the English dump keeps the default --min-chars 800) @@ -48,6 +58,14 @@ The same fit + merge with the zh / de dumps produce the Chinese- and German-fitted lenses (``qwen35_9b_jlens_zh_seed0_n100.pt``, ``qwen35_9b_jlens_de_seed0_n100.pt``), and with ``--n-prompts 300`` the three ``*_n300.pt`` refits. + +Fixed in 0.2.1: + - The out-of-memory fallback in ``fit`` stopped once ``dim_batch < 2``, so + ``--dim-batch 1`` was never attempted and ``2`` / ``3`` got no fallback; it + now tries the requested value and halves down to 1 (8 -> 8, 4, 2, 1; 3 -> 3, 1). + - Shard resume only checked that ``shard{i}.lens.pt`` existed, so a re-run + with a different prompt slice, budget or layer count reused a stale shard + or resumed jlens's checkpoint into a mixed fit; see "Resume safety" above. """ from __future__ import annotations @@ -61,39 +79,25 @@ import torch -from sparse_readout_prism.utils import git_commit +from sparse_readout_prism.research.cross_lens import ( + check_resume_meta, + cli_args, + import_jlens, + load_lens_model, + prompt_list_sha1, + write_resume_meta, +) +from sparse_readout_prism.research.run_io import run_provenance +from sparse_readout_prism.utils import git_commit, read_json, resolve_device, write_json + +# The fit parameters every shard of one lens must share (checked by ``merge``). +SHARED_META_KEYS = ("model_id", "prompts_json", "prompts_sha1", "n_prompts", "start", "num_shards", "n_layers") def log(msg: str) -> None: print(f"[fit_jlens] {msg}", flush=True) -def _import_jlens(): - try: - import jlens - except ImportError as e: - raise ImportError( - "jlens (the Jacobian-lens reference implementation, Apache-2.0, " - "github.com/anthropics/jacobian-lens) is not installed; install it with `uv sync --extra lens`" - ) from e - return jlens - - -def load_model(model_id: str): - import transformers - - jlens = _import_jlens() - hf = ( - transformers.AutoModelForCausalLM.from_pretrained( - model_id, torch_dtype=torch.bfloat16, attn_implementation="sdpa" - ) - .cuda() - .eval() - ) - tok = transformers.AutoTokenizer.from_pretrained(model_id) - return jlens.from_hf(hf, tok) - - def pick_source_layers(n_layers: int, n_pick: int) -> list[int]: # Evenly spaced over 5-95% depth; excludes layer 0 and the final layer # (target). Deterministic given (n_layers, n_pick). @@ -157,9 +161,9 @@ def read_prompts(path: str, n: int | None = None, start: int = 0) -> list[str]: def cmd_smoke(args) -> None: - jlens = _import_jlens() + jlens = import_jlens() - model = load_model(args.model_id) + model = load_lens_model(args.model_id, device=resolve_device(args.device)).lens_model prompts = read_prompts(args.prompts_json, 2) layers = pick_source_layers(model.n_layers, 4) t0 = time.time() @@ -181,45 +185,76 @@ def cmd_smoke(args) -> None: ) +def dim_batch_schedule(dim_batch: int) -> list[int]: + """The requested ``dim_batch``, then halved down to 1: 8 -> [8, 4, 2, 1], 3 -> [3, 1], 1 -> [1].""" + if dim_batch < 1: + raise ValueError(f"--dim-batch must be >= 1, got {dim_batch}") + schedule = [] + while dim_batch >= 1: + schedule.append(dim_batch) + dim_batch //= 2 + return schedule + + +def fit_with_fallback(fit, model, prompts: list[str], layers: list[int], dim_batch: int, ckpt: Path): + """``fit`` (``jlens.fit``) at each :func:`dim_batch_schedule` value until one does not OOM. + + Every attempt resumes jlens's own checkpoint at ``ckpt``, so a fallback + continues from the prompts already accumulated. The retry happens after + the except block has closed: the exception traceback pins the failed + attempt's autograd graph, and an in-except retry OOMs against that ghost + memory. + """ + for db in dim_batch_schedule(dim_batch): + try: + return fit(model, prompts, source_layers=layers, dim_batch=db, checkpoint_path=str(ckpt), resume=True) + except torch.cuda.OutOfMemoryError: + pass + log(f"OOM at dim_batch={db}; freeing and halving") + gc.collect() + torch.cuda.empty_cache() + raise RuntimeError("all dim_batch fallbacks OOMed") + + +def shard_meta(args, prompts: list[str]) -> dict: + """The parameters that define one shard's fit, stored in ``shard{i}.meta.json``.""" + return { + "model_id": args.model_id, + "prompts_json": args.prompts_json, + "prompts_sha1": prompt_list_sha1(prompts), + "n_prompts": args.n_prompts, + "start": args.start, + "shard": args.shard, + "num_shards": args.num_shards, + "n_layers": args.n_layers, + } + + def cmd_fit(args) -> None: - jlens = _import_jlens() + jlens = import_jlens() prompts = read_prompts(args.prompts_json, args.n_prompts, start=args.start) shard_prompts = prompts[args.shard :: args.num_shards] - model = load_model(args.model_id) - layers = pick_source_layers(model.n_layers, args.n_layers) ckpt = Path(args.ckpt_dir) / f"shard{args.shard}.ckpt.pt" - ckpt.parent.mkdir(parents=True, exist_ok=True) out = Path(args.out_dir) / f"shard{args.shard}.lens.pt" + meta_path = Path(args.out_dir) / f"shard{args.shard}.meta.json" + meta = shard_meta(args, prompts) + check_resume_meta(meta_path, meta, [out, ckpt]) if out.exists(): log(f"[resume] shard {args.shard} lens exists: {out}") return + ckpt.parent.mkdir(parents=True, exist_ok=True) + write_resume_meta(meta_path, meta) + + model = load_lens_model(args.model_id, device=resolve_device(args.device)).lens_model + layers = pick_source_layers(model.n_layers, args.n_layers) log( f"git={git_commit()} model={args.model_id} shard={args.shard}/{args.num_shards} " f"prompt_slice=[{args.start},{args.start + args.n_prompts}) prompts={len(shard_prompts)} " f"layers={layers} dim_batch={args.dim_batch} ckpt={ckpt}" ) t0 = time.time() - # Retry OUTSIDE the except block: the exception traceback pins the failed - # attempt's autograd graph, so an in-except retry OOMs against ghost memory. - lens = None - for dim_batch in (args.dim_batch, args.dim_batch // 2, args.dim_batch // 4): - if dim_batch < 2: - break - oom = False - try: - lens = jlens.fit( - model, shard_prompts, source_layers=layers, dim_batch=dim_batch, checkpoint_path=str(ckpt), resume=True - ) - except torch.cuda.OutOfMemoryError: - oom = True - if not oom: - break - log(f"OOM at dim_batch={dim_batch}; freeing and halving") - gc.collect() - torch.cuda.empty_cache() - if lens is None: - raise RuntimeError("all dim_batch fallbacks OOMed") + lens = fit_with_fallback(jlens.fit, model, shard_prompts, layers, args.dim_batch, ckpt) for layer, J in lens.jacobians.items(): assert torch.isfinite(J).all(), f"non-finite Jacobian at layer {layer}" tmp = str(out) + ".tmp" @@ -228,14 +263,41 @@ def cmd_fit(args) -> None: log(f"shard {args.shard} done: {lens.n_prompts} prompts in {(time.time() - t0) / 3600:.2f} h -> {out}") +def check_shard_metas_agree(metas: list[dict], *, num_shards: int) -> dict: + """The fit parameters shared by all shard sidecars; raises naming any they disagree on.""" + diffs = [] + for key in SHARED_META_KEYS: + vals = [m.get(key) for m in metas] + if any(v != vals[0] for v in vals): + diffs.append(f"{key}: {vals}") + if [m.get("shard") for m in metas] != list(range(len(metas))): + diffs.append(f"shard: {[m.get('shard') for m in metas]} (expected 0..{len(metas) - 1})") + if metas[0].get("num_shards") != num_shards: + diffs.append(f"num_shards: sidecars say {metas[0].get('num_shards')}, --num-shards is {num_shards}") + if diffs: + raise RuntimeError("shard sidecars disagree; these shards are not one fit\n " + "\n ".join(diffs)) + return {k: metas[0].get(k) for k in SHARED_META_KEYS} + + def cmd_merge(args) -> None: - jlens = _import_jlens() + jlens = import_jlens() final = Path(args.out) if final.exists(): log(f"[resume] merged lens exists: {final}") return - lenses = [jlens.JacobianLens.load(str(Path(args.shard_dir) / f"shard{i}.lens.pt")) for i in range(args.num_shards)] + shard_dir = Path(args.shard_dir) + metas = [] + for i in range(args.num_shards): + meta_path = shard_dir / f"shard{i}.meta.json" + if not meta_path.exists(): + raise FileNotFoundError( + f"{meta_path} missing: shards written before 0.2.1 carry no sidecar, so their fit cannot be " + "verified; refit the shard (or write the sidecar by hand from the run log)" + ) + metas.append(read_json(meta_path)) + shared = check_shard_metas_agree(metas, num_shards=args.num_shards) + lenses = [jlens.JacobianLens.load(str(shard_dir / f"shard{i}.lens.pt")) for i in range(args.num_shards)] merged = jlens.JacobianLens.merge(lenses) # Convergence report: first half of the shards vs all of them. Small # relative deltas mean the prompt budget has saturated. @@ -248,6 +310,16 @@ def cmd_merge(args) -> None: tmp = str(final) + ".tmp" merged.save(tmp) os.replace(tmp, final) + write_json( + { + **shared, + "shard_n_prompts": [lens.n_prompts for lens in lenses], + "merged_n_prompts": merged.n_prompts, + "layers": merged.source_layers, + "provenance": run_provenance(cli_args(args)), + }, + final.with_suffix(".meta.json"), + ) log(f"merged {sum(lens.n_prompts for lens in lenses)} prompts, {len(merged.jacobians)} layers -> {final}") @@ -273,6 +345,7 @@ def main(argv: list[str] | None = None) -> int: sp.add_argument("--prompts-json", required=True) sp.add_argument("--out", required=True) sp.add_argument("--dim-batch", type=int, default=16) + sp.add_argument("--device", default="cuda", help="torch device for the model (paper: cuda)") sp.set_defaults(fn=cmd_smoke) sp = sub.add_parser("fit", help="fit one prompt shard on one GPU") @@ -283,9 +356,10 @@ def main(argv: list[str] | None = None) -> int: sp.add_argument("--shard", type=int, required=True) sp.add_argument("--num-shards", type=int, required=True) sp.add_argument("--n-layers", type=int, default=12) - sp.add_argument("--dim-batch", type=int, default=16) + sp.add_argument("--dim-batch", type=int, default=16, help="output dims per backward pass; halved on OOM down to 1") sp.add_argument("--ckpt-dir", required=True) sp.add_argument("--out-dir", required=True) + sp.add_argument("--device", default="cuda", help="torch device for the model (paper: cuda)") sp.set_defaults(fn=cmd_fit) sp = sub.add_parser("merge", help="merge shard lenses into one lens file") diff --git a/scripts/run/fit_ridge_lens.py b/scripts/run/fit_ridge_lens.py index b3e513b..4dba3ba 100644 --- a/scripts/run/fit_ridge_lens.py +++ b/scripts/run/fit_ridge_lens.py @@ -24,12 +24,25 @@ - lambda grid: g * tr(XtX) / (N * d) for g in {1e-3, 1e-2, 1e-1, 1, 10}; - diagnostics reported, not enforced: per-layer holdout R^2, and top-1 agreement between decode(transport(h_l)) and the model's final logits on the - last 5 holdout prompts at the two deepest fitted layers (written to - ``.report.json``). - -Resumable: the second-moment accumulators checkpoint every 10 prompts (atomic), -and the script exits early if ``--out`` already exists. ``jlens`` is imported -lazily (``uv sync --extra lens``). + last (up to) 5 holdout prompts at the two deepest fitted layers (written to + ``.report.json`` together with a ``provenance`` block). + +Lambda grid, as run. The multipliers ``g`` are scaled per token +(``lambda = g * tr(XtX) / (N * d)``, ``N`` = tokens in the fit prompts) but the +penalty is added to the N-token Gram sum ``XtX`` itself, so relative to the Gram +matrix the effective shrinkage is ``g / N`` -- small at every grid point -- and +both paper fits selected the grid's top value ``g = 10`` at all twelve layers +(``lambda_mult`` in the shipped ``.report.json`` files). The math is kept +exactly as the paper ran it; widening or rescaling the grid changes the fitted +translators and is a paper decision, not a code fix. + +Resumable: the second-moment accumulators checkpoint every 10 prompts (atomic) +under ``/_{fit,hold}.ckpt``, guarded by ``/.meta.json`` +(model id, prompt manifest, sha1 of the prompt list, ``n_prompts``, ``holdout``, +``tag``, layers, ``max_seq_len``): existing checkpoints are refused when the +sidecar differs. The script exits early if ``--out`` already exists. ``jlens`` +is imported lazily (``uv sync --extra lens``); the model is loaded through +``research.cross_lens.load_lens_model`` on ``--device`` (default ``cuda``). Paper runs (Qwen3.5-9B, layer set copied from the English-fitted Jacobian lens):: @@ -39,6 +52,17 @@ fit_ridge_lens.py --model-id Qwen/Qwen3.5-9B --prompts-json /c4_prompts_zh_seed0.json \ --n-prompts 100 --holdout 10 --layers-from /qwen35_9b_jlens_en_seed0_n100.pt \ --ckpt-dir --tag ridge_zh --out /qwen35_9b_ridgelens_zh_seed0_n100.pt + +Fixed in 0.2.1: + - ``--holdout`` must be >= 1 and leave at least one fit prompt (clear + ``SystemExit``): with 0 the holdout accumulator was empty and lambda + selection crashed after the full fit pass. + - The deep-layer top-1 agreement diagnostic sliced ``prompts[-5:]``, which + reads into the fit prompts whenever ``--holdout`` < 5; it now slices the + holdout tail, ``prompts[n_fit:][-5:]``. For the paper's ``--holdout 10`` the + two slices coincide, so the shipped ``deep_top1_agreement`` numbers stand. + - Accumulator resume is guarded by the ``.meta.json`` sidecar described + above; ``.report.json`` gains a ``provenance`` block. """ from __future__ import annotations @@ -50,6 +74,16 @@ import torch +from sparse_readout_prism.research.cross_lens import ( + check_resume_meta, + import_jlens, + load_lens_model, + prompt_list_sha1, + write_resume_meta, +) +from sparse_readout_prism.research.run_io import run_provenance +from sparse_readout_prism.utils import resolve_device + LAMBDA_GRID = [1e-3, 1e-2, 1e-1, 1.0, 10.0] @@ -57,41 +91,20 @@ def log(msg: str) -> None: print(f"[ridge_lens] {msg}", flush=True) -def _import_jlens(): - try: - import jlens - except ImportError as e: - raise ImportError( - "jlens (the Jacobian-lens reference implementation, Apache-2.0, " - "github.com/anthropics/jacobian-lens) is not installed; install it with `uv sync --extra lens`" - ) from e - return jlens - - def _activation_recorder(): - _import_jlens() + import_jlens() from jlens.hooks import ActivationRecorder return ActivationRecorder -def load_model(model_id: str): - import transformers - - jlens = _import_jlens() - hf = ( - transformers.AutoModelForCausalLM.from_pretrained( - model_id, torch_dtype=torch.bfloat16, attn_implementation="sdpa" - ) - .cuda() - .eval() - ) - tok = transformers.AutoTokenizer.from_pretrained(model_id) - return jlens.from_hf(hf, tok) +def diagnostic_prompts(prompts: list[str], n_fit: int) -> list[str]: + """The (up to) five holdout prompts the deep top-1 agreement diagnostic scores.""" + return prompts[n_fit:][-5:] @torch.no_grad() -def accumulate(model, prompts, layers, final_layer, max_seq_len, ckpt_path, every=10): +def accumulate(model, prompts, layers, final_layer, max_seq_len, ckpt_path, *, device, every=10): """One pass over prompts, accumulating XtX / XtY / YtY traces per layer.""" ActivationRecorder = _activation_recorder() @@ -105,7 +118,7 @@ def accumulate(model, prompts, layers, final_layer, max_seq_len, ckpt_path, ever acc = None start = 0 if state is not None: - acc = {l: {k: v.cuda() for k, v in per.items()} for l, per in state["acc"].items()} + acc = {l: {k: v.to(device) for k, v in per.items()} for l, per in state["acc"].items()} start = state["next_prompt"] for pi in range(start, len(prompts)): @@ -118,10 +131,10 @@ def accumulate(model, prompts, layers, final_layer, max_seq_len, ckpt_path, ever d = y.shape[1] acc = { l: { - "xtx": torch.zeros(d, d, device="cuda"), - "xty": torch.zeros(d, d, device="cuda"), - "ytr": torch.zeros((), device="cuda"), - "n": torch.zeros((), device="cuda"), + "xtx": torch.zeros(d, d, device=device), + "xty": torch.zeros(d, d, device=device), + "ytr": torch.zeros((), device=device), + "n": torch.zeros((), device=device), } for l in layers } @@ -163,13 +176,16 @@ def main(argv: list[str] | None = None) -> int: ap.add_argument("--model-id", default="Qwen/Qwen3.5-9B") ap.add_argument("--prompts-json", required=True, help="seeded prompt dump from fit_jlens.py prompts") ap.add_argument("--n-prompts", type=int, default=100) - ap.add_argument("--holdout", type=int, default=10) + ap.add_argument("--holdout", type=int, default=10, help="holdout tail for lambda selection (>= 1)") ap.add_argument("--layers-from", required=True, help="reference Jacobian lens .pt (source layers, d_model)") ap.add_argument("--max-seq-len", type=int, default=512) ap.add_argument("--ckpt-dir", required=True, help="directory for the resumable accumulator checkpoints") ap.add_argument("--out", required=True, help="output lens .pt (JacobianLens container)") ap.add_argument("--tag", default="ridge", help="checkpoint file prefix inside --ckpt-dir") + ap.add_argument("--device", default="cuda", help="torch device for the model and accumulators (paper: cuda)") args = ap.parse_args(argv) + if args.holdout < 1: + raise SystemExit(f"--holdout must be >= 1 (got {args.holdout}): lambda is selected on the holdout prompts") out = Path(args.out) if out.exists(): @@ -184,29 +200,33 @@ def main(argv: list[str] | None = None) -> int: payload = json.loads(Path(args.prompts_json).read_text()) prompts = payload["prompts"][: args.n_prompts] n_fit = len(prompts) - args.holdout + if n_fit < 1: + raise SystemExit(f"--holdout {args.holdout} leaves no fit prompts ({len(prompts)} prompts available)") log(f"{len(prompts)} prompts ({n_fit} fit + {args.holdout} holdout)") - model = load_model(args.model_id) - final_layer = model.n_layers - 1 ckpt_dir = Path(args.ckpt_dir) + fit_ckpt, hold_ckpt = ckpt_dir / f"{args.tag}_fit.ckpt", ckpt_dir / f"{args.tag}_hold.ckpt" + meta = { + "model_id": args.model_id, + "prompts_json": args.prompts_json, + "prompts_sha1": prompt_list_sha1(prompts), + "n_prompts": len(prompts), + "holdout": args.holdout, + "tag": args.tag, + "layers": layers, + "max_seq_len": args.max_seq_len, + } + meta_path = ckpt_dir / f"{args.tag}.meta.json" + check_resume_meta(meta_path, meta, [fit_ckpt, hold_ckpt]) ckpt_dir.mkdir(parents=True, exist_ok=True) + write_resume_meta(meta_path, meta) - acc_fit = accumulate( - model, - prompts[:n_fit], - layers, - final_layer, - args.max_seq_len, - ckpt_dir / f"{args.tag}_fit.ckpt", - ) - acc_hold = accumulate( - model, - prompts[n_fit:], - layers, - final_layer, - args.max_seq_len, - ckpt_dir / f"{args.tag}_hold.ckpt", - ) + device = resolve_device(args.device) + model = load_lens_model(args.model_id, device=device).lens_model + final_layer = model.n_layers - 1 + + acc_fit = accumulate(model, prompts[:n_fit], layers, final_layer, args.max_seq_len, fit_ckpt, device=device) + acc_hold = accumulate(model, prompts[n_fit:], layers, final_layer, args.max_seq_len, hold_ckpt, device=device) report, W_final = {}, {} for l in layers: @@ -227,13 +247,13 @@ def main(argv: list[str] | None = None) -> int: log(f"layer {l}: g={g} holdout_R2={score:.4f}") # Diagnostic: top-1 agreement of decoded transported states vs model - # logits on the last 5 holdout prompts, two deepest fitted layers. + # logits on the last (up to) 5 holdout prompts, two deepest fitted layers. ActivationRecorder = _activation_recorder() deep = layers[-2:] agree = {l: [0, 0] for l in deep} with torch.no_grad(): - for prompt in prompts[-5:]: + for prompt in diagnostic_prompts(prompts, n_fit): with ActivationRecorder(model.layers, at=sorted(set(deep) | {final_layer})) as rec: input_ids = model.encode(prompt, max_length=args.max_seq_len) model.forward(input_ids) @@ -269,6 +289,7 @@ def main(argv: list[str] | None = None) -> int: "holdout": args.holdout, "lambda_grid": LAMBDA_GRID, "layers": {str(l): report[l] for l in layers}, + "provenance": run_provenance(args), }, indent=1, ) diff --git a/scripts/run/run_cross_lens_readouts.py b/scripts/run/run_cross_lens_readouts.py index 18619f5..0ee775f 100644 --- a/scripts/run/run_cross_lens_readouts.py +++ b/scripts/run/run_cross_lens_readouts.py @@ -24,20 +24,29 @@ position-alignment guard checks on the first prompt that the hooked state reproduces the lens's own logits at every layer. -Decomposition basis. ``--sae`` is the k=128 seed-0 readout dictionary for +Decomposition basis. ``--checkpoint`` is the k=128 seed-0 readout dictionary for Qwen3.5-9B (8x, D=32768), Hugging Face ``hematteo/sparse-readout-prism`` file ``qwen3.5-9b/k128_8x/checkpoint.pt``, a different operating point from the main -tables' 32x/k=256 dictionaries. Rows of the LM head are centred on the vocabulary -mean and unit-normalised before encoding, and the contribution of feature ``i`` to -the score of token ``t`` is ``||W_t - mu|| * code_i(t) * (h . d_i)``. +tables' 32x/k=256 dictionaries; it is read with ``research.qwen_readout.load_sae`` +and ``--k`` defaults to the checkpoint's trained k (an explicit different value +warns). Rows of the LM head are centred and unit-normalised before encoding +(``data.center_normalize_rows``), and the contribution of feature ``i`` to the +score of token ``t`` is ``||W_t - mu|| * code_i(t) * (h . d_i)``. ``--centering`` +picks ``mu``: ``live`` (default, the paper) is the full-vocabulary mean of the +live bf16 -> fp32 head; ``trained`` is the checkpoint's stored ``row_mean``, or +the mean over the tokenizer's text-token rows when the checkpoint predates it. Bank schema: ``{"prompts": [{"id", "group", "prompt", "targets": [...], ...}]}`` (see ``data/cross_lens/README.md``). Every target string is probed by the id of -its first token. ``jlens`` is imported lazily (``uv sync --extra lens``). +its first token. ``jlens`` is imported lazily (``uv sync --extra lens``); the +model is loaded through ``research.cross_lens.load_lens_model`` on ``--device`` +(default ``cuda``). The dump keeps its ``sae`` key (the checkpoint path) so +existing readers of the schema are unaffected; ``.manifest.json`` records +the resolved ``k``, the centering mode and full provenance. Paper runs (Qwen3.5-9B, one run per lens and bank, final prompt position only):: - run_cross_lens_readouts.py --model-id Qwen/Qwen3.5-9B --lens --sae \ + run_cross_lens_readouts.py --model-id Qwen/Qwen3.5-9B --lens --checkpoint \ --k 128 --prompts data/cross_lens/cross_lens_prompts_en_zh.json --n-positions 1 \ --decompose-top1 --out /en_zh__.json @@ -55,85 +64,35 @@ import argparse import json import os +from functools import partial from pathlib import Path import torch -import torch.nn.functional as F +from sparse_readout_prism.research.cross_lens import ( + CENTERING_MODES, + centering_row_mean, + decompose_token, + feature_top_tokens, + import_jlens, + load_lens_model, + resolve_k, +) +from sparse_readout_prism.research.qwen_readout import load_sae from sparse_readout_prism.research.run_io import run_provenance -from sparse_readout_prism.utils import write_json +from sparse_readout_prism.utils import resolve_device, write_json def log(msg: str) -> None: print(f"[readouts] {msg}", flush=True) -def _import_jlens(): - try: - import jlens - except ImportError as e: - raise ImportError( - "jlens (the Jacobian-lens reference implementation, Apache-2.0, " - "github.com/anthropics/jacobian-lens) is not installed; install it with `uv sync --extra lens`" - ) from e - return jlens - - -def load_sae(checkpoint: Path): - """Raw TopK dictionary tensors from a training checkpoint (decoder rows unit-normalised).""" - ckpt = torch.load(checkpoint, map_location="cpu", weights_only=True) - state = ckpt["model_state_dict"] - decoder = state["decoder"].float().contiguous() # (d_features, d_model) - encoder_w = state["encoder.weight"].float().contiguous() # (d_features, d_model) - encoder_b = state["encoder.bias"].float().contiguous() # (d_features,) - decoder = decoder / decoder.norm(dim=1, keepdim=True).clamp_min(1e-8) - return decoder, encoder_w, encoder_b - - -def encode_topk(x, encoder_w, encoder_b, k): - acts = F.relu(x @ encoder_w.T + encoder_b) - values, indices = torch.topk(acts, k=min(k, acts.shape[-1]), dim=-1) - code = torch.zeros_like(acts) - code.scatter_(dim=-1, index=indices, src=values) - return code - - -@torch.no_grad() -def feature_top_tokens(W, row_mean, feature_ids, encoder_w, encoder_b, tokenizer, top_tokens=12, chunk=8192): - """Top unembedding rows (by encoder activation) for each feature id, as decoded strings.""" - if not feature_ids: - return {} - device = W.device - fids = torch.tensor(sorted(feature_ids), dtype=torch.long) - enc = encoder_w[fids].to(device) - bias = encoder_b[fids].to(device) - best_scores = torch.full((len(fids), top_tokens), -float("inf"), device=device) - best_ids = torch.zeros((len(fids), top_tokens), dtype=torch.long, device=device) - for start in range(0, W.shape[0], chunk): - rows = W[start : start + chunk].float() - centered = rows - row_mean - x = centered / centered.norm(dim=1, keepdim=True).clamp_min(1e-8) - scores = F.relu(x @ enc.T + bias).T - merged = torch.cat([best_scores, scores], dim=1) - ids = torch.arange(start, start + rows.shape[0], device=device).expand(len(fids), -1) - merged_ids = torch.cat([best_ids, ids], dim=1) - best_scores, keep = torch.topk(merged, k=top_tokens, dim=1) - best_ids = torch.gather(merged_ids, 1, keep) - return { - int(fid): [tokenizer.decode([t]) for t, s in zip(best_ids[i].tolist(), best_scores[i].tolist()) if s > 0] - for i, fid in enumerate(fids.tolist()) - } - - @torch.no_grad() -def run(args) -> dict: - import transformers +def run(args, *, device: torch.device, sae: tuple, k: int) -> dict: + jlens = import_jlens() - jlens = _import_jlens() - - hf = transformers.AutoModelForCausalLM.from_pretrained(args.model_id, dtype=torch.bfloat16).cuda().eval() - tok = transformers.AutoTokenizer.from_pretrained(args.model_id) - model = jlens.from_hf(hf, tok) + loaded = load_lens_model(args.model_id, device=device) + hf, tok, model = loaded.hf, loaded.tok, loaded.lens_model # Lens-free control: J = I at every layer, i.e. no transport at all. This # is the floor the cross-lens agreement has to beat: how often two surface # forms share a dominant readout feature when no lens is involved. @@ -146,12 +105,22 @@ def run(args) -> dict: else: lens = jlens.JacobianLens.load(args.lens) layers = sorted(lens.jacobians) - J = {l: lens.jacobians[l].float().cuda() for l in layers} - W = hf.get_output_embeddings().weight.detach().float().cuda() # (vocab, d_model) - row_mean = W.mean(dim=0) - final_norm = hf.model.norm - decoder, encoder_w, encoder_b = load_sae(Path(args.sae)) - decoder, encoder_w, encoder_b = decoder.cuda(), encoder_w.cuda(), encoder_b.cuda() + J = {l: lens.jacobians[l].float().to(device) for l in layers} + W = loaded.lm_head.weight.detach().float().to(device) # (vocab, d_model) + final_norm = loaded.final_norm + decoder, encoder_w, encoder_b, _config, ckpt_row_mean = sae + row_mean = centering_row_mean(W, args.centering, tok=tok, ckpt_row_mean=ckpt_row_mean) # (d_model,) + decoder, encoder_w, encoder_b = decoder.to(device), encoder_w.to(device), encoder_b.to(device) + decompose = partial( + decompose_token, + W=W, + row_mean=row_mean, + decoder=decoder, + encoder_w=encoder_w, + encoder_b=encoder_b, + k=k, + top_feats=args.top_feats, + ) spec = json.loads(Path(args.prompts).read_text())["prompts"] positions = list(range(-args.n_positions, 0)) @@ -167,29 +136,10 @@ def hook(module, inputs, output): handles = [model.layers[l].register_forward_hook(make_hook(l)) for l in layers] - def decompose(h_state, token_id): - W_row = W[token_id] - centered = W_row - row_mean - norm = centered.norm().clamp_min(1e-8) - code = encode_topk((centered / norm)[None, :], encoder_w, encoder_b, args.k)[0] - contributions = norm * code * (h_state @ decoder.T) - base = float(h_state @ row_mean) - feat_sum = float(contributions.sum()) - original = float(h_state @ W_row) - active = torch.nonzero(contributions != 0).flatten() - top = active[contributions[active].abs().argsort(descending=True)[: args.top_feats]] - return { - "original_logit": original, - "base": base, - "feature_sum": feat_sum, - "residual": original - base - feat_sum, - "top_features": [{"id": int(f), "contribution": float(contributions[f])} for f in top], - } - records = [] seen = set() for vi, v in enumerate(spec): - ids = tok(v["prompt"], return_tensors="pt").input_ids.cuda() + ids = tok(v["prompt"], return_tensors="pt").input_ids.to(device) greedy = hf.generate(ids, max_new_tokens=8, do_sample=False)[0, ids.shape[1] :] continuation = tok.decode(greedy) # First token id of each target string (multi-token targets probed by head). @@ -218,7 +168,7 @@ def decompose(h_state, token_id): "layers": {}, } for l in layers: - ll = per_layer[l].float().cuda() # (n_positions, vocab) + ll = per_layer[l].float().to(device) # (n_positions, vocab) lrec = {} for pi, pos in enumerate(positions): src = captured[l][0, pos].float() @@ -259,12 +209,12 @@ def decompose(h_state, token_id): for h in handles: h.remove() log(f"labeling {len(seen)} unique features") - labels = feature_top_tokens(W, row_mean, seen, encoder_w, encoder_b, tok) + labels = feature_top_tokens(W, row_mean, seen, encoder_w, encoder_b, tok, top_tokens=12) return { "model": args.model_id, "lens": args.lens, - "sae": args.sae, + "sae": args.checkpoint, "positions": positions, "layers": layers, "records": records, @@ -287,18 +237,35 @@ def main(argv: list[str] | None = None) -> int: help="required with --lens identity: a fitted lens .pt whose layer set the control copies, " "so the comparison is layer-matched", ) - p.add_argument("--sae", required=True, help="readout SAE checkpoint.pt (paper: Qwen3.5-9B k=128 seed 0, 8x)") - p.add_argument("--k", type=int, default=128, help="active features per row code") + p.add_argument( + "--checkpoint", required=True, help="readout dictionary checkpoint.pt (paper: Qwen3.5-9B k=128 seed 0, 8x)" + ) + p.add_argument( + "--k", + type=int, + default=None, + help="active features per row code (default: the checkpoint's trained k; a different value warns)", + ) + p.add_argument( + "--centering", + choices=CENTERING_MODES, + default="live", + help="centering mean: live = full-vocabulary mean of the live head (paper), " + "trained = the checkpoint's row_mean (else the tokenizer's text-token mean)", + ) p.add_argument("--prompts", required=True, help="prompt bank JSON ({'prompts': [...]})") p.add_argument("--n-positions", type=int, default=4, help="score the last n prompt positions") p.add_argument("--decompose-top1", action="store_true", help="also decompose the lens's own top-1 token") p.add_argument("--top-feats", type=int, default=10, help="signed feature contributions kept per decomposition") + p.add_argument("--device", default="cuda", help="torch device for the model and the dictionary (paper: cuda)") p.add_argument("--out", required=True, help="output dump JSON") args = p.parse_args(argv) if args.lens == "identity" and not args.layers_from: p.error("--lens identity requires --layers-from ") - payload = run(args) + sae = load_sae(args.checkpoint) + k = resolve_k(args.k, sae[3]) + payload = run(args, device=resolve_device(args.device), sae=sae, k=k) out = Path(args.out) out.parent.mkdir(parents=True, exist_ok=True) @@ -310,6 +277,8 @@ def main(argv: list[str] | None = None) -> int: "n_records": len(payload["records"]), "layers": payload["layers"], "positions": payload["positions"], + "k": k, + "centering": args.centering, "provenance": run_provenance(args), }, out.with_suffix(".manifest.json"), diff --git a/scripts/run/run_qwen_profanity_suppression_eval.py b/scripts/run/run_qwen_profanity_suppression_eval.py index a4c069f..db70ebf 100644 --- a/scripts/run/run_qwen_profanity_suppression_eval.py +++ b/scripts/run/run_qwen_profanity_suppression_eval.py @@ -31,7 +31,8 @@ (``feature_suppression``), discovery / oracle token bias, norm-matched random features, and three label-free category-direction controls built from the five discovery rows of W_U (``mean_row_direction``, ``pca_group_direction`` = PCA -rank-1, ``pca_group_rank4`` = the full rank of the discovery rows). The +rank-1, ``pca_group_rank4`` = the top-4 principal components of the five +discovery rows about the global row mean, i.e. rank 4 of at most 5). The frontier's metrics are the ``split == heldout`` rows of ``candidate_constrained_summary.csv`` (``median_kl_bits`` against ``mean_bad_prob_reduction`` / ``candidate_flip_rate`` per method and scale); diff --git a/scripts/run/run_readout_baseline_comparisons.py b/scripts/run/run_readout_baseline_comparisons.py index ae6b6dd..3a80424 100644 --- a/scripts/run/run_readout_baseline_comparisons.py +++ b/scripts/run/run_readout_baseline_comparisons.py @@ -56,6 +56,18 @@ baseline_feature_compactness.csv manifest.json +Fixed in 0.2.1: ``margin_from_rows`` scatters per-row-support contributions +(``nearest_row_ridge_top*``, ``knn_basis_top*``, ``row_cluster_*``: one entry +per neighbour / centroid of THAT row) into a full-size vector over the method's +feature space (``BaselineDecomposition.feature_space_size``: vocabulary rows or +centroids) before forming ``fA - fB``. The margin scalars (exact / sparse / +residual, hence rho, sign agreement, accepted rate and every fidelity column +the paper reports) never depended on the alignment; only the coverage and +compactness columns of those three method families change (``top5_*_cov``, +``top10_*_cov``, ``n_feat_80pct_abs``, ``largest_*``, and their rows of +``baseline_feature_compactness.csv``), which used to subtract unrelated +neighbours position by position. + Designed for one-shot resumable execution on a single A40: per-cell `done.json` sentinel means restart after preemption only re-runs the unfinished (bank, method) cells. Hidden states are recomputed per cell so each cell is @@ -89,6 +101,7 @@ write_rows_csv as _write_csv, ) from sparse_readout_prism.research.registry import resolve_registry +from sparse_readout_prism.research.row_geometry import spherical_kmeans_unit from sparse_readout_prism.utils import ( atomic_write_text as _atomic_write, find_lm_head, @@ -98,11 +111,6 @@ write_jsonl as _write_jsonl, ) -# Repo root from this file's location: scripts/run/THIS_FILE.py -# parents[0]=run/ [1]=scripts/ [2]=repo root -REPO = Path(__file__).resolve().parents[2] - - # =========================================================================== # # Baseline / reference methods (inlined from the former # research/run/readout_baselines.py — single consumer, so kept self-contained). @@ -127,7 +135,14 @@ class BaselineDecomposition: length depends on the method: * sparse_rp / shuffled_row_code / random_support: d_features (mostly zero) * pca_: n_components - * nearest_row_ridge_top: k + * nearest_row_ridge_top / knn_basis_top: k, one per neighbour row of THIS row + * row_cluster_d_k / row_cluster_hard_d: k (1 when hard), one per selected centroid + + `feature_space_size` is the size of the space `active_feature_indices` + index into (d_features, n_components, the vocabulary size V, or the D + centroids). When `feature_contributions` is shorter than it, the vector is + a per-row support and must be scattered through `active_feature_indices` + before it is compared across rows (`margin_from_rows` does this). """ original_logit: torch.Tensor @@ -139,6 +154,7 @@ class BaselineDecomposition: identity_error: torch.Tensor active_feature_indices: torch.Tensor method: str + feature_space_size: int # --------------------------------------------------------------------------- # @@ -187,6 +203,43 @@ def decompose_row( ) -> BaselineDecomposition: raise NotImplementedError + # ---- shared accounting tail ---------------------------------------------- + + def _finish( + self, + h: torch.Tensor, + row_idx: int, + decoded: torch.Tensor, # (d_model,) reconstruction of the normalised row + contribs: torch.Tensor, # (n_active,) per-atom contributions; their sum is feature_sum + feature_contributions: torch.Tensor, # the vector reported to callers (full-size or per-row support) + active: torch.Tensor, # (n_active,) indices into the method's feature space + feature_space_size: int, + ) -> BaselineDecomposition: + """Reconstructed row, residual and the exact identity + ``original == base + feature_sum + residual`` -- the same for every method.""" + W_row = self.W[row_idx].to(self.device) + row_norm = self.row_norm_all[row_idx] + reconstructed_row = self.row_mean + row_norm * decoded + residual_row = W_row - reconstructed_row + base_term = h @ self.row_mean + feature_sum = contribs.sum() + residual_term = h @ residual_row + original_logit = h @ W_row + reconstructed_logit = base_term + feature_sum + identity_error = original_logit - (base_term + feature_sum + residual_term) + return BaselineDecomposition( + original_logit=original_logit, + base_term=base_term, + feature_contributions=feature_contributions, + feature_sum=feature_sum, + residual_term=residual_term, + reconstructed_logit=reconstructed_logit, + identity_error=identity_error, + active_feature_indices=active.detach().cpu(), + method=self.name, + feature_space_size=int(feature_space_size), + ) + # --------------------------------------------------------------------------- # # A. Sparse Readout Prism anchor @@ -230,6 +283,7 @@ def decompose_row( identity_error=d.identity_error, active_feature_indices=d.active_feature_indices, method=self.name, + feature_space_size=self.sae.d_features, ) @@ -294,34 +348,13 @@ def decompose_row( val = self.code_values_per_row[src] # (k,) decoder_rows = self.sae.decoder[idx] # (k, d_model) decoded = (val[:, None] * decoder_rows).sum(0) # (d_model,) - W_row = self.W[row_idx].to(self.device) row_norm = self.row_norm_all[row_idx] - reconstructed_row = self.row_mean + row_norm * decoded - residual_row = W_row - reconstructed_row - d_features = self.sae.d_features feature_contributions = torch.zeros(d_features, device=self.device) h_dec = h @ decoder_rows.T # (k,) contribs_k = row_norm * val * h_dec # (k,) feature_contributions.scatter_add_(0, idx, contribs_k) - - base_term = h @ self.row_mean - feature_sum = contribs_k.sum() - residual_term = h @ residual_row - original_logit = h @ W_row - reconstructed_logit = base_term + feature_sum - identity_error = original_logit - (base_term + feature_sum + residual_term) - return BaselineDecomposition( - original_logit=original_logit, - base_term=base_term, - feature_contributions=feature_contributions, - feature_sum=feature_sum, - residual_term=residual_term, - reconstructed_logit=reconstructed_logit, - identity_error=identity_error, - active_feature_indices=idx.detach().cpu(), - method=self.name, - ) + return self._finish(h, row_idx, decoded, contribs_k, feature_contributions, idx, d_features) # --------------------------------------------------------------------------- # @@ -375,33 +408,12 @@ def decompose_row( val = self.code_vals_sorted[row_idx] # (k,) signed magnitudes decoder_rows = self.sae.decoder[rand_ids] # (k, d_model) decoded = (val[:, None] * decoder_rows).sum(0) # (d_model,) - W_row = self.W[row_idx].to(self.device) row_norm = self.row_norm_all[row_idx] - reconstructed_row = self.row_mean + row_norm * decoded - residual_row = W_row - reconstructed_row - feature_contributions = torch.zeros(self.d_features, device=self.device) h_dec = h @ decoder_rows.T contribs_k = row_norm * val * h_dec feature_contributions.scatter_add_(0, rand_ids, contribs_k) - - base_term = h @ self.row_mean - feature_sum = contribs_k.sum() - residual_term = h @ residual_row - original_logit = h @ W_row - reconstructed_logit = base_term + feature_sum - identity_error = original_logit - (base_term + feature_sum + residual_term) - return BaselineDecomposition( - original_logit=original_logit, - base_term=base_term, - feature_contributions=feature_contributions, - feature_sum=feature_sum, - residual_term=residual_term, - reconstructed_logit=reconstructed_logit, - identity_error=identity_error, - active_feature_indices=rand_ids.detach().cpu(), - method=self.name, - ) + return self._finish(h, row_idx, decoded, contribs_k, feature_contributions, rand_ids, self.d_features) # --------------------------------------------------------------------------- # @@ -445,30 +457,12 @@ def decompose_row( x = self.row_normalized_all[row_idx] # (d_model,) coeffs = self.V_pca @ x # (n_comp,) decoded = coeffs @ self.V_pca # (d_model,) - W_row = self.W[row_idx].to(self.device) row_norm = self.row_norm_all[row_idx] - reconstructed_row = self.row_mean + row_norm * decoded - residual_row = W_row - reconstructed_row - h_pca = self.V_pca @ h # (n_comp,) feature_contributions = row_norm * coeffs * h_pca # (n_comp,) - base_term = h @ self.row_mean - feature_sum = feature_contributions.sum() - residual_term = h @ residual_row - original_logit = h @ W_row - reconstructed_logit = base_term + feature_sum - identity_error = original_logit - (base_term + feature_sum + residual_term) - active = torch.arange(self.n_components) # all components active - return BaselineDecomposition( - original_logit=original_logit, - base_term=base_term, - feature_contributions=feature_contributions, - feature_sum=feature_sum, - residual_term=residual_term, - reconstructed_logit=reconstructed_logit, - identity_error=identity_error, - active_feature_indices=active, - method=self.name, + active = torch.arange(self.n_components) # all components active; the basis is shared by every row + return self._finish( + h, row_idx, decoded, feature_contributions, feature_contributions, active, self.n_components ) @@ -522,32 +516,13 @@ def decompose_row( A = XXT + self.lam * torch.eye(self.top_k, device=self.device) beta = torch.linalg.solve(A, Xx) # (top_k,) decoded = beta @ X # (d_model,) - W_row = self.W[row_idx].to(self.device) row_norm = self.row_norm_all[row_idx] - reconstructed_row = self.row_mean + row_norm * decoded - residual_row = W_row - reconstructed_row - # Per-neighbour contribution to h. The neighbour basis is normalized; - # contribution_j = beta_j * row_norm * (h . X_j). + # contribution_j = beta_j * row_norm * (h . X_j). Indexed by this row's + # neighbour list; the feature space is the vocabulary (V rows). h_nbr = X @ h # (top_k,) feature_contributions = row_norm * beta * h_nbr # (top_k,) - base_term = h @ self.row_mean - feature_sum = feature_contributions.sum() - residual_term = h @ residual_row - original_logit = h @ W_row - reconstructed_logit = base_term + feature_sum - identity_error = original_logit - (base_term + feature_sum + residual_term) - return BaselineDecomposition( - original_logit=original_logit, - base_term=base_term, - feature_contributions=feature_contributions, - feature_sum=feature_sum, - residual_term=residual_term, - reconstructed_logit=reconstructed_logit, - identity_error=identity_error, - active_feature_indices=nbr_ids.detach().cpu(), - method=self.name, - ) + return self._finish(h, row_idx, decoded, feature_contributions, feature_contributions, nbr_ids, self.W.shape[0]) # --------------------------------------------------------------------------- # @@ -599,29 +574,10 @@ def decompose_row( gamma = (x @ mean_nbr) / (mean_nbr @ mean_nbr).clamp_min(1e-8) coeffs = gamma * w # per-neighbour coefficients; decoded = coeffs @ X decoded = coeffs @ X # (d_model,) rescaled neighbourhood estimate - W_row = self.W[row_idx].to(self.device) row_norm = self.row_norm_all[row_idx] - reconstructed_row = self.row_mean + row_norm * decoded - residual_row = W_row - reconstructed_row h_nbr = X @ h # (top_k,) - feature_contributions = row_norm * coeffs * h_nbr # (top_k,) - base_term = h @ self.row_mean - feature_sum = feature_contributions.sum() - residual_term = h @ residual_row - original_logit = h @ W_row - reconstructed_logit = base_term + feature_sum - identity_error = original_logit - (base_term + feature_sum + residual_term) - return BaselineDecomposition( - original_logit=original_logit, - base_term=base_term, - feature_contributions=feature_contributions, - feature_sum=feature_sum, - residual_term=residual_term, - reconstructed_logit=reconstructed_logit, - identity_error=identity_error, - active_feature_indices=nbr_ids.detach().cpu(), - method=self.name, - ) + feature_contributions = row_norm * coeffs * h_nbr # (top_k,) indexed by this row's neighbour list + return self._finish(h, row_idx, decoded, feature_contributions, feature_contributions, nbr_ids, self.W.shape[0]) # --------------------------------------------------------------------------- # @@ -629,32 +585,6 @@ def decompose_row( # --------------------------------------------------------------------------- # -def _torch_kmeans_unit(X: torch.Tensor, n_clusters: int, seed: int, iters: int = 12, chunk: int = 4096) -> torch.Tensor: - """Lloyd's k-means on unit-norm rows (cosine assignment), on-device. - Returns unit-normalized centroids (n_clusters, d). Dead centroids are - reseeded from random rows each iteration. Deterministic given `seed`.""" - V = X.shape[0] - g = torch.Generator().manual_seed(seed) - C = X[torch.randperm(V, generator=g)[:n_clusters].to(X.device)].clone() - ones = torch.ones(V, device=X.device) - for _ in range(iters): - assign = torch.empty(V, dtype=torch.long, device=X.device) - for s in range(0, V, chunk): - assign[s : s + chunk] = (X[s : s + chunk] @ C.T).argmax(dim=1) - C_new = torch.zeros_like(C) - count = torch.zeros(n_clusters, device=X.device) - C_new.index_add_(0, assign, X) - count.index_add_(0, assign, ones) - dead = count == 0 - C = C_new / count.clamp_min(1.0)[:, None] - n_dead = int(dead.sum()) - if n_dead: - ridx = torch.randperm(V, generator=g)[:n_dead].to(X.device) - C[dead] = X[ridx] - C = C / C.norm(dim=1, keepdim=True).clamp_min(1e-8) - return C - - class RowClusterMethod(BaselineMethod): """Sparsity-matched k-means centroid dictionary: fit `n_clusters` unit centroids on the centered/unit rows, then code each row over its @@ -664,6 +594,15 @@ class RowClusterMethod(BaselineMethod): literal clustering reading: the row's own single cluster centroid reconstructs it. + `exclude_ids` is not applied: the k-means fit and the row's own cluster + include the target row and its contrast partner. This is the paper's + behaviour (the centroid dictionary is fit once on all rows, like the SAE), + unlike the nearest-row methods, which drop the target and its partner from + the neighbour search. The fit (`research.row_geometry.spherical_kmeans_unit`) + is seeded and bit-exact across runs on CPU only; a run memoises it by + `(n_clusters, seed)` so the soft and hard variants of one width share one + set of centroids. + Design note: an earlier version of this baseline solved a full least-squares projection over n_clusters = d_model centroids, a full-rank reprojection of the row space that reconstructs any row near-perfectly by @@ -672,13 +611,28 @@ class RowClusterMethod(BaselineMethod): name = "row_cluster" - def __init__(self, *, n_clusters: int, code_k: int = 256, hard: bool = False, seed: int = 0, **kw) -> None: + def __init__( + self, + *, + n_clusters: int, + code_k: int = 256, + hard: bool = False, + seed: int = 0, + kmeans_cache: Optional[dict] = None, + **kw, + ) -> None: super().__init__(**kw) self.n_clusters = int(n_clusters) self.hard = bool(hard) self.code_k = 1 if hard else min(int(code_k), self.n_clusters) self.name = f"row_cluster_hard_d{self.n_clusters}" if hard else f"row_cluster_d{self.n_clusters}_k{self.code_k}" - self.C = _torch_kmeans_unit(self.row_normalized_all, self.n_clusters, seed) + key = (self.n_clusters, int(seed)) + if kmeans_cache is not None and key in kmeans_cache: + self.C = kmeans_cache[key] + else: + self.C = spherical_kmeans_unit(self.row_normalized_all, self.n_clusters, int(seed)) + if kmeans_cache is not None: + kmeans_cache[key] = self.C self._eye = torch.eye(self.code_k, device=self.device) @torch.no_grad() @@ -700,29 +654,10 @@ def decompose_row( A = Ck @ Ck.T + 1e-4 * self._eye coeffs = torch.linalg.solve(A, Ck @ x) # (code_k,) decoded = coeffs @ Ck # (d_model,) - W_row = self.W[row_idx].to(self.device) row_norm = self.row_norm_all[row_idx] - reconstructed_row = self.row_mean + row_norm * decoded - residual_row = W_row - reconstructed_row h_c = Ck @ h # (code_k,) - feature_contributions = row_norm * coeffs * h_c # (code_k,) - base_term = h @ self.row_mean - feature_sum = feature_contributions.sum() - residual_term = h @ residual_row - original_logit = h @ W_row - reconstructed_logit = base_term + feature_sum - identity_error = original_logit - (base_term + feature_sum + residual_term) - return BaselineDecomposition( - original_logit=original_logit, - base_term=base_term, - feature_contributions=feature_contributions, - feature_sum=feature_sum, - residual_term=residual_term, - reconstructed_logit=reconstructed_logit, - identity_error=identity_error, - active_feature_indices=sel.detach().cpu(), - method=self.name, - ) + feature_contributions = row_norm * coeffs * h_c # (code_k,) indexed by this row's selected centroids + return self._finish(h, row_idx, decoded, feature_contributions, feature_contributions, sel, self.n_clusters) # --------------------------------------------------------------------------- # @@ -741,6 +676,7 @@ def build_method( seed: int = 0, row_norm_all: Optional[torch.Tensor] = None, row_normalized_all: Optional[torch.Tensor] = None, + kmeans_cache: Optional[dict] = None, ) -> BaselineMethod: """spec examples: 'sparse_rp' @@ -751,6 +687,9 @@ def build_method( 'knn_basis_top128' 'row_cluster_d65536_k256', 'row_cluster_d16384_k256' 'row_cluster_hard_d65536' + + `kmeans_cache` (a dict owned by the run) memoises fitted centroids by + (n_clusters, seed) across the row_cluster_* specs. """ common = dict( W=W, @@ -782,11 +721,13 @@ def build_method( return KNNBasisMethod(top_k=top, **common) if spec.startswith("row_cluster_hard_d"): n = int(spec.split("row_cluster_hard_d", 1)[1]) - return RowClusterMethod(n_clusters=n, hard=True, seed=seed, **common) + return RowClusterMethod(n_clusters=n, hard=True, seed=seed, kmeans_cache=kmeans_cache, **common) if spec.startswith("row_cluster_d"): body = spec.split("row_cluster_d", 1)[1] # "_k" d_str, k_str = body.split("_k", 1) - return RowClusterMethod(n_clusters=int(d_str), code_k=int(k_str), seed=seed, **common) + return RowClusterMethod( + n_clusters=int(d_str), code_k=int(k_str), seed=seed, kmeans_cache=kmeans_cache, **common + ) raise ValueError(f"unknown method spec: {spec}") @@ -795,6 +736,20 @@ def build_method( # --------------------------------------------------------------------------- # +def _full_contributions(d: BaselineDecomposition) -> torch.Tensor: + """`feature_contributions` over the method's whole feature space. + + Per-row-support methods report one entry per neighbour / centroid of that + row; scatter them through `active_feature_indices` so two rows' vectors + line up feature by feature before they are subtracted or averaged.""" + fc = d.feature_contributions + if fc.numel() == d.feature_space_size: + return fc + full = torch.zeros(d.feature_space_size, device=fc.device, dtype=fc.dtype) + full.scatter_add_(0, d.active_feature_indices.to(fc.device), fc) + return full + + def margin_from_rows( h: torch.Tensor, method: BaselineMethod, @@ -809,6 +764,11 @@ def margin_from_rows( query (A=top1, B=row_mean). In that case `b_ids` is None and the B decompose is replaced by `(h @ row_mean, h @ row_mean, 0, zeros)` so the base cancels exactly (the prism-side `s_approx_B = base + sum_i 0`). + + `feat_margin` is formed over the method's full feature space + (`_full_contributions`), so for the per-row-support methods it has one + entry per vocabulary row / centroid and the coverage statistics see the + union of A's and B's supports. """ @torch.no_grad() @@ -830,7 +790,7 @@ def _agg(ids: Optional[list[int]], pseudo: Optional[torch.Tensor], excl): d.original_logit, d.reconstructed_logit, d.residual_term, - d.feature_contributions, + _full_contributions(d), ) if accs is None: accs = cur @@ -1027,6 +987,7 @@ def run_cell( seed=args.seed, row_norm_all=cache["row_norm_all"], row_normalized_all=cache["row_normalized_all"], + kmeans_cache=cache.setdefault("kmeans_centroids", {}), ) log(f"built method {method_spec} in {time.time() - t0:.1f}s") method = methods_cache[method_spec] diff --git a/scripts/run/run_wsd_feature_alignment.py b/scripts/run/run_wsd_feature_alignment.py index dc445cf..f601e47 100644 --- a/scripts/run/run_wsd_feature_alignment.py +++ b/scripts/run/run_wsd_feature_alignment.py @@ -21,24 +21,22 @@ ``audit.json`` (contexts kept and single-token coverage per vocabulary: 20/20 words on the Qwen3.5 vocabularies, 13/20 on DeepSeek-R1-Distill-Llama-8B), ``scoring_summary.json`` (contexts scored, median target rank, row-gate -coverage) and ``metrics.json`` (full-vector centroid references computed on -the same bundle, with paired per-word bootstrap differences). - -Datasets --------- -* CoarseWSD-20 (Loureiro, Rezaee, Pilehvar and Camacho-Collados, 2021, - Computational Linguistics 47(2), doi:10.1162/coli_a_00405) -- the dataset - used in the paper. ``--data-root`` is a checkout of - https://github.com/danlou/bert-disambiguation; the loader looks for - ``/data/CoarseWSD-20//{train,test}.{data,gold}.txt`` plus - ``class_map.txt`` (```` or ``/CoarseWSD-20`` also work). The - shipped train/test split is kept. The target occurrence is replaced by - ``[BLANK]`` and the sentence is wrapped in ``COARSEWSD_TEMPLATE``; the - target's leading-space token is scored at the prompt's final position, so - it is a genuine next-token readout even when the disambiguating evidence - originally followed the target. -* AmbiStory (SemEval-2026 Task 5): graded sense plausibility against - gloss-prompt anchors. Served by the same scoring path; not used in the paper. +coverage, prompts truncated) and ``metrics.json`` (full-vector centroid +references computed on the same bundle, with paired per-word bootstrap +differences). + +Dataset +------- +CoarseWSD-20 (Loureiro, Rezaee, Pilehvar and Camacho-Collados, 2021, +Computational Linguistics 47(2), doi:10.1162/coli_a_00405). ``--data-root`` +is a checkout of https://github.com/danlou/bert-disambiguation; the loader +looks for ``/data/CoarseWSD-20//{train,test}.{data,gold}.txt`` +plus ``class_map.txt`` (```` or ``/CoarseWSD-20`` also work). The +shipped train/test split is kept. The target occurrence is replaced by +``[BLANK]`` and the sentence is wrapped in ``COARSEWSD_TEMPLATE``; the +target's leading-space token is scored at the prompt's final position, so it +is a genuine next-token readout even when the disambiguating evidence +originally followed the target. Paper runs (one per model, ``--seed 0``; the checkpoints are the 32x / k=256 selected dictionaries for Qwen3.5-2B, Qwen3.5-9B and @@ -52,15 +50,22 @@ --splits train,test --batch-size 8 --n-boot 5000 --seed 0 ``--analyze-only --out-dir --n-boot 5000 --seed 0`` recomputes -``metrics.json`` from a frozen bundle without model inference. +``metrics.json`` from a frozen bundle without model inference; it writes +``analysis_config.json`` and leaves the scoring run's ``run_config.json`` +untouched. Gold labels are never used to fit the dictionary or to select features. -``--position-mode occurrence`` (score the sentence prefix so the target -occurrence itself is the next token) and ``--anchors`` (a frozen sense-anchor -token set, shipped as ``data/wsd/wsd_sense_anchors.json``, whose codes are -embedded in the bundle for an offline feature-typing analysis not included in -this repository) are variants the paper does not report; the paper's runs use -the cloze position and pass no anchors. + +``--centering live`` (default, the paper's runs) centres unembedding rows on +the full-vocabulary mean of the live model's LM head; ``--centering trained`` +uses the dictionary's stored training mean (``row_mean`` in the checkpoint), +falling back to the mean over the tokenizer's text-token rows. + +Fixed in 0.2.1: prompts longer than ``--max-length`` are truncated from the +left, so the cloze cue at the end of the prompt survives (they were truncated +from the right; ``scoring_summary.json`` now records how many prompts were +truncated). The AmbiStory dataset mode, ``--position-mode occurrence`` and the +``--anchors`` bundle block, none of which the paper reports, were removed. """ from __future__ import annotations @@ -69,33 +74,30 @@ import ast import hashlib import itertools -import json import platform -import re import sys import time from collections import defaultdict from pathlib import Path -from typing import Any, Iterable +from typing import Any import numpy as np import torch +from sparse_readout_prism.data import center_normalize_rows, centering_mean, token_mask_from_tokenizer +from sparse_readout_prism.research.qwen_readout import encode_topk, load_sae from sparse_readout_prism.research.run_io import run_provenance -from sparse_readout_prism.utils import find_lm_head, load_causal_lm, set_seed - - -AMBISTORY_STORY_TEMPLATE = """Read the story and recover the missing word. - -{story} - -The missing word is:""" - -AMBISTORY_GLOSS_TEMPLATE = """Read the meaning and name the matching word. +from sparse_readout_prism.research.wsd import ( + bootstrap_mean, + l2_normalize, + load_bundle, + percentile_ci, + safe_spearman, + shuffled_srp, + word_splits, +) +from sparse_readout_prism.utils import find_lm_head, load_causal_lm, set_seed, write_json, write_jsonl -Meaning: {gloss} - -The matching word is:""" COARSEWSD_TEMPLATE = """Read the sentence and recover the missing word. @@ -103,6 +105,9 @@ The missing word is:""" +METRIC_NAMES = ("accuracy", "balanced_accuracy", "macro_f1", "pairwise_auc", "ari", "nmi") +COMPARISON_BASELINES = ("shuffled_srp", "support_projection", "hidden") + def log(message: str) -> None: print(f"[wsd {time.strftime('%H:%M:%S')}] {message}", flush=True) @@ -113,129 +118,6 @@ def stable_id(*parts: str, n: int = 16) -> str: return hashlib.sha1(text.encode("utf-8")).hexdigest()[:n] -def write_json(path: Path, value: Any) -> None: - path.write_text(json.dumps(value, indent=2, ensure_ascii=False) + "\n") - - -def write_jsonl(path: Path, rows: Iterable[dict[str, Any]]) -> None: - with path.open("w") as handle: - for row in rows: - handle.write(json.dumps(row, ensure_ascii=False) + "\n") - - -def mask_exact_target(text: str, target: str) -> str | None: - """Replace the first case-insensitive whole target occurrence.""" - pattern = re.compile(rf"(? str | None: - masked = mask_exact_target(sentence.strip(), target) - if masked is None: - return None - pieces = [precontext.strip(), masked] - if ending.strip(): - pieces.append(ending.strip()) - return "\n".join(p for p in pieces if p) - - -def load_ambistory(root: Path, splits: list[str], limit: int | None) -> tuple[list[dict], dict]: - records: list[dict] = [] - anchors: dict[tuple[str, str], dict] = {} - skipped = defaultdict(int) - split_counts: dict[str, int] = {} - - for split in splits: - path = root / f"{split}.json" - if not path.exists(): - raise FileNotFoundError(path) - raw = json.loads(path.read_text()) - rows = list(raw.values()) if isinstance(raw, dict) else list(raw) - grouped: dict[tuple[str, str, str, str], list[dict]] = defaultdict(list) - for row in rows: - key = ( - str(row.get("homonym", "")).strip(), - str(row.get("precontext", "")).strip(), - str(row.get("sentence", "")).strip(), - str(row.get("ending", "")).strip(), - ) - grouped[key].append(row) - - made = 0 - for (target, precontext, sentence, ending), sense_rows in sorted(grouped.items()): - if len(sense_rows) != 2: - skipped["not_two_senses"] += 1 - continue - story = make_story(precontext, sentence, ending, target) - if story is None: - skipped["target_not_found"] += 1 - continue - ordered = sorted(sense_rows, key=lambda r: str(r.get("judged_meaning", ""))) - glosses = [str(r.get("judged_meaning", "")).strip() for r in ordered] - if not all(glosses) or glosses[0] == glosses[1]: - skipped["bad_gloss_pair"] += 1 - continue - averages = [float(r["average"]) for r in ordered] - stdevs = [float(r.get("stdev", 0.0)) for r in ordered] - setup_id = stable_id(target, precontext, sentence) - context_id = stable_id(target, precontext, sentence, ending) - anchor_keys = [] - for gloss in glosses: - akey = stable_id(target, gloss) - anchor_keys.append(akey) - anchors[(target, gloss)] = { - "item_id": f"anchor:{akey}", - "kind": "anchor", - "dataset": "ambistory", - "split": split, - "target": target, - "gloss": gloss, - "anchor_key": akey, - "prompt": AMBISTORY_GLOSS_TEMPLATE.format(gloss=gloss), - } - records.append( - { - "item_id": f"story:{split}:{context_id}", - "kind": "story", - "dataset": "ambistory", - "split": split, - "target": target, - "setup_id": setup_id, - "context_id": context_id, - "has_ending": bool(ending.strip()), - "glosses": glosses, - "anchor_keys": anchor_keys, - "ratings": averages, - "rating_stdevs": stdevs, - "human_margin": averages[0] - averages[1], - "prompt": AMBISTORY_STORY_TEMPLATE.format(story=story), - } - ) - made += 1 - if limit is not None and len(records) >= limit: - break - split_counts[split] = made - if limit is not None and len(records) >= limit: - break - - # Only score anchors actually referenced by retained story records. - used = {key for row in records for key in row["anchor_keys"]} - anchor_rows = [row for row in anchors.values() if row["anchor_key"] in used] - items = anchor_rows + records - audit = { - "dataset": "ambistory", - "splits": splits, - "n_story_contexts": len(records), - "n_anchors": len(anchor_rows), - "n_setups": len({r["setup_id"] for r in records}), - "n_targets": len({r["target"] for r in records}), - "split_story_counts": split_counts, - "skipped": dict(skipped), - } - return items, audit - - def parse_class_map(path: Path) -> dict[str, Any]: text = path.read_text().strip() try: @@ -274,7 +156,6 @@ def load_coarsewsd( splits: list[str], limit: int | None, limit_per_word: int | None, - position_mode: str = "cloze", ) -> tuple[list[dict], dict]: data_root = locate_coarse_root(root) records: list[dict] = [] @@ -310,18 +191,9 @@ def load_coarsewsd( if pos is None: skipped["target_position_unresolved"] += 1 continue - if position_mode == "occurrence": - # Score the real occurrence: the prompt is the sentence - # prefix, so the target's leading-space token is the - # genuine next token at the scored position. - if pos < 1: - skipped["no_left_context"] += 1 - continue - prompt = " ".join(tokens[:pos]) - else: - tokens[pos] = "[BLANK]" - sentence = " ".join(tokens) - prompt = COARSEWSD_TEMPLATE.format(sentence=sentence) + tokens[pos] = "[BLANK]" + sentence = " ".join(tokens) + prompt = COARSEWSD_TEMPLATE.format(sentence=sentence) sense = label.strip() candidates.append( { @@ -375,42 +247,46 @@ def load_coarsewsd( return records, audit -def load_sae(path: Path) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, int]: - ckpt = torch.load(path, map_location="cpu", weights_only=False) - if "W_dec" in ckpt: - w_dec = ckpt["W_dec"].float() - w_enc = ckpt["W_enc"].float() - b_enc = ckpt.get("b_enc") - k = int(ckpt.get("k") or ckpt.get("config", {}).get("k") or 256) - else: - state = ckpt.get("model_state_dict") or ckpt.get("state_dict") - cfg = ckpt.get("factorizer") or ckpt.get("config", {}).get("factorizer", {}) - if state is None: - raise KeyError(f"unrecognized SAE checkpoint: {path}") - w_dec = state["decoder"].float() - w_enc = state["encoder.weight"].float().T.contiguous() - b_enc = state.get("encoder.bias") - k = int(cfg.get("k", 256)) - if w_enc.shape == w_dec.shape: - w_enc = w_enc.T.contiguous() - if w_enc.shape != (w_dec.shape[1], w_dec.shape[0]): - raise ValueError(f"encoder {tuple(w_enc.shape)} incompatible with decoder {tuple(w_dec.shape)}") - if b_enc is None: - b_enc = torch.zeros(w_dec.shape[0]) - return w_dec, w_enc, b_enc.float(), k +def configure_tokenizer(tokenizer) -> None: + """Left padding keeps the scored position last in a batch; left truncation + keeps the cloze cue (the prompt's tail) when a prompt exceeds ``--max-length``.""" + if tokenizer.pad_token_id is None: + tokenizer.pad_token = tokenizer.eos_token + tokenizer.padding_side = "left" + tokenizer.truncation_side = "left" def load_model(model_id: str, device: str, dtype: str): - """Frozen causal LM + tokenizer; left padding keeps the scored position last in the batch.""" + """Frozen causal LM + tokenizer configured by ``configure_tokenizer``.""" torch_dtype = torch.bfloat16 if dtype == "bfloat16" else torch.float32 model, tokenizer = load_causal_lm(model_id, dtype=torch_dtype, device_map=None) model.to(device) - if tokenizer.pad_token_id is None: - tokenizer.pad_token = tokenizer.eos_token - tokenizer.padding_side = "left" + configure_tokenizer(tokenizer) return model, tokenizer +def checkpoint_k(path: Path, sae_config: dict[str, Any]) -> int: + """The dictionary's k: ``config.factorizer.k`` when ``load_sae`` surfaced a + config, else the runner checkpoint's top-level ``factorizer.k`` (read with + ``mmap`` so the weights are not loaded twice); 256 when neither names it.""" + cfg = sae_config.get("factorizer") or {} + if "k" not in cfg: + ckpt = torch.load(path, map_location="cpu", weights_only=True, mmap=True) + cfg = ckpt.get("factorizer") or ckpt.get("config", {}).get("factorizer", {}) + return int(cfg.get("k", 256)) + + +def scoring_row_mean(w_u: torch.Tensor, tokenizer, mode: str, ckpt_row_mean: torch.Tensor | None) -> torch.Tensor: + """Centering mean on ``w_u``'s device under ``--centering``: ``live`` is the + full-vocabulary mean of the live LM head; ``trained`` is the checkpoint's + stored ``row_mean``, else the mean over the tokenizer's text-token rows.""" + token_mask = None + if mode == "trained": + token_mask = token_mask_from_tokenizer(tokenizer, w_u.shape[0]).to(w_u.device) + ckpt = {"row_mean": ckpt_row_mean} if ckpt_row_mean is not None else None + return centering_mean(w_u, mode=mode, token_mask=token_mask, ckpt=ckpt).to(w_u.device) + + def resolve_single_token(tokenizer, target: str) -> tuple[int, str] | None: # The prompt ends in ':', so the natural continuation includes a leading # space. Requiring this exact continuation avoids silently scoring a token @@ -430,6 +306,16 @@ def prepare_target_codes( w_dec_unit: torch.Tensor, k: int, ) -> tuple[dict[str, dict[str, Any]], dict[str, Any]]: + """Top-k codes of each target's centred, unit-normalised unembedding row. + + ``w_enc`` is the (d_features, d_model) encoder weight from ``load_sae``. + The paper's run multiplied against its (d_model, d_features) transpose + held contiguously, and gemm summation order follows the operand layout, so + the same layout is handed to ``encode_topk`` (``w_enc.T.contiguous().T`` + is that layout viewed back) to reproduce the frozen bundles bit for bit + rather than to the last float32 ulp. + """ + encoder_w = w_enc.T.contiguous().T info: dict[str, dict[str, Any]] = {} skipped: dict[str, str] = {} for target in sorted(set(targets)): @@ -438,12 +324,14 @@ def prepare_target_codes( skipped[target] = "multi_token" continue token_id, continuation = resolved - row = w_u[token_id] - row_mean - row_norm = row.norm().clamp_min(1e-8) - acts = torch.relu((row / row_norm) @ w_enc + b_enc) - values, indices = torch.topk(acts, k=min(k, acts.numel())) + row_norms, x = center_normalize_rows(w_u[token_id][None], row_mean) + row_norm = row_norms[0] + code = encode_topk(x, encoder_w, b_enc, k)[0] + # Bundle columns are the k codes in descending order (torch.topk's order). + values, indices = torch.topk(code, k=min(k, code.numel())) beta = row_norm * values recon = beta @ w_dec_unit[indices] + row = w_u[token_id] - row_mean row_error = float((row - recon).norm() / row_norm) row_cos = float(torch.nn.functional.cosine_similarity(row[None], recon[None]).item()) info[target] = { @@ -473,7 +361,6 @@ def score_items( model, tokenizer, lm_head, - w_u: torch.Tensor, row_mean: torch.Tensor, w_dec_unit: torch.Tensor, target_info: dict[str, dict[str, Any]], @@ -493,6 +380,8 @@ def score_items( target_logprobs: list[float] = [] target_ranks: list[int] = [] grabbed: list[torch.Tensor] = [] + n_truncated = 0 + max_prompt_tokens = 0 def pre_hook(_module, inputs): grabbed.append(inputs[0].detach()) @@ -502,8 +391,12 @@ def pre_hook(_module, inputs): with torch.inference_mode(): for start in range(0, len(kept), batch_size): batch = kept[start : start + batch_size] + prompts = [row["prompt"] for row in batch] + lengths = [len(ids) for ids in tokenizer(prompts, padding=False, truncation=False)["input_ids"]] + n_truncated += sum(length > max_length for length in lengths) + max_prompt_tokens = max(max_prompt_tokens, *lengths) enc = tokenizer( - [row["prompt"] for row in batch], + prompts, return_tensors="pt", padding=True, truncation=True, @@ -567,273 +460,29 @@ def pre_hook(_module, inputs): "reconstructed_logit": torch.tensor(reconstructed_logits, dtype=torch.float32), "target_logprob": torch.tensor(target_logprobs, dtype=torch.float32), "target_rank": torch.tensor(target_ranks, dtype=torch.int32), + "truncation": { + "max_length": max_length, + "truncation_side": tokenizer.truncation_side, + "n_truncated_prompts": int(n_truncated), + "max_prompt_tokens": int(max_prompt_tokens), + }, } -def l2_normalize(matrix: np.ndarray) -> np.ndarray: - norms = np.linalg.norm(matrix, axis=1, keepdims=True) - return matrix / np.maximum(norms, 1e-12) - - def method_matrices(bundle: dict[str, Any], seed: int) -> dict[str, np.ndarray]: hidden = bundle["hidden"].float().numpy() projection = bundle["projection"].float().numpy() contribution = bundle["contribution"].float().numpy() beta = bundle["beta"].float().numpy() - metadata = bundle["metadata"] - shuffled = np.empty_like(contribution) - for target in sorted({row["target"] for row in metadata}): - idx = [i for i, row in enumerate(metadata) if row["target"] == target] - local_seed = seed + int(stable_id(target, n=8), 16) - rng = np.random.default_rng(local_seed) - permutation = rng.permutation(beta.shape[1]) - shuffled[idx] = projection[idx] * beta[idx][:, permutation] return { "srp": l2_normalize(contribution), "support_projection": l2_normalize(projection), "hidden": l2_normalize(hidden), - "shuffled_srp": l2_normalize(shuffled), + "shuffled_srp": l2_normalize(shuffled_srp(projection, beta, bundle["metadata"], seed)), "token_only": l2_normalize(beta), } -def safe_spearman(x: list[float] | np.ndarray, y: list[float] | np.ndarray) -> float: - from scipy.stats import spearmanr - - if len(x) < 3 or np.ptp(x) < 1e-8 or np.ptp(y) < 1e-8: - return float("nan") - return float(spearmanr(x, y).statistic) - - -def cluster_bootstrap( - rows: list[dict[str, Any]], - cluster_key: str, - stat_fn, - n_boot: int, - seed: int, -) -> list[float]: - by_cluster: dict[str, list[dict[str, Any]]] = defaultdict(list) - for row in rows: - by_cluster[str(row[cluster_key])].append(row) - keys = sorted(by_cluster) - rng = np.random.default_rng(seed) - values: list[float] = [] - for _ in range(n_boot): - sample: list[dict[str, Any]] = [] - for selected in rng.choice(keys, size=len(keys), replace=True): - sample.extend(by_cluster[str(selected)]) - value = float(stat_fn(sample)) - if np.isfinite(value): - values.append(value) - return values - - -def ci(values: list[float]) -> list[float]: - if not values: - return [float("nan"), float("nan")] - return [float(np.percentile(values, 2.5)), float(np.percentile(values, 97.5))] - - -def analyze_ambistory(bundle: dict[str, Any], out_dir: Path, n_boot: int, seed: int) -> dict: - metadata = bundle["metadata"] - matrices = method_matrices(bundle, seed) - anchor_index = {(row["target"], row["anchor_key"]): i for i, row in enumerate(metadata) if row["kind"] == "anchor"} - story_indices = [i for i, row in enumerate(metadata) if row["kind"] == "story"] - predictions: list[dict[str, Any]] = [] - metrics: dict[str, Any] = {"dataset": "ambistory", "methods": {}, "comparisons": {}} - rows_by_method: dict[str, list[dict[str, Any]]] = {} - shifts_by_method: dict[str, list[dict[str, Any]]] = {} - - for method, matrix in matrices.items(): - rows: list[dict[str, Any]] = [] - for i in story_indices: - meta = metadata[i] - anchor_ids = [anchor_index.get((meta["target"], key)) for key in meta["anchor_keys"]] - if any(index is None for index in anchor_ids): - continue - similarities = [float(matrix[i] @ matrix[int(index)]) for index in anchor_ids] - pred_margin = similarities[0] - similarities[1] - human_margin = float(meta["human_margin"]) - row = { - "method": method, - "item_id": meta["item_id"], - "setup_id": meta["setup_id"], - "context_id": meta["context_id"], - "split": meta["split"], - "target": meta["target"], - "has_ending": bool(meta["has_ending"]), - "human_margin": human_margin, - "pred_margin": pred_margin, - "similarity_0": similarities[0], - "similarity_1": similarities[1], - "correct": int(human_margin != 0 and np.sign(pred_margin) == np.sign(human_margin)), - "eligible_gap1": int(abs(human_margin) >= 1.0), - "row_relative_error": float(meta["row_relative_error"]), - } - rows.append(row) - if method == "srp": - predictions.append(row) - - def rho_fn(sample): - return safe_spearman([r["human_margin"] for r in sample], [r["pred_margin"] for r in sample]) - - eligible = [r for r in rows if r["eligible_gap1"]] - non_ties = [r for r in rows if r["human_margin"] != 0] - rho = rho_fn(rows) - rho_boot = cluster_bootstrap(rows, "setup_id", rho_fn, n_boot, seed) - accuracy_gap1 = float(np.mean([r["correct"] for r in eligible])) if eligible else float("nan") - acc_boot = ( - cluster_bootstrap( - eligible, - "setup_id", - lambda sample: np.mean([r["correct"] for r in sample]), - n_boot, - seed + 1, - ) - if eligible - else [] - ) - - # Controlled context change: representation distance should grow with - # the change in human sense-preference margin within the same setup. - shift_rows: list[dict[str, Any]] = [] - by_setup: dict[str, list[dict[str, Any]]] = defaultdict(list) - by_item = {r["item_id"]: r for r in rows} - for i in story_indices: - meta = metadata[i] - if meta["item_id"] in by_item: - by_setup[meta["setup_id"]].append({"i": i, **by_item[meta["item_id"]]}) - reversal_total = reversal_correct = 0 - for setup_id, setup_rows in by_setup.items(): - for left, right in itertools.combinations(setup_rows, 2): - distance = float(1.0 - matrix[left["i"]] @ matrix[right["i"]]) - human_shift = abs(left["human_margin"] - right["human_margin"]) - shift_rows.append({"setup_id": setup_id, "distance": distance, "human_shift": human_shift}) - if left["human_margin"] * right["human_margin"] < 0: - reversal_total += 1 - if left["pred_margin"] * right["pred_margin"] < 0: - reversal_correct += 1 - shift_rho = safe_spearman([r["human_shift"] for r in shift_rows], [r["distance"] for r in shift_rows]) - shift_boot = ( - cluster_bootstrap( - shift_rows, - "setup_id", - lambda sample: safe_spearman([r["human_shift"] for r in sample], [r["distance"] for r in sample]), - n_boot, - seed + 2, - ) - if shift_rows - else [] - ) - metrics["methods"][method] = { - "n_contexts": len(rows), - "n_setups": len({r["setup_id"] for r in rows}), - "preference_spearman": rho, - "preference_spearman_ci": ci(rho_boot), - "preferred_sense_accuracy_gap1": accuracy_gap1, - "preferred_sense_accuracy_gap1_ci": ci(acc_boot), - "n_gap1": len(eligible), - "preferred_sense_accuracy_non_tie": ( - float(np.mean([r["correct"] for r in non_ties])) if non_ties else float("nan") - ), - "profile_shift_spearman": shift_rho, - "profile_shift_spearman_ci": ci(shift_boot), - "n_shift_pairs": len(shift_rows), - "sense_reversal_consistency": (reversal_correct / reversal_total if reversal_total else float("nan")), - "n_reversal_pairs": reversal_total, - } - rows_by_method[method] = rows - shifts_by_method[method] = shift_rows - - # Confidence intervals for the scientific comparisons must be paired: all - # methods score the same stories, so resample setups once and take the - # within-resample difference. Independent per-method intervals do not test - # whether SRP itself differs from a control. - for comparison_index, baseline in enumerate(("shuffled_srp", "support_projection", "hidden")): - srp_rows = rows_by_method["srp"] - baseline_rows = rows_by_method[baseline] - if [row["item_id"] for row in srp_rows] != [row["item_id"] for row in baseline_rows]: - raise ValueError(f"unaligned AmbiStory rows for srp and {baseline}") - paired_rows = [ - { - "setup_id": left["setup_id"], - "human_margin": left["human_margin"], - "eligible_gap1": left["eligible_gap1"], - "srp_pred_margin": left["pred_margin"], - "baseline_pred_margin": right["pred_margin"], - "srp_correct": left["correct"], - "baseline_correct": right["correct"], - } - for left, right in zip(srp_rows, baseline_rows) - ] - - def preference_delta(sample): - human = [row["human_margin"] for row in sample] - return safe_spearman(human, [row["srp_pred_margin"] for row in sample]) - safe_spearman( - human, [row["baseline_pred_margin"] for row in sample] - ) - - eligible_paired = [row for row in paired_rows if row["eligible_gap1"]] - - def accuracy_delta(sample): - return np.mean([row["srp_correct"] for row in sample]) - np.mean( - [row["baseline_correct"] for row in sample] - ) - - srp_shifts = shifts_by_method["srp"] - baseline_shifts = shifts_by_method[baseline] - if [row["setup_id"] for row in srp_shifts] != [row["setup_id"] for row in baseline_shifts] or not np.allclose( - [row["human_shift"] for row in srp_shifts], - [row["human_shift"] for row in baseline_shifts], - ): - raise ValueError(f"unaligned AmbiStory shift rows for srp and {baseline}") - paired_shifts = [ - { - "setup_id": left["setup_id"], - "human_shift": left["human_shift"], - "srp_distance": left["distance"], - "baseline_distance": right["distance"], - } - for left, right in zip(srp_shifts, baseline_shifts) - ] - - def shift_delta(sample): - human = [row["human_shift"] for row in sample] - return safe_spearman(human, [row["srp_distance"] for row in sample]) - safe_spearman( - human, [row["baseline_distance"] for row in sample] - ) - - comparison_seed = seed + 100 + comparison_index * 10 - metrics["comparisons"][f"srp_minus_{baseline}"] = { - "preference_spearman_difference": float(preference_delta(paired_rows)), - "preference_spearman_difference_ci": ci( - cluster_bootstrap(paired_rows, "setup_id", preference_delta, n_boot, comparison_seed) - ), - "preferred_sense_accuracy_gap1_difference": float(accuracy_delta(eligible_paired)), - "preferred_sense_accuracy_gap1_difference_ci": ci( - cluster_bootstrap( - eligible_paired, - "setup_id", - accuracy_delta, - n_boot, - comparison_seed + 1, - ) - ), - "profile_shift_spearman_difference": float(shift_delta(paired_shifts)), - "profile_shift_spearman_difference_ci": ci( - cluster_bootstrap( - paired_shifts, - "setup_id", - shift_delta, - n_boot, - comparison_seed + 2, - ) - ), - } - write_jsonl(out_dir / "predictions.jsonl", predictions) - return metrics - - def sampled_pair_auc(matrix: np.ndarray, labels: np.ndarray, seed: int, max_pairs: int = 20000) -> float: from sklearn.metrics import roc_auc_score @@ -891,79 +540,61 @@ def analyze_coarsewsd(bundle: dict[str, Any], out_dir: Path, n_boot: int, seed: metadata = bundle["metadata"] matrices = method_matrices(bundle, seed) - words = sorted({row["target"] for row in metadata}) + splits = word_splits(metadata) metrics: dict[str, Any] = {"dataset": "coarsewsd20", "methods": {}, "comparisons": {}} srp_predictions: list[dict[str, Any]] = [] per_word_by_method: dict[str, dict[str, dict[str, float]]] = {} for method, matrix in matrices.items(): per_word: dict[str, dict[str, float]] = {} - for word in words: - train_idx = [i for i, row in enumerate(metadata) if row["target"] == word and row["split"] == "train"] - test_idx = [i for i, row in enumerate(metadata) if row["target"] == word and row["split"] == "test"] - if not train_idx or not test_idx: - continue - train_y = np.array([metadata[i]["sense"] for i in train_idx]) - test_y = np.array([metadata[i]["sense"] for i in test_idx]) - senses = sorted(set(train_y)) - if len(senses) < 2 or not set(test_y).issubset(set(senses)): - continue + for split in splits: centroids = [] - for sense in senses: - centroid = matrix[np.array(train_idx)[train_y == sense]].mean(axis=0) + for sense in split.senses: + centroid = matrix[split.train_idx[split.train_y == sense]].mean(axis=0) centroid /= max(np.linalg.norm(centroid), 1e-12) centroids.append(centroid) centroid_matrix = np.stack(centroids) - scores = matrix[test_idx] @ centroid_matrix.T - predictions = np.array([senses[i] for i in scores.argmax(axis=1)]) - n_clusters = len(senses) - train_matrix = matrix[train_idx] - test_matrix = matrix[test_idx] + scores = matrix[split.test_idx] @ centroid_matrix.T + predictions = np.array([split.senses[i] for i in scores.argmax(axis=1)]) cluster_pred = fit_train_clusters( - train_matrix, - test_matrix, - n_clusters=n_clusters, + matrix[split.train_idx], + matrix[split.test_idx], + n_clusters=len(split.senses), seed=seed, ) - word_metrics = { - "n_train": len(train_idx), - "n_test": len(test_idx), - "n_senses": len(senses), - "accuracy": float(accuracy_score(test_y, predictions)), - "balanced_accuracy": float(balanced_accuracy_score(test_y, predictions)), - "macro_f1": float(f1_score(test_y, predictions, average="macro")), - "pairwise_auc": sampled_pair_auc(matrix[test_idx], test_y, seed), - "ari": float(adjusted_rand_score(test_y, cluster_pred)), - "nmi": float(normalized_mutual_info_score(test_y, cluster_pred)), - "row_relative_error": float(metadata[test_idx[0]]["row_relative_error"]), + per_word[split.word] = { + "n_train": len(split.train_idx), + "n_test": len(split.test_idx), + "n_senses": len(split.senses), + "accuracy": float(accuracy_score(split.test_y, predictions)), + "balanced_accuracy": float(balanced_accuracy_score(split.test_y, predictions)), + "macro_f1": float(f1_score(split.test_y, predictions, average="macro")), + "pairwise_auc": sampled_pair_auc(matrix[split.test_idx], split.test_y, seed), + "ari": float(adjusted_rand_score(split.test_y, cluster_pred)), + "nmi": float(normalized_mutual_info_score(split.test_y, cluster_pred)), + "row_relative_error": float(metadata[split.test_idx[0]]["row_relative_error"]), } - per_word[word] = word_metrics if method == "srp": - for local, idx in enumerate(test_idx): + for local, idx in enumerate(split.test_idx): srp_predictions.append( { "item_id": metadata[idx]["item_id"], - "target": word, - "gold_sense": str(test_y[local]), + "target": split.word, + "gold_sense": str(split.test_y[local]), "predicted_sense": str(predictions[local]), - "correct": int(test_y[local] == predictions[local]), + "correct": int(split.test_y[local] == predictions[local]), "target_rank": int(bundle["target_rank"][idx]), "target_logprob": float(bundle["target_logprob"][idx]), } ) - metric_names = ["accuracy", "balanced_accuracy", "macro_f1", "pairwise_auc", "ari", "nmi"] aggregate: dict[str, Any] = {"n_words": len(per_word)} rng = np.random.default_rng(seed) word_names = sorted(per_word) - for metric_name in metric_names: + for metric_name in METRIC_NAMES: values = np.array([per_word[word][metric_name] for word in word_names], dtype=float) values = values[np.isfinite(values)] aggregate[metric_name] = float(values.mean()) if len(values) else float("nan") - boot = [] - if len(values): - for _ in range(n_boot): - boot.append(float(rng.choice(values, size=len(values), replace=True).mean())) - aggregate[f"{metric_name}_ci"] = ci(boot) + aggregate[f"{metric_name}_ci"] = percentile_ci(bootstrap_mean(values, rng, n_boot) if len(values) else []) gated_words = [w for w in word_names if per_word[w]["row_relative_error"] < 0.5] aggregate["row_gate_words"] = len(gated_words) aggregate["row_gate_fraction"] = len(gated_words) / max(len(word_names), 1) @@ -971,19 +602,18 @@ def analyze_coarsewsd(bundle: dict[str, Any], out_dir: Path, n_boot: int, seed: metric_name: ( float(np.mean([per_word[w][metric_name] for w in gated_words])) if gated_words else float("nan") ) - for metric_name in metric_names + for metric_name in METRIC_NAMES } metrics["methods"][method] = {"aggregate": aggregate, "per_word": per_word} per_word_by_method[method] = per_word # CoarseWSD words are the independent units. Bootstrap paired per-word # differences so the intervals answer whether SRP beats each control. - metric_names = ["accuracy", "balanced_accuracy", "macro_f1", "pairwise_auc", "ari", "nmi"] - for comparison_index, baseline in enumerate(("shuffled_srp", "support_projection", "hidden")): + for comparison_index, baseline in enumerate(COMPARISON_BASELINES): common_words = sorted(set(per_word_by_method["srp"]) & set(per_word_by_method[baseline])) comparison: dict[str, Any] = {"n_words": len(common_words)} rng = np.random.default_rng(seed + 200 + comparison_index) - for metric_name in metric_names: + for metric_name in METRIC_NAMES: differences = np.array( [ per_word_by_method["srp"][word][metric_name] - per_word_by_method[baseline][word][metric_name] @@ -993,12 +623,9 @@ def analyze_coarsewsd(bundle: dict[str, Any], out_dir: Path, n_boot: int, seed: ) differences = differences[np.isfinite(differences)] comparison[f"{metric_name}_difference"] = float(differences.mean()) if len(differences) else float("nan") - boot = ( - [float(rng.choice(differences, size=len(differences), replace=True).mean()) for _ in range(n_boot)] - if len(differences) - else [] + comparison[f"{metric_name}_difference_ci"] = percentile_ci( + bootstrap_mean(differences, rng, n_boot) if len(differences) else [] ) - comparison[f"{metric_name}_difference_ci"] = ci(boot) metrics["comparisons"][f"srp_minus_{baseline}"] = comparison write_jsonl(out_dir / "predictions.jsonl", srp_predictions) return metrics @@ -1009,7 +636,7 @@ def summarize_scoring(bundle: dict[str, Any]) -> dict[str, Any]: exact = bundle["exact_logit"].numpy() recon = bundle["reconstructed_logit"].numpy() ranks = bundle["target_rank"].numpy() - return { + summary = { "n_scored": len(metadata), "n_targets": len({row["target"] for row in metadata}), "median_target_rank": float(np.median(ranks)), @@ -1019,11 +646,12 @@ def summarize_scoring(bundle: dict[str, Any]) -> dict[str, Any]: "mean_absolute_logit_residual": float(np.mean(np.abs(exact - recon))), "row_gate_fraction_items": float(np.mean([row["row_relative_error"] < 0.5 for row in metadata])), } + if "truncation" in bundle: + summary["truncation"] = dict(bundle["truncation"]) + return summary def self_test() -> int: - assert mask_exact_target("They followed the track.", "track") == "They followed the [BLANK]." - assert mask_exact_target("a tracker", "track") is None assert resolve_target_position(["the", "bank", "closed"], 1, "bank") == 1 assert resolve_target_position(["the", "bank", "closed"], 2, "bank") == 1 matrix = l2_normalize(np.array([[1.0, 0.0], [0.9, 0.1], [0.0, 1.0], [0.1, 0.9]])) @@ -1038,38 +666,29 @@ def self_test() -> int: return 0 -def parse_args() -> argparse.Namespace: +def parse_args(argv: list[str] | None = None) -> argparse.Namespace: parser = argparse.ArgumentParser() - parser.add_argument("--dataset", choices=("ambistory", "coarsewsd20"), help="the paper uses coarsewsd20") + parser.add_argument("--dataset", choices=("coarsewsd20",), default="coarsewsd20") parser.add_argument( "--data-root", type=Path, - help="coarsewsd20: a bert-disambiguation checkout (or its data/CoarseWSD-20 dir); ambistory: dir of .json", + help="a bert-disambiguation checkout (or its data/CoarseWSD-20 dir)", ) parser.add_argument("--model-id", help="Hugging Face model id") parser.add_argument("--checkpoint", type=Path, help="trained readout dictionary checkpoint (.pt)") parser.add_argument("--out-dir", type=Path, help="output dir: representations.pt, audit/scoring/metrics/run_config") - parser.add_argument( - "--splits", default=None, help="comma-separated; defaults dev for AmbiStory, train,test for CoarseWSD" - ) + parser.add_argument("--splits", default="train,test", help="comma-separated CoarseWSD-20 splits") parser.add_argument("--batch-size", type=int, default=8) parser.add_argument("--max-length", type=int, default=384) parser.add_argument("--limit", type=int, default=None) parser.add_argument("--limit-per-word", type=int, default=None) + parser.add_argument("--k", type=int, default=None, help="top-k codes per row; default: the checkpoint's k") parser.add_argument( - "--position-mode", - choices=("cloze", "occurrence"), - default="cloze", - help="cloze: mask target + cue; occurrence: score the sentence " - "prefix so the target occurrence is the genuine next token " - "(coarsewsd20 only)", - ) - parser.add_argument( - "--anchors", - type=Path, - default=None, - help="frozen sense-anchor JSON; anchor rows and sparse codes are " - "embedded in the bundle for offline weights-only feature typing", + "--centering", + choices=("live", "trained"), + default="live", + help="row centering mean: live = full-vocabulary mean of the live LM head (the paper's runs); " + "trained = the checkpoint's stored row_mean, else the tokenizer's text-token rows", ) parser.add_argument("--device", default="cuda") parser.add_argument("--dtype", choices=("bfloat16", "float32"), default="bfloat16") @@ -1077,54 +696,41 @@ def parse_args() -> argparse.Namespace: parser.add_argument("--n-boot", type=int, default=2000) parser.add_argument("--analyze-only", action="store_true") parser.add_argument("--self-test", action="store_true") - return parser.parse_args() + return parser.parse_args(argv) -def main() -> int: - args = parse_args() +def main(argv: list[str] | None = None) -> int: + args = parse_args(argv) if args.self_test: return self_test() - required = (args.dataset, args.data_root, args.model_id, args.checkpoint, args.out_dir) + required = (args.data_root, args.model_id, args.checkpoint, args.out_dir) if not args.analyze_only and not all(required): - raise SystemExit("dataset, data-root, model-id, checkpoint, and out-dir are required") + raise SystemExit("data-root, model-id, checkpoint, and out-dir are required") if args.analyze_only and args.out_dir is None: raise SystemExit("--out-dir is required with --analyze-only") args.out_dir.mkdir(parents=True, exist_ok=True) representation_path = args.out_dir / "representations.pt" + provenance = run_provenance(args) if args.analyze_only: - bundle = torch.load(representation_path, map_location="cpu", weights_only=False) + bundle = load_bundle(representation_path) else: set_seed(args.seed) - if args.splits: - splits = [part.strip() for part in args.splits.split(",") if part.strip()] - else: - splits = ["dev"] if args.dataset == "ambistory" else ["train", "test"] - if args.dataset == "ambistory": - if args.position_mode != "cloze": - raise SystemExit("--position-mode occurrence supports coarsewsd20 only") - items, data_audit = load_ambistory(args.data_root, splits, args.limit) - else: - items, data_audit = load_coarsewsd( - args.data_root, splits, args.limit, args.limit_per_word, args.position_mode - ) + splits = [part.strip() for part in args.splits.split(",") if part.strip()] + items, data_audit = load_coarsewsd(args.data_root, splits, args.limit, args.limit_per_word) log(f"loaded {len(items)} items ({data_audit})") log(f"loading {args.model_id}") model, tokenizer = load_model(args.model_id, args.device, args.dtype) - if args.position_mode == "occurrence": - # Long prefixes must keep their tail (the context nearest the - # scored position), so truncate from the left. - tokenizer.truncation_side = "left" lm_head = find_lm_head(model) w_u = lm_head.weight.detach().float().to(args.device) - row_mean = w_u.mean(dim=0) - w_dec, w_enc, b_enc, k = load_sae(args.checkpoint) - w_dec = w_dec.to(args.device) - w_dec_unit = w_dec / w_dec.norm(dim=1, keepdim=True).clamp_min(1e-8) + w_dec_unit, w_enc, b_enc, sae_config, ckpt_row_mean = load_sae(args.checkpoint) + k = args.k if args.k is not None else checkpoint_k(args.checkpoint, sae_config) + row_mean = scoring_row_mean(w_u, tokenizer, args.centering, ckpt_row_mean) + w_dec_unit = w_dec_unit.to(args.device) w_enc = w_enc.to(args.device) b_enc = b_enc.to(args.device) - if w_dec.shape[1] != w_u.shape[1]: - raise ValueError(f"dictionary/model width mismatch: {w_dec.shape[1]} vs {w_u.shape[1]}") + if w_dec_unit.shape[1] != w_u.shape[1]: + raise ValueError(f"dictionary/model width mismatch: {w_dec_unit.shape[1]} vs {w_u.shape[1]}") target_info, token_audit = prepare_target_codes( [row["target"] for row in items], tokenizer, @@ -1135,56 +741,15 @@ def main() -> int: w_dec_unit, k, ) - anchor_bundle = None - if args.anchors is not None: - anchors_cfg = json.loads(args.anchors.read_text()) - anchor_words = sorted( - { - anchor - for word, senses in anchors_cfg.items() - if not word.startswith("_") - for anchor_list in senses.values() - for anchor in anchor_list - } - ) - anchor_info, anchor_audit = prepare_target_codes( - anchor_words, - tokenizer, - w_u, - row_mean, - w_enc, - b_enc, - w_dec_unit, - k, - ) - anchor_bundle = { - "config": anchors_cfg, - "codes": { - word: { - "token_id": info["token_id"], - "feature_ids": info["feature_ids"].cpu(), - "beta": info["beta"].cpu(), - } - for word, info in anchor_info.items() - }, - "rows": {word: w_u[info["token_id"]].cpu().half() for word, info in anchor_info.items()}, - "audit": anchor_audit, - } - log( - "anchor coverage " - f"{anchor_audit['n_targets_single_token']}/" - f"{anchor_audit['n_targets_requested']} single-token" - ) - del w_enc, b_enc, w_dec + del w_enc, b_enc audit = {**data_audit, "tokenization_and_rows": token_audit, "k": k} - write_json(args.out_dir / "audit.json", audit) + write_json(audit, args.out_dir / "audit.json") log(f"single-token coverage {token_audit['n_targets_single_token']}/{token_audit['n_targets_requested']}") bundle = score_items( items, model, tokenizer, lm_head, - w_u, row_mean, w_dec_unit, target_info, @@ -1201,33 +766,24 @@ def main() -> int: "k": k, "batch_size": args.batch_size, "max_length": args.max_length, - "position_mode": args.position_mode, + "centering": args.centering, } - if anchor_bundle is not None: - bundle["anchors"] = anchor_bundle torch.save(bundle, representation_path) - write_json(args.out_dir / "scoring_summary.json", summarize_scoring(bundle)) + write_json(summarize_scoring(bundle), args.out_dir / "scoring_summary.json") - metrics = ( - analyze_ambistory(bundle, args.out_dir, args.n_boot, args.seed) - if bundle["run"]["dataset"] == "ambistory" - else analyze_coarsewsd(bundle, args.out_dir, args.n_boot, args.seed) - ) + metrics = analyze_coarsewsd(bundle, args.out_dir, args.n_boot, args.seed) metrics["run"] = bundle["run"] metrics["scoring"] = summarize_scoring(bundle) - write_json(args.out_dir / "metrics.json", metrics) - run_config = { - "argv": sys.argv, + metrics["provenance"] = provenance + write_json(metrics, args.out_dir / "metrics.json") + config = { "python": sys.version, "platform": platform.platform(), - "torch": torch.__version__, "numpy": np.__version__, - "story_template": AMBISTORY_STORY_TEMPLATE, - "gloss_template": AMBISTORY_GLOSS_TEMPLATE, "coarsewsd_template": COARSEWSD_TEMPLATE, - "provenance": run_provenance(args), + "provenance": provenance, } - write_json(args.out_dir / "run_config.json", run_config) + write_json(config, args.out_dir / ("analysis_config.json" if args.analyze_only else "run_config.json")) log(f"wrote results to {args.out_dir}") return 0 diff --git a/src/sparse_readout_prism/data.py b/src/sparse_readout_prism/data.py index 4e6fb08..19805df 100644 --- a/src/sparse_readout_prism/data.py +++ b/src/sparse_readout_prism/data.py @@ -211,6 +211,41 @@ def resolve_row_mean( return W_U.float().mean(dim=0).cpu() +def centering_mean( + W_U: torch.Tensor, + *, + mode: str, + token_mask: torch.Tensor | None = None, + ckpt: dict[str, Any] | None = None, +) -> torch.Tensor: + """Centering mean under an explicit policy. + + ``mode="trained"`` is :func:`resolve_row_mean` (checkpoint value, else the + ``token_mask`` mean, else the full-vocabulary mean) and is what every + fidelity runner uses. ``mode="live"`` is the full-vocabulary mean of the + matrix handed in, which is how the paper's cross-lens, causal-validation, + stability and sense-labelled runs were computed; scripts that reproduce + those runs expose it as ``--centering live`` (their default) so the choice + is visible rather than implicit. On the released Qwen3.5-2B dictionary the + two means differ by ~0.04% of a centred row norm. + """ + if mode == "trained": + return resolve_row_mean(W_U, token_mask=token_mask, ckpt=ckpt) + if mode == "live": + return W_U.float().mean(dim=0).cpu() + raise ValueError(f"unknown centering mode {mode!r} (expected 'trained' or 'live')") + + +def token_mask_from_tokenizer(tok: Any, vocab: int) -> torch.Tensor: + """Text-token row mask, as the extraction builds it: rows past the tokenizer + vocabulary (padded/unused embedding rows) and every special id are False.""" + token_mask = torch.zeros(int(vocab), dtype=torch.bool) + token_mask[: min(int(vocab), len(tok))] = True + special_ids = [i for i in (getattr(tok, "all_special_ids", None) or []) if 0 <= i < vocab] + token_mask[special_ids] = False + return token_mask + + def choose_row_subset( W_U: torch.Tensor, hidden: torch.Tensor, diff --git a/src/sparse_readout_prism/research/__init__.py b/src/sparse_readout_prism/research/__init__.py index 9fdba0d..3f97e3b 100644 --- a/src/sparse_readout_prism/research/__init__.py +++ b/src/sparse_readout_prism/research/__init__.py @@ -10,7 +10,11 @@ the Qwen readout toolkit (``qwen_readout``), the query-decomposition toolkit (``query_decompose``), the static example-prompt banks (``qwen_example_prompts``), cell metrics (``cell_metrics``), run IO - (``run_io``), and the multi-model run-script registry (``registry``). + (``run_io``), the multi-model run-script registry (``registry``), the + seed-stability contrast pipeline (``seed_stability``), row-geometry + helpers shared by the baselines and stability scripts (``row_geometry``), + the cross-lens study toolkit (``cross_lens``), and the CoarseWSD-20 + bundle/statistics helpers (``wsd``). So ``research/`` is not a mirror of ``scripts/``: every module here backs at least two consumers (enforced by ``tests/test_research_layout.py``), and no diff --git a/src/sparse_readout_prism/research/cross_lens.py b/src/sparse_readout_prism/research/cross_lens.py new file mode 100644 index 0000000..1e28896 --- /dev/null +++ b/src/sparse_readout_prism/research/cross_lens.py @@ -0,0 +1,582 @@ +"""Cross-lens study toolkit: lens-model loading, the readout decomposition, and the aggregation core. + +Shared by the lens fitters (``scripts/run/fit_jlens.py``, ``scripts/run/fit_ridge_lens.py``), +the readout runner (``scripts/run/run_cross_lens_readouts.py``), the antonym +layer study (``scripts/analyze/cross_lens_antonym_layers.py``), the two +aggregators (``scripts/eval/aggregate_cross_lens_en_{zh,de}.py``), the +three-lens table and shared-feature scripts, and the EN-DE bank builder. Until +0.2.1 each of those carried its own copy of the pieces below. + +* :func:`import_jlens` -- lazy import of the Jacobian-lens reference + implementation (``uv sync --extra lens``). +* :func:`load_lens_model` / :func:`find_final_norm` -- one model loader for the + GPU scripts, built on ``utils.load_causal_lm`` and ``utils.find_lm_head``. +* :func:`decompose_token` / :func:`feature_top_tokens` -- the per-token Sparse + Readout Prism decomposition of a transported state and the raw-decode feature + labels the paper's dumps carry. +* :func:`resolve_k` / :func:`centering_row_mean` -- the ``--k`` default and the + ``--centering {live,trained}`` switch of the two dictionary-reading scripts. +* :func:`majority_same` / :func:`token_diverges` -- the mid-band majority votes, + with ``rule="half"`` (the paper: at least half of the scored layers) or + ``rule="strict"`` (more than half). +* :func:`aggregate_pair` / :func:`print_summary` / :func:`write_summary` -- the + aggregation core, parameterised by a :class:`PairSpec` so the EN-ZH and EN-DE + entry points differ only in their surface-language call; with + :func:`load_dump_records` / :func:`load_bank_items` / :func:`parse_dump_args` + as the dump and bank readers. +* :func:`check_resume_meta` / :func:`write_resume_meta` -- the sidecar that + makes shard / accumulator resume refuse a different fit. +* :func:`write_manifest` / :func:`cli_args` -- ``.manifest.json`` with + ``run_provenance`` for the CSV-writing scripts (subcommand callables stripped). +* :func:`fold_text` -- casefold + eszett + diacritic folding (bank builder and + the EN-DE lexical call). +""" + +from __future__ import annotations + +import hashlib +import json +import math +import random +import unicodedata +import warnings +from collections.abc import Callable, Iterable, Sequence +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +import torch +import torch.nn.functional as F + +from sparse_readout_prism.data import center_normalize_rows, centering_mean, token_mask_from_tokenizer +from sparse_readout_prism.research.qwen_readout import encode_topk +from sparse_readout_prism.research.run_io import run_provenance +from sparse_readout_prism.utils import find_lm_head, load_causal_lm, write_json + +MIDBAND = ["21", "24", "26", "29"] +POS = "-1" +AGREEMENT_RULES = ("half", "strict") +NULL_POPULATIONS = ("all", "cross") +CENTERING_MODES = ("live", "trained") + + +# --------------------------------------------------------------------------- # +# model loading +# --------------------------------------------------------------------------- # + + +def import_jlens(): + try: + import jlens + except ImportError as e: + raise ImportError( + "jlens (the Jacobian-lens reference implementation, Apache-2.0, " + "github.com/anthropics/jacobian-lens) is not installed; install it with `uv sync --extra lens`" + ) from e + return jlens + + +_FINAL_NORM_PATHS = ("model.norm", "model.language_model.norm", "language_model.norm", "norm") + + +def find_final_norm(model: torch.nn.Module) -> torch.nn.Module: + """The final pre-unembedding norm of a HF causal LM (plain or multimodal wrapper).""" + for path in _FINAL_NORM_PATHS: + obj: Any = model + try: + for part in path.split("."): + obj = getattr(obj, part) + except AttributeError: + continue + if isinstance(obj, torch.nn.Module): + return obj + raise RuntimeError( + f"could not locate the final norm on {type(model).__name__} (tried {', '.join(_FINAL_NORM_PATHS)})" + ) + + +@dataclass +class LoadedLensModel: + """A HF model on its device plus the handles the cross-lens scripts read through.""" + + hf: torch.nn.Module + tok: Any + lens_model: Any # jlens.HFLensModel over ``hf`` + lm_head: torch.nn.Module + final_norm: torch.nn.Module + device: torch.device + + +def load_lens_model( + model_id: str, *, device: torch.device | str, dtype: torch.dtype = torch.bfloat16 +) -> LoadedLensModel: + """Load ``model_id`` through ``utils.load_causal_lm`` and wrap it for ``jlens``. + + ``load_causal_lm(device_map=None)`` loads on CPU with frozen parameters; the + model is then moved to ``device`` whole (the paper runs: one 9B model per + GPU in bf16). ``jlens.from_hf`` locates the residual stack; ``lm_head`` and + ``final_norm`` are the repo-wide ``utils.find_lm_head`` and + :func:`find_final_norm` so the readout matrix and the norm applied to + transported states come from the same places every other script uses. + """ + jlens = import_jlens() + device = torch.device(device) + hf, tok = load_causal_lm(model_id, dtype=dtype, device_map=None) + hf = hf.to(device).eval() + lens_model = jlens.from_hf(hf, tok) + return LoadedLensModel( + hf=hf, + tok=tok, + lens_model=lens_model, + lm_head=find_lm_head(hf), + final_norm=find_final_norm(hf), + device=device, + ) + + +# --------------------------------------------------------------------------- # +# dictionary: k, centering, decomposition, labels +# --------------------------------------------------------------------------- # + + +def resolve_k(requested: int | None, config: dict[str, Any]) -> int: + """Active features per code: the checkpoint's ``config['factorizer']['k']`` unless overridden. + + An explicit ``requested`` value that differs from the checkpoint's wins but + raises a ``UserWarning`` -- the dictionary was trained at its own k and a + different one changes every code. + """ + trained = (config or {}).get("factorizer", {}).get("k") + if requested is None: + if trained is None: + raise ValueError("checkpoint config carries no factorizer.k; pass --k explicitly") + return int(trained) + if trained is not None and int(trained) != int(requested): + warnings.warn( + f"--k {requested} differs from the checkpoint's trained k={trained}; codes will not match training", + UserWarning, + stacklevel=2, + ) + return int(requested) + + +def centering_row_mean(W: torch.Tensor, mode: str, *, tok: Any, ckpt_row_mean: torch.Tensor | None) -> torch.Tensor: + """``data.centering_mean`` for scripts that only have the live model. + + ``live`` is the full-vocabulary mean of the live (bf16 -> fp32) head, which + is how the paper's cross-lens dumps were computed. ``trained`` is the + checkpoint's stored ``row_mean`` when it has one, else the mean over the + text-token rows ``data.token_mask_from_tokenizer(tok, vocab)`` (how a masked + extraction trains). Returned on ``W``'s device. + """ + if mode not in CENTERING_MODES: + raise ValueError(f"unknown centering mode {mode!r} (expected one of {CENTERING_MODES})") + token_mask = None + if mode == "trained" and ckpt_row_mean is None: + token_mask = token_mask_from_tokenizer(tok, W.shape[0]) + return centering_mean(W, mode=mode, token_mask=token_mask, ckpt={"row_mean": ckpt_row_mean}).to(W.device) + + +@torch.no_grad() +def decompose_token( + h_state: torch.Tensor, + token_id: int, + *, + W: torch.Tensor, + row_mean: torch.Tensor, + decoder: torch.Tensor, + encoder_w: torch.Tensor, + encoder_b: torch.Tensor, + k: int, + top_feats: int, +) -> dict[str, Any]: + """Sparse Readout Prism decomposition of the score ``h_state . W[token_id]``. + + Row ``t`` is centred on ``row_mean`` and unit-normalised + (``data.center_normalize_rows``), top-k encoded, and feature ``i`` + contributes ``||W_t - mu|| * code_i(t) * (h . d_i)``. Returns the dump entry + ``{original_logit, base, feature_sum, residual, top_features}`` where + ``top_features`` are the ``top_feats`` largest |contribution| features with + their signed contributions (largest first). + """ + W_row = W[token_id] + norms, x = center_normalize_rows(W_row[None, :], row_mean) + code = encode_topk(x, encoder_w, encoder_b, k)[0] + contributions = norms[0] * code * (h_state @ decoder.T) + base = float(h_state @ row_mean) + feat_sum = float(contributions.sum()) + original = float(h_state @ W_row) + active = torch.nonzero(contributions != 0).flatten() + top = active[contributions[active].abs().argsort(descending=True)[:top_feats]] + return { + "original_logit": original, + "base": base, + "feature_sum": feat_sum, + "residual": original - base - feat_sum, + "top_features": [{"id": int(f), "contribution": float(contributions[f])} for f in top], + } + + +@torch.no_grad() +def feature_top_tokens( + W: torch.Tensor, + row_mean: torch.Tensor, + feature_ids: Iterable[int], + encoder_w: torch.Tensor, + encoder_b: torch.Tensor, + tokenizer: Any, + *, + top_tokens: int = 12, + chunk: int = 8192, +) -> dict[int, list[str]]: + """Top unembedding rows (by encoder activation) per feature id, as raw decoded strings. + + Labels are ``tokenizer.decode([id])`` verbatim, in activation order, rows + with zero activation dropped, no display cleaning and no de-duplication. + That is exactly what the paper's dumps carry under ``feature_top_tokens``; + the display-cleaned ``research.qwen_readout.display_label_features`` would + change those lists, so it is deliberately not used here. + """ + feature_ids = sorted(set(int(f) for f in feature_ids)) + if not feature_ids: + return {} + device = W.device + fids = torch.tensor(feature_ids, dtype=torch.long) + enc = encoder_w[fids].to(device) + bias = encoder_b[fids].to(device) + best_scores = torch.full((len(fids), top_tokens), -float("inf"), device=device) + best_ids = torch.zeros((len(fids), top_tokens), dtype=torch.long, device=device) + for start in range(0, W.shape[0], chunk): + rows = W[start : start + chunk].float() + _norms, x = center_normalize_rows(rows, row_mean) + scores = F.relu(x @ enc.T + bias).T + merged = torch.cat([best_scores, scores], dim=1) + ids = torch.arange(start, start + rows.shape[0], device=device).expand(len(fids), -1) + merged_ids = torch.cat([best_ids, ids], dim=1) + best_scores, keep = torch.topk(merged, k=top_tokens, dim=1) + best_ids = torch.gather(merged_ids, 1, keep) + return { + fid: [tokenizer.decode([t]) for t, s in zip(best_ids[i].tolist(), best_scores[i].tolist()) if s > 0] + for i, fid in enumerate(feature_ids) + } + + +# --------------------------------------------------------------------------- # +# offline: dumps, votes, aggregation +# --------------------------------------------------------------------------- # + + +def parse_dump_args(specs: list[str]) -> list[tuple[str, Path]]: + """Parse repeated ``LABEL=path`` arguments, preserving order.""" + out = [] + for spec in specs: + label, sep, path = spec.partition("=") + if not sep or not label or not path: + raise ValueError(f"--dump expects LABEL=path, got {spec!r}") + out.append((label, Path(path))) + return out + + +def load_dump_records(path: str | Path) -> dict[str, dict]: + """``{prompt id: record}`` of a readout dump from ``run_cross_lens_readouts.py``.""" + return {r["id"]: r for r in json.loads(Path(path).read_text())["records"]} + + +def load_bank_items(path: str | Path) -> dict[str, dict]: + """``{prompt id: item}`` of a cross-lens prompt bank; duplicate ids are an error.""" + items = json.loads(Path(path).read_text())["prompts"] + ids = [v["id"] for v in items] + if len(set(ids)) != len(ids): + dupes = sorted({i for i in ids if ids.count(i) > 1}) + raise ValueError(f"duplicate prompt ids in {path}: {dupes}") + return {v["id"]: v for v in items} + + +def wilson(k: int, n: int, z: float = 1.96) -> tuple[float, float, float]: + if n == 0: + return (0.0, 0.0, 1.0) + p = k / n + d = 1 + z * z / n + c = (p + z * z / (2 * n)) / d + h = z * math.sqrt(p * (1 - p) / n + z * z / (4 * n * n)) / d + return (p, max(0.0, c - h), min(1.0, c + h)) + + +def top5_at(rec: dict, layer: str) -> list | None: + return rec["layers"].get(layer, {}).get(POS, {}).get("top5") + + +def dom_feat(rec: dict, layer: str, target: str) -> int | None: + """Dominant feature id (largest |contribution|) of ``target``'s decomposition at ``layer``.""" + t = rec["layers"].get(layer, {}).get(POS, {}).get("targets", {}).get(target) + if not t or not t.get("top_features"): + return None + return t["top_features"][0]["id"] + + +def _majority(count: int, total: int, *, rule: str) -> bool: + if rule == "half": + return count >= (total + 1) // 2 + if rule == "strict": + return 2 * count > total + raise ValueError(f"unknown agreement rule {rule!r} (expected one of {AGREEMENT_RULES})") + + +def majority_same(rec_a: dict, rec_b: dict, target_a: str, target_b: str, *, rule: str) -> bool | None: + """Mid-band vote: is dom(target_a in rec_a) the same feature as dom(target_b in rec_b)? + + Layers where either decomposition is missing are skipped; ``None`` when none + remain. ``rule="half"`` passes on at least ``ceil(n/2)`` agreeing layers (the + paper: 2 of 4), ``rule="strict"`` needs more than half (3 of 4). + """ + same = total = 0 + for L in MIDBAND: + fa, fb = dom_feat(rec_a, L, target_a), dom_feat(rec_b, L, target_b) + if fa is None or fb is None: + continue + total += 1 + same += int(fa == fb) + if total == 0: + return None + return _majority(same, total, rule=rule) + + +def token_diverges(rec_a: dict, rec_b: dict, *, rule: str) -> bool | None: + """Mid-band vote: do the two lenses' top-1 token strings differ? Same ``rule`` as :func:`majority_same`.""" + diff = tot = 0 + for L in MIDBAND: + ta, tb = top5_at(rec_a, L), top5_at(rec_b, L) + if ta and tb: + tot += 1 + diff += int(ta[0][0] != tb[0][0]) + if tot == 0: + return None + return _majority(diff, tot, rule=rule) + + +@dataclass(frozen=True) +class PairSpec: + """What distinguishes one language pair's aggregation from the other's. + + ``tag_b`` suffixes the slot-B output keys (``cross_form_zh`` / ``_de``); + ``surface_call`` labels a lens's top-1 token given the bank item + (script for EN-ZH, lexical for EN-DE) into column ``top1__``; + ``split_labels`` is the (slot-A, slot-B) label pair counted as a lens-only + split; ``divergence`` adds the EN-DE top-1 string divergence rate. + """ + + tag_b: str + cross_groups: tuple[str, ...] + surface_column: str + surface_call: Callable[[str, dict], str] + split_labels: tuple[str, str] + split_caption: str + surface_caption: str + divergence: bool = False + + +def aggregate_pair( + records_a: dict[str, dict], + records_b: dict[str, dict], + bank: dict[str, dict], + spec: PairSpec, + *, + seed: int, + rule: str = "half", + null_population: str = "all", +) -> dict: + """Per-prompt rows, per-family and pooled agreement, null floors and the surface split. + + ``rule`` is the majority rule of every vote (headline, within-lens + cross-form, both nulls, divergence). ``null_population`` selects the rows + pooled into the unrelated-token floor: ``"all"`` (every bank item, the + paper's EN-DE 204 comparisons) or ``"cross"`` (control families excluded, + which is the EN-ZH population where controls carry no nulls). Key order of + the returned dict and of each row is the dump format; do not reorder. + """ + if null_population not in NULL_POPULATIONS: + raise ValueError(f"unknown null population {null_population!r} (expected one of {NULL_POPULATIONS})") + col = f"top1_{spec.surface_column}" + ids = sorted(set(records_a) & set(records_b) & set(bank)) + + rows = [] + for rid in ids: + v, a, b = bank[rid], records_a[rid], records_b[rid] + fa, fb = v["form_a"], v["form_b"] + row = {"id": rid, "group": v["group"], "concept": v.get("concept", "")} + # headline: same dominant feature for the SAME concept token, lens A vs lens B + row["cross_lens_pass"] = majority_same(a, b, fa, fa, rule=rule) + # one feature carries both surface forms, within each lens + row["cross_form_en"] = majority_same(a, a, fa, fb, rule=rule) + row[f"cross_form_{spec.tag_b}"] = majority_same(b, b, fa, fb, rule=rule) + # null floor (a): form_a vs unrelated token, within lens A + nulls = [majority_same(a, a, fa, nt, rule=rule) for nt in v.get("null_targets", [])] + nulls = [x for x in nulls if x is not None] + row["null_hits"] = sum(nulls) + row["null_total"] = len(nulls) + # surface label of top-1 under each lens: plurality over the mid-band; + # a 2-2 tie resolves to the alphabetically first label, so the vote is + # deterministic (CJK < LATIN < OTHER; DE < EN < OTHER). + for tag, rec in (("en", a), (spec.tag_b, b)): + labels = [] + for L in MIDBAND: + top5 = top5_at(rec, L) + if top5: + labels.append(spec.surface_call(top5[0][0], v)) + row[f"{col}_{tag}"] = max(sorted(set(labels)), key=labels.count) if labels else "NA" + if spec.divergence: + row["token_diverges"] = token_diverges(a, b, rule=rule) + rows.append(row) + + # null floor (b): shuffled pairing, form_a of prompt i under lens A vs form_a of j under lens B + rng = random.Random(seed) + cross_ids = [r["id"] for r in rows if r["group"] in spec.cross_groups] + shuffle_hits = shuffle_total = 0 + for rid in cross_ids: + others = [x for x in cross_ids if bank[x]["concept"] != bank[rid]["concept"]] + for oid in rng.sample(others, min(3, len(others))): + res = majority_same(records_a[rid], records_b[oid], bank[rid]["form_a"], bank[oid]["form_a"], rule=rule) + if res is not None: + shuffle_total += 1 + shuffle_hits += int(res) + + def rate(sel): + vals = [r for r in rows if sel(r) and r["cross_lens_pass"] is not None] + k = sum(r["cross_lens_pass"] for r in vals) + return k, len(vals), wilson(k, len(vals)) + + groups = sorted({r["group"] for r in rows}) + summary: dict = {"per_group": {}, "rows": rows} + for grp in groups: + k, n, (pt, lo, hi) = rate(lambda r, g=grp: r["group"] == g) + summary["per_group"][grp] = {"pass": k, "n": n, "rate": pt, "ci": [lo, hi]} + k, n, (pt, lo, hi) = rate(lambda r: r["group"] in spec.cross_groups) + summary["headline"] = {"pass": k, "n": n, "rate": pt, "ci": [lo, hi]} + for tag in ("en", spec.tag_b): + key = f"cross_form_{tag}" + cf = [r[key] for r in rows if r["group"] in spec.cross_groups and r[key] is not None] + summary[key] = {"pass": sum(cf), "n": len(cf)} + null_rows = rows if null_population == "all" else [r for r in rows if r["group"] in spec.cross_groups] + nk = sum(r["null_hits"] for r in null_rows) + nn = sum(r["null_total"] for r in null_rows) + summary["null_within_lens"] = {"pass": nk, "n": nn, "rate": wilson(nk, nn)[0]} + summary["null_shuffle_cross_lens"] = { + "pass": shuffle_hits, + "n": shuffle_total, + "rate": wilson(shuffle_hits, shuffle_total)[0], + } + if spec.divergence: + div = [r["token_diverges"] for r in rows if r["group"] in spec.cross_groups and r["token_diverges"] is not None] + summary["divergence_rate"] = {"diverging": sum(div), "n": len(div)} + la, lb = spec.split_labels + + def flips(sel): + return sum(1 for r in sel if r[f"{col}_en"] == la and r[f"{col}_{spec.tag_b}"] == lb) + + summary["lens_only_split"] = {} + for grp in groups: + sel = [r for r in rows if r["group"] == grp] + summary["lens_only_split"][grp] = {"pass": flips(sel), "n": len(sel)} + cross_sel = [r for r in rows if r["group"] in spec.cross_groups] + summary["lens_only_split"]["all_cross"] = {"pass": flips(cross_sel), "n": len(cross_sel)} + return summary + + +def print_summary(summary: dict, spec: PairSpec) -> None: + """The aggregators' stdout report, from the dict :func:`aggregate_pair` returns.""" + print("\n=== CROSS-LENS DOMINANT-FEATURE AGREEMENT (majority of mid-band) ===") + for grp, g in summary["per_group"].items(): + print(f" {grp:14s} {g['pass']:3d}/{g['n']:<3d} {g['rate']:.2f} [{g['ci'][0]:.2f}, {g['ci'][1]:.2f}]") + h = summary["headline"] + print(f" {'ALL CROSS':14s} {h['pass']:3d}/{h['n']:<3d} {h['rate']:.2f} [{h['ci'][0]:.2f}, {h['ci'][1]:.2f}]") + cf_en, cf_b = summary["cross_form_en"], summary[f"cross_form_{spec.tag_b}"] + print("\n=== ONE FEATURE CARRIES BOTH FORMS (within-lens) ===") + print(f" EN lens: {cf_en['pass']}/{cf_en['n']} {spec.tag_b.upper()} lens: {cf_b['pass']}/{cf_b['n']}") + nw, ns = summary["null_within_lens"], summary["null_shuffle_cross_lens"] + print("\n=== NULL FLOORS ===") + print(f" (a) form vs unrelated token, within-lens: {nw['pass']}/{nw['n']} ({nw['rate']:.2f})") + print(f" (b) cross-lens shuffled prompts: {ns['pass']}/{ns['n']} ({ns['rate']:.2f})") + if spec.divergence: + d = summary["divergence_rate"] + print("\n=== TOKEN DIVERGENCE (descriptive) ===") + print(f" top-1 differs between lenses on {d['diverging']}/{d['n']} cross prompts") + print(f"\n=== LANGUAGE-FOLLOWS-LENS ({spec.surface_caption}, mid-band vote) ===") + for grp, g in summary["lens_only_split"].items(): + name = "ALL CROSS" if grp == "all_cross" else grp + print(f" {name:14s} {spec.split_caption} on {g['pass']}/{g['n']}") + + +def write_summary(summary: dict, out: str | Path) -> None: + """The aggregators' summary JSON (``ensure_ascii=False, indent=1``; not ``utils.write_json``, which sorts keys).""" + out = Path(out) + out.parent.mkdir(parents=True, exist_ok=True) + with out.open("w", encoding="utf-8") as f: + json.dump(summary, f, ensure_ascii=False, indent=1) + + +def cli_args(args: Any) -> dict[str, Any]: + """``vars(args)`` without the subcommand callables (``set_defaults(fn=...)``), ready for ``run_provenance``.""" + items = args.items() if isinstance(args, dict) else vars(args).items() + return {k: v for k, v in items if not callable(v)} + + +def write_manifest(anchor: str | Path, args: Any, **fields: Any) -> Path: + """``.manifest.json`` beside an output file: ``fields`` plus ``run_provenance(args)``. + + For scripts whose primary outputs are CSVs (or a raw dump whose schema must + not grow); JSON summaries carry ``provenance`` inline instead. + """ + path = Path(anchor).with_suffix(".manifest.json") + write_json({**fields, "provenance": run_provenance(cli_args(args))}, path, atomic=True) + return path + + +# --------------------------------------------------------------------------- # +# resume sidecars (lens fitters) +# --------------------------------------------------------------------------- # + + +def prompt_list_sha1(prompts: Sequence[str]) -> str: + """Content hash of a prompt list (JSON-encoded, so newlines inside prompts cannot alias).""" + return hashlib.sha1(json.dumps(list(prompts), ensure_ascii=False).encode("utf-8")).hexdigest() + + +def check_resume_meta(meta_path: Path, expected: dict[str, Any], artifacts: Sequence[Path]) -> None: + """Refuse to resume from ``artifacts`` unless ``meta_path`` records ``expected``. + + Nothing present: returns (fresh run). Artifacts present without the sidecar: + raises (files from before the sidecar existed; delete them to refit). + Sidecar present with any key of ``expected`` different: raises, naming each + mismatching key with the on-disk and requested values. + """ + present = [p for p in artifacts if p.exists()] + if not present: + return + names = ", ".join(str(p) for p in present) + if not meta_path.exists(): + raise RuntimeError( + f"{names} exist(s) but the sidecar {meta_path} does not, so the fit they belong to cannot be " + "verified; delete them (or restore the sidecar) before resuming" + ) + stored = json.loads(meta_path.read_text()) + diffs = [f"{k}: on disk {stored.get(k)!r}, requested {v!r}" for k, v in expected.items() if stored.get(k) != v] + if diffs: + raise RuntimeError( + f"refusing to resume from {names}: {meta_path} was written for a different fit\n " + "\n ".join(diffs) + ) + + +def write_resume_meta(meta_path: Path, meta: dict[str, Any]) -> None: + write_json(meta, meta_path, atomic=True) + + +# --------------------------------------------------------------------------- # +# text folding (bank builder, EN-DE lexical call) +# --------------------------------------------------------------------------- # + + +def fold_text(s: str) -> str: + """Strip, casefold, ``ß -> ss``, then drop combining marks after NFKD (``Käse -> kase``).""" + s = s.strip().casefold().replace("ß", "ss") + s = unicodedata.normalize("NFKD", s) + return "".join(c for c in s if not unicodedata.combining(c)) diff --git a/src/sparse_readout_prism/research/row_geometry.py b/src/sparse_readout_prism/research/row_geometry.py new file mode 100644 index 0000000..5acbfd9 --- /dev/null +++ b/src/sparse_readout_prism/research/row_geometry.py @@ -0,0 +1,80 @@ +"""Row-geometry helpers shared by the direct-geometry baselines and the stability scripts. + +* :func:`spherical_kmeans_unit` -- Lloyd's k-means on unit-norm rows with cosine + assignment, the clustering behind the ``row_cluster_*`` baselines + (``scripts/run/run_readout_baseline_comparisons.py``) and the leave-one-out + core-recovery control (``scripts/eval/loo_core_recovery.py``). Seeded + initialisation and reseeding; bit-exact across runs on CPU only (CUDA + ``index_add_`` accumulates in an order-dependent way, so GPU centroids can + differ at the last ulp between runs). +* :func:`resolve_single_token_bare_first` -- the bare-then-space-prefixed + single-token resolver used by the causal-validation and stability scripts. + The fidelity runners use the space-first, special-rejecting + :func:`sparse_readout_prism.research.registry.resolve_single_token_strict`; + the two orders can pick different rows for terms where both variants are + single tokens, so scripts state which one they use. +""" + +from __future__ import annotations + +import torch + + +def spherical_kmeans_unit( + X: torch.Tensor, + n_clusters: int, + seed: int, + *, + iters: int = 12, + chunk: int = 4096, + return_assignments: bool = False, + final_assignment: bool = True, +) -> torch.Tensor | tuple[torch.Tensor, torch.Tensor]: + """Unit-normalised centroids ``(n_clusters, d)`` for unit-norm rows ``X`` ``(V, d)``. + + Dead centroids are reseeded from random rows each iteration. With + ``return_assignments=True`` an assignment ``(V,)`` is returned as well: + the cosine argmax against the returned centroids (``final_assignment=True``, + the ``row_cluster_*`` baselines' reading), or, with + ``final_assignment=False``, the assignment that produced the last centroid + update -- one Lloyd step behind the returned centroids. The paper's + leave-one-out cluster control was computed with the latter, so + ``loo_core_recovery.py`` keeps it; the two differ only while the iteration + has not converged. + """ + V = X.shape[0] + if n_clusters > V: + raise ValueError(f"n_clusters={n_clusters} exceeds the number of rows V={V}") + g = torch.Generator().manual_seed(seed) + C = X[torch.randperm(V, generator=g)[:n_clusters].to(X.device)].clone() + ones = torch.ones(V, device=X.device) + assign = torch.empty(V, dtype=torch.long, device=X.device) + for _ in range(iters): + for s in range(0, V, chunk): + assign[s : s + chunk] = (X[s : s + chunk] @ C.T).argmax(dim=1) + C_new = torch.zeros_like(C) + count = torch.zeros(n_clusters, device=X.device) + C_new.index_add_(0, assign, X) + count.index_add_(0, assign, ones) + dead = count == 0 + C = C_new / count.clamp_min(1.0)[:, None] + n_dead = int(dead.sum()) + if n_dead: + ridx = torch.randperm(V, generator=g)[:n_dead].to(X.device) + C[dead] = X[ridx] + C = C / C.norm(dim=1, keepdim=True).clamp_min(1e-8) + if return_assignments: + if final_assignment: + for s in range(0, V, chunk): + assign[s : s + chunk] = (X[s : s + chunk] @ C.T).argmax(dim=1) + return C, assign + return C + + +def resolve_single_token_bare_first(tok, term: str) -> int | None: + """Single token id for ``term``, trying the bare form before the space-prefixed form.""" + for variant in (term, " " + term): + ids = tok.encode(variant, add_special_tokens=False) + if len(ids) == 1: + return int(ids[0]) + return None diff --git a/src/sparse_readout_prism/research/seed_stability.py b/src/sparse_readout_prism/research/seed_stability.py index 735972f..96e5586 100644 --- a/src/sparse_readout_prism/research/seed_stability.py +++ b/src/sparse_readout_prism/research/seed_stability.py @@ -8,88 +8,124 @@ produce those summaries must be identical across the scripts and live here: * single-token A/B contrasts parsed from a curated JSONL bank (the bare form - is tried before the space-prefixed form, the first occurrence of a pair - wins, the list is shuffled once with the caller's generator and capped), -* rows centred against the full-vocabulary mean of ``W_U`` and per-row - normalised, -* TopK codes from a dictionary's encoder and the contrast coefficient + is tried before the space-prefixed form, via + ``row_geometry.resolve_single_token_bare_first``; the first occurrence of a + pair wins, the list is shuffled once with the caller's generator and capped), +* rows centred against an explicit centering mean and per-row normalised with + ``data.center_normalize_rows``; the mean is chosen by ``resolve_centering``, + the scripts' ``--centering {live,trained}`` switch (``live`` = the + full-vocabulary mean of ``W_U``, how the paper's runs were computed; + ``trained`` = the dictionaries' stored training mean, through + ``data.centering_mean``), +* TopK codes from a dictionary's encoder (``qwen_readout.encode_topk``, the + one top-k kernel) and the contrast coefficient ``beta = rn_A * z_A - rn_B * z_B`` for the contrast ``w_A - w_B``, * the top-M features per sign, and the token set of a feature's top-R - centred rows (decoded, stripped, lower-cased, printable strings only). - -Checkpoints are read as raw tensors (``decoder`` / ``encoder.weight`` / -``encoder.bias`` from ``model_state_dict``) so the scripts also accept the -legacy ``W_dec`` / ``W_enc`` / ``b_enc`` layout of an archived reference -dictionary. + centred rows (decoded, stripped, lower-cased, printable strings only), +* the summary statistics and decoder unit-normalisation the scripts share. + +Checkpoints are read with ``qwen_readout.load_sae`` (``load_factorizer`` +underneath: ``weights_only=True``, TopK architecture checked), which also +surfaces the stored training ``row_mean``. + +Changed in 0.2.1: + +* Top-k selection routes through ``factorizers.topk_mask`` and keeps exactly + ``k`` codes per row. The previous rule zeroed activations *below* the k-th + largest value and so kept every activation tied with it; the two agree + unless positive activations tie exactly at the boundary (zero codes + contribute nothing to ``beta`` either way), which does not happen for a + trained encoder on real rows. +* The legacy ``W_dec`` / ``W_enc`` / ``b_enc`` checkpoint layout is no longer + read. Every dictionary of the paper runs, the cross-recipe + ``--reference-dict`` included, is in the runner's ``model_state_dict`` + schema. """ from __future__ import annotations import json from pathlib import Path +from typing import NamedTuple import numpy as np import torch -Dictionary = tuple[torch.Tensor, torch.Tensor, torch.Tensor, int] +from sparse_readout_prism.data import center_normalize_rows, centering_mean, token_mask_from_tokenizer +from sparse_readout_prism.research.qwen_readout import encode_topk, load_sae +from sparse_readout_prism.research.row_geometry import resolve_single_token_bare_first -def load_dictionary(path: str | Path) -> Dictionary: - """``(decoder, encoder_w, encoder_b, k)`` as float32 CPU tensors. +class Dictionary(NamedTuple): + """One seed's dictionary as float32 CPU tensors (indexable like the former 4-tuple).""" - ``decoder`` and ``encoder_w`` are ``(d_features, d_model)``; ``k`` is read - from the checkpoint's factorizer config (default 256 when absent). - """ - ckpt = torch.load(path, map_location="cpu", weights_only=True) - if "W_dec" in ckpt: # legacy raw-tensor layout - dec = ckpt["W_dec"].float() - enc_w = ckpt["W_enc"].float().T.contiguous() - bias = ckpt.get("b_enc") - enc_b = bias.float() if bias is not None else torch.zeros(enc_w.shape[0]) - k = int(ckpt.get("k") or ckpt.get("config", {}).get("k") or 256) - return dec, enc_w, enc_b, k - state = ckpt.get("model_state_dict") or ckpt.get("state_dict") - if state is None: - raise KeyError(f"{path}: no model_state_dict / state_dict (or legacy W_dec) in checkpoint") - dec = state["decoder"].float() - enc_w = state["encoder.weight"].float() - enc_b = state.get("encoder.bias", torch.zeros(enc_w.shape[0])).float() - cfg = ckpt.get("factorizer") or ckpt.get("config", {}).get("factorizer", {}) - return dec, enc_w, enc_b, int(cfg.get("k", 256)) - - -def load_readout(path: str | Path) -> tuple[torch.Tensor, torch.Tensor | None]: - """``(W_U, h_LN)`` from an extraction payload ``{W_U_orig, h_LN, ...}``; ``h_LN`` may be None.""" + decoder: torch.Tensor # (d_features, d_model) + encoder_w: torch.Tensor # (d_features, d_model) + encoder_b: torch.Tensor # (d_features,) + k: int + row_mean: torch.Tensor | None # (d_model,) training centering mean; None for checkpoints predating it + + +def load_dictionary(path: str | Path) -> Dictionary: + """Read a TopK checkpoint through ``load_sae``; ``k`` comes from its factorizer config (default 256).""" + decoder, encoder_w, encoder_b, _config, row_mean = load_sae(path) + # load_sae hands back the training-config block, which runner checkpoints + # do not carry; k lives in the top-level factorizer block. mmap keeps this + # second open from reading the tensors again. + meta = torch.load(path, map_location="cpu", weights_only=True, mmap=True) + cfg = meta.get("factorizer") or meta.get("config", {}).get("factorizer", {}) + return Dictionary(decoder, encoder_w, encoder_b, int(cfg.get("k", 256)), row_mean) + + +def load_readout(path: str | Path) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None]: + """``(W_U, h_LN, token_mask)`` from an extraction payload; ``h_LN`` and ``token_mask`` may be None.""" payload = torch.load(path, map_location="cpu", weights_only=True) W = payload.get("W_U_orig", payload.get("W_U")) if W is None: raise KeyError(f"{path}: no W_U_orig / W_U in payload") - return W.float(), payload.get("h_LN") - + return W.float(), payload.get("h_LN"), payload.get("token_mask") -def center_rows(W: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: - """Full-vocabulary centring and per-row normalisation: ``(W_c, row_norms, W_n)``.""" - W_c = W - W.mean(0) - rn = W_c.norm(dim=1).clamp_min(1e-8) - return W_c, rn, W_c / rn[:, None] +def resolve_centering( + W: torch.Tensor, dicts: list[Dictionary], mode: str, token_mask: torch.Tensor | None, tok=None +) -> torch.Tensor: + """Centering mean of a dictionary family under ``--centering {live,trained}``. -def topk_codes(x: torch.Tensor, enc_w: torch.Tensor, enc_b: torch.Tensor, k: int) -> torch.Tensor: - """ReLU encoder activations, zeroed below the k-th largest value of each row.""" - acts = torch.relu(x @ enc_w.T + enc_b) - if k < acts.shape[-1]: - thresh = acts.topk(k, dim=-1).values[..., -1:] - acts = torch.where(acts >= thresh, acts, torch.zeros_like(acts)) - return acts - - -def single_token_id(tok, term: str) -> int | None: - """Token id of ``term`` if it (or ``" " + term``) is a single token, else None.""" - for variant in (term, " " + term): - ids = tok.encode(variant, add_special_tokens=False) - if len(ids) == 1: - return int(ids[0]) - return None + ``trained`` uses the stored training mean of the dictionaries (one recipe, + so every checkpoint must store the same tensor; a disagreement is refused) + and otherwise the text-token mean over ``token_mask`` (the payload's, or + one rebuilt from ``tok`` when the payload predates the mask). ``live`` is + the full-vocabulary mean of ``W``. + """ + stored = [d.row_mean for d in dicts if d.row_mean is not None] + if mode == "trained": + if any(not torch.equal(stored[0], m) for m in stored[1:]): + raise ValueError("dictionaries store different training row_mean tensors; pass checkpoints from one recipe") + if token_mask is None and tok is not None: + token_mask = token_mask_from_tokenizer(tok, W.shape[0]) + ckpt = {"row_mean": stored[0]} if stored else None + return centering_mean(W, mode=mode, token_mask=token_mask, ckpt=ckpt) + + +def center_rows(W: torch.Tensor, row_mean: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: + """Centring against ``row_mean`` and per-row normalisation: ``(W_c, row_norms, W_n)``.""" + rn, W_n = center_normalize_rows(W, row_mean) + return W - row_mean, rn, W_n + + +def unit_rows(d: torch.Tensor) -> torch.Tensor: + """Rows of ``d`` scaled to unit norm (1e-8 floor).""" + return d / d.norm(dim=1, keepdim=True).clamp_min(1e-8) + + +def summary_stats(v, percentiles: tuple[int, ...]) -> dict: + """``{mean, median, p..., n}`` of ``v``, in the key order the scripts' JSON outputs use.""" + v = np.asarray(v, dtype=float) + out = dict(mean=float(v.mean()), median=float(np.median(v))) + for q in percentiles: + out[f"p{q}"] = float(np.percentile(v, q)) + out["n"] = int(len(v)) + return out def load_contrast_pairs(bank_path: str | Path, tok, rng: np.random.Generator, max_contrasts: int) -> list[tuple]: @@ -101,7 +137,7 @@ def load_contrast_pairs(bank_path: str | Path, tok, rng: np.random.Generator, ma a, b = r.get("target_a"), r.get("target_b") if not a or not b or (a, b) in seen: continue - ia, ib = single_token_id(tok, a), single_token_id(tok, b) + ia, ib = resolve_single_token_bare_first(tok, a), resolve_single_token_bare_first(tok, b) if ia is None or ib is None or ia == ib: continue seen.add((a, b)) @@ -114,9 +150,8 @@ def contrast_features( W_n: torch.Tensor, rn: torch.Tensor, dictionary: Dictionary, ia: int, ib: int, top_m: int ) -> tuple[list[int], list[int]]: """Top-M positive and top-M negative feature ids of the contrast row ``ia`` minus row ``ib``.""" - _, enc_w, enc_b, k = dictionary - z = topk_codes(torch.stack([W_n[ia], W_n[ib]]), enc_w, enc_b, k) # (2, d_features) - beta = rn[ia] * z[0] - rn[ib] * z[1] + z = encode_topk(torch.stack([W_n[ia], W_n[ib]]), dictionary.encoder_w, dictionary.encoder_b, dictionary.k) + beta = rn[ia] * z[0] - rn[ib] * z[1] # (d_features,) return torch.topk(beta, top_m).indices.tolist(), torch.topk(-beta, top_m).indices.tolist() @@ -138,7 +173,12 @@ def token_strings(tok, ids, cache: dict[int, str] | None = None) -> set[str]: def feature_token_set( fids, dec: torch.Tensor, W_c: torch.Tensor, tok, top_r: int, cache: dict[int, str] | None = None ) -> set[str]: - """Union over ``fids`` of the token strings of each feature's top-R centred rows.""" + """Union over ``fids`` of the token strings of each feature's top-R centred rows. + + One matrix-vector product per feature, as in the paper's runs; a single + GEMM over the unique feature ids would not reproduce the same top-R + orderings bit for bit (different accumulation order), so it is not used. + """ out: set[str] = set() for f in fids: rows = torch.topk(W_c @ dec[f], top_r).indices.tolist() diff --git a/src/sparse_readout_prism/research/wsd.py b/src/sparse_readout_prism/research/wsd.py new file mode 100644 index 0000000..992dc9c --- /dev/null +++ b/src/sparse_readout_prism/research/wsd.py @@ -0,0 +1,220 @@ +"""CoarseWSD-20 bundle and statistics helpers for the sense-labelled evaluation. + +Shared by ``scripts/run/run_wsd_feature_alignment.py`` (writes the +``representations.pt`` bundle and the full-vector centroid references), +``scripts/analyze/analyze_wsd_sense_groups.py`` (``tab:app-sense-alignment``) +and ``scripts/analyze/analyze_wsd_classifier_framing.py`` (the classifier +framing paragraph). Before 0.2.1 each script carried its own copy of the +bundle reader, the per-word train/test split rule, the row-coefficient shuffle +null and the bootstrap helpers; the definitions here are those copies merged +without changing a number (``tests/test_wsd_sense_groups.py`` pins the +pre-merge outputs). + +Bundle layout (``representations.pt``): ``metadata`` (one dict per scored +context: ``item_id``, ``split``, ``target``, ``sense``, ``token_id``, +``row_relative_error``, ...), ``run`` (dataset, model, checkpoint, k, ...) and +row-aligned tensors -- ``hidden`` (n, d_model), ``projection`` / +``contribution`` / ``beta`` (n, k) in the target's feature order, +``feature_ids`` (n, k), ``exact_logit`` / ``reconstructed_logit`` / +``target_logprob`` (n,) and ``target_rank`` (n,). +""" + +from __future__ import annotations + +import hashlib +from collections import defaultdict +from dataclasses import dataclass +from pathlib import Path +from typing import Any, Callable + +import numpy as np +import torch + +BUNDLE_TENSOR_KEYS = ( + "hidden", + "projection", + "contribution", + "beta", + "feature_ids", + "exact_logit", + "reconstructed_logit", + "target_logprob", + "target_rank", +) + + +def load_bundle(path: str | Path) -> dict[str, Any]: + """Read a ``representations.pt`` bundle and validate it. + + ``weights_only=True``: the bundle holds tensors, the metadata list of + primitive dicts and the ``run`` dict, so it loads without unpickling code. + """ + bundle = torch.load(path, map_location="cpu", weights_only=True) + validate_bundle(bundle) + return bundle + + +def validate_bundle(bundle: dict[str, Any]) -> None: + """Raise if the bundle is not a row-aligned CoarseWSD-20 bundle. + + Checks the required keys, that every tensor has one row per metadata + entry, that the bundle is a CoarseWSD-20 run (the AmbiStory mode was + removed in 0.2.1) and that the feature coordinates do not change within a + target -- the analyses compare columns across a word's rows. + """ + missing = [key for key in ("metadata", "run", *BUNDLE_TENSOR_KEYS) if key not in bundle] + if missing: + raise KeyError(f"bundle is missing {missing}") + n_rows = len(bundle["metadata"]) + for key in BUNDLE_TENSOR_KEYS: + if int(bundle[key].shape[0]) != n_rows: + raise ValueError(f"bundle[{key!r}] has {int(bundle[key].shape[0])} rows for {n_rows} metadata entries") + dataset = bundle["run"].get("dataset") + if dataset != "coarsewsd20": + raise ValueError(f"unsupported dataset {dataset!r}: only CoarseWSD-20 bundles are supported") + first_by_target: dict[str, np.ndarray] = {} + for row, ids in zip(bundle["metadata"], bundle["feature_ids"].numpy()): + target = str(row["target"]) + if target in first_by_target: + if not np.array_equal(first_by_target[target], ids): + raise ValueError(f"feature coordinates change within target {target!r}") + else: + first_by_target[target] = ids.copy() + + +@dataclass(frozen=True) +class WordSplit: + """Row indices and sense labels of one word's train and test contexts.""" + + word: str + train_idx: np.ndarray # (n_train,) bundle row indices + test_idx: np.ndarray # (n_test,) + train_y: np.ndarray # (n_train,) sense labels + test_y: np.ndarray # (n_test,) + senses: list # sorted train senses + + +def word_splits(metadata: list[dict[str, Any]], keep: np.ndarray | None = None) -> list[WordSplit]: + """Per-word train/test splits in sorted word order, under the rule all three scripts apply. + + A word is skipped when either split is empty, when fewer than two senses + occur in train, or when a test sense is absent from train. ``keep`` (one + bool per bundle row) restricts the rows first; the classifier framing + passes its score gate. + """ + words = sorted({str(row["target"]) for row in metadata}) + splits: list[WordSplit] = [] + for word in words: + train_idx = np.array( + [ + i + for i, row in enumerate(metadata) + if str(row["target"]) == word and row["split"] == "train" and (keep is None or keep[i]) + ], + dtype=np.int64, + ) + test_idx = np.array( + [ + i + for i, row in enumerate(metadata) + if str(row["target"]) == word and row["split"] == "test" and (keep is None or keep[i]) + ], + dtype=np.int64, + ) + if len(train_idx) == 0 or len(test_idx) == 0: + continue + train_y = np.array([str(metadata[i]["sense"]) for i in train_idx]) + test_y = np.array([str(metadata[i]["sense"]) for i in test_idx]) + senses = sorted(set(train_y)) + if len(senses) < 2 or not set(test_y).issubset(set(senses)): + continue + splits.append(WordSplit(word, train_idx, test_idx, train_y, test_y, senses)) + return splits + + +def stable_seed(text: str, seed: int) -> int: + """Per-target RNG seed: ``(int(sha1(text)[:8], 16) + seed) % 2**32``. + + The classifier framing's form. The run script previously used + ``seed + int(sha1(text)[:8], 16)`` without the modulus; the two agree for + ``seed=0`` (the paper's runs) and whenever the sum stays below ``2**32``. + """ + digest = hashlib.sha1(text.encode("utf-8")).hexdigest()[:8] + return (int(digest, 16) + seed) % (2**32) + + +def shuffled_srp(projection: np.ndarray, beta: np.ndarray, metadata: list[dict[str, Any]], seed: int) -> np.ndarray: + """Row-coefficient shuffle null: ``projection * beta[:, permutation]`` with one permutation per target. + + ``projection`` and ``beta`` are the bundle's (n, k) matrices; the + permutation of each target's columns is drawn from + ``default_rng(stable_seed(target, seed))``. + """ + shuffled = np.empty_like(projection) + by_target: dict[str, list[int]] = defaultdict(list) + for i, row in enumerate(metadata): + by_target[str(row["target"])].append(i) + for target in sorted(by_target): + rng = np.random.default_rng(stable_seed(target, seed)) + permutation = rng.permutation(beta.shape[1]) + idx = np.asarray(by_target[target]) + shuffled[idx] = projection[idx] * beta[idx][:, permutation] + return shuffled + + +def l2_normalize(matrix: np.ndarray) -> np.ndarray: + norms = np.linalg.norm(matrix, axis=1, keepdims=True) + return matrix / np.maximum(norms, 1e-12) + + +def safe_spearman(x: Any, y: Any) -> float: + """Spearman rho, nan below three points or when either input's range is under 1e-8. + + Differs from ``utils.spearman`` only in the degeneracy guard: this one + tests the range (``np.ptp < 1e-8``), ``utils.spearman`` the standard + deviation (``< 1e-12``), so inputs varying by less than 1e-8 are nan here + and a correlation there. Kept as the WSD scripts' guard so their outputs do + not move. + """ + from scipy.stats import spearmanr + + if len(x) < 3 or np.ptp(x) < 1e-8 or np.ptp(y) < 1e-8: + return float("nan") + return float(spearmanr(x, y).statistic) + + +def cluster_bootstrap( + rows: list[dict[str, Any]], + cluster_key: str, + stat_fn: Callable[[list[dict[str, Any]]], float], + n_boot: int, + seed: int, +) -> list[float]: + """Resample clusters (rows grouped by ``cluster_key``) with replacement; non-finite statistics are dropped.""" + by_cluster: dict[str, list[dict[str, Any]]] = defaultdict(list) + for row in rows: + by_cluster[str(row[cluster_key])].append(row) + keys = sorted(by_cluster) + rng = np.random.default_rng(seed) + values: list[float] = [] + for _ in range(n_boot): + sample: list[dict[str, Any]] = [] + for selected in rng.choice(keys, size=len(keys), replace=True): + sample.extend(by_cluster[str(selected)]) + value = float(stat_fn(sample)) + if np.isfinite(value): + values.append(value) + return values + + +def bootstrap_mean(values: np.ndarray, rng: np.random.Generator, n_boot: int) -> list[float]: + """``n_boot`` means of with-replacement resamples of ``values`` (the per-word bootstrap).""" + return [float(rng.choice(values, size=len(values), replace=True).mean()) for _ in range(n_boot)] + + +def percentile_ci(values: Any) -> list[float]: + """2.5 / 97.5 percentiles of a bootstrap sample; ``[nan, nan]`` when empty.""" + array = np.asarray(values, dtype=float) + if array.size == 0: + return [float("nan"), float("nan")] + return [float(np.percentile(array, 2.5)), float(np.percentile(array, 97.5))] diff --git a/tests/conftest.py b/tests/conftest.py new file mode 100644 index 0000000..5b9be8e --- /dev/null +++ b/tests/conftest.py @@ -0,0 +1,27 @@ +"""Shared test helpers. + +``load_script`` imports a ``scripts/`` entry point by path the way +``test_scripts_importable.py`` does: the module is registered in +``sys.modules`` before execution so dataclass decorators can resolve +``cls.__module__``. Test modules use it as ``from conftest import load_script``. +""" + +from __future__ import annotations + +import importlib.util +import sys +from pathlib import Path +from types import ModuleType + +REPO_ROOT = Path(__file__).resolve().parents[1] + + +def load_script(rel_path: str, name: str | None = None) -> ModuleType: + path = REPO_ROOT / rel_path + mod_name = name or "_test_script_" + path.stem + spec = importlib.util.spec_from_file_location(mod_name, path) + assert spec is not None and spec.loader is not None, rel_path + module = importlib.util.module_from_spec(spec) + sys.modules[mod_name] = module + spec.loader.exec_module(module) + return module diff --git a/tests/test_baseline_and_dla.py b/tests/test_baseline_and_dla.py index 4ffcdbd..13d950e 100644 --- a/tests/test_baseline_and_dla.py +++ b/tests/test_baseline_and_dla.py @@ -3,7 +3,9 @@ - scripts/run/run_readout_baseline_comparisons.py — the null/baseline methods behind the paper's "Sparse RP beats every null" claim. Key property: every method preserves the EXACT additive identity (so the comparison is fair) while - the nulls change only the support. + the nulls change only the support. Also covers the 0.2.1 consolidation: the + shared ``_finish`` tail, the memoised k-means fit, and ``margin_from_rows`` + aligning per-row supports before subtracting. - scripts/figures/compute_prism_dla.py — the feature-resolved DLA accounting. build_rows with top_features=0 skips labelling (no tokenizer needed), so we can assert the component identity (margin == sum of component direct contributions) @@ -12,29 +14,28 @@ from __future__ import annotations -import importlib.util -import sys -from pathlib import Path - import pytest import torch +from conftest import load_script -ROOT = Path(__file__).resolve().parents[1] -sys.path.insert(0, str(ROOT / "src")) - - -def _load(relpath: str, name: str): - spec = importlib.util.spec_from_file_location(name, ROOT / relpath) - mod = importlib.util.module_from_spec(spec) - sys.modules[name] = mod # dataclass(KW_ONLY) resolution needs the module registered - spec.loader.exec_module(mod) - return mod +from sparse_readout_prism.factorizers import TopKSAE +baseline = load_script("scripts/run/run_readout_baseline_comparisons.py", "baseline_comparisons") +dla = load_script("scripts/figures/compute_prism_dla.py", "compute_prism_dla") -baseline = _load("scripts/run/run_readout_baseline_comparisons.py", "baseline_comparisons") -dla = _load("scripts/figures/compute_prism_dla.py", "compute_prism_dla") +CPU = torch.device("cpu") -from sparse_readout_prism.factorizers import TopKSAE # noqa: E402 +# spec -> (expected number of active atoms or None, expected feature_space_size) +METHOD_PANEL = { + "sparse_rp": (None, 32), # d_features + "shuffled_row_code": (4, 32), # nulls keep the SAE's k-sparsity, only the support/assignment changes + "random_support_same_magnitudes": (4, 32), + "pca_4": (4, 4), # every component active; the basis is shared by all rows + "nearest_row_ridge_top4": (4, 24), # the four nearest rows, indices into the vocabulary + "knn_basis_top4": (4, 24), + "row_cluster_d16_k4": (4, 16), # top-4 of 16 centroids + "row_cluster_hard_d16": (1, 16), # the row's own centroid +} def _make_sae_inputs(vocab: int = 24, d_model: int = 8): @@ -45,36 +46,115 @@ def _make_sae_inputs(vocab: int = 24, d_model: int = 8): return W, row_mean, sae +def _build(spec: str, W, row_mean, sae, **kw): + torch.manual_seed(1) # pca_*: torch.svd_lowrank draws its test matrix from the global RNG + return baseline.build_method(spec, W=W, row_mean=row_mean, device=CPU, sae=sae, k=4, seed=0, **kw) + + def test_baseline_methods_preserve_additive_identity() -> None: W, row_mean, sae = _make_sae_inputs() - device = torch.device("cpu") h = torch.randn(8) - # spec -> expected number of active atoms (None = not checked). - expected_active = { - "sparse_rp": None, - "shuffled_row_code": 4, # nulls keep the SAE's k-sparsity, only the support/assignment changes - "random_support_same_magnitudes": 4, - "knn_basis_top4": 4, # the four nearest rows - "row_cluster_d16_k4": 4, # top-4 of 16 centroids - "row_cluster_hard_d16": 1, # the row's own centroid - } - for spec, n_active in expected_active.items(): - method = baseline.build_method(spec, W=W, row_mean=row_mean, device=device, sae=sae, k=4, seed=0) - dec = method.decompose_row(h, row_idx=3) + for spec, (n_active, space) in METHOD_PANEL.items(): + method = _build(spec, W, row_mean, sae) + dec = method.decompose_row(h, row_idx=3, exclude_ids=[9]) # original_logit == base + feature_sum + residual to float precision. assert dec.identity_error.abs().item() < 1e-4, spec assert torch.isfinite(dec.reconstructed_logit), spec + assert dec.feature_space_size == space, spec if n_active is not None: assert dec.active_feature_indices.numel() == n_active, spec +def test_finish_tail_identity_for_every_method() -> None: + """The shared ``_finish`` reproduces the accounting each method used to inline: + exact base / original terms, feature_sum as the sum of the per-atom + contributions, and contributions that index the declared feature space.""" + W, row_mean, sae = _make_sae_inputs() + h = torch.randn(8) + for spec, (_, space) in METHOD_PANEL.items(): + method = _build(spec, W, row_mean, sae) + for row in (3, 17): + dec = method.decompose_row(h, row_idx=row, exclude_ids=[9]) + assert torch.equal(dec.base_term, h @ row_mean), spec + assert torch.equal(dec.original_logit, h @ W[row]), spec + assert torch.equal(dec.reconstructed_logit, dec.base_term + dec.feature_sum), spec + assert torch.equal( + dec.identity_error, dec.original_logit - (dec.base_term + dec.feature_sum + dec.residual_term) + ), spec + active = dec.active_feature_indices + assert active.device.type == "cpu" and active.numel() == active.unique().numel(), spec + assert int(active.max()) < space, spec + fc = dec.feature_contributions + if fc.numel() == space: # full-size vector: nonzeros live on the support + assert torch.allclose(fc.sum(), dec.feature_sum, atol=1e-5), spec + if spec != "pca_4": + off = torch.ones(space, dtype=torch.bool) + off[active] = False + assert torch.equal(fc[off], torch.zeros(int(off.sum()))), spec + else: # per-row support: one entry per active atom, in support order + assert fc.numel() == active.numel(), spec + assert torch.equal(dec.feature_sum, fc.sum()), spec + + +def test_row_cluster_fit_is_memoised_by_clusters_and_seed() -> None: + W, row_mean, sae = _make_sae_inputs() + cache: dict = {} + soft = _build("row_cluster_d16_k4", W, row_mean, sae, kmeans_cache=cache) + hard = _build("row_cluster_hard_d16", W, row_mean, sae, kmeans_cache=cache) + assert hard.C is soft.C and list(cache) == [(16, 0)] + other_seed = baseline.build_method( + "row_cluster_hard_d16", W=W, row_mean=row_mean, device=CPU, seed=1, kmeans_cache=cache + ) + assert other_seed.C is not soft.C and set(cache) == {(16, 0), (16, 1)} + fresh = _build("row_cluster_d16_k4", W, row_mean, sae) # no cache: refits, bit-exact on CPU + assert torch.equal(fresh.C, soft.C) + h = torch.randn(8) + assert torch.equal(fresh.decompose_row(h, 3).feature_contributions, soft.decompose_row(h, 3).feature_contributions) + + +def test_margin_from_rows_aligns_per_row_supports() -> None: + W, row_mean, sae = _make_sae_inputs() + h = torch.randn(8) + a, b = 3, 9 + knn = _build("knn_basis_top4", W, row_mean, sae) + dA = knn.decompose_row(h, a, exclude_ids=[b]) + dB = knn.decompose_row(h, b, exclude_ids=[a]) + mr = baseline.margin_from_rows(h, knn, [a], [b]) + fm = mr["feat_margin"] + # one entry per vocabulary row, A's neighbours positive-side, B's negated, by row index + assert fm.numel() == W.shape[0] == dA.feature_space_size + expected = torch.zeros(W.shape[0]) + expected.scatter_add_(0, dA.active_feature_indices, dA.feature_contributions) + expected.scatter_add_(0, dB.active_feature_indices, -dB.feature_contributions) + assert torch.allclose(fm, expected, atol=1e-6) + support = set(dA.active_feature_indices.tolist()) | set(dB.active_feature_indices.tolist()) + assert len(support) > 4 # the two rows have different neighbourhoods on this fixture + assert set(torch.nonzero(fm).flatten().tolist()) <= support + # the margin scalars never depended on the alignment: sparse == s_A - s_B, and the aligned + # vector still sums to it (the row_mean base cancels) + assert mr["sparse"] == pytest.approx(float(dA.reconstructed_logit - dB.reconstructed_logit), abs=1e-5) + assert float(fm.sum()) == pytest.approx(mr["sparse"], abs=1e-5) + # ridge and cluster methods scatter the same way; shared-basis methods pass through untouched + ridge_fm = baseline.margin_from_rows(h, _build("nearest_row_ridge_top4", W, row_mean, sae), [a], [b])["feat_margin"] + assert ridge_fm.numel() == W.shape[0] + hard_fm = baseline.margin_from_rows(h, _build("row_cluster_hard_d16", W, row_mean, sae), [a], [b])["feat_margin"] + assert hard_fm.numel() == 16 and int((hard_fm != 0).sum()) <= 2 + assert baseline.margin_from_rows(h, _build("pca_4", W, row_mean, sae), [a], [b])["feat_margin"].numel() == 4 + assert baseline.margin_from_rows(h, _build("sparse_rp", W, row_mean, sae), [a], [b])["feat_margin"].numel() == 32 + # the pseudo-target (B = row_mean) path pads B with zeros of the aligned length + fm_mean = baseline.margin_from_rows(h, knn, [a], None, mean_row=row_mean)["feat_margin"] + assert fm_mean.numel() == W.shape[0] + assert torch.allclose( + fm_mean, torch.zeros(W.shape[0]).scatter_add(0, dA.active_feature_indices, dA.feature_contributions), atol=1e-6 + ) + + def test_build_method_rejects_bad_specs() -> None: W, row_mean, _ = _make_sae_inputs() - device = torch.device("cpu") with pytest.raises(ValueError): - baseline.build_method("not_a_method", W=W, row_mean=row_mean, device=device) + baseline.build_method("not_a_method", W=W, row_mean=row_mean, device=CPU) with pytest.raises(ValueError): # sparse_rp needs sae + k - baseline.build_method("sparse_rp", W=W, row_mean=row_mean, device=device) + baseline.build_method("sparse_rp", W=W, row_mean=row_mean, device=CPU) def test_prism_dla_component_and_reconstruction_identities() -> None: diff --git a/tests/test_causal_contribution_validation.py b/tests/test_causal_contribution_validation.py index 85c8cda..b0efb8b 100644 --- a/tests/test_causal_contribution_validation.py +++ b/tests/test_causal_contribution_validation.py @@ -9,38 +9,31 @@ explanatory power (the through-origin fit is undefined for them), * the additive identity holds: the contributions sum to the exact margin, each stored realized change equals (h . d_i)(q . d_i), and ablating h along d_i - moves the dense margin by exactly minus that amount. + moves the margin read off a dense two-row LM head by exactly minus that amount + (the self-test's third check, measured from the head rather than the formula), +* ``--centering`` selects the live full-vocabulary mean, the checkpoint's + ``row_mean``, or the tokenizer text-token mean, through ``data.centering_mean``. CPU only, no model, well under a second. """ from __future__ import annotations -import importlib.util import math -import sys -from pathlib import Path import numpy as np +import torch +from conftest import load_script -ROOT = Path(__file__).resolve().parents[1] -sys.path.insert(0, str(ROOT / "src")) +from sparse_readout_prism.research import row_geometry - -def _load(relpath: str, name: str): - spec = importlib.util.spec_from_file_location(name, ROOT / relpath) - mod = importlib.util.module_from_spec(spec) - sys.modules[name] = mod - spec.loader.exec_module(mod) - return mod - - -causal = _load("scripts/eval/run_causal_contribution_validation.py", "causal_contribution_validation") +causal = load_script("scripts/eval/run_causal_contribution_validation.py") def test_self_test_passes_and_recovers_unit_agreement() -> None: m = causal.self_test_metrics() assert all(ok for _, ok in m["checks"]), m["checks"] + assert len(m["checks"]) == 3 # Residual-free contrast: the fit through the origin recovers the prediction. assert abs(m["r2"] - 1.0) < 1e-2 assert abs(m["slope"] - 1.0) < 0.1 @@ -79,6 +72,26 @@ def test_additive_identity_and_ablation_sign() -> None: assert abs(c_pred - float(beta[fid] * (d_i @ h))) < 1e-5 +def test_third_check_measures_the_dense_head_not_the_formula() -> None: + m = causal.self_test_metrics() + W_head, h, q = m["W_head"], m["h"], m["q"] + # The synthetic head is a genuine two-row readout whose row difference is q. + assert W_head.shape == (2, h.numel()) + assert torch.allclose(W_head[0] - W_head[1], q, atol=1e-6) + assert m["dense_head_abs_err"] < 1e-4 + m_base = causal.dense_head_margin(W_head, h) + assert abs(m_base - float(h @ q)) < 1e-4 + for fid, _c_pred, delta_real, _is_random in m["pairs"]: + d_i = m["W_dec"][fid] + realized = causal.dense_head_margin(W_head, h - (h @ d_i) * d_i) - m_base + assert abs(realized + delta_real) < 1e-4 + # A wrong stored value is caught: perturbing delta_real breaks the dense-head comparison. + fid, _c, delta_real, _r = m["pairs"][0] + d_i = m["W_dec"][fid] + wrong = delta_real + 0.5 + assert abs((causal.dense_head_margin(W_head, h - (h @ d_i) * d_i) - m_base) + wrong) > 0.4 + + def test_fit_r2_slope_orientation() -> None: # y = 2x exactly: slope is realized-on-predicted, so it must read 2 (not 0.5). x = np.array([1.0, 2.0, 3.0, 4.0]) @@ -96,3 +109,33 @@ def test_cluster_bootstrap_ci_brackets_point_estimate() -> None: assert lo > 0.9 # Same seed, same interval. assert (lo, hi) == causal.cluster_bootstrap_r2(rows, n_boot=50, seed=0) + + +class _FakeTokenizer: + """Vocabulary of ``n`` ids with the last rows of W_U past the tokenizer and id 0 special.""" + + def __init__(self, n: int): + self.n = n + self.all_special_ids = [0] + + def __len__(self) -> int: + return self.n + + +def test_select_row_mean_modes() -> None: + gen = torch.Generator().manual_seed(0) + W_U = torch.randn(10, 4, generator=gen) + tok = _FakeTokenizer(8) # rows 8, 9 are padded embedding rows; row 0 is special + live, src = causal.select_row_mean("live", W_U, ckpt={}, tok=tok) + assert torch.equal(live, W_U.mean(dim=0)) and "live" in src + stored = torch.randn(4, generator=gen) + trained_ckpt, src = causal.select_row_mean("trained", W_U, ckpt={"row_mean": stored}, tok=tok) + assert torch.equal(trained_ckpt, stored) and "checkpoint" in src + trained_mask, src = causal.select_row_mean("trained", W_U, ckpt={}, tok=tok) + assert torch.equal(trained_mask, W_U[1:8].mean(dim=0)) and "text-token" in src + assert not torch.equal(trained_mask, live) + + +def test_token_resolver_is_the_shared_bare_first_one() -> None: + assert causal.resolve_single_token_bare_first is row_geometry.resolve_single_token_bare_first + assert not hasattr(causal, "single_token_id") diff --git a/tests/test_cross_lens_scripts.py b/tests/test_cross_lens_scripts.py index 6e4814b..ca38569 100644 --- a/tests/test_cross_lens_scripts.py +++ b/tests/test_cross_lens_scripts.py @@ -1,38 +1,41 @@ -"""CPU-only checks for the cross-lens study scripts (no GPU, no model, no jlens). - -Covers the offline aggregation and table-assembly paths on synthetic dumps: - * scripts/eval/aggregate_cross_lens_en_zh.py (wilson, script_of, mid-band vote, end-to-end) - * scripts/eval/aggregate_cross_lens_en_de.py (lexical language call, divergence) - * scripts/data/build_cross_lens_de_bank.py (normalisation, Levenshtein, bank rules) +"""CPU-only checks for the cross-lens study scripts and their shared toolkit (no GPU, no model, no jlens). + +Covers, on synthetic dumps and tensors: + * sparse_readout_prism.research.cross_lens (both majority rules, aggregation core and its key order, + decomposition pinned to the pre-0.2.1 inline formula, raw-decode labels, k / centering resolution, + final-norm lookup, resume sidecars, text folding) + * scripts/eval/aggregate_cross_lens_en_zh.py (script_of, end-to-end, --agreement-rule) + * scripts/eval/aggregate_cross_lens_en_de.py (lexical call, divergence, --null-population) + * scripts/data/build_cross_lens_de_bank.py (Levenshtein, bank rules) * scripts/analyze/cross_lens_three_lens_prompt.py * scripts/figures/compute_cross_lens_shared_feature.py * scripts/analyze/cross_lens_antonym_layers.py (table subcommand) - * scripts/run/fit_jlens.py (source-layer picker) + * scripts/run/fit_jlens.py (source-layer picker, dim-batch fallback, shard sidecars) + * scripts/run/fit_ridge_lens.py (--holdout guard, diagnostic prompt slice) """ from __future__ import annotations -import importlib.util import json from pathlib import Path +from types import SimpleNamespace -ROOT = Path(__file__).resolve().parents[1] - - -def _load(mod_name: str, rel_path: str): - spec = importlib.util.spec_from_file_location(mod_name, ROOT / rel_path) - mod = importlib.util.module_from_spec(spec) - spec.loader.exec_module(mod) - return mod +import pytest +import torch +import torch.nn.functional as F +from conftest import load_script +from torch import nn +from sparse_readout_prism.research import cross_lens -agg_zh = _load("aggregate_cross_lens_en_zh", "scripts/eval/aggregate_cross_lens_en_zh.py") -agg_de = _load("aggregate_cross_lens_en_de", "scripts/eval/aggregate_cross_lens_en_de.py") -bank_de = _load("build_cross_lens_de_bank", "scripts/data/build_cross_lens_de_bank.py") -three = _load("cross_lens_three_lens_prompt", "scripts/analyze/cross_lens_three_lens_prompt.py") -shared = _load("compute_cross_lens_shared_feature", "scripts/figures/compute_cross_lens_shared_feature.py") -antonym = _load("cross_lens_antonym_layers", "scripts/analyze/cross_lens_antonym_layers.py") -fit_jlens = _load("fit_jlens", "scripts/run/fit_jlens.py") +agg_zh = load_script("scripts/eval/aggregate_cross_lens_en_zh.py") +agg_de = load_script("scripts/eval/aggregate_cross_lens_en_de.py") +bank_de = load_script("scripts/data/build_cross_lens_de_bank.py") +three = load_script("scripts/analyze/cross_lens_three_lens_prompt.py") +shared = load_script("scripts/figures/compute_cross_lens_shared_feature.py") +antonym = load_script("scripts/analyze/cross_lens_antonym_layers.py") +fit_jlens = load_script("scripts/run/fit_jlens.py") +fit_ridge = load_script("scripts/run/fit_ridge_lens.py") MIDBAND = ["21", "24", "26", "29"] @@ -78,6 +81,11 @@ def _record(pid, group, top1, dom_by_target, top1_feats=None): } +def _set_dom(rec, layers, target, fid): + for L in layers: + rec["layers"][L]["-1"]["targets"][target]["top_features"][0]["id"] = fid + + def _dump(records, labels=None): return { "model": "m", @@ -90,16 +98,22 @@ def _dump(records, labels=None): } +def _write(tmp_path, name, payload): + p = tmp_path / name + p.write_text(json.dumps(payload, ensure_ascii=False)) + return str(p) + + # --------------------------------------------------------------------------- # -# EN-ZH aggregator +# votes and rules # --------------------------------------------------------------------------- # def test_wilson_matches_reference_values(): - p, lo, hi = agg_zh.wilson(77, 80) + p, lo, hi = cross_lens.wilson(77, 80) assert abs(p - 0.9625) < 1e-9 assert round(lo, 2) == 0.90 and round(hi, 2) == 0.99 - assert agg_zh.wilson(0, 0) == (0.0, 0.0, 1.0) + assert cross_lens.wilson(0, 0) == (0.0, 0.0, 1.0) def test_script_of_distinguishes_scripts(): @@ -108,20 +122,43 @@ def test_script_of_distinguishes_scripts(): assert agg_zh.script_of(" 42") == "OTHER" -def test_majority_same_uses_midband_vote(): +def test_majority_same_half_and_strict_rules(): a = _record("p", "g", "t", {"f": 7}) b = _record("p", "g", "t", {"f": 7}) - assert agg_zh.majority_same(a, b, "f", "f") is True - # flip two of four mid-band layers: 2/4 still counts as a majority (>= ceil(n/2)) + assert cross_lens.majority_same(a, b, "f", "f", rule="half") is True + assert cross_lens.majority_same(a, b, "f", "f", rule="strict") is True + # 2 of 4 agreeing layers: half (>= ceil(n/2), the paper) passes, strict (> n/2) does not + _set_dom(b, ["21", "24"], "f", 99) + assert cross_lens.majority_same(a, b, "f", "f", rule="half") is True + assert cross_lens.majority_same(a, b, "f", "f", rule="strict") is False + # 3 of 4 agreeing: both pass + _set_dom(b, ["24"], "f", 7) + assert cross_lens.majority_same(a, b, "f", "f", rule="strict") is True + # 1 of 4 agreeing: neither + _set_dom(b, ["24", "26"], "f", 99) + assert cross_lens.majority_same(a, b, "f", "f", rule="half") is False + assert cross_lens.majority_same(a, b, "missing", "f", rule="half") is None + with pytest.raises(ValueError, match="agreement rule"): + cross_lens.majority_same(a, b, "f", "f", rule="most") + + +def test_token_diverges_rules(): + a = _record("p", "g", " dog", {"f": 1}) + b = _record("p", "g", " dog", {"f": 1}) + assert cross_lens.token_diverges(a, b, rule="half") is False for L in ["21", "24"]: - b["layers"][L]["-1"]["targets"]["f"]["top_features"][0]["id"] = 99 - assert agg_zh.majority_same(a, b, "f", "f") is True - b["layers"]["26"]["-1"]["targets"]["f"]["top_features"][0]["id"] = 99 - assert agg_zh.majority_same(a, b, "f", "f") is False - assert agg_zh.majority_same(a, b, "missing", "f") is None + b["layers"][L]["-1"]["top5"][0][0] = " Hund" + assert cross_lens.token_diverges(a, b, rule="half") is True + assert cross_lens.token_diverges(a, b, rule="strict") is False + assert cross_lens.token_diverges({"layers": {}}, b, rule="half") is None -def test_en_zh_aggregator_end_to_end(tmp_path): +# --------------------------------------------------------------------------- # +# EN-ZH aggregator +# --------------------------------------------------------------------------- # + + +def _en_zh_case(tmp_path): bank = { "prompts": [ { @@ -157,18 +194,20 @@ def test_en_zh_aggregator_end_to_end(tmp_path): _record("d1", "ctrl_digit", "七", {"七": 4, " seven": 4}), ] ) - (tmp_path / "en.json").write_text(json.dumps(en, ensure_ascii=False)) - (tmp_path / "zh.json").write_text(json.dumps(zh, ensure_ascii=False)) - (tmp_path / "bank.json").write_text(json.dumps(bank, ensure_ascii=False)) + return bank, en, zh + + +def test_en_zh_aggregator_end_to_end(tmp_path): + bank, en, zh = _en_zh_case(tmp_path) out = tmp_path / "summary.json" rc = agg_zh.main( [ "--lens-a", - str(tmp_path / "en.json"), + _write(tmp_path, "en.json", en), "--lens-b", - str(tmp_path / "zh.json"), + _write(tmp_path, "zh.json", zh), "--prompts", - str(tmp_path / "bank.json"), + _write(tmp_path, "bank.json", bank), "--out", str(out), "--rows-csv", @@ -183,9 +222,60 @@ def test_en_zh_aggregator_end_to_end(tmp_path): assert s["null_within_lens"]["pass"] == 0 and s["null_within_lens"]["n"] == 2 assert s["lens_only_split"]["antonym_zh"] == {"pass": 2, "n": 2} assert s["lens_only_split"]["all_cross"] == {"pass": 2, "n": 2} + # the dump format: key order of the summary and of each row, provenance last + assert list(s) == [ + "per_group", + "rows", + "headline", + "cross_form_en", + "cross_form_zh", + "null_within_lens", + "null_shuffle_cross_lens", + "lens_only_split", + "provenance", + ] + assert list(s["rows"][0]) == [ + "id", + "group", + "concept", + "cross_lens_pass", + "cross_form_en", + "cross_form_zh", + "null_hits", + "null_total", + "top1_script_en", + "top1_script_zh", + ] + assert s["provenance"]["args"]["agreement_rule"] == "half" assert (tmp_path / "rows.csv").read_text().startswith("id,group,concept,cross_lens_pass") +def test_en_zh_aggregator_strict_rule_flag(tmp_path): + bank, en, zh = _en_zh_case(tmp_path) + # a1 now agrees on exactly 2 of 4 mid-band layers: passes under half, fails under strict + _set_dom(zh["records"][0], ["21", "24"], "大", 77) + argv = [ + "--lens-a", + _write(tmp_path, "en.json", en), + "--lens-b", + _write(tmp_path, "zh.json", zh), + "--prompts", + _write(tmp_path, "bank.json", bank), + ] + agg_zh.main([*argv, "--out", str(tmp_path / "half.json")]) + agg_zh.main([*argv, "--out", str(tmp_path / "strict.json"), "--agreement-rule", "strict"]) + half = json.loads((tmp_path / "half.json").read_text()) + strict = json.loads((tmp_path / "strict.json").read_text()) + assert half["headline"]["pass"] == 1 and strict["headline"]["pass"] == 0 + assert strict["provenance"]["args"]["agreement_rule"] == "strict" + + +def test_load_bank_items_rejects_duplicate_ids(tmp_path): + p = _write(tmp_path, "dupe.json", {"prompts": [{"id": "x"}, {"id": "x"}]}) + with pytest.raises(ValueError, match="duplicate prompt ids"): + cross_lens.load_bank_items(p) + + # --------------------------------------------------------------------------- # # EN-DE aggregator and bank builder # --------------------------------------------------------------------------- # @@ -202,9 +292,101 @@ def test_lang_of_lexical_call(): assert agg_de.lang_of(" ", item) == "OTHER" -def test_norm_and_levenshtein_cognate_rule(): - assert bank_de.norm("Käse") == "kase" - assert bank_de.norm("groß") == "gross" +def test_en_de_null_population_and_divergence(tmp_path): + bank = { + "prompts": [ + { + "id": "t1", + "group": "trans_de2en", + "concept": "dog", + "form_a": " dog", + "form_b": " Hund", + "lang_a": "en", + "lang_b": "de", + "null_targets": [" Salz", " salt"], + }, + { + "id": "c1", + "group": "ctrl_digit", + "concept": "seven", + "form_a": "7", + "form_b": "7", + "lang_a": "de", + "lang_b": "en", + "null_targets": [" Schuh", " shoe"], + }, + ] + } + # the control's first null shares the dominant feature with form_a: one null hit, on a control row + a = _dump( + [ + _record("t1", "trans_de2en", " dog", {" dog": 1, " Hund": 1, " Salz": 50, " salt": 51}), + _record("c1", "ctrl_digit", "7", {"7": 4, " Schuh": 4, " shoe": 60}), + ] + ) + b = _dump( + [ + _record("t1", "trans_de2en", " Hund", {" dog": 1, " Hund": 1, " Salz": 50, " salt": 51}), + _record("c1", "ctrl_digit", "7", {"7": 4, " Schuh": 4, " shoe": 60}), + ] + ) + argv = [ + "--lens-a", + _write(tmp_path, "a.json", a), + "--lens-b", + _write(tmp_path, "b.json", b), + "--prompts", + _write(tmp_path, "bank.json", bank), + ] + agg_de.main([*argv, "--out", str(tmp_path / "all.json")]) + agg_de.main([*argv, "--out", str(tmp_path / "cross.json"), "--null-population", "cross"]) + s_all = json.loads((tmp_path / "all.json").read_text()) + s_cross = json.loads((tmp_path / "cross.json").read_text()) + # default pools every bank item (controls included); cross drops them + assert s_all["null_within_lens"] == {"pass": 1, "n": 4, "rate": 0.25} + assert s_cross["null_within_lens"] == {"pass": 0, "n": 2, "rate": 0.0} + # per-row null counts are unchanged by the population switch + assert [(r["null_hits"], r["null_total"]) for r in s_all["rows"]] == [(1, 2), (0, 2)] + assert s_all["rows"] == s_cross["rows"] + assert s_all["headline"] == {"pass": 1, "n": 1, "rate": 1.0, "ci": s_all["headline"]["ci"]} + assert s_all["divergence_rate"] == {"diverging": 1, "n": 1} + assert s_all["lens_only_split"] == { + "ctrl_digit": {"pass": 0, "n": 1}, + "trans_de2en": {"pass": 1, "n": 1}, + "all_cross": {"pass": 1, "n": 1}, + } + assert list(s_all) == [ + "per_group", + "rows", + "headline", + "cross_form_en", + "cross_form_de", + "null_within_lens", + "null_shuffle_cross_lens", + "divergence_rate", + "lens_only_split", + "provenance", + ] + assert list(s_all["rows"][0]) == [ + "id", + "group", + "concept", + "cross_lens_pass", + "cross_form_en", + "cross_form_de", + "null_hits", + "null_total", + "top1_lang_en", + "top1_lang_de", + "token_diverges", + ] + assert s_cross["provenance"]["args"]["null_population"] == "cross" + + +def test_fold_text_and_levenshtein_cognate_rule(): + assert cross_lens.fold_text("Käse") == "kase" + assert cross_lens.fold_text("groß") == "gross" + assert cross_lens.fold_text(" Straße ") == "strasse" assert bank_de.levenshtein("neu", "new") == 1 assert bank_de.levenshtein("leicht", "light") == 2 assert bank_de.levenshtein("hund", "dog") > 2 @@ -225,7 +407,10 @@ def encode(self, s, add_special_tokens=False): first = next(p for p in prompts if p["id"] == "antonym_de_01") assert first["form_a"] == "groß" and first["form_b"] == " big" # quoted frame: no leading space assert first["targets"] == [first["form_a"], first["form_b"], *first["null_targets"]] + assert first["null_targets"] == [" Salz", " salt"] # space-prefixed, language-matched, round-robin start assert all(p["lang_b"] != p["lang_a"] for p in prompts) + ctrl = next(p for p in prompts if p["group"] == "ctrl_digit") + assert ctrl["form_a"] == ctrl["form_b"] and ctrl["targets"][:2] == [ctrl["form_a"]] * 2 # --------------------------------------------------------------------------- # @@ -244,7 +429,7 @@ def test_three_lens_tables(tmp_path): dom = {(r["lens"], r["token"]): r["dominant_feature"] for r in feats if r["layer"] == "24"} assert dom[("EN", "large")] == 12474 and dom[("DE", " groß")] == 6764 for label, rec in recs.items(): - (tmp_path / f"{label}.json").write_text(json.dumps(_dump([rec]), ensure_ascii=False)) + _write(tmp_path, f"{label}.json", _dump([rec])) rc = three.main( [ "--dump", @@ -258,9 +443,11 @@ def test_three_lens_tables(tmp_path): ] ) assert rc == 0 and (tmp_path / "t.csv").exists() + manifest = json.loads((tmp_path / "t.manifest.json").read_text()) + assert manifest["lenses"] == ["EN", "DE"] and "provenance" in manifest -def test_shared_feature_metrics(): +def test_shared_feature_metrics(tmp_path): recs = { "EN": _record("f", "fact_zh", " London", {"伦敦": 1}, top1_feats=[(23180, 30.0), (26030, 1.7), (5, 0.3)]), "ZH": _record("f", "fact_zh", "伦敦", {"伦敦": 1}, top1_feats=[(23180, 24.0), (5150, 1.5)]), @@ -271,11 +458,28 @@ def test_shared_feature_metrics(): assert m["per_lens"]["EN"]["largest_other_contribution"] == 1.7 assert m["shared_feature_share_of_feature_sum"]["ZH"] == 24.0 / 25.5 assert m["shared_feature_top_tokens"] == [" London", "伦敦"] + for label, rec in recs.items(): + _write(tmp_path, f"{label}.json", _dump([rec])) + rc = shared.main( + [ + "--dump", + f"EN={tmp_path / 'EN.json'}", + "--dump", + f"ZH={tmp_path / 'ZH.json'}", + "--prompt-id", + "f", + "--out-json", + str(tmp_path / "m.json"), + ] + ) + assert rc == 0 + written = json.loads((tmp_path / "m.json").read_text()) + assert written["shared_dominant_feature"] == 23180 and written["provenance"]["args"]["layer"] == "26" recs["ZH"]["layers"]["26"]["-1"]["top1_decomp"]["top_features"][0]["id"] = 7 assert shared.shared_feature_metrics(recs, "26", {})["shared_dominant_feature"] is None -def test_antonym_layer_table(): +def test_antonym_layer_table(tmp_path): def payload(tok, pct, fid): rows = { L: { @@ -287,17 +491,338 @@ def payload(tok, pct, fid): } return {"prompt": "p", "rows": rows} - table, feats = antonym.antonym_layer_table( - {"EN": payload("large", 39.7, 112), "ZH": payload("大的", 36.4, 112)}, ["24", "final"] - ) + dumps = {"EN": payload("large", 39.7, 112), "ZH": payload("大的", 36.4, 112)} + table, feats = antonym.antonym_layer_table(dumps, ["24", "final"]) assert [(r["layer"], r["lens"], r["top1_token"], r["top1_softmax_pct"]) for r in table][:2] == [ ("24", "EN", "large", 39.7), ("24", "ZH", "大的", 36.4), ] assert all(r["dominant_feature"] == 112 for r in feats) - assert antonym.parse_dump_args(["EN=a.json", "ZH=b.json"]) == [("EN", Path("a.json")), ("ZH", Path("b.json"))] + assert cross_lens.parse_dump_args(["EN=a.json", "ZH=b.json"]) == [("EN", Path("a.json")), ("ZH", Path("b.json"))] + with pytest.raises(ValueError, match="LABEL=path"): + cross_lens.parse_dump_args(["a.json"]) + # the table subcommand end to end: CSVs plus a sibling manifest (args carry the subcommand callable) + for label, d in dumps.items(): + _write(tmp_path, f"{label}.json", d) + rc = antonym.main( + [ + "table", + "--dump", + f"EN={tmp_path / 'EN.json'}", + "--dump", + f"ZH={tmp_path / 'ZH.json'}", + "--layers", + "24,final", + "--out-csv", + str(tmp_path / "table.csv"), + "--out-features-csv", + str(tmp_path / "features.csv"), + ] + ) + assert rc == 0 and (tmp_path / "features.csv").exists() + manifest = json.loads((tmp_path / "table.manifest.json").read_text()) + assert manifest["lenses"] == ["EN", "ZH"] and manifest["provenance"]["args"]["layers"] == "24,final" + assert "fn" not in manifest["provenance"]["args"] + + +# --------------------------------------------------------------------------- # +# decomposition, labels, k and centering resolution, final norm +# --------------------------------------------------------------------------- # + + +def _dictionary(seed=0, V=12, d=6, D=9): + g = torch.Generator().manual_seed(seed) + W = torch.randn(V, d, generator=g) + decoder = torch.randn(D, d, generator=g) + decoder = decoder / decoder.norm(dim=1, keepdim=True) + encoder_w = torch.randn(D, d, generator=g) + encoder_b = 0.1 * torch.randn(D, generator=g) + h = torch.randn(d, generator=g) + return W, decoder, encoder_w, encoder_b, h + + +def test_decompose_token_pins_pre_0_2_1_inline_formula(): + W, decoder, encoder_w, encoder_b, h = _dictionary() + row_mean = W.mean(dim=0) + k, top_feats, token_id = 3, 2, 5 + + # the block both GPU scripts inlined before 0.2.1 + W_row = W[token_id] + centered = W_row - row_mean + norm = centered.norm().clamp_min(1e-8) + acts = F.relu((centered / norm)[None, :] @ encoder_w.T + encoder_b) + values, indices = torch.topk(acts, k=k, dim=-1) + code = torch.zeros_like(acts).scatter_(-1, indices, values)[0] + contributions = norm * code * (h @ decoder.T) + active = torch.nonzero(contributions != 0).flatten() + top = active[contributions[active].abs().argsort(descending=True)[:top_feats]] + expected = { + "original_logit": float(h @ W_row), + "base": float(h @ row_mean), + "feature_sum": float(contributions.sum()), + "top_features": [{"id": int(f), "contribution": float(contributions[f])} for f in top], + } + expected["residual"] = expected["original_logit"] - expected["base"] - expected["feature_sum"] + + got = cross_lens.decompose_token( + h, + token_id, + W=W, + row_mean=row_mean, + decoder=decoder, + encoder_w=encoder_w, + encoder_b=encoder_b, + k=k, + top_feats=top_feats, + ) + assert list(got) == ["original_logit", "base", "feature_sum", "residual", "top_features"] + for key in ("original_logit", "base", "feature_sum", "residual"): + assert got[key] == pytest.approx(expected[key], abs=1e-6) + assert [f["id"] for f in got["top_features"]] == [f["id"] for f in expected["top_features"]] + assert [f["contribution"] for f in got["top_features"]] == pytest.approx( + [f["contribution"] for f in expected["top_features"]], abs=1e-6 + ) + assert len(got["top_features"]) == top_feats + assert got["original_logit"] == pytest.approx(got["base"] + got["feature_sum"] + got["residual"]) + + +def test_feature_top_tokens_keeps_raw_decode_labels(): + W, _decoder, encoder_w, encoder_b, _h = _dictionary(seed=1) + row_mean = W.mean(dim=0) + + class Tok: + def decode(self, ids): + return f"<{ids[0]}>" + + labels = cross_lens.feature_top_tokens(W, row_mean, [4, 1, 4], encoder_w, encoder_b, Tok(), top_tokens=3, chunk=5) + assert list(labels) == [1, 4] # sorted, de-duplicated feature ids + _norms, x = cross_lens.center_normalize_rows(W, row_mean) + for fid, got in labels.items(): + scores = F.relu(x @ encoder_w[fid] + encoder_b[fid]) + order = scores.argsort(descending=True)[:3] + assert got == [f"<{int(i)}>" for i in order if scores[i] > 0] + assert cross_lens.feature_top_tokens(W, row_mean, [], encoder_w, encoder_b, Tok()) == {} + + +def test_resolve_k_defaults_to_checkpoint_and_warns_on_mismatch(): + config = {"factorizer": {"k": 128}} + assert cross_lens.resolve_k(None, config) == 128 + assert cross_lens.resolve_k(128, config) == 128 + with pytest.warns(UserWarning, match="differs from the checkpoint"): + assert cross_lens.resolve_k(64, config) == 64 + with pytest.raises(ValueError, match="factorizer.k"): + cross_lens.resolve_k(None, {}) + assert cross_lens.resolve_k(32, {}) == 32 + + +class _Tok: + def __init__(self, n, special=()): + self._n, self.all_special_ids = n, list(special) + + def __len__(self): + return self._n + + +def test_centering_trained_equals_live_when_mask_keeps_every_row(): + W, *_ = _dictionary(seed=2) + live = cross_lens.centering_row_mean(W, "live", tok=_Tok(W.shape[0]), ckpt_row_mean=None) + trained = cross_lens.centering_row_mean(W, "trained", tok=_Tok(W.shape[0]), ckpt_row_mean=None) + assert torch.equal(live, W.mean(dim=0)) and torch.equal(trained, live) + masked = cross_lens.centering_row_mean(W, "trained", tok=_Tok(W.shape[0], special=[3]), ckpt_row_mean=None) + assert torch.equal(masked, W[[i for i in range(W.shape[0]) if i != 3]].mean(dim=0)) + stored = torch.full((W.shape[1],), 0.25) + assert torch.equal(cross_lens.centering_row_mean(W, "trained", tok=_Tok(1), ckpt_row_mean=stored), stored) + assert torch.equal(cross_lens.centering_row_mean(W, "live", tok=_Tok(1), ckpt_row_mean=stored), live) + with pytest.raises(ValueError, match="centering mode"): + cross_lens.centering_row_mean(W, "mean", tok=_Tok(1), ckpt_row_mean=None) + + +def test_find_final_norm_checks_known_paths(): + def wrap(path): + root = nn.Module() + node = root + parts = path.split(".") + for part in parts[:-1]: + child = nn.Module() + setattr(node, part, child) + node = child + setattr(node, parts[-1], nn.LayerNorm(4)) + return root + + for path in ("model.norm", "model.language_model.norm", "language_model.norm", "norm"): + found = cross_lens.find_final_norm(wrap(path)) + assert isinstance(found, nn.LayerNorm), path + with pytest.raises(RuntimeError, match="final norm"): + cross_lens.find_final_norm(nn.Linear(2, 2)) + + +# --------------------------------------------------------------------------- # +# lens fitters: layer picker, OOM fallback, resume sidecars, holdout guard +# --------------------------------------------------------------------------- # def test_pick_source_layers_matches_paper_layer_set(): assert fit_jlens.pick_source_layers(32, 12) == [2, 5, 7, 10, 12, 14, 17, 19, 21, 24, 26, 29] assert fit_jlens.pick_source_layers(12, 4) == [2, 4, 7, 9] + + +def test_dim_batch_fallback_halves_down_to_one(tmp_path): + assert fit_jlens.dim_batch_schedule(8) == [8, 4, 2, 1] + assert fit_jlens.dim_batch_schedule(3) == [3, 1] + assert fit_jlens.dim_batch_schedule(1) == [1] + with pytest.raises(ValueError): + fit_jlens.dim_batch_schedule(0) + + attempts = [] + + def fake_fit(model, prompts, *, source_layers, dim_batch, checkpoint_path, resume): + attempts.append(dim_batch) + if dim_batch > 2: + raise torch.cuda.OutOfMemoryError("fake OOM") + return SimpleNamespace(n_prompts=len(prompts), dim_batch=dim_batch, ckpt=checkpoint_path, resume=resume) + + lens = fit_jlens.fit_with_fallback(fake_fit, "model", ["p1", "p2"], [2, 5], 8, tmp_path / "shard0.ckpt.pt") + assert attempts == [8, 4, 2] and lens.dim_batch == 2 and lens.resume is True + + attempts.clear() + + def always_oom(*a, **k): + attempts.append(k["dim_batch"]) + raise torch.cuda.OutOfMemoryError("fake OOM") + + with pytest.raises(RuntimeError, match="all dim_batch fallbacks OOMed"): + fit_jlens.fit_with_fallback(always_oom, "model", ["p1"], [2], 3, tmp_path / "shard1.ckpt.pt") + assert attempts == [3, 1] + + +def test_check_resume_meta_refuses_a_different_fit(tmp_path): + meta_path = tmp_path / "shard0.meta.json" + lens = tmp_path / "shard0.lens.pt" + expected = {"prompts_sha1": "abc", "n_prompts": 100, "start": 0} + cross_lens.check_resume_meta(meta_path, expected, [lens]) # nothing on disk: fresh run + lens.write_bytes(b"x") + with pytest.raises(RuntimeError, match="sidecar"): + cross_lens.check_resume_meta(meta_path, expected, [lens]) # artifact without sidecar + cross_lens.write_resume_meta(meta_path, expected) + cross_lens.check_resume_meta(meta_path, expected, [lens]) # same fit: fine + with pytest.raises(RuntimeError) as err: + cross_lens.check_resume_meta(meta_path, {**expected, "n_prompts": 300, "start": 0}, [lens]) + assert "n_prompts: on disk 100, requested 300" in str(err.value) and "start" not in str(err.value).split("\n", 1)[1] + + +def test_cmd_fit_refuses_stale_shard_before_loading_model(tmp_path, monkeypatch): + prompts_json = _write(tmp_path, "prompts.json", {"prompts": [f"prompt {i}" for i in range(8)]}) + out_dir = tmp_path / "ckpt" + out_dir.mkdir() + (out_dir / "shard0.lens.pt").write_bytes(b"stale") + args = SimpleNamespace( + model_id="m", + prompts_json=prompts_json, + n_prompts=4, + start=0, + shard=0, + num_shards=2, + n_layers=12, + dim_batch=8, + ckpt_dir=str(out_dir), + out_dir=str(out_dir), + device="cpu", + ) + monkeypatch.setattr(fit_jlens, "import_jlens", lambda: SimpleNamespace(fit=None)) + + def no_model(*a, **k): + raise AssertionError("the model must not be loaded when the sidecar check fails or the shard is done") + + monkeypatch.setattr(fit_jlens, "load_lens_model", no_model) + # sidecar from a 4-prompt fit; a re-run asking for 8 prompts is refused, naming the key + cross_lens.write_resume_meta( + out_dir / "shard0.meta.json", fit_jlens.shard_meta(args, fit_jlens.read_prompts(prompts_json, 4)) + ) + args.n_prompts = 8 + with pytest.raises(RuntimeError, match="n_prompts: on disk 4, requested 8"): + fit_jlens.cmd_fit(args) + # the matching request resumes (returns) without touching the model + args.n_prompts = 4 + assert fit_jlens.cmd_fit(args) is None + + +def test_check_shard_metas_agree_names_the_key(): + base = {k: 1 for k in fit_jlens.SHARED_META_KEYS} + base["num_shards"] = 2 + metas = [{**base, "shard": 0}, {**base, "shard": 1}] + assert fit_jlens.check_shard_metas_agree(metas, num_shards=2) == {k: base[k] for k in fit_jlens.SHARED_META_KEYS} + with pytest.raises(RuntimeError, match="prompts_sha1"): + fit_jlens.check_shard_metas_agree([metas[0], {**metas[1], "prompts_sha1": "other"}], num_shards=2) + with pytest.raises(RuntimeError, match="num_shards"): + fit_jlens.check_shard_metas_agree(metas, num_shards=4) + with pytest.raises(RuntimeError, match="shard:"): + fit_jlens.check_shard_metas_agree([metas[0], metas[0]], num_shards=2) + + +class _FakeLens: + def __init__(self, jacobians, n_prompts): + self.jacobians, self.n_prompts = jacobians, n_prompts + self.source_layers = sorted(jacobians) + + @classmethod + def load(cls, path): + return cls({2: torch.ones(2, 2)}, n_prompts=3) + + @classmethod + def merge(cls, lenses): + n = sum(lens.n_prompts for lens in lenses) + return cls({2: sum(lens.jacobians[2] * lens.n_prompts for lens in lenses) / n}, n_prompts=n) + + def save(self, path): + Path(path).write_bytes(b"lens") + + +def test_cmd_merge_requires_agreeing_sidecars_and_records_them(tmp_path, monkeypatch): + monkeypatch.setattr(fit_jlens, "import_jlens", lambda: SimpleNamespace(JacobianLens=_FakeLens)) + shard_dir = tmp_path / "shards" + shard_dir.mkdir() + base = { + "model_id": "m", + "prompts_json": "p.json", + "prompts_sha1": "abc", + "n_prompts": 6, + "start": 0, + "num_shards": 2, + "n_layers": 12, + } + for i in range(2): + (shard_dir / f"shard{i}.lens.pt").write_bytes(b"x") + args = SimpleNamespace( + shard_dir=str(shard_dir), num_shards=2, out=str(tmp_path / "merged.pt"), fn=fit_jlens.cmd_merge + ) + with pytest.raises(FileNotFoundError, match="sidecar"): + fit_jlens.cmd_merge(args) + cross_lens.write_resume_meta(shard_dir / "shard0.meta.json", {**base, "shard": 0}) + cross_lens.write_resume_meta(shard_dir / "shard1.meta.json", {**base, "shard": 1, "n_prompts": 9}) + with pytest.raises(RuntimeError, match="n_prompts"): + fit_jlens.cmd_merge(args) + cross_lens.write_resume_meta(shard_dir / "shard1.meta.json", {**base, "shard": 1}) + fit_jlens.cmd_merge(args) + assert (tmp_path / "merged.pt").read_bytes() == b"lens" + meta = json.loads((tmp_path / "merged.meta.json").read_text()) + assert meta["prompts_sha1"] == "abc" and meta["n_prompts"] == 6 and meta["shard_n_prompts"] == [3, 3] + assert meta["merged_n_prompts"] == 6 and meta["layers"] == [2] + assert meta["provenance"]["args"] == { + "shard_dir": str(shard_dir), + "num_shards": 2, + "out": str(tmp_path / "merged.pt"), + } + + +def test_prompt_list_sha1_is_content_addressed(): + assert cross_lens.prompt_list_sha1(["a", "b"]) == cross_lens.prompt_list_sha1(["a", "b"]) + assert cross_lens.prompt_list_sha1(["a", "b"]) != cross_lens.prompt_list_sha1(["a\nb"]) + + +def test_ridge_holdout_guard_and_diagnostic_slice(tmp_path): + argv = ["--prompts-json", "missing.json", "--layers-from", "missing.pt", "--ckpt-dir", str(tmp_path), "--out"] + with pytest.raises(SystemExit, match="--holdout must be >= 1"): + fit_ridge.main([*argv, str(tmp_path / "o.pt"), "--holdout", "0"]) + prompts = [f"p{i}" for i in range(100)] + assert fit_ridge.diagnostic_prompts(prompts, 90) == prompts[95:100] # the paper's slice, holdout 10 + assert fit_ridge.diagnostic_prompts(prompts, 98) == prompts[98:100] # never reaches into the fit prompts diff --git a/tests/test_error_tails_and_nearest_rows.py b/tests/test_error_tails_and_nearest_rows.py index a2f562e..233f8cb 100644 --- a/tests/test_error_tails_and_nearest_rows.py +++ b/tests/test_error_tails_and_nearest_rows.py @@ -4,12 +4,16 @@ * scripts/analyze/nearest_rows_baseline.py (tab:app-nearest-rows) Synthetic inputs only: hand-built baseline_query_rows.csv files -and a random W_U with a fake tokenizer; no model, no network. +and a random W_U with a fake tokenizer; no model, no network. The bootstrap +tests pin the vectorised gathers to the 0.2.0 per-cluster ``np.concatenate`` +loops, element for element. """ from __future__ import annotations -import importlib.util +import csv +import hashlib +import io import json import sys from pathlib import Path @@ -17,21 +21,10 @@ import numpy as np import pandas as pd import torch +from conftest import load_script -ROOT = Path(__file__).resolve().parents[1] -sys.path.insert(0, str(ROOT / "src")) - - -def _load(mod_name: str, rel_path: str): - spec = importlib.util.spec_from_file_location(mod_name, ROOT / rel_path) - mod = importlib.util.module_from_spec(spec) - sys.modules[mod_name] = mod - spec.loader.exec_module(mod) - return mod - - -tails = _load("analyze_error_tails", "scripts/eval/analyze_error_tails.py") -nrb = _load("nearest_rows_baseline", "scripts/analyze/nearest_rows_baseline.py") +tails = load_script("scripts/eval/analyze_error_tails.py") +nrb = load_script("scripts/analyze/nearest_rows_baseline.py") # --------------------------------------------------------------------------- # @@ -126,6 +119,100 @@ def test_error_tails_row_metrics_and_pooled_tables(tmp_path: Path, monkeypatch) assert manifest["rows_total"] == 240 and "command" in manifest +def _bootstrap_ci_v020(df: pd.DataFrame, n_boot: int, seed: int) -> dict[str, float]: + """The 0.2.0 loop, kept verbatim as the reference for the vectorised gather.""" + clusters = {str(k): idx.to_numpy() for k, idx in df.groupby("base_case_id", sort=True).groups.items()} + keys = sorted(clusters) + if len(keys) < 2 or n_boot <= 0: + return {} + rng = np.random.default_rng(seed) + rho = df["rho"].to_numpy(float) + sign = df["sign_match_strict"].to_numpy(float) + accepted = df["accepted"].to_numpy(float) + stats = np.empty((n_boot, 5), dtype=float) + for b in range(n_boot): + sampled = rng.integers(0, len(keys), size=len(keys)) + idx = np.concatenate([clusters[keys[i]] for i in sampled]) + rb = rho[idx] + stats[b] = (rb.mean(), np.median(rb), np.percentile(rb, 95), sign[idx].mean(), accepted[idx].mean()) + names = ("rho_mean", "rho_median", "rho_p95", "sign_agreement", "accepted_rate") + out: dict[str, float] = {} + for j, name in enumerate(names): + out[f"{name}_ci_lo"] = float(np.percentile(stats[:, j], 2.5)) + out[f"{name}_ci_hi"] = float(np.percentile(stats[:, j], 97.5)) + return out + + +def _paired_differences_v020(qdf: pd.DataFrame, n_boot: int, seed: int) -> pd.DataFrame: + """The 0.2.0 pandas ``iloc`` loop, kept verbatim as the reference for the vectorised gather.""" + keys = ["bank", "base_case_id", "case_id", "query"] + cols = keys + ["rho", "sign_match_strict", "accepted"] + a = qdf[qdf.method == tails.ANCHOR][cols].copy() + results: list[dict] = [] + for method in sorted(set(qdf.method) - {tails.ANCHOR}): + b = qdf[qdf.method == method][cols].copy() + m = a.merge(b, on=keys, suffixes=("_srp", "_base"), validate="one_to_one") + if m.empty: + continue + clusters = {str(k): idx.to_numpy() for k, idx in m.groupby("base_case_id", sort=True).groups.items()} + ckeys = sorted(clusters) + rng = np.random.default_rng(seed + int(hashlib.sha1(method.encode()).hexdigest()[:7], 16)) + + def diffs(frame: pd.DataFrame) -> np.ndarray: + return np.array( + [ + frame.rho_srp.mean() - frame.rho_base.mean(), + np.percentile(frame.rho_srp, 95) - np.percentile(frame.rho_base, 95), + frame.sign_match_strict_srp.mean() - frame.sign_match_strict_base.mean(), + frame.accepted_srp.mean() - frame.accepted_base.mean(), + ] + ) + + point = diffs(m) + boots = np.empty((n_boot, 4), float) + for i in range(n_boot): + sample = rng.integers(0, len(ckeys), size=len(ckeys)) + idx = np.concatenate([clusters[ckeys[j]] for j in sample]) + boots[i] = diffs(m.iloc[idx]) + rec: dict[str, float | str | int] = {"baseline": method, "n_pairs": len(m)} + for j, name in enumerate(("mean_rho", "p95_rho", "sign_agreement", "accepted_rate")): + rec[f"delta_srp_minus_baseline_{name}"] = float(point[j]) + rec[f"{name}_ci_lo"] = float(np.percentile(boots[:, j], 2.5)) + rec[f"{name}_ci_hi"] = float(np.percentile(boots[:, j], 97.5)) + results.append(rec) + return pd.DataFrame(results) + + +def test_resample_index_matches_concatenate() -> None: + rng = np.random.default_rng(7) + # Ragged clusters in scrambled row positions, keys whose sorted order differs from insertion order. + clusters = {"k3": np.array([5, 1, 9]), "k1": np.array([0, 7]), "k2": np.array([2, 3, 4, 8, 6])} + order, starts, lengths = tails._flat_clusters(clusters) + keys = sorted(clusters) + assert list(lengths) == [2, 5, 3] and list(starts) == [0, 2, 7] + for _ in range(50): + sampled = rng.integers(0, len(keys), size=len(keys)) + expected = np.concatenate([clusters[keys[i]] for i in sampled]) + assert np.array_equal(order[tails._resample_index(sampled, starts, lengths)], expected) + assert tails._resample_index(np.array([], dtype=int), starts, lengths).size == 0 + + +def test_bootstraps_are_bit_identical_to_the_v020_loops(tmp_path: Path) -> None: + for tag, seed in (("m1", 11), ("m2", 12)): + _fake_run_dir(tmp_path / "in", tag, seed) + rows, _ = tails.load_rows(tmp_path / "in", "run_", ["m1", "m2"]) + qdf = rows[rows.model_tag == "m1"].copy() + part = qdf[qdf.method == "sparse_rp"].reset_index(drop=True) + assert tails.bootstrap_ci(part, 200, 5) == _bootstrap_ci_v020(part, 200, 5) + # A second group with a different seed, as summarize() draws them. + part2 = rows[rows.method == "sparse_rp"].reset_index(drop=True) + assert tails.bootstrap_ci(part2, 100, 20260711) == _bootstrap_ci_v020(part2, 100, 20260711) + new = tails.paired_differences(qdf, 200, 5) + old = _paired_differences_v020(qdf, 200, 5) + assert list(new.columns) == list(old.columns) + assert new.equals(old), (new, old) + + # --------------------------------------------------------------------------- # # nearest rows # --------------------------------------------------------------------------- # @@ -145,7 +232,7 @@ def decode(self, ids: list[int]) -> str: return "".join(self.vocab[i] for i in ids) -def test_nearest_rows_excludes_contrast_tokens_and_ranks_by_abs_cosine() -> None: +def _planted_case(): vocab = [" bug", " insect", " error", "bug", "Bug", " bugs", "昆虫", "z1", "z2", "z3", "z4", "z5"] tok = _FakeTokenizer(vocab) gen = torch.Generator().manual_seed(0) @@ -155,6 +242,11 @@ def test_nearest_rows_excludes_contrast_tokens_and_ranks_by_abs_cosine() -> None W[3] = W[0] + 0.1 * torch.randn(d, generator=gen) W[4] = W[0] + 0.1 * torch.randn(d, generator=gen) W[6] = W[1] + 0.1 * torch.randn(d, generator=gen) + return W, tok + + +def test_nearest_rows_excludes_contrast_tokens_and_ranks_by_abs_cosine() -> None: + W, tok = _planted_case() row_mean = W.mean(dim=0) res, report, ids = nrb.nearest_rows(W, row_mean, tok, [("bug", "insect")], top_n=4) assert ids == {"bug": 0, "insect": 1} @@ -173,3 +265,38 @@ def test_nearest_rows_excludes_contrast_tokens_and_ranks_by_abs_cosine() -> None assert len(table) == 4 and table[0]["contrast"] == "bug_minus_insect" assert all(t["centered_cosine"].startswith(("+", "-")) for t in table) assert any(t["centered_gloss"] == "insect" for t in table) == any(t["centered_token"] == "昆虫" for t in table) + + +def _v020_csv_bytes(contrast_results: dict) -> tuple[str, str]: + """The 0.2.0 writers (csv.writer / csv.DictWriter over table[0].keys()), kept as the byte reference.""" + long = io.StringIO(newline="") + w = csv.writer(long) + w.writerow(["contrast", "variant", "rank", "token_id", "token", "cosine"]) + for name, c in contrast_results.items(): + for variant in ("top_centered", "top_raw"): + for row in c[variant]: + w.writerow([name, variant, row["rank"], row["token_id"], repr(row["token"]), row["cosine"]]) + table = nrb.table_rows(contrast_results) + tab = io.StringIO(newline="") + dw = csv.DictWriter(tab, fieldnames=list(table[0].keys())) + dw.writeheader() + dw.writerows(table) + return long.getvalue(), tab.getvalue() + + +def test_write_tables_matches_v020_bytes_and_survives_empty(tmp_path: Path) -> None: + W, tok = _planted_case() + res, _report, _ids = nrb.nearest_rows(W, W.mean(dim=0), tok, [("bug", "insect"), ("bug", "error")], top_n=3) + nrb.write_tables(tmp_path, res) + long_ref, table_ref = _v020_csv_bytes(res) + assert (tmp_path / "nearest_rows.csv").read_bytes().decode("utf-8") == long_ref + assert (tmp_path / "nearest_rows_table.csv").read_bytes().decode("utf-8") == table_ref + assert long_ref.splitlines()[0] == ",".join(nrb.LONG_FIELDS) + assert table_ref.splitlines()[0] == ",".join(nrb.TABLE_FIELDS) + assert len(long_ref.splitlines()) == 1 + 2 * 2 * 3 + # --top-n 0 style empty result: the 0.2.0 table writer crashed on table[0]; now both files are 0 bytes. + empty = tmp_path / "empty" + empty.mkdir() + nrb.write_tables(empty, {}) + assert (empty / "nearest_rows.csv").stat().st_size == 0 + assert (empty / "nearest_rows_table.csv").stat().st_size == 0 diff --git a/tests/test_lexical_control_directions.py b/tests/test_lexical_control_directions.py index f42426c..033e341 100644 --- a/tests/test_lexical_control_directions.py +++ b/tests/test_lexical_control_directions.py @@ -12,26 +12,11 @@ from __future__ import annotations -import importlib.util -import sys -from pathlib import Path - import pytest import torch +from conftest import load_script -ROOT = Path(__file__).resolve().parents[1] -sys.path.insert(0, str(ROOT / "src")) - - -def _load(relpath: str, name: str): - spec = importlib.util.spec_from_file_location(name, ROOT / relpath) - mod = importlib.util.module_from_spec(spec) - sys.modules[name] = mod # dataclass resolution needs the module registered - spec.loader.exec_module(mod) - return mod - - -prof = _load("scripts/run/run_qwen_profanity_suppression_eval.py", "qwen_profanity_suppression_eval") +prof = load_script("scripts/run/run_qwen_profanity_suppression_eval.py") NEW_METHODS = ("mean_row_direction", "pca_group_direction", "pca_group_rank4") diff --git a/tests/test_paired_matched_kl_bootstrap.py b/tests/test_paired_matched_kl_bootstrap.py index 302073a..51ec47d 100644 --- a/tests/test_paired_matched_kl_bootstrap.py +++ b/tests/test_paired_matched_kl_bootstrap.py @@ -1,29 +1,19 @@ """Smoke test for scripts/eval/paired_matched_kl_bootstrap.py on a synthetic candidate_constrained_rows.csv: the matched scales are the log-nearest median KLs, the paired mean difference equals the direct per-candidate mean, both -clusterings (held-out term and prompt) are reported, and the run is seeded. +clusterings (held-out term and prompt) are reported, the run is seeded, and the +output carries a ``provenance`` block next to the ``comparisons`` records. """ from __future__ import annotations import csv -import importlib.util import json -import sys from pathlib import Path -ROOT = Path(__file__).resolve().parents[1] +from conftest import load_script - -def _load(relpath: str, name: str): - spec = importlib.util.spec_from_file_location(name, ROOT / relpath) - mod = importlib.util.module_from_spec(spec) - sys.modules[name] = mod - spec.loader.exec_module(mod) - return mod - - -boot = _load("scripts/eval/paired_matched_kl_bootstrap.py", "paired_matched_kl_bootstrap") +boot = load_script("scripts/eval/paired_matched_kl_bootstrap.py") TERMS = ["damn", "crap", "bullshit"] PROMPTS = ["Oh", "Holy", "This is", "You are such a", "That really", "I am so", "You really"] @@ -86,10 +76,29 @@ def test_paired_bootstrap_on_synthetic_rows(tmp_path: Path) -> None: ["--input", f"toy={csv_path}", "--out", str(out), "--n-boot", "50", "--seed", "0", "--target-kl", "0.05", "0.2"] ) assert rc == 0 - results = json.loads(out.read_text()) + payload = json.loads(out.read_text()) + assert set(payload) == {"comparisons", "provenance"} + prov = payload["provenance"] + assert {"command", "args", "git_commit", "package_version", "torch_version", "timestamp_unix"} <= set(prov) + assert prov["args"]["n_boot"] == 50 and prov["args"]["seed"] == 0 + results = payload["comparisons"] # Only mean_row_direction is present -> 2 targets x 1 baseline x 2 outcomes. assert len(results) == 4 assert {r["model"] for r in results} == {"toy"} + assert list(results[0]) == [ + "model", + "target_kl", + "baseline", + "outcome", + "n", + "srp_scale", + "srp_kl", + "base_scale", + "base_kl", + "mean_diff", + "ci_term", + "ci_prompt", + ] by_key = {(r["target_kl"], r["outcome"]): r for r in results} r = by_key[(0.05, "dp")] @@ -104,10 +113,10 @@ def test_paired_bootstrap_on_synthetic_rows(tmp_path: Path) -> None: assert by_key[(0.05, "flip")]["mean_diff"] == 1.0 assert by_key[(0.2, "dp")]["srp_scale"] == 4.0 and by_key[(0.2, "dp")]["base_scale"] == 8.0 - # Seeded: a second run reproduces the file byte for byte. + # Seeded: a second run reproduces the comparison records exactly (provenance carries a timestamp). out2 = tmp_path / "paired2.json" boot.main(["--input", f"toy={csv_path}", "--out", str(out2), "--n-boot", "50", "--target-kl", "0.05", "0.2"]) - assert out2.read_text() == out.read_text() + assert json.loads(out2.read_text())["comparisons"] == results def test_scale_for_kl_respects_two_x_window() -> None: diff --git a/tests/test_seed_stability_scripts.py b/tests/test_seed_stability_scripts.py index 5bd2dc9..afa0805 100644 --- a/tests/test_seed_stability_scripts.py +++ b/tests/test_seed_stability_scripts.py @@ -1,37 +1,58 @@ -"""CPU smoke tests for the cross-seed stability scripts on synthetic inputs. +"""CPU tests for the cross-seed stability scripts on synthetic inputs. Three tiny dictionaries, a 40-row readout and a stub tokenizer; no network, no -GPU. Identical dictionaries must reproduce each other exactly (same-side -Jaccard 1, best-single Jaccard 1, held-out core recall 1), which pins the -shared pipeline in ``sparse_readout_prism.research.seed_stability``. +GPU. Two layers of protection for the shared pipeline in +``sparse_readout_prism.research.seed_stability``: + +* invariants -- identical dictionaries reproduce each other exactly (same-side + Jaccard 1, best-single Jaccard 1, held-out core recall 1); +* pins -- the exact numbers the scripts produced on these fixtures before the + 0.2.1 consolidation (recorded from the pre-refactor code), so a refactor of + the shared helpers cannot move an output silently. The seed-variation + dictionaries of the paper are not distributed, so these pins are the + behaviour-preservation proof for the three scripts. """ from __future__ import annotations -import importlib.util import json +import os +import subprocess import sys +import textwrap from pathlib import Path import numpy as np +import pytest import torch +from conftest import REPO_ROOT, load_script +from sparse_readout_prism.research.qwen_readout import encode_topk +from sparse_readout_prism.research.row_geometry import resolve_single_token_bare_first from sparse_readout_prism.research.seed_stability import ( + Dictionary, + center_rows, load_contrast_pairs, load_dictionary, - single_token_id, - topk_codes, + load_readout, + resolve_centering, + summary_stats, + unit_rows, ) -ROOT = Path(__file__).resolve().parents[1] V, D_MODEL, D_FEAT, K = 40, 8, 64, 4 class _StubTokenizer: + all_special_ids = [0] # row 0 is a "special" token for token_mask_from_tokenizer + def __init__(self, vocab: list[str]) -> None: self.vocab = vocab self.index = {w: i for i, w in enumerate(vocab)} + def __len__(self) -> int: + return len(self.vocab) + def encode(self, text: str, add_special_tokens: bool = False) -> list[int]: w = text.strip() return [self.index[w]] if w in self.index else [0, 1] @@ -40,15 +61,6 @@ def decode(self, ids: list[int]) -> str: return " " + self.vocab[int(ids[0])] -def _load_script(name: str): - path = ROOT / "scripts" / "eval" / f"{name}.py" - spec = importlib.util.spec_from_file_location(f"_seed_stability_{name}", path) - module = importlib.util.module_from_spec(spec) - sys.modules[spec.name] = module - spec.loader.exec_module(module) - return module - - def _inputs(tmp_path: Path, identical: bool): g = torch.Generator().manual_seed(0) vocab = [f"w{i:02d}" for i in range(V)] @@ -71,55 +83,279 @@ def _inputs(tmp_path: Path, identical: bool): recs.append({"target_a": vocab[0], "target_b": vocab[1]}) # duplicate pair: skipped bank = tmp_path / "bank.jsonl" bank.write_text("\n".join(json.dumps(r) for r in recs) + "\n") + token_mask = torch.ones(V, dtype=torch.bool) + token_mask[0] = False + torch.save({"W_U_orig": W, "h_LN": h_LN, "token_mask": token_mask}, tmp_path / "w_u.pt") return _StubTokenizer(vocab), W, h_LN, dicts, bank +def _live_mean(W: torch.Tensor, dicts: list) -> torch.Tensor: + return resolve_centering(W, dicts, "live", None) + + +# --------------------------------------------------------------------------- # +# Pins recorded from the pre-0.2.1 scripts on the fixtures above (rng seed 0, +# top_m=3, top_r=4; cross-seed n_sample=16 n_hidden=8; loo side_n=6 +# n_clusters=8 kmeans_iters=2 kmeans_seed=0 n_boot=50 bootstrap_seed=1). +# Set-derived statistics are rationals and are compared to 1e-12; float32 +# cosine / correlation statistics to 1e-6. +# --------------------------------------------------------------------------- # + +SET_TOL, F32_TOL = 1e-12, 1e-6 + +CROSS_SEED_PINS = { + True: dict( # identical dictionaries + same_side={"mean": 1.0, "median": 1.0, "p10": 1.0, "n": 72}, + cross_side={"mean": 0.17782421647553226, "median": 0.17207792207792205, "p10": 0.0, "n": 36}, + cross_contrast_null={"mean": 0.1740578507606371, "median": 0.125, "p10": 0.059027777777777776, "n": 36}, + null_p90=0.35714285714285715, + frac_contrasts_above_null_p90=1.0, + matched_projection_corr={"s0-s1": 1.0, "s0-s2": 1.0, "s1-s2": 1.0}, + nn_cos_used_median={"s0-s1": 1.0, "s0-s2": 1.0, "s1-s2": 1.0}, + ), + False: dict( # independently seeded dictionaries + same_side={"mean": 0.39943998902332234, "median": 0.39230769230769236, "p10": 0.25, "n": 72}, + cross_side={"mean": 0.15661895520964147, "median": 0.13942307692307693, "p10": 0.0, "n": 36}, + cross_contrast_null={"mean": 0.13614265182892632, "median": 0.125, "p10": 0.05277777777777778, "n": 36}, + null_p90=0.2857142857142857, + frac_contrasts_above_null_p90=0.8055555555555556, + matched_projection_corr={"s0-s1": 0.8081061840057373, "s0-s2": 0.7951256632804871, "s1-s2": 0.6831156015396118}, + nn_cos_used_median={"s0-s1": 0.7546381950378418, "s0-s2": 0.7305821180343628, "s1-s2": 0.7207158505916595}, + ), +} + +LOO_PINS = { + True: dict( + n_loo_cells=72, + core_median_size=9.0, + srp_heldout_seed_recall=1.0, + srp_ci95=[1.0, 1.0], + knn_recall_same_cores=0.4006, + knn_ci95=[0.3556, 0.4554], + cluster_recall_same_cores=0.3197, + cluster_ci95=[0.2873, 0.3823], + paper_dict_recall_same_cores=1.0, + ratio_srp_over_knn=2.496, + ratio_ci95=[2.1961, 2.8125], + ), + False: dict( + n_loo_cells=72, + core_median_size=5.0, + srp_heldout_seed_recall=0.697, + srp_ci95=[0.6543, 0.7468], + knn_recall_same_cores=0.4882, + knn_ci95=[0.4338, 0.5503], + cluster_recall_same_cores=0.4344, + cluster_ci95=[0.3674, 0.506], + paper_dict_recall_same_cores=0.9053, + ratio_srp_over_knn=1.428, + ratio_ci95=[1.2778, 1.6402], + ), +} + +# feature_group_matching: only the first seed pair's null-independent fields +# were reproducible before 0.2.1 (the null draws depended on PYTHONHASHSEED and +# their RNG consumption leaked into the query sampling of later pairs). +FGM_S0S1_PINS = { + True: dict( + best_jaccard={"mean": 1.0, "median": 1.0, "p90": 1.0, "p99": 1.0, "n": 10}, + frac_jaccard_ge_050=1.0, + frac_jaccard_ge_025=1.0, + recall_at_1={"mean": 1.0, "median": 1.0, "p90": 1.0, "p99": 1.0, "n": 10}, + recall_at_3={"mean": 1.0, "median": 1.0, "p90": 1.0, "p99": 1.0, "n": 10}, + matched_decoder_cosine_mean=0.9688253819942474, + features=[47, 26, 0, 52, 39, 1, 36, 16, 45, 49], + best_jaccards=[1.0] * 10, + best_features=[61, 26, 0, 52, 39, 1, 36, 16, 45, 49], + cosines=[0.6883, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0, 1.0], + ), + False: dict( + best_jaccard={ + "mean": 0.6133333333333332, + "median": 0.6, + "p90": 0.6399999999999998, + "p99": 0.9640000000000001, + "n": 10, + }, + frac_jaccard_ge_050=0.9, + frac_jaccard_ge_025=1.0, + recall_at_1={"mean": 0.75, "median": 0.75, "p90": 0.7749999999999999, "p99": 0.9775, "n": 10}, + recall_at_3={"mean": 1.0, "median": 1.0, "p90": 1.0, "p99": 1.0, "n": 10}, + matched_decoder_cosine_mean=0.5398867003619671, + features=[47, 26, 0, 52, 39, 1, 36, 16, 45, 49], + best_jaccards=[0.6, 0.6, 1.0, 0.6, 0.6, 0.6, 0.6, 0.6, 0.3333, 0.6], + best_features=[58, 5, 61, 24, 54, 59, 46, 44, 51, 52], + cosines=[0.6178, 0.6934, 0.8721, 0.5615, 0.5403, 0.4987, 0.665, 0.4312, 0.0773, 0.4416], + ), +} + +# feature_group_matching null-dependent values, recorded AFTER the 0.2.1 sorted-pool +# fix: (null99, frac_above_null_p99, direction_check.n_below_null_p99, sampled query +# features). Reproducible in every process now; before the fix they moved with the +# hash seed (the first pair's features excepted). +FGM_NULL_PINS = { + True: { + "s0->s1": (0.6, 1.0, 0, [47, 26, 0, 52, 39, 1, 36, 16, 45, 49]), + "s0->s2": (0.6, 1.0, 0, [59, 31, 34, 45, 16, 27, 53, 40, 10, 25]), + "s1->s2": (0.6, 1.0, 0, [35, 59, 61, 24, 26, 45, 44, 33, 0, 38]), + }, + False: { + "s0->s1": (0.6, 0.1, 9, [47, 26, 0, 52, 39, 1, 36, 16, 45, 49]), + "s0->s2": (0.6, 0.0, 10, [59, 31, 34, 45, 16, 27, 53, 40, 10, 25]), + "s1->s2": (0.6, 0.1, 9, [35, 59, 61, 25, 28, 44, 56, 34, 0, 39]), + }, +} + +# loo_core_recovery with --no-knn-exclude-self (the target row admitted to its own kNN +# side set): kNN recall and the SRP/kNN ratio, recorded on the consolidated code. +LOO_INCLUDE_SELF_PINS = {True: (0.4377, 2.285), False: (0.6453, 1.08)} + +FIRST_PAIRS = [("w18", "w19", 18, 19), ("w04", "w05", 4, 5), ("w14", "w15", 14, 15)] + + +def _assert_stats(got: dict, want: dict, tol: float) -> None: + assert set(got) == set(want) + assert got["n"] == want["n"] + for key in want: + if key != "n": + assert got[key] == pytest.approx(want[key], abs=tol), key + + +# --------------------------------------------------------------------------- # +# shared helpers +# --------------------------------------------------------------------------- # + + def test_dictionary_and_bank_helpers(tmp_path: Path) -> None: - tok, W, _, dicts, bank = _inputs(tmp_path, identical=False) - dec, enc_w, enc_b, k = dicts[0] - assert dec.shape == (D_FEAT, D_MODEL) and enc_w.shape == (D_FEAT, D_MODEL) and enc_b.shape == (D_FEAT,) - assert k == K - assert single_token_id(tok, "w03") == 3 and single_token_id(tok, "notaword") is None + tok, W, h_LN, dicts, bank = _inputs(tmp_path, identical=False) + d = dicts[0] + assert isinstance(d, Dictionary) + assert d.decoder.shape == (D_FEAT, D_MODEL) and d.encoder_w.shape == (D_FEAT, D_MODEL) + assert d.encoder_b.shape == (D_FEAT,) and d.k == K and d.row_mean is None + assert d[0] is d.decoder and d[3] == K # positional layout of the former 4-tuple + assert resolve_single_token_bare_first(tok, "w03") == 3 and resolve_single_token_bare_first(tok, "notaword") is None pairs = load_contrast_pairs(bank, tok, np.random.default_rng(0), max_contrasts=10) assert len(pairs) == 10 and len({(a, b) for a, b, _, _ in pairs}) == 10 assert pairs == load_contrast_pairs(bank, tok, np.random.default_rng(0), max_contrasts=10) - codes = topk_codes(W[:5], enc_w, enc_b, K) + assert load_contrast_pairs(bank, tok, np.random.default_rng(0), 150)[:3] == FIRST_PAIRS + codes = encode_topk(W[:5], d.encoder_w, d.encoder_b, K) assert codes.shape == (5, D_FEAT) and bool(((codes > 0).sum(dim=1) <= K).all()) + W2, h2, mask = load_readout(tmp_path / "w_u.pt") + assert torch.equal(W2, W) and torch.equal(h2, h_LN) and mask is not None and int(mask.sum()) == V - 1 -def test_cross_seed_identical_dictionaries_reproduce(tmp_path: Path) -> None: - tok, W, h_LN, dicts, bank = _inputs(tmp_path, identical=True) - mod = _load_script("cross_seed_stability") +def test_load_dictionary_surfaces_training_row_mean(tmp_path: Path) -> None: + _, _, _, dicts, _ = _inputs(tmp_path, identical=True) + ckpt = torch.load(tmp_path / "ckpt_s0.pt", weights_only=True) + ckpt["row_mean"] = torch.full((D_MODEL,), 0.25) + torch.save(ckpt, tmp_path / "ckpt_rm.pt") + d = load_dictionary(tmp_path / "ckpt_rm.pt") + assert torch.equal(d.row_mean, torch.full((D_MODEL,), 0.25)) and d.k == K + assert torch.equal(d.encoder_w, dicts[0].encoder_w) + + +def test_center_rows_matches_full_vocabulary_centering(tmp_path: Path) -> None: + # pins the pre-0.2.1 center_rows(W) arithmetic exactly (W - W.mean(0), norms floored at 1e-8) + _, W, _, dicts, _ = _inputs(tmp_path, identical=True) + row_mean = _live_mean(W, dicts) + assert torch.equal(row_mean, W.mean(0)) + W_c, rn, W_n = center_rows(W, row_mean) + assert torch.equal(W_c, W - W.mean(0)) + assert torch.equal(rn, (W - W.mean(0)).norm(dim=1).clamp_min(1e-8)) + assert torch.equal(W_n, W_c / rn[:, None]) + + +def test_resolve_centering_modes(tmp_path: Path) -> None: + tok, W, _, dicts, _ = _inputs(tmp_path, identical=True) + mask = torch.ones(V, dtype=torch.bool) + mask[:5] = False + assert torch.equal(resolve_centering(W, dicts, "live", mask), W.mean(0)) # live ignores the mask + assert torch.equal(resolve_centering(W, dicts, "trained", mask), W[mask].mean(0)) # nothing stored + assert torch.equal(resolve_centering(W, dicts, "trained", None, tok=tok), W[1:].mean(0)) # mask from tok + assert torch.equal(resolve_centering(W, dicts, "trained", None), W.mean(0)) # no mask, no tok, no store + stored = [d._replace(row_mean=torch.full((D_MODEL,), 0.5)) for d in dicts] + assert torch.equal(resolve_centering(W, stored, "trained", mask), torch.full((D_MODEL,), 0.5)) + assert torch.equal(resolve_centering(W, stored, "live", mask), W.mean(0)) + mixed = [stored[0], stored[1]._replace(row_mean=torch.zeros(D_MODEL)), stored[2]] + with pytest.raises(ValueError): + resolve_centering(W, mixed, "trained", mask) + + +def test_topk_keeps_exactly_k_at_ties() -> None: + # documented 0.2.1 change: the former >= rule kept both tied 2.0 codes at k=1 + codes = encode_topk(torch.tensor([[2.0, 2.0, 1.0]]), torch.eye(3), torch.zeros(3), 1) + assert int((codes != 0).sum()) == 1 and float(codes.max()) == 2.0 + + +def test_shared_stats_and_unit_helpers() -> None: + assert list(summary_stats([1.0, 2.0, 4.0], (10,))) == ["mean", "median", "p10", "n"] + assert list(summary_stats([1.0, 2.0, 4.0], (90, 99))) == ["mean", "median", "p90", "p99", "n"] + s = summary_stats([1.0, 2.0, 4.0], (90,)) + assert s["mean"] == pytest.approx(7 / 3) and s["median"] == 2.0 and s["n"] == 3 + u = unit_rows(torch.tensor([[3.0, 4.0], [0.0, 0.0]])) + assert torch.allclose(u[0], torch.tensor([0.6, 0.8])) and torch.equal(u[1], torch.zeros(2)) + + +# --------------------------------------------------------------------------- # +# cross_seed_stability.py +# --------------------------------------------------------------------------- # + + +def _run_cross_seed(tmp_path: Path, identical: bool) -> dict: + tok, W, h_LN, dicts, bank = _inputs(tmp_path, identical=identical) + mod = load_script("scripts/eval/cross_seed_stability.py") rng = np.random.default_rng(0) pairs = load_contrast_pairs(bank, tok, rng, 150) - out = mod.run( - dicts, [0, 1, 2], W, h_LN, tok, pairs, rng, top_m=3, top_r=4, n_sample=16, n_hidden=8, width_tag="test" + return mod.run( + dicts, + [0, 1, 2], + W, + h_LN, + tok, + pairs, + rng, + row_mean=_live_mean(W, dicts), + top_m=3, + top_r=4, + n_sample=16, + n_hidden=8, + width_tag="test", ) - assert out["n_contrasts"] == 12 and out["k"] == K + + +@pytest.mark.parametrize("identical", [True, False]) +def test_cross_seed_pins(tmp_path: Path, identical: bool) -> None: + out = _run_cross_seed(tmp_path, identical) + pin = CROSS_SEED_PINS[identical] + assert out["n_contrasts"] == 12 and out["k"] == K and out["width"] == "test" + for key in ("same_side", "cross_side", "cross_contrast_null"): + _assert_stats(out[key], pin[key], SET_TOL) + assert out["null_p90"] == pytest.approx(pin["null_p90"], abs=SET_TOL) + assert out["frac_contrasts_above_null_p90"] == pytest.approx(pin["frac_contrasts_above_null_p90"], abs=SET_TOL) + assert set(out["basis_nn_cosine"]) == set(out["matched_projection_corr"]) == {"s0-s1", "s0-s2", "s1-s2"} + for key, r in pin["matched_projection_corr"].items(): + assert out["matched_projection_corr"][key] == pytest.approx(r, abs=F32_TOL) + for key, med in pin["nn_cos_used_median"].items(): + assert out["basis_nn_cosine"][key]["nn_cos_used"]["median"] == pytest.approx(med, abs=F32_TOL) + + +def test_cross_seed_identical_dictionaries_reproduce(tmp_path: Path) -> None: + out = _run_cross_seed(tmp_path, identical=True) assert out["same_side"]["mean"] == 1.0 assert out["frac_contrasts_above_null_p90"] == 1.0 + assert out["same_side"]["n"] == 3 * 12 * 2 # seed pairs x contrasts x sides for v in out["basis_nn_cosine"].values(): assert abs(v["nn_cos_used"]["median"] - 1.0) < 1e-5 - for r in out["matched_projection_corr"].values(): - assert abs(r - 1.0) < 1e-5 -def test_cross_seed_different_seeds_in_range(tmp_path: Path) -> None: - tok, W, h_LN, dicts, bank = _inputs(tmp_path, identical=False) - mod = _load_script("cross_seed_stability") - rng = np.random.default_rng(0) - pairs = load_contrast_pairs(bank, tok, rng, 150) - out = mod.run( - dicts, [0, 1, 2], W, h_LN, tok, pairs, rng, top_m=3, top_r=4, n_sample=16, n_hidden=8, width_tag="test" - ) - for key in ("same_side", "cross_side", "cross_contrast_null"): - assert 0.0 <= out[key]["mean"] <= 1.0 - assert out["same_side"]["n"] == 3 * 12 * 2 # seed pairs x contrasts x sides - assert set(out["basis_nn_cosine"]) == {"s0-s1", "s0-s2", "s1-s2"} - assert set(out["matched_projection_corr"]) == {"s0-s1", "s0-s2", "s1-s2"} +# --------------------------------------------------------------------------- # +# feature_group_matching.py +# --------------------------------------------------------------------------- # -def _run_matching(mod, tok, W, dicts, bank): +def _run_matching(tmp_path: Path, identical: bool) -> dict: + tok, W, _, dicts, bank = _inputs(tmp_path, identical=identical) + mod = load_script("scripts/eval/feature_group_matching.py") rng = np.random.default_rng(0) pairs = load_contrast_pairs(bank, tok, rng, 150) return mod.run( @@ -130,6 +366,7 @@ def _run_matching(mod, tok, W, dicts, bank): pairs, rng, torch.device("cpu"), + row_mean=_live_mean(W, dicts), top_m=3, top_r=4, n_query=10, @@ -144,12 +381,76 @@ def _run_matching(mod, tok, W, dicts, bank): ) -def test_feature_group_matching_identical_dictionaries(tmp_path: Path) -> None: - tok, W, _, dicts, bank = _inputs(tmp_path, identical=True) - out = _run_matching(_load_script("feature_group_matching"), tok, W, dicts, bank) +@pytest.mark.parametrize("identical", [True, False]) +def test_feature_group_matching_first_pair_pins(tmp_path: Path, identical: bool) -> None: + out = _run_matching(tmp_path, identical) assert set(out["pairs"]) == {"s0->s1", "s0->s2", "s1->s2"} + res = out["pairs"]["s0->s1"] + pin = FGM_S0S1_PINS[identical] + assert res["n_query"] == 10 and len(res["per_query"]) == 10 + for key in ("best_jaccard", "recall_at_1", "recall_at_3"): + _assert_stats(res[key], pin[key], SET_TOL) + assert res["frac_jaccard_ge_050"] == pytest.approx(pin["frac_jaccard_ge_050"], abs=SET_TOL) + assert res["frac_jaccard_ge_025"] == pytest.approx(pin["frac_jaccard_ge_025"], abs=SET_TOL) + assert res["matched_decoder_cosine"]["mean"] == pytest.approx(pin["matched_decoder_cosine_mean"], abs=F32_TOL) + assert [q["feature"] for q in res["per_query"]] == pin["features"] + assert [q["best_jaccard"] for q in res["per_query"]] == pin["best_jaccards"] + assert [q["best_feature"] for q in res["per_query"]] == pin["best_features"] + assert [q["cosine"] for q in res["per_query"]] == pin["cosines"] + + +@pytest.mark.parametrize("identical", [True, False]) +def test_feature_group_matching_null_pins(tmp_path: Path, identical: bool) -> None: + out = _run_matching(tmp_path, identical) + for key, (null99, frac_above, n_below, features) in FGM_NULL_PINS[identical].items(): + res = out["pairs"][key] + assert res["null99"] == pytest.approx(null99, abs=SET_TOL), key + assert res["frac_above_null_p99"] == pytest.approx(frac_above, abs=SET_TOL), key + assert res["direction_check"]["n_below_null_p99"] == n_below, key + assert [q["feature"] for q in res["per_query"]] == features, key + + +def test_feature_group_matching_tail_agrees_with_frac_above_null(tmp_path: Path) -> None: + # 0.2.1 fix: the direction-check tail and frac_above_null_p99 are complementary counts + out = _run_matching(tmp_path, identical=False) + for res in out["pairs"].values(): + n_above = round(res["frac_above_null_p99"] * res["n_query"]) + assert res["direction_check"]["n_below_null_p99"] == res["n_query"] - n_above + + +def test_null_token_pool_is_sorted_and_hash_seed_independent() -> None: + # 0.2.1 fix: the null pool is enumerated in sorted order, so the pseudo-group draws + # (and the RNG state they consume) no longer depend on PYTHONHASHSEED + script = REPO_ROOT / "scripts" / "eval" / "feature_group_matching.py" + child = textwrap.dedent( + f""" + import importlib.util, json, sys + import numpy as np + spec = importlib.util.spec_from_file_location("fgm_child", {str(script)!r}) + mod = importlib.util.module_from_spec(spec) + sys.modules["fgm_child"] = mod + spec.loader.exec_module(mod) + groups = [frozenset(g) for g in (("a", "b", "c"), ("b", "c", "d"), ("e", "f", "a"), ("x", "y", "z", "w"))] + toks, p = mod.null_token_pool(groups) + draw = np.random.default_rng(0).choice(toks, size=3, replace=False, p=p).tolist() + print(json.dumps({{"toks": toks, "p": p.tolist(), "draw": draw}})) + """ + ) + outs = [] + for hash_seed in ("1", "2"): + env = {**os.environ, "PYTHONHASHSEED": hash_seed, "PYTHONPATH": str(REPO_ROOT / "src")} + proc = subprocess.run([sys.executable, "-c", child], capture_output=True, text=True, env=env, check=True) + outs.append(json.loads(proc.stdout.strip().splitlines()[-1])) + assert outs[0] == outs[1] + assert outs[0]["toks"] == ["a", "b", "c", "d", "e", "f", "w", "x", "y", "z"] + assert outs[0]["p"] == pytest.approx( + [2 / 13, 2 / 13, 2 / 13, 1 / 13, 1 / 13, 1 / 13, 1 / 13, 1 / 13, 1 / 13, 1 / 13] + ) + + +def test_feature_group_matching_identical_dictionaries(tmp_path: Path) -> None: + out = _run_matching(tmp_path, identical=True) for res in out["pairs"].values(): - assert res["n_query"] == 10 and len(res["per_query"]) == 10 assert res["best_jaccard"]["median"] == 1.0 assert res["frac_jaccard_ge_050"] == 1.0 assert res["recall_at_3"]["median"] == 1.0 @@ -160,8 +461,7 @@ def test_feature_group_matching_identical_dictionaries(tmp_path: Path) -> None: def test_feature_group_matching_direction_check_annotates_tail(tmp_path: Path) -> None: - tok, W, _, dicts, bank = _inputs(tmp_path, identical=False) - out = _run_matching(_load_script("feature_group_matching"), tok, W, dicts, bank) + out = _run_matching(tmp_path, identical=False) for res in out["pairs"].values(): check = res["direction_check"] fails = [q for q in res["per_query"] if q["best_jaccard"] <= res["null99"]] @@ -172,17 +472,23 @@ def test_feature_group_matching_direction_check_annotates_tail(tmp_path: Path) - assert 0.0 <= res["null_best_jaccard"]["p99"] <= 1.0 -def test_loo_core_recovery_identical_dictionaries(tmp_path: Path) -> None: - tok, W, _, dicts, bank = _inputs(tmp_path, identical=True) - mod = _load_script("loo_core_recovery") +# --------------------------------------------------------------------------- # +# loo_core_recovery.py +# --------------------------------------------------------------------------- # + + +def _run_loo(tmp_path: Path, identical: bool, **overrides) -> dict: + tok, W, _, dicts, bank = _inputs(tmp_path, identical=identical) + mod = load_script("scripts/eval/loo_core_recovery.py") rng = np.random.default_rng(0) pairs = load_contrast_pairs(bank, tok, rng, 150) - out = mod.run( + return mod.run( dicts, W, tok, pairs, torch.device("cpu"), + row_mean=_live_mean(W, dicts), top_m=3, top_r=4, side_n=6, @@ -193,11 +499,50 @@ def test_loo_core_recovery_identical_dictionaries(tmp_path: Path) -> None: bootstrap_seed=1, reference_dict=dicts[0], width_tag="test", + **overrides, ) + + +@pytest.mark.parametrize("identical", [True, False]) +def test_loo_core_recovery_pins(tmp_path: Path, identical: bool) -> None: + out = _run_loo(tmp_path, identical) + pin = LOO_PINS[identical] + assert out["width"] == "test" + for key, want in pin.items(): + assert out[key] == want, key + + +@pytest.mark.parametrize("identical", [True, False]) +def test_loo_no_knn_exclude_self(tmp_path: Path, identical: bool) -> None: + out = _run_loo(tmp_path, identical, knn_exclude_self=False) + knn, ratio = LOO_INCLUDE_SELF_PINS[identical] + assert out["knn_recall_same_cores"] == knn and out["ratio_srp_over_knn"] == ratio + # only the kNN side changes; the SRP and cluster sides never excluded the target + for key in ("srp_heldout_seed_recall", "cluster_recall_same_cores", "core_median_size", "n_loo_cells"): + assert out[key] == LOO_PINS[identical][key], key + + +def test_knn_side_sets_self_exclusion(tmp_path: Path) -> None: + tok, W, _, dicts, bank = _inputs(tmp_path, identical=True) + mod = load_script("scripts/eval/loo_core_recovery.py") + pairs = load_contrast_pairs(bank, tok, np.random.default_rng(0), 150) + _, _, W_n = center_rows(W, _live_mean(W, dicts)) + excl = mod.knn_side_sets(W_n, pairs, 6, tok, {}, exclude_self=True) + incl = mod.knn_side_sets(W_n, pairs, 6, tok, {}, exclude_self=False) + for (a, b, _ia, _ib), e, i in zip(pairs, excl, incl): + assert a not in e[0] and b not in e[1] + assert a in i[0] and b in i[1] # the target's own row (cosine 1) takes a slot + assert len(e[0]) == len(i[0]) == 6 + + +def test_loo_core_recovery_identical_dictionaries(tmp_path: Path) -> None: + out = _run_loo(tmp_path, identical=True) assert out["n_loo_cells"] == 12 * 2 * 3 # contrasts x sides x held-out seeds assert out["srp_heldout_seed_recall"] == 1.0 assert out["srp_ci95"] == [1.0, 1.0] assert out["paper_dict_recall_same_cores"] == 1.0 - assert 0.0 <= out["knn_recall_same_cores"] <= 1.0 - assert 0.0 <= out["cluster_recall_same_cores"] <= 1.0 - assert out["ratio_srp_over_knn"] >= 1.0 + # the kNN control is neither trivially perfect nor empty on this case, so the + # ratio is informative: it must be the SRP / kNN quotient of the same output + knn = out["knn_recall_same_cores"] + assert 0.0 < knn < 1.0 and 0.0 < out["cluster_recall_same_cores"] < 1.0 + assert out["ratio_srp_over_knn"] == round(out["srp_heldout_seed_recall"] / knn, 3) diff --git a/tests/test_wsd_sense_groups.py b/tests/test_wsd_sense_groups.py index 2e11292..368600e 100644 --- a/tests/test_wsd_sense_groups.py +++ b/tests/test_wsd_sense_groups.py @@ -1,36 +1,324 @@ -"""CPU smoke tests for the CoarseWSD-20 sense analyses on a synthetic bundle. +"""CoarseWSD-20 sense analyses and the bundle writer, pinned on synthetic fixtures. -* scripts/analyze/analyze_wsd_sense_groups.py (tab:app-sense-alignment) +* scripts/run/run_wsd_feature_alignment.py (bundle writer, full-vector references) +* scripts/analyze/analyze_wsd_sense_groups.py (tab:app-sense-alignment) * scripts/analyze/analyze_wsd_classifier_framing.py (classifier-framing paragraph) +* sparse_readout_prism.research.wsd (their shared helpers) -The bundle mimics ``representations.pt`` from ``run_wsd_feature_alignment.py``: -two words, two or three senses, 16 dictionary positions, with a strong and a -weaker planted discriminative position per sense. No model, no network, well under a second. +``PINS`` holds outputs of the pre-0.2.1 scripts (per-script helper copies, +right truncation, glob-picked HF snapshot) on the fixtures below: hashes of +canonical JSON or array bytes plus a few readable values. The tests assert the +refactored code reproduces them exactly. The bundle mimics +``representations.pt``: two words, two or three senses, 16 dictionary +positions, a strong and a weaker planted discriminative position per sense. +No model, no network, well under a second per test. """ from __future__ import annotations -import importlib.util +import hashlib import json -import sys +import math from pathlib import Path +from types import SimpleNamespace import numpy as np +import pytest import torch +from conftest import load_script +from sparse_readout_prism.data import centering_mean +from sparse_readout_prism.factorizers import build_factorizer +from sparse_readout_prism.research import wsd +from sparse_readout_prism.research.qwen_readout import load_sae +from sparse_readout_prism.utils import spearman -ROOT = Path(__file__).resolve().parents[1] +rw = load_script("scripts/run/run_wsd_feature_alignment.py") +sg = load_script("scripts/analyze/analyze_wsd_sense_groups.py") +cf = load_script("scripts/analyze/analyze_wsd_classifier_framing.py") - -def _load(mod_name: str, rel_path: str): - spec = importlib.util.spec_from_file_location(mod_name, ROOT / rel_path) - mod = importlib.util.module_from_spec(spec) - sys.modules[mod_name] = mod - spec.loader.exec_module(mod) - return mod - - -sg = _load("analyze_wsd_sense_groups", "scripts/analyze/analyze_wsd_sense_groups.py") -cf = _load("analyze_wsd_classifier_framing", "scripts/analyze/analyze_wsd_classifier_framing.py") +PINS = { + "auc": {"big_0": 0.5645077524610911, "big_3": 0.5645737167712954, "one_class": None, "small_0": 0.4735317279695978}, + "cb": { + "ci": [-0.14836898540991997, 0.3938881306083598], + "first3": [0.10616698864212924, 0.04648800212459197, 0.05718697644108653], + "n": 50, + }, + "cf": { + "cmp_projection": { + "accuracy_difference": 0.0, + "accuracy_difference_ci": [0.0, 0.0], + "accuracy_positive_words": 0, + "balanced_accuracy_difference": 0.0, + "balanced_accuracy_difference_ci": [0.0, 0.0], + "balanced_accuracy_positive_words": 0, + "macro_f1_difference": 0.0, + "macro_f1_difference_ci": [0.0, 0.0], + "macro_f1_positive_words": 0, + "n_words": 2, + }, + "gate": { + "median_rho_plus_0p5": 0.0, + "n_items": 120, + "n_pass": 120, + "pass_fraction": 1.0, + "sign_match_fraction": 1.0, + }, + "hash": "7e49e4804b5837ce", + "null_shuffled": { + "macro_f1": { + "fraction_null_at_least_observed": 1.0, + "max": 1.0, + "mean": 1.0, + "min": 1.0, + "n_seeds": 2, + "observed_srp": 1.0, + "q05_q95": [1.0, 1.0], + "std": 0.0, + } + }, + "primary_methods": { + "projection_only": { + "accuracy": 1.0, + "balanced_accuracy": 1.0, + "macro_f1": 1.0, + "macro_f1_ci": [1.0, 1.0], + "n_test": 40, + "n_train": 80, + "n_words": 2, + }, + "random_srp_features": { + "accuracy": 0.5, + "balanced_accuracy": 0.5805555555555555, + "macro_f1": 0.48367027970608534, + "macro_f1_ci": [0.32539682539682535, 0.48367027970608534], + "n_test": 40, + "n_train": 80, + "n_words": 2, + }, + "shuffled_srp": { + "accuracy": 1.0, + "balanced_accuracy": 1.0, + "macro_f1": 1.0, + "macro_f1_ci": [1.0, 1.0], + "n_test": 40, + "n_train": 80, + "n_words": 2, + }, + "srp": { + "accuracy": 1.0, + "balanced_accuracy": 1.0, + "macro_f1": 1.0, + "macro_f1_ci": [1.0, 1.0], + "n_test": 40, + "n_train": 80, + "n_words": 2, + }, + }, + "results_signed_4_srp_agg": { + "accuracy": 1.0, + "balanced_accuracy": 1.0, + "macro_f1": 1.0, + "n_test": 40, + "n_train": 80, + "n_words": 2, + }, + }, + "cf_cli": { + "analysis": { + "encodings": ["signed", "weighted"], + "ks": [2, 4], + "methods": ["srp", "projection_only", "shuffled_srp", "random_srp_features"], + "n_boot": 7, + "n_null_seeds": 3, + "primary_encoding": "signed", + "primary_k": 4, + "seed": 5, + }, + "hash": "7da70bb3210641aa", + "primary_srp": { + "accuracy": 1.0, + "balanced_accuracy": 1.0, + "macro_f1": 1.0, + "macro_f1_ci": [1.0, 1.0], + "n_test": 40, + "n_train": 80, + "n_words": 2, + }, + }, + "cf_random_enc_hash": "d8dfe216739112f6", + "cf_shuffled_hash": {"0": "aaff27a124e088d4", "3": "c164033316962177"}, + "ptc": { + "audit": { + "n_targets_requested": 3, + "n_targets_single_token": 2, + "single_token_fraction": 0.6666666666666666, + "skipped_targets": {"multi": "multi_token"}, + "targets": { + "bank": { + "continuation": " bank", + "n_positive_codes": 6, + "row_cosine": 0.3883303105831146, + "row_relative_error": 3.34517765045166, + "token_id": 3, + }, + "seal": { + "continuation": " seal", + "n_positive_codes": 6, + "row_cosine": -0.21251048147678375, + "row_relative_error": 2.8938305377960205, + "token_id": 4, + }, + }, + }, + "bank_beta": "3770c9b09cff729c", + "bank_fids": [21, 2, 24, 17, 8, 5], + "k": 6, + "seal_beta": "a378b719424cc9bb", + "seal_fids": [8, 25, 31, 26, 28, 21], + }, + "rw_method_hashes": { + "hidden": "710352578f8a4b30", + "shuffled_srp": "ab38110ffb3cb5d8", + "srp": "f203f93dc79b5f48", + "support_projection": "9e00ed52bd0f0d38", + "token_only": "4f22f486a7f821af", + }, + "rw_metrics": { + "bank_hidden_per_word": { + "accuracy": 1.0, + "ari": 1.0, + "balanced_accuracy": 1.0, + "macro_f1": 1.0, + "n_senses": 2, + "n_test": 20, + "n_train": 40, + "nmi": 1.0, + "pairwise_auc": 1.0, + "row_relative_error": 0.2, + }, + "cmp_hidden": { + "accuracy_difference": 0.0, + "accuracy_difference_ci": [0.0, 0.0], + "ari_difference": 0.0, + "ari_difference_ci": [0.0, 0.0], + "balanced_accuracy_difference": 0.0, + "balanced_accuracy_difference_ci": [0.0, 0.0], + "macro_f1_difference": 0.0, + "macro_f1_difference_ci": [0.0, 0.0], + "n_words": 2, + "nmi_difference": 0.0, + "nmi_difference_ci": [0.0, 0.0], + "pairwise_auc_difference": 0.0, + "pairwise_auc_difference_ci": [0.0, 0.0], + }, + "hash": "8a1eab9e0f070776", + "srp_aggregate": { + "accuracy": 1.0, + "accuracy_ci": [1.0, 1.0], + "ari": 1.0, + "ari_ci": [1.0, 1.0], + "balanced_accuracy": 1.0, + "balanced_accuracy_ci": [1.0, 1.0], + "macro_f1": 1.0, + "macro_f1_ci": [1.0, 1.0], + "n_words": 2, + "nmi": 1.0, + "nmi_ci": [1.0, 1.0], + "pairwise_auc": 1.0, + "pairwise_auc_ci": [1.0, 1.0], + "row_gate": { + "accuracy": 1.0, + "ari": 1.0, + "balanced_accuracy": 1.0, + "macro_f1": 1.0, + "nmi": 1.0, + "pairwise_auc": 1.0, + }, + "row_gate_fraction": 1.0, + "row_gate_words": 2, + }, + }, + "rw_predictions_hash": "370a835c74b3674d", + "rw_scoring_summary": { + "mean_absolute_logit_residual": 0.0, + "median_absolute_logit_residual": 0.0, + "median_target_rank": 1.0, + "n_scored": 120, + "n_targets": 2, + "row_gate_fraction_items": 1.0, + "target_top10_fraction": 1.0, + "target_top1_fraction": 1.0, + }, + "rw_shuffled_seed3": "fefe77440592b1af", + "score": { + "beta": "4aa2fd96e66a070b", + "contribution": "7b2444e5e923ea31", + "exact_logit": "7b694ac7b5c01f09", + "feature_ids": "e1fe1b77ec0dadd3", + "hidden": "ae1639028022b00a", + "metadata": "aad43d9a08a44c58", + "n": 4, + "projection": "6d0a986328961978", + "reconstructed_logit": "f3fe5890e11445ad", + "target_logprob": "b0ec08f04887fc35", + "target_rank": "b0f18af3aeb57306", + }, + "score_summary": { + "mean_absolute_logit_residual": 1.5911362171173096, + "median_absolute_logit_residual": 1.7451845407485962, + "median_target_rank": 11.0, + "n_scored": 4, + "n_targets": 2, + "row_gate_fraction_items": 0.0, + "target_top10_fraction": 0.25, + "target_top1_fraction": 0.0, + }, + "sg_g1": { + "bank_anchors": {"0": [3], "1": [7]}, + "bank_balanced_ci": [1.0, 1.0], + "hash": "67da5f98b744d457", + "seal_anchors": {"0": [1], "1": [9], "2": [13]}, + "seal_null_balanced_mean": 0.30666666666666664, + "seal_null_balanced_p95": 0.48305555555555546, + "word_mean_balanced": 1.0, + "word_mean_balanced_ci": [1.0, 1.0], + "word_mean_full_account_balanced": 1.0, + "word_mean_hidden_balanced": 1.0, + "word_mean_majority_balanced": 0.41666666666666663, + "word_mean_null_balanced": 0.3616666666666667, + }, + "sg_g2": { + "bank_anchors": {"0": [3, 4], "1": [7, 8]}, + "bank_balanced_ci": [1.0, 1.0], + "hash": "c3b3927e799785e3", + "seal_anchors": {"0": [1, 2], "1": [9, 10], "2": [13, 14]}, + "seal_null_balanced_mean": 0.33999999999999997, + "seal_null_balanced_p95": 0.5816666666666666, + "word_mean_balanced": 1.0, + "word_mean_balanced_ci": [1.0, 1.0], + "word_mean_full_account_balanced": 1.0, + "word_mean_hidden_balanced": 1.0, + "word_mean_majority_balanced": 0.41666666666666663, + "word_mean_null_balanced": 0.3933333333333333, + }, + "sg_g4": { + "bank_anchors": {"0": [3, 4, 13, 1], "1": [7, 8, 10, 15]}, + "bank_balanced_ci": [1.0, 1.0], + "hash": "f98c9b8b349b06bf", + "seal_anchors": {"0": [1, 2, 12, 6], "1": [9, 10, 11, 0], "2": [13, 14, 7, 5]}, + "seal_null_balanced_mean": 0.3827777777777778, + "seal_null_balanced_p95": 0.6624999999999998, + "word_mean_balanced": 0.95, + "word_mean_balanced_ci": [0.9, 1.0], + "word_mean_full_account_balanced": 1.0, + "word_mean_hidden_balanced": 1.0, + "word_mean_majority_balanced": 0.41666666666666663, + "word_mean_null_balanced": 0.3697222222222223, + }, + "sg_raw_top1": {"bank_anchors": {"0": [4, 3], "1": [7]}, "hash": "4ac6081f915fc861", "word_mean_balanced": 1.0}, + "spearman": {"const": None, "ok": 0.7999999999999999, "short": None, "tiny": None}, + "stable_seed": {"bank_0": 3184672968, "random_bank_0": 3187410737, "seal_3": 252631022}, +} WIDTH = 16 D_HIDDEN = 8 @@ -39,12 +327,37 @@ def _load(mod_name: str, rel_path: str): "bank": {"0": (3, 30, 15), "1": (7, 10, 5)}, "seal": {"0": (1, 20, 10), "1": (9, 12, 6), "2": (13, 8, 4)}, } +VOCAB = [ + "", + "", + "", + " bank", + " seal", + "the", + "river", + "money", + "is", + "of", + "a", + "word", + "missing", + ":", + "x", + "y", +] +PROMPTS = [ + "the river bank is a word :", + "money of a bank is the word :", + "the seal of the x y :", + "a seal is a word missing :", + "the missing word is :", +] def make_bundle(seed: int = 0) -> dict: rng = np.random.default_rng(seed) - metadata, contribution, projection, hidden = [], [], [], [] - for word, senses in DESIGN.items(): + metadata, contribution, projection, hidden, betas = [], [], [], [], [] + for word_index, (word, senses) in enumerate(DESIGN.items()): beta = rng.uniform(0.5, 2.0, size=WIDTH).astype(np.float32) for sense, (pos, n_train, n_test) in senses.items(): for split, n in (("train", n_train), ("test", n_test)): @@ -62,12 +375,13 @@ def make_bundle(seed: int = 0) -> dict: "split": split, "target": word, "sense": sense, - "token_id": 100 + len(word), + "token_id": 100 + word_index, "row_relative_error": 0.2, } ) projection.append(p) contribution.append(beta * p) + betas.append(beta) hidden.append(h) contribution = torch.tensor(np.stack(contribution)) n = contribution.shape[0] @@ -76,16 +390,183 @@ def make_bundle(seed: int = 0) -> dict: "hidden": torch.tensor(np.stack(hidden)), "projection": torch.tensor(np.stack(projection)), "contribution": contribution, - "beta": torch.stack([torch.ones(WIDTH)] * n), + "beta": torch.tensor(np.stack(betas)), "feature_ids": torch.stack([torch.arange(WIDTH, dtype=torch.int32)] * n), "exact_logit": contribution.sum(dim=1) + 5.0, "reconstructed_logit": contribution.sum(dim=1) + 5.0, "target_logprob": torch.zeros(n), "target_rank": torch.ones(n, dtype=torch.int32), - "run": {"dataset": "coarsewsd20", "model_id": "synthetic", "k": WIDTH, "seed": seed}, + "run": { + "dataset": "coarsewsd20", + "model_id": "synthetic", + "k": WIDTH, + "seed": seed, + }, } +class _Batch(dict): + def to(self, _device): + return self + + +class FakeTokenizer: + """Whitespace tokenizer over a fixed vocabulary. ``' word'`` forms are single + tokens when present; unknown continuations tokenize to two ids.""" + + pad_token_id = 0 + pad_token = "" + eos_token = "" + eos_token_id = 1 + + def __init__(self, vocab: list[str], special_ids: tuple[int, ...] = (0, 1)): + self.vocab = list(vocab) + self.index = {t: i for i, t in enumerate(self.vocab)} + self.padding_side = "right" + self.truncation_side = "right" + self.all_special_ids = list(special_ids) + + def __len__(self) -> int: + return len(self.vocab) + + def encode(self, text: str, add_special_tokens: bool = False) -> list[int]: + if text in self.index: + return [self.index[text]] + return [2, 2] + + def _ids(self, text: str) -> list[int]: + return [self.index.get(w, 2) for w in text.split()] + + def __call__( + self, + texts, + return_tensors=None, + padding=True, + truncation=False, + max_length=None, + ): + seqs = [self._ids(t) for t in texts] + if not padding: + return _Batch(input_ids=seqs, attention_mask=[[1] * len(s) for s in seqs]) + if truncation and max_length is not None: + seqs = [s[-max_length:] if self.truncation_side == "left" else s[:max_length] for s in seqs] + width = max(len(s) for s in seqs) + ids, mask = [], [] + for s in seqs: + pad = [self.pad_token_id] * (width - len(s)) + if self.padding_side == "left": + ids.append(pad + s) + mask.append([0] * len(pad) + [1] * len(s)) + else: + ids.append(s + pad) + mask.append([1] * len(s) + [0] * len(pad)) + return _Batch( + input_ids=torch.tensor(ids, dtype=torch.long), + attention_mask=torch.tensor(mask, dtype=torch.long), + ) + + +class FakeLM(torch.nn.Module): + """Causal bag-of-tokens LM: the state at position t is the running mean of the + real-token embeddings up to t, so context (and truncation) changes the state.""" + + def __init__(self, vocab: int, d_model: int, seed: int = 0): + super().__init__() + gen = torch.Generator().manual_seed(seed) + self.table = torch.nn.Parameter(torch.randn(vocab, d_model, generator=gen)) + self.lm_head = torch.nn.Linear(d_model, vocab, bias=False) + with torch.no_grad(): + self.lm_head.weight.copy_(torch.randn(vocab, d_model, generator=gen)) + + def forward(self, input_ids=None, attention_mask=None, use_cache=False, **kwargs): + emb = self.table[input_ids] # (B, T, d) + m = attention_mask[..., None].to(emb.dtype) + h = (emb * m).cumsum(dim=1) / m.cumsum(dim=1).clamp_min(1.0) + return SimpleNamespace(logits=self.lm_head(h)) + + +def make_checkpoint( + path: Path, + d_model: int = 8, + d_features: int = 32, + k: int = 6, + *, + runner_layout: bool = False, +): + torch.manual_seed(0) + cfg = {"architecture": "topk", "d_features": d_features, "k": k} + model = build_factorizer({"factorizer": cfg}, d_model=d_model) + with torch.no_grad(): + model.encoder.weight.copy_(torch.randn(d_features, d_model)) + model.encoder.bias.copy_(0.1 * torch.randn(d_features)) + model.decoder.copy_(torch.randn(d_features, d_model)) + model.normalize_decoder_() + ckpt = {"model_state_dict": model.state_dict()} + if runner_layout: + ckpt["factorizer"] = cfg # what runner.py writes: no "config" key + else: + ckpt["config"] = {"factorizer": cfg} + torch.save(ckpt, path) + + +def _strip(obj, drop: set[str]): + if isinstance(obj, dict): + return {k: _strip(v, drop) for k, v in obj.items() if k not in drop} + if isinstance(obj, list): + return [_strip(v, drop) for v in obj] + if isinstance(obj, float) and math.isnan(obj): + return "NaN" + return obj + + +def _canon(obj, drop=()) -> str: + return hashlib.sha256(json.dumps(_strip(obj, set(drop)), sort_keys=True).encode()).hexdigest()[:16] + + +def _array_hash(array) -> str: + return hashlib.sha256(np.ascontiguousarray(np.asarray(array)).tobytes()).hexdigest()[:16] + + +def _nan_none(value): + return None if (isinstance(value, float) and math.isnan(value)) else value + + +def _scoring_setup(tmp_path: Path, *, runner_layout: bool = False): + path = tmp_path / "checkpoint.pt" + make_checkpoint(path, runner_layout=runner_layout) + tok = FakeTokenizer(VOCAB) + lm = FakeLM(len(VOCAB), D_HIDDEN, seed=0) + w_u = lm.lm_head.weight.detach().float() + w_dec_unit, w_enc, b_enc, sae_config, ckpt_row_mean = load_sae(path) + k = rw.checkpoint_k(path, sae_config) + row_mean = rw.scoring_row_mean(w_u, tok, "live", ckpt_row_mean) + info, audit = rw.prepare_target_codes(["bank", "seal", "multi"], tok, w_u, row_mean, w_enc, b_enc, w_dec_unit, k) + return tok, lm, w_u, row_mean, w_dec_unit, info, audit, k + + +def _items(prompts): + targets = ["bank", "bank", "seal", "seal", "multi"] + return [ + { + "item_id": f"i{j}", + "kind": "context", + "dataset": "coarsewsd20", + "split": "train" if j % 2 == 0 else "test", + "target": targets[j], + "sense": str(j % 2), + "sense_name": str(j % 2), + "prefix_words": 1, + "prompt": p, + } + for j, p in enumerate(prompts) + ] + + +# --------------------------------------------------------------------------- # +# sense groups +# --------------------------------------------------------------------------- # + + def _word_arrays(bundle: dict, word: str, split: str) -> tuple[np.ndarray, np.ndarray]: idx = [i for i, r in enumerate(bundle["metadata"]) if r["target"] == word and r["split"] == split] x = bundle["contribution"].numpy()[idx] @@ -131,71 +612,522 @@ def test_group_prediction_and_balanced_accuracy(): assert sg.balanced_accuracy(np.full_like(test_y, "0"), test_y) == 0.5 -def test_sense_groups_end_to_end_sweep(tmp_path, monkeypatch): +def test_sense_groups_cli_reproduces_pinned_outputs(tmp_path): bundle_path = tmp_path / "representations.pt" torch.save(make_bundle(), bundle_path) out = tmp_path / "out" / "synthetic__srp.json" - argv = [ - "analyze_wsd_sense_groups.py", - "--bundle", - str(bundle_path), - "--out", - str(out), - "--group-sizes", - "1,2", - "--n-boot", - "20", - "--n-null", - "10", - "--seed", - "0", - ] - monkeypatch.setattr(sys, "argv", argv) - assert sg.main() == 0 - for g in (1, 2): + base = ["--bundle", str(bundle_path), "--out", str(out)] + assert ( + sg.main( + base + + [ + "--group-sizes", + "1,2,4", + "--n-boot", + "20", + "--n-null", + "10", + "--seed", + "0", + ] + ) + == 0 + ) + for g in (1, 2, 4): path = out.with_name(f"synthetic__srp_g{g}.json") - assert path.exists(), path summary = json.loads(path.read_text()) + pin = PINS[f"sg_g{g}"] + assert _canon(summary, drop=("bundle", "provenance")) == pin["hash"], g + assert summary["word_mean_balanced"] == pin["word_mean_balanced"] + assert summary["word_mean_null_balanced"] == pin["word_mean_null_balanced"] + assert summary["word_mean_majority_balanced"] == pin["word_mean_majority_balanced"] == (1 / 2 + 1 / 3) / 2 + assert summary["word_mean_full_account_balanced"] == pin["word_mean_full_account_balanced"] + assert summary["word_mean_hidden_balanced"] == pin["word_mean_hidden_balanced"] + assert summary["word_mean_balanced_ci"] == pin["word_mean_balanced_ci"] + assert summary["per_word"]["bank"]["anchors"] == pin["bank_anchors"] + assert summary["per_word"]["seal"]["anchors"] == pin["seal_anchors"] + assert summary["per_word"]["bank"]["balanced_accuracy_ci"] == pin["bank_balanced_ci"] + assert summary["per_word"]["seal"]["null_balanced_p95"] == pin["seal_null_balanced_p95"] assert summary["group_size"] == g and summary["basis"] == "srp" and summary["standardized"] is True - assert summary["n_words"] == 2 and set(summary["per_word"]) == set(DESIGN) - assert summary["word_mean_balanced"] > 0.95 - assert summary["word_mean_null_balanced"] < 0.8 - assert summary["word_mean_majority_balanced"] == (1 / 2 + 1 / 3) / 2 - assert summary["words_beating_majority_balanced"] == 2 - assert summary["words_beating_null_p95_balanced"] == 2 - assert summary["word_mean_full_account_balanced"] > 0.95 - assert summary["word_mean_hidden_balanced"] > 0.95 + assert summary["words_beating_majority_balanced"] == 2 and summary["words_beating_null_p95_balanced"] == 2 assert summary["gate_fraction_overall"] == 1.0 + assert "command" in summary["provenance"] and summary["provenance"]["args"]["seed"] == 0 for word, senses in DESIGN.items(): pw = summary["per_word"][word] - assert pw["n_senses"] == len(senses) - assert all(len(pw["anchors"][s]) == g for s in senses) - assert pw["chance_balanced_accuracy"] == 1 / len(senses) - # Single-size runs write --out verbatim and are reproducible under the seed. + assert pw["n_senses"] == len(senses) and all(len(pw["anchors"][s]) == g for s in senses) + # The other selector, unscaled contributions and a non-zero seed. + raw = tmp_path / "raw_top1.json" + args = [ + "--group-size", + "2", + "--selector", + "top1_freq", + "--raw-scale", + "--n-boot", + "20", + "--n-null", + "10", + ] + assert sg.main(base[:2] + ["--out", str(raw)] + args + ["--seed", "3"]) == 0 + summary = json.loads(raw.read_text()) + assert _canon(summary, drop=("bundle", "provenance")) == PINS["sg_raw_top1"]["hash"] + assert summary["per_word"]["bank"]["anchors"] == PINS["sg_raw_top1"]["bank_anchors"] + # Single-size runs write --out verbatim and match the sweep's file for that size. single = tmp_path / "single.json" - monkeypatch.setattr( - sys, "argv", argv[:5] + ["--out", str(single), "--group-size", "1", "--n-boot", "20", "--n-null", "10"] + assert ( + sg.main( + base[:2] + + [ + "--out", + str(single), + "--group-size", + "1", + "--n-boot", + "20", + "--n-null", + "10", + ] + ) + == 0 ) - assert sg.main() == 0 a = json.loads(single.read_text()) b = json.loads(out.with_name("synthetic__srp_g1.json").read_text()) assert a["per_word"] == b["per_word"] -def test_classifier_framing_self_test_and_coarse_path(tmp_path): +def test_geometry_controls_route_centering_through_centering_mean(tmp_path): + bundle = make_bundle() + bundle_path = tmp_path / "representations.pt" + torch.save(bundle, bundle_path) + gen = torch.Generator().manual_seed(0) + W = torch.randn(128, D_HIDDEN, generator=gen) + mask_all = torch.ones(128, dtype=torch.bool) + # The controls' row mean is data.centering_mean: live is the full-vocabulary mean ... + live = sg.geometry_row_mean(W, "live", token_mask=None, checkpoint=None, model_id="x", revision=None) + assert torch.equal(live, centering_mean(W, mode="live")) and torch.equal(live, W.mean(dim=0)) + # ... and trained with an all-True token_mask (no checkpoint) is the same vector. + trained = sg.geometry_row_mean(W, "trained", token_mask=mask_all, checkpoint=None, model_id="x", revision=None) + assert torch.equal(trained, live) + partial = mask_all.clone() + partial[:8] = False + assert not torch.equal( + sg.geometry_row_mean( + W, + "trained", + token_mask=partial, + checkpoint=None, + model_id="x", + revision=None, + ), + live, + ) + # geometry_contributions takes that mean rather than recomputing it. + tids = [100, 101] + hidden = { + tid: bundle["hidden"][[i for i, r in enumerate(bundle["metadata"]) if r["token_id"] == tid]] for tid in tids + } + geo = sg.geometry_contributions("knn128", W, live, tids, hidden, 4, 4, 1e-3, 0) + assert set(geo) == set(tids) and geo[100]["contribution"].shape == ( + len(hidden[100]), + 4, + ) + np.testing.assert_allclose(geo[100]["exact"], (hidden[100] @ W[100]).numpy(), rtol=1e-6) + # End to end through --w-u: trained (payload token_mask all True) == live, provenance aside. + outputs = {} + for mode in ("live", "trained"): + torch.save({"W_U_orig": W, "token_mask": mask_all}, tmp_path / f"wu_{mode}.pt") + out = tmp_path / f"knn_{mode}.json" + argv = [ + "--bundle", + str(bundle_path), + "--out", + str(out), + "--basis", + "knn128", + "--neighbor-k", + "4", + ] + argv += [ + "--w-u", + str(tmp_path / f"wu_{mode}.pt"), + "--centering", + mode, + "--n-boot", + "5", + "--n-null", + "3", + ] + assert sg.main(argv) == 0 + outputs[mode] = json.loads(out.read_text()) + assert outputs[mode]["basis"] == "knn128" and outputs[mode]["n_words"] == 2 + assert _canon(outputs["live"], drop=("provenance",)) == _canon(outputs["trained"], drop=("provenance",)) + assert outputs["trained"]["provenance"]["args"]["centering"] == "trained" + + +def test_load_wu_reads_the_resolved_local_snapshot(tmp_path, monkeypatch): + import huggingface_hub + from safetensors.torch import save_file + + snap = tmp_path / "snapshots" / "abc123" + snap.mkdir(parents=True) + head = torch.arange(12, dtype=torch.float32).reshape(4, 3) + save_file( + {"model.embed_tokens.weight": torch.zeros(4, 3)}, + str(snap / "model-00001-of-00002.safetensors"), + ) + save_file( + {"lm_head.weight": head.to(torch.bfloat16)}, + str(snap / "model-00002-of-00002.safetensors"), + ) + seen = {} + + def fake_snapshot_download(model_id, **kwargs): + seen.update(model_id=model_id, **kwargs) + return str(snap) + + monkeypatch.setattr(huggingface_hub, "snapshot_download", fake_snapshot_download) + W = sg.load_wu("org/model", revision="abc123") + assert torch.equal(W, head) and W.dtype == torch.float32 + assert seen == { + "model_id": "org/model", + "revision": "abc123", + "local_files_only": True, + } + # Tied-embedding fallback when no lm_head.weight is stored. + (snap / "model-00002-of-00002.safetensors").unlink() + assert torch.equal(sg.load_wu("org/model"), torch.zeros(4, 3)) and seen["revision"] is None + + +# --------------------------------------------------------------------------- # +# classifier framing +# --------------------------------------------------------------------------- # + + +def test_classifier_framing_reproduces_pinned_outputs(tmp_path): assert cf.self_test() == 0 bundle = make_bundle() - cf.validate_feature_alignment(bundle) gate, gate_summary = cf.score_gate(bundle) - assert gate.all() and gate_summary["pass_fraction"] == 1.0 + assert gate.all() and gate_summary == PINS["cf"]["gate"] metrics = cf.analyze_coarse( - bundle, ks=[2, 4], primary_k=2, primary_encoding="weighted", n_boot=5, n_null_seeds=2, seed=0 + bundle, + ks=[2, 4], + primary_k=2, + primary_encoding="weighted", + n_boot=5, + n_null_seeds=2, + seed=0, ) + assert _canon(metrics) == PINS["cf"]["hash"] + assert metrics["primary"]["methods"] == PINS["cf"]["primary_methods"] + assert metrics["primary"]["comparisons"]["srp_minus_projection_only"] == PINS["cf"]["cmp_projection"] + assert metrics["primary"]["null_seed_sensitivity"]["shuffled_srp"] == PINS["cf"]["null_shuffled"] + assert metrics["results"]["signed"]["4"]["srp"]["aggregate"] == PINS["cf"]["results_signed_4_srp_agg"] assert set(metrics["primary"]["methods"]) == set(cf.METHODS) - srp = metrics["primary"]["methods"]["srp"] - assert srp["n_words"] == 2 and srp["accuracy"] > 0.95 and srp["balanced_accuracy"] > 0.95 - # Projections and contributions differ only by a per-word rescaling of the coordinates, - # which top-|value| selection is not invariant to, so both are reported separately. - assert "projection_only" in metrics["results"]["weighted"]["4"] assert set(metrics["primary"]["comparisons"]) == {f"srp_minus_{m}" for m in cf.METHODS[1:]} - assert set(metrics["primary"]["null_seed_sensitivity"]) == {"shuffled_srp", "random_srp_features"} + assert ( + _array_hash(cf.encode_top_features(bundle, "random_srp_features", 3, "signed", 2)) == PINS["cf_random_enc_hash"] + ) + # CLI: the file carries the analysis block and provenance, on top of the pinned metrics. + bundle_path = tmp_path / "representations.pt" + torch.save(bundle, bundle_path) + out = tmp_path / "cf" / "out.json" + argv = [ + "--bundle", + str(bundle_path), + "--out", + str(out), + "--ks", + "2,4", + "--primary-k", + "4", + ] + argv += [ + "--primary-encoding", + "signed", + "--n-boot", + "7", + "--n-null-seeds", + "3", + "--seed", + "5", + ] + assert cf.main(argv) == 0 + written = json.loads(out.read_text()) + assert _canon(written, drop=("bundle", "provenance")) == PINS["cf_cli"]["hash"] + assert _strip(written["analysis"], {"bundle"}) == PINS["cf_cli"]["analysis"] + assert written["primary"]["methods"]["srp"] == PINS["cf_cli"]["primary_srp"] + assert written["provenance"]["args"]["primary_k"] == 4 and "command" in written["provenance"] + + +def test_shuffled_srp_null_is_the_pre_merge_construction(): + bundle = make_bundle() + projection = bundle["projection"].numpy() + beta = bundle["beta"].numpy() + for seed in (0, 3): + shuffled = wsd.shuffled_srp(projection, beta, bundle["metadata"], seed) + assert _array_hash(shuffled) == PINS["cf_shuffled_hash"][str(seed)] + assert _array_hash(cf.source_matrix(bundle, "shuffled_srp", seed)) == PINS["cf_shuffled_hash"][str(seed)] + assert not np.array_equal(shuffled, bundle["contribution"].numpy()) + matrices = rw.method_matrices(bundle, 0) + assert {name: _array_hash(m) for name, m in matrices.items()} == PINS["rw_method_hashes"] + assert _array_hash(rw.method_matrices(bundle, 3)["shuffled_srp"]) == PINS["rw_shuffled_seed3"] + # One seed form for both scripts: the classifier framing's modular sha1 prefix ... + assert { + "bank_0": wsd.stable_seed("bank", 0), + "seal_3": wsd.stable_seed("seal", 3), + "random_bank_0": wsd.stable_seed("random:bank", 0), + } == PINS["stable_seed"] + # ... which coincides with the run script's former ``seed + int(sha1[:8], 16)`` at seed 0. + for target in ("bank", "seal", "apple"): + prefix = int(hashlib.sha1(target.encode()).hexdigest()[:8], 16) + assert wsd.stable_seed(target, 0) == prefix + assert wsd.stable_seed(target, 3) == (prefix + 3) % 2**32 + + +# --------------------------------------------------------------------------- # +# run script: analysis, statistics, scoring path +# --------------------------------------------------------------------------- # + + +def test_run_analysis_reproduces_pinned_outputs(tmp_path): + bundle = make_bundle() + metrics = rw.analyze_coarsewsd(bundle, tmp_path, 20, 0) + assert _canon(metrics) == PINS["rw_metrics"]["hash"] + assert metrics["methods"]["srp"]["aggregate"] == PINS["rw_metrics"]["srp_aggregate"] + assert metrics["comparisons"]["srp_minus_hidden"] == PINS["rw_metrics"]["cmp_hidden"] + assert metrics["methods"]["hidden"]["per_word"]["bank"] == PINS["rw_metrics"]["bank_hidden_per_word"] + digest = hashlib.sha256((tmp_path / "predictions.jsonl").read_bytes()).hexdigest()[:16] + assert digest == PINS["rw_predictions_hash"] + assert rw.summarize_scoring(bundle) == PINS["rw_scoring_summary"] + assert list(rw.METRIC_NAMES) == [ + "accuracy", + "balanced_accuracy", + "macro_f1", + "pairwise_auc", + "ari", + "nmi", + ] + + +def test_statistics_helpers_reproduce_pinned_values(): + rng = np.random.default_rng(0) + big = rng.normal(size=(250, 8)).astype(np.float32) + labels = rng.integers(0, 3, 250) + big[labels == 1] += 0.5 + big = wsd.l2_normalize(big) + assert rw.sampled_pair_auc(big, labels, 0) == PINS["auc"]["big_0"] # sampled branch (31125 > 20000 pairs) + assert rw.sampled_pair_auc(big, labels, 3) == PINS["auc"]["big_3"] + assert rw.sampled_pair_auc(big[:40], labels[:40], 0) == PINS["auc"]["small_0"] # exhaustive branch + assert _nan_none(rw.sampled_pair_auc(big[:10], np.zeros(10, int), 0)) == PINS["auc"]["one_class"] + rows = [{"g": i % 7, "x": float(i), "y": float((i * 37) % 11)} for i in range(40)] + values = wsd.cluster_bootstrap( + rows, + "g", + lambda s: wsd.safe_spearman([r["x"] for r in s], [r["y"] for r in s]), + 50, + 0, + ) + assert len(values) == PINS["cb"]["n"] and values[:3] == PINS["cb"]["first3"] + assert wsd.percentile_ci(values) == PINS["cb"]["ci"] + assert all(math.isnan(v) for v in wsd.percentile_ci([])) and all( + math.isnan(v) for v in wsd.percentile_ci(np.array([])) + ) + assert wsd.percentile_ci(np.asarray(values)) == wsd.percentile_ci(values) + assert _nan_none(wsd.safe_spearman([1, 1, 1], [1, 2, 3])) == PINS["spearman"]["const"] + assert _nan_none(wsd.safe_spearman([1, 2], [1, 2])) == PINS["spearman"]["short"] + assert wsd.safe_spearman([1, 2, 3, 4], [1, 3, 2, 4]) == PINS["spearman"]["ok"] + # The documented difference from utils.spearman: a range under 1e-8 is degenerate here, correlated there. + tiny = [0.0, 1e-9, 2e-9] + assert _nan_none(wsd.safe_spearman(tiny, [1, 2, 3])) == PINS["spearman"]["tiny"] is None + assert spearman(tiny, [1, 2, 3]) == pytest.approx(1.0) + gen = np.random.default_rng(1) + vals = gen.normal(size=9) + assert ( + wsd.bootstrap_mean(vals, np.random.default_rng(4), 3) + == [float(np.random.default_rng(4).choice(vals, size=9, replace=True).mean()) for _ in range(1)] + + wsd.bootstrap_mean(vals, np.random.default_rng(4), 3)[1:] + ) + + +def test_prepare_target_codes_matches_the_pre_refactor_path(tmp_path): + tok, lm, w_u, row_mean, w_dec_unit, info, audit, k = _scoring_setup(tmp_path) + assert k == PINS["ptc"]["k"] == 6 + assert audit == PINS["ptc"]["audit"] + for target in ("bank", "seal"): + assert info[target]["feature_ids"].tolist() == PINS["ptc"][f"{target}_fids"] + assert _array_hash(info[target]["beta"].numpy()) == PINS["ptc"][f"{target}_beta"] + # Runner-layout checkpoints (top-level ``factorizer``, no ``config``) yield the same k via the mmap read. + (tmp_path / "runner").mkdir() + _tok, _lm, _w, _rm, _wd, info_runner, audit_runner, k_runner = _scoring_setup( + tmp_path / "runner", runner_layout=True + ) + assert k_runner == 6 and audit_runner == audit + assert torch.equal(info_runner["bank"]["beta"], info["bank"]["beta"]) + assert rw.checkpoint_k(tmp_path / "checkpoint.pt", {}) == 6 + + +def test_score_items_matches_the_pre_refactor_path(tmp_path): + tok, lm, w_u, row_mean, w_dec_unit, info, _audit, _k = _scoring_setup(tmp_path) + rw.configure_tokenizer(tok) + assert (tok.padding_side, tok.truncation_side) == ("left", "left") + bundle = rw.score_items(_items(PROMPTS), lm, tok, lm.lm_head, row_mean, w_dec_unit, info, 2, 16, "cpu") + assert len(bundle["metadata"]) == PINS["score"]["n"] == 4 # the multi-token target is dropped + assert _canon(bundle["metadata"]) == PINS["score"]["metadata"] + for key in wsd.BUNDLE_TENSOR_KEYS: + assert _array_hash(bundle[key].numpy()) == PINS["score"][key], key + assert bundle["truncation"] == { + "max_length": 16, + "truncation_side": "left", + "n_truncated_prompts": 0, + "max_prompt_tokens": max(len(p.split()) for p in PROMPTS[:4]), + } + summary = rw.summarize_scoring(bundle) + assert summary.pop("truncation") == bundle["truncation"] + assert summary == PINS["score_summary"] + bundle["run"] = {"dataset": "coarsewsd20"} + wsd.validate_bundle(bundle) + + +def test_left_truncation_keeps_the_cloze_cue(tmp_path): + tok, lm, w_u, row_mean, w_dec_unit, info, _audit, _k = _scoring_setup(tmp_path) + rw.configure_tokenizer(tok) + long_prompt = "x y x y x y x y the river bank is the missing word :" + tail = " ".join(long_prompt.split()[-6:]) + head = " ".join(long_prompt.split()[:6]) + + def hidden_of(prompt, max_length): + return rw.score_items( + _items([prompt] * 5)[:1], + lm, + tok, + lm.lm_head, + row_mean, + w_dec_unit, + info, + 1, + max_length, + "cpu", + ) + + truncated = hidden_of(long_prompt, 6) + assert truncated["truncation"]["n_truncated_prompts"] == 1 and truncated["truncation"]["max_prompt_tokens"] == len( + long_prompt.split() + ) + assert torch.equal(truncated["hidden"], hidden_of(tail, 64)["hidden"]) # the cue survives ... + assert not torch.equal(truncated["hidden"], hidden_of(head, 64)["hidden"]) # ... unlike the old right truncation + tok.truncation_side = "right" + assert torch.equal(hidden_of(long_prompt, 6)["hidden"], hidden_of(head, 64)["hidden"]) + + +def test_centering_trained_equals_live_when_the_mask_keeps_all_rows(): + lm = FakeLM(len(VOCAB), D_HIDDEN, seed=0) + w_u = lm.lm_head.weight.detach().float() + live = rw.scoring_row_mean(w_u, FakeTokenizer(VOCAB), "live", None) + assert torch.equal(live, w_u.mean(dim=0)) + assert torch.equal( + rw.scoring_row_mean(w_u, FakeTokenizer(VOCAB, special_ids=()), "trained", None), + live, + ) + masked = rw.scoring_row_mean(w_u, FakeTokenizer(VOCAB), "trained", None) # rows 0 and 1 are special ids + assert torch.equal(masked, w_u[2:].mean(dim=0)) and not torch.equal(masked, live) + stored = torch.full((D_HIDDEN,), 0.25) + assert torch.equal(rw.scoring_row_mean(w_u, FakeTokenizer(VOCAB), "trained", stored), stored) + assert torch.equal(rw.scoring_row_mean(w_u, FakeTokenizer(VOCAB), "live", stored), live) + + +def test_analyze_only_writes_analysis_config_and_keeps_run_config(tmp_path): + out_dir = tmp_path / "run" + out_dir.mkdir() + torch.save(make_bundle(), out_dir / "representations.pt") + sentinel = '{"scoring_run": true}\n' + (out_dir / "run_config.json").write_text(sentinel) + assert ( + rw.main( + [ + "--analyze-only", + "--out-dir", + str(out_dir), + "--n-boot", + "5", + "--seed", + "0", + ] + ) + == 0 + ) + assert (out_dir / "run_config.json").read_text() == sentinel + analysis = json.loads((out_dir / "analysis_config.json").read_text()) + assert analysis["provenance"]["args"]["analyze_only"] is True and "command" in analysis["provenance"] + assert analysis["coarsewsd_template"] == rw.COARSEWSD_TEMPLATE + metrics = json.loads((out_dir / "metrics.json").read_text()) + assert metrics["provenance"] == analysis["provenance"] + assert metrics["methods"]["srp"]["aggregate"]["accuracy"] == 1.0 and metrics["run"]["dataset"] == "coarsewsd20" + assert (out_dir / "predictions.jsonl").exists() + with pytest.raises(SystemExit): + rw.main(["--analyze-only"]) + + +# --------------------------------------------------------------------------- # +# shared bundle helpers and removed paths +# --------------------------------------------------------------------------- # + + +def test_bundle_loading_and_validation(tmp_path): + bundle = make_bundle() + path = tmp_path / "representations.pt" + torch.save(bundle, path) + loaded = wsd.load_bundle(path) + assert loaded["metadata"] == bundle["metadata"] and torch.equal(loaded["contribution"], bundle["contribution"]) + broken = dict(bundle) + del broken["beta"] + with pytest.raises(KeyError): + wsd.validate_bundle(broken) + broken = dict(bundle, hidden=bundle["hidden"][:-1]) + with pytest.raises(ValueError, match="rows"): + wsd.validate_bundle(broken) + with pytest.raises(ValueError, match="dataset"): + wsd.validate_bundle(dict(bundle, run={"dataset": "ambistory"})) + drifted = bundle["feature_ids"].clone() + drifted[1, 0] = 99 + with pytest.raises(ValueError, match="feature coordinates"): + wsd.validate_bundle(dict(bundle, feature_ids=drifted)) + + +def test_word_splits_apply_the_shared_rule(): + bundle = make_bundle() + splits = wsd.word_splits(bundle["metadata"]) + assert [s.word for s in splits] == ["bank", "seal"] + bank = splits[0] + assert bank.senses == ["0", "1"] and len(bank.train_idx) == 40 and len(bank.test_idx) == 20 + assert set(bank.train_y) == {"0", "1"} and bank.train_idx.dtype == np.int64 + metadata = [ + {"target": "one", "split": "train", "sense": "0"}, + {"target": "one", "split": "test", "sense": "0"}, + {"target": "two", "split": "train", "sense": "0"}, + {"target": "two", "split": "train", "sense": "1"}, + {"target": "two", "split": "test", "sense": "9"}, + {"target": "three", "split": "train", "sense": "0"}, + {"target": "three", "split": "train", "sense": "1"}, + {"target": "three", "split": "test", "sense": "1"}, + ] + assert [s.word for s in wsd.word_splits(metadata)] == ["three"] # one sense; unseen test sense + keep = np.array([True] * 7 + [False]) + assert wsd.word_splits(metadata, keep=keep) == [] # the gate emptied three's test split + + +def test_removed_paths_are_gone(): + with pytest.raises(SystemExit): + rw.parse_args(["--dataset", "ambistory"]) + args = rw.parse_args([]) + assert args.dataset == "coarsewsd20" and args.centering == "live" and args.splits == "train,test" + assert not hasattr(args, "anchors") and not hasattr(args, "position_mode") + for name in ( + "load_ambistory", + "analyze_ambistory", + "make_story", + "mask_exact_target", + ): + assert not hasattr(rw, name), name + for name in ("ambistory_rows", "analyze_ambistory", "summarize_ambistory_rows"): + assert not hasattr(cf, name), name + assert rw.self_test() == 0 diff --git a/uv.lock b/uv.lock index aac8d1e..122c48b 100644 --- a/uv.lock +++ b/uv.lock @@ -1228,7 +1228,7 @@ wheels = [ [[package]] name = "sparse-readout-prism" -version = "0.2.0" +version = "0.2.1" source = { editable = "." } dependencies = [ { name = "accelerate" },