文章摘要
文章回溯从GPT - 2到Kimi K3的架构进化路径。GPT - 2存在自回归解码效率问题,后续依次出现KV缓存、线性注意力、DeltaNet、门控DeltaNet等优化机制,解决了生成效率、计算复杂度、信息干扰、旧记忆清除等问题。Kimi K3基于Kimi Linear深度优化,融合多种记忆机制,还引入AttnRes选择性复用历史表征,技术演进本质是大模型记忆管理能力的进化。

如果把2026年的Kimi K3模型参数量换算成2019年的GPT-2,大约相当于22580个初代GPT-2——七年时间里,大模型的参数量扩张了两万多倍,但技术演进绝不仅仅是单纯的参数堆叠。这篇文章将回溯从GPT-2到Kimi K3的架构进化路径,解析大模型记忆管理机制的迭代逻辑,以及Kimi K3如何通过全新的记忆设计突破传统Transformer的限制。

初代GPT-2:奠定大模型基础框架

GPT-2是首个大规模落地的仅解码器架构大语言模型,其核心流程是将输入的词元转换为词嵌入与位置嵌入,经过多层Transformer块处理后,通过语言模型头输出预测的下一个词元概率。每个Transformer块由归一化层、因果自注意力模块和MLP模块组成,通过残差连接串联起来,保证梯度在深层网络中能够有效传播。

GPT-2的注意力计算会为每个注意力头生成查询、键和值向量,通过点积计算注意力分数,经过掩码和Softmax归一化后,对值向量进行加权求和,最终拼接所有注意力头的输出完成投影。不过在自回归解码过程中,这种架构存在明显的效率问题:每生成一个新token都需要重新计算整个输入序列的注意力表示,大量计算被冗余消耗。

以下是GPT-2的核心实现代码片段:

CODE

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

以及单个Transformer块的结构代码:

CODE

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

以12个注意力头、嵌入维度768、词表大小50304的配置为例,初代GPT-2的参数量约为1.24亿个。

KV缓存:解决生成效率的基础优化

为了避免自回归生成中的重复计算,KV缓存机制应运而生。其核心思路是将历史token的键(Key)和值(Value)向量预先存储下来,在后续生成新token时直接复用这些缓存的向量,无需重新计算整个输入序列的注意力表示,大幅提升了解码效率。不过随着输入序列长度增加,KV缓存的内存占用会线性增长,甚至成为内存带宽的瓶颈。

线性注意力:将复杂度从平方降至线性

传统Transformer的注意力计算复杂度为O(N²),与输入序列长度的平方成正比,这限制了模型处理超长上下文的能力。线性注意力机制通过改变归一化方式,将点积注意力转化为可结合的矩阵运算,从而将计算复杂度降至线性级别。

具体来说,线性注意力不再通过Softmax对注意力分数进行全局归一化,而是分别对查询(Query)和键(Key)应用ELU+1的特征映射,让Q和K的点积可以被重新结合,从而将不断增长的K和V向量集合压缩为一个固定大小的D×D状态矩阵。不过这种近似归一化的方式会牺牲一定的模型保真度,实际的精度损失取决于具体的模型架构和应用场景。

以下是传统KV缓存的注意力实现代码:

CODE

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

而线性注意力的实现则更为精简,仅需维护一个固定大小的状态矩阵:

CODE

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

DeltaNet:解决缓存的信息干扰问题

线性注意力通过累加的方式更新缓存状态,虽然实现了线性复杂度,但存在一个致命缺陷:当序列长度超过缓存容量后,新旧的键值关联会互相干扰,无法单独检索某个历史token的独立表示。DeltaNet正是为了解决这个信息可恢复性损失而设计的。

DeltaNet引入了写入强度参数,在写入新的键值关联时,先读取当前缓存中对应位置的旧记忆,仅保留真正的新增信息进行更新,避免了旧有记忆被无差别覆盖。其核心代码如下:

CODE

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还支持分块并行计算,将输入序列划分为多个固定大小的块,通过数学重参数化实现块内和块间的并行计算,大幅提升了硬件利用率。

门控DeltaNet:加入动态遗忘机制

DeltaNet虽然解决了单个键值对的精准更新问题,但仍然无法在上下文切换时批量清除旧记忆。门控DeltaNet结合了Mamba的门控更新机制,为缓存引入了全局衰减参数,既保留了Delta规则的精准更新能力,又能通过衰减机制整体释放内存容量。

通过添加一个0到1之间的衰减参数,模型可以在每个时间步对之前的缓存状态进行统一衰减,再写入新的键值关联,防止缓存状态无限增长。这种设计让模型能够在保留关键信息的同时,主动遗忘过时的内容,更好地适配长上下文任务。

Kimi Linear:细粒度的记忆管理

Kimi Linear在门控DeltaNet的基础上进一步优化,不再使用单一标量控制全局衰减,而是为每个通道独立学习衰减参数,实现了更精细的记忆管理。同时它还引入了多头潜在注意力层和混合专家模块,将稀疏激活的专家容量与高效的循环记忆结合起来。

与基础的DeltaNet架构相比,Kimi Linear主要有三个改进:一是在架构中交错使用多头潜在注意力层,兼顾循环记忆的高效性和全注意力的上下文检索能力;二是用混合专家层替代传统的MLP,通过稀疏激活提升模型容量和效率;三是通过α投影为记忆模块增加了细粒度的控制能力,让模型能够更精准地管理每个通道的记忆衰减。

在受控对比实验中,Kimi Linear的性能超过了全注意力架构,在保持更高质量的同时,实现了最高6倍的解码吞吐量提升。

Kimi K3:融合多种记忆机制的终极架构

Kimi K3的语言骨干网络正是基于Kimi Linear架构进行了深度优化。它包含23个四层宏循环,每个循环中三层使用Kimi Delta Attention(KDA),第四层使用多头潜在注意力机制,兼顾了循环记忆的高效性和全注意力的上下文检索能力。

除了架构上的混合设计,Kimi K3还引入了多项改进:

  • 大幅提升了整体模型规模
  • 每12层设置一次块级AttnRes
  • 加入MLA查询LoRA和输出门控
  • 采用潜在空间混合专家模型(Latent-space MoE)
  • 使用SiTU激活函数替代传统的SiLU
  • 加入门控MLA模块

KDA模块提供了恒定状态的循环记忆,而周期性的MLA层则保留了对上下文进行完整Softmax检索的能力,两者结合让模型能够同时高效处理长上下文和精准检索局部信息。

其中SiTU激活函数的实现代码如下:

CODE

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)

Kimi K3总共包含898个专家网络,其中两个共享专家处理每个Token,剩余的896个专家中,路由器会为每个Token选择16个专家进行激活,通过稀疏激活大幅降低了计算量。同时,专家网络运行在压缩后的潜在空间中,进一步将FLOPs减半,提升了推理效率。

AttnRes:选择性复用历史表征

传统的残差连接会将所有层的输出平等累加,导致后续层难以有效影响累积的残差,甚至引发训练不稳定。AttnRes则为每个历史层的残差输出分配了可学习的权重,通过注意力机制动态选择最相关的历史表征,让模型能够选择性地复用更早层的深度表示。

具体来说,AttnRes为每个历史层的输出学习一个查询向量,通过点积计算与当前键的相似度,将相似度归一化后作为权重,对所有历史层的输出进行加权求和,作为当前层的输入。这种设计让模型不必只依赖于紧邻的前一层,而是可以动态选择最相关的历史表征,缓解了深层网络的残差稀释问题。

Kimi K3将AttnRes设置在每12个解码器层之后,在23个宏循环中总共生成8个AttnRes块,在仅增加约2%推理延迟的前提下,获得了显著的性能提升,同时提供了1.25倍的计算优势。

技术演进的核心逻辑

从GPT-2到Kimi K3的七年技术演进,本质上是大模型记忆管理能力的进化:从最初全量存储所有信息的基础架构,到通过KV缓存提升生成效率,再到线性注意力实现线性复杂度的长期记忆,接着通过DeltaNet和门控机制实现动态记忆更新,最终到Kimi K3将循环记忆、全注意力检索、稀疏专家和选择性残差访问结合在一起。

每一次架构演进都不是单纯的参数扩张,而是针对前一代系统的核心限制,通过增加特定形式的容量来优化记忆的存储、更新和检索能力。固定容量的联想记忆必然需要淘汰策略,而注意力机制则为这种选择性读取提供了最有效的方式,让大模型能够在高效运行的同时,更好地处理复杂的上下文任务。

以上内容不代表本平台立场,仅供读者参考