delta-slotstack-236m / patch_test.py
Aurov's picture
initial upload
bbb9a9b verified
Raw
History Blame Contribute Delete
14.4 kB
#!/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()