这篇文章用于沉淀和分析阅读苏神(苏剑林)《Attention Residuals 回忆录》 时涉及的相关知识。内容会随阅读逐步增加,每一章集中处理一个需要补充推导或实现背景的部分。
RMSNorm 的数学原理
RMSNorm 用输入向量在特征维度上的二阶原点矩衡量整体幅度,再按这个幅度缩放整个向量。
设一个 token 的隐藏状态为
x=(x1,x2,…,xd)∈Rd.
先定义带数值稳定项的均方根
rε(x)=d1i=1∑dxi2+ε,
再对所有维度使用同一个标量完成归一化,并加入可学习的逐维缩放参数 γ:
RMSNorm(x)=γ⊙rε(x)x.
这里的均值发生在同一个 token 的 d 个特征维度上。它是对当前向量的确定性计算,无需估计数据集或 batch 的统计量。RMSNorm 原论文 将其作为 LayerNorm 的简化形式,保留重缩放不变性并省去均值中心化。
从均方根到欧氏范数
暂时令 ε=0,并假设 x=0。均方根可以直接写成欧氏范数:
r0(x)=d1i=1∑dxi2=d∥x∥2.
记不含可学习参数的归一化结果为 x~,则
x~=r0(x)x=d∥x∥2x.
x/∥x∥2 只保留输入方向,前面的 d 将长度设为固定值。直接计算可得
∥x~∥22=d1∑j=1dxj2∑i=1dxi2=d,
因此
∥x~∥2=d.
在 ε=0 的理想条件下,RMSNorm 的归一化部分沿径向把非零输入映射到半径为 d 的超球面。这个映射保留方向并移除长度。
目标长度取 d,等价于让归一化结果的平均平方值为 1:
d1i=1∑dx~i2=1.
因此 RMSNorm 控制的是整个隐藏状态的平均平方幅度。单个坐标仍然可以远大于 1,各个维度也不会分别获得单位方差。
二阶原点矩与方差
RMSNorm 和 LayerNorm 的差异可以从二阶原点矩与方差的关系看出。定义当前向量在特征维度上的均值和方差:
μ=d1i=1∑dxi,σ2=d1i=1∑d(xi−μ)2.
二阶原点矩为
m2=d1i=1∑dxi2.
展开方差后得到
m2=σ2+μ2.
RMSNorm 使用 m2,所以均值偏移 μ2 也会进入归一化尺度。LayerNorm 先减去均值,再使用中心二阶矩 σ2。忽略可学习参数时,两者分别为
xRMS=σ2+μ2+εx,
xLN=σ2+εx−μ1.
其中,1=(1,1,…,1)T∈Rd 是 d 维全一列向量,因此 μ1 表示每个坐标都取值为 μ 的向量。
当 μ=0、两者采用相同的 ε 和缩放参数,并且 LayerNorm 的偏置为零时,两种归一化给出相同结果。当 μ 只接近零时,它们使用的尺度接近,RMSNorm 仍会保留沿 1 方向的分量。
这个差异也有直接的几何表示。令
P=I−d111T,
则 Px=x−μ1。LayerNorm 先把输入投影到与 1 正交的均值为零子空间,再在该子空间内归一化;RMSNorm 直接在完整的 Rd 中按长度归一化。LayerNorm 原论文 给出的定义同时包含中心化、缩放以及归一化后的可学习增益和偏置。
| 性质 | RMSNorm | LayerNorm |
|---|
| 归一化统计量 | 二阶原点矩 m2 | 中心二阶矩 σ2 |
| 是否减去特征均值 | 否 | 是 |
| 正比例尺度变化 | 在 ε=0 时消除 | 在 ε=0 时消除 |
| 整体平移 x+c1 | 会改变结果 | 归一化结果不变 |
| 几何作用 | 在完整空间中按半径归一化 | 先进入均值为零子空间,再按半径归一化 |
尺度变化与平移
尺度不敏感性是 RMSNorm 在 AttnRes 推导中承担的关键性质。令 a>0,则
rε(ax)=a2m2+ε=am2+a2ε.
当 ε=0 时,正比例缩放会被精确消除:
RMSNorm(ax)=RMSNorm(x),a>0.
实际实现采用 ε>0,此时等式只近似成立。当 a2m2≫ε 时,稳定项的影响很小。对于 a<0,在 ε=0 时有
RMSNorm(ax)=−RMSNorm(x),
因为输入方向发生了符号翻转。
RMSNorm 对整体平移没有这项不变性。令 x′=x+c1,其二阶原点矩变为
d1i=1∑d(xi+c)2=m2+2cμ+c2.
LayerNorm 会在中心化时消除新增的 c1,所以整体平移不会改变其归一化结果。RMSNorm 保留这部分均值信息,并让它参与尺度计算。
梯度沿切向传播
为了观察归一化怎样改变梯度,先去掉 γ,记
r=d1xTx+ε,y=rx.
它的 Jacobian 为
∂x∂y=r1I−dr31xxT.
如果上游梯度为 g=∂L/∂y,输入梯度就是
∂x∂L=r1(g−ydyTg).
第一项按输入 RMS 统一调整梯度尺度,第二项减去与输入径向方向相关的分量。当 ε=0 时,Jacobian 在径向上的作用为
∂x∂yx=0.
这与正比例尺度不变性一致:沿 x 方向只改变长度,归一化结果保持不变。对任意满足 xTv=0 的切向量 v,Jacobian 的作用为 v/r。因此在理想条件下,归一化核心保留切向变化并消除径向变化。
加入 ε 后,径向特征值变为 ε/r3,径向梯度会被显著减小,但不会严格归零。加入 γ 后,上述分析仍适用于归一化核心;反向传播时,上游梯度会先逐维乘以 γ。
ε 与 γ 改变了什么
ε 在输入接近零向量时保持分母为正。它也会让输出长度略低于 d:
∥x~∥22=∥x∥22/d+ε∥x∥22=∥x∥22+dεd∥x∥22<d.
只有在 ∥x∥22≫dε 时,输出长度才接近 d。
γ 为每个特征维度恢复可学习的尺度。归一化核心在理想条件下把输入映射到超球面,逐维乘以非均匀的 γ 后,这个超球面的像成为轴对齐的椭球面。模型由此可以调整各个特征维度的有效范围。
这些性质共同提供整体尺度控制:后续层看到的输入 RMS 保持在相对稳定的范围,输入的全局放大对归一化结果影响较小,径向梯度也受到抑制。它们不保证隐藏状态均值为零、不保证每个维度具有单位方差,也不保证 Pre-Norm 架构中的残差流本身具有有界范数。
回到 AttnRes
原文在层间注意部分使用了 RMSNorm 的正比例尺度不变性。假设层函数可以写成
f(z)=F(RMSNorm(z)),
那么在 ε=0、c>0 的条件下有
f(cz)=f(z).
对于一组非负权重 bs,令 B=∑sbs>0,则
s∑bsys=Bs∑Bbsys.
进入 In Norm 后,外部的正标量 B 会被消除。因此,把权重归一化到 ∑sas=1 不会改变后续层接收到的归一化方向;实际实现中的 ε 使该结论成为高信号幅度下的近似。
AttnRes 还使用
at+1,s∝exp(wt+1TRMSNorm(ys))
计算层间注意力。这里 RMSNorm 在计算相似度前控制 Key 的整体尺度,使注意力分数主要取决于归一化后的方向和可学习 Query wt+1。这两个位置分别对应 RMSNorm 的尺度不敏感性和方向保留性质,也是后续理解 AttnRes 数学形式与系统实现的基础。
Full AttnRes
Full AttnRes 让每一层直接对所有历史层输出做 Attention。它保留了最完整的层间选择能力,也把资源开销直接关联到网络深度 L。下面分别分析它在推理和训练中的成本。
沿用 Attention Residuals 技术报告的记号:L 是 AttnRes 的执行位置数,Self-Attention 和 MLP 分别计作一层;隐藏维度为 d;B、T 分别表示 batch size 和序列长度。为突出主项,以下复杂度暂时省略归一化、Softmax 标量运算和硬件利用率差异。
推理开销
Full AttnRes 的推理开销需要分别考察算术量、显存容量、HBM 访存量和端到端延迟。额外 FLOPs 较少只能说明计算单元的工作量较小;如果新增操作需要反复读取历史表示,实际延迟仍可能受到显存带宽与算子调度限制。
算术量
Full AttnRes 的第 l 层需要访问此前的 l 个表示。一共存在
Npair=l=1∑Ll=2L(L+1)
个 source-target 对。每一对需要一次 Query-Key 点积和一次加权 Value 累加;如果将一次乘法和一次加法分别计作一个 FLOP,主项约为 4d FLOPs。因此每个 token 的额外前向算术量约为
FFull,fwd≈4dNpair=2L(L+1)d.
这个 O(L2d) 项中的 L 是网络深度,通常远小于序列长度。它会随深度平方增长,但在当前大模型中往往仍小于主干线性层的计算量。
以 Kimi K2 的公开规格作量级参照:模型包含 61 个 Transformer Decoder Block,隐藏维度为 7168,每个 token 激活约 32B 参数。由于 Attention 和 MLP 各有一个 AttnRes 执行位置,可取 L=122,于是
Npair=7503,
FFull,fwd≈4×7168×7503≈0.215 GFLOPs/token.
使用 2Pactive 粗略估计主干线性层前向计算,可得约 64 GFLOPs/token,两者之比约为 0.34%。这个估算没有把序列 Attention 加入分母,因此只能用于判断量级。它来自 K2 配置与 AttnRes 公式的独立外推,不是 K2 上运行 Full AttnRes 的实测结果。每层新增一个 Query 向量和一组 RMSNorm 缩放参数时,参数增量约为
2Ld=2×122×7168≈1.75 M,
相对 K2 的总参数量同样很小。K2 开销估算对话提供了这一外推口径,具体数字在这里按公开配置重新计算。
静态 Query 与两阶段计算
朴素实现会让每一层重新扫描此前的 Key 和 Value,使残差模块的访存量随 L2d 增长。AttnRes 把第 l 层的 Query 设为与输入无关的可学习参数 wl。为了重排计算,将 L 层划为 G 个调度组,每组包含
S=GL
层。同一组中的 S 个 Query 在该组开始计算前已经全部确定,因此可以改写为两个阶段:
- Phase 1 把同一组中的 S 个 Query 组成矩阵,一次读取此前的层表示,批量计算历史部分的 Attention,同时保留最大值、指数和与加权 Value 和等 Softmax 统计量。
- Phase 2 按层处理组内新产生的表示,再通过 Online Softmax 合并两个阶段的结果。
这里减少的 KV 访问特指从 HBM 读取历史表示的次数。每个 Key 和 Value 仍然参与 S 个 Query 的计算。设历史来源数为 R,K,V∈RR×d,逐层执行需要分别计算 qsK⊤。把 S 个 Query 纵向堆叠后,有
q1K⊤q2K⊤⋮qSK⊤=q1q2⋮qSK⊤=QK⊤.
这个等式把 S 次矩阵—向量乘改写成一次矩阵—矩阵乘。高性能 GEMM 会把 K/V tile 从 HBM 载入共享内存和寄存器,再让同一 tile 服务多个 Query 行。若每个元素占 b Byte,只统计历史 K/V 的理想读流量,批处理前后的主项为
2SRdb⟶2Rdb.
乘法次数没有改变,变化来自片上数据复用和更高的算术强度。CUDA C++ Best Practices Guide说明了在线程块内将 GEMM tile 从全局内存载入共享内存后复用的通用机制;CUTLASS 的 GEMM 文档进一步说明了共享内存与寄存器层级的分块和流水化。
“每组读取一次”是算法级 I/O 模型。Query 行被拆到多个线程块、算子没有融合或片上容量不足时,同一 K/V tile 仍可能被多次载入。实际收益需要通过 DRAM 读取量、L2 命中率和内核延迟验证。静态 Query 提供批量计算的前提,GEMM tiling 实现片上复用,Online Softmax 负责把历史部分与当前组内的顺序部分精确合并。
设两个阶段分别处理互不重叠的来源集合 A 和 B。对 X∈{A,B},维护
mX=i∈Xmaxzi,ℓX=i∈X∑ezi−mX,oX=i∈X∑ezi−mXvi.
令 m=max(mA,mB),合并结果为
h=emA−mℓA+emB−mℓBemA−moA+emB−moB.
这个表达式等于在 A∪B 上直接计算 Softmax Attention,所以两阶段调度只改变计算顺序,不改变结果。
批处理让历史层表示从“组内每层读取一次”变为“每个调度组读取一次”。报告给出的每层访存复杂度由朴素 Full AttnRes 的 O(Ld) 降至
O((S+G)d).
在 L=128、G=8、S=16 的典型设置下,如果只统计残差机制自身的读写,优化后的 Full AttnRes 总 I/O 为 24d,标准 Residuals 为 3d。这些数字不能直接换算为端到端延迟,因为 Attention、MLP、MoE 和通信仍占据模型的大部分执行时间。
历史表示的显存占用
Full AttnRes 需要维护全部历史层表示或与之等价的未来层累加器,BF16 Prefill 的逻辑容量为
MFull,prefill=2BT(L+1)d Byte.
以 B=1、T=32K、L=128、d=7168 为例,这项容量约为 60 GB。推理架构分析指出,这些状态可以沿序列维度切分;8 路分片后,每个设备约为 7.5 GB。这解决了单设备容量问题,仍需通过两阶段调度降低反复读取产生的 HBM 流量。
Decode 阶段一次只处理本轮生成的 token。AttnRes 在同一个 token 的深度方向计算 Attention,已经完成的历史 token 不需要保留这组深度表示供后续 token 使用;序列 Attention 的 KV Cache 是另一项独立状态。因此单序列的深度工作区为 O(Ld),不随已有上下文长度 T 累积。代入 K2 维度和 BF16,一个 token 的 123 个深度状态约占
123×7168×2≈1.68 MiB.
训练开销
训练会为相同的 Full AttnRes 前向公式增加两项约束:所有历史层表示都要支持反向传播;大规模模型还要把这些表示传过 Pipeline Parallel 的阶段边界。原文将训练侧的工程判断链接到一篇训练侧讨论。下面的复杂度公式和实测数字以公开技术报告为准,K2 数字继续作为独立外推。
反向传播的算术量
Full AttnRes 的反向传播需要计算 Value、Key、Query 和 Softmax 分数的梯度。按矩阵乘与点积的常用 FLOPs 口径,反向计算约为前向的两倍,前向与反向合计约为前向的三倍:
FFull,train≈3FFull,fwd=6L(L+1)d.
代入前面的 K2 维度,结果约为 0.645 GFLOPs/token。同样用 6Pactive 估计主干线性层训练计算,可得约 192 GFLOPs/token,比例仍约为 0.34%。这个比例省略序列 Attention、RMSNorm 和 Softmax 等项,只说明 Full AttnRes 的核心向量算术并非主要增量。实际训练时间还受算子利用率、激活读写和分布式通信影响,不能由 FLOPs 比例直接预测。
激活重计算与流水线通信
在不使用激活重计算的普通训练中,各层输出原本就会为反向传播保留。Full AttnRes 复用这些输出,技术报告据此判断其额外显存接近零。大规模训练通常启用 Activation Checkpointing:中间输出在前向后释放,反向时再重算。Full AttnRes 中的每个层输出还会作为后续所有层的 Key 和 Value,因此这些输出必须长期存活,或采用新的重计算与调度方案。
设一个 microbatch 一共包含 M=BT 个 token,元素存储宽度为 b Byte。Full AttnRes 的深度历史逻辑容量为
MFull,history≈M(L+1)db.
继续使用 K2 的维度作量级参照,在 BF16、M=4096、d=7168 时,一个隐藏状态张量为 56 MiB,Embedding 和 122 个层输出合计约为
123×56 MiB=6.73 GiB.
这个数值是未分片、单 microbatch 的逻辑容量,未计入 Attention 与 MoE 中间激活、梯度、通信缓冲区和工作区。它也不等于训练净增显存;净增量取决于原有 Checkpointing、序列分片和流水线调度。
标准残差连接在相邻 Pipeline Stage 之间只需传递当前隐藏状态。Full AttnRes 要让下游阶段访问所有历史层输出,使跨阶段表示数增长到 O(L)。本地两阶段批处理可以降低 HBM 读取,无法消除这些表示的跨阶段传输。技术报告因此把大规模 Full AttnRes 训练的主要约束定位在 O(Ld) 的激活存活和 Pipeline Parallel 通信。
Block AttnRes
Block AttnRes 将 L 层划分为 N 个块,每块包含 S=L/N 层。块内继续累加层输出,块间只对压缩后的块表示做 Attention。它把 Full AttnRes 的历史表示数从 L 降到 N;当 N≈8 时,技术报告观察到 Block 版本能够保留 Full 版本的大部分收益。
推理开销
Block AttnRes 复用静态 Query 与两阶段计算:Phase 1 批量处理此前完成的块表示,Phase 2 顺序处理当前块的部分和,并通过前文给出的 Online Softmax 公式精确合并。语义层面的块压缩进一步减少了需要保存和读取的来源数。
在 L=128、N=8、S=16 的典型设置下,技术报告统计的 Block AttnRes 残差机制总 I/O 为 5.5d,低于优化后 Full AttnRes 的 24d,接近标准 Residuals 的 3d。这些统计均排除了 Attention、MLP 和 MoE 等层函数内部的读写。
历史表示的显存占用
Block AttnRes 在 Prefill 阶段需要保存 N 个块表示,BF16 下的容量约为
Mprefill=2BNTd Byte.
当 B=1、N=8、T=128K、d=7168 时,总容量约为 15 GB。技术报告沿序列维度把这些表示切分到 P 个 Tensor Parallel 设备,使单设备容量下降为
Mdevice=2BNPTd Byte.
在 P=8 时约为 1.9 GB;再采用 16K Chunked Prefill 后,报告给出的单设备额外容量低于 0.3 GB。
Decode 阶段只处理当前新 token 的深度表示。AttnRes 在同一个 token 的深度方向计算 Attention,已经完成的历史 token 不需要保留这组深度表示供后续 token 使用;序列 Attention 的 KV Cache 是另一项独立状态。因此 AttnRes 的 Decode 工作区按 O(BNd) 增长,不随已有上下文长度 T 累积。并行验证多个 token 或增大 Decode batch 时,这项容量按本轮同时处理的 token 数线性增长。
如何解释小于 2% 的延迟
技术报告在典型推理工作负载上测得 Block AttnRes 的端到端延迟增量低于 2%。推理架构分析进一步报告:在 batch size 为 128、包含 64 个 Transformer Decoder Block 的配置中,Decode 延迟增量低于 0.5 ms;32K Prefill 中的增量低于测试波动的可分辨范围。
这些结果依赖静态 Query、两阶段批处理、Online Softmax、算子融合、计算重叠和序列分片共同成立。2% 是报告工作负载下的端到端测量,不是所有硬件、batch size、并行策略和上下文长度下的固定常数。Full AttnRes 的 24d 残差 I/O 也明显高于 Block AttnRes 的 5.5d,所以 Block 版本承担了推理效率约束下的主要落地路径。
训练开销
Block AttnRes 只让后续层访问压缩后的块表示,把需要长期存活和跨阶段传播的表示数从 L 降到 N。训练侧仍要处理这些块表示的激活存储和 Pipeline Parallel 通信,但资源规模由网络总层数改为固定在约 8 个块。
激活显存
设一个 microbatch 一共包含 M=BT 个 token,元素存储宽度为 b Byte。Block AttnRes 的深度历史逻辑容量为
MBlock,history≈MNdb.
继续使用 K2 的维度作量级参照,在 BF16、M=4096、d=7168 时,一个隐藏状态张量为
4096×7168×2=56 MiB.
按 8 至 9 个块表示估算,运行容量约为 448 至 504 MiB,明显低于 Full AttnRes 的 6.73 GiB。这些数值是未分片、单 microbatch 的表示容量,未计入 Attention 与 MoE 中间激活、梯度、通信缓冲区和工作区。它们也不等于训练净增显存;净增量取决于原有 Checkpointing、序列分片和流水线调度。
Pipeline Parallel 通信
标准残差连接在相邻 Pipeline Stage 之间只需传递当前隐藏状态,通信张量大小不随模型深度增长。Full AttnRes 要让下游阶段访问所有历史层输出,跨阶段表示数为 O(L);Block AttnRes 只传递完成的块表示,把它降为 O(N)。
技术报告进一步分析了 Interleaved Pipeline。设物理阶段数为 P,每个物理阶段包含 V 个 Virtual Stage,总 Chunk 数为
C=PV.
如果每个物理阶段平均产生 Np 个块表示,朴素方案在每次切换时重发全部历史,每个 token 的通信元素数为
Commnaive=j=1∑C−1jNpd=2C(C−1)Npd.
跨阶段缓存让每个物理阶段保留此前接收的块,后续 Virtual Stage 只发送新增表示。通信量变为
Commcached=2P(P−1)Npd+(V−1)2P2Npd.
这里的公式按每个 token 的元素数计量;乘以 token 数和数据类型字节数后才得到实际 Payload。缓存把单次阶段切换的峰值从 O(C) 降到 O(P),在报告的模型中使峰值降低约 V 倍,并允许稳态 1F1B 调度把通信与计算重叠。反向传播可以复用同一缓存策略。
实测开销与结论边界
跨阶段缓存后,每个块只需在全部 Virtual Stage 中保存一次。Activation Checkpointing 可以消除块间 Attention 的中间量,Checkpoint 输入与它替代的普通隐藏状态大小相同。技术报告测得:未启用 Pipeline Parallel 时,Block AttnRes 的训练时间增量接近可忽略;启用 Pipeline Parallel 时,端到端训练时间增量低于 4%。
原文概括的“约 5% 以内开销换取 25% 收益”需要按 Scaling Law 口径理解:Block AttnRes 达到的验证损失,基线模型需要约 1.25 倍训练计算量才能达到。25% 指等效的基线训练计算量差异,不能解释为所有下游评测绝对提升 25%。
报告中的大模型实验证据来自 48B 总参数、3B 激活参数的 Kimi Linear,预训练量为 1.4T tokens;报告没有提供 K2 加入 Full AttnRes 后的训练实测。上面的 K2 数字用于把公式换算成可感知的容量和算术量级。最终可以支持的结论是:Full AttnRes 的核心算术量在大模型总计算中占比较小,推理端可以利用静态 Query 重排计算;训练端仍受激活存活时间与 Pipeline Parallel 通信约束。Block AttnRes 用块级压缩把相关状态数量从 L 降到约 8,在保留 Full 版本大部分收益的同时进入当前训练和推理系统的可承受范围。