写点什么

从 GPT-2 到 Kimi K3:七年规模扩大 2.26 万倍,大模型架构主线是在建立一套“记忆操作系统”

waterloo_intern
  • 2026-07-30
    北京
  • 本文字数:9769 字

    阅读完需:约 32 分钟

AI摘要

Kimi K3 通过重构记忆机制,突破传统 Transformer 的 KV Cache 瓶颈,在模型规模扩大2.2万倍的背景下,将架构演进焦点从单纯扩容转向选择性记忆与动态信息管理。

KV Cache 优化生成效率;Linear Attention 改进长期依赖建模;DeltaNet 与 Kimi Linear 实现记忆的更新、遗忘与调度。

适合大模型架构师、推理引擎开发者、AI 系统工程师阅读。

摘要:从 GPT-2 到 Kimi K3,大模型技术演进路线展示了 AI 架构如何从“记住一切”,走向“选择性记忆”。KV Cache 解决生成效率问题,Linear Attention 探索更高效的长期记忆,DeltaNet 和 Kimi Linear 则进一步让模型学会更新、遗忘和管理信息。

这篇文章将沿着七年的架构演进,解析大模型规模增长背后的技术变化,以及 Kimi K3 如何通过新的记忆机制突破传统 Transformer 的限制。

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

GPT-2GPT-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 和 value,并将注意力头维度移到前面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)源于一个直观的观察:将新生成的词元追加到输入后,模型原本需要重新计算此前所有词元的投影。存储这些词元的键(Key)和值(Value)向量,可以避免这种重复计算。

这个存储空间就是 KV 缓存。它保存了此前 N−1 个词元的向量,并且可能变得非常庞大,甚至造成内存带宽瓶颈。

总体而言,在词表包含约 5 万个词元、网络包含 12 个 Block 和 12 个注意力头、嵌入维度为 768 的情况下,我们的基准模型大约拥有 1.24 亿个参数。

vocab_size: int = 50304 # GPT-2 的词表大小为 50257,为提升效率而填充至最接近的 64 的倍数n_layer: int = 12n_head: int = 12n_embd: int = 768
复制代码

Kimi K3 拥有 2.8 万亿(2.8T)个参数。就参数量而言,一个 Kimi K3 模型大约相当于 22,580 个 GPT-2 模型。

线性注意力(Linear Attention)

Softmax 注意力机制会在 q·k 点积之后应用非线性变换,使每个 Query 与每个 Key 相互关联。而线性注意力则分别对 q 和 k 应用特征映射(例如 ELU + 1)。这使得矩阵乘法具备重新结合的特性,从而可以将不断增长的 K 和 V 向量集合压缩为一个固定大小的 D×D 状态矩阵。

论文中关于 O(N²) 的描述曾让我感到困惑:“Transformer 每个时间步的计算成本与当前序列长度的平方成正比。”这种说法并不完全准确。这其实是 Flash Attention 所解决的问题……后来我发现,Flash Attention 是在 2020 年发布的。

在当时,训练过程中通常会显式构造完整的 N×N 注意力矩阵。Flash Attention 尚未出现,而参考实现中的自回归生成通常会在没有 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)执行两次二维读取和两次一维写入,而 KV 缓存则会随着序列长度线性增长,复杂度为 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 和 k 的 ELU + 1,然后再让二者进行交互。两种方法都会对最终得到的分数进行归一化,但线性注意力使用的特征映射只是 Softmax 核的一种表达能力较弱的近似。这种近似可能会降低模型的保真度,不过实际精度损失取决于具体架构和工作负载。

注意,我们仍然需要除以 qk 的总和(为了简化图示,后续图中省略了这一归一化项)。从宏观角度来看,注意力机制包含三个步骤:使 qk 分数变为非负值。线性注意力使用 ELU + 1,而 Softmax 使用指数运算。除以总和进行归一化。计算这些值对应的加权平均。

这样既保留了注意力机制的基本约束,又使用表达能力较弱的特征映射,使 QK 分数保持非负。

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

有限大小的缓存必须覆盖或合并已经存储的信息。来自第 i−1 个 Token 的状态不会拥有自己的独立存储位置;它会被累加到同一个 D×D 矩阵中。因此,新的 Query 无法再检索每个此前 Token 的完全独立表示。

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

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

当 N≫DN 时,线性注意力机制展现出很强的吸引力,但这也暴露了它的主要局限。一旦状态超过其有效容量,关联就会开始相互干扰,因为更新方式是纯累加的,缓存中不会淘汰旧信息。

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
复制代码

用图示说明会更容易理解:

考虑一个关联关系:

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

这是本文中最难理解的部分。我花了大约七个小时才真正理解它,因此我会从实现角度展开解释。简单来说,DeltaNet 实现了一种带有广义 Householder 转换矩阵的一阶线性递归,使其能够进行分块并行(Chunkwise Parallel)的前向计算,从而实现硬件高效的线性时间训练。它将输入和输出划分为多个大小为 C 的块,并根据前一个块的最终状态,以及当前块中的 Query、Key、Value 块来计算每个块的输出。

实际问题出现在 Prefill(预填充)阶段。对于长度为 T 的 Token 序列,直接应用 Delta 规则的顺序实现如下:

S = torch.zeros(b, h, dh, dh) if cache is None else cacheouts = []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 规则,直接的线性注意力 Prefill 过程仍然是顺序执行的:

S = torch.zeros(b, h, dh, dh) if cache is None else cacheouts = []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 设置为标准的 O(N²) 注意力机制,而 C=1 则对应普通线性注意力机制。我们可以在两者之间进行权衡:通过增加块内计算量,换取更好的硬件利用率。实际应用中,C 通常取 64 或 128,因为张量核心(Tensor Core)指令在这一粒度下运行效率更高。

中间块会在状态更新过程中被折叠进 S:

S = torch.zeros(b, h, dh, dh) if cache is None else cacheouts = []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)
复制代码

在一个块内,我们执行:

因此,计算成本被分成两部分:

从纯 FLOPs 角度来看,C=1 是最节省计算量的选择,但实际运行时间未必最低。当计算任务能够高效映射到 GPU 的矩阵乘法硬件上时,GPU 可以更快完成更多计算。

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

核心问题很简单:用于纯加性注意力的分块方法,并不能直接应用于 Delta 更新:

v_old = k_i @ S                  u_i = b_i * (v_i - v_old)
复制代码

我们需要知道每一个状态,才能计算需要被减去的信息。如果不进行某种数学上的重新参数化,就无法以相同方式实现并行化。因此,作者将 Delta 更新从:

u=v_new-v_oldS_t= S_(t-1)+K.T@uo=q@S_T
复制代码

重新表示为:

这里,顺序循环会在每次迭代中计算一个 Delta。

重新参数化后的形式,使分块代码能够一次性计算所有 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)                # 使用向量化前向替换计算公式 (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(带修正量)                S += k_i.t() @ u_i # 使用加法更新状态                O[i] = o_intra + o_inter # 更新输出:Flash + Recurrent        return O.reshape(L, d)
复制代码

这就引出了我们的第一个比较点:MHA 与 DeltaNet Transformer 的对比:

门控 DeltaNet(Gated DeltaNet)

我们现在已经拥有了一种精确修改缓存的方法。对于每个新的事实(每个新的 Key 向量),我们都可以准确查看当前存储的旧信息,并将其替换为我们希望存储的新信息。

然而,这种机制只能遗忘那些存在明确替代项的关联。它无法在上下文切换期间有效清除多个关联,也无法通过整体衰减记忆来释放容量。

如果我们使用纯粹的加性线性注意力:

添加遗忘机制非常简单。我们只需要一个控制遗忘状态的参数:

S_old=cacheS_new=k@v# cache=S_old+S_newcache=alpha * S_old + S_new
复制代码

这就是 Mamba-2 的贡献。我们先对之前的缓存进行衰减,然后再以完整强度添加新的缓存,从而防止状态无限增长。

在每个时间步中,以动态比例统一衰减所有键值关联是一种可行的方法,这也是 Mamba 所采用的方法。但它没有考虑不同键值关联的重要性差异。

也就是说,如果模型需要遗忘某一个特定关联,那么所有关联都会被同等程度地遗忘。相比之下,Delta 规则可以更新单个事实,但无法让其他事实整体衰减。

因此,门控 Delta 规则(Gated Delta Rule)将 Mamba 的门控更新规则与 Delta 规则结合起来。它增加了一个参数 α:当 α=1 时,切换到纯 Delta 规则;当 α=0 时,清空记忆。难点在于,如何使用相同的分块并行方法实现这一机制。

该实现采用了上一节介绍的相同 DeltaNet 重参数化方法。其数学形式几乎完全一致,只增加了一个与数据相关的 0 到 1 之间的标量,用于控制之前状态的衰减。这将高效的键值关联学习与自适应记忆管理结合在了一起。

相应的代码修改如下:

γʳ/γⁱ 项表示累积衰减。一个在时间步 x 写入、并在 x+t 时读取的 Token,会经过:

的连续乘法衰减。这相当于前缀和计算的乘法形式。

最终的架构如下所示:

KDA / Kimi Linear

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

Kimi Linear 因一个核心优势而受到关注:在受控对比实验中,它的性能超过了全注意力机制。作者将其描述为一种可以直接替代的架构,在保持更高质量的同时,实现了最高 6 倍的解码吞吐量提升。

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

KDA 更新规则保持类似,但代码现在更接近下面这样:

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

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

  • 它采用混合系统,在架构中交错使用多头潜在注意力(MLA,Multi-Head Latent Attention)层。

  • 它使用混合专家(MoE,Mixture of Experts)层替代 MLP。

  • 它通过 α 投影为 DeltaNet 增加了容量。

后续章节将更详细地介绍 MLA 和 MoE。目前,重要的是理解这并非盲目的规模扩展。新增容量具有明确的数学目的:按通道控制衰减,使模型能够更精细地管理记忆衰减。

扩展规律(Scaling Laws)仍然有效,但容量必须增加在正确的位置,并以系统能够有效利用的形式存在。这一演进过程中的每种架构,都通过增加特定形式的容量,来解决前一个系统中的具体限制。

Kimi K3

最终,Kimi K3 的语言骨干网络与上述 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(Gated MLA)

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

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

门控 MLA 决定从 MLA 检索到的每个特征,有多少能够传递到残差流中。它通过将每个元素与一个由输入投影得到的门控值进行逐元素相乘来实现。

在传统 MoE 中,学习得到的路由器会利用点积相似度,将每个 Token 分配给部分专家网络。Kimi K3 总共有 898 个专家。其中两个共享专家会处理每个 Token;在剩余的 896 个专家中,路由器会为每个 Token 选择 16 个专家。

Kimi K3 还改变了专家激活方式。它不再像传统方式那样,对上投影结果应用 SiLU,将其与门控值逐元素相乘,然后再进行下投影,而是使用 SiTU:

d = x.shape[-1] // 2gate = 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 倍。一种补偿性优化方式是,专家模型运行在压缩后的潜在空间中,这使得它们的前向传播速度更快,并且几乎将 FLOPs 减半。

其余变化包括 MLA 查询 LoRA、输出门控,以及每 12 层使用一次块级注意力残差(Blockwise AttnRes)。注意力残差会增加约 2% 的推理延迟,但带来两个重要优势:

  • 选择性检索更早的表征,从而缓解残差稀释(Residual Dilution)和隐藏状态增长。

  • 提供 1.25 倍的计算优势。

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

AttnRes(注意力残差)

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

问题在于缺少选择性访问能力。不同类型的层会接收到相同的聚合状态,即使它们可能需要不同的权重组合。由于这种累加过程完全是加性的,后续层必须学习越来越大的输出,才能影响已经累积的残差,这可能导致训练不稳定。AttnRes 不再平等地对待所有层,而是为总和中的每一项乘以一个专门的权重,使模型能够根据上下文,为最有用的层赋予更高权重:

每个权重都通过 Query 与 Key 之间的点积计算得到。Query 是针对每一层学习得到的,而 Key 和 Value 则来自此前残差流中的状态。得分会被归一化,使其总和为 1,然后用于形成这些状态的加权组合。

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

下面的伪代码以块(block)为单位实现了相同的思想。一个块是 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 检索、稀疏专家容量以及选择性深度残差访问结合在一起。最终,该系统能够将额外的容量分配到具有明确功能作用的位置。

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

原文连接:

https://x.com/waterloo_intern/status/2081762065392541951