File size: 14,413 Bytes
bbb9a9b
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
#!/usr/bin/env python3
"""
================================================================================
patch_test.py — 定向操控: 把"必要性"升级成"因果"
================================================================================

一. 为什么消融不够
------------------
head_test.py 的 C 表只能说"挖掉 L2h8, b128 就崩"。这证明它是【必需的一环】,
不证明【答案就存在它里面】—— 它也可能只是给真正的检索头供货的上游。

定向操控回答后者: 造两条只差一个 token 的句子 (A 的答案是"猫", B 是"狗"),
跑 B, 但把某个头在探针位置的输出【换成 A 的】。若 B 的预测跟着翻成"猫",
则答案确实在这个头里, 且是因果的。

★ 这一步是本架构独有的。因为 (a) 没有 Wo, head h 永远物理占据
  [h·hd,(h+1)·hd); (b) 槽写一次不覆盖 —— 所以"替换第 k 层第 h 头的输出"
  是一次精确赋值。标准 transformer 里 Wo 一乘、残差一加, 得先做近似的
  "解叠加", 而近似准不准本身还得另外论证。

二. 最小对
----------
    序列 = c[0..p-1] + [探针] + c[p..]     探针插在 c[p-1] 之后
    prev 的答案 = c[p-1];  b_k 的答案 = c[p-k]
  A 与 B 逐位相同, 只有答案那一位不同 —— 其余上下文完全一致, 所以任何
  预测差异都只能来自那一位。

三. 指标 (标准的归一化 patching 效应)
--------------------------------------
    d = logit(A的答案) - logit(B的答案)   在探针位置上
    d_B    干净跑 B          (负, B 偏好自己的答案)
    d_A    干净跑 A          (正)
    d_patch  跑 B 但把某头换成 A 的
    效应 = (d_patch - d_B) / (d_A - d_B)
      1.0 = 完全翻转成 A 的答案, 该头独自携带答案
      0.0 = 毫无影响
      中间值 = 部分携带 (答案分散在多个头里)
    可以 >1 或 <0, 属于过冲/反向, 少量出现是正常噪声。

四. 用法
--------
    python patch_test.py --model delta_d1280_h10_l10.py \\
        --ckpt /workspace/data/d1280_h10_l10_out/latest.pt \\
        --probe prev --pairs 32
    --probe 可选 prev / b16 / b128 / b512   (答案位置确定的探针才能造最小对)
================================================================================
"""
import argparse, importlib.util
import numpy as np
import torch


def load_module(path):
    spec = importlib.util.spec_from_file_location("deltamod", path)
    m = importlib.util.module_from_spec(spec)
    spec.loader.exec_module(m)
    return m


def build_pairs(m, n_pairs, T, p, probe, back, rng):
    """造 n_pairs 组最小对。返回 (idx_A, idx_B, ansA, ansB), 探针都在位置 p。

    答案位在序列中的下标 = p - back  (prev: back=1; b_k: back=k)
    A 与 B 只有这一位不同, 其余完全一致。
    """
    tok = m._PROBE_TOK[probe]
    if m.SMOKE:
        src = rng.integers(0, 50257, size=(n_pairs, T + 8))
    else:
        ld = m.ShardLoader(m.DATA_DIR, "val", 1, T + 8, insert_p=0.0)
        rows = []
        while len(rows) < n_pairs:
            need = ld.B * ld.row
            buf = np.asarray(ld.tokens[ld.pos:ld.pos + need], dtype=np.int64)
            ld.pos += need
            if len(buf) < T + 8:
                ld.pos = 0
                continue
            rows.append(buf[:T + 8])
        src = np.stack(rows)
    src = np.clip(src, 0, 50256)

    A = np.empty((n_pairs, T), dtype=np.int64)
    B = np.empty((n_pairs, T), dtype=np.int64)
    ai = p - back                      # 答案在序列中的下标
    for i in range(n_pairs):
        c = src[i]
        seq = np.concatenate([c[:p], [tok], c[p:T - 1]])
        A[i] = seq
        b = seq.copy()
        alt = int(rng.integers(0, 50257))
        while alt == int(seq[ai]):
            alt = int(rng.integers(0, 50257))
        b[ai] = alt
        B[i] = b
    return (torch.from_numpy(A), torch.from_numpy(B),
            torch.from_numpy(A[:, ai].copy()), torch.from_numpy(B[:, ai].copy()))


@torch.no_grad()
def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--model", default="delta_d1280_h10_l10.py")
    ap.add_argument("--ckpt", default=None)
    ap.add_argument("--probe", default="prev", choices=["prev", "b16", "b128", "b512"])
    ap.add_argument("--pairs", type=int, default=32)
    ap.add_argument("--seqlen", type=int, default=None)
    ap.add_argument("--seed", type=int, default=0)
    ap.add_argument("--direct", action="store_true",
                    help="额外跑'只重算最后一个ffn'的直接效应, 与全量重跑对照")
    ap.add_argument("--joint", type=int, default=0,
                    help="贪心累加前 N 个头做联合替换, 给出'要几个头才能完全控制答案'的曲线")
    a = ap.parse_args()

    m = load_module(a.model)
    model = m.DeltaLM().to(m.device)
    tag = "随机初始化 (仅验证通路)"
    if a.ckpt:
        ck = torch.load(a.ckpt, map_location=m.device, weights_only=False)
        model.load_state_dict(ck["model"])
        tag = f"{a.ckpt}  step={ck.get('step','?')}"
    model.eval()
    L, H = model.L, model.H
    hd = model.d // H

    back = 1 if a.probe == "prev" else m._BACK_OFF[a.probe]
    T = a.seqlen or min(m.SEQ_LEN, 256 if m.SMOKE else 1024)
    p = T - 8                                     # 探针靠后, 保证 back 不越界
    assert p - back >= 0, f"序列太短, {a.probe} 需要 T > {back + 8}"

    rng = np.random.default_rng(a.seed)
    A, B, ansA, ansB = build_pairs(m, a.pairs, T, p, a.probe, back, rng)
    A, B = A.to(m.device), B.to(m.device)
    ansA, ansB = ansA.to(m.device), ansB.to(m.device)

    print("=" * 100)
    print(f"  定向操控 (activation patching)   {tag}")
    print(f"  探针={a.probe}  往回={back}  最小对={a.pairs}组  T={T}  探针位置={p}")
    print("=" * 100)

    def dgap(logits):
        """在探针位置上: logit(A的答案) - logit(B的答案)"""
        lg = logits[:, p, :].float()
        return (lg.gather(1, ansA[:, None]) - lg.gather(1, ansB[:, None])).squeeze(1)

    def argmax_rate(logits):
        """★ 这才是'模型真的改口了吗'。上面的 dgap 只比较两个候选词之间的分差,
        不保证在全部 VOCAB 个词里最高分的就是 A 的答案 —— 可能冒出第三方词。
        返回 (吐出A答案的比例, 吐出B答案的比例, 吐出其他词的比例)。"""
        am = logits[:, p, :].float().argmax(-1)
        ra = (am == ansA).float().mean().item()
        rb = (am == ansB).float().mean().item()
        return ra, rb, 1.0 - ra - rb

    with m.amp_ctx():
        dA = dgap(model.intervene(A, []))
        dB = dgap(model.intervene(B, []))
        _, slA = model.run(A, keep_slots=True)
        _, slB = model.run(B, keep_slots=True)

    denom = (dA - dB).mean().item()
    with m.amp_ctx():
        rA = argmax_rate(model.intervene(A, []))
        rB = argmax_rate(model.intervene(B, []))
    print(f"\n  干净基线: d_A={dA.mean():+.3f}  d_B={dB.mean():+.3f}  "
          f"可翻转区间={denom:.3f}")
    print(f"  实际吐出的词: 跑A时 A答案{rA[0]*100:.0f}% / 其他{rA[2]*100:.0f}%"
          f"  |  跑B时 B答案{rB[1]*100:.0f}% / A答案{rB[0]*100:.0f}% / 其他{rB[2]*100:.0f}%")
    if rB[1] < 0.5:
        print("  ⚠ 干净跑 B 时模型本来就答不对, 翻转实验无意义。等训练更久。")
    if denom < 1.0:
        print("  ⚠ 区间太小, 模型可能没学会这个探针, 结果无意义。换个探针或等训练更久。")

    def direct_logits(slots):
        """只重算最后一个 ffn + norm_f + 选词。前面的 att/ffn 一律不动。

        ★ 这一步标准架构给不了。那里第 8 层的注意力输出被 Wo 乘过、加进残差流、
          又被后面两层反复读写和 LayerNorm —— 到最后一层时那个向量在状态里已经
          不存在了, 你没有一个"槽8"可以改, 只能从第 8 层往后整个重跑。
          这里槽写一次不覆盖, 槽8 在 f10 眼里就是个输入, 与它怎么来的无关。
        ★ 语义也不同: 只重算 f10 = 【直接效应】(只走 f10 这一条路到 logits);
          全量重跑 = 直接 + 间接 (还包括 改槽k → 影响 att(k+1..L) → 再影响 f10)。
        """
        final = model.ffns[L - 1](slots[:L], slots[L])
        return model.head(model.norm_f(final)).float()

    def patched_slots(k, h):
        sl = slice(h * hd, (h + 1) * hd)
        out = list(slB)
        t = slB[k].clone()
        t[:, p, sl] = slA[k][:, p, sl]
        out[k] = t
        return out

    grid = np.zeros((L, H))
    gridd = np.zeros((L, H)) if a.direct else None
    with m.amp_ctx():
        if a.direct:                      # 自检: 用干净槽重算 f10 应复现干净 logits
            chk = (dgap(direct_logits(list(slB))) - dB).abs().max().item()
            print(f"  [自检] 干净槽重算 f10 与全量前向的最大偏差 {chk:.2e} (应≈0)")
        for k in range(1, L + 1):
            for h in range(H):
                sl = slice(h * hd, (h + 1) * hd)
                val = slB[k][..., sl].clone()
                val[:, p, :] = slA[k][:, p, sl]        # 只换探针那一个位置
                dp = dgap(model.intervene(B, [(k, h, val)]))
                grid[k-1, h] = ((dp - dB).mean() / (dA - dB).mean()).item()
                if a.direct:
                    dd = dgap(direct_logits(patched_slots(k, h)))
                    gridd[k-1, h] = ((dd - dB).mean() / (dA - dB).mean()).item()

    print(f"\n  归一化操控效应 (1.0=完全翻转成A的答案, 0=无影响)")
    print("  " + "-" * 96)
    print("  " + "层\\头".ljust(8) + "".join(f"h{h+1:<7}" for h in range(H)))
    for k in range(L):
        print(f"  L{k+1:<7}" + "".join(f"{grid[k,h]:<8.3f}" for h in range(H)))

    flat = sorted(((grid[i, j], i + 1, j + 1) for i in range(L) for j in range(H)),
                  reverse=True)
    if a.direct:
        print(f"\n  【直接效应】只重算最后一个 ffn (前面的 att/ffn 一律不动)")
        print("  " + "-" * 96)
        print("  " + "层\\头".ljust(8) + "".join(f"h{h+1:<7}" for h in range(H)))
        for k in range(L):
            print(f"  L{k+1:<7}" + "".join(f"{gridd[k,h]:<8.3f}" for h in range(H)))
        print(f"\n  【对照】总效应 vs 直接效应   (直接/总 = 该头的输出被选词直接用掉的比例)")
        print("  " + "-" * 96)
        print(f"  {'头':<10}{'总效应':<12}{'直接效应':<12}{'直接占比':<12}{'判读'}")
        for v, k, h in flat[:8]:
            dv = gridd[k-1, h-1]
            r = dv / v if abs(v) > 1e-3 else float("nan")
            note = ("几乎全是直接: 输出被选词直接用" if r > 0.8 else
                    "几乎全是间接: 主要喂给下游 att" if r < 0.2 else
                    "直接+间接各占一部分")
            print(f"  L{k}h{h:<7}{v:<12.3f}{dv:<12.3f}{r:<12.2f}{note}")
        print(f"\n  合计: 总 {grid.sum():.3f} | 直接 {gridd.sum():.3f}"
              f" | 直接占比 {gridd.sum()/max(grid.sum(),1e-6):.2f}")
        print("  ★ L{}(最后一层)的直接效应必然等于总效应 —— 它后面只剩 f10, "
              "没有间接路径可走。".format(L))
        print("    这是本对照的天然自检: 最后一行两张表应逐位相同。")

    print(f"\n  最强的 5 个: " + " | ".join(f"L{k}h{h} {v:.3f}" for v, k, h in flat[:5]))
    print(f"  合计效应 {grid.sum():.3f}   最大单头 {flat[0][0]:.3f} (L{flat[0][1]}h{flat[0][2]})")
    if flat[0][0] > 0.5:
        print("  ⇒ 单个头就能把答案改掉一半以上: 答案在该头里, 因果定位成立。")
    elif flat[0][0] > 0.15:
        print("  ⇒ 部分携带: 答案分散在几个头里, 需要多头联合替换才能完全翻转。")
    else:
        print("  ⇒ 无单头因果定位: 答案不在任何单个头的输出里 (可能在 gate/up 的")
        print("     组合里, 或分散得太开)。这是真实的负结果。")
    if a.joint > 0:
        print(f"\n  联合替换 (贪心累加, 每步加入当前边际效应最大的头)")
        print("  " + "-" * 96)
        chosen, pool = [], [(i + 1, j + 1) for i in range(L) for j in range(H)]
        with m.amp_ctx():
            for nstep in range(min(a.joint, L * H)):
                best = None
                for (k, h) in pool:
                    pts = []
                    for (kk, hh) in chosen + [(k, h)]:
                        sl = slice((hh - 1) * hd, hh * hd)
                        v = slB[kk][..., sl].clone()
                        v[:, p, :] = slA[kk][:, p, sl]
                        pts.append((kk, hh - 1, v))
                    e = ((dgap(model.intervene(B, pts)) - dB).mean()
                         / (dA - dB).mean()).item()
                    if best is None or e > best[0]:
                        best = (e, k, h)
                e, k, h = best
                chosen.append((k, h))
                pool.remove((k, h))
                pts = []
                for (kk, hh) in chosen:
                    sl2 = slice((hh - 1) * hd, hh * hd)
                    v2 = slB[kk][..., sl2].clone()
                    v2[:, p, :] = slA[kk][:, p, sl2]
                    pts.append((kk, hh - 1, v2))
                ra, rb, ro = argmax_rate(model.intervene(B, pts))
                print(f"  +{nstep+1:<3} L{k}h{h:<4} 分差效应 {e:.3f}  |  "
                      f"★实际吐出: A答案 {ra*100:.0f}%  B答案 {rb*100:.0f}%  其他 {ro*100:.0f}%")
                if e > 0.95:
                    break
        print("\n  ★ 看【实际吐出】那一列, 不是看分差效应。分差只说明两个候选词之间")
        print("    的相对位置变了, 模型最终吐哪个词要看 argmax 在全 5 万词上的结果。")
        print("  ★ 'A答案' 比例从 0% 爬到接近干净跑A的水平 = 真的改口了。")
        print("    若分差效应很高但 A答案 比例上不去 ⇒ 冒出了第三方词, 说明替换")
        print("    破坏了模型状态而非定向改写, 结论要打折。")

    print("\n  ★ 与 head_test 的 C 表对照: 消融排名(必要性) 与 操控排名(携带性)")
    print("    若不一致, 说明'挖掉就崩'的头是上游供货者, 而非答案的持有者。")
    print("=" * 100)


if __name__ == "__main__":
    main()