← 返回文章档案

Attention Residuals 阅读补充

围绕 Attention Residuals 整理相关知识、数学推导与系统分析,包括 RMSNorm 的数学原理,以及 Full 与 Block AttnRes 的推理和训练开销。

这篇文章用于沉淀和分析阅读苏神(苏剑林)《Attention Residuals 回忆录》 时涉及的相关知识。内容会随阅读逐步增加,每一章集中处理一个需要补充推导或实现背景的部分。

RMSNorm 的数学原理

RMSNorm 用输入向量在特征维度上的二阶原点矩衡量整体幅度,再按这个幅度缩放整个向量。

设一个 token 的隐藏状态为

x=(x1,x2,,xd)Rd.\boldsymbol{x}=(x_1,x_2,\ldots,x_d)\in\mathbb{R}^d.

先定义带数值稳定项的均方根

rε(x)=1di=1dxi2+ε,r_\varepsilon(\boldsymbol{x}) =\sqrt{\frac{1}{d}\sum_{i=1}^d x_i^2+\varepsilon},

再对所有维度使用同一个标量完成归一化,并加入可学习的逐维缩放参数 γ\boldsymbol{\gamma}

RMSNorm(x)=γxrε(x).\operatorname{RMSNorm}(\boldsymbol{x}) =\boldsymbol{\gamma}\odot \frac{\boldsymbol{x}}{r_\varepsilon(\boldsymbol{x})}.

这里的均值发生在同一个 token 的 dd 个特征维度上。它是对当前向量的确定性计算,无需估计数据集或 batch 的统计量。RMSNorm 原论文 将其作为 LayerNorm 的简化形式,保留重缩放不变性并省去均值中心化。

从均方根到欧氏范数

暂时令 ε=0\varepsilon=0,并假设 x0\boldsymbol{x}\neq\boldsymbol{0}。均方根可以直接写成欧氏范数:

r0(x)=1di=1dxi2=x2d.r_0(\boldsymbol{x}) =\sqrt{\frac{1}{d}\sum_{i=1}^d x_i^2} =\frac{\lVert\boldsymbol{x}\rVert_2}{\sqrt d}.

记不含可学习参数的归一化结果为 x~\tilde{\boldsymbol{x}},则

x~=xr0(x)=dxx2.\tilde{\boldsymbol{x}} =\frac{\boldsymbol{x}}{r_0(\boldsymbol{x})} =\sqrt d\frac{\boldsymbol{x}}{\lVert\boldsymbol{x}\rVert_2}.

x/x2\boldsymbol{x}/\lVert\boldsymbol{x}\rVert_2 只保留输入方向,前面的 d\sqrt d 将长度设为固定值。直接计算可得

x~22=i=1dxi21dj=1dxj2=d,\lVert\tilde{\boldsymbol{x}}\rVert_2^2 =\frac{\sum_{i=1}^d x_i^2} {\frac{1}{d}\sum_{j=1}^d x_j^2} =d,

因此

x~2=d.\lVert\tilde{\boldsymbol{x}}\rVert_2=\sqrt d.

ε=0\varepsilon=0 的理想条件下,RMSNorm 的归一化部分沿径向把非零输入映射到半径为 d\sqrt d 的超球面。这个映射保留方向并移除长度。

目标长度取 d\sqrt d,等价于让归一化结果的平均平方值为 11

1di=1dx~i2=1.\frac{1}{d}\sum_{i=1}^d \tilde{x}_i^2=1.

因此 RMSNorm 控制的是整个隐藏状态的平均平方幅度。单个坐标仍然可以远大于 11,各个维度也不会分别获得单位方差。

二阶原点矩与方差

RMSNorm 和 LayerNorm 的差异可以从二阶原点矩与方差的关系看出。定义当前向量在特征维度上的均值和方差:

μ=1di=1dxi,σ2=1di=1d(xiμ)2.\mu=\frac{1}{d}\sum_{i=1}^d x_i, \qquad \sigma^2=\frac{1}{d}\sum_{i=1}^d(x_i-\mu)^2.

二阶原点矩为

m2=1di=1dxi2.m_2=\frac{1}{d}\sum_{i=1}^d x_i^2.

展开方差后得到

m2=σ2+μ2.m_2=\sigma^2+\mu^2.

RMSNorm 使用 m2m_2,所以均值偏移 μ2\mu^2 也会进入归一化尺度。LayerNorm 先减去均值,再使用中心二阶矩 σ2\sigma^2。忽略可学习参数时,两者分别为

x~RMS=xσ2+μ2+ε,\widetilde{\boldsymbol{x}}_{\mathrm{RMS}} =\frac{\boldsymbol{x}}{\sqrt{\sigma^2+\mu^2+\varepsilon}}, x~LN=xμ1σ2+ε.\widetilde{\boldsymbol{x}}_{\mathrm{LN}} =\frac{\boldsymbol{x}-\mu\boldsymbol{1}} {\sqrt{\sigma^2+\varepsilon}}.

其中,1=(1,1,,1)TRd\boldsymbol{1}=(1,1,\ldots,1)^{\mathsf T}\in\mathbb{R}^ddd 维全一列向量,因此 μ1\mu\boldsymbol{1} 表示每个坐标都取值为 μ\mu 的向量。

μ=0\mu=0、两者采用相同的 ε\varepsilon 和缩放参数,并且 LayerNorm 的偏置为零时,两种归一化给出相同结果。当 μ\mu 只接近零时,它们使用的尺度接近,RMSNorm 仍会保留沿 1\boldsymbol{1} 方向的分量。

这个差异也有直接的几何表示。令

P=I1d11T,\boldsymbol{P} =\boldsymbol{I}-\frac{1}{d}\boldsymbol{1}\boldsymbol{1}^{\mathsf T},

Px=xμ1\boldsymbol{P}\boldsymbol{x}=\boldsymbol{x}-\mu\boldsymbol{1}。LayerNorm 先把输入投影到与 1\boldsymbol{1} 正交的均值为零子空间,再在该子空间内归一化;RMSNorm 直接在完整的 Rd\mathbb{R}^d 中按长度归一化。LayerNorm 原论文 给出的定义同时包含中心化、缩放以及归一化后的可学习增益和偏置。

性质RMSNormLayerNorm
归一化统计量二阶原点矩 m2m_2中心二阶矩 σ2\sigma^2
是否减去特征均值
正比例尺度变化ε=0\varepsilon=0 时消除ε=0\varepsilon=0 时消除
整体平移 x+c1\boldsymbol{x}+c\boldsymbol{1}会改变结果归一化结果不变
几何作用在完整空间中按半径归一化先进入均值为零子空间,再按半径归一化

尺度变化与平移

尺度不敏感性是 RMSNorm 在 AttnRes 推导中承担的关键性质。令 a>0a>0,则

rε(ax)=a2m2+ε=am2+εa2.r_\varepsilon(a\boldsymbol{x}) =\sqrt{a^2m_2+\varepsilon} =a\sqrt{m_2+\frac{\varepsilon}{a^2}}.

ε=0\varepsilon=0 时,正比例缩放会被精确消除:

RMSNorm(ax)=RMSNorm(x),a>0.\operatorname{RMSNorm}(a\boldsymbol{x}) =\operatorname{RMSNorm}(\boldsymbol{x}), \qquad a>0.

实际实现采用 ε>0\varepsilon>0,此时等式只近似成立。当 a2m2εa^2m_2\gg\varepsilon 时,稳定项的影响很小。对于 a<0a<0,在 ε=0\varepsilon=0 时有

RMSNorm(ax)=RMSNorm(x),\operatorname{RMSNorm}(a\boldsymbol{x}) =-\operatorname{RMSNorm}(\boldsymbol{x}),

因为输入方向发生了符号翻转。

RMSNorm 对整体平移没有这项不变性。令 x=x+c1\boldsymbol{x}'=\boldsymbol{x}+c\boldsymbol{1},其二阶原点矩变为

1di=1d(xi+c)2=m2+2cμ+c2.\frac{1}{d}\sum_{i=1}^d(x_i+c)^2 =m_2+2c\mu+c^2.

LayerNorm 会在中心化时消除新增的 c1c\boldsymbol{1},所以整体平移不会改变其归一化结果。RMSNorm 保留这部分均值信息,并让它参与尺度计算。

梯度沿切向传播

为了观察归一化怎样改变梯度,先去掉 γ\boldsymbol{\gamma},记

r=1dxTx+ε,y=xr.r=\sqrt{\frac{1}{d}\boldsymbol{x}^{\mathsf T}\boldsymbol{x}+\varepsilon}, \qquad \boldsymbol{y}=\frac{\boldsymbol{x}}{r}.

它的 Jacobian 为

yx=1rI1dr3xxT.\frac{\partial\boldsymbol{y}}{\partial\boldsymbol{x}} =\frac{1}{r}\boldsymbol{I} -\frac{1}{dr^3}\boldsymbol{x}\boldsymbol{x}^{\mathsf T}.

如果上游梯度为 g=L/y\boldsymbol{g}=\partial\mathcal{L}/\partial\boldsymbol{y},输入梯度就是

Lx=1r(gyyTgd).\frac{\partial\mathcal{L}}{\partial\boldsymbol{x}} =\frac{1}{r}\left( \boldsymbol{g} -\boldsymbol{y}\frac{\boldsymbol{y}^{\mathsf T}\boldsymbol{g}}{d} \right).

第一项按输入 RMS 统一调整梯度尺度,第二项减去与输入径向方向相关的分量。当 ε=0\varepsilon=0 时,Jacobian 在径向上的作用为

yxx=0.\frac{\partial\boldsymbol{y}}{\partial\boldsymbol{x}}\boldsymbol{x} =\boldsymbol{0}.

这与正比例尺度不变性一致:沿 x\boldsymbol{x} 方向只改变长度,归一化结果保持不变。对任意满足 xTv=0\boldsymbol{x}^{\mathsf T}\boldsymbol{v}=0 的切向量 v\boldsymbol{v},Jacobian 的作用为 v/r\boldsymbol{v}/r。因此在理想条件下,归一化核心保留切向变化并消除径向变化。

加入 ε\varepsilon 后,径向特征值变为 ε/r3\varepsilon/r^3,径向梯度会被显著减小,但不会严格归零。加入 γ\boldsymbol{\gamma} 后,上述分析仍适用于归一化核心;反向传播时,上游梯度会先逐维乘以 γ\boldsymbol{\gamma}

ε 与 γ 改变了什么

ε\varepsilon 在输入接近零向量时保持分母为正。它也会让输出长度略低于 d\sqrt d

x~22=x22x22/d+ε=dx22x22+dε<d.\lVert\tilde{\boldsymbol{x}}\rVert_2^2 =\frac{\lVert\boldsymbol{x}\rVert_2^2} {\lVert\boldsymbol{x}\rVert_2^2/d+\varepsilon} =\frac{d\lVert\boldsymbol{x}\rVert_2^2} {\lVert\boldsymbol{x}\rVert_2^2+d\varepsilon} <d.

只有在 x22dε\lVert\boldsymbol{x}\rVert_2^2\gg d\varepsilon 时,输出长度才接近 d\sqrt d

γ\boldsymbol{\gamma} 为每个特征维度恢复可学习的尺度。归一化核心在理想条件下把输入映射到超球面,逐维乘以非均匀的 γ\boldsymbol{\gamma} 后,这个超球面的像成为轴对齐的椭球面。模型由此可以调整各个特征维度的有效范围。

这些性质共同提供整体尺度控制:后续层看到的输入 RMS 保持在相对稳定的范围,输入的全局放大对归一化结果影响较小,径向梯度也受到抑制。它们不保证隐藏状态均值为零、不保证每个维度具有单位方差,也不保证 Pre-Norm 架构中的残差流本身具有有界范数。

回到 AttnRes

原文在层间注意部分使用了 RMSNorm 的正比例尺度不变性。假设层函数可以写成

f(z)=F(RMSNorm(z)),\boldsymbol{f}(\boldsymbol{z}) =\boldsymbol{F}(\operatorname{RMSNorm}(\boldsymbol{z})),

那么在 ε=0\varepsilon=0c>0c>0 的条件下有

f(cz)=f(z).\boldsymbol{f}(c\boldsymbol{z})=\boldsymbol{f}(\boldsymbol{z}).

对于一组非负权重 bsb_s,令 B=sbs>0B=\sum_s b_s>0,则

sbsys=BsbsBys.\sum_s b_s\boldsymbol{y}_s =B\sum_s\frac{b_s}{B}\boldsymbol{y}_s.

进入 In Norm 后,外部的正标量 BB 会被消除。因此,把权重归一化到 sas=1\sum_s a_s=1 不会改变后续层接收到的归一化方向;实际实现中的 ε\varepsilon 使该结论成为高信号幅度下的近似。

AttnRes 还使用

at+1,sexp(wt+1TRMSNorm(ys))a_{t+1,s}\propto \exp\left( \boldsymbol{w}_{t+1}^{\mathsf T} \operatorname{RMSNorm}(\boldsymbol{y}_s) \right)

计算层间注意力。这里 RMSNorm 在计算相似度前控制 Key 的整体尺度,使注意力分数主要取决于归一化后的方向和可学习 Query wt+1\boldsymbol{w}_{t+1}。这两个位置分别对应 RMSNorm 的尺度不敏感性和方向保留性质,也是后续理解 AttnRes 数学形式与系统实现的基础。

Full AttnRes

Full AttnRes 让每一层直接对所有历史层输出做 Attention。它保留了最完整的层间选择能力,也把资源开销直接关联到网络深度 LL。下面分别分析它在推理和训练中的成本。

沿用 Attention Residuals 技术报告的记号:LL 是 AttnRes 的执行位置数,Self-Attention 和 MLP 分别计作一层;隐藏维度为 ddBBTT 分别表示 batch size 和序列长度。为突出主项,以下复杂度暂时省略归一化、Softmax 标量运算和硬件利用率差异。

推理开销

Full AttnRes 的推理开销需要分别考察算术量、显存容量、HBM 访存量和端到端延迟。额外 FLOPs 较少只能说明计算单元的工作量较小;如果新增操作需要反复读取历史表示,实际延迟仍可能受到显存带宽与算子调度限制。

算术量

Full AttnRes 的第 ll 层需要访问此前的 ll 个表示。一共存在

Npair=l=1Ll=L(L+1)2N_{\mathrm{pair}} =\sum_{l=1}^{L}l =\frac{L(L+1)}{2}

个 source-target 对。每一对需要一次 Query-Key 点积和一次加权 Value 累加;如果将一次乘法和一次加法分别计作一个 FLOP,主项约为 4d4d FLOPs。因此每个 token 的额外前向算术量约为

FFull,fwd4dNpair=2L(L+1)d.F_{\mathrm{Full,fwd}} \approx 4dN_{\mathrm{pair}} =2L(L+1)d.

这个 O(L2d)O(L^2d) 项中的 LL 是网络深度,通常远小于序列长度。它会随深度平方增长,但在当前大模型中往往仍小于主干线性层的计算量。

Kimi K2 的公开规格作量级参照:模型包含 61 个 Transformer Decoder Block,隐藏维度为 7168,每个 token 激活约 32B 参数。由于 Attention 和 MLP 各有一个 AttnRes 执行位置,可取 L=122L=122,于是

Npair=7503,N_{\mathrm{pair}}=7503, FFull,fwd4×7168×75030.215 GFLOPs/token.F_{\mathrm{Full,fwd}} \approx 4\times7168\times7503 \approx 0.215\ \mathrm{GFLOPs/token}.

使用 2Pactive2P_{\mathrm{active}} 粗略估计主干线性层前向计算,可得约 64 GFLOPs/token64\ \mathrm{GFLOPs/token},两者之比约为 0.34%0.34\%。这个估算没有把序列 Attention 加入分母,因此只能用于判断量级。它来自 K2 配置与 AttnRes 公式的独立外推,不是 K2 上运行 Full AttnRes 的实测结果。每层新增一个 Query 向量和一组 RMSNorm 缩放参数时,参数增量约为

2Ld=2×122×71681.75 M,2Ld=2\times122\times7168\approx1.75\ \text{M},

相对 K2 的总参数量同样很小。K2 开销估算对话提供了这一外推口径,具体数字在这里按公开配置重新计算。

静态 Query 与两阶段计算

朴素实现会让每一层重新扫描此前的 Key 和 Value,使残差模块的访存量随 L2dL^2d 增长。AttnRes 把第 ll 层的 Query 设为与输入无关的可学习参数 wl\boldsymbol{w}_l。为了重排计算,将 LL 层划为 GG 个调度组,每组包含

S=LGS=\frac{L}{G}

层。同一组中的 SS 个 Query 在该组开始计算前已经全部确定,因此可以改写为两个阶段:

  1. Phase 1 把同一组中的 SS 个 Query 组成矩阵,一次读取此前的层表示,批量计算历史部分的 Attention,同时保留最大值、指数和与加权 Value 和等 Softmax 统计量。
  2. Phase 2 按层处理组内新产生的表示,再通过 Online Softmax 合并两个阶段的结果。

这里减少的 KV 访问特指从 HBM 读取历史表示的次数。每个 Key 和 Value 仍然参与 SS 个 Query 的计算。设历史来源数为 RRK,VRR×dK,V\in\mathbb{R}^{R\times d},逐层执行需要分别计算 qsKq_sK^\top。把 SS 个 Query 纵向堆叠后,有

[q1Kq2KqSK]=[q1q2qS]K=QK.\begin{bmatrix} q_1K^\top\\ q_2K^\top\\ \vdots\\ q_SK^\top \end{bmatrix} = \begin{bmatrix} q_1\\ q_2\\ \vdots\\ q_S \end{bmatrix} K^\top =QK^\top.

这个等式把 SS 次矩阵—向量乘改写成一次矩阵—矩阵乘。高性能 GEMM 会把 K/V tile 从 HBM 载入共享内存和寄存器,再让同一 tile 服务多个 Query 行。若每个元素占 bb Byte,只统计历史 K/V 的理想读流量,批处理前后的主项为

2SRdb2Rdb.2SRdb \quad\longrightarrow\quad 2Rdb.

乘法次数没有改变,变化来自片上数据复用和更高的算术强度。CUDA C++ Best Practices Guide说明了在线程块内将 GEMM tile 从全局内存载入共享内存后复用的通用机制;CUTLASS 的 GEMM 文档进一步说明了共享内存与寄存器层级的分块和流水化。

“每组读取一次”是算法级 I/O 模型。Query 行被拆到多个线程块、算子没有融合或片上容量不足时,同一 K/V tile 仍可能被多次载入。实际收益需要通过 DRAM 读取量、L2 命中率和内核延迟验证。静态 Query 提供批量计算的前提,GEMM tiling 实现片上复用,Online Softmax 负责把历史部分与当前组内的顺序部分精确合并。

设两个阶段分别处理互不重叠的来源集合 AABB。对 X{A,B}X\in\{A,B\},维护

mX=maxiXzi,X=iXezimX,oX=iXezimXvi.m_X=\max_{i\in X}z_i, \qquad \ell_X=\sum_{i\in X}e^{z_i-m_X}, \qquad \boldsymbol{o}_X=\sum_{i\in X}e^{z_i-m_X}\boldsymbol{v}_i.

m=max(mA,mB)m=\max(m_A,m_B),合并结果为

h=emAmoA+emBmoBemAmA+emBmB.\boldsymbol{h} =\frac{ e^{m_A-m}\boldsymbol{o}_A+e^{m_B-m}\boldsymbol{o}_B }{ e^{m_A-m}\ell_A+e^{m_B-m}\ell_B }.

这个表达式等于在 ABA\cup B 上直接计算 Softmax Attention,所以两阶段调度只改变计算顺序,不改变结果。

批处理让历史层表示从“组内每层读取一次”变为“每个调度组读取一次”。报告给出的每层访存复杂度由朴素 Full AttnRes 的 O(Ld)O(Ld) 降至

O((S+G)d).O((S+G)d).

L=128L=128G=8G=8S=16S=16 的典型设置下,如果只统计残差机制自身的读写,优化后的 Full AttnRes 总 I/O 为 24d24d,标准 Residuals 为 3d3d。这些数字不能直接换算为端到端延迟,因为 Attention、MLP、MoE 和通信仍占据模型的大部分执行时间。

历史表示的显存占用

Full AttnRes 需要维护全部历史层表示或与之等价的未来层累加器,BF16 Prefill 的逻辑容量为

MFull,prefill=2BT(L+1)d Byte.M_{\mathrm{Full,prefill}} =2BT(L+1)d\ \text{Byte}.

B=1B=1T=32KT=32\mathrm{K}L=128L=128d=7168d=7168 为例,这项容量约为 60 GB60\ \mathrm{GB}推理架构分析指出,这些状态可以沿序列维度切分;8 路分片后,每个设备约为 7.5 GB7.5\ \mathrm{GB}。这解决了单设备容量问题,仍需通过两阶段调度降低反复读取产生的 HBM 流量。

Decode 阶段一次只处理本轮生成的 token。AttnRes 在同一个 token 的深度方向计算 Attention,已经完成的历史 token 不需要保留这组深度表示供后续 token 使用;序列 Attention 的 KV Cache 是另一项独立状态。因此单序列的深度工作区为 O(Ld)O(Ld),不随已有上下文长度 TT 累积。代入 K2 维度和 BF16,一个 token 的 123 个深度状态约占

123×7168×21.68 MiB.123\times7168\times2 \approx1.68\ \mathrm{MiB}.

训练开销

训练会为相同的 Full AttnRes 前向公式增加两项约束:所有历史层表示都要支持反向传播;大规模模型还要把这些表示传过 Pipeline Parallel 的阶段边界。原文将训练侧的工程判断链接到一篇训练侧讨论。下面的复杂度公式和实测数字以公开技术报告为准,K2 数字继续作为独立外推。

反向传播的算术量

Full AttnRes 的反向传播需要计算 Value、Key、Query 和 Softmax 分数的梯度。按矩阵乘与点积的常用 FLOPs 口径,反向计算约为前向的两倍,前向与反向合计约为前向的三倍:

FFull,train3FFull,fwd=6L(L+1)d.F_{\mathrm{Full,train}} \approx 3F_{\mathrm{Full,fwd}} =6L(L+1)d.

代入前面的 K2 维度,结果约为 0.645 GFLOPs/token0.645\ \mathrm{GFLOPs/token}。同样用 6Pactive6P_{\mathrm{active}} 估计主干线性层训练计算,可得约 192 GFLOPs/token192\ \mathrm{GFLOPs/token},比例仍约为 0.34%0.34\%。这个比例省略序列 Attention、RMSNorm 和 Softmax 等项,只说明 Full AttnRes 的核心向量算术并非主要增量。实际训练时间还受算子利用率、激活读写和分布式通信影响,不能由 FLOPs 比例直接预测。

激活重计算与流水线通信

在不使用激活重计算的普通训练中,各层输出原本就会为反向传播保留。Full AttnRes 复用这些输出,技术报告据此判断其额外显存接近零。大规模训练通常启用 Activation Checkpointing:中间输出在前向后释放,反向时再重算。Full AttnRes 中的每个层输出还会作为后续所有层的 Key 和 Value,因此这些输出必须长期存活,或采用新的重计算与调度方案。

设一个 microbatch 一共包含 M=BTM=BT 个 token,元素存储宽度为 bb Byte。Full AttnRes 的深度历史逻辑容量为

MFull,historyM(L+1)db.M_{\mathrm{Full,history}} \approx M(L+1)db.

继续使用 K2 的维度作量级参照,在 BF16、M=4096M=4096d=7168d=7168 时,一个隐藏状态张量为 56 MiB56\ \mathrm{MiB},Embedding 和 122 个层输出合计约为

123×56 MiB=6.73 GiB.123\times56\ \mathrm{MiB} =6.73\ \mathrm{GiB}.

这个数值是未分片、单 microbatch 的逻辑容量,未计入 Attention 与 MoE 中间激活、梯度、通信缓冲区和工作区。它也不等于训练净增显存;净增量取决于原有 Checkpointing、序列分片和流水线调度。

标准残差连接在相邻 Pipeline Stage 之间只需传递当前隐藏状态。Full AttnRes 要让下游阶段访问所有历史层输出,使跨阶段表示数增长到 O(L)O(L)。本地两阶段批处理可以降低 HBM 读取,无法消除这些表示的跨阶段传输。技术报告因此把大规模 Full AttnRes 训练的主要约束定位在 O(Ld)O(Ld) 的激活存活和 Pipeline Parallel 通信。

Block AttnRes

Block AttnRes 将 LL 层划分为 NN 个块,每块包含 S=L/NS=L/N 层。块内继续累加层输出,块间只对压缩后的块表示做 Attention。它把 Full AttnRes 的历史表示数从 LL 降到 NN;当 N8N\approx8 时,技术报告观察到 Block 版本能够保留 Full 版本的大部分收益。

推理开销

Block AttnRes 复用静态 Query 与两阶段计算:Phase 1 批量处理此前完成的块表示,Phase 2 顺序处理当前块的部分和,并通过前文给出的 Online Softmax 公式精确合并。语义层面的块压缩进一步减少了需要保存和读取的来源数。

L=128L=128N=8N=8S=16S=16 的典型设置下,技术报告统计的 Block AttnRes 残差机制总 I/O 为 5.5d5.5d,低于优化后 Full AttnRes 的 24d24d,接近标准 Residuals 的 3d3d。这些统计均排除了 Attention、MLP 和 MoE 等层函数内部的读写。

历史表示的显存占用

Block AttnRes 在 Prefill 阶段需要保存 NN 个块表示,BF16 下的容量约为

Mprefill=2BNTd Byte.M_{\mathrm{prefill}} =2BNTd\ \text{Byte}.

B=1B=1N=8N=8T=128KT=128\mathrm{K}d=7168d=7168 时,总容量约为 15 GB15\ \mathrm{GB}。技术报告沿序列维度把这些表示切分到 PP 个 Tensor Parallel 设备,使单设备容量下降为

Mdevice=2BNTPd Byte.M_{\mathrm{device}} =2BN\frac{T}{P}d\ \text{Byte}.

P=8P=8 时约为 1.9 GB1.9\ \mathrm{GB};再采用 16K Chunked Prefill 后,报告给出的单设备额外容量低于 0.3 GB0.3\ \mathrm{GB}

Decode 阶段只处理当前新 token 的深度表示。AttnRes 在同一个 token 的深度方向计算 Attention,已经完成的历史 token 不需要保留这组深度表示供后续 token 使用;序列 Attention 的 KV Cache 是另一项独立状态。因此 AttnRes 的 Decode 工作区按 O(BNd)O(BNd) 增长,不随已有上下文长度 TT 累积。并行验证多个 token 或增大 Decode batch 时,这项容量按本轮同时处理的 token 数线性增长。

如何解释小于 2% 的延迟

技术报告在典型推理工作负载上测得 Block AttnRes 的端到端延迟增量低于 2%2\%推理架构分析进一步报告:在 batch size 为 128、包含 64 个 Transformer Decoder Block 的配置中,Decode 延迟增量低于 0.5 ms0.5\ \mathrm{ms};32K Prefill 中的增量低于测试波动的可分辨范围。

这些结果依赖静态 Query、两阶段批处理、Online Softmax、算子融合、计算重叠和序列分片共同成立。2%2\% 是报告工作负载下的端到端测量,不是所有硬件、batch size、并行策略和上下文长度下的固定常数。Full AttnRes 的 24d24d 残差 I/O 也明显高于 Block AttnRes 的 5.5d5.5d,所以 Block 版本承担了推理效率约束下的主要落地路径。

训练开销

Block AttnRes 只让后续层访问压缩后的块表示,把需要长期存活和跨阶段传播的表示数从 LL 降到 NN。训练侧仍要处理这些块表示的激活存储和 Pipeline Parallel 通信,但资源规模由网络总层数改为固定在约 8 个块。

激活显存

设一个 microbatch 一共包含 M=BTM=BT 个 token,元素存储宽度为 bb Byte。Block AttnRes 的深度历史逻辑容量为

MBlock,historyMNdb.M_{\mathrm{Block,history}} \approx MNdb.

继续使用 K2 的维度作量级参照,在 BF16、M=4096M=4096d=7168d=7168 时,一个隐藏状态张量为

4096×7168×2=56 MiB.4096\times7168\times2 =56\ \mathrm{MiB}.

按 8 至 9 个块表示估算,运行容量约为 448448504 MiB504\ \mathrm{MiB},明显低于 Full AttnRes 的 6.73 GiB6.73\ \mathrm{GiB}。这些数值是未分片、单 microbatch 的表示容量,未计入 Attention 与 MoE 中间激活、梯度、通信缓冲区和工作区。它们也不等于训练净增显存;净增量取决于原有 Checkpointing、序列分片和流水线调度。

Pipeline Parallel 通信

标准残差连接在相邻 Pipeline Stage 之间只需传递当前隐藏状态,通信张量大小不随模型深度增长。Full AttnRes 要让下游阶段访问所有历史层输出,跨阶段表示数为 O(L)O(L);Block AttnRes 只传递完成的块表示,把它降为 O(N)O(N)

技术报告进一步分析了 Interleaved Pipeline。设物理阶段数为 PP,每个物理阶段包含 VV 个 Virtual Stage,总 Chunk 数为

C=PV.C=PV.

如果每个物理阶段平均产生 NpN_p 个块表示,朴素方案在每次切换时重发全部历史,每个 token 的通信元素数为

Commnaive=j=1C1jNpd=C(C1)2Npd.\operatorname{Comm}_{\mathrm{naive}} =\sum_{j=1}^{C-1}jN_pd =\frac{C(C-1)}{2}N_pd.

跨阶段缓存让每个物理阶段保留此前接收的块,后续 Virtual Stage 只发送新增表示。通信量变为

Commcached=P(P1)2Npd+(V1)P22Npd.\operatorname{Comm}_{\mathrm{cached}} =\frac{P(P-1)}{2}N_pd +(V-1)\frac{P^2}{2}N_pd.

这里的公式按每个 token 的元素数计量;乘以 token 数和数据类型字节数后才得到实际 Payload。缓存把单次阶段切换的峰值从 O(C)O(C) 降到 O(P)O(P),在报告的模型中使峰值降低约 VV 倍,并允许稳态 1F1B 调度把通信与计算重叠。反向传播可以复用同一缓存策略。

实测开销与结论边界

跨阶段缓存后,每个块只需在全部 Virtual Stage 中保存一次。Activation Checkpointing 可以消除块间 Attention 的中间量,Checkpoint 输入与它替代的普通隐藏状态大小相同。技术报告测得:未启用 Pipeline Parallel 时,Block AttnRes 的训练时间增量接近可忽略;启用 Pipeline Parallel 时,端到端训练时间增量低于 4%4\%

原文概括的“约 5%5\% 以内开销换取 25%25\% 收益”需要按 Scaling Law 口径理解:Block AttnRes 达到的验证损失,基线模型需要约 1.251.25 倍训练计算量才能达到。25%25\% 指等效的基线训练计算量差异,不能解释为所有下游评测绝对提升 25%25\%

报告中的大模型实验证据来自 48B 总参数、3B 激活参数的 Kimi Linear,预训练量为 1.4T tokens;报告没有提供 K2 加入 Full AttnRes 后的训练实测。上面的 K2 数字用于把公式换算成可感知的容量和算术量级。最终可以支持的结论是:Full AttnRes 的核心算术量在大模型总计算中占比较小,推理端可以利用静态 Query 重排计算;训练端仍受激活存活时间与 Pipeline Parallel 通信约束。Block AttnRes 用块级压缩把相关状态数量从 LL 降到约 8,在保留 Full 版本大部分收益的同时进入当前训练和推理系统的可承受范围。