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()
|