diff --git a/.github/workflows/checks.yml b/.github/workflows/checks.yml index e5acdab05f..4a1fb4664e 100644 --- a/.github/workflows/checks.yml +++ b/.github/workflows/checks.yml @@ -443,6 +443,7 @@ jobs: - "Patchscopes_Generation_Demo" - "Santa_Coder" # - "stable_lm" + - "SVD_Circuits_Demo" - "T5" requires_hf_token: [false] include: diff --git a/demos/SVD_Circuits_Demo.ipynb b/demos/SVD_Circuits_Demo.ipynb new file mode 100644 index 0000000000..7354d5f280 --- /dev/null +++ b/demos/SVD_Circuits_Demo.ipynb @@ -0,0 +1,1009 @@ +{ + "cells": [ + { + "cell_type": "markdown", + "id": "3aed1a91", + "metadata": {}, + "source": [ + "\"Open" + ] + }, + { + "cell_type": "markdown", + "id": "ef1e2b03", + "metadata": {}, + "source": [ + "# SVD Circuits Demo\n", + "\n", + "This notebook decomposes one attention head's QK ($W_Q W_K^\\top$) and OV ($W_V W_O$) weight maps into orthogonal singular directions. It displays descriptive vocab and activation projections, then measures the effects of retaining or ablating eligible OV directions on one IOI prompt's logit difference. Mathematically distinct directions need not be semantically or causally distinct subfunctions.\n", + "\n", + "Based on Areeb Ahmad, Abhinav Joshi, Ashutosh Modi, \"Beyond Components: Singular Vector-Based Interpretability of Transformer Circuits\", [arXiv 2511.20273](https://arxiv.org/abs/2511.20273)." + ] + }, + { + "cell_type": "markdown", + "id": "120f870e", + "metadata": {}, + "source": [ + "## Setup (Ignore)" + ] + }, + { + "cell_type": "code", + "execution_count": 1, + "id": "51b1c187", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-28T03:32:15.180735Z", + "iopub.status.busy": "2026-09-28T03:32:15.180507Z", + "iopub.status.idle": "2026-09-28T03:32:15.249892Z", + "shell.execute_reply": "2026-09-28T03:32:15.249432Z" + } + }, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "Running as a Jupyter notebook - intended for development only!\n" + ] + } + ], + "source": [ + "# NBVAL_IGNORE_OUTPUT\n", + "# Janky code to do different setup when run in a Colab notebook vs VSCode\n", + "import os\n", + "\n", + "DEVELOPMENT_MODE = False\n", + "IN_GITHUB = os.getenv(\"GITHUB_ACTIONS\") == \"true\"\n", + "try:\n", + " import google.colab\n", + "\n", + " IN_COLAB = True\n", + " print(\"Running as a Colab notebook\")\n", + "except ImportError:\n", + " IN_COLAB = False\n", + " print(\"Running as a Jupyter notebook - intended for development only!\")\n", + " DEVELOPMENT_MODE = True\n", + " from IPython import get_ipython\n", + "\n", + " ipython = get_ipython()\n", + " ipython.run_line_magic(\"load_ext\", \"autoreload\")\n", + " ipython.run_line_magic(\"autoreload\", \"2\")\n", + "\n", + "if IN_COLAB or IN_GITHUB:\n", + " %pip install transformer_lens" + ] + }, + { + "cell_type": "code", + "execution_count": 2, + "id": "b260d8c2", + "metadata": { + "execution": { + "iopub.execute_input": "2026-09-28T03:32:15.251062Z", + "iopub.status.busy": "2026-09-28T03:32:15.250983Z", + "iopub.status.idle": "2026-09-28T03:32:24.853702Z", + "shell.execute_reply": "2026-09-28T03:32:24.853179Z" + } + }, + "outputs": [ + { + "data": { + "application/vnd.jupyter.widget-view+json": { + "model_id": "0d7d9821c27948d385b373deb184db49", + "version_major": 2, + "version_minor": 0 + }, + "text/plain": [ + "Loading weights: 0%| | 0/148 [00:004} {'sigma':>10} {'sigma/sigma_max':>16} {'degenerate':>11} {'null':>6}\")\n", + "for row in ov.rank_report[:12]:\n", + " print(\n", + " f\"{row.idx:>4} {row.sigma:>10.4f} {row.sigma_ratio:>16.4f} \"\n", + " f\"{str(row.is_degenerate):>11} {str(row.is_null):>6}\"\n", + " )" + ] + }, + { + "cell_type": "markdown", + "id": "876a71ee", + "metadata": {}, + "source": [ + "### Reading the rank report\n", + "\n", + "Near-equal consecutive singular values define a subspace whose basis is only defined up to a rotation, so individual directions inside such a block are not attributable. Numerically null directions are arbitrary vectors from the map's null space, and are equally unusable.\n", + "\n", + "`decompose_head` marks degenerate and null directions in the rank report. `logit_signature` refuses individual directions that are degenerate or null. `patch_along_directions` refuses selections that split a degenerate block, but permits complete blocks as subspaces; empty or full-span retained sets require an explicit threshold.\n", + "\n", + "`vocab_readout` and `project_activations` return raw numerical projections, including degenerate columns, without replacing them with block summaries. The displays below therefore select only isolated, non-null directions from the rank report. Without this filtering, \"direction 3 is the surname subfunction\" would be a claim about an arbitrary basis choice." + ] + }, + { + "cell_type": "markdown", + "id": "d7082dbb", + "metadata": {}, + "source": [ + "## 2. Vocab projections of eligible directions" + ] + }, + { + "cell_type": "code", + "execution_count": 4, + "id": "d553bfc3", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "direction 0: ['Sav', 'AV', 'av', 'adh', 'avour', 'aw', ' Gaw', 'tein', ' Sav', 'amac']\n", + "direction 1: ['NI', 'ANI', 'irin', 'Ir', ' NI', 'tro', ' Irwin', 'iq', 'ani', ' Ir']\n", + "direction 2: [' Lindsay', ' McKenna', 'etr', ' 339', 'Lind', ' copper', ' seiz', ' Ack', 'gae', ' Instr']\n", + "direction 23: [' Sieg', ' Mahjong', ' Buk', 'abal', 'tery', 'Mp', 'jp', 'hap', 'YN', 'ambo']\n", + "direction 48: ['û', 'Rand', ' Pupp', ' DN', 'Doc', 'iris', ' Blog', ' Ning', ' Wade', ' Ashe']\n" + ] + } + ], + "source": [ + "# NBVAL_IGNORE_OUTPUT\n", + "# Token rankings can change across kernels when float32 projection magnitudes nearly tie.\n", + "assert all(not ov.rank_report[i].is_degenerate and not ov.rank_report[i].is_null for i in display_ids)\n", + "if display_ids:\n", + " # Readout columns retain the original direction IDs, so request the full prefix.\n", + " readout = vocab_readout(model, ov, k=max(display_ids) + 1)\n", + " for i in display_ids:\n", + " top = torch.topk(readout[:, i].abs(), 10).indices\n", + " tokens = [model.tokenizer.decode([int(t)]) for t in top]\n", + " print(f\"direction {i}: {tokens}\")\n", + "else:\n", + " print(\"No attributable directions to display.\")\n" + ] + }, + { + "cell_type": "markdown", + "id": "ea384ef5", + "metadata": {}, + "source": [ + "Vocab projections can motivate hypotheses but do not establish semantic labels. `SVDInterpreter`'s own docstring warns that singular directions are not reliably interpretable. The next section describes activation coefficients; the intervention section measures changes in the selected prompt's logit difference, without confirming a named subfunction." + ] + }, + { + "cell_type": "markdown", + "id": "b8bba160", + "metadata": {}, + "source": [ + "## 3. Activation projections on a prompt" + ] + }, + { + "cell_type": "code", + "execution_count": 5, + "id": "ec0c9cdc", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "prompt tokens: ['<|endoftext|>', 'When', ' Mary', ' and', ' John', ' went', ' to', ' the', ' store', ',', ' John', ' gave', ' a', ' drink', ' to']\n", + "\n", + " token dir0 dir1 dir2 dir23 dir48\n", + "<|endoftext|> 0.34 0.20 0.17 0.14 0.27\n", + " When 0.34 0.21 0.19 0.12 0.29\n", + " Mary 0.31 0.21 0.19 0.15 0.27\n", + " and 0.20 0.37 0.51 0.35 0.20\n", + " John 0.26 0.19 0.18 0.14 0.30\n", + " went 0.34 0.20 0.17 0.13 0.28\n", + " to 0.27 0.18 0.16 0.15 0.31\n", + " the 0.34 0.22 0.20 0.13 0.27\n", + " store 0.31 0.24 0.23 0.11 0.31\n", + " , 3.46 0.04 0.65 1.49 0.77\n", + " John 0.27 0.17 0.16 0.10 0.36\n", + " gave 5.59 1.41 2.72 2.42 0.07\n", + " a 0.25 0.27 0.22 0.11 0.29\n", + " drink 0.29 0.19 0.24 0.07 0.32\n", + " to 7.31 1.74 3.41 3.24 0.24\n" + ] + } + ], + "source": [ + "# NBVAL_IGNORE_OUTPUT\n", + "# The coefficients are float32 reductions over the head's output, so the second\n", + "# decimal can differ between CPU kernels (1.75 vs 1.74 on the same input). The\n", + "# table is here to show which directions fire at which positions, not to pin a\n", + "# numeric value, so its output is not compared.\n", + "projection = project_activations(model, ov, PROMPT)\n", + "\n", + "print(\"prompt tokens:\", projection.str_tokens)\n", + "print()\n", + "assert all(not ov.rank_report[i].is_degenerate and not ov.rank_report[i].is_null for i in display_ids)\n", + "if display_ids:\n", + " print(f\"{'token':>10} \" + \" \".join(f\"{f'dir{i}':>8}\" for i in display_ids))\n", + " for pos, token in enumerate(projection.str_tokens):\n", + " coeffs = \" \".join(f\"{abs(float(projection.coefficients[pos, i])):>8.2f}\" for i in display_ids)\n", + " print(f\"{token:>10} {coeffs}\")\n", + "else:\n", + " print(\"No attributable directions to display.\")\n" + ] + }, + { + "cell_type": "markdown", + "id": "2e9d78d1", + "metadata": {}, + "source": [ + "## 4. Intervention comparisons\n", + "\n", + "`patch_along_directions` projects the head's output onto a chosen singular subspace and measures the resulting change in the selected metric, including downstream computation. It reports:\n", + "\n", + "- `delta_metric`: patched metric minus original metric.\n", + "- `baseline_delta_metric`: the mean magnitude over random same-width subspaces drawn inside the head's own OV span, so the control is an arbitrary subspace of *this head's* output rather than an unrelated residual-stream direction.\n", + "- `gated`: whether the metric change passed the mode-specific threshold comparison.\n", + "\n", + "The meaning of `gated` depends on the mode:\n", + "\n", + "- With `keep=[i]` (retain only direction $i$), `gated` means `abs(delta_metric) < baseline`: the selected metric changes less than the sampled mean control magnitude.\n", + "- With `ablate=[i]` (zero direction $i$), `gated` means `abs(delta_metric) > baseline`: the selected metric changes more than that mean.\n", + "\n", + "`keep=S` and `ablate=complement(S)` resolve to the same retained set, which is why the comparison must depend on the mode. The mean control magnitude is not a confidence bound or p-value; a passing comparison does not establish reconstruction of the head's entire behaviour or statistical separation from arbitrary directions.\n", + "\n", + "The whole-head ablation below provides context for the direction tables. It reports the unmodified metric and the signed change after zeroing the head's output. Its empty-span control coincides with the intervention, so its gate flag is not interpreted." + ] + }, + { + "cell_type": "code", + "execution_count": 6, + "id": "70138572", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + "original logit difference: 3.362\n", + "whole-head ablation delta: +0.129\n", + "\n", + " dir sigma delta baseline gated\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 0 8.4096 0.154 0.127 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 1 8.2600 0.141 0.127 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 2 8.1019 0.132 0.127 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 23 6.8420 0.193 0.127 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 48 5.5535 0.131 0.127 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 49 5.4845 0.113 0.127 True\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 50 5.3609 0.131 0.127 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 51 5.2535 0.163 0.127 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 55 5.0464 0.078 0.127 True\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 56 4.9738 0.114 0.127 True\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 59 4.7041 0.131 0.127 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 60 4.6103 0.136 0.127 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 61 4.4351 0.172 0.127 False\n" + ] + } + ], + "source": [ + "# NBVAL_IGNORE_OUTPUT\n", + "# The deltas are float32 reductions, and a value sitting on a rounding boundary\n", + "# can print one digit apart across CPU kernels. The gate's polarity is asserted\n", + "# by the test suite; this table is the illustration.\n", + "mary_token = model.to_single_token(\" Mary\")\n", + "john_token = model.to_single_token(\" John\")\n", + "metric = lambda logits: float(logits[0, -1, mary_token] - logits[0, -1, john_token])\n", + "\n", + "whole_head = patch_along_directions(\n", + " model, ov, PROMPT, metric, keep=[], threshold=0.0, n_baseline=1,\n", + " rng=torch.Generator().manual_seed(0),\n", + ")\n", + "print(f\"original logit difference: {whole_head.original_metric:.3f}\")\n", + "print(f\"whole-head ablation delta: {whole_head.delta_metric:+.3f}\")\n", + "print()\n", + "\n", + "# Sweep every attributable direction, not just the first few: the gate is a\n", + "# comparison against a random control, so which directions pass is the result.\n", + "# Deltas print at three decimals because the fourth differs between CPU kernels.\n", + "print(f\"{'dir':>4} {'sigma':>9} {'delta':>10} {'baseline':>10} {'gated':>6}\")\n", + "for idx in eligible_ids:\n", + " ov.require_isolated(idx)\n", + " row = ov.rank_report[idx]\n", + " result = patch_along_directions(\n", + " model, ov, PROMPT, metric, keep=[idx], rng=torch.Generator().manual_seed(0)\n", + " )\n", + " print(\n", + " f\"{row.idx:>4} {row.sigma:>9.4f} {result.delta_metric:>10.3f} \"\n", + " f\"{result.baseline_delta_metric:>10.3f} {str(result.gated):>6}\"\n", + " )\n" + ] + }, + { + "cell_type": "code", + "execution_count": 7, + "id": "8044e885", + "metadata": {}, + "outputs": [ + { + "name": "stdout", + "output_type": "stream", + "text": [ + " dir sigma delta baseline gated\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 0 8.4096 -0.012 0.041 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 1 8.2600 -0.015 0.041 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 2 8.1019 -0.014 0.041 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 23 6.8420 -0.058 0.041 True\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 48 5.5535 -0.001 0.041 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 49 5.4845 0.016 0.041 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 50 5.3609 -0.005 0.041 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 51 5.2535 -0.036 0.041 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 55 5.0464 0.040 0.041 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 56 4.9738 0.025 0.041 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 59 4.7041 0.002 0.041 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 60 4.6103 -0.016 0.041 False\n" + ] + }, + { + "name": "stdout", + "output_type": "stream", + "text": [ + " 61 4.4351 -0.032 0.041 False\n" + ] + } + ], + "source": [ + "# NBVAL_IGNORE_OUTPUT\n", + "# Same reason as the keep-mode table: the deltas are float32 reductions, and a\n", + "# value sitting on a rounding boundary can print one ulp apart across CPU kernels.\n", + "print(f\"{'dir':>4} {'sigma':>9} {'delta':>10} {'baseline':>10} {'gated':>6}\")\n", + "for idx in eligible_ids:\n", + " ov.require_isolated(idx)\n", + " row = ov.rank_report[idx]\n", + " result = patch_along_directions(\n", + " model, ov, PROMPT, metric, ablate=[idx], rng=torch.Generator().manual_seed(0)\n", + " )\n", + " print(\n", + " f\"{row.idx:>4} {row.sigma:>9.4f} {result.delta_metric:>10.3f} \"\n", + " f\"{result.baseline_delta_metric:>10.3f} {str(result.gated):>6}\"\n", + " )\n" + ] + }, + { + "cell_type": "markdown", + "id": "a6ffb744", + "metadata": {}, + "source": [ + "## 5. What this does and does not show\n", + "\n", + "This example decomposes one GPT-2-small head and measures prompt-specific effects of retaining or ablating eligible OV directions. The reported gate verdicts are comparisons with an averaged random in-span control. They do not establish named subfunctions, statistical separation from arbitrary directions, or replication of the paper's causal subfunction taxonomy.\n", + "\n", + "- **Single head, single model, single prompt.** Everything above concerns layer 9 head 9 of `gpt2-small` on one IOI prompt. Nothing here assembles a multi-head circuit.\n", + "- **Descriptive projections.** Orthogonal weight-space directions are mathematically distinct; vocab readouts and activation coefficients do not supply semantic labels.\n", + "- **Comparisons, not significance.** The tables report metric changes and mean control magnitudes, not control distributions or confidence bounds. They do not establish that any direction is distinguishable from an arbitrary one.\n", + "- **Downstream responses.** The whole-head ablation delta measures the final logit difference after downstream computation, not the head's direct write contribution. A small final-logit effect need not imply a small direct contribution; these measurements do not identify a specific compensating component.\n", + "- **Repeatability, not confirmation.** Verdicts near the mean control magnitude can change with the seed or draw count. Fixing the seed makes the comparison repeatable, not scientifically conclusive.\n", + "- **No automated labelling.** A passing gate does not name what a direction computes. Named-subfunction claims would need additional prompts, counterfactuals, and suitable control-distribution analyses." + ] + } + ], + "metadata": { + "language_info": { + "codemirror_mode": { + "name": "ipython", + "version": 3 + }, + "file_extension": ".py", + "mimetype": "text/x-python", + "name": "python", + "nbconvert_exporter": "python", + "pygments_lexer": "ipython3", + "version": "3.12.14" + }, + "widgets": { + "application/vnd.jupyter.widget-state+json": { + "state": { + "0d7d9821c27948d385b373deb184db49": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "2.0.0", + "model_name": "HBoxModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "2.0.0", + "_model_name": "HBoxModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "2.0.0", + "_view_name": "HBoxView", + "box_style": "", + "children": [ + "IPY_MODEL_24f11b35f964473192adb75bfbdbf493", + "IPY_MODEL_b22b885ac6564eea9cc8335688983a09", + "IPY_MODEL_e21e1c9765e94501a539b796224d928e" + ], + "layout": "IPY_MODEL_6f36e52abc6b437abf8702731770d747", + "tabbable": null, + "tooltip": null + } + }, + "12e5daca51bc4f4ab0830e9f3c6d11b6": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "2.0.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "2.0.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "2.0.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border_bottom": null, + "border_left": null, + "border_right": null, + "border_top": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "24f11b35f964473192adb75bfbdbf493": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "2.0.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "2.0.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "2.0.0", + "_view_name": "HTMLView", + "description": "", + "description_allow_html": false, + "layout": "IPY_MODEL_340f81a8ccbc41a2b4eb96ff2d49c5e1", + "placeholder": "​", + "style": "IPY_MODEL_3ad1b21f0cee48caa22019b226fde81a", + "tabbable": null, + "tooltip": null, + "value": "Loading weights: 100%" + } + }, + "340f81a8ccbc41a2b4eb96ff2d49c5e1": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "2.0.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "2.0.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "2.0.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border_bottom": null, + "border_left": null, + "border_right": null, + "border_top": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "3ad1b21f0cee48caa22019b226fde81a": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "2.0.0", + "model_name": "HTMLStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "2.0.0", + "_model_name": "HTMLStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "2.0.0", + "_view_name": "StyleView", + "background": null, + "description_width": "", + "font_size": null, + "text_color": null + } + }, + "6f36e52abc6b437abf8702731770d747": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "2.0.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "2.0.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "2.0.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border_bottom": null, + "border_left": null, + "border_right": null, + "border_top": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "a81ad8e18e5345619986fa93b1de0f4e": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "2.0.0", + "model_name": "HTMLStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "2.0.0", + "_model_name": "HTMLStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "2.0.0", + "_view_name": "StyleView", + "background": null, + "description_width": "", + "font_size": null, + "text_color": null + } + }, + "af26ae76f36443f59561c3359e95da17": { + "model_module": "@jupyter-widgets/base", + "model_module_version": "2.0.0", + "model_name": "LayoutModel", + "state": { + "_model_module": "@jupyter-widgets/base", + "_model_module_version": "2.0.0", + "_model_name": "LayoutModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "2.0.0", + "_view_name": "LayoutView", + "align_content": null, + "align_items": null, + "align_self": null, + "border_bottom": null, + "border_left": null, + "border_right": null, + "border_top": null, + "bottom": null, + "display": null, + "flex": null, + "flex_flow": null, + "grid_area": null, + "grid_auto_columns": null, + "grid_auto_flow": null, + "grid_auto_rows": null, + "grid_column": null, + "grid_gap": null, + "grid_row": null, + "grid_template_areas": null, + "grid_template_columns": null, + "grid_template_rows": null, + "height": null, + "justify_content": null, + "justify_items": null, + "left": null, + "margin": null, + "max_height": null, + "max_width": null, + "min_height": null, + "min_width": null, + "object_fit": null, + "object_position": null, + "order": null, + "overflow": null, + "padding": null, + "right": null, + "top": null, + "visibility": null, + "width": null + } + }, + "b22b885ac6564eea9cc8335688983a09": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "2.0.0", + "model_name": "FloatProgressModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "2.0.0", + "_model_name": "FloatProgressModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "2.0.0", + "_view_name": "ProgressView", + "bar_style": "success", + "description": "", + "description_allow_html": false, + "layout": "IPY_MODEL_af26ae76f36443f59561c3359e95da17", + "max": 148.0, + "min": 0.0, + "orientation": "horizontal", + "style": "IPY_MODEL_d015b8b53ca44792a760b887c651702c", + "tabbable": null, + "tooltip": null, + "value": 148.0 + } + }, + "d015b8b53ca44792a760b887c651702c": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "2.0.0", + "model_name": "ProgressStyleModel", + "state": { + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "2.0.0", + "_model_name": "ProgressStyleModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/base", + "_view_module_version": "2.0.0", + "_view_name": "StyleView", + "bar_color": null, + "description_width": "" + } + }, + "e21e1c9765e94501a539b796224d928e": { + "model_module": "@jupyter-widgets/controls", + "model_module_version": "2.0.0", + "model_name": "HTMLModel", + "state": { + "_dom_classes": [], + "_model_module": "@jupyter-widgets/controls", + "_model_module_version": "2.0.0", + "_model_name": "HTMLModel", + "_view_count": null, + "_view_module": "@jupyter-widgets/controls", + "_view_module_version": "2.0.0", + "_view_name": "HTMLView", + "description": "", + "description_allow_html": false, + "layout": "IPY_MODEL_12e5daca51bc4f4ab0830e9f3c6d11b6", + "placeholder": "​", + "style": "IPY_MODEL_a81ad8e18e5345619986fa93b1de0f4e", + "tabbable": null, + "tooltip": null, + "value": " 148/148 [00:00<00:00, 6931.73it/s]" + } + } + }, + "version_major": 2, + "version_minor": 0 + } + } + }, + "nbformat": 4, + "nbformat_minor": 5 +} diff --git a/docs/make_docs.py b/docs/make_docs.py index 074975db40..5428107e4d 100644 --- a/docs/make_docs.py +++ b/docs/make_docs.py @@ -43,6 +43,7 @@ def copy_demos(_app: Optional[Any] = None): "Jacobian_Lens_Coordinate_Patch_Benchmark_Demo.ipynb", "Jacobian_Lens_Decomposition_Demo.ipynb", "Main_Demo.ipynb", + "SVD_Circuits_Demo.ipynb", ] if copy_to_dir.exists(): diff --git a/docs/source/content/analysis_tools.md b/docs/source/content/analysis_tools.md index fcb334a2fb..1ebfbd94ef 100644 --- a/docs/source/content/analysis_tools.md +++ b/docs/source/content/analysis_tools.md @@ -139,8 +139,12 @@ For SVD head decomposition, inspect the rank report before assigning meaning to individual direction. Near-equal singular values define a subspace whose basis can rotate; numerically null directions are also unsuitable for individual attribution. The weight decomposition alone is not a causal validation of a proposed subfunction. +See [SVD Circuits](svd_circuits.md) for the degeneracy guard, the causal gate, and a +worked example, and the [SVD Circuits demo](../generated/demos/SVD_Circuits_Demo.html) +for a runnable walkthrough. -API: {func}`~transformer_lens.tools.analysis.svd_circuits.decompose_head`. +API: {func}`~transformer_lens.tools.analysis.svd_circuits.decompose_head`, +{func}`~transformer_lens.tools.analysis.svd_circuits.patch_along_directions`. ## Try a geometry question without downloading a model diff --git a/docs/source/content/svd_circuits.md b/docs/source/content/svd_circuits.md new file mode 100644 index 0000000000..61eac1b3e8 --- /dev/null +++ b/docs/source/content/svd_circuits.md @@ -0,0 +1,147 @@ +# SVD Circuits + +SVD Circuits decomposes a single attention head's QK and OV weight maps into orthogonal +singular directions. Vocab readouts and activation projections describe those directions; +patching measures prompt-specific changes in a caller-selected output metric. + +Mathematically distinct directions need not be semantically or causally distinct +subfunctions. A weight-space decomposition or a plausible token projection does not +establish a direction's role. Interventions provide additional measurements, but the +reported gate comparison is not by itself a validation of a named mechanism. + +## Definition + +For a head at layer $\ell$ with head index $h$, the two maps are + +$$ +\mathrm{QK} = W_Q W_K^\top, \qquad \mathrm{OV} = W_V W_O, +$$ + +each of shape $d_{\text{model}} \times d_{\text{model}}$ but of rank at most +$d_{\text{head}}$. Their SVDs are + +$$ +W_V W_O = U \Sigma V^\top, \qquad \operatorname{rank} \le d_{\text{head}}. +$$ + +The convention matters, because both factors are $d_{\text{model}}$-wide and a swap +raises no shape error: + +- For OV, the columns of $V$ are the residual-stream **output/write** directions, the + ones projected through $W_U$ for a vocab readout. The columns of $U$ span the + value-computation **input** space. +- For QK, the columns of $V$ are the **source/key-read** directions and the columns of + $U$ the **destination/query-read** directions. QK produces no write direction, so it + has no vocab readout. + +`FactoredMatrix` computes both SVDs without materialising the $d_{\text{model}}^2$ +product. + +## The degeneracy guard + +Singular directions are unique only when the singular values are distinct. Equal or +near-equal consecutive singular values leave the corresponding subspace defined only up +to an arbitrary rotation, so any statement of the form "direction 3 is the surname +subfunction" is a statement about a basis choice rather than about the model. + +`decompose_head` returns a rank report marking degenerate and numerically null +directions. Null directions are arbitrary vectors from the map's null space and are +also unsuitable for individual attribution. The consumers have different guards: + +- `logit_signature` requires an isolated, non-null direction and refuses individual + directions inside a degenerate block. +- `patch_along_directions` refuses selections that split a degenerate block. Complete + blocks can be retained or removed as subspaces; empty or full-span retained sets + require an explicit `threshold` because their random controls coincide with them. +- `vocab_readout` and `project_activations` return raw numerical projections, including + columns inside degenerate blocks. They neither refuse these columns nor replace + them with block summaries. Consult the rank report and exclude both `is_degenerate` + and `is_null` before interpreting individual directions. + +## The intervention comparison + +`patch_along_directions` reconstructs the head's output onto a chosen singular subspace +and reports the resulting change in a caller-supplied metric: + +- `delta_metric`: patched metric minus original metric. +- `baseline_delta_metric`: the mean magnitude over random same-width subspaces drawn + inside the head's own OV span. Drawing the control in-span rather than from the full + residual stream makes it the effect of an arbitrary subspace of *this head's* output, + which is the comparison the gate needs. +- `gated`: whether the metric change passed the mode-specific threshold comparison. + +The comparison depends on the mode the caller expressed, because `keep=S` and +`ablate=complement(S)` resolve to the same retained set: + +| Mode | Meaning | `gated` is | +|---|---|---| +| `keep=S` | retain only $S$ | `abs(delta_metric) < threshold` | +| `ablate=S` | zero $S$, retain the rest | `abs(delta_metric) > threshold` | + +With `keep`, a passing comparison means retaining the selected subspace changes the +metric less than the threshold. With `ablate`, it means removing the selected subspace +changes the metric more than the threshold. These describe effects on the chosen metric, +not reconstruction of the head's entire behaviour or identification of its function. + +The threshold defaults to `baseline_delta_metric`, a sampled mean magnitude rather than +a confidence bound or p-value. It does not establish statistical separation from +arbitrary directions. The sampled mean can vary with the seed and draw count, so verdicts +near it can change. Pass an explicit `rng` for repeatability and increase `n_baseline` to +sample the mean more thoroughly; neither makes a verdict scientifically conclusive. + +## Compatibility mode + +The vocab readout projects through the final LayerNorm folded into $W_U$, which requires +`enable_compatibility_mode()` on an adapter that supports folding. A decomposition +records the folded-LayerNorm state it was taken under, and the readout and patch +consumers refuse a decomposition whose state no longer matches the model. See +[Compatibility Mode](compatibility_mode.md). + +## Worked example + +```python +import torch + +from transformer_lens.model_bridge import TransformerBridge +from transformer_lens.tools.analysis.svd_circuits import ( + decompose_head, + patch_along_directions, + vocab_readout, +) + +model = TransformerBridge.boot_transformers("gpt2", dtype=torch.float32, device="cpu") +model.enable_compatibility_mode() +model.eval() + +prompt = "When Mary and John went to the store, John gave a drink to" +ov = decompose_head(model, layer=9, head=9, which=("OV",)).OV + +readout = vocab_readout(model, ov, k=10) + +mary, john = model.to_single_token(" Mary"), model.to_single_token(" John") +metric = lambda logits: float(logits[0, -1, mary] - logits[0, -1, john]) + +result = patch_along_directions( + model, ov, prompt, metric, keep=[0], rng=torch.Generator().manual_seed(0) +) +print(result.delta_metric, result.baseline_delta_metric, result.gated) +``` + +## What this does not establish + +- A passing gate reports a metric comparison for one intervention on one prompt. Vocab + readouts and activation coefficients are descriptive projections, not semantic labels. +- Results are single-head and prompt-specific. Nothing here assembles a multi-head circuit. +- The worked example and demo do not establish named subfunctions, statistical separation + from arbitrary directions, or replication of the paper's causal subfunction taxonomy. +- Output-metric changes include downstream responses. A small final-logit effect does not + imply a small direct write contribution, nor identify which downstream components alter + the effect. +- Verdicts near the mean control magnitude can change with the seed or draw count. Fixing + a seed makes the comparison repeatable, not scientifically conclusive. + +## Links + +- [SVD Circuits demo](../generated/demos/SVD_Circuits_Demo.ipynb) +- Areeb Ahmad, Abhinav Joshi, Ashutosh Modi, "Beyond Components: Singular Vector-Based + Interpretability of Transformer Circuits", [arXiv 2511.20273](https://arxiv.org/abs/2511.20273) diff --git a/docs/source/index.md b/docs/source/index.md index e81b4cf4cf..3a513e97ec 100644 --- a/docs/source/index.md +++ b/docs/source/index.md @@ -67,6 +67,8 @@ generated/demos/Jacobian_Lens_Coordinate_Patch_Benchmark_Demo content/backward_lens content/debugging_numerical_divergence content/sparse_probing +content/svd_circuits +generated/demos/SVD_Circuits_Demo generated/demos/Main_Demo generated/demos/Exploratory_Analysis_Demo content/special_cases diff --git a/makefile b/makefile index 2f0c951949..8958a4f118 100644 --- a/makefile +++ b/makefile @@ -58,6 +58,7 @@ notebook-test: $(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/Qwen.ipynb $(RERUN_ARGS) $(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/Santa_Coder.ipynb $(RERUN_ARGS) $(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/stable_lm.ipynb $(RERUN_ARGS) + $(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/SVD_Circuits_Demo.ipynb $(RERUN_ARGS) $(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/SVD_Interpreter_Demo.ipynb $(RERUN_ARGS) $(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/Tracr_to_Transformer_Lens_Demo.ipynb $(RERUN_ARGS) diff --git a/tests/integration/test_svd_circuits_oracle_parity.py b/tests/integration/test_svd_circuits_oracle_parity.py new file mode 100644 index 0000000000..9a9bba3dc2 --- /dev/null +++ b/tests/integration/test_svd_circuits_oracle_parity.py @@ -0,0 +1,190 @@ +"""Slow qualitative OV sweep on GPT-2-small layer 9 head 9. + +Reports prompt-specific gate counts without requiring a minimum or claiming a causal +subfunction split. There is no pinned external numerical oracle: the Beyond Components +paper publishes no numeric table for this head. The repository's ``SVDInterpreter`` +cross-check is covered separately in ``test_svd_circuits.py``. + +Keep-mode gates compare the metric change with a sampled mean random in-span control +magnitude, not a statistical significance threshold. Counts may vary with seed and draw +count, including zero. The checks cover eligibility, finite results, gate polarity, and +fixed-seed reproducibility. Each sweep costs ``k * (n_baseline + 2)`` forward passes, +so the model tests are marked ``slow`` and excluded from the default tiers. +""" + +from __future__ import annotations + +import math +from dataclasses import dataclass +from typing import Callable, List, Sequence, Tuple + +import pytest +import torch + +from transformer_lens.model_bridge import TransformerBridge +from transformer_lens.tools.analysis.svd_circuits import ( + HeadSVD, + decompose_head, + patch_along_directions, +) + +CLEAN_PROMPT = "When Mary and John went to the store, John gave a drink to" + +# Wang et al. (2022), arXiv 2211.00593, identifies L9H9 as a name mover in the IOI circuit. +LAYER, HEAD = 9, 9 + +TOP_K_DIRECTIONS = 8 +BASELINE_DRAWS = 32 +BASELINE_SEED = 0 + + +@pytest.fixture(scope="module") +def gpt2_bridge(): + model = TransformerBridge.boot_transformers("gpt2", device="cpu", dtype=torch.float32) + model.enable_compatibility_mode() + # Forward-pass tools require evaluation mode. + model.eval() + return model + + +def _logit_diff_metric(model) -> Callable[[torch.Tensor], float]: + mary_token = model.to_single_token(" Mary") + john_token = model.to_single_token(" John") + + def metric(logits: torch.Tensor) -> float: + return float(logits[0, -1, mary_token] - logits[0, -1, john_token]) + + return metric + + +@dataclass(frozen=True) +class _GateRow: + """One swept direction's gate verdict.""" + + idx: int + sigma: float + delta_metric: float + baseline_delta_metric: float + gated: bool + + +def _eligible_directions(head_svd: HeadSVD, k: int) -> List[int]: + """The first ``k`` directions that are attributable on their own. + + Fewer than ``k`` may qualify: the degeneracy guard is a feature, so a short sweep + is used as-is rather than treated as an error. + """ + eligible = [ + row.idx for row in head_svd.rank_report if not row.is_degenerate and not row.is_null + ] + return eligible[:k] + + +def _gated_directions( + model, + head_svd: HeadSVD, + prompt: torch.Tensor, + metric: Callable[[torch.Tensor], float], + *, + k: int, + n_baseline: int, + seed: int, +) -> List[_GateRow]: + """Report keep-mode changes against mean same-width random in-span control magnitudes.""" + rows: List[_GateRow] = [] + for idx in _eligible_directions(head_svd, k): + result = patch_along_directions( + model, + head_svd, + prompt, + metric, + keep=[idx], + rng=torch.Generator().manual_seed(seed), + n_baseline=n_baseline, + ) + rows.append( + _GateRow( + idx=idx, + sigma=head_svd.rank_report[idx].sigma, + delta_metric=result.delta_metric, + baseline_delta_metric=result.baseline_delta_metric, + gated=result.gated, + ) + ) + return rows + + +def _format_rows(rows: Sequence[_GateRow], *, seed: int, n_baseline: int) -> str: + lines = [ + f"seed={seed}, draws={n_baseline}, swept={len(rows)}, " + f"gated={sum(row.gated for row in rows)} (descriptive count)", + f"{'dir':>4} {'sigma':>10} {'delta':>12} {'baseline':>12} {'gated':>6}", + ] + for row in rows: + lines.append( + f"{row.idx:>4} {row.sigma:>10.4f} {row.delta_metric:>12.6f} " + f"{row.baseline_delta_metric:>12.6f} {str(row.gated):>6}" + ) + return "\n".join(lines) + + +def _sweep( + gpt2_bridge, *, seed: int = BASELINE_SEED, n_baseline: int = BASELINE_DRAWS +) -> Tuple[_GateRow, ...]: + decomposition = decompose_head(gpt2_bridge, layer=LAYER, head=HEAD, which=("OV",)) + ov = decomposition.OV + assert ov is not None + return tuple( + _gated_directions( + gpt2_bridge, + ov, + gpt2_bridge.to_tokens(CLEAN_PROMPT), + _logit_diff_metric(gpt2_bridge), + k=TOP_K_DIRECTIONS, + n_baseline=n_baseline, + seed=seed, + ) + ) + + +@pytest.fixture(scope="module") +def default_sweep(gpt2_bridge) -> Tuple[_GateRow, ...]: + """Share immutable default rows without replacing the fresh reproducibility sweep.""" + return _sweep(gpt2_bridge) + + +@pytest.mark.slow +@pytest.mark.parametrize("seed", [0, 1, 2]) +@pytest.mark.parametrize("n_baseline", [16, BASELINE_DRAWS]) +def test_qualitative_sweep_has_finite_results_and_consistent_gates( + gpt2_bridge, default_sweep, seed: int, n_baseline: int +) -> None: + rows = ( + default_sweep + if (seed, n_baseline) == (BASELINE_SEED, BASELINE_DRAWS) + else _sweep(gpt2_bridge, seed=seed, n_baseline=n_baseline) + ) + report = _format_rows(rows, seed=seed, n_baseline=n_baseline) + print(report) + + ov = decompose_head(gpt2_bridge, layer=LAYER, head=HEAD, which=("OV",)).OV + assert ov is not None + expected_ids = [row.idx for row in ov.rank_report if not row.is_degenerate and not row.is_null][ + :TOP_K_DIRECTIONS + ] + actual_ids = [row.idx for row in rows] + assert rows, "no attributable OV directions to sweep" + assert actual_ids == expected_ids, report + assert len(set(actual_ids)) == len(actual_ids), report + for row in rows: + assert math.isfinite(row.delta_metric), report + assert math.isfinite(row.baseline_delta_metric), report + assert row.baseline_delta_metric >= 0, report + assert row.gated == (abs(row.delta_metric) < row.baseline_delta_metric), report + + +@pytest.mark.slow +def test_sweep_is_reproducible_under_a_fixed_seed(gpt2_bridge, default_sweep) -> None: + second = _sweep(gpt2_bridge) + + assert default_sweep == second diff --git a/tests/unit/test_make_docs.py b/tests/unit/test_make_docs.py index 3066ce940c..b13dfb32ef 100644 --- a/tests/unit/test_make_docs.py +++ b/tests/unit/test_make_docs.py @@ -42,4 +42,5 @@ def test_copy_demos_creates_generated_dir_when_absent(tmp_path, monkeypatch): "Jacobian_Lens_Coordinate_Patch_Benchmark_Demo.ipynb", "Jacobian_Lens_Decomposition_Demo.ipynb", "Main_Demo.ipynb", + "SVD_Circuits_Demo.ipynb", ] diff --git a/tests/unit/tools/test_svd_circuits.py b/tests/unit/tools/test_svd_circuits.py index 384f87b419..eb56cca1cb 100644 --- a/tests/unit/tools/test_svd_circuits.py +++ b/tests/unit/tools/test_svd_circuits.py @@ -988,6 +988,53 @@ def run_with_hooks(self, prompt, fwd_hooks): return self._logits(activation) +@pytest.mark.parametrize( + "spectrum, guarded_ids", + [([8.0, 4.0, 4.0, 1.0], [1, 2]), ([8.0, 4.0, 2.0, 0.0], [3])], + ids=["degenerate", "null"], +) +def test_raw_projections_preserve_columns_rejected_by_attribution_guards(spectrum, guarded_ids): + """Raw consumers preserve numerical columns; signature and patch guards remain distinct.""" + + class ProjectionModel(_PatchStubModel): + def run_with_cache(self, prompt, names_filter): + hook_name = "blocks.0.attn.hook_result" + assert names_filter(hook_name) + return self(prompt), {hook_name: self._result} + + def to_str_tokens(self, prompt): + return [str(i) for i in range(self._result.shape[1])] + + ov = _factored_head_svd( + *_factored_with_spectrum(spectrum), which="OV", layer=0, head=0, eps=1e-2 + ) + model = ProjectionModel(d_model=D_MODEL, n_heads=1) + model.W_U = torch.randn(D_MODEL, 16) + readout = vocab_readout(model, ov, k=len(ov.S)) + projection = project_activations(model, ov, "prompt") + + assert readout.shape == (16, len(ov.S)) + assert projection.coefficients.shape == (model._result.shape[1], len(ov.S)) + assert torch.isfinite(readout[:, guarded_ids]).all() + assert torch.isfinite(projection.coefficients[:, guarded_ids]).all() + assert model.cfg.use_attn_result is False + for idx in guarded_ids: + assert ov.rank_report[idx].is_degenerate or ov.rank_report[idx].is_null + with pytest.raises(DegenerateDirectionError): + logit_signature(model, ov, direction=idx, tokens=[0]) + + block = ov.block_of(guarded_ids[0]) + if len(block) > 1: + with pytest.raises(DegenerateDirectionError, match="split block"): + patch_along_directions( + model, ov, "prompt", lambda logits: float(logits.sum()), keep=[block[0]] + ) + result = patch_along_directions( + model, ov, "prompt", lambda logits: float(logits.sum()), keep=block, n_baseline=1 + ) + assert result.retained == block + + def test_patch_along_directions_which_guard(): """QK has no write direction to reconstruct onto, so a QK HeadSVD is refused.""" qk = _factored_head_svd( @@ -1197,6 +1244,34 @@ def test_patch_threshold_above_delta_gates_false(): assert raised.gated is False +@pytest.mark.parametrize("mode", ["keep", "ablate"]) +def test_patch_threshold_equal_to_delta_gates_false(mode): + """Equality with an explicit threshold passes neither mode's strict comparison.""" + ov = _factored_head_svd( + *_factored_with_spectrum([8.0, 4.0, 2.0, 1.0]), which="OV", layer=0, head=0, eps=1e-2 + ) + stub = _span_aligned_stub(ov, [1.0, 0.5, 0.25, 0.125]) + metric = lambda logits: float(logits.sum()) + keep = [0] if mode == "keep" else None + ablate = [0] if mode == "ablate" else None + reference = patch_along_directions( + stub, ov, "prompt", metric, keep=keep, ablate=ablate, n_baseline=1 + ) + tied = patch_along_directions( + stub, + ov, + "prompt", + metric, + keep=keep, + ablate=ablate, + threshold=abs(reference.delta_metric), + n_baseline=1, + ) + + assert tied.delta_metric == reference.delta_metric + assert tied.gated is False + + def test_patch_baseline_is_reproducible(): """The averaged baseline is drawn from the passed generator, so the same seed reproduces it bit-for-bit and a different seed gives a different average.""" @@ -1460,3 +1535,170 @@ def test_patch_along_directions_restores_use_attn_result(tiny_bridge): assert tiny_bridge.cfg.use_attn_result == initial finally: tiny_bridge.set_use_attn_result(original) + + +# --------------------------------------------------------------------------- # +# Top-k recovery through activation patching +# --------------------------------------------------------------------------- # +def _block_aligned_ladder(head_svd, *, geometric=True): + """Cumulative direction counts that never split a degenerate block. + + ``patch_along_directions`` refuses a retained set that splits a block, so a + k-ladder has to advance block by block rather than one direction at a time. + Geometric spacing keeps the sweep to roughly ``log2(rank)`` calls instead of + one per direction, which matters because each call costs ``n_baseline + 2`` + forward passes. + """ + starts = [] + last = None + for row in head_svd.rank_report: + if row.block_id != last: + starts.append(row.idx) + last = row.block_id + cumulative = starts[1:] + [len(head_svd.rank_report)] + if not geometric: + return cumulative + + ladder = [] + target = 1 + for end in cumulative: + if end >= target: + ladder.append(end) + while target <= end: + target *= 2 + if ladder[-1] != cumulative[-1]: + ladder.append(cumulative[-1]) + return ladder + + +class _RecoveryStubModel(_PatchStubModel): + """Expose identity-probe head outputs without a downstream scalar readout.""" + + def __init__(self, clean): + super().__init__(d_model=clean.shape[-1], n_heads=1, pos=clean.shape[0]) + self._result = clean[None, :, None, :] + + def _logits(self, result): + return result[:, :, 0, :] + + +def _measure_patch_recovery(ov, clean, retained): + """Measure relative output error through the public patch and capture its output.""" + model = _RecoveryStubModel(clean) + outputs = [] + + def metric(output): + outputs.append(output.detach().clone()) + return float((output - clean.unsqueeze(0)).norm() / clean.norm()) + + result = patch_along_directions( + model, + ov, + "prompt", + metric, + keep=retained, + n_baseline=1, + rng=torch.Generator().manual_seed(0), + # Empty/full-span controls coincide with the intervention and need an explicit threshold. + threshold=1e-4, + ) + assert result.original_metric == 0.0 + assert result.retained == retained + assert model.cfg.use_attn_result is False + return result, outputs[1].squeeze(0) + + +@pytest.mark.parametrize( + "spectrum", [(8.0, 4.0, 2.0, 1.0), (8.0, 4.0, 4.0, 1.0)], ids=["isolated", "block"] +) +def test_top_k_recovery_converges_monotonically(spectrum): + """Patched identity probes match the top-k map and its singular-value tail error.""" + W_V, W_O = _factored_with_spectrum(spectrum) + ov = _factored_head_svd(W_V, W_O, which="OV", layer=0, head=0, eps=1e-2) + clean = W_V @ W_O + ladder = [0] + _block_aligned_ladder(ov, geometric=False) + if spectrum[1] == spectrum[2]: + assert ov.block_of(1) == [1, 2] + assert ladder == [0, 1, 3, 4] + else: + assert ladder == [0, 1, 2, 3, 4] + tolerance = 64 * torch.finfo(clean.dtype).eps + + residuals = [] + print("\nk patched residual expected residual max output error") + for k in ladder: + result, patched = _measure_patch_recovery(ov, clean, list(range(k))) + expected_output = W_V[:, :k] @ W_O[:k, :] + expected_residual = float((ov.S[k:].square().sum() / ov.S.square().sum()).sqrt()) + torch.testing.assert_close(patched, expected_output, rtol=tolerance, atol=tolerance) + assert result.patched_metric == pytest.approx(expected_residual, abs=tolerance) + residuals.append(result.patched_metric) + print( + f"{k:3d} {result.patched_metric:>18.8f} {expected_residual:>19.8f} " + f"{float((patched - expected_output).abs().max()):>18.8f}" + ) + + for prev, cur in zip(residuals, residuals[1:]): + assert cur <= prev + tolerance + assert residuals[0] == pytest.approx(1.0, abs=tolerance) + assert residuals[-1] < tolerance + + +@pytest.mark.parametrize("control", ["bottom", "rotated", "noop"]) +def test_top_k_recovery_distinguishes_wrong_subspaces(control, monkeypatch): + """Same-width wrong subspaces and a no-op cannot reproduce the top-k output.""" + W_V, W_O = _factored_with_spectrum([8.0, 4.0, 2.0, 1.0]) + ov = _factored_head_svd(W_V, W_O, which="OV", layer=0, head=0, eps=1e-2) + clean = W_V @ W_O + k = 2 + expected_output = W_V[:, :k] @ W_O[:k, :] + top, top_output = _measure_patch_recovery(ov, clean, list(range(k))) + torch.testing.assert_close(top_output, expected_output) + + if control == "bottom": + wrong, wrong_output = _measure_patch_recovery( + ov, clean, list(range(len(ov.S) - k, len(ov.S))) + ) + else: + if control == "rotated": + random_basis, _ = torch.linalg.qr( + torch.randn(len(ov.S), len(ov.S), generator=torch.Generator().manual_seed(17)) + ) + directions = ov.V @ random_basis[:, :k] + projector = directions @ directions.T + else: + projector = torch.eye(clean.shape[-1], dtype=clean.dtype) + + def control_hook(head, requested_projector): + return _make_subspace_hook(head, projector) + + monkeypatch.setattr( + "transformer_lens.tools.analysis.svd_circuits._make_subspace_hook", control_hook + ) + wrong, wrong_output = _measure_patch_recovery(ov, clean, list(range(k))) + + assert not torch.allclose(wrong_output, expected_output, rtol=1e-5, atol=1e-6) + if control == "noop": + assert wrong.patched_metric == 0.0 + else: + assert wrong.patched_metric > top.patched_metric + 0.1 + + +def test_top_k_recovery_full_span_is_noop_on_bridge(tiny_bridge): + """Real Bridge hook wiring recovers full-span logits without assuming metric monotonicity.""" + ov = decompose_head(tiny_bridge, 0, 0, which=("OV",)).OV + assert ov is not None + rank = ov.V.shape[1] + partial_k = _block_aligned_ladder(ov)[0] + assert partial_k < rank + prompt = torch.tensor([[5, 63, 7, 9]]) + metric = lambda logits: float(logits[0, -1, 0] - logits[0, -1, 1]) + + partial = patch_along_directions( + tiny_bridge, ov, prompt, metric, keep=list(range(partial_k)), n_baseline=1 + ) + full = patch_along_directions( + tiny_bridge, ov, prompt, metric, keep=list(range(rank)), n_baseline=1, threshold=1e-4 + ) + assert abs(full.delta_metric) < 1e-4 + assert abs(partial.delta_metric) > 100 * abs(full.delta_metric)