File size: 4,813 Bytes
ca6c0e5
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""Cross-step causal patching with correct KL divergence.

FIX #3 — KL direction was backwards in v1:
  v1: F.kl_div(patched.log(), baseline) = KL(baseline || patched)  ← WRONG
  v2: F.kl_div(baseline.log(), patched) = KL(patched || baseline)  ← CORRECT
  Correct interpretation: how much does patching source into target change
  the output distribution? KL(patched || baseline) measures divergence FROM
  the baseline TOWARDS the patched distribution.

FIX #6 — Patch length mismatch:
  v1 patched src[:, :min_len, :] into target — misaligns token semantics when
  the source prompt is shorter than the target (step N+1 includes N's history).
  v2 patches only the LAST TOKEN position: src[:, -1:, :].
  The last token represents the model's "current reasoning state" and is
  semantically comparable across steps regardless of prompt length.
"""
from __future__ import annotations

import numpy as np
import torch
import torch.nn.functional as F

from .storage import load_step_activations


def cross_step_patch(
    model,
    source_activation_path: str,
    target_prompt: str,
    patch_layer: int,
) -> dict:
    """Patch source step's last-token residual into target step's forward pass.

    Patches ONLY the last token position of the source into the corresponding
    last-token position in the target. This avoids the token-length mismatch
    problem: regardless of prompt length, the last token position represents
    the model's 'current state' at that reasoning step.

    KL divergence: KL(patched_dist || baseline_dist)
    Interpretation: how much does injecting source state redirect target's next token?
    """
    src = load_step_activations(source_activation_path)
    hook_name = f"blocks.{patch_layer}.hook_resid_post"

    if hook_name not in src:
        available = [k for k in src if "resid" in k]
        return {
            "patch_layer": patch_layer,
            "kl": 0.0,
            "significant": False,
            "token_changes": [],
            "error": (
                f"Layer {patch_layer} not in source activations. "
                f"Captured layers: {available}"
            ),
        }

    # Extract ONLY the last-token residual from source
    # src[hook_name]: float16 ndarray [1, seq_len, d_model] -> [1, 1, d_model]
    src_last_token = torch.tensor(
        src[hook_name].astype(np.float32)
    )[:, -1:, :]  # shape [1, 1, d_model]

    tgt_tokens = model.to_tokens(target_prompt)

    # Baseline: unpatched forward pass through target prompt
    with torch.no_grad():
        baseline_logits = model(tgt_tokens, return_type="logits")
    baseline_probs = F.softmax(baseline_logits[0, -1], dim=-1)

    # Patched: inject source's last-token residual into target's last-token position
    def patch_last_token(value, hook):
        # value: [1, seq_len, d_model] — replace ONLY the final token position
        value[:, -1:, :] = src_last_token
        return value

    with torch.no_grad():
        with model.hooks(fwd_hooks=[(hook_name, patch_last_token)]):
            patched_logits = model(tgt_tokens, return_type="logits")
    patched_probs = F.softmax(patched_logits[0, -1], dim=-1)

    # KL(patched || baseline) — correct direction (FIX #3)
    # F.kl_div(log_reference, target) = KL(target || reference)
    # So: F.kl_div(baseline.log(), patched) = KL(patched || baseline)
    kl = float(
        F.kl_div(
            (baseline_probs + 1e-10).log(),  # log of reference (baseline)
            patched_probs,                    # target distribution (patched)
            reduction="sum",
        ).item()
    )

    # Top-5 token probability shifts
    delta = (patched_probs - baseline_probs).abs()
    top_idx = delta.topk(5).indices.tolist()
    token_changes = []
    for idx in top_idx:
        tok = model.tokenizer.decode([idx])
        token_changes.append({
            "token": tok,
            "token_id": int(idx),
            "baseline_p": round(float(baseline_probs[idx].item()), 4),
            "patched_p": round(float(patched_probs[idx].item()), 4),
            "delta": round(float((patched_probs[idx] - baseline_probs[idx]).item()), 4),
        })

    significant = kl > 0.05

    return {
        "patch_layer": patch_layer,
        "kl": round(kl, 4),
        "significant": significant,
        "token_changes": token_changes,
        "interpretation": (
            f"Patching source step's last-token residual (layer {patch_layer}) "
            f"into the target step "
            f"{'significantly changes' if significant else 'does not significantly change'} "
            f"the output distribution (KL(patched||baseline)={kl:.4f}). "
            f"{'Causal influence from source step confirmed at this layer.' if significant else 'No strong causal influence at this layer.'}"
        ),
    }