22580:从 GPT-2 到 Kimi K3,全解析
七年放大 22580 倍,架构到底变了多少?这篇工程师 worklog 从 GPT-2 的 decoder-only 出发,逐步推导线性注意力、DeltaNet、门控 DeltaNet、KDA 与 AttnRes,答案远不只是「放大」。
本文译自 X 平台用户 @waterloo_intern 的技术长文《22580: From GPT2 to Kimi3, Explained》,发布于 2026 年 7 月。
- 发布方:一位工程师的个人 worklog(学习 + 复现笔记),不是任何公司的官方文档
- 主角:从 GPT-2 (2019) 到 Kimi K3 (2026) 的大模型架构演进线
- 文档性质:一手推导记录 = 干货 + 个人理解,作者会当场自我修正;结尾落在月之暗面的 Kimi K3 上,相关性能结论引自厂商自己的论文
- 关键节点:GPT-2 (2019) → 线性注意力 (2020) → DeltaNet → 门控 DeltaNet / Mamba-2 → Kimi Linear → Kimi K3 (2026 年 7 月)
▍背景:七年里架构在解决什么问题 2019 年的 GPT-2 用 softmax 注意力:每生成一个词,都要回看历史里所有词的键值向量(KV 缓存),缓存随对话变长而无限变大。此后七年架构演进的主线,就是给这个「越背越重的书包」找出路:把历史折叠进固定大小的状态(线性注意力)、学会精准改写(Delta 规则)、学会选择性遗忘(门控),最后把所有工具组合成混合架构(Kimi K3)。
▍三句话看懂这条演进线 ① GPT-2:全量回看,记得最全,但缓存随长度线性膨胀。 ② 线性注意力家族:把历史压缩进一块固定大小的「黑板」,缓存不再增长,但写满了会互相覆盖——后续每一步都在解决「怎么写、怎么擦、怎么忘」。 ③ Kimi K3:不再二选一,把恒定状态记忆、周期性全量检索、稀疏专家、深度方向的选择性访问组合在一起;22580 倍的规模增长背后,是一整套记忆管理设计。
▍内容地图 原文按「GPT-2 地基 → 线性注意力 → DeltaNet 两节 → 门控 → Kimi Linear → K3 → AttnRes」推进:GPT-2 节逐行讲代码,是后面一切的地基;DeltaNet 并行化一节是全文最硬的推导(作者自称花了七个小时);门控到 K3 节奏明显加快;结尾三段是全文总纲,值得读两遍。本文代码即正文——读代码就是读论点。
两万两千五百八十。这就是一个 Kimi K3 (2026) 里面能装下多少个 GPT-2 (2019)。七年时间,我们整整放大了 22580 倍。但这真的只是……规模吗?
在这篇 worklog 里,我会带你走一遍我们是怎么走到今天的,以及这期间到底变了多少、又有多少其实没变。我们会顺着通往 Kimi K3 的几步关键架构演进,一步一步追下去。
GPT-2
GPT-2 是一个 decoder-only 架构:
tok_emb = self.transformer.wte(idx) # token 嵌入,形状为 (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输入会拿到 token 嵌入和位置嵌入:
把每个 transformer block 放大来看,长这样:
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注意力 (attention) 的过程:
B, T, C = x.size() # 批大小、序列长度、嵌入维度 (n_embd)
# 为所有 head 批量计算 query、key、value,并把 head 维提到前面当作 batch 维
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) # 把所有 head 的输出并排拼回去
# 输出投影
y = self.resid_dropout(self.c_proj(y))
return y得到最终的隐状态 (hidden-state) 矩阵之后,语言模型头 (language-model head) 会把它映射成词表上的 logits(词表上每个候选词的分数)。在自回归 (autoregressive) 解码时,只需要最后一个位置的 logits 就能选出下一个 token。
这正是 decoder-only 生成方式的一个低效之处:模型为每个输入位置都计算了表示,但每一步解码只消费最后一个位置的 logits。如果没有缓存,下一个 token 到来时,前面的大量计算都要重来一遍。
Decoder-only 指模型只有「解码器」一摞层:它能看到自己已经写出的所有词,据此预测下一个词,再把新词接在后面重复这一过程(自回归)。原文第一个代码块说的就是这件事的完整流水线:把词变成向量(嵌入)→ 过 12 层 block → 用最后一个位置的输出给词表打分(logits),从分数最高的候选里挑下一个词。
KV 缓存 (KV cache) 来自一个很直接的观察:把生成出来的 token 拼回输入之后,模型本来要为之前所有 token 重新算一遍投影。把它们的键 (key) 和值 (value) 向量存下来,就能省掉这些重复劳动。
这个存储就是 KV 缓存。它保存着之前 N-1 个 token 的向量,而且可能会大到成为内存带宽瓶颈。
每生成一个新词,模型只需要历史各词的键 (key) 和值 (value) 向量参与计算,而这些向量一经算出就不再改变。把它们存下来(KV 缓存),解码就从「每步重算全部历史」变成「每步只算新词 + 读缓存」。
代价是缓存随对话长度线性增长:聊到几万 token 时,读写缓存比计算本身更耗内存带宽。后面整条线性注意力路线,出发点就是消灭这个不断变大的缓存。
总结一下:词表里约 5 万个候选 token、12 个 block、12 个 head、嵌入维度 768,我们的基线模型大约是 1.24 亿参数。
vocab_size: int = 50304 # GPT-2 的词表大小是 50257,为了效率向上取整到 64 的倍数
n_layer: int = 12
n_head: int = 12
n_embd: int = 768而一个 2.8 万亿参数的 Kimi K3,参数量大约相当于 22580 个 GPT-2。
GPT-2 约 1.24 亿参数,Kimi K3 约 2.8 万亿参数。形象点说:把 GPT-2 比作一本 300 页的书,22580 本摞起来约有 1.4 公里高——七年时间,前沿模型的块头从「一本书」变成了「一座书塔」。但参数多 22580 倍不等于能力强 22580 倍,本文的主线正是「除了变大,还变了什么」。
线性注意力 (Linear Attention)
Softmax 注意力是在 q·k 相乘之后才施加非线性,于是每个查询 (query) 都和每个键 (key) 耦合在一起。线性注意力则换了个做法:对 q 和 k 分别施加一个特征映射 (feature map),比如 ELU+1(一种让分数非负的激活函数)。这样一来乘法就可以重新结合 (re-associable),不断增长的 K 和 V 向量集合就能被折叠进一个固定大小的 D×D 状态 (state) 里。
Softmax 注意力把每个词和每个词两两比较,历史越长比较越多。线性注意力换个顺序:利用乘法结合律,先把所有键值向量压缩(折叠)进一个固定大小的 D×D 矩阵 S——不管历史是 100 个词还是 10 万个词,S 的尺寸都不变,查询时只跟 S 打交道。缓存因此从「随长度增长」变成「恒定大小」。
这就是后文一切取舍的起点:省下了内存,却把无数历史挤进了有限的格子。
论文里 O(N²) 的说法一开始误导了我。所谓「transformer 每个时间步的开销随当前序列长度的平方增长」并不成立——那是 FlashAttention 要解决的问题……后来我一看,这篇论文是 2020 年发的。
在那个年代,训练时通常会把完整的 N×N 注意力矩阵实实在在地物化 (materialize) 出来;FlashAttention 还不存在;而参考实现的自回归代码往往不带 KV 缓存,每个 token 都要把整段历史重算一遍。
作者被论文的 O(N²) 说法误导,本身就是一个时间坐标问题:线性注意力论文发表于 2020 年,彼时 FlashAttention (2022) 还没出现,参考实现普遍不带 KV 缓存,「每步重算历史、物化完整 N×N 矩阵」确实是常态。放到 2026 年读旧论文,很多「动机」已时过境迁。
读架构论文先核对发表年份——这是作者用自己的困惑换来的阅读习惯。
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
# 所以我在 t 维(也就是 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'))
#得到注意力 (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
# 用 x 得到 qkv
o_proj=self.o_proj(o)
past_kv=(k, v)
return o_proj, past_kv同一个过程,画成图会更好懂。每一步解码都要对 HBM(显卡的高带宽内存)做两次 ND(N 维数据块)读取和两次 1D 写入,而 KV 缓存随序列长度线性增长,也就是 O(N)。
注意这些过量的读写——而这篇论文把它们替换成了:
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 #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 和 k 相互作用之前分别施加的 ELU+1。两种做法都会对算出来的分数做归一化,但线性注意力用的特征映射是对 softmax 核 (kernel) 的一种表达能力更弱的近似。这个近似会损失保真度,不过实际掉多少精度,取决于具体架构和负载。
注意我们仍然要除以 qk 的总和——示意图里为了简洁把这一步省掉了。从高层看,注意力就三步:
- 让 qk 分数非负。线性注意力用 ELU+1,softmax 用指数。
- 除以总和。
- 对值 (value) 做加权平均。
这保留了注意力的基本约定,只是换用一个表达能力更弱的特征映射来让 QK 分数非负。
DeltaNet(快速权重程序员)
有限的缓存必然会覆盖或融合已存的信息。来自第 i-1 个 token 的状态并没有自己的独立槽位——它被加进了同一个 D×D 矩阵里。于是新的查询再也无法取回每个早期 token 完美隔离的表示。
但这个「加法」也正是效率提升的来源。用相加而不是拼接来更新缓存,缓存就不会随 O(N) 增长;可同样是这个操作,让信息之间互相干扰。DeltaNet 要解决的就是这种「取不回来」的问题。
Schlag 的论文(快速权重程序员,Fast Weight Programmers)说得很到位:「当序列长度超过存储容量时,模型可能进入过载 (overcapacity) 状态。要在这种状态下正常工作,模型应该学会动态地与记忆内容交互,有选择地决定保留哪些键值关联、删除哪些。纯粹的累加指令可能并不适合这个目的……像公式 17 那样,无休止地往一个容量有限的记忆里添加新关联,迟早会触到极限。」
让线性注意力变得诱人的那个前提——N 远大于 D——同时也暴露了它最主要的局限。一旦状态超过有效容量,关联之间就会开始互相干扰,因为更新是累加的,而且没有任何东西会离开缓存。
把固定大小的状态想象成一块黑板:每来一个新事实就往上写一行。黑板够大时相安无事;写满之后,新字只能叠在旧字上——再去读,前后两行糊在一起,谁也认不全。这就是「关联互相干扰」。
纯累加的线性注意力只会往上写、从不擦除。Schlag 论文指出的正是:容量有限的黑板必须配一套「决定保留什么、擦掉什么」的规则。DeltaNet 就是第一块带板擦的黑板。
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) # 增量:只写真正新增的部分
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配一个可视化的例子会更好理解。
拿一条单独的关联来说,写作 S = k.T @ v。如果用同一个 key 去读,得到的是 k @ (k.T @ v),也就是 (k @ k.T) v,也就是 k 的范数(向量长度)平方乘以 v。所以读回来的结果会带上 key 范数平方这个缩放因子;如果把 k 归一化到单位长度,或者直接把结果除以范数,就能精确地拿回 v。
Q 也是一个学出来的指针。Wq 和 Wk 读的是同一条残差流 (residual stream),某个事实的查询会指向当初写入这个事实的那个键的方向。更新的过程是:先问当前的 key 能从缓存里取出什么信息,从我们想存的值里减去这份已有信息,把 key 乘上这个差值,再把结果加回去。旧信息被移除,新信息写进原来的位置。
DeltaNet 的更新分三步:先用这个 key 去黑板上读出当前存着的旧值 (v_old),算出新值与旧值的差(delta,增量),再只把差值写上去。效果是「精准改写」:同一个 key 再读时拿到的就是新值,旧值被恰好覆盖,黑板上其他内容不受影响。
代价是每一步都要先读再写——这正是下一节「难以并行化」的根源。
DeltaNet(用 Delta 规则并行化线性 Transformer)
这是整篇文章里最难的一节。我大概花了七个小时才算真正搞懂它,所以我会从实现出发来搭建解释。一句话概括:DeltaNet 实现了一个带有广义 Householder 转移矩阵(一种经典的矩阵变换,知道名字即可)的一阶线性递归,从而支持按块 (chunk-wise) 并行的前向传播,实现硬件高效的线性时间训练。它把输入输出切成若干大小为 C 的块,每个块的输出基于上一块的最终状态和当前块的 query、key、value 块来计算。
实际的问题出在 prefill(预填充)阶段。对一段 T 个 token 的序列直接实现 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 规则,直接的线性注意力 prefill 也依然是顺序执行的:
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)分块 (chunked) 的写法提供了一条更高效的路。其中的机制,通过一个例子来看最容易懂:
设 C=N 就退化回标准的 O(N^2) 注意力,而 C=1 就是普通的线性注意力。取中间值,相当于用块内的额外计算换取更好的硬件利用率。实践中 C 常取 64 或 128,因为张量核心 (tensor core) 指令在这个粒度上效率最高;UMMA(一种矩阵乘指令)就是一个例子。
中间的那些小分块 (tile) 会作为状态更新的一部分被折叠进 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ᵀv)。这是「先算分数」,就是带掩码的正常注意力顺序。跨块的时候,我们走的是 (kᵀv)q,也就是递归顺序,状态优先。注意力的开销按 O(N²) 增长,而这个不会。块内我做的是真注意力(带掩码的 QKᵀ 乘 V),跨块我把所有东西折叠进状态,再用一次矩阵乘法读回来。所以开销分成两块:一块是固定的 2Ld²,这是状态更新的活儿,跟 C 完全无关;另一块是增长的 2LCd,来自对角线上的那些分数矩阵。完整注意力就是 C 等于 L 的特例,此时第二项变成 2L²d,二次方。所以 C 取得越小,我花的 FLOPs 就越少。
纯按 FLOP 算,C=1 是最省的选择,但按墙上时钟 (wall-clock) 时间算不一定。当计算能高效映射到矩阵乘法硬件上时,GPU 能更快完成更多算术运算。
这一节最容易带走的错误结论是「C=1 最省 FLOPs,所以块越小越好」。作者自己泼了冷水:按 FLOPs 算 C=1 最省,按墙上时钟算未必——GPU 为矩阵乘法而生,块太小,算术是省了,硬件却在空转。C 取 64/128,本质是「用一点多余的计算,换硬件满负荷运转」。
读任何「线性复杂度」宣传都该问一句:省的是 FLOPs 还是时间?这两者在 GPU 时代经常不是一回事。
下一步,就是把同样的思路推广到 DeltaNet 上。
底层的问题很简单:用于纯累加注意力的分块方法,没法直接套用到 delta 更新上:
v_old = k_i @ S
u_i = b_i * (v_i - v_old)为了算出需要减掉的那些信息,我们需要每一个中间状态。不做一些数学上的重参数化 (re-parameterization),就没法用同样的方式并行化。于是作者们把 delta 更新从这种形式:
u=v_new-v_old
S_t= S_(t-1)+K.T@u
o=q@S_T这里是一个顺序循环,每次迭代算出一个 delta。重参数化之后的形式是:
S_t = S_{t-1}(I − β_t k_t k_tᵀ) + β_t v_t k_tᵀ
o_t = S_t q_t这种写法让分块代码可以一次性算出全部 C 个 delta:
def chunk_delta_rule_forward(Q, K, V, beta, C):
# L: 序列长度, d: head 维度
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 # 注意力 @ v(带修正,所以是 u)
S += k_i.t() @ u_i # 用加法更新状态
O[i] = o_intra + o_inter #用 flash 部分加递归部分更新输出
return O.reshape(L, d)这就带我们来到第一个对比点:MHA 对阵 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 规则 (Gated Delta rule) 把 Mamba 的门控更新规则和 Delta 规则结合了起来。它增加了一个参数 alpha:取 1 时退化为纯 Delta 规则,取 0 时清空记忆。难点在于要用同样的按块并行方法把它实现出来。
实现上沿用了上一节讲的 DeltaNet 重参数化。数学上几乎一模一样,只多了一样东西:一个介于 0 和 1 之间、由数据决定的标量,用来控制旧状态的衰减。这就把有效的键值关联学习和自适应的记忆管理结合到了一起。
到 DeltaNet 为止,黑板只会精准改写、不会遗忘:想清空一块区域,只能拿新内容逐条覆盖。Mamba-2 补上另一半:每步给整块黑板乘一个小于 1 的系数,让所有字迹统一变淡(衰减),但不分轻重。门控 DeltaNet 把两者合起来——Delta 规则负责「精准擦这条」,门控负责「整体淡一点」。一个管改写,一个管遗忘,黑板才真正可管理。
对应的代码改动如下:
γʳ/γⁱ 这一项负责累计衰减。一个在时间步 x 写入、在 x+t 被读出的 token,已经乘上了 αₓαₓ₊₁αₓ₊₂…αₓ₊ₜ。这就是前缀和 (prefix-sum) 计算的乘法版。
最终的架构长这样:
KDA/Kimi Linear
到这一步,研究者们开始尝试混合架构——在一个模型里组合多种注意力形式,比如把 Gated DeltaNet 和 Mamba 搭在一起。
Kimi Linear 因为一个核心主张而备受关注:在受控对比下,它超过了完整注意力 (full attention)。作者把它定位为一个可以直接替换的架构方案,质量更好,解码吞吐量最高提升 6 倍。
Kimi Linear 对 Gated DeltaNet 的改进在于引入了细粒度门控 (fine-grained gating):不再用单一标量做衰减,而是为每个通道 (channel) 学一个单独的衰减值。
「受控对比下超过 full attention、解码吞吐最高提升 6 倍」——这句话的出题人是 Kimi 自己(Kimi Linear 论文)。受控对比意味着变量由作者设定:基线实现、序列长度、评测任务,都可能是对己方有利的组合;「最高 6 倍」通常也只对特定解码长度成立。
这不是说结论为假,而是说在第三方独立复现之前,它应被当作「厂商主张」而非「行业共识」。读模型论文时,先把每个对比句子的主语换成「作者声称」,再决定信几成。
KDA 的更新规则保持相似,但代码现在更像这样:
这里的 alpha.reshape(nb, C, d) 正是这篇论文最重要的贡献:对记忆衰减的细粒度控制。
把 Kimi Linear 架构和 DeltaNet Transformer 并排看,它引入了三大改动:
- 使用混合架构,交错插入多头潜在注意力 (Multi-head Latent Attention, MLA) 层。
- 用混合专家 (Mixture-of-Experts, MoE) 层替换了 MLP。
- 通过 alpha 投影为 DeltaNet 增加了容量。
后面的章节会更详细地介绍 MLA 和 MoE。现在重点要记住:这不是盲目堆规模。新增的容量有明确的数学目的——逐通道的缩放让模型对记忆衰减有了更精细的控制。
规模定律 (scaling law) 依然成立,但容量必须加在正确的位置、以系统能用得上的形式加进去。这条演进线上的每一种架构,都是为了解决前一个系统的某个具体局限而增加容量的。
Kimi K3
最终,Kimi K3 的语言主干和上面的 Kimi Linear 模型长得很像。它包含 23 个四层宏循环 (macrocycle)。每个宏循环里,三层用 Kimi Delta Attention,第四层用多头潜在注意力。第一层用稠密前馈网络 (dense FFN);其余每一层都用潜在混合专家 (latent MoE)。
乍一看,相比 Kimi Linear 的改动似乎不大:
- 规模大幅增加
- 每 12 层一次的块级 AttnRes
- MLA 查询 LoRA 和输出门控
- 潜在空间 MoE
- SiTU 激活
- 门控 MLA
KDA 提供恒定状态的递归记忆,而周期性的 MLA 层保留了对上下文的完整 softmax 检索。下面这张简化的示意图,可以作为理解后文改动的参考。
我们先从比较直接的几个改动说起:门控 MLA、潜在空间 MoE 和 SiTU 激活。
门控 MLA (Gated MLA) 决定 MLA 检索出的每个特征有多少能进入残差流。做法是把特征与一个从输入投影出来的门 (gate) 做逐元素相乘。
在传统的 MoE 里,一个学出来的路由器 (router) 用点积相似度把每个 token 分配给一小部分专家网络。Kimi K3 一共有 898 个专家。其中 2 个是共享专家,处理每个 token;剩下 896 个里面,路由器为每个 token 选 16 个。
Kimi K3 有 898 个「专家」小网络,但每个 token 只动用其中 18 个(2 个共享专家 + 路由器选出的 16 个)。就像一家有 898 个科室的大医院,每个病人只看 18 个科室——总参数可以做到 2.8 万亿(医院大),每次推理却只激活约 320 亿(出诊少)。这就是「规模大幅增加」却不至于慢到不可用的原因。
Kimi K3 还改了专家的激活函数。不再是对上投影 (up projection) 施加 SiLU(一种常用的激活函数)、与门逐元素相乘、再做下投影 (down projection),而是改用 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 倍。一个可以抵消这项开销的优化是:专家们在压缩后的潜在空间里工作,前向传播快得多,FLOPs 几乎减半。
剩下的改动是 MLA 查询 LoRA、输出门控,以及每 12 层一次的块级注意力残差 (Attention Residuals, AttnRes)。AttnRes 增加约 2% 的推理延迟,但带来两个重要收益:
- 有选择地检索更早的表示,缓解残差稀释 (residual dilution) 和隐状态膨胀
- 1.25 倍的计算优势(译者注:原文此处表述含糊,未说明是节省算力还是等效算力增益)
AttnRes 和 MLA 从不同方向解决的是同一个底层局限。KDA 层使用恒定大小的状态,信息被丢弃不可避免。MLA 从 token 上下文里检索,而 AttnRes 从更早的深度方向表示里检索。
AttnRes
感谢 @chloey3k 对本节的帮助。每次前向传播,输入都要穿过一摞层。在这里,每一层由一个注意力块(KDA 或 MLA)和一个 MLP 或 MoE 块组成。正常情况下,每一层的输入是原始嵌入与之前所有层输出的求和——所有项权重相等。
这里 h_i 是第 i 层的输入,h_1 是当前 token(目前为止序列里的最后一个 token)的嵌入,f_i(h_i) 是第 i 层(一个注意力块或 MLP 块)的输出。
问题在于缺少有选择的访问。不同类型的层收到的是同一份聚合状态,尽管它们可能受益于不同的权重配比。而且因为这个递归是纯累加的,越靠后的层越要学出越来越大的输出才能影响累积起来的残差,这可能让训练失稳。AttnRes 不再一视同仁地对待所有层,而是给求和式里的每一项乘上一个专门的权重,让模型能把更重要的位置留给当前上下文里最有用的那些层。
普通残差连接像传话时默认「每个人说的都一样重要」:每层的输入是之前所有层输出的等权重累加——层数越深,最早的信息被稀释得越厉害(残差稀释)。AttnRes 让每层用一组学出来的权重,对更早各层的输出做加权检索:注意力不再只沿「序列」方向看历史,也沿「深度」方向看过去。只在每 12 层的块边界上做,约 2% 的延迟代价,换来训练更稳、远程信息更可取。
每个权重 α_i 都由一个查询-键点积算出来。查询是每一层各自学出来的,键和值则来自更早的残差流状态。这些分数被归一化到总和为 1,然后用来对那些状态做加权组合。
这样一来,模型就不必只依赖它的直接前驱了。AttnRes 让每一层都能有选择地访问更早层的输出,用它学出来的查询,检索出对当前计算最有用的表示。
下面的伪代码把同样的思路用在块 (block) 粒度上。一个块是 12 个 decoder 层上累积的注意力输出与 MLP 输出的逐元素求和,作为一份单独的深度表示存起来,供后面的 AttnRes 混合使用。
如果每一层都做残差注意力,训练和推理成本都太高了。只在固定的块边界上做,能以更低的成本拿到大部分收益。在 Kimi K3 里,每个边界出现在 12 个 decoder 层之后。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 的整条演进线就走完了。
「一个固定容量的关联记忆需要一个淘汰策略……学出来的选择机制是必需的,而注意力正是最有效的选择性读取机制。」
全文一句话总纲:七年架构演进,表面是 22580 倍的规模,实质是记忆管理的三次升级——存什么(固定状态)、怎么改(Delta 改写与门控遗忘)、怎么取(注意力检索)。
冷静注脚:这是作者通篇推导后的个人总结,逻辑自洽且与主流文献一致,但「最有效」一类的断言仍属一家之言。
核心的变化不只是规模。每一步架构演进,改变的都是模型存储什么、如何更新这个状态、或者如何检索那些固定大小状态留不住的信息。
Kimi K3 把恒定状态的递归记忆、周期性的 softmax 检索、稀疏的专家容量、以及有选择的深度方向残差访问组合在了一起。最终得到的系统,会把新增的容量花在它有明确功能角色的地方。
本质上,一个固定容量的关联记忆(维度固定)需要一个淘汰策略 (eviction policy),因为纯累加的线性操作一旦到达容量上限,就必然会不断引入干扰。为此,学出来的选择机制——比如门控、路由或者衰减——是必需的,而注意力正是最有效的选择性读取机制。
- 普通读者:下次看到「新架构更快更强」的新闻,先问三个问题——缓存怎么变了?状态是固定大小还是随长度增长?谁来决定遗忘什么?这三问能拆穿大多数架构叙事的包装。
- 学生 / 入门者:把本文 GPT-2 那几段代码亲手敲一遍(它们来自广为流传的 nanoGPT 教学实现),再对照线性注意力的 forward 找差异——两版代码 diff 出来的那几行,就是七年架构演进的最小切面。
- 开发者:评估线性注意力类方案时别只看 FLOPs——用目标硬件实测 prefill 与 decode 的墙上时钟;块大小 (chunk size) 与算子融合 (fused kernel) 往往比理论复杂度更决定成败。
读同类架构文章的检查清单:① 复杂度声明对应的是训练还是推理、FLOPs 还是墙钟时间?② 对比实验的基线与评测条件由谁设定?③ 记忆/缓存的容量上限在哪里,到顶之后发生什么?④ 新机制解决了上一个方案的哪个具体局限——还是只是更大?