从 GPT-2 到 Kimi K3:大语言模型架构演进全解

参考来源:https://x.com/waterloo_intern/status/2081762065392541951

本文为原英文技术报告的完整中文翻译与博客化整理。原文图片链接、代码块、代码内容和公式顺序均保持不变;专业名称首次出现时保留英文原名,便于对照学习。


二万二千五百八十。一个 Kimi K3(2026)的参数量,大致相当于 22,580 个 GPT-2(2019)模型。七年间,我们将模型规模扩大了 22,580 倍。但这一切真的只是……规模变大了吗?

在这篇开发记录中,我会带你回顾我们是如何走到今天的,并分析自 GPT-2 以来,模型架构究竟发生了多少变化——或者说,实际上有多少部分并未改变。我们将沿着通往 Kimi K3 的路线,梳理其中最重要的架构演进。

GPT-2

GPT-2 采用仅解码器(Decoder-only)架构:

tok_emb = self.transformer.wte(idx) # token embeddings of shape (b, t, n_embd)
pos_emb = self.transformer.wpe(pos) # position embeddings of shape (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

输入首先经过词元嵌入(Token Embedding)和位置嵌入(Positional 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

注意力计算过程如下:

B, T, C = x.size() # batch size, sequence length, embedding dimensionality (n_embd)

        # calculate query, key, values for all heads in batch and move head forward to be the batch dim
        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)

        # manual implementation of attention
        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) # re-assemble all head outputs side by side

        # output projection
        y = self.resid_dropout(self.c_proj(y))
        return y

得到最终的隐藏状态矩阵后,语言模型头会将其映射为整个词表上的 Logits。在自回归解码过程中,模型只需要使用最后一个位置的 Logits 来选择下一个词元。

这体现了仅解码器生成方式的一项低效之处:模型会为输入中的每一个位置计算表示,但每一步解码实际只使用最后一个位置的 Logits。如果没有缓存机制,在生成下一个词元时,其中大量计算都需要重新执行。

KV 缓存源于一个非常直接的观察:将新生成的词元追加到输入后,如果不使用缓存,模型就必须重新计算此前所有词元的投影。保存这些词元对应的 Key 和 Value 向量,可以避免这部分重复计算。

这块存储空间就是 KV 缓存。它保存此前 N − 1 N-1 N1 个词元的向量,并且可能增长到足以形成显存带宽瓶颈的规模。

总体来看,在词表规模约为 5 万、包含 12 个模块、12 个注意力头、嵌入维度为 768 的情况下,这个基线模型约有 1.24 亿个参数。

vocab_size: int = 50304 # GPT-2 vocab_size of 50257, padded up to nearest multiple of 64 for efficiency
n_layer: int = 12
n_head: int = 12
n_embd: int = 768

Kimi K3 拥有 2.8 万亿个参数,其参数总量大致相当于 22,580 个 GPT-2 模型。

线性注意力(Linear Attention)

Softmax 注意力在计算 q ⋅ k q\cdot k qk 乘积之后才施加非线性,因此每一个 Query 都会与每一个 Key 耦合。线性注意力则分别对 q q q k k k 应用特征映射,例如 E L U + 1 \mathrm{ELU}+1 ELU+1。这样一来,矩阵乘法就可以重新结合,持续增长的 K K K V V V 向量集合也就能够被折叠进一个固定大小的 D × D D\times D D×D 状态矩阵中。

论文中关于 O ( N 2 ) O(N^2) O(N2) 的表述一开始让我有些困惑。“Transformer 每个时间步的计算成本会随着当前序列长度的平方增长”并不完全准确——这正是 FlashAttention 要解决的问题……随后我才注意到,这篇论文发表于 2020 年。

在当时的训练实践中,人们通常会显式构造完整的 N × N N\times N N×N 注意力矩阵;FlashAttention 尚未出现,而许多参考级自回归实现也经常在没有 KV 缓存的情况下重复计算整个词元历史。

def forward(self, x, mask=None, past_kv=None):
  # x is 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)

  # at prefill, q,k,v have shapes b,h,t,d
  # at decode, shape is b, h, 1, d
  # so i cat at the t dimension, 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: #we're in prefill and need to mask
    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'))

  #get attn (bhtt x bhtd)
  attn=scores.softmax(-1)#bhtt
  o=attn@v #bhtd
  o=o.transpose(1,2).contiguous().view(b,t,d)  #b,t,d

  # use x to get qkv
  o_proj=self.o_proj(o)
  past_kv=(k, v)
  return o_proj, past_kv

把这一过程画出来会更容易理解。每一步解码都会从高带宽显存(HBM)执行两次 ND 读取和两次一维写入,同时 KV 缓存会随着序列长度以 O ( N ) O(N) O(N) 的速度线性增长。

可以注意到其中存在大量读写操作,而这篇论文将其替换为下面的形式:

def forward(self, x, mask=None, cache=None):
  # x is 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 #bhtd
 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 相互作用之前,分别对它们应用 E L U + 1 \mathrm{ELU}+1 ELU+1。两种方法都会对最终得到的分数进行归一化,但线性注意力使用的特征映射只是对 Softmax 核的一种表达能力较弱的近似。这种近似可能会降低结果保真度,不过实际精度损失取决于具体架构和工作负载。

需要注意的是,我们仍然会除以 q k qk qk 分数之和,只是为了简化图示,该步骤没有画出来。从更高层次看,注意力计算可以分为三个步骤:

  1. q k qk qk 分数变为非负数。线性注意力使用 E L U + 1 \mathrm{ELU}+1 ELU+1,而 Softmax 使用指数运算。
  2. 除以所有分数之和,完成归一化。
  3. 计算 Value 的加权平均值。

这种做法保留了注意力机制的基本功能约定,只是使用了表达能力较弱的特征映射来保证 Q K QK QK 分数非负。

DeltaNet(快速权重程序器,Fast Weight Programmers)

容量有限的缓存必须覆盖已有信息,或者将新信息与已有信息合并。词元 i − 1 i-1 i1 的状态不会获得一个独立存储槽位,而是被加入同一个 D × D D\times D D×D 矩阵中。因此,新的 Query 将无法再从中检索出每一个早期词元彼此完全隔离的表示。

不过,这种相加操作也正是效率提升的来源。通过加法而不是拼接来更新缓存,可以避免缓存按照 O ( N ) O(N) O(N) 持续增长;但同一个操作也会造成信息之间相互干扰。DeltaNet 试图解决的,正是这种信息难以恢复的问题。


Schlag 在论文《快速权重程序器》(Fast Weight Programmers)中对此有一段非常精炼的描述:“当序列长度超过存储容量时,模型可能进入一种容量超载状态。为了在这种状态下正常工作,模型应当学会与记忆内容进行动态交互,并有选择地决定保留哪些键值关联、删除哪些键值关联。纯粹的加法式写入指令可能并不适合这一目的……正如公式 17 所示,在容量有限的记忆中无休止地加入新关联,最终必然会触及极限。”

使线性注意力具有吸引力的场景——也就是 N ≫ D N\gg D ND——同时也暴露了它最主要的局限。当状态超过其有效容量后,由于更新方式是纯加法,并且没有任何内容会离开缓存,不同关联就会开始彼此干扰。

def forward(self, x, mask=None, cache=None):
  # x is 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)   
  # new: per-token write strength

  S = cache if cache is not None else 0.0  

  v_old = k @ S # read the board at this key
  u = beta * (v - v_old) # the delta: only what's actually new
  S = S + k.transpose(-1, -2) @ u # same outer-product write as before

  o = q @ S # read, no denominator
  o = o.transpose(1, 2).contiguous().view(b, t, d)
  return self.o_proj(o), S

通过一个可视化示例,会更容易理解这一过程。

假设我们将一个关联写成 S = k ⊤ v S=k^\top v S=kv。如果使用同一个 Key 将它读出,就会得到 k ( k ⊤ v ) k(k^\top v) k(kv),也就是 ( k k ⊤ ) v (kk^\top)v (kk)v;其中 k k ⊤ kk^\top kk 等于 k k k 的范数平方。因此,读取得到的结果会被 Key 的范数平方所缩放。如果将 k k k 归一化为单位长度,或者直接用结果除以该范数,就能够精确恢复 v v v

Q Q Q 同样是一个通过学习得到的指针。 W q W_q Wq W k W_k Wk 读取同一条残差流,而针对某条事实的 Query 会指向写入该事实时所使用的 Key 方向。更新操作首先询问:当前 Key 能够从缓存中读出什么信息?随后,它从我们希望写入的 Value 中减去这部分已有信息,将 Key 与这个差值相乘,再把结果加回状态中。这样,旧信息会被移除,新信息则被写入原来的位置。

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

这是全文最难理解的一部分。为了形成一个能够真正用于解释的理解,我大约花了七个小时,因此下面会从具体实现出发逐步展开。简要来说,DeltaNet 使用广义 Householder 转移矩阵实现一阶线性递归,从而支持按块并行的前向传播,并实现对硬件友好的线性时间训练。它将输入和输出划分为若干个大小为 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 # write
    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 = q[:, :, i:i+1]  
    k = k[:, :, i:i+1]  
    v = v[:, :, i:i+1]

    S=S_old+k@v
      o=q@S #bhtd
      o=self.norm(o)
    o=o.transpose(1, 2).contiguous().view(b, t, d)

    out=self.o_proj(o)
    cache=S
    outs.append(out)

o = torch.cat(outs, dim=2)

分块形式能够提供一种更高效的实现方式。通过下面的示例可以更直观地理解其运行机制:

C = N C=N C=N 时,就会退化为标准的 O ( N 2 ) O(N^2) O(N2) 注意力;当 C = 1 C=1 C=1 时,则得到普通的线性注意力。介于两者之间的 C C C 值,会以增加块内计算量为代价,换取更高的硬件利用率。实践中, C C C 通常取 64 或 128,因为 Tensor Core 指令能够在这一粒度上高效运行,UMMA 就是其中一个例子。

在状态更新过程中,中间计算得到的各个 Tile 会被折叠进 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 #this is everything up to this block
      
      attn=(q_c@k_c.transpose(-1,-2)).tril() #masked attention 
      o_curr=attn@v_c
          
        o=o_prev+o_curr
    
    S_new=k_c.transpose(-1,-2)@v_c #recurrent attention 
    S=S+S_new
    outs.append(o)

o = torch.cat(outs, dim=2)

在一个分块内部,我们计算 q ( k ⊤ v ) q(k^\top v) q(kv)。这里先计算分数,采用带掩码的常规注意力顺序;而在不同分块之间,我们按照 ( k ⊤ v ) q (k^\top v)q (kv)q 的顺序计算,也就是先形成状态,再执行递归读取。标准注意力的成本会按照 O ( N 2 ) O(N^2) O(N2) 增长,而这种方法不会。在分块内部,我计算真实的注意力,即带掩码的 Q K ⊤ V QK^\top V QKV;在分块之间,则把所有信息折叠进状态,再通过一次矩阵乘法将其读出。因此,成本可以拆成两部分:第一部分是固定成本 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)就越少。

从纯 FLOP 数量来看, C = 1 C=1 C=1 是成本最低的选择,但它在实际运行时间上未必最快。当计算任务能够高效映射到 GPU 的矩阵乘法硬件时,GPU 往往可以在更短时间内完成更多算术运算。

下一步,就是将同样的方法扩展到 DeltaNet。

根本问题很简单:用于纯加法注意力的分块方法,不能直接应用到 Delta 更新上:

v o l d = k i @ S u i = b i ∗ ( v i − v o l d ) v_old = k_i @ S u_i = b_i * (v_i - v_old) vold=ki@Sui=bi(vivold)

为了计算需要减去的信息,我们必须获得每一个中间状态。如果不进行某种数学重参数化,就无法用同样的方式实现并行化。因此,作者将 Delta 更新从下面的形式:

u = v n e w − v o l d S t = S ( t − 1 ) + K . T @ u o = q @ S T u=v_new-v_old S_t= S_(t-1)+K.T@u o=q@S_T u=vnewvoldSt=S(t1)+K.T@uo=q@ST

在这种形式中,串行循环每次迭代只能计算一个 Delta。经过重参数化后,形式变为:

S t = S t − 1 ( I − β t k t k t T ) + β t v t k t T o t = S t q t S_t = S_{t-1}(I − β_t k_t k_tᵀ) + β_t v_t k_tᵀ o_t = S_t q_t St=St1(IβtktktT)+βtvtktTot=Stqt

这种形式使分块代码能够一次性计算当前分块中的全部 C C C 个 Delta:

def chunk_delta_rule_forward(Q, K, V, beta, C):
        # L: sequence length, d: head dimension
        L, d = Q.shape
        # chunking
        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)
        
        # compute eq. 10 with vectorized forward substitution for fast inverse
        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

        # chunkwise parallel. Eq. 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 # the corrections, all of one chunk
                o_inter = q_i @ S
                A_i = (q_i @ k_i.t()).tril() #qk.t
                o_intra = A_i @ u_i # attention @ v (with corrections, so u)
                S += k_i.t() @ u_i # update state with addition 
                O[i] = o_intra + o_inter #update output with flash + recurrent
        return O.reshape(L, d)

至此,我们得到了第一个可以直接比较的节点:多头注意力(MHA)Transformer 与 DeltaNet Transformer:

门控 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 规则能够更新单独一条事实,却无法让其余事实自然衰减。

因此,门控 Delta 规则将 Mamba 的门控更新规则与 Delta 规则结合起来。它引入参数 α \alpha α:当 α = 1 \alpha=1 α=1 时,机制退化为纯 Delta 规则;当 α = 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}\cdots\alpha_{x+t} αxαx+1αx+2αx+t。这可以看作前缀和计算在乘法形式下的对应版本。

最终得到的架构如下:

KDA / Kimi Linear

发展到这里,研究人员开始探索在同一架构中结合多种注意力形式的混合模型,例如将 Gated DeltaNet 与 Mamba 结合。

Kimi Linear 因一个核心结论而受到关注:在受控对比实验中,它的表现超过了完整注意力。作者将其描述为一种可直接替换现有注意力模块的架构方案,不仅质量更高,解码吞吐量最高还能提升到 6 倍。

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

KDA 的更新规则仍然相似,但代码形式更接近下面这样:

其中,alpha.reshape(nb, C, d) 体现了论文最重要的贡献:对记忆衰减进行细粒度控制。

与 DeltaNet Transformer 对比来看,Kimi Linear 架构引入了三项主要变化:

  1. 采用混合系统,在网络中交错插入多头潜在注意力(Multi-head Latent Attention,MLA)层。
  2. 用混合专家(Mixture-of-Experts,MoE)层替换 MLP。
  3. 通过 α \alpha α 投影提升 DeltaNet 的容量。

后文会更详细地介绍 MLA 和 MoE。此处首先要理解的是:这并不是盲目扩大模型规模。新增容量具有明确的数学目的——按通道设置缩放因子,使模型能够更细致地控制记忆衰减。

扩展定律依然重要,但容量必须被添加在正确的位置,并采用系统真正能够利用的形式。这条演进路线中的每一种架构,都通过增加特定能力来解决前一代系统中一个具体的限制。

Kimi K3

最终,Kimi K3 的语言模型主干与上面的 Kimi Linear 模型较为相似。它包含 23 个由四层组成的宏循环。在每个宏循环中,前三层使用 Kimi Delta Attention,第四层使用多头潜在注意力。第一层采用稠密前馈网络,其余所有层都采用潜在空间混合专家结构。

乍看之下,与 Kimi Linear 相比,Kimi K3 的变化似乎并不算多:

  • 模型规模大幅提升
  • 每 12 层加入一次分块式 AttnRes
  • MLA Query LoRA 与输出门控
  • 潜在空间 MoE
  • SiTU 激活函数
  • 门控 MLA

KDA 提供状态大小恒定的递归记忆,而周期性出现的 MLA 层则保留了面向完整上下文的 Softmax 检索能力。下面这张简化架构图可以作为后续理解各项改动的参考。


我们先从更直接的几项变化讲起:门控 MLA、潜在空间 MoE,以及 SiTU 激活函数。

门控 MLA 用于决定通过 MLA 检索到的每一种特征,有多少能够进入残差流。其实现方式,是将这些特征与一个由输入投影得到的门控向量逐元素相乘。

在传统 MoE 中,一个通过学习得到的路由器会利用点积相似度,将每个词元发送到部分专家网络。Kimi K3 一共包含 898 个专家。其中 2 个是共享专家,会处理每一个词元;在其余 896 个专家中,路由器会为每个词元选择 16 个。

Kimi K3 还修改了专家网络中的激活方式。传统做法是:对升维投影应用 SiLU,与门控分支逐元素相乘,再执行降维投影;Kimi K3 则改用 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 Query LoRA、输出门控,以及每 12 层执行一次的分块式注意力残差(Attention Residual)。AttnRes 大约会增加 2% 的推理延迟,但能够带来两项重要收益:

  • 选择性检索早期表示,从而缓解残差信息稀释和隐藏状态持续增大的问题
  • 获得 1.25 倍的计算优势

AttnRes 与 MLA 从不同方向解决了同一个底层限制。KDA 层使用固定大小的状态,因此不可避免地必须丢弃一部分信息。MLA 从词元上下文中检索信息,而 AttnRes 则从网络深度方向上的早期表示中检索信息。

AttnRes(注意力残差)

感谢

@chloey3k 对本节内容提供的帮助。在每一次前向传播中,输入都会依次通过一叠网络层。这里的每一层都由一个注意力模块(KDA 或 MLA)以及一个 MLP 或 MoE 模块组成。通常情况下,每一层的输入,等于原始嵌入与此前所有网络层输出之和,并且所有输出的权重完全相同。

h l = h 1 + ∑ i = 1 l − 1 f i ( h i ) h_l = h_1 + \sum_{i=1}^{l-1} f_i(h_i) hl=h1+i=1l1fi(hi)

其中, h i h_i hi 是第 i i i 层的输入, h 1 h_1 h1 是当前词元的嵌入,也就是目前序列中最后一个词元的嵌入; f i ( h i ) f_i(h_i) fi(hi) 则表示第 i i i 层的输出,即某个注意力模块或 MLP 模块的输出。

问题在于,这种结构缺少选择性访问能力。不同类型的网络层接收到的是同一个聚合状态,尽管它们可能更适合采用不同的权重组合。由于递归过程是纯加法形式,越靠后的网络层还必须学会生成幅度越来越大的输出,才能对不断累积的残差产生足够影响,这可能导致训练不稳定。AttnRes 不再平等对待所有网络层,而是为求和式中的每一项乘上一个专门的权重,使模型能够根据当前上下文,提高最有用网络层的贡献。

h l = α 0 ⋅ h 1 + ∑ i = 1 l − 1 α i ⋅ f i ( h i ) h_l = \alpha_0 \cdot h_1 + \sum_{i=1}^{l-1} \alpha_i \cdot f_i(h_i) hl=α0h1+i=1l1αifi(hi)

每一个权重 α i \alpha_i αi 都由 Query-Key 点积计算得到。每一层拥有一个通过学习得到的 Query,而 Key 和 Value 则来自此前的残差流状态。所有分数会被归一化,使其总和等于 1,随后用于对这些状态进行加权组合。

因此,模型不再只能依赖紧邻的上一层。AttnRes 让每一层都能够选择性访问更早的网络层输出,并允许其学习到的 Query 检索当前计算最需要的表示。

下面的伪代码在分块粒度上应用了同样的思路。一个分块由连续 12 个解码器层中累计得到的注意力输出与 MLP 输出逐元素相加而成,并作为单个深度表示保存下来,供后续 AttnRes 混合使用。

如果在每一层都应用残差注意力,会显著增加训练和推理成本。只在固定分块边界处使用它,则可以用更低成本获得大部分收益。在 Kimi K3 中,每一个边界都位于连续 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 到 Kimi K3 的整条架构演进梳理。

其中最核心的变化并不只是规模扩大。每一次架构演进,都改变了模型存储什么信息、如何更新状态,或者如何重新检索那些固定大小状态无法完整保留的信息。

Kimi K3 将固定状态的递归记忆、周期性的 Softmax 检索、稀疏专家容量,以及沿网络深度进行的选择性残差访问结合在一起。最终得到的系统,会把新增容量投入具有明确功能作用的位置。

从本质上说,固定容量、固定维度的关联记忆必须具备一套淘汰策略,因为纯加法线性操作在达到容量上限后,最终一定会引入信息干扰。因此,门控、路由或衰减等通过学习获得的选择机制是必需的,而注意力则是目前最有效的选择性读取机制。

更多推荐