从 124M 参数的 GPT-2,到 2.8 万亿参数的 KimiK3,七年时间,AI 模型在规模上膨胀了整整 22,580 倍。
但如果只把大模型的进化归结为“暴力堆算力和数据”,未免低估了这几年技术演进的精妙。从 KV Cache 瓶颈到 Linear Attention,从 DeltaNet 到 Kimi Linear 的细粒度内存管理,架构演进的本质是在不断回答同一个难题:如何让模型在有限的计算与内存下,记住最重要的信息,忘掉冗余的干扰?

今天分享的这篇 worklog,梳理了通往 KimiK3 的技术演进脉络,强烈推荐给所有对大模型底层架构感兴趣的朋友。

在这里插入图片描述

Twenty-two thousand five hundred and eighty. That’s how many GPT-2 (2019) models fit inside KimiK3 (2026).
这里是原文章的完整中文翻译,保留了原有的文本目录结构、伪代码实现以及图片插图的文字标记。


22580:从 GPT-2 到 Kimi3,详解

两万两千五百八十个。这是 KimiK3 (2026) 能容纳的 GPT-2 (2019) 模型数量。七年间,我们的规模扩大了 22,580 倍。但这仅仅是……规模扩大吗?在这篇工作日志中,我将回顾我们是如何走到今天这一步的,以及自那时以来究竟发生了多少变化(或者说变化有多小)。我们将追溯促成 KimiK3 诞生的主要架构发展历程。


1. GPT-2

GPT-2 是一种仅包含解码器(Decoder-only)的架构:

tok_emb = self.transformer.wte(idx) # 形状为 (b, t, n_embd) 的词元嵌入
pos_emb = self.transformer.wpe(pos) # 形状为 (t, n_embd) 的位置嵌入
x = self.transformer.drop(tok_emb + pos_emb)
for block in self.transformer.h:
    x = block(x)
x = self.transformer.ln_f(x)
logits = self.lm_head(x)
return logits

输入接收词元嵌入和位置嵌入:

> 📷 **[图片插图]**:
GPT-2 输入嵌入(Token Embedding + Position Embedding)结构图

每个 Transformer 模块放大后看起来是这样的:

class Block(nn.Module):
    def __init__(self, config):
        super().__init__()
        self.ln_1 = LayerNorm(config.n_embd, bias=config.bias)
        self.attn = CausalSelfAttention(config)
        self.ln_2 = LayerNorm(config.n_embd, bias=config.bias)
        self.mlp = MLP(config)

    def forward(self, x):
        x = x + self.attn(self.ln_1(x))
        x = x + self.mlp(self.ln_2(x))
        return x

> 📷 **[图片插图]**:
GPT-2 Block 结构示意图(包含 Attention 与 MLP 的残差连接)

注意力计算过程:

B, T, C = x.size() # 批次大小, 序列长度, 嵌入维度 (n_embd)

# 为批次中的所有头计算 query, key, values,并将头维度移到前面
q, k, v = self.c_attn(x).split(self.n_embd, dim=2)
k = k.view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs)
q = q.view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs)
v = v.view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, nh, T, hs)

# 注意力计算的手动实现
att = (q @ k.transpose(-2, -1)) * (1.0 / math.sqrt(k.size(-1)))
att = att.masked_fill(self.bias[:, :, :T, :T] == 0, float('-inf'))
att = F.softmax(att, dim=-1)
att = self.attn_dropout(att)
y = att @ v # (B, nh, T, T) x (B, nh, T, hs) -> (B, nh, T, hs)
y = y.transpose(1, 2).contiguous().view(B, T, C) # 重新拼接所有头的输出

# 输出投影
y = self.resid_dropout(self.c_proj(y))
return y

最终的隐藏状态矩阵生成后,语言模型头部(LM Head)会将其映射到词汇表 Logits。在自回归解码过程中,只需要最后一个位置的 Logits 即可选择下一个词元。

这是仅解码器生成方式的低效之处:模型会计算每个输入位置的表示,但每次解码步骤仅消耗最后一个位置的 Logits。如果没有缓存,很多工作都会在处理下一个词元时重复进行。

> 📷 **[图片插图]**:
自回归解码过程及最后一个位置 Logits 的提取示意图

KV 缓存(KV Cache)的出现源于一个简单的观察:在将生成的 Token 添加到输入后,模型原本需要重新计算所有先前 Token 的投影。存储它们的键(Key)和值(Value)向量可以避免这种冗余工作。

该存储空间就是键值缓存。它保留了前 N − 1 N-1 N1 个 Token 的向量,但随着序列变长,它可能变得非常庞大,从而造成内存带宽瓶颈。

总体而言,对于大约 5 万个 Token 的词表、12 个 Block、12 个头以及 768 的嵌入维度,我们的基准模型大约有 1.24 亿(124M)个参数。

vocab_size: int = 50304 # GPT-2 的词表大小为 50257,为了计算效率补齐至 64 的倍数
n_layer: int = 12
n_head: int = 12
n_embd: int = 768

而 KimiK3 模型包含 2.8 万亿(2.8T)个参数,一个模型的参数数量大约相当于 22,580 个 GPT-2 模型。


2. 线性注意力 (Linear Attention)

Softmax 注意力机制在 q ⋅ k q \cdot k qk 点积之后应用其非线性映射,将每个 Query 与每个 Key 关联起来。而线性注意力机制则分别对 q q q k k k 应用特征图(例如 ELU + 1 \text{ELU} + 1 ELU+1)。这使得乘积具备可结合律,从而可以将不断增长的 K K K V V V 向量集合折叠成一个固定的 D × D D \times D D×D 状态矩阵。

论文中关于 O ( N 2 ) \mathcal{O}(N^2) O(N2) 的表述曾让我感到困惑:“Transformer 的每个时间步的成本与当前序列长度的平方成正比”,这种说法并不完全准确。Flash Attention 正是解决了存储开销的问题……后来我发现那篇论文是在 2020 年发布的。
在当时,训练通常会具体实例化完整的 N × N N \times N N×N 注意力矩阵,FlashAttention 尚未问世,参考自回归实现通常会在没有 KV 缓存的情况下重新计算 Token 历史记录。

def forward(self, x, mask=None, past_kv=None):
    # x 形状为 b, t, d
    b, t, d = x.shape
    d_head = d // self.num_heads
    h = self.num_heads
    qkv = self.qkv_proj(x)
    q = qkv[:, :, :d].view(b, t, h, d_head).transpose(1, 2)
    k = qkv[:, :, d:2*d].view(b, t, h, d_head).transpose(1, 2)
    v = qkv[:, :, 2*d:].view(b, t, h, d_head).transpose(1, 2)

    # 在 Prefill 阶段,q, k, v 的形状为 b, h, t, d
    # 在 Decode 阶段,形状为 b, h, 1, d
    # 因此在时间维度(dim 2)进行拼接
    if past_kv is not None:
        k_past = past_kv[0]
        v_past = past_kv[1]
        k = torch.cat((k_past, k), dim=2)
        v = torch.cat((v_past, v), dim=2)

    scores = (q @ k.transpose(-1, -2)) / math.sqrt(d_head)
    if past_kv is None:
        # 处于 Prefill 阶段,需要进行因果掩码
        causal_mask = torch.ones(t, t, dtype=bool, device=q.device)
        causal_mask = torch.triu(causal_mask, diagonal=1)
        scores = scores.masked_fill(causal_mask, float('-inf'))
    if mask is not None:
        scores = scores.masked_fill(~mask, float('-inf'))

    attn = scores.softmax(-1)
    o = attn @ v
    o = o.transpose(1, 2).contiguous().view(b, t, d)
    o_proj = self.o_proj(o)
    past_kv = (k, v)
    return o_proj, past_kv

同样的流程更容易用图像来表示。每个解码步骤都会对高带宽内存(HBM)执行两次 2D 读取和两次 1D 写入,而 KV 缓存则随序列长度线性增长,时间复杂度为 O ( N ) \mathcal{O}(N) O(N)

📷 **[图片插图]**:
传统 KV Cache 在解码过程中的内存读写流转示意图

请注意过多的读写操作,线性注意力将其替换为:

def forward(self, x, mask=None, cache=None):
    # x 形状为 b, t, d
    b, t, d = x.shape
    d_head = d // self.num_heads
    h = self.num_heads
    qkv = self.qkv_proj(x)
    q = qkv[:, :, :d].view(b, t, h, d_head).transpose(1, 2)
    k = qkv[:, :, d:2*d].view(b, t, h, d_head).transpose(1, 2)
    v = qkv[:, :, 2*d:].view(b, t, h, d_head).transpose(1, 2)

    k = F.elu(k) + 1
    k = k.transpose(-1, -2)
    q = F.elu(q) + 1

    S, z = cache if cache is not None else (0.0, 0.0)
    S = S + k @ v
    z = z + k

    o = q @ S
    denom = q @ z
    o_scaled = o / denom
    o_scaled = o_scaled.transpose(1, 2).contiguous().view(b, t, d)
    o_proj = self.o_proj(o_scaled)
    cache = (S, z)
    return o_proj, cache

这其中有利有弊。这里我们将 Softmax 中使用的指数函数替换为分别应用于 q q q k k k ELU + 1 \text{ELU}+1 ELU+1,然后再进行交互。两种方法都会对结果分数进行归一化,但线性注意力机制使用的特征图是对 Softmax 核的一种表达能力较弱的近似。这种近似可能会降低模型的保真度,但实际的精度损失取决于网络架构和工作负载。

注意,我们仍然要除以 q k qk qk 的总和(为了简化图示,后续图中省略了该归一化项)。从宏观层面来看,注意力包含三个步骤:

  1. 使 q k qk qk 分数非负。线性注意力机制使用 ELU + 1 \text{ELU}+1 ELU+1,而 Softmax 机制使用指数运算。
  2. 除以总和归一化。
  3. 计算这些值的加权平均值。

这样既保留了基本的注意力契约,又使用了表达力较弱的特征图,使得 QK 分数为非负值。


3. DeltaNet(快速权重程序员 / Fast Weight Programmers)

有限的缓存必须覆盖或合并已存储的信息。来自标记 i − 1 i-1 i1 的状态不会拥有自己的独立槽位;它会被累加到同一个 D × D D \times D D×D 矩阵中。因此,新的查询无法再检索到每个先前标记的完全独立的表示。

这一改进也是效率提升的来源。采用累加式而非拼接式的方式更新缓存,可以避免缓存以 O ( N ) \mathcal{O}(N) O(N) 的速度增长,但同样的操作也会导致信息相互干扰。DeltaNet 解决了这种可恢复性的损失。

📷 **[图片插图]**:

Schlag 在其论文《Fast Weight Programmers》中精辟地指出:“当序列长度超过存储容量时,模型可能会陷入容量过载状态。为了在这种状态下正确运行,模型应该学习动态地与内存内容交互,并有选择地决定保留哪些键值关联,删除哪些关联。纯粹的加法指令可能不适用于此目的……不断地向有限大小的内存中添加新的关联,必然会达到一个极限。”

N ≫ D N \gg D ND 时,线性注意力机制便会变得极具吸引力,但这同时也暴露了它的主要局限性。一旦状态超过其有效容量,关联就会开始相互干扰,因为更新是纯累加的,缓存中不会淘汰旧数据。

def forward(self, x, mask=None, cache=None):
    # x 形状为 b, t, d
    b, t, d = x.shape
    d_head = d // self.num_heads
    h = self.num_heads
    qkv = self.qkv_proj(x)
    q = qkv[:, :, :d].view(b, t, h, d_head).transpose(1, 2)
    k = qkv[:, :, d:2*d].view(b, t, h, d_head).transpose(1, 2)
    v = qkv[:, :, 2*d:].view(b, t, h, d_head).transpose(1, 2)

    q = F.normalize(F.silu(q), dim=-1)
    k = F.normalize(F.silu(k), dim=-1)
    beta = torch.sigmoid(self.w_beta(x)).view(b, 1, t, 1) # 新增:每个 Token 的写入强度

    S = cache if cache is not None else 0.0
    v_old = k @ S                   # 在当前 key 位置读取已有记忆
    u = beta * (v - v_old)          # 计算 Delta:仅保留真正的新信息
    S = S + k.transpose(-1, -2) @ u # 外积写入矩阵状态
    o = q @ S                       # 读取,无需分母归一化
    o = o.transpose(1, 2).contiguous().view(b, t, d)
    return self.o_proj(o), S

用图示说明更容易理解:

📷 **[图片插图]**:*Delta 规则的擦除与新写入机制(从已有记忆中减去旧值  后写入新值 )*

考虑一个关联关系 S = k T ⋅ v S = k^T \cdot v S=kTv。如果使用相同的键读取,则得到 k ⋅ ( k T ⋅ v ) = ( k ⋅ k T ) v k \cdot (k^T \cdot v) = (k \cdot k^T) v k(kTv)=(kkT)v,即 k k k 乘以 v v v 的范数。因此,读取结果会按键的范数进行缩放,如果将 k k k 归一化为单位长度,就能精确地得到 v v v

Q Q Q 也是一个已学习的指针。 W q W_q Wq W k W_k Wk 读取同一个残差流,对某个事实的查询指向该事实写入的键方向。更新操作首先查询当前键从缓存中检索到的信息。它从要存储的值中减去现有信息,将键乘以差值,然后将结果加回去。旧信息被移除,新信息被写入其位置。


4. DeltaNet(使用 Delta 规则并行化线性 Transformer)

这是本文中最难理解的部分。我花了大约七个小时才弄明白,所以我会从实现层面来解释。

简而言之,DeltaNet 实现了一个带有广义 Householder 变换矩阵的一阶线性递归,从而能够进行分块并行(Chunkwise Parallel)前向传播,实现硬件高效的线性时间训练。它将输入和输出分割成若干个大小为 C C C 的块,并根据前一个块的最终状态和当前块的 Query/Key/Value 块来计算每个块的输出。

实际问题在于 Prefill(预填充阶段)。对 T T T 个标记序列直接应用 Delta 规则的顺序代码如下所示:

S = torch.zeros(b, h, dh, dh) if cache is None else cache
outs = []
for i in range(t):
    k_i = k[:, :, i:i+1]
    v_i = v[:, :, i:i+1]
    b_i = beta[:, :, i:i+1]
    v_old = k_i @ S
    u_i = b_i * (v_i - v_old)
    S = S + k_i.transpose(-1, -2) @ u_i # 写入状态
    outs.append(q[:, :, i:i+1] @ S)
o = torch.cat(outs, dim=2)

与标准注意力机制不同,这种公式需要对每个 Key 向量进行修正,因此并行矩阵乘法的实现路径并不显而易见。即使没有 Delta 规则,直接的线性注意力预填充仍然是顺序执行的:

S = torch.zeros(b, h, dh, dh) if cache is None else cache
outs = []
for i in range(t):
    q_i = q[:, :, i:i+1]
    k_i = k[:, :, i:i+1]
    v_i = v[:, :, i:i+1]
    S = S + k_i @ v_i
    o = q_i @ S
    o = self.norm(o)
    o = o.transpose(1, 2).contiguous().view(b, t, d)
    out = self.o_proj(o)
    outs.append(out)
o = torch.cat(outs, dim=2)

分块式公式提供了一种更高效的方法。通过示例可以更容易地理解其机制:

在这里插入图片描述

设置 C = N C=N C=N 可还原标准的 O ( N 2 ) \mathcal{O}(N^2) O(N2) 注意力机制,而 C = 1 C=1 C=1 则提供常规的线性注意力机制。我们通过插值的方式,在中间值之间进行权衡,以牺牲额外的块内工作量来换取更好的硬件利用率。实际上, C C C 通常取 64 或 128,因为张量核心(Tensor Core)指令在该粒度下运行效率更高。

中间图块在状态更新过程中被折叠成 S S S

在这里插入图片描述

S = torch.zeros(b, h, dh, dh) if cache is None else cache
outs = []
for i in range(t // C):
    q_c = q[:, :, i*C:(i+1)*C]
    k_c = k[:, :, i*C:(i+1)*C]
    v_c = v[:, :, i*C:(i+1)*C]

    o_prev = q_c @ S                            # 当前块之前累积的所有状态输出
    attn = (q_c @ k_c.transpose(-1, -2)).tril() # 块内的因果掩码注意力
    o_curr = attn @ v_c
    o = o_prev + o_curr

    S_new = k_c.transpose(-1, -2) @ v_c         # 循环更新状态
    S = S + S_new
    outs.append(o)
o = torch.cat(outs, dim=2)

在每个块内,我们执行 q ( k T v ) q(k^T v) q(kTv)。这是先处理分数,也就是带掩码的常规注意力顺序。跨块执行时,我们遵循 ( k T v ) q (k^T v)q (kTv)q,也就是先处理状态,遵循循环顺序。注意力机制的复杂度为 O ( N 2 ) \mathcal{O}(N^2) O(N2),而此操作的复杂度则不会。在每个块内,我执行真正的注意力操作(掩码后的 Q K T × V QK^T \times V QKT×V),而在跨块执行时,我将所有内容折叠到状态中,然后通过一次矩阵乘法将其读出。

因此,成本分为两部分:

  • 一部分是固定的 2 L d 2 2Ld^2 2Ld2,用于处理状态,完全不考虑 C C C
  • 另一部分是增长的 2 L C d 2LCd 2LCd,用于处理位于对角线上的分数矩阵。

完全注意力机制恰好是 C = L C=L C=L 的情况,此时第二项变为 2 L 2 d 2L^2d 2L2d(即二次方)。因此, C C C 越小,执行的 FLOPs 就越少。从纯粹的浮点运算次数(FLOPs)来看, C = 1 C=1 C=1 是最经济的选择,但实际运行时间未必最短。当运算任务能够高效地映射到 GPU 的矩阵乘法硬件上时,GPU 可以更快地完成更多运算。

下一步是将相同的方法推广到 DeltaNet。

在这里插入图片描述

根本问题很简单:用于纯粹加性注意力的分组方法并不直接适用于增量更新:

v_old = k_i @ S                  
u_i = b_i * (v_i - v_old)

我们需要每个状态才能计算出需要减去的信息。如果不进行一些数学上的重新参数化,我们就无法以相同的方式并行化它。因此,作者重写了以下增量更新:

u = v_new - v_old
S_t = S_{t-1} + K^T @ u
o = q @ S_T

这里,顺序循环每次迭代计算一个增量。重新参数化后的形式为:

S t = S t − 1 ( I − β t k t k t T ) + β t v t k t T S_t = S_{t-1}(I - \beta_t k_t k_t^T) + \beta_t v_t k_t^T St=St1(IβtktktT)+βtvtktT

o t = S t q t o_t = S_t q_t ot=Stqt

这种写法使得分块代码能够一次性计算所有 C C C 个增量:

def chunk_delta_rule_forward(Q, K, V, beta, C):
    # L: 序列长度, d: 头维度
    L, d = Q.shape
    # 分块处理
    Q, K, V = map(lambda x: x.reshape(-1, C, d), [Q, K, V])
    beta = beta.reshape(-1, C)
    K_beta = K * beta.unsqueeze(-1)
    V_beta = V * beta.unsqueeze(-1)
    
    # 使用向量化的前向替换计算方程 (10),以便快速求逆
    T = -(K_beta @ K.t()).tril(-1)
    for i in range(1, C):
        T[i, :i] = T[i, :i] + (T[i, :, None] * T[:, :i]).sum(-2)
    
    T += torch.eye(C)
    W = T @ K_beta
    U = T @ V_beta

    # 块间并行,计算方程 (8-9)
    S = torch.zeros(d, d)
    O = torch.empty_like(V)
    
    for i in range(L // C):
        q_i, k_i, w_i = Q[i], K[i], W[i]
        u_i = U[i] - w_i @ S          # 计算当前块的所有修正量
        o_inter = q_i @ S
        A_i = (q_i @ k_i.t()).tril()   # qk.t
        o_intra = A_i @ u_i            # attention @ v (带有修正量 u)
        S += k_i.t() @ u_i             # 使用加法更新状态 
        O[i] = o_intra + o_inter       # 更新输出:Flash + Recurrent 混合
    return O.reshape(L, d)

这就引出了我们的第一个比较点:MHA 与 DeltaNet Transformer 的对比:
在这里插入图片描述


5. 门控 DeltaNet (Gated DeltaNet)

我们现在有了一种对缓存进行精确修改的方法。对于每个新的事实(每个新的 Key 向量),我们都可以精确地查看当时存储的旧信息,并将其替换为我们想要关注的新信息。

然而,这种机制只能遗忘那些有特定替代项的关联。它无法在上下文切换期间有效地清除多个关联,也无法普遍地衰减内存以释放容量。

如果我们只进行纯粹的加性线性注意力,添加遗忘功能很简单。我们只需要一个控制遗忘状态的参数:

S_old = cache
S_new = k @ v
# cache = S_old + S_new
cache = alpha * S_old + S_new

在这里插入图片描述

这是 Mamba-2 的贡献。我们先将之前的缓存失效,然后再添加新的缓存,从而防止状态无限增长。

在每个时间步长内,以动态比例均匀衰减所有键值关联是一种可行的方法,Mamba 就是这么做的。但这并没有考虑到不同键值关联重要性的差异。也就是说,如果模型需要遗忘某个特定的关联,那么所有关联都会被同等地遗忘。相比之下,Delta 规则可以更新单个事实,但无法使其余事实失效。

因此,门控增量规则(Gated Delta Rule)结合了 Mamba 的门控更新规则和 Delta 规则。它增加了一个参数 α \alpha α,当 α = 1 \alpha=1 α=1 时切换到纯增量规则,当 α = 0 \alpha=0 α=0 时清除内存。难点在于如何使用相同的并行分块方法来实现这一点。

该实现采用了与上一节所述相同的 DeltaNet 重参数化方法。其数学原理几乎完全相同,仅增加了一个与数据相关的标量(介于 0 和 1 之间),用于控制先前状态的衰减。这使得高效的键值关联学习与自适应记忆管理相结合。

相应的代码更改如下所示:

在这里插入图片描述

γ r / γ i \gamma^r / \gamma^i γr/γi 项用于描述累积衰减。在时间步 x x x 写入并在 x + t x+t x+t 读取的标记已乘以 α x α x + 1 α x + 2 … α x + t \alpha_x \alpha_{x+1} \alpha_{x+2} \dots \alpha_{x+t} αxαx+1αx+2αx+t。这相当于前缀和计算的乘法形式。

最终的架构如下所示:
在这里插入图片描述


6. KDA / Kimi Linear

此时,研究人员开始尝试将多种形式的注意力机制结合在一个架构中的混合模型,例如 Gated DeltaNet 与 Mamba。

Kimi Linear 因其一个核心优势而备受关注:在受控对比条件下,它的性能优于完全注意力机制。作者将其定位为一种可直接替换的架构,具有更高的质量和高达 6 倍的解码吞吐量。

Kimi Linear 通过引入细粒度门控改进了 Gated DeltaNet。它不再使用单一的标量衰减,而是为每个通道学习一个单独的衰减值。

在这里插入图片描述

KDA 更新规则保持不变,但代码现在看起来更像这样:

在这里,alpha.reshape(nb, C, d) 体现了本文最重要的贡献:对内存衰减进行精细控制。

与 DeltaNet Transformer 并列,Kimi Linear 架构引入了三个主要变化:

  1. 它采用混合系统,交错使用多头潜在注意力(MLA, Multi-Head Latent Attention)层。
  2. 它用混合专家(MoE, Mixture of Experts)层代替了 MLP。
  3. 它通过 α \alpha α 预测为 DeltaNet 增加了容量。

在这里插入图片描述

后续章节将更详细地介绍 MLA 和 MoE。目前,重点在于这并非盲目扩展。新增容量具有特定的数学目的:按通道扩展使模型能够更精细地控制内存衰减。

扩展规律(Scaling Laws)依然适用,但容量的增加必须在正确的位置,并以系统能够利用的形式进行。这一发展过程中的每一种架构都增加了容量,以解决前一系统中的具体限制。


7. Kimi K3

最终,KimiK3 语言骨干网络与上述 Kimi Linear 模型类似。它包含 23 个四层宏环(Macro-loops)。在每个宏环中,三层使用 Kimi Delta Attention (KDA),第四层使用多头潜在注意力机制 (MLA)。第一层使用密集前馈网络(Dense FFN);其余各层均使用潜在混合专家模型(Latent MoE)。

乍一看,Kimi Linear 的变化似乎并不大:

  • 规模大幅增加
  • 每 12 层进行块状 AttnRes 计算
  • MLA 查询 LoRA 和输出门控
  • 潜在空间 MoE (Latent-space MoE)
  • SiTU 激活函数
  • 封闭式 MLA

KDA 提供恒定状态的循环记忆,而周期性 MLA 层则保留了对上下文的完整 Softmax 检索。以下简化的可视化图为下文讨论的变化提供了有用的参考:

在这里插入图片描述

我们将从更直接的变化开始:门控 MLA、潜在空间 MoE 和 SiTU 激活。

门控 MLA 决定了每个提取特征有多少从 MLA 传递到残差流中。它通过将每个元素与从输入投影出的门进行逐元素乘法来实现这一点。

在传统的路由算法(MoE)中,学习到的路由器利用点积相似度将每个词元发送给一部分专家网络。KimiK3 总共有 898 位专家。其中两位专家共享并处理每个词元;在剩余的 896 位专家中,路由器为每个词元选择 16 位专家。

KimiK3 还改变了专家激活算法。它不再像以前那样对上投影应用 SiLU,将其逐元素乘以门函数,然后再应用下投影,而是使用 SiTU

d = x.shape[-1] // 2
gate = x[..., :d].to(torch.float32)
up = x[..., d:].to(torch.float32)

situ_a = self.beta * torch.tanh(gate / self.beta) * torch.sigmoid(gate)
if self.linear_beta is not None:
    up = self.linear_beta * torch.tanh(up / self.linear_beta)

return (situ_a * up).to(x.dtype)

该模型还会向下预测分配给共享专家的输入值,并向上预测他们的最终总和:

在这里插入图片描述

这揭示了模型推理中一个反复出现的挑战:如果没有融合算子(Fused Kernel),新的激活函数比原始路径慢近 3 倍。一种补偿性的优化方法是,专家模型在一个压缩的潜在空间中运行,这使得它们的前向传播速度更快,并且几乎将浮点运算次数减半。

其余的改动包括 MLA 查询 LoRA、输出门控以及每 12 层使用分块注意力残差(Blockwise AttnRes)。注意力残差会增加大约 2% 的推理延迟,但带来两个重要的好处:

  • 选择性地检索早期表征,可以减轻残余稀释(Residual Dilution)和隐藏态增长。
  • 带来 1.25 倍的计算优势。

AttnRes 和 MLA 从不同的角度解决了同一个根本限制。KDA 层使用固定大小的状态,因此不可避免地会丢弃信息。MLA 从词元上下文中检索信息,而 AttnRes 则从更早的深度表示中检索信息。


8. AttnRes(注意力残差)

在每次前向传播过程中,输入会依次经过一系列层。每一层都包含一个注意力模块(KDA 或 MLA)和一个 MLP 或 MoE 模块。通常,每一层的输入是原始嵌入向量与之前所有层的输出之和,所有权重相等:

h i = h 1 + ∑ j = 1 i − 1 f j ( h j ) h_i = h_1 + \sum_{j=1}^{i-1} f_j(h_j) hi=h1+j=1i1fj(hj)

这里, h i h_i hi 是第 i i i 层的输入, h 1 h_1 h1 是当前标记(到目前为止序列中的最后一个标记)的嵌入, f j ( h j ) f_j(h_j) fj(hj) 是第 j j j 层的输出(一个注意力或 MLP 模块)。

问题在于缺乏选择性访问。不同类型的层接收到相同的聚合状态,即使它们可能受益于不同的权重。由于循环是纯粹的加性运算,后续层必须学习越来越大的输出才能影响累积残差,这可能会破坏训练的稳定性。AttnRes 并没有平等对待所有层,而是将总和中的每一项乘以一个特定的权重,这使得模型能够根据上下文赋予最有用的层更高的权重:

h i = a 0 ⋅ h 1 + ∑ j = 1 i − 1 a j ⋅ f j ( h j ) h_i = a_0 \cdot h_1 + \sum_{j=1}^{i-1} a_j \cdot f_j(h_j) hi=a0h1+j=1i1ajfj(hj)

每个权重 a j a_j aj 由查询与键值之间的点积计算得出。查询是针对每一层学习得到的,而键和值则来自之前的残差流状态。得分被归一化为总和为 1,然后用于形成这些状态的加权组合。在这里插入图片描述

因此,该模型不必仅依赖于其直接前一层。AttnRes 允许每一层选择性地访问先前层的输出,从而使其学习到的查询能够检索出对当前计算最有用的表示。

以下伪代码以块为单位应用了相同的思想。一个块是注意力机制和多层感知器(MLP)输出在 12 个解码器层上的逐元素求和,并存储为一个单一的深度表示,以便后续进行 AttnRes 混合。

在每一层都应用残差注意力会增加过多的训练和推理成本。仅在固定的块边界应用残差注意力,可以在较低的成本下获得大部分收益。在 KimiK3 中,每个边界出现在 12 个解码器层之后。在 23 个四层宏循环中,这会产生 8 个 AttnRes 块,从而提高推理速度。

这可能是 block_attn_res 函数中最重要的部分:

V = torch.stack(blocks + [partial_block]) # [N+1, B, T, D]
K = norm(V)
logits = torch.einsum('d, n b t d -> n b t', proj.weight.squeeze(), K)
h = torch.einsum('n b t, n b t d -> b t d', logits.softmax(0), V)
return h

结语

至此,从 GPT-2 到 KimiK3 的演进过程就完成了。

核心变化不仅仅在于规模。架构上的每一个步骤都会改变模型存储的内容、更新状态的方式,或者检索固定大小状态无法保存的信息的方式。

KimiK3 结合了常态循环记忆、周期性 Softmax 检索、稀疏专家容量和选择性深度残差访问。其结果是,该系统能够将额外的容量用于具有特定功能作用的地方。

本质上,容量固定的联想记忆(固定维度)需要一种淘汰策略(Eviction Policy),因为纯粹的线性加法运算一旦达到容量上限,最终会引入干扰。为此,学习到的选择机制(例如门控、路由或衰减)是必要的,而注意力则是最有效的选择性读取机制。


Scaling laws 依然奏效,但单纯的规模扩张已经不再是唯一的解法。正如文章所展示的,KimiK3 并非盲目堆砌参数,而是将每一份容量精确分配到了特定的功能组件上——结合常数状态循环记忆、周期性 Softmax 检索、稀疏 MoE 以及深度 Residual 访问。

了解这些架构变迁,能让我们在面对“更大”的模型时,看清真正支撑起它能力跃迁的技术内核。全文包含大量数学推导与 PyTorch 伪代码,非常值得精读与收藏!

原文链接:
22580: From GPT2 to Kimi3, Explained

更多推荐