File size: 18,334 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 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 | #!/usr/bin/env python3
"""
================================================================================
additive_test.py — 检验槽位栈的 gate 读出是否超出"加性残差可表达类"
================================================================================
一. 要证的命题
--------------
加性残差流里 consumer 读到的是 W(Σ_j α_j s_j) = Σ_j α_j·W s_j
展开成分块矩阵就是 [α_0 W, α_1 W, …, α_{k-1} W]
→ 每一行都是同一个 u 按同一组 α 复制 k 份
→ 全部 dff 个输出坐标共用一组混合比 α
槽位栈是 [W^(0), W^(1), …, W^(k-1)] 各块彼此无关
→ 每个输出坐标可以有自己的混合比
★ 所以"加性可表达"的精确含义 = W_g 沿块维度是秩1的。
这给出一个可判定的检验, 而不是靠聚合占比去猜。
★ 为何"聚合贡献占比偏离范数比例"证明不了:
1) 标准残差里 W 可以挑子空间, 把某个源投到零 —— 占比与范数比例无关是
加性残差本来就能做到的事;
2) 聚合占比是在全部输出坐标上平均的, 恰好把"逐坐标混合比不同"这个真正
的差别平均掉。两个模型可以聚合占比完全相同, 一个逐坐标高度分化,
一个完全均一。
二. 三个测试
------------
T1 逐坐标混合比离散度 纯权重分析。对每个输出坐标算它对各槽的行范数占比
p_m ∈ 单纯形; 加性预测所有 p_m 相同。测散度。
(这是"聚合占比"的正确版本: 看离散度不看均值)
T2 加性投影 沿块维度 SVD, 秩r 截断后量 VAL。r=1 即"投影到加性
可表达类"。★必须配同 Frobenius 幅度的随机扰动对照,
否则分不清"结构被破坏"与"权重被改动太多"。
T3 块置换 各块读的槽做错排。加性下各块只差标量, 置换近乎无损。
三. 判读
--------
T2 是决定性的: 若 VAL(秩1) ≫ VAL(同幅度随机扰动), 则被破坏的是结构本身,
加性残差确实表达不了当前的读出方式。若两者相当, 说明这些层的 gate 实际上
停留在加性可表达的范围内 —— 那是个真实的负结果, 该认。
T1/T3 是支持性证据: 它们能证明"分化了", 但不能单独证明"加性做不到"。
四. 用法
--------
python additive_test.py --ckpt /workspace/data/d1280_h10_l10_out/latest.pt \
--model delta_d1280_h10_l10.py --batches 20
SMOKE=1 python additive_test.py --model delta_d1280_h10_l10.py --batches 2
(无 --ckpt 时用随机初始化, 只验证脚本通路; 随机初始化下各块本就正交,
结论无意义)
================================================================================
"""
import argparse, importlib.util, math, os, sys
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
# ═══════════════════════ T1: 逐坐标混合比离散度 ═══════════════════════
def coordwise_mixing(W, k, d):
"""W: (dff, k*d)。返回 (每坐标混合比矩阵 (dff,k), 平均占比 (k,), 离散度)。
对输出坐标 m, 它对槽 j 的读出强度 = 该行在块 j 上的 L2 范数。
归一化成单纯形上的点 p_m。加性残差预测: 所有 p_m 完全相同。
离散度用平均总变差距离 (到全局均值) 表示, 单位 pp:
disp = mean_m (1/2)·Σ_j |p_m[j] - p̄[j]| × 100
0pp = 所有坐标混合比一致 (加性形态); 越大 = 逐坐标越分化。
"""
B = W.reshape(W.shape[0], k, d) # (dff, k, d) [:, j, :] = 块 j
s = B.norm(dim=-1) # (dff, k)
p = s / s.sum(-1, keepdim=True).clamp_min(1e-12)
pbar = p.mean(0)
disp = (0.5 * (p - pbar).abs().sum(-1)).mean().item() * 100
return p, pbar, disp
# ═══════════════════════ T2: 加性投影 (块维 SVD) ═══════════════════════
def block_spectrum(W, k, d):
"""把 k 个块各自拉直成向量, 求 k×k Gram 的特征谱。
加性可表达 ⟺ 这些向量共线 ⟺ 谱只有一个非零值。
返回 (特征值降序, 秩1能量占比)。
随机初始化基线: 各块独立随机 → 谱近似平坦 → 秩1占比 ≈ 1/k。
"""
V = W.reshape(W.shape[0], k, d).permute(1, 0, 2).reshape(k, -1).double()
G = V @ V.T
ev = torch.linalg.eigvalsh(G).flip(0).clamp_min(0)
return ev, (ev[0] / ev.sum().clamp_min(1e-30)).item()
def matrix_rank_matched(W, eps):
"""沿【普通矩阵秩】截断, 挑 rank 使 ‖ΔW‖_F 最接近 eps。
★ 这是比"同幅度随机扰动"严得多的对照。各向同性噪声在 dff×(k·d) 的大空间
里落到模型实际使用的子空间上的分量极小 —— 它匹配了幅度, 却没匹配"命中
率", 于是任何结构化删除都天然更疼, 与块结构无关。
本对照同样是结构化低秩删除、同样的 Frobenius 幅度, 唯一区别是沿哪个轴,
从而把"块维度"这一个因素单独隔离出来。
"""
U, S, Vh = torch.linalg.svd(W.double(), full_matrices=False)
tail = torch.flip(torch.cumsum(torch.flip(S**2, [0]), 0), [0]).sqrt() # ‖Δ‖(rank r)
r = int((tail - eps).abs().argmin().item())
r = max(1, min(r, len(S) - 1))
Wr = (U[:, :r] * S[:r]) @ Vh[:r]
return Wr.to(W.dtype), r, tail[r].item()
def random_subspace_keep(W, frac_keep, gen):
"""对照D: 只保留输入空间里一个【随机】子空间, 丢掉其余。
为何需要它: 对照B(普通矩阵秩截断)砍掉的是奇异值最小的方向 —— 同等能量下
最良性的一种删除, 只给出损伤下界。对照D 同样是结构化删除、同样的能量,
但方向随机挑, 不偏袒"最没用的"。它才是"一般性结构化损伤"的公允基线。
W P, P = G(GᵀG)⁻¹Gᵀ 投影到随机 q 维子空间, q = frac_keep · (k·d)。
各向同性下期望保留能量 ≈ q/(k·d), 故 frac_keep 直接取"块维秩1 保留的能量比"。
"""
kd = W.shape[1]
q = max(1, min(kd - 1, int(round(frac_keep * kd))))
G = torch.randn(kd, q, generator=gen).to(W.device).float()
A = W.float() @ G # (dff, q)
M = G.T @ G # (q, q)
B = torch.linalg.solve(M + 1e-6 * torch.eye(q, device=W.device), A.T).T
return (B @ G.T).to(W.dtype)
def rotated_block_rank1(W, k, d, gen):
"""对照C: 每块输入空间先乘一个随机正交阵 R_j, 做块维秩1, 再转回来。
结果形如 W^(j) = α_j · U R_jᵀ
· 仍是"秩1 + 全局标量 α" —— 逐坐标幅度混合比与加性残差同样受限
· 但每个槽被读的方向不同 —— 加性残差表达不了这一点
于是它把两个自由度拆开了:
对照C 也很疼 ⇒ 关键自由度是"逐坐标幅度混合比"
对照C 良性 ⇒ 关键自由度是"从各源读不同方向", 与幅度分配无关
⚠ 实测判定本对照【无效】, 保留仅供复现: 随机旋转等于"先把各槽内容打乱
再做加性合并", 是"加性残差 + 输入被搅烂", 严格劣于加性本身, 拿它对照
证明不了任何事。更根本的问题是"加性可表达" ⟺ "块维秩1" 是同一个集合
(V 秩1 ⟹ W^(j)=α_j U 强制成立), 无法固定秩而只改加性与否。
公允对照请用 random_subspace_keep (对照D)。
"""
dff = W.shape[0]
B = W.reshape(dff, k, d)
Rs, Bt = [], torch.empty_like(B)
for j in range(k):
Q, _ = torch.linalg.qr(torch.randn(d, d, generator=gen).to(W.device))
Rs.append(Q.to(W.dtype))
Bt[:, j, :] = B[:, j, :] @ Rs[j]
B1 = rank_truncate(Bt.reshape(dff, k * d), k, d, 1).reshape(dff, k, d)
out = torch.empty_like(B)
for j in range(k):
out[:, j, :] = B1[:, j, :] @ Rs[j].T
return out.reshape(dff, k * d)
def rank_truncate(W, k, d, r):
"""沿块维度做秩 r 截断。r=1 即投影到加性残差可表达类。"""
dff = W.shape[0]
V = W.reshape(dff, k, d).permute(1, 0, 2).reshape(k, -1).double()
U, S, Vh = torch.linalg.svd(V, full_matrices=False)
Vr = (U[:, :r] * S[:r]) @ Vh[:r]
return Vr.reshape(k, dff, d).permute(1, 0, 2).reshape(dff, k * d).to(W.dtype)
# ═══════════════════════ 评估 ═══════════════════════
@torch.no_grad()
def val_loss(m, model, loader, nb):
model.eval()
loader.reset()
t = 0.0
with m.amp_ctx():
for _ in range(nb):
x, y = loader.next_batch(m.device)
t += model(x, y)[1].item()
return t / nb
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--model", default="delta_d1280_h10_l10.py")
ap.add_argument("--ckpt", default=None)
ap.add_argument("--batches", type=int, default=20)
ap.add_argument("--ranks", default="1,2,3")
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--per-layer", action="store_true",
help="每次只对一层做块维秩1截断, 跳出饱和区并做逐层归因")
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()
if m.SMOKE:
loader = m.SyntheticLoader(m.MICRO_BATCH, m.SEQ_LEN, 2)
else:
loader = m.ShardLoader(m.DATA_DIR, "val", m.MICRO_BATCH, m.SEQ_LEN,
insert_p=0.0)
L, d, dff = model.L, model.d, m.DFF_
print("=" * 100)
print(f" 加性可表达性检验 {tag}")
print(f" d={d} dff={dff} L={L} VAL 用 {a.batches} 个 batch")
print("=" * 100)
base = val_loss(m, model, loader, a.batches)
print(f"\n 基线 VAL = {base:.4f}\n")
# ── T1 + 谱 ──
print(" 【T1】逐坐标混合比离散度 & 块维谱")
print(" " + "-" * 96)
print(f" {'层':<5}{'k':<4}{'离散度(pp)':<13}{'秩1能量占比':<14}"
f"{'随机基线1/k':<13}{'谱(前4, 归一化)'}")
for i in range(1, L): # ffn_0 只有 1 块, 平凡加性
W = model.ffns[i].Wg.weight.data.float()
k = i + 1
_, pbar, disp = coordwise_mixing(W, k, d)
ev, r1 = block_spectrum(W, k, d)
evn = (ev / ev.sum()).tolist()[:4]
print(f" f{i+1:<4}{k:<4}{disp:<13.1f}{r1:<14.4f}{1.0/k:<13.4f}"
+ " ".join(f"{v:.3f}" for v in evn))
print("\n 离散度 0pp = 所有输出坐标共用一组混合比 = 加性形态")
print(" 秩1能量占比 →1 = 各块共线 = 加性可表达; →1/k = 各块正交")
print(" ★ 随机初始化时秩1占比本就≈1/k, 故该列低不能单独当作学到了东西")
# ── T2 ──
ranks = [int(r) for r in a.ranks.split(",")]
print(f"\n 【T2】加性投影 (沿块维秩r截断) + 同幅度随机扰动对照")
print(" " + "-" * 96)
orig = {i: model.ffns[i].Wg.weight.data.clone() for i in range(L)}
g = torch.Generator(device="cpu").manual_seed(a.seed)
for r in ranks:
# 秩 r 截断
dfro = 0.0
for i in range(1, L):
W = orig[i].float()
Wr = rank_truncate(W, i + 1, d, r)
dfro += (Wr - W).pow(2).sum().item()
model.ffns[i].Wg.weight.data.copy_(Wr.to(orig[i].dtype))
lo = val_loss(m, model, loader, a.batches)
for i in range(L):
model.ffns[i].Wg.weight.data.copy_(orig[i])
# 对照: 同 Frobenius 幅度的随机扰动, 逐层按各自幅度分配
for i in range(1, L):
W = orig[i].float()
Wr = rank_truncate(W, i + 1, d, r)
eps = (Wr - W).norm()
n = torch.randn(W.shape, generator=g).to(W.device)
n = n / n.norm() * eps
model.ffns[i].Wg.weight.data.copy_((W + n).to(orig[i].dtype))
lc = val_loss(m, model, loader, a.batches)
for i in range(L):
model.ffns[i].Wg.weight.data.copy_(orig[i])
# 严对照: 同幅度的【普通矩阵秩】截断 (结构化删除, 只是换了个轴)
rs = []
for i in range(1, L):
W = orig[i].float()
eps = (rank_truncate(W, i + 1, d, r) - W).norm()
Wm, rr, _ = matrix_rank_matched(W, eps)
rs.append(rr)
model.ffns[i].Wg.weight.data.copy_(Wm.to(orig[i].dtype))
lm = val_loss(m, model, loader, a.batches)
for i in range(L):
model.ffns[i].Wg.weight.data.copy_(orig[i])
print(f" 块维秩{r}: VAL {lo:.4f} (Δ{lo-base:+.4f}) ‖ΔW‖_F={math.sqrt(dfro):.1f}")
print(f" 对照A 同幅度随机扰动 VAL {lc:.4f} (Δ{lc-base:+.4f}) [弱对照]")
print(f" 对照B 同幅度普通矩阵秩截断 VAL {lm:.4f} (Δ{lm-base:+.4f}) "
f"[严对照, matrix rank≈{int(np.mean(rs))}]")
if lo > lm + 0.05:
print(f" ⇒ 块维删除比同幅度普通秩删除更疼 (Δ差 {lo-lm:+.3f}), "
f"块结构被单独隔离出来了")
else:
print(f" ⇒ 两者相当, 损伤不能归因于块结构本身")
print("\n ★ 判据以【对照B】为准: 块维秩1 显著疼于同幅度普通矩阵秩截断")
print(" ⇒ 损伤来自块结构本身, 加性残差表达不了当前读出方式。")
print(" ★ 注意各截断若都落在'已损坏'区间(VAL≈ln(vocab)), 秩阶梯之间的排序")
print(" 是噪声, 只有'秩1 vs 对照'这一比较有效。")
print(" ★ 本测试证明的是: 该已训练模型重度使用了加性表达不了的成分。它")
print(" 没有证明'原生训练一个加性模型也达不到同样 loss' —— 后者需要重训")
print(" tied-block 基线, 且有严重参数量混淆。这是不重训能拿到的最强证据。")
# ── T2b: 逐层归因 ──
if a.per_layer:
print(f"\n 【T2b】逐层块维秩1截断 (一次只动一层, 跳出饱和区)")
print(" " + "-" * 96)
print(f" {'层':<6}{'k':<4}{'块维秩1(加性)':<16}{'对照B 最良性删除':<20}"
f"{'对照D 随机子空间':<20}{'‖Δ‖块':<9}{'‖Δ‖D':<9}{'保留能量'}")
rows = []
for i in range(1, L):
W = orig[i].float()
Wr = rank_truncate(W, i + 1, d, 1)
eps = (Wr - W).norm()
model.ffns[i].Wg.weight.data.copy_(Wr.to(orig[i].dtype))
lb = val_loss(m, model, loader, a.batches)
Wm, rr, _ = matrix_rank_matched(W, eps)
model.ffns[i].Wg.weight.data.copy_(Wm.to(orig[i].dtype))
lc2 = val_loss(m, model, loader, a.batches)
model.ffns[i].Wg.weight.data.copy_(orig[i])
fk = 1.0 - (eps / W.norm()) ** 2 # 块维秩1 保留的能量比
Wd_ = random_subspace_keep(W, float(fk), g)
epsd = (Wd_ - W).norm().item()
model.ffns[i].Wg.weight.data.copy_(Wd_.to(orig[i].dtype))
lc4 = val_loss(m, model, loader, a.batches)
model.ffns[i].Wg.weight.data.copy_(orig[i])
rows.append((i + 1, lb - base, lc2 - base, lc4 - base))
print(f" f{i+1:<5}{i+1:<4}{lb-base:<16.4f}{lc2-base:<20.4f}"
f"{lc4-base:<20.4f}{eps.item():<9.1f}{epsd:<9.1f}{fk*100:.0f}%")
bb = np.array([r[1] for r in rows]); dd = np.array([r[3] for r in rows])
print(f"\n 合计 ΔVAL: 块维秩1(加性) {bb.sum():.2f} | 对照B(最良性) "
f"{np.array([r[2] for r in rows]).sum():.2f} | 对照D(随机子空间) {dd.sum():.2f}")
print("\n ★ 判据看【对照D】: 同样丢掉同等能量, 但丢的方向随机挑。")
if bb.sum() > 1.5 * dd.sum():
print(" 块维秩1 明显更疼 ⇒ 损伤不是'丢了这么多能量'的一般后果,")
print(" 而是特定于'塌成加性形态'这个结构。结论成立。")
elif dd.sum() > 1.5 * bb.sum():
print(" 随机子空间反而更疼 ⇒ 加性形态其实保住了不少功能,")
print(" 之前对照B 得出的差距是假象。结论需推翻。")
else:
print(" 两者相当 ⇒ 损伤主要来自'丢了这么多能量', 与是否加性无关。")
print(" 这是真实的负结果, 该认。")
print(" ★ 对照B 砍的是奇异值最小的方向, 是同能量下最良性的删除, 只给")
print(" 损伤下界; 与它的差距被系统性放大, 倍数别当真。")
# ── T3 ──
print(f"\n 【T3】块置换 (错排, 各块改读别的槽)")
print(" " + "-" * 96)
rng = np.random.default_rng(a.seed)
for trial in range(3):
for i in range(1, L):
k = i + 1
while True: # 错排: 无不动点
perm = rng.permutation(k)
if not (perm == np.arange(k)).any():
break
W = orig[i].reshape(dff, k, d)
model.ffns[i].Wg.weight.data.copy_(
W[:, torch.from_numpy(perm.copy()), :].reshape(dff, k * d))
lp = val_loss(m, model, loader, a.batches)
for i in range(L):
model.ffns[i].Wg.weight.data.copy_(orig[i])
print(f" 错排 #{trial+1}: VAL {lp:.4f} (Δ{lp-base:+.4f})")
print("\n ★ 加性形态下各块只差一个标量, 置换应近乎无损。VAL 大幅上升")
print(" ⇒ 每块已专门化到自己那个槽。(支持性证据, 不能单独证明加性做不到)")
print("=" * 100)
if __name__ == "__main__":
main()
|