File size: 12,085 Bytes
bb9d913
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
"""
Activation Patching β€” causal localization of behavior over model components.

Activation patching runs the model on a BASE prompt while splicing in cached
activations from a SOURCE prompt at one site (layer Γ— position, or layer Γ—
head), then measures how much a behavioral metric moves. Sweeping sites yields
a map of where the computation that distinguishes the two prompts lives.

The two directions answer different questions and can disagree; both are
exposed and neither is labelled "the circuit":

- **denoising** (patch clean β†’ corrupted run): does restoring this site
  SUFFICE to recover the clean behavior?
- **noising** (patch corrupted β†’ clean run): is this site NECESSARY β€” does
  corrupting it alone destroy the clean behavior?

Metric: logit difference between two single-token answers at the final
position (Wang et al. 2022, IOI). `normalized` rescales it so 0 = the base
run's own value and 1 = the source run's value; in denoising 1 means full
restoration, in noising 1 means full destruction. Deliberately NOT clamped:
values outside [0, 1] mean the patch overshot or backfired, which is the
surprising result worth seeing.

Like logit_lens / attention, this probes raw representations β€” prompts are
used as-is, with no chat template.
"""

from typing import Any, Callable, Dict, List, Optional

import torch
import torch.nn.functional as F

from model import get_model

VALID_DIRECTIONS = ("denoising", "noising")
VALID_COMPONENTS = ("resid_post", "head_z")


# -----------------------------------------------------------------------------
# Metric helpers
# -----------------------------------------------------------------------------

def logit_diff(final_logits: torch.Tensor, answer_id: int, baseline_id: int) -> float:
    """logits[answer] - logits[baseline] at one position. final_logits: [d_vocab]."""
    return float((final_logits[answer_id] - final_logits[baseline_id]).item())


def kl_divergence(p_logits: torch.Tensor, q_logits: torch.Tensor) -> float:
    """KL(P || Q) in nats between the next-token distributions of two logit vectors."""
    log_p = F.log_softmax(p_logits, dim=-1)
    log_q = F.log_softmax(q_logits, dim=-1)
    return float(torch.sum(log_p.exp() * (log_p - log_q)).item())


def normalized_recovery(patched: float, base: float, source: float) -> Optional[float]:
    """
    Rescale a patched metric: 0 = base run's value, 1 = source run's value.

    Returns None when |source - base| is numerically zero β€” the two runs don't
    disagree on the metric, so "fraction of the gap crossed" is undefined.
    """
    denom = source - base
    if abs(denom) < 1e-9:
        return None
    return (patched - base) / denom


# -----------------------------------------------------------------------------
# Core: token-level patching sweep
# -----------------------------------------------------------------------------

def _make_position_patch_hook(source_act: torch.Tensor, pos: int) -> Callable:
    """Hook that overwrites one position of the residual stream with `source_act`."""
    def hook_fn(activation, hook):
        # activation: [batch, seq_len, d_model]
        activation[:, pos, :] = source_act[pos, :]
        return activation
    return hook_fn


def _make_head_patch_hook(source_z: torch.Tensor, head: int) -> Callable:
    """Hook that overwrites one head's output (all positions) with `source_z`."""
    def hook_fn(activation, hook):
        # activation: [batch, seq_len, n_heads, d_head]
        activation[:, :, head, :] = source_z[:, head, :]
        return activation
    return hook_fn


def patch_grid(
    base_tokens: torch.Tensor,
    source_tokens: torch.Tensor,
    answer_id: int,
    baseline_id: int,
    component: str = "resid_post",
    layers: Optional[List[int]] = None,
    positions: Optional[List[int]] = None,
    heads: Optional[List[int]] = None,
    max_runs: Optional[int] = None,
) -> Dict[str, Any]:
    """
    Sweep single-site patches of `source_tokens` activations into `base_tokens` runs.

    The metric at every cell is logit_diff(answer_id, baseline_id) at the final
    position. Returns raw per-cell metrics plus both unpatched baselines;
    direction semantics (which prompt is base vs source) are the caller's.

    Raises ValueError on shape mismatch, out-of-range sites, or a sweep larger
    than `max_runs` (each cell is a full forward pass β€” callers exposing this
    publicly should cap it).
    """
    model = get_model()

    if base_tokens.shape != source_tokens.shape:
        raise ValueError(
            f"base and source prompts must tokenize to the same shape; got "
            f"{tuple(base_tokens.shape)} vs {tuple(source_tokens.shape)}. "
            f"Pick prompts that differ only in same-token-length spans."
        )
    seq_len = base_tokens.shape[1]

    if component not in VALID_COMPONENTS:
        raise ValueError(f"component must be one of {VALID_COMPONENTS}, got {component!r}")

    n_layers, n_heads = model.cfg.n_layers, model.cfg.n_heads
    layers = list(range(n_layers)) if layers is None else list(layers)
    for layer in layers:
        if not 0 <= layer < n_layers:
            raise ValueError(f"layer {layer} out of range for n_layers={n_layers}")

    if component == "resid_post":
        positions = list(range(seq_len)) if positions is None else [
            p if p >= 0 else seq_len + p for p in positions
        ]
        for pos in positions:
            if not 0 <= pos < seq_len:
                raise ValueError(f"position {pos} out of range for seq_len={seq_len}")
        cols = positions
    else:
        heads = list(range(n_heads)) if heads is None else list(heads)
        for head in heads:
            if not 0 <= head < n_heads:
                raise ValueError(f"head {head} out of range for n_heads={n_heads}")
        cols = heads

    n_runs = len(layers) * len(cols)
    if max_runs is not None and n_runs > max_runs:
        raise ValueError(
            f"sweep of {len(layers)} layers x {len(cols)} sites = {n_runs} forward "
            f"passes exceeds the cap of {max_runs}; restrict layers/positions/heads"
        )

    hook_of_layer = (
        (lambda layer: f"blocks.{layer}.hook_resid_post")
        if component == "resid_post"
        else (lambda layer: f"blocks.{layer}.attn.hook_z")
    )
    wanted = {hook_of_layer(layer) for layer in layers}

    with torch.no_grad():
        source_logits, source_cache = model.run_with_cache(
            source_tokens, names_filter=lambda name: name in wanted
        )
        base_logits = model(base_tokens)

    base_ld = logit_diff(base_logits[0, -1], answer_id, baseline_id)
    source_ld = logit_diff(source_logits[0, -1], answer_id, baseline_id)

    rows = []
    for layer in layers:
        hook_name = hook_of_layer(layer)
        source_act = source_cache[hook_name][0]  # [seq, d_model] or [seq, n_heads, d_head]
        cells = []
        for col in cols:
            if component == "resid_post":
                hook_fn = _make_position_patch_hook(source_act, col)
                cell_key = "position"
            else:
                hook_fn = _make_head_patch_hook(source_act, col)
                cell_key = "head"
            with torch.no_grad():
                patched_logits = model.run_with_hooks(
                    base_tokens, fwd_hooks=[(hook_name, hook_fn)]
                )
            final = patched_logits[0, -1]
            patched_ld = logit_diff(final, answer_id, baseline_id)
            cells.append({
                cell_key: col,
                "logit_diff": round(patched_ld, 6),
                "normalized": _round_opt(normalized_recovery(patched_ld, base_ld, source_ld)),
                "kl_from_base": round(kl_divergence(final, base_logits[0, -1]), 6),
                "kl_to_source": round(kl_divergence(final, source_logits[0, -1]), 6),
            })
        rows.append({"layer": layer, "cells": cells})

    return {
        "component": component,
        "layers": layers,
        ("positions" if component == "resid_post" else "heads"): cols,
        "base_logit_diff": round(base_ld, 6),
        "source_logit_diff": round(source_ld, 6),
        "grid": rows,
    }


def _round_opt(x: Optional[float]) -> Optional[float]:
    return None if x is None else round(x, 6)


# -----------------------------------------------------------------------------
# String-level wrapper (API surface)
# -----------------------------------------------------------------------------

def _resolve_single_token(answer: str) -> int:
    """Map an answer string to exactly one token id, or fail with a usable message."""
    model = get_model()
    try:
        return int(model.to_single_token(answer))
    except Exception:
        try:
            pieces = model.to_str_tokens(answer, prepend_bos=False)
            detail = f" (splits into {pieces})"
        except Exception:
            detail = ""
        raise ValueError(
            f"answer {answer!r} is not a single token for this tokenizer"
            f"{detail}; pick a single-token answer β€” for GPT-2 that usually "
            f"means a leading space, e.g. ' Paris'"
        )


def run_activation_patching(
    clean_prompt: str,
    corrupted_prompt: str,
    clean_answer: str,
    corrupted_answer: str,
    direction: str = "denoising",
    component: str = "resid_post",
    layers: Optional[List[int]] = None,
    positions: Optional[List[int]] = None,
    heads: Optional[List[int]] = None,
    max_runs: Optional[int] = None,
) -> Dict[str, Any]:
    """
    Full activation-patching sweep between a clean/corrupted prompt pair.

    The metric is always logit_diff = logits[clean_answer] - logits[corrupted_answer]
    at the final position, so `clean_logit_diff` should be positive and
    `corrupted_logit_diff` negative (or at least smaller) when the pair is
    well-formed; a note is attached when it isn't.

    direction:
        "denoising" β€” base = corrupted run, source = clean activations.
        "noising"   β€” base = clean run, source = corrupted activations.
    """
    if direction not in VALID_DIRECTIONS:
        raise ValueError(f"direction must be one of {VALID_DIRECTIONS}, got {direction!r}")

    model = get_model()
    clean_tokens = model.to_tokens(clean_prompt)
    corrupted_tokens = model.to_tokens(corrupted_prompt)
    answer_id = _resolve_single_token(clean_answer)
    baseline_id = _resolve_single_token(corrupted_answer)
    if answer_id == baseline_id:
        raise ValueError(
            f"clean_answer and corrupted_answer resolve to the same token id "
            f"({answer_id}); the logit-diff metric would be identically zero"
        )

    if direction == "denoising":
        base_tokens, source_tokens = corrupted_tokens, clean_tokens
    else:
        base_tokens, source_tokens = clean_tokens, corrupted_tokens

    result = patch_grid(
        base_tokens,
        source_tokens,
        answer_id,
        baseline_id,
        component=component,
        layers=layers,
        positions=positions,
        heads=heads,
        max_runs=max_runs,
    )

    # Re-express the direction-relative baselines in clean/corrupted terms.
    if direction == "denoising":
        corrupted_ld, clean_ld = result["base_logit_diff"], result["source_logit_diff"]
    else:
        clean_ld, corrupted_ld = result["base_logit_diff"], result["source_logit_diff"]

    notes = []
    if clean_ld <= corrupted_ld:
        notes.append(
            "clean prompt does not favor clean_answer over corrupted_answer "
            f"(clean logit_diff {clean_ld} <= corrupted {corrupted_ld}); "
            "check the answers aren't swapped"
        )

    return {
        "direction": direction,
        "clean_prompt": clean_prompt,
        "corrupted_prompt": corrupted_prompt,
        "clean_answer": clean_answer,
        "corrupted_answer": corrupted_answer,
        "tokens": model.to_str_tokens(base_tokens[0]),
        "clean_logit_diff": clean_ld,
        "corrupted_logit_diff": corrupted_ld,
        "notes": notes,
        **result,
    }