diff --git a/CLAUDE.md b/CLAUDE.md index 9be1e53..9801745 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -1,6 +1,6 @@ # CLAUDE.md -Guidance for Claude Code (claude.ai/code) when working in this repository. +This file provides guidance to Claude Code (claude.ai/code) when working with code in this repository. ## Project overview @@ -23,8 +23,8 @@ Token Search"). Pivotal tokens are one of three scales it searches. ## The invariant that matters most -**A latent event's `score` is a J-lens readout probability. It is NOT a -probability delta.** +**A latent event's `score` is a J-lens readout score (a probability, or a cosine +for `jlens_cosine`). It is NOT a probability delta.** - Latent events have `prob_delta`, `prob_before`, `prob_after`, and `is_positive` set to `None`, deliberately. Do not fill them in. @@ -41,6 +41,13 @@ probability delta.** A single 0.5 floor across both scales keeps the filler meta-token `" the"` (readout 0.92) and throws away a pivotal token worth +0.45. This has already been shipped once by accident; do not reintroduce it. + `tests/test_regressions.py` pins this and the other bugs listed under "Things + that will bite you". +- The same rule holds one level down: `jlens_cosine` scores are cosines, while + `jlens`/`logit_lens` scores are probabilities. `min_readout_score`, + `most_surfaced()` and exporter `min_score` raise on a mix, and `summary()` + reports `readout_score_by_scale`. Scales are defined by `READOUT_SCORE_SCALE` + in `pts/events.py`; register any new readout method there. - Never count latent events as positive or negative — they have no valence. - `logit_lens` readouts are weaker evidence than `jlens` ones; keep `readout_method` visible so they can be filtered apart. @@ -57,7 +64,8 @@ which is the main way this project can mislead people. Tests in pip install -e . # core + model deps pip install -e '.[all]' # + sentence-transformers, math-verify, pytest -pytest tests/ -q # 66 tests, ~15s, uses a tiny random model +pytest tests/ -q # 94 tests, ~20s, uses a tiny random model +pytest tests/test_jlens.py::test_causal_mask_holds -q # a single test pts run --granularity token|sentence|latent|all --model M --output-path events.jsonl pts fit-jlens --model M --output-path ./jlens/m # calibrate the Jacobian lens @@ -84,6 +92,8 @@ pts push --input-path X --hf-repo-id user/repo | `pts/searchers/base.py` | Model loading, prompt formatting, the probability cache. | | `pts/searchers/{token,sentence,latent,reasoning}.py` | The four searchers. | | `pts/oracle.py`, `pts/dataset.py` | Success evaluation and dataset loading (largely unchanged). | +| `pts/verification.py` | Arithmetic CoT verification used by Sentence PTS. Imports torch at module level. | +| `pts/cli.py` | Every `pts` subcommand and its flags. | | `pts/exporters.py` | All output formats + dataset cards. | | `pts/core.py`, `pts/storage.py`, `pts/thought_anchors.py` | legacy compatibility shims. | @@ -91,6 +101,8 @@ pts push --input-path X --hf-repo-id user/repo classification, and linking layers are pure Python; model-touching code is imported lazily via `__getattr__`. Keep it that way. +`research/` and `visualizer/` are standalone and excluded from the package. + ## The J-lens `J_l = E[∂h_final,t' / ∂h_l,t]`, averaged over source positions `t`, all later @@ -108,9 +120,18 @@ row — the paper's `O(n × d_model)` cost. `torch.autograd.functional.jacobian`, and `test_causal_mask_holds` checks the assumption directly. Do not weaken those tests. -No reference code was released with the workspace paper — this is written from -the equations and has **not** been validated against the authors' results. Say so -when writing docs. "Meta-token" is our term, not the paper's. +`fit` was written from the equations before Anthropic released reference code +(`anthropics/jacobian-lens`). Its estimator differs: the reference skips the first +16 positions and the last one, and sums over `t' >= t` where we average. Our +fitted lenses have **not** been validated against theirs. Say so when writing +docs. "Meta-token" is our term, not the paper's. + +`JLens.load` also reads reference `.pt` lenses (`{"J": {layer: [d,d]}, ...}`), +locally or as `hf:////`. Neuronpedia hosts ~40 at +`neuronpedia/jacobian-lens`. They use our convention (`J @ h` on decoder-block +outputs, target = final block), and `lm_head(final_norm(J h))` is the reference +readout. Loading checks width and depth against the model. `jlens_cosine` is +WorkspaceBench's readout, not the paper's. ## Things that will bite you diff --git a/README.md b/README.md index 33a010d..fedcaad 100644 --- a/README.md +++ b/README.md @@ -89,7 +89,9 @@ pts run --granularity all --model Qwen/Qwen3-0.6B \ --readout-method jlens --jlens-path ./jlens --output-path events.jsonl ``` -No J-lens yet? `--readout-method logit_lens` needs no calibration. It is the same +No J-lens yet? `--jlens-path` also takes a reference lens straight from the Hub, +e.g. `hf://neuronpedia/jacobian-lens/qwen3-1.7b/jlens/Salesforce-wikitext/Qwen3-1.7B_jacobian_lens.pt`. +Or use `--readout-method logit_lens`, which needs no calibration. It is the same readout with `J = I`, and it is a weaker signal. See [docs/latent_pts.md](docs/latent_pts.md). @@ -162,8 +164,6 @@ and compatible with that work, not the same as it. ## Datasets - [codelion/Qwen3-0.6B-pts](https://huggingface.co/datasets/codelion/Qwen3-0.6B-pts) -- [codelion/Qwen3-0.6B-pts-thought-anchors](https://huggingface.co/datasets/codelion/Qwen3-0.6B-pts-thought-anchors) -- [codelion/Qwen3-0.6B-pts-steering-vectors](https://huggingface.co/datasets/codelion/Qwen3-0.6B-pts-steering-vectors) - [codelion/DeepSeek-R1-Distill-Qwen-1.5B-pts](https://huggingface.co/datasets/codelion/DeepSeek-R1-Distill-Qwen-1.5B-pts) ## References diff --git a/docs/dataset_schema.md b/docs/dataset_schema.md index 1d1de22..f951fa9 100644 --- a/docs/dataset_schema.md +++ b/docs/dataset_schema.md @@ -96,7 +96,7 @@ a hypothetical. It is what the first implementation did. |---|---|---| | `search_method` | string | `token_pts` \| `sentence_pts` \| `latent_pts` \| `latent_pts_enrichment` | | `intervention_type` | string? | `append_token` \| `replace_sentence` \| `remove_sentence` | -| `readout_method` | string? | Latent only: `jlens` \| `logit_lens`. **`logit_lens` is weaker evidence**, so filter on this. | +| `readout_method` | string? | Latent only: `jlens` \| `jlens_cosine` \| `logit_lens`. **`logit_lens` is weaker evidence**, so filter on this. `jlens_cosine` scores a cosine, not a probability, so never compare its `score` with the others. | ### Classification and links diff --git a/docs/latent_pts.md b/docs/latent_pts.md index b5c6df2..5ac7199 100644 --- a/docs/latent_pts.md +++ b/docs/latent_pts.md @@ -57,6 +57,14 @@ Jacobian for all source positions at the same time. That is the paper's stated cost of `O(n x d_model)` backward passes. Dividing row `t` by the count of `t' >= t` turns the sum into the mean the definition asks for. +Anthropic's reference code, [`anthropics/jacobian-lens`](https://github.com/anthropics/jacobian-lens), +came out after `fit` was written, and its estimator differs in two ways: it +skips the first 16 positions (attention sinks) and the last one, and it keeps +the sum over `t' >= t` instead of dividing it into a mean. So a lens from +`pts fit-jlens` is not the same matrix as a reference lens, and ours has not +been validated against theirs. If one exists for your model, prefer the +reference lens (below). + `tests/test_jlens.py` checks this against a brute-force `torch.autograd.functional.jacobian`. The two agree to about 1e-8, and a separate test checks the causal-mask assumption directly. @@ -70,6 +78,17 @@ It is also a weaker signal. A workspace claim is a claim about future influence, and the logit lens does not measure future influence. Events read this way are tagged `readout_method: "logit_lens"` so you can filter them out. +### The cosine readout + +[WorkspaceBench](https://github.com/camilablank/workspace-bench) reads its J-lens +arm differently: it ranks tokens by the cosine between `h_l` and each token's +J-lens vector `J_l^T w_t`, with no final norm and no softmax. That stops tokens +with large J-lens vectors from dominating every readout. +`--readout-method jlens_cosine` does the same. Its `score` is a cosine in +[-1, 1], not a probability, so `EventStorage` refuses to rank or threshold it +together with `jlens` or `logit_lens` events. Filter to one `readout_method` +first. + ## Usage ### 1. Fit a J-lens (once per model) @@ -88,6 +107,24 @@ the lens is a property of the model rather than the task, so the source dataset works fine with no labels. Pass `--calibration-file` for your own text, one sequence per line. Lower `--basis-chunk` if you run out of memory. +### 1b. Or use a reference lens + +Neuronpedia hosts reference-code lenses for about 40 models (GPT-2, Gemma 2/3/4, +Llama 3.1/3.3, Qwen3/3.5/3.6, OLMo 3 and others) at +[`neuronpedia/jacobian-lens`](https://huggingface.co/neuronpedia/jacobian-lens). +`--jlens-path` takes one directly as `hf:////`, or a local +`.pt` file in the same format: + +```bash +pts enrich --input-path events.jsonl --output-path events_latent.jsonl \ + --model Qwen/Qwen3-1.7B --with-latent --readout-method jlens \ + --jlens-path hf://neuronpedia/jacobian-lens/qwen3-1.7b/jlens/Salesforce-wikitext/Qwen3-1.7B_jacobian_lens.pt +``` + +These lenses cover every layer below the last, so any workspace layer works. A +lens whose width or depth does not fit the model is refused. When the +`config.yaml` next to it names a different model, you get a warning. + ### 2. Enrich a dataset you already have This is the main path. It reuses curated token and sentence datasets instead of @@ -113,52 +150,16 @@ pts run --granularity all --model Qwen/Qwen3-0.6B \ --output-path events.jsonl ``` -## What to keep in mind - -Each of these is a way the output can mislead you if you forget it. - -**A latent score is a readout probability, not a Δ-probability.** It says how -strongly the lens surfaces a token, not how much that token changed the answer. -It is not comparable to the `prob_delta` on token and sentence events. -`prob_delta` and `is_positive` stay `null` on latent events on purpose. - -**Latent events are observational.** Enrichment reports what the lens sees. It -does not show that a meta-token caused an emitted event. That would take an -intervention: steer or ablate the direction and re-measure success. - -**Readouts are noisy.** Neither lens is guaranteed to be faithful to what the -model actually represents. A high-scoring `verify` might be an artifact of the -unembedding geometry rather than a real concept. - -**This is an independent reimplementation.** No code was released with the -workspace paper. The math here comes from the published equations and is checked -for internal correctness, but it has not been validated against the authors' -results, so do not report PTS numbers as reproducing theirs. - -**Links are heuristics.** `linked_event_ids` come from a weighted score (query -match, context overlap, category agreement, timing), not a verified causal path. -Run `--shuffle-control`: if the observed link scores do not clearly beat the -shuffled baseline, the structure is not above chance. - -**Records are model-specific.** A pivotal token in one model tells you nothing -about another model's workspace. Enriching with a different model is refused -unless you pass `--allow-model-mismatch`, and even then both model ids are stored. - -## What would make it convincing - -The claim under test is that emitted pivotal tokens and thought-anchor sentences -are often preceded by latent meta-tokens in the workspace. In rough order of -strength, the evidence that would back it up: - -1. **Intervention.** Steer or ablate a meta-token direction and show success - probability moves. -2. **Lead time with a control.** Show a category-matching meta-token appears `k` - tokens before the emitted event more often than chance, using - `--shuffle-control` as the baseline. -3. **Consistent chains.** Show `latent: verification -> token: " Wait" -> - sentence: "Let me check..."` holds across many queries, not just anecdotes. -4. **J-lens beats logit-lens.** If the effect is just as strong with - `--readout-method logit_lens`, it is not about future influence, and so not - about a workspace. - -Until at least (2) holds with a control, treat the output as exploratory. +## Notes + +- A latent event's `score` is a readout score, not a Δ-probability: a + probability for `jlens` and `logit_lens`, a cosine for `jlens_cosine`. It is + not comparable to the `prob_delta` on token and sentence events, which is why + `prob_delta` and `is_positive` are `null` on latent events. +- Latent events are observational. The lens shows what a meta-token leans toward, + not that it caused anything downstream. A causal claim would need an + intervention (steer or ablate, then re-measure). +- Run `--shuffle-control`. If the link scores do not beat the shuffled baseline, + the structure is not above chance. +- Records are model-specific. Enriching with a different model needs + `--allow-model-mismatch`. diff --git a/pts/__init__.py b/pts/__init__.py index 8a905ea..28d4b73 100644 --- a/pts/__init__.py +++ b/pts/__init__.py @@ -75,6 +75,7 @@ "TokenExporter": ("pts.exporters", "TokenExporter"), "EventExporter": ("pts.exporters", "EventExporter"), "JLens": ("pts.latent.jlens", "JLens"), + "JLensCosine": ("pts.latent.jlens", "JLensCosine"), "Readout": ("pts.latent.jlens", "Readout"), } diff --git a/pts/cli.py b/pts/cli.py index b154f6c..36668df 100644 --- a/pts/cli.py +++ b/pts/cli.py @@ -538,13 +538,21 @@ def _add_latent_args(p) -> None: p.add_argument( "--readout-method", default="logit_lens", - choices=["jlens", "logit_lens"], + choices=["jlens", "jlens_cosine", "logit_lens"], help="How to read meta-tokens out of the workspace. 'jlens' needs " - "--jlens-path (fit one with `pts fit-jlens`). 'logit_lens' needs no " + "--jlens-path. 'jlens_cosine' ranks by cosine to each token's J-lens " + "vector, as WorkspaceBench does; its score is a cosine, not a " + "probability. 'logit_lens' needs no " "calibration but is weaker evidence: it reads what an activation " "would say now, not what it pushes the model to say later.", ) - p.add_argument("--jlens-path", default=None, help="Directory holding fitted J-lens matrices") + p.add_argument( + "--jlens-path", + default=None, + help="A `pts fit-jlens` directory, a reference jacobian-lens .pt file, or " + "hf://// (e.g. a Neuronpedia lens from " + "hf://neuronpedia/jacobian-lens/...)", + ) p.add_argument( "--workspace-layers", nargs="+", @@ -565,7 +573,9 @@ def _add_latent_args(p) -> None: # Not --top-k: `pts run` already uses that for sampling. p.add_argument("--readout-top-k", type=int, default=25, help="How many vocabulary tokens to read out per position") - p.add_argument("--min-score", type=float, default=0.01, help="Minimum readout score to keep") + p.add_argument("--min-score", type=float, default=0.01, + help="Minimum readout score to keep (a probability, or a cosine " + "for jlens_cosine)") p.add_argument("--keep-per-position", type=int, default=3, help="Max meta-token events kept per position/layer") p.add_argument("--link-threshold", type=float, default=0.5, help="Minimum link score") diff --git a/pts/event_storage.py b/pts/event_storage.py index 94fd39f..a0a08b9 100644 --- a/pts/event_storage.py +++ b/pts/event_storage.py @@ -17,6 +17,7 @@ EVENT_SENTENCE, EVENT_TOKEN, from_any_record, + readout_score_scale, ) logger = logging.getLogger(__name__) @@ -175,7 +176,13 @@ def filter( So the two scales get their own thresholds: ``min_prob_delta`` for emitted events, ``min_readout_score`` for latent ones. Each applies only to the scale it belongs to and leaves the other untouched. + + Latent scores are not all on one scale either: ``jlens_cosine`` scores a + cosine, the others a probability. ``min_readout_score`` raises if the + latent events here mix the two; filter by ``readout_method`` first. """ + if min_readout_score is not None: + self._single_readout_scale("min_readout_score") result = EventStorage() for evt in self.events: @@ -235,8 +242,25 @@ def most_surfaced(self, n: int = 10) -> List[CausalReasoningEvent]: the token matters causally. """ latent = self.by_event_type(EVENT_LATENT) + self._single_readout_scale("most_surfaced()") return sorted(latent, key=lambda e: e.score, reverse=True)[:n] + def _latent_scales(self) -> Dict[str, List[CausalReasoningEvent]]: + scales: Dict[str, List[CausalReasoningEvent]] = {} + for e in self.by_event_type(EVENT_LATENT): + scales.setdefault(readout_score_scale(e.readout_method), []).append(e) + return scales + + def _single_readout_scale(self, operation: str) -> None: + scales = self._latent_scales() + if len(scales) > 1: + raise ValueError( + f"{operation} would compare latent scores on different scales " + f"({', '.join(sorted(scales))}): a jlens_cosine score is a cosine, " + "not a probability. Filter to one readout_method first, e.g. " + "filter(criteria={'readout_method': 'jlens'})." + ) + def queries(self) -> List[str]: seen = [] for e in self.events: @@ -262,6 +286,7 @@ def summary(self) -> Dict[str, Any]: emitted = [e for e in self.events if e.prob_delta is not None] latent = self.by_event_type(EVENT_LATENT) + scales = self._latent_scales() summary = { "total_events": len(self.events), @@ -281,11 +306,23 @@ def summary(self) -> Dict[str, Any]: "max_abs_prob_delta": ( max(abs(e.prob_delta) for e in emitted) if emitted else None ), - # Latent only: these are readout probabilities, on a different scale. + # Latent only, and only when every readout is on one scale: these + # are readout scores, not probability deltas. "average_readout_score": ( - sum(e.score for e in latent) / len(latent) if latent else None + sum(e.score for e in latent) / len(latent) + if latent and len(scales) == 1 else None + ), + "max_readout_score": ( + max(e.score for e in latent) if latent and len(scales) == 1 else None ), - "max_readout_score": max((e.score for e in latent), default=None), + "readout_score_by_scale": { + scale: { + "count": len(evts), + "average": sum(e.score for e in evts) / len(evts), + "max": max(e.score for e in evts), + } + for scale, evts in scales.items() + }, } # the legacy format's get_anchor_summary reads these names; they mean emitted-only. summary["average_score"] = summary["average_abs_prob_delta"] diff --git a/pts/events.py b/pts/events.py index eb05da3..046068d 100644 --- a/pts/events.py +++ b/pts/events.py @@ -54,6 +54,29 @@ EVENT_SENTENCE: VISIBILITY_EMITTED, } +# Readout methods for latent events, and the scale each one's ``score`` is on. +# ``jlens`` and ``logit_lens`` score a softmax probability; ``jlens_cosine`` +# scores a cosine in [-1, 1]. Scores on different scales must never be ranked +# or thresholded together. +READOUT_JLENS = "jlens" +READOUT_LOGIT_LENS = "logit_lens" +READOUT_JLENS_COSINE = "jlens_cosine" + +SCORE_SCALE_PROBABILITY = "probability" +SCORE_SCALE_COSINE = "cosine" + +READOUT_SCORE_SCALE = { + READOUT_JLENS: SCORE_SCALE_PROBABILITY, + READOUT_LOGIT_LENS: SCORE_SCALE_PROBABILITY, + READOUT_JLENS_COSINE: SCORE_SCALE_COSINE, +} + + +def readout_score_scale(readout_method: Optional[str]) -> str: + """The scale a latent event's ``score`` is on. Records that predate the + field were all softmax readouts, so a missing method means probability.""" + return READOUT_SCORE_SCALE.get(readout_method or READOUT_JLENS, SCORE_SCALE_PROBABILITY) + def _now() -> str: return time.strftime("%Y-%m-%dT%H:%M:%S") diff --git a/pts/exporters.py b/pts/exporters.py index bfaae67..1684ae6 100644 --- a/pts/exporters.py +++ b/pts/exporters.py @@ -132,20 +132,27 @@ def export_causal_events( different scales, so applying one threshold to both silently deletes every latent event (readout scores are routinely well below a 0.1 prob-delta floor) while looking like a principled filter. + + ``min_score`` refuses latent events on mixed score scales (see + ``EventStorage.filter``), and 0 means no floor. """ + if min_score: + self.storage._single_readout_scale("--min-score") events = [] for e in self.storage: if e.event_type == EVENT_LATENT: - if e.score >= min_score: + if not min_score or e.score >= min_score: events.append(e) elif e.prob_delta is None or abs(e.prob_delta) >= min_prob_delta: events.append(e) self._write(output_path, [e.to_dict() for e in events]) def export_metatokens(self, output_path: str, min_score: float = 0.0) -> None: + if min_score: + self.storage._single_readout_scale("--min-score") events = [ e for e in self.storage - if e.event_type == EVENT_LATENT and e.score >= min_score + if e.event_type == EVENT_LATENT and (not min_score or e.score >= min_score) ] if not events: logger.warning( diff --git a/pts/latent/__init__.py b/pts/latent/__init__.py index e4e8d98..3304a5f 100644 --- a/pts/latent/__init__.py +++ b/pts/latent/__init__.py @@ -25,7 +25,7 @@ get_layer_modules, resolve_workspace_layers, ) -from .jlens import JLens, LogitLens, Readout, ReadoutResult, load_readout +from .jlens import JLens, JLensCosine, LogitLens, Readout, ReadoutResult, load_readout from .metatokens import MetaTokenExtractor, enrich_events_with_latent __all__ = [ @@ -34,6 +34,7 @@ "get_layer_modules", "resolve_workspace_layers", "JLens", + "JLensCosine", "LogitLens", "Readout", "ReadoutResult", diff --git a/pts/latent/jlens.py b/pts/latent/jlens.py index 3d42c45..4f9bdf2 100644 --- a/pts/latent/jlens.py +++ b/pts/latent/jlens.py @@ -17,15 +17,28 @@ token, each the average causal influence of that direction on eventually producing that token. -No reference implementation was released with the paper, so this is written -from the equations. It has not been validated against the authors' results -- -treat the readouts as hypotheses. See ``docs/latent_pts.md``. +``fit`` was written from the paper's equations before Anthropic released +reference code (``anthropics/jacobian-lens``), and its estimator is not +identical to theirs: the reference skips the first 16 positions and the last +one, and sums over later targets ``t' >= t`` where we average over them. Our +fitted matrices have not been validated against theirs -- treat readouts from +them as hypotheses. ``JLens.load`` also reads the reference ``.pt`` format, +including the fitted lenses Neuronpedia hosts on the Hugging Face Hub (see +``load`` below), so you can skip fitting and use theirs. See +``docs/latent_pts.md``. Note that the **logit lens is exactly this construction with J = I**: it asks what the activation would say if emitted right now, rather than what it is pushing the model to say later. That makes it a principled zero-cost baseline rather than a hack, and it is what ``--readout-method logit_lens`` uses. +``--readout-method jlens_cosine`` ranks tokens the way WorkspaceBench's J-lens +arm does: by the cosine between ``h`` and each token's J-lens vector +``J_l^T w_t``, with no final norm. That removes the advantage tokens get from +having large J-lens vectors. Its ``score`` is a cosine in [-1, 1], **not a +probability**, so it must never be thresholded or ranked together with +``jlens`` / ``logit_lens`` scores; ``EventStorage`` refuses to. + Computing J efficiently ----------------------- A naive Jacobian would need one backward pass per (source position, output @@ -45,11 +58,13 @@ import json import logging import os +import re from dataclasses import dataclass -from typing import Any, Dict, List, Optional, Sequence +from typing import Any, Dict, List, Optional, Sequence, Tuple import torch +from ..events import READOUT_JLENS, READOUT_JLENS_COSINE, READOUT_LOGIT_LENS from .activations import ( ResidualCapture, get_final_norm, @@ -60,9 +75,9 @@ logger = logging.getLogger(__name__) -READOUT_JLENS = "jlens" -READOUT_LOGIT_LENS = "logit_lens" -READOUT_METHODS = (READOUT_JLENS, READOUT_LOGIT_LENS) +READOUT_METHODS = (READOUT_JLENS, READOUT_JLENS_COSINE, READOUT_LOGIT_LENS) + +HUB_PREFIX = "hf://" @dataclass @@ -110,6 +125,13 @@ def transform(self, h: torch.Tensor, layer: int) -> torch.Tensor: """Map a layer-l activation into the final-residual-stream basis.""" raise NotImplementedError + def scores(self, h: torch.Tensor, layer: int) -> torch.Tensor: + """One score per vocabulary token. By default a probability distribution.""" + projected = self.transform(h, layer) + if self.final_norm is not None: + projected = self.final_norm(projected) + return torch.softmax((self.W_U @ projected).float(), dim=-1) + def read( self, activations: torch.Tensor, @@ -125,16 +147,9 @@ def read( with torch.no_grad(): h = activations.to(self.W_U.dtype).to(self.W_U.device) - projected = self.transform(h, layer) - - if self.final_norm is not None: - projected = self.final_norm(projected) - - logits = self.W_U @ projected - probs = torch.softmax(logits.float(), dim=-1) - - k = min(top_k, probs.shape[-1]) - top = torch.topk(probs, k=k) + scores = self.scores(h, layer) + k = min(top_k, scores.shape[-1]) + top = torch.topk(scores, k=k) results = [] for rank, (score, token_id) in enumerate(zip(top.values.tolist(), top.indices.tolist()), 1): @@ -383,37 +398,181 @@ def save(self, path: str) -> None: @classmethod def load(cls, path: str, model: Any, tokenizer: Any) -> "JLens": - import numpy as np + """Load J-lens matrices from any of: - npz_path = os.path.join(path, "jlens.npz") - if not os.path.exists(npz_path): - raise FileNotFoundError( - f"No J-lens at {path} (expected {npz_path}). " - f"Fit one with: pts fit-jlens --model --output-path {path}" - ) + - a directory written by ``save`` (``jlens.npz`` + ``config.json``); + - a ``.pt`` file in the reference ``JacobianLens`` layout + (``{"J": {layer: [d, d]}, "n_prompts", "source_layers", "d_model"}``); + - ``hf:////``, the same file on the Hub, e.g. + ``hf://neuronpedia/jacobian-lens/gpt2-small/jlens/Salesforce-wikitext/gpt2_jacobian_lens.pt``. - data = np.load(npz_path) - matrices = { - int(key.split("_")[1]): torch.from_numpy(data[key]) for key in data.files - } + Reference lenses map a decoder block's output at layer ``l`` into the + final block's output basis, as ``J_l @ h`` -- the same convention as + ``fit``, so they drop in unchanged. + """ + if path.startswith(HUB_PREFIX): + path = _download_from_hub(path) + if os.path.isfile(path): + matrices, config = _load_reference_pt(path) + else: + matrices, config = _load_npz_dir(path) + + _check_lens_fits_model(matrices, config, model, path) + logger.info(f"Loaded J-lens for layers {sorted(matrices)} from {path}") + return cls(model, tokenizer, matrices=matrices, config=config) + + +class JLensCosine(JLens): + """J-lens ranked by cosine to each token's J-lens vector, as WorkspaceBench does. + + ``score_t = / (||J_l^T w_t|| ||h||)``. There is no final norm + and no softmax, so ``score`` is a cosine in [-1, 1], not a probability; see + the module docstring for why that matters downstream. Dividing by ``||h||`` + does not change the ranking WorkspaceBench uses; it only makes scores + comparable across positions. + """ + + method = READOUT_JLENS_COSINE - config = {} - config_path = os.path.join(path, "config.json") - if os.path.exists(config_path): - with open(config_path) as f: - config = json.load(f) - - fitted_for = config.get("model_id") - current = getattr(model.config, "_name_or_path", None) - if fitted_for and current and fitted_for != current: - logger.warning( - f"This J-lens was fitted on {fitted_for} but is being used with " - f"{current}. J-lens matrices are model-specific; readouts will be " - "meaningless across models." + def __init__(self, *args: Any, **kwargs: Any): + super().__init__(*args, **kwargs) + self._vector_norms: Dict[int, torch.Tensor] = {} + + def _jlens_vector_norms(self, layer: int, J: torch.Tensor) -> torch.Tensor: + """``||J_l^T w_t||`` for every token, cached per layer. + + ``W_U @ J`` is ``[vocab, d_model]`` -- about 3 GB in fp32 for a + 150k-token vocabulary at d_model 5120 -- so build it in chunks and keep + only the norms. + """ + if layer not in self._vector_norms: + W_U = self.W_U + norms = torch.empty(W_U.shape[0], dtype=torch.float32, device=W_U.device) + for start in range(0, W_U.shape[0], 8192): + chunk = W_U[start:start + 8192].float() @ J + norms[start:start + 8192] = chunk.norm(dim=1) + self._vector_norms[layer] = norms.clamp_min(1e-9) + return self._vector_norms[layer] + + def scores(self, h: torch.Tensor, layer: int) -> torch.Tensor: + if layer not in self.matrices: + self.transform(h, layer) # raises the helpful KeyError + J = self.matrices[layer].to(torch.float32).to(h.device) + h = h.to(torch.float32) + raw = (self.W_U @ (J @ h).to(self.W_U.dtype)).float() + return raw / (self._jlens_vector_norms(layer, J) * h.norm().clamp_min(1e-9)) + + +# -- loading reference lenses ------------------------------------------------- + +def parse_hub_path(path: str) -> Tuple[str, str]: + """Split ``hf://org/repo/sub/dir/file.pt`` into ``("org/repo", "sub/dir/file.pt")``.""" + parts = path[len(HUB_PREFIX):].split("/", 2) + if len(parts) < 3 or not all(parts): + raise ValueError( + f"Bad Hub J-lens path {path!r}: expected hf:////" + ) + return f"{parts[0]}/{parts[1]}", parts[2] + + +def _download_from_hub(path: str) -> str: + from huggingface_hub import hf_hub_download + + repo_id, filename = parse_hub_path(path) + local = hf_hub_download(repo_id, filename) + # Neuronpedia puts a config.yaml naming the model beside each lens. Fetch it + # too so the model-mismatch check has something to compare against. + try: + hf_hub_download(repo_id, f"{os.path.dirname(filename)}/config.yaml") + except Exception: + pass + return local + + +def _load_reference_pt(path: str) -> Tuple[Dict[int, torch.Tensor], Dict[str, Any]]: + # weights_only: a lens file is tensors and ints, so refuse anything that + # would need arbitrary unpickling. + checkpoint = torch.load(path, map_location="cpu", weights_only=True) + if not isinstance(checkpoint, dict) or "J" not in checkpoint: + found = sorted(checkpoint) if isinstance(checkpoint, dict) else type(checkpoint).__name__ + raise ValueError( + f"{path} is not a reference JacobianLens file (expected a 'J' key, found {found!r})" + ) + matrices = {int(layer): J.to(torch.float32) for layer, J in checkpoint["J"].items()} + config: Dict[str, Any] = { + "format": "reference_pt", + "source": path, + "layers": sorted(matrices), + "d_model": checkpoint.get("d_model"), + "num_sequences": checkpoint.get("n_prompts"), + "method": "averaged_jacobian (reference estimator)", + } + model_id = _model_id_from_sidecar(os.path.join(os.path.dirname(path), "config.yaml")) + if model_id: + config["model_id"] = model_id + return matrices, config + + +def _model_id_from_sidecar(yaml_path: str) -> Optional[str]: + """``hf_model_name`` from a Neuronpedia ``config.yaml``, without needing PyYAML.""" + if not os.path.exists(yaml_path): + return None + with open(yaml_path) as f: + match = re.search(r'^hf_model_name:\s*"?([^"\n]+)"?\s*$', f.read(), re.MULTILINE) + return match.group(1).strip() if match else None + + +def _load_npz_dir(path: str) -> Tuple[Dict[int, torch.Tensor], Dict[str, Any]]: + import numpy as np + + npz_path = os.path.join(path, "jlens.npz") + if not os.path.exists(npz_path): + raise FileNotFoundError( + f"No J-lens at {path} (expected {npz_path}, a reference .pt file, or " + f"an hf://// path). " + f"Fit one with: pts fit-jlens --model --output-path {path}" + ) + + data = np.load(npz_path) + matrices = { + int(key.split("_")[1]): torch.from_numpy(data[key]) for key in data.files + } + + config: Dict[str, Any] = {} + config_path = os.path.join(path, "config.json") + if os.path.exists(config_path): + with open(config_path) as f: + config = json.load(f) + return matrices, config + + +def _check_lens_fits_model( + matrices: Dict[int, torch.Tensor], config: Dict[str, Any], model: Any, path: str +) -> None: + """Refuse a lens whose shape cannot belong to this model; warn on a name mismatch.""" + d_model = getattr(model.config, "hidden_size", None) + for layer, J in matrices.items(): + if d_model is not None and tuple(J.shape) != (d_model, d_model): + raise ValueError( + f"J-lens at {path} has layer-{layer} matrix {tuple(J.shape)}, but this " + f"model's d_model is {d_model}. It was fitted for a different model." ) + num_layers = len(get_layer_modules(model)) + bad = [l for l in matrices if not 0 <= l < num_layers - 1] + if bad: + raise ValueError( + f"J-lens at {path} has layers {sorted(bad)}, outside 0..{num_layers - 2} " + f"for this {num_layers}-layer model. It was fitted for a different model." + ) - logger.info(f"Loaded J-lens for layers {sorted(matrices)} from {path}") - return cls(model, tokenizer, matrices=matrices, config=config) + fitted_for = config.get("model_id") + current = getattr(model.config, "_name_or_path", None) + if fitted_for and current and fitted_for != current: + logger.warning( + f"This J-lens was fitted on {fitted_for} but is being used with " + f"{current}. J-lens matrices are model-specific; readouts will be " + "meaningless across models." + ) def load_readout( @@ -431,16 +590,18 @@ def load_readout( if method == READOUT_LOGIT_LENS: return LogitLens(model, tokenizer) - if method == READOUT_JLENS: + if method in (READOUT_JLENS, READOUT_JLENS_COSINE): if not jlens_path: raise ValueError( - "--readout-method jlens requires --jlens-path pointing at fitted " - "matrices. Fit them with `pts fit-jlens`, or use " + f"--readout-method {method} requires --jlens-path pointing at fitted " + "matrices: a `pts fit-jlens` directory, a reference .pt file, or " + "hf:////. Or use " "`--readout-method logit_lens` for the zero-cost baseline " "(weaker evidence: it reads what the activation would say now, " "not what it pushes the model to say later)." ) - return JLens.load(jlens_path, model, tokenizer) + lens_cls = JLensCosine if method == READOUT_JLENS_COSINE else JLens + return lens_cls.load(jlens_path, model, tokenizer) raise ValueError( f"Unknown readout method {method!r}. Available: {', '.join(READOUT_METHODS)}." diff --git a/tests/test_jlens.py b/tests/test_jlens.py index 460cfc6..40380f5 100644 --- a/tests/test_jlens.py +++ b/tests/test_jlens.py @@ -20,7 +20,13 @@ get_unembedding, resolve_workspace_layers, ) -from pts.latent.jlens import JLens, LogitLens, load_readout # noqa: E402 +from pts.latent.jlens import ( # noqa: E402 + JLens, + JLensCosine, + LogitLens, + load_readout, + parse_hub_path, +) TINY = "hf-internal-testing/tiny-random-LlamaForCausalLM" @@ -230,3 +236,109 @@ def test_reading_an_unfitted_layer_is_an_error(model_and_tokenizer): jl = JLens(model, tok, matrices={0: torch.eye(model.config.hidden_size)}) with pytest.raises(KeyError, match="No J-lens matrix fitted for layer 1"): jl.read(torch.randn(model.config.hidden_size), 1) + + +# -- reference (anthropics/jacobian-lens) lenses --------------------------- + +def _save_reference_lens(path, matrices, n_prompts=7): + """Write the layout ``jlens.JacobianLens.save`` produces: fp16, int layer keys.""" + d_model = next(iter(matrices.values())).shape[0] + torch.save( + { + "J": {l: J.to(torch.float16) for l, J in matrices.items()}, + "n_prompts": n_prompts, + "source_layers": sorted(matrices), + "d_model": d_model, + }, + path, + ) + + +def test_loads_the_reference_pt_layout(model_and_tokenizer, tmp_path): + model, tok = model_and_tokenizer + d_model = model.config.hidden_size + J = torch.randn(d_model, d_model).to(torch.float16).float() # exact in fp16 + _save_reference_lens(tmp_path / "lens.pt", {0: J}) + + loaded = JLens.load(str(tmp_path / "lens.pt"), model, tok) + assert loaded.layers == [0] + assert loaded.matrices[0].dtype == torch.float32 + assert torch.equal(loaded.matrices[0], J) + assert loaded.config["num_sequences"] == 7 + assert loaded.config["format"] == "reference_pt" + + # Same convention as our own fit: the loaded lens reads exactly like one + # built from the same matrix. + h = torch.randn(d_model) + direct = JLens(model, tok, matrices={0: J}).read(h, 0, top_k=5) + assert [r.token_id for r in loaded.read(h, 0, top_k=5)] == [r.token_id for r in direct] + + +def test_reference_lens_picks_up_the_model_from_its_sidecar(model_and_tokenizer, tmp_path): + model, tok = model_and_tokenizer + d_model = model.config.hidden_size + _save_reference_lens(tmp_path / "lens.pt", {0: torch.eye(d_model)}) + (tmp_path / "config.yaml").write_text( + '# Neuronpedia fit\nnp_model_id: "x"\nhf_model_name: "org/some-model"\n' + ) + loaded = JLens.load(str(tmp_path / "lens.pt"), model, tok) + assert loaded.config["model_id"] == "org/some-model" + + +def test_a_lens_for_another_model_is_refused(model_and_tokenizer, tmp_path): + """A wrong-width lens would fail deep inside a matmul, or worse, not at all + if the widths happened to line up with a transposed read.""" + model, tok = model_and_tokenizer + d_model = model.config.hidden_size + _save_reference_lens(tmp_path / "wide.pt", {0: torch.eye(d_model + 1)}) + with pytest.raises(ValueError, match="d_model"): + JLens.load(str(tmp_path / "wide.pt"), model, tok) + + num_layers = len(get_layer_modules(model)) + _save_reference_lens(tmp_path / "deep.pt", {num_layers + 3: torch.eye(d_model)}) + with pytest.raises(ValueError, match="outside"): + JLens.load(str(tmp_path / "deep.pt"), model, tok) + + +def test_a_non_lens_pt_file_is_refused(model_and_tokenizer, tmp_path): + model, tok = model_and_tokenizer + torch.save({"jacobian_sum": {}, "n_done": 0}, tmp_path / "ckpt.pt") + with pytest.raises(ValueError, match="not a reference JacobianLens file"): + JLens.load(str(tmp_path / "ckpt.pt"), model, tok) + + +def test_hub_paths_split_into_repo_and_file(): + assert parse_hub_path( + "hf://neuronpedia/jacobian-lens/gpt2-small/jlens/Salesforce-wikitext/gpt2_jacobian_lens.pt" + ) == ( + "neuronpedia/jacobian-lens", + "gpt2-small/jlens/Salesforce-wikitext/gpt2_jacobian_lens.pt", + ) + with pytest.raises(ValueError, match="hf://"): + parse_hub_path("hf://neuronpedia/jacobian-lens") + + +# -- the cosine readout ---------------------------------------------------- + +def test_cosine_readout_matches_the_workspacebench_formula(model_and_tokenizer): + """score_t = / (||J^T w_t|| ||h||): no final norm, no softmax.""" + model, tok = model_and_tokenizer + d_model = model.config.hidden_size + J = torch.randn(d_model, d_model) + h = torch.randn(d_model) + W = get_unembedding(model).float() + + expected = (W @ (J @ h)) / ((W @ J).norm(dim=1) * h.norm()) + top = torch.topk(expected, 5) + + results = JLensCosine(model, tok, matrices={0: J}).read(h, 0, top_k=5) + assert [r.token_id for r in results] == top.indices.tolist() + assert all(abs(r.score - e) < 1e-4 for r, e in zip(results, top.values.tolist())) + assert all(-1.0 <= r.score <= 1.0 for r in results) + assert all(r.readout_method == "jlens_cosine" for r in results) + + +def test_cosine_readout_needs_a_lens_too(model_and_tokenizer): + model, tok = model_and_tokenizer + with pytest.raises(ValueError, match="requires --jlens-path"): + load_readout(model, tok, method="jlens_cosine", jlens_path=None) diff --git a/tests/test_regressions.py b/tests/test_regressions.py index 2800e35..630a796 100644 --- a/tests/test_regressions.py +++ b/tests/test_regressions.py @@ -133,6 +133,53 @@ def test_filter_has_no_cross_scale_min_score(): ) + +@pytest.fixture +def two_readout_scales(): + """The same meta-token read two ways: a softmax probability and a cosine.""" + s = EventStorage() + s.add_event(make_latent_event( + query="q", context="c", metatoken=" the", token_id=2, + score=0.92, layer=8, model_id="m", position=-1, readout_method="jlens", + )) + s.add_event(make_latent_event( + query="q", context="c", metatoken=" Paris", token_id=3, + score=0.09, layer=8, model_id="m", position=-1, readout_method="jlens_cosine", + )) + return s + + +def test_latent_scores_on_different_scales_are_never_ranked_together(two_readout_scales): + """A cosine of 0.09 can be the strongest readout there is; a probability of + 0.92 can be filler. Sorting or thresholding them together is the prob-delta + vs readout bug again, one level down.""" + with pytest.raises(ValueError, match="different scales"): + two_readout_scales.most_surfaced(5) + with pytest.raises(ValueError, match="different scales"): + two_readout_scales.filter(min_readout_score=0.5) + + one_scale = two_readout_scales.filter(criteria={"readout_method": "jlens_cosine"}) + assert [e.label for e in one_scale.most_surfaced(5)] == [" Paris"] + + +def test_summary_reports_each_readout_scale_separately(two_readout_scales): + s = two_readout_scales.summary() + assert s["average_readout_score"] is None + assert s["max_readout_score"] is None + assert s["readout_score_by_scale"]["probability"]["max"] == pytest.approx(0.92) + assert s["readout_score_by_scale"]["cosine"]["max"] == pytest.approx(0.09) + + +def test_export_min_score_refuses_mixed_readout_scales(two_readout_scales, tmp_path): + from pts.exporters import EventExporter + + exporter = EventExporter(two_readout_scales) + with pytest.raises(ValueError, match="different scales"): + exporter.export_metatokens(str(tmp_path / "m.jsonl"), min_score=0.05) + # No floor, no comparison: exporting everything is fine. + exporter.export_metatokens(str(tmp_path / "m.jsonl")) + assert len((tmp_path / "m.jsonl").read_text().splitlines()) == 2 + # -- 4. shuffle_control key consistency ------------------------------------ def test_shuffle_control_has_the_same_keys_on_every_path():