Reference 对齐
Day 11-12 我们手写了 attention 与 decoder block,单测证明了它「causal 且 shape 对」。但「不泄露未来 + 形状对」不等于「数值算得对」——一个把 √d 缩放写错、或 LayerNorm 漏了除标准差的实现,照样能通过 causal 单测却给出错误的 logits。今天 Day 13 补上 B2 这块拼图最关键的一环:数值对齐——固定权重与种子,把我们
阶段: B2 · attention/decoder/KV + judge 校准(Day 11-20) 标签: #numerical-alignment #reference-test #floating-point #fixture
今日导引(由浅入深)
Day 11-12 我们手写了 attention 与 decoder block,单测证明了它「causal 且 shape 对」。但「不泄露未来 + 形状对」不等于「数值算得对」——一个把 √d 缩放写错、或 LayerNorm 漏了除标准差的实现,照样能通过 causal 单测却给出错误的 logits。今天 Day 13 补上 B2 这块拼图最关键的一环:数值对齐——固定权重与种子,把我们的输出与权威实现(PyTorch 或预存 reference 张量)逐层 element-wise 比对,取 max abs diff,选合理容差判对错。这是从「能跑」到「可信」的台阶,也是后续 Day 14 KV cache 加速「不许改变数值结果」的验收基线。今天的「最小可判定产出」:一份 max abs diff < 1e-4 的对齐报告(fp32 容差)。
1. 机理精读
为什么需要数值对齐:手写实现的 bug 常常是「正确得似是而非」——结果有限、形状对、甚至能训练出点东西,但与参考实现系统性偏差。光靠单元行为断言(如 causal)抓不住缩放因子、归一化、算子顺序这类「数值层面」的错。对齐法的核心:控制一切随机性(固定权重、固定输入、固定种子),让两个实现成为同一函数的两次求值,差异只可能来自浮点细节,于是任何超出浮点噪声的 diff 都暴露 bug。
行为测试 vs 数值对齐,互补不互替:Day 11-12 写的是「行为测试」——causal 性质、shape 守恒、确定性。它们能抓「结构性」错误(看了未来、维度错位),但抓不住「数值标定」错误。举三个能通过全部行为测试、却数值错的 bug:
- √d 缩放写成 √(2d):causal 仍成立、shape 不变、仍确定,但每个 logit 系统性偏小,softmax 偏平——只有与 reference 比 max abs diff 才暴露。
- LayerNorm 漏除标准差(只 center):行为测试若不查 variance 就放过,输出尺度全错。
- FFN 用了错的激活(如 GELU 当 ReLU):行为全对,数值偏差遍布每个元素。
所以「行为对 + 数值对」缺一不可,今天补的是后者。
容差怎么选:浮点不是实数,a+b+c 的累加顺序不同、matmul 用不用 fused-multiply-add、SIMD 归约树形状不同,都会让「数学上相等」的两条路径产生微小差异。经验法则:
- fp32:max abs diff 在
1e-4量级是合理的,差异主要来自算子顺序与浮点累加。 - bf16/fp16:尾数位少得多(bf16 仅 7 位尾数),需放宽到
1e-2量级,否则会把正常精度损失误判为 bug。 - fp8(e4m3/e5m2):尾数只剩 2-3 位,需放到
1e-1量级甚至看相对误差。 选容差是「既不放过 bug、又不误报噪声」的平衡——太严(如 fp32 下要求 1e-9)会被累加顺序噪声触发误报,太松(fp32 下 1e-1)会放过真 bug。
绝对容差 vs 相对容差:纯 abs(a-b) 在数值很大时会失真(1e6 量级的 1e-4 相对误差 abs diff 就有 100)。工程上常用 abs(a-b) <= atol + rtol·abs(b)(PyTorch allclose 的语义),同时给绝对和相对两个旋钮。本日教学因张量值都在 O(1) 量级,用纯 abs + 1e-4 即可。
对齐的前提是配置一致:两个实现必须用完全相同的超参与算子选择才能比——同样的 RoPE base、同样的 LayerNorm vs RMSNorm、同样的激活、同样的 head 划分。拿不同 RoPE/LN 实现直接对齐必然 diff 巨大,但那不是 bug 而是「比错了对象」。所以对齐前先核对配置表。
离线 fixture 替代在线 reference:没有本地 PyTorch 也能做对齐——预先用权威实现跑一次、把 reference 张量存成 committed JSON fixture,TS 侧只读 fixture 做 diff。好处:CI 可复现、无需 GPU/Python 环境、diff 管道随时能跑通。代价:fixture 一旦权重/配置变就要重生成。参考 Karpathy Zero-to-Hero(micrograd/makemore)的「与 PyTorch 逐层比对」方法论。
2. 推导 / 手算 / 代码走读
seed 计划的脚本 scripts/align-decoder.ts 当前仓库中不存在(已用 Glob 核实 scripts/{align-decoder,...}.ts 无匹配)——属 seed 标注的「待建」,下文是其设计走读而非既有代码走读,不编造函数名。
对齐管道的最小设计(待建):
- 导出固定权重:用
src/agent/transformer/transformer.ts的randomParams(d, heads, hidden, mulberry32(seed))生成确定性权重,序列化为 JSON(含Wq/Wk/Wv/Wo/g1/b1/g2/b2/W1/bff1/W2/bff2全部字段)。 - 生成 reference:把同一份权重 JSON 喂给 PyTorch 等价实现跑 forward,导出输出张量;或离线预存为 committed reference fixture(
reference.json)。 - TS 侧 diff:读 fixture + 用
decoderBlock(X, p)跑同输入,逐元素Math.abs(ours[i][j] - ref[i][j]),取max。 - 判定:
maxAbsDiff < 1e-4(fp32)→ 对齐通过;否则记录 max diff 与「出问题的层」(attention 输出?LN 后?FFN 后?逐层定位)。
手算一个微型对齐思路(验证 LayerNorm 这一最易错算子):对 [1,2,3,4],正确 LayerNorm(gamma=1,beta=0)应给零均值、单位方差。逐步算:
- mean =
(1+2+3+4)/4 = 2.5。 - var =
((1-2.5)²+(2-2.5)²+(3-2.5)²+(4-2.5)²)/4 = (2.25+0.25+0.25+2.25)/4 = 1.25。 inv = 1/√(1.25+eps) ≈ 0.8944。- 输出 =
[(1-2.5)·0.8944, (2-2.5)·0.8944, (3-2.5)·0.8944, (4-2.5)·0.8944] ≈ [-1.342, -0.447, 0.447, 1.342]。 - 验证:新 mean ≈ 0 ✓;新 var =
(1.342²+0.447²)·2/4 ≈ 1.0✓。
本仓 transformer.test.ts 已有等价断言(mean≈0,容差 1e-8;variance≈1,容差 1e-4),它正是「center-only 漏除标准差」这类 bug 的捕手——若实现忘了 1/√var(只减均值不除标准差),输出会是 [-1.5,-0.5,0.5,1.5],其 var = 1.25 ≠ 1,单测立刻红。这就是「数值对齐」精神的最小化身:用已知正确量化值卡住实现。
逐层定位的二分思路:若最终输出 diff 大,不要盲改。按 block 的数据流切点逐个比对——LN(x) 后、attention 后、attention 残差后、FFN-LN 后、FFN 后——第一个 diff 超容差的切点就是 bug 所在层。这把「整块错」缩小到「某算子错」,是对齐脚本最值钱的能力。
一份对齐报告该长什么样(待产出的目标格式):
=== decoder-block alignment report ===
config: d=8 heads=2 hidden=16 dtype=fp32 seed=3
input: [4,8] (committed fixture)
layer-by-layer max abs diff vs reference:
LN(x) : 3.1e-7 OK
attn out : 8.7e-6 OK
attn residual : 8.9e-6 OK
FFN-LN out : 1.2e-6 OK
FFN out : 2.4e-5 OK
final maxAbsDiff : 2.4e-5 < 1e-4 => ALIGNED ✓
报告的两个要素:(1) 逐层 diff(定位用),(2) 最终判定(< 1e-4 => ALIGNED)。上面的数字是示意格式而非实测——seed 标注本日「待建/待跑」,真实 diff 待脚本与 reference 就位后产出。
「逐层」要比哪些切点(与本仓 decoderBlock 数据流对应):LN(x) → attention 输出 → attention 残差(x1)→ FFN-LN → FFN 输出。任一切点先超容差,bug 就在该算子。
3. 今日实战
seed:「写 scripts/align-decoder.ts:导出固定权重 JSON,用同权重在 PyTorch(或离线 reference 张量)跑 forward,将 decoderBlock.ts 输出与 reference 做 element-wise diff,记录 max abs diff 与出问题的层。」可执行步骤(待建脚本):
- 新建
scripts/align-decoder.ts,import { decoderBlock, randomParams } from '../src/agent/transformer/transformer'(注意实际函数在transformer.ts,非 seed 名义的decoderBlock.ts)。 - 用固定
mulberry32(seed)生成权重并JSON.stringify落盘agent-evals/align/params.json。 - 若有本地 PyTorch:写等价 nn 实现读这份 JSON 跑 forward 存
reference.json;否则预存 committed fixture。 - TS 侧读 reference,对同输入跑
decoderBlock,算maxAbsDiff,console.logmax diff + 触发层;maxAbsDiff < 1e-4则打印「ALIGNED」。 - 先在无 reference 时用 self-fixture(自己跑两次)跑通 diff 管道,再接真实 PyTorch reference。
self-fixture 自检的价值(即便没有外部 reference):
- 跑两次同输入应得
maxAbsDiff == 0——验证实现确定性(本仓decoderBlock已有确定性单测佐证)。 - 故意把 √d 缩放改错再跑,确认 diff 管道能报红——验证对齐脚本本身没写成「永远绿」的假阳性。
- 这两步保证:当真实 PyTorch reference 接入时,diff 管道是可信的测量工具,而非摆设。
为什么离线 fixture 是务实选择:本仓的工程纪律是「纯 TS 确定性可测试、无 GPU/无 key」。在线连 PyTorch 做对齐违背这条——CI 跑不了、复现要装 Python 环境。预存 committed reference fixture(一次性在权威实现上生成、存成 JSON 进仓)则让对齐 diff 管道变成纯 TS 单测:读 fixture、跑 decoderBlock、比 diff,全程无外部依赖。代价是 fixture 与配置耦合,变更需重生成(见误区 6)。这与 Day 5 的 eval、Day 15 的采样「math 可离线、真实调用才需 key」是同一套离线优先哲学。
4. 今日实测 / 产出
- 状态(按 seed):seed 标注「待建/待跑」。目标产出「max abs diff < 1e-4 的对齐报告」。
- 诚实标注:
scripts/align-decoder.ts尚未创建(Glob 核实不存在),对齐报告未产出——保持「待建/待跑」,不升级为已完成。 - 可先跑通的部分:seed 明确「若无本地 PyTorch,可预存 reference fixture(committed JSON)做 TS 侧比对,先跑通 diff 管道」。被对齐的
decoderBlock本身已存在且单测绿(transformer.test.ts),意味着对齐脚本一旦建成、reference 一旦就位,diff 管道即可运行——但当前对齐报告仍属待产出。
5. 常见误区 / 陷阱
- 容差和 dtype 不匹配:fp32 用 1e-4,bf16/fp16 需放宽到 1e-2,fp8 更松。在 bf16 reference 上要求 1e-4 会被精度噪声误报。
- 拿不同配置的实现对齐:不同 RoPE base / LN 类型 / 激活直接比,diff 巨大但非 bug。先确认两侧超参一字不差。
- 不控制随机性就比:权重/输入/dropout/种子任一不固定,diff 来自随机性而非实现差异,对齐无意义。对齐脚本里务必关 dropout、设 eval 模式。
- 只比最终输出、不逐层定位:最终 diff 大时无法知道 bug 在 attention 还是 FFN。应逐层(attention 后 / LN 后 / FFN 后)记录,二分定位。
- 纯用绝对容差:张量值很大时 abs diff 失真,应同时给相对容差(
atol + rtol·abs(ref))。 - fixture 与代码漂移:committed reference fixture 在权重/配置变更后必须重生成,否则对齐的是「旧实现」,假绿。
6.5 对齐之于整条能力曲线的位置
今天这一步看似工程琐碎,却是 B2→B12 一切「复现 SOTA」工作的通行证:
- B12 复现 FlashAttention/MLA 时,验证「我写的高效 kernel 与朴素版数值一致」靠的就是今天的对齐法(只是容差按 bf16/fp8 放宽)。
- Day 14 KV cache 的「加速不改数值」是对齐法的一个特例(reference = cache-off 版本)。
- 任何算子移植/量化(fp32→fp8)的验收都是「对齐 + 容差」。
所以把今天的「fix-seed → diff → 容差」这条管道吃透,比记住任何单个数字都重要——它是可复用的验证方法论。
6. 学习资源(每条带 YYYY-MM)
- Karpathy《Neural Networks: Zero to Hero》(micrograd/makemore 系列, 2022-08 起)——seed 指定的「与 PyTorch 逐层对齐」方法论来源。
- PyTorch
torch.allclose/rtol,atol文档 (持续更新, 2024)——容差判定的标准接口语义参照。 - David Goldberg《What Every Computer Scientist Should Know About Floating-Point Arithmetic》(1991-03, ACM Computing Surveys)——浮点累加误差的经典底层依据。
- Karpathy nanoGPT (2023-01 起)——被对齐的 decoder 参考实现范式。
- 本仓
src/agent/__tests__/transformer.test.ts(AICAP-180, 2026-06)——LayerNorm「零均值单位方差」断言即「用已知正确值卡实现」的对齐精神最小化身。
SOTA检查 (2026-06 更新)
- 当前主流:数值对齐(fix weights/seed → element-wise diff → 容差判定)是恒定工程实践,无过时风险,所有严肃的内核/算子移植都靠它验收。
- 是否仍 SOTA:方法本身不被「替代」,只随 dtype 演进调容差——2026 训练/推理大量用 bf16/fp8/fp4,对齐容差需相应放宽(fp8 比 fp16 更松)。一个实用对齐容差速查:fp32→
1e-4、bf16/fp16→1e-2、fp8(e4m3)→~1e-1且看相对误差、fp4/MXFP→主要看相对误差与分布统计(逐元素已不可靠)。低精度下「对齐」更像「分布一致性检查」而非「逐元素相等」。 - 过时黑名单:避免拿不同 RoPE/LN 实现直接对齐(先确认配置一致);避免在低精度 reference 上套 fp32 容差(误报)。
- 下次复查点:B2 末(Day 20);当 B12 引入 FlashAttention/MLA 真实实现做对齐时,按其 dtype(bf16/fp8)重设容差并复验。
衔接
- 昨天:Day 12 — Multi-head + decoder block 结构(手写完整 decoder block,shape 对、causal)。
- 今天:用固定权重 + 数值 diff 验证这块手写 block「算得对」,确立 max abs diff < 1e-4 的对齐基线。
- 明天:Day 14 — KV cache 原理(在已验证正确的 block 上做加速,且加速不许改变数值结果)。