Kimi k3 技术报告详解
前言
在看到kimi k3的技术报告后,再次被狠狠震惊到,体系架构的更新并集成到LLM中,成功实现了open frontier intelligence! 本人由于没啥机会无法体验到k3带来的震撼感,但已然从技术报告中感受到了一种大招释放的爽度,因此打算详细梳理技术报告里包含的各种细节。不足之处尽请讨论~
模型架构
架构的整体更新,取决于三个维度:
- sequence length 这一部分主要是混合注意力机制的应用,DeltaAttention的硬件成功应用使得混合注意力的效能再次拔高了一个层次,k3主要是三层KDA和一层GatedMLA,而这里面的如gate、conv等要件都已经被充分验证的有效性,因此这次属于在scale的层次上再次验证了hybird的有效性;
- network depth 这是对模型深度的探索,实际是对残余信息的充分利用,使得整体的信息量使用效率提高,且对于分布式训练环境取得了系统级的提升;
- model width 则是对MoE的进一步优化,实现Stable LatentMoE以做到高效激活对应专家信息的token
而其余如Per-Head Muon优化器或加入视觉能力,本质上是对scale能力的提升。接下来将详细讲解这些模块的细节和功能
Hybird Attention
混合注意力很显然是线性注意力和全量注意力和混合,线性注意力的O(n)使得模型能够进一步处理长序列信息,而全量注意力保留了完整的信息,通常学习全局信息。
Kimi Delta Attention
我们从全量注意力,也就是softmax开始推导: \(\begin{equation} \color{purple}{ \boldsymbol{O} = \mathop{\text{softmax}}(\boldsymbol{Q}\boldsymbol{K}^{\top} + \log \boldsymbol{M})\boldsymbol{V} } \end{equation}\)
这里省略了缩放因子,而$M$是个下三角掩码矩阵,log后对里面的元素注意取对数,以实现causal的效果。而softmax需要进行$n*n$的矩阵运算计算$exp(QK^T)$,导致了平方复杂度:
\[\begin{equation} \color{purple}{ \boldsymbol{O} = \exp(\boldsymbol{Q}\boldsymbol{K}^{\top} + \log \boldsymbol{M})\boldsymbol{V} = (\exp(\boldsymbol{Q}\boldsymbol{K}^{\top})\odot \boldsymbol{M})\boldsymbol{V} } \end{equation}\]而线性注意力就是为了近似softmax并取得线性复杂度的效果,最简单的方案就是去到指数$exp$:
\[\begin{equation} \color{purple}{ \boldsymbol{O} = (\boldsymbol{Q}\boldsymbol{K}^{\top}\odot \boldsymbol{M})\boldsymbol{V} } \end{equation}\]从分量进一步理解: \(\begin{equation} \color{purple}{ \boldsymbol{o}_t = \sum_{j=1}^t \boldsymbol{v}_j (\boldsymbol{k}_j^{\top} \boldsymbol{q}_t) = \sum_{j=1}^t (\boldsymbol{v}_j \boldsymbol{k}_j^{\top}) \boldsymbol{q}_t = \left(\sum_{j=1}^t \boldsymbol{v}_j \boldsymbol{k}_j^{\top}\right) \boldsymbol{q}_t } \end{equation}\)
我们将括号记为$S_t$,这样就有
\[\begin{equation} \color{purple}{ \boldsymbol{o}_t = \boldsymbol{S}_t \boldsymbol{q}_t, \qquad \boldsymbol{S}_t = \boldsymbol{S}_{t-1} + \boldsymbol{v}_t \boldsymbol{k}_t^{\top} } \end{equation}\]这样就是以$S_t$作为State的线性RNN形态了
接下来,我们引入delta-rule,其目的在于解决上述等权叠加导致信息稀疏的问题,核心仍然在对State $S_t$的更新:
\[\begin{equation} \color{purple}{ \begin{aligned} \boldsymbol{S}_t &= \boldsymbol{S}_{t-1} - \underbrace{(\boldsymbol{S}_{t-1} \boldsymbol{k}_t)}_{\boldsymbol{v}_t^{\text{old}}} \boldsymbol{k}_t^{\top} + \underbrace{(\beta_t \boldsymbol{v}_t + (1-\beta_t)\boldsymbol{S}_{t-1}\boldsymbol{k}_t)}_{\boldsymbol{v}_t^{\text{new}}} \boldsymbol{k}_t^{\top} \\ &= \boldsymbol{S}_{t-1}(\mathbf{I} - \beta_t \boldsymbol{k}_t \boldsymbol{k}_t^{\top}) + \beta_t \boldsymbol{v}_t \boldsymbol{k}_t^{\top} \end{aligned} } \end{equation}\]尽管形式上复杂,但概念上好理解。更新的状态就是删去旧状态并加入新状态,而里面的$\beta_t$表示写入的权重大小。整体简化后就是给旧状态$S_{t-1}$加入了Householder传输矩阵的权重,并更新新状态。
进一步,可以加入遗忘门稀疏等权信息的问题,而从SLiCE这篇文章中我们可知,低秩对角矩阵可以拥有最大的表达力,因此KDA从形式上就很清晰了:
\[\begin{equation} \color{purple}{ \mathbf{S}_t = (\mathbf{I} - \beta_t \boldsymbol{k}_t \boldsymbol{k}_t^{\top}) \operatorname{Diag}(\boldsymbol{\alpha}_t) \mathbf{S}_{t-1} + \beta_t \boldsymbol{k}_t \boldsymbol{v}_t^{\top}, \qquad \tilde{\boldsymbol{o}}_t = \mathbf{S}_t^{\top} \boldsymbol{q}_t. } \end{equation}\]Chunkwise parallel form
要实现与GPU硬件对齐的算法设计,分块并行方法已经是公认的有效方式。核心思路就是块与块之间串行递归,块内部进行并行计算。
首先,将长序列切成长度为 $C$ 的小段, $\mathbf{X}_{[t]}$ 表示第 $t$ 个块中所有向量的堆叠; $\mathbf{S}_{[t]}$ 表示进入第 $t$ 块时的递归状态。我们就会有通道级的累计衰减:
\[\begin{equation} \color{purple}{ \boldsymbol{\gamma}_{[t]}^{i \to j} := \prod_{r=i}^{j} \boldsymbol{\alpha}_{[t]}^{r}, \qquad \boldsymbol{\gamma}_{[t]}^{r} := \boldsymbol{\gamma}_{[t]}^{1 \to r}. } \end{equation}\]这里指的是从位置 i 到 j 的累积衰减,从遗忘门的对角矩阵 $\boldsymbol{\alpha}_{[t]}^r$ 得到。我们将从1到C所有的衰减按行堆叠得到 $\boldsymbol{\Gamma}_{[t]}^{1 \to C} \in \mathbb{R}^{C \times d_k}$ ,之后利用上三角变换 (UT变换),定义辅助矩阵 $\mathbf{M}_{[t]}$ 将块内所有 token 之间的衰减和 delta 修正关系编码,然后
\[\begin{equation} \color{purple}{ \mathbf{W}_{[t]} = \mathbf{M}_{[t]} (\boldsymbol{\Gamma}_{[t]}^{1 \to C} \odot \mathbf{K}_{[t]}), \qquad \mathbf{U}_{[t]} = \mathbf{M}_{[t]} \mathbf{V}_{[t]} } \end{equation}\]这样通过一次矩阵求逆就可以把原本需要逐个 token 计算的依赖关系批量处理, $\mathbf{U}_{[t]}$ 和 $\mathbf{W}_{[t]}$ 也可以在块内计算。而对于伪值项的定义
\[\color{purple}{ \widetilde{\mathbf{V}}_{[t]} := \mathbf{U}_{[t]} - \mathbf{W}_{[t]} \mathbf{S}_{[t]} }\]它就能将输出拆分成两部分:
\[\begin{equation} \color{purple}{ \begin{aligned} \mathbf{A}_{[t]} &= \operatorname{Tril}\left[(\mathbf{Q}_{[t]} \odot \boldsymbol{\Gamma}_{[t]}^{1 \to C})(\mathbf{K}_{[t]} / \boldsymbol{\Gamma}_{[t]}^{1 \to C})^{\top}\right], \\ \mathbf{O}_{[t]} &= \underbrace{(\boldsymbol{\Gamma}_{[t]}^{1 \to C} \odot \mathbf{Q}_{[t]}) \mathbf{S}_{[t]}}_{\text{inter-chunk}} + \underbrace{\mathbf{A}_{[t]} \widetilde{\mathbf{V}}_{[t]}}_{\text{intra-chunk}}. \end{aligned} } \end{equation}\]这里的$\mathbf{A}_{[t]}$是带衰减的因果注意力矩阵,用于块内注意力打分;而输出明显分成块与块之间和块内部,实现了块之间串行递归,块内部并行计算的特点,且保持数学等价性,结果和逐 token 递归完全一致。
进一步观察,$\boldsymbol{\Gamma}_{[t]}^{1 \to C}$给每个键值都进行了缩放,如果某些$\alpha$很小,整体的倒数就会趋向无穷大导致溢出。以前的做法是把每个 chunk 再细分为 16-token 的 tiles,但由于衰减的因果性对角线部分仍需要逐个位置对计算。
为了解决对角线这个瓶颈,衰减因子的映射从
\[\begin{equation} \color{purple}{ g_t^h = -e^{A_h} \text{Softplus}(z_t^h) \in (-\infty, 0) } \end{equation}\] \[\begin{equation} \color{purple}{ \alpha_t^h = \exp(g_t^h) \in (0, 1) } \end{equation}\]变化为
\[\begin{equation} \color{purple}{ g_t^h = g_{\min} \cdot \text{Sigmoid}\left(e^{A_h} z_t^h\right) \in (g_{\min}, 0) } \end{equation}\] \[\begin{equation} \color{purple}{ \alpha_t^h = \exp(g_t^h) \in (e^{g_{\min}}, 1) } \end{equation}\]Sigmoid 映射给衰减因子$\alpha$加了硬下界,倒数被锁定在安全范围内。这样所有的元素都可以通过 Tensor Core 的 dense matmul进行矩阵计算,速度更快。
Gated MLA
Multi-head Latent Attention (MLA)在我以前的博客有所涉及:DeepSeek-V2,其本质就是压缩了键值的表征,以缓解kv cache带来的大量缓存问题,而kimi k3将其应用到最后一层的全局表征中,并不使用位置编码(由于全局特征无需受到位置限制的交互),以及加入了门控机制,门控机制在Qwen的研究中已经完全验证其有效性。
最终的输出如下所示:
\[\begin{equation} \color{purple}{ \mathbf{y}_t = \mathbf{W}_o \left[ \mathrm{Sigmoid}(\mathbf{W}_g \mathbf{x}_t) \odot \mathrm{RMSNorm}(\tilde{\boldsymbol{o}}_t) \right] } \end{equation}\]其中$\mathbf{W}_g$就是门控机制,用于调制全局信息的token强度
Attention Residuals
我们首先看一下全量注意力的残差是如何设计的:
\[\begin{equation} \color{purple}{ \boldsymbol{k}_i = \boldsymbol{v}_i = \begin{cases} \boldsymbol{h}_1 & i = 0 \\ f_i(\boldsymbol{h}_i) & 1 \leq i \leq l-1 \end{cases} } \end{equation}\]其中键$k_i$值$v_i$来自前面所有层的输出,$f_i$表示第$i$个输出层,$h_1$是token的嵌入embedding,Softmax kernel被表示为$\phi(\boldsymbol{q}, \boldsymbol{k}) = \exp\left(\boldsymbol{q}^\top \mathrm{RMSNorm}(\boldsymbol{k})\right)$,全局注意力可以表示为如下形式:
\[\begin{equation} \color{purple}{ \alpha_{i \to l} = \frac{\phi(\boldsymbol{q}_l, \boldsymbol{k}_i)}{\sum_{j=0}^{l-1} \phi(\boldsymbol{q}_l, \boldsymbol{k}_j)}, \quad \boldsymbol{h}_l = \sum_{i=0}^{l-1} \alpha_{i \to l} \cdot \boldsymbol{v}_i } \end{equation}\]显然,标准注意力是在层与层之间进行,而上述的KDA这是块(block)之间的相互作用,也就是序列维度而非层维度,所以需要设计块之间的注意力残差(Block Attention Residuals)
具体而言,我们将L层分成N个块,每块$S=L/N$层,块内部通过求和压缩为单个表示为\(\color{purple}{\boldsymbol{b}_n = \sum_{j \in \mathcal{B}_n} f_j(\boldsymbol{h}_j)}\)
块间就只在N个块级做注意力:
\[\begin{equation} \color{purple}{ \mathbf{V} = \begin{cases} [\boldsymbol{b}_0, \boldsymbol{b}_1, \ldots, \boldsymbol{b}_{n-1}]^\top & \text{if } i = 1 \\ [\boldsymbol{b}_0, \boldsymbol{b}_1, \ldots, \boldsymbol{b}_{n-1}, \boldsymbol{b}_n^{i-1}]^\top & \text{if } i \geq 2 \end{cases} } \end{equation}\]总结来说,标准残差是每层只能看前一层,这里是每层可以看所有前面层,并用注意力决定谁更重要。Block 版本通过分块压缩进一步将复杂度压缩到N级别。

