| |
| """ |
| ================================================================================ |
| 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 |
| 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: |
| 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() |
|
|