内容提要
本文介绍Kimi K3模型序列维度的chunkwise并行算法:将序列分块,块间递推状态、块内用矩阵乘并行计算。K3将每步log-decay加−5下界,使16-token对角tile也能用BF16矩阵乘,消除FP32逐位置对瓶颈。相比通用DPLR内核,KDA内核因结构简化速度快约一倍。训练和prefill用chunkwise形式,decode用递推形式。
延伸解读
下界衰减的数值动机
K3 将每步 log-decay 的下界设为 -5,与 16-token 的 tile 宽度配套,使 16 步累计衰减不超过 e^-80,落在 BF16 动态范围内。这一改动让原本只能逐位置 FP32 计算的对角 tile 也能使用 BF16 矩阵乘,从而消除块内计算的主要瓶颈。
衰减下界对记忆的影响
下界 -5 意味着每个通道每步最多衰减到原来的约 0.67%,因此记忆不会因衰减而被瞬间清空。若需精确擦除,仍依赖 delta rule 的 (I - βkk^T) 项。论文未讨论此下界对表达力的影响,但指出类似下界门在 RWKV-7、Griffin、HGRN2 中已有先例。
KDA 内核为何更快
相比通用 DPLR 内核,KDA 因转移矩阵的特殊结构(a_t 和 b_t 均绑定 k_t),在 chunkwise 形式中只需两张 tile 矩阵(A_qk, A_kk),而 DPLR 需要四张;输出阶段也简化为一次输出和一次状态更新。这使得 KDA 内核在 2K 到 64K 长度上速度约为 DPLR 的两倍。
Q&A
Kimi K3 的 chunkwise 并行算法是如何将序列分块并实现块间递推、块内并行的?
Kimi K3 将序列切成长度为 C(C=64)的块。块间只更新一次状态,即从 S[t] 递推到 S[t+1];块内则通过矩阵乘并行计算 C 个输出。具体地,块内利用 WY 表示将一串 Householder 乘积转化为对角减低秩和,并通过 UT 变换(一次 C×C 三角求逆)解出辅助向量,从而将顺序计算转化为矩阵乘。
Kimi K3 为什么给每步的 log-decay 加一个 -5 的下界?
K3 将每步的 log-decay 限制在 -5 到 0 之间,这样 16 步累计的 log-decay 最小为 -80,对应的 exp 值 e^80 约 5.5×10^34,在 BF16 动态范围(最大约 3.4×10^38)内。这使得 16-token 的对角 tile 也能拆成两个因子用 BF16 矩阵乘,从而消除 FP32 逐位置对的瓶颈,提升硬件利用率。
KDA 内核相比通用 DPLR 内核为什么快约一倍?
因为 KDA 的转移矩阵结构更简单:通用 DPLR 内核需要四张二级 tile 矩阵(Aab, Aak, Aqb, Aqk),而 KDA 只需要两张(Aqk, Akk);输出阶段 DPLR 有三项输出和两次状态更新,KDA 合成一行输出和一次状态更新。因此计算量减少,速度提升约一倍。
Kimi K3 在训练、prefill 和 decode 阶段分别使用哪种计算形式?
训练和 prefill 阶段使用 chunkwise 形式(chunk_kda),decode 阶段(每次生成一个 token)使用递推形式(fused_recurrent_kda)。代码中通过判断 use_cache 和 q_len 是否为 1 来选择模式。
Kimi K3 的衰减函数与 Kimi Linear 的负 softplus 有何不同?
Kimi Linear 使用负 softplus:g = -e^A * softplus(z),取值范围为 (-∞, 0),没有下界。K3 改用带尺度的 sigmoid:g = g_min * sigmoid(e^A z),其中 g_min = -5,取值范围为 (-5, 0),有下界。这保证了数值稳定性,并允许使用 BF16 矩阵乘。
Kimi K3 中下界衰减对记忆保持有什么影响?
下界衰减意味着每个通道每步最多衰减到原来的 e^-5 ≈ 0.0067(即 0.67%),因此记忆不会因衰减而被瞬间清空。如果需要精确擦除,仍可通过 delta rule 的 (I - βkk^T) 实现。论文未讨论对表达力的影响,但指出这种下界门在 RWKV-7、Griffin、HGRN2 中有先例。
Kimi K3 中块内注意力矩阵 A 的计算公式是什么?
块内注意力矩阵 A 的计算公式为:A[t] = Tril[(Q[t] ⊙ Γ[t]^{1→C}) (K[t] / Γ[t]^{1→C})^T],其中 Γ 是累积衰减,Tril 表示取下三角。该矩阵用于计算块内输出 O[t] = (Γ ⊙ Q) S[t] + A[t] V~[t]。