22580:从 GPT-2 到 Kimi K3,完整解析
七年间,大模型参数规模增长了 22580 倍。从 GPT-2 的标准注意力,到线性注意力、DeltaNet、Kimi Linear 与 Kimi K3,这是一份完整的架构演进拆解。


二万二千五百八十。这就是一个 KimiK3(2026)所能容纳的 GPT-2(2019)模型数量。七年间,我们把规模扩大了 22,580 倍。但这一切真的只是……规模变大了吗?
在这篇工作日志中,我会梳理我们如何走到今天,以及从那时起,实际发生的变化究竟有多少——或有多么少。我们将沿着通往 KimiK3 的脉络,回顾其中主要的架构演进。

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
输入会获得词元嵌入和位置嵌入:

放大来看,每个 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

注意力(attention)的计算过程:
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 缓存(KV cache)源于一个很直接的观察:将生成的词元追加到输入之后,模型原本需要重新计算之前所有词元的投影。把它们的键向量和值向量保存下来,就能避免这些重复工作。
这块存储就是 KV 缓存。它保留前 N-1 个词元的向量,体积可能大到造成内存带宽瓶颈。
总体来看,在候选词元约 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
KimiK3 拥有 2.8 万亿参数,因此一个 KimiK3 模型的参数量大致相当于 22,580 个 GPT-2 模型。
线性注意力
Softmax 注意力在 q·k 乘积之后施加非线性,因此每个查询都会与每个键耦合。线性注意力则分别对 q 和 k 应用特征映射,例如 ELU+1 。这样一来,乘法就可以重新结合,于是不断增长的 K 、 V 向量集合可以被折叠进一个固定大小的 D×D 状态中。
论文用 O(N²) 来描述问题,一开始让我有些困惑。“Transformer 每个时间步的成本会随当前序列长度的平方增长”并不准确。那不是 FlashAttention 所解决的问题吗……随后我才注意到,这篇论文发表于 2020 年。
在当时,训练通常会显式生成完整的 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 读取和两次 1D 写入,而 KV 缓存会随序列长度以 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 和 k 交互之前分别对它们应用 ELU+1 。两种方法都会对最终分数进行归一化,但线性注意力所用的特征映射只是对 softmax 核的一种表达能力较弱的近似。这种近似可能降低保真度,不过实际精度损失取决于具体架构和工作负载。
注意,我们仍然会除以 qk 之和,只是为了简化,图中省略了这一步。从高层来看,注意力由三个步骤组成:
- 让 qk 分数变为非负。线性注意力使用 ELU+1,而 softmax 使用指数运算。
- 除以总和。
- 计算值的加权平均。 这保留了注意力的基本约定,但为了让 QK 分数非负,使用了表达能力较弱的特征映射。
DeltaNet(快速权重编程器)
有限的缓存必然需要覆盖或合并其中已有的信息。来自词元 i-1 的状态不会获得一个独立槽位,而是被加到同一个 D×D 矩阵中。因此,新的查询再也无法取回每个早期词元完全隔离的表示。
这种相加操作同时也是效率提升的来源。通过加法而不是拼接来更新缓存,可以避免缓存以 O(N) 增长,但同一操作也会导致信息相互干扰。 DeltaNet 正是为了解决这种可恢复性损失。

Schlag 的论文《 Fast Weight Programmers 》对此有一段精辟的表述:“当序列长度超过存储容量时,模型可能会进入容量过载状态。要在这种状态下正常运行,模型应当学会与内存内容进行动态交互,并有选择地决定保留哪些键值关联、删除哪些关联。纯加法指令可能并不适合这一目的……像公式 17 那样,不断向容量有限的内存添加新关联,最终必然会触及极限。”
线性注意力最具吸引力的场景——N 远大于 D——也恰好暴露了它的主要局限。一旦状态超出有效容量,各种关联就会开始相互干扰,因为更新采用加法,而且没有任何信息会离开缓存。
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.T @ v 的单个关联。如果用同一个键将它读回,就会得到 k @ (k.T @ v),也就是 (k @ k.T) v,等于 k 的范数平方乘以 v 。因此,读取结果会按键的范数平方进行缩放;如果把 k 归一化为单位长度,或者直接将结果除以该范数,就能精确恢复 v 。
Q 同样是一个学习得到的指针。 Wq 和 Wk 读取同一条残差流,而针对某项事实的查询会指向写入该事实时所使用的键方向。更新时,模型首先询问当前键能从缓存中取回什么信息。接着,它从想要存储的值中减去已有信息,用这个差值乘以键,再把结果加回去。旧信息由此被移除,新信息则写入原处。
DeltaNet(使用 Delta 规则并行化线性 Transformer)
这是整条演进线里最难的一节。我花了大约七个小时才形成一套切实可用的理解,因此下面会从实现出发逐步解释。简而言之,DeltaNet 借助广义 Householder 转移矩阵实现一阶线性递归,从而支持按分块并行的前向传播,实现硬件高效的线性时间训练。它把输入和输出拆成多个大小为 C 的分块,并根据前一个分块的最终状态,以及当前分块的查询、键、值块,计算每个分块的输出。
实际问题出在预填充(prefill)阶段。若直接在长度为 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)
与标准注意力不同,这种写法要求在每个键向量处执行一次校正,因此如何把它转化成并行矩阵乘法并不直观。即便没有 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,就会恢复标准的 O(N²) 注意力;令 C=1,则得到普通的线性注意力。取中间值时,我们通过增加分块内部的计算,换取更高的硬件利用率。在实践中,C 往往取 64 或 128,因为张量核心指令在这种粒度下运行效率很高,UMMA 就是一例。
在状态更新过程中,中间的分块会被折叠进 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)。这里先计算分数,采用带掩码的普通注意力顺序。跨块时,我们则遵循 (kᵀv)q,也就是先处理状态,采用递归顺序。注意力的计算量会按 O(N²) 增长,而这种方法不会。在块内,我执行完整的注意力计算,也就是用带掩码的 QKᵀ 乘以 V;跨块时,则把所有内容折叠到状态中,再用一次矩阵乘法读回。因此,成本分成两部分:固定部分是 2Ld²,对应状态计算,与 C 完全无关;增长部分是 2LCd,对应位于对角线上的分数矩阵。完整注意力只是 C 等于 L 的特例,此时第二项变为 2L²d,呈二次增长。所以 C 越小,所需 FLOPs 越少。
从纯 FLOP 数量来看,C=1 成本最低,但实际耗时未必最短。如果任务能高效映射到 GPU 的矩阵乘法硬件上,GPU 可以用更短时间完成更多算术运算。
下一步,是把同样的方法扩展到 DeltaNet 。

根本问题很简单:用于纯加法注意力的分块方法无法直接应用于 Delta 更新:
v_old = k_i @ S
u_i = b_i * (v_i - v_old)
为了计算需要减去的信息,我们必须按顺序获得每一个状态。如果不进行某种数学重参数化,就无法用同样的方式将其并行化。因此,论文把 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: 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)与 DeltaNet Transformer:

门控 DeltaNet
现在,我们已经有了一种能精确修改缓存的方法。每当出现一项新事实(即一个新的键向量),我们都能准确查看该位置原先存储的信息,并将其替换为我们希望注意的新信息。
这套机制只能忘记那些有明确替代项的关联。它无法在上下文切换时高效清除多个关联,也无法让记忆整体衰减以释放容量。
如果我们使用的是纯加法线性注意力:
加入遗忘能力并不困难,只需要一个参数来控制要遗忘的状态:
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:当 alpha 设为 1 时,切换为纯 Delta 规则;设为 0 时,则清空记忆。难点在于,如何用同样的并行分块方法实现它。
具体实现沿用了上一节介绍的 DeltaNet 重参数化。数学形式几乎完全相同,只增加了一项:一个由数据决定、取值在 0 到 1 之间的标量,用于控制旧状态的衰减。这就把有效的键值关联学习与自适应记忆管理结合到了一起。
相应的代码改动如下:

γʳ/γⁱ 项用于处理累积衰减。在时间步 x 写入、到 x+t 时读取的词元,已经依次乘上了 αₓαₓ₊₁αₓ₊₂…αₓ₊ₜ。这相当于前缀和计算的乘法版本。
最终得到的架构如下:

KDA/Kimi Linear
发展到这里,研究者开始尝试混合模型:在同一架构中组合多种注意力形式,例如把门控 DeltaNet 与 Mamba 结合起来。
Kimi Linear 因一项核心主张而引起关注:在受控对比中,它的表现优于完整注意力。研究团队将其描述为一种可直接替换的架构方案,质量更好,解码吞吐量最高还能提升 6 倍。
Kimi Linear 通过引入细粒度门控改进了门控 DeltaNet 。它不再只使用单个标量衰减值,而是为每个通道分别学习一个衰减值。

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

这里的 alpha.reshape(nb, C, d) 体现了论文最重要的贡献:对记忆衰减进行细粒度控制。
与 DeltaNet Transformer 并列比较时,Kimi Linear 架构引入了三项主要改动:
- 使用混合系统,交错插入多头潜在注意力(Multi-head Latent Attention,MLA)层。
- 用混合专家(Mixture-of-Experts,MoE)层替代 MLP 。
- 通过 alpha 投影扩展 DeltaNet 的容量。

后文会更详细地介绍 MLA 和 MoE 。目前最重要的一点是:这并不是盲目扩大规模。新增容量有明确的数学用途——逐通道缩放让模型可以更精细地控制记忆衰减。
缩放定律依然重要,但容量必须加在正确的位置,并以系统能够利用的形式加入。这条演进路线上的每一种架构,都通过增加容量来解决前一套系统中的某个具体局限。
Kimi K3
最终,KimiK3 的语言主干与上文的 Kimi Linear 模型相似。它包含 23 个四层宏周期。在每个宏周期中,三层使用 Kimi Delta Attention,第四层使用多头潜在注意力。第一层采用稠密前馈网络,其余所有层都使用潜在空间混合专家。
乍看之下,相比 Kimi Linear,改动似乎不算大:
- 显著扩大模型规模
- 每 12 层加入一次块级 AttnRes
- MLA 查询 LoRA 与输出门控
- 潜在空间 MoE
- SiTU 激活函数
- 门控 MLA KDA 提供状态大小恒定的递归记忆,而周期性出现的 MLA 层则保留对上下文进行完整 softmax 检索的能力。下面这张简化示意图可以作为理解后续改动的参考。

我们先从几项较为直接的改动讲起:门控 MLA 、潜在空间 MoE 和 SiTU 激活函数。
门控 MLA 决定从 MLA 检索出的每项特征,有多少能够进入残差流。具体做法是,将这些特征与一个由输入投影得到的门逐元素相乘。
在传统 MoE 中,学习得到的路由器通过点积相似度,把每个词元发送到一部分专家网络。 KimiK3 共有 898 个专家。其中两个是共享专家,会处理每个词元;在剩余的 896 个专家中,路由器会为每个词元选择 16 个。
KimiK3 还修改了专家所用的激活函数。传统方式是对上投影应用 SiLU,再将结果与门逐元素相乘,最后进行下投影;KimiK3 则使用 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)
模型还会先对共享专家的输入进行下投影,再对这些专家的最终求和结果进行上投影:

这体现了模型推理中一个反复出现的难题:如果没有融合内核,新激活路径的速度几乎比原始路径慢 3 倍。作为补偿性优化,专家会在压缩后的潜在空间中运行,这能让它们的前向传播快得多,并使 FLOPs 几乎减半。
其余改动包括 MLA 查询 LoRA 、输出门控,以及每 12 层设置一次块级注意力残差(Attention Residuals)。 AttnRes 会增加约 2% 的推理延迟,但带来两项重要收益:
- 选择性检索较早的表示,从而缓解残差稀释和隐藏状态增长
- 1.25 倍的计算优势 AttnRes 与 MLA 从不同方向解决了同一项根本局限。 KDA 层使用固定大小的状态,因此不可避免地需要丢弃信息。 MLA 从词元上下文中检索信息,而 AttnRes 则从深度方向上更早的表示中检索信息。
AttnRes
感谢 @chloey3k 对本节的帮助。在每次前向传播中,输入会依次通过一组堆叠层。这里,每一层都由一个注意力块(KDA 或 MLA)和一个 MLP 或 MoE 块构成。通常,每一层的输入是原始嵌入与之前所有层输出之和,且各项权重完全相同。
$$ h_l = h_1 + \sum_{i=1}^{l-1} f_i(h_i) $$
这里,h_i 是第 i 层的输入,h_1 是当前词元的嵌入(即截至当前序列的最后一个词元),而 f_i(h_i) 是第 i 层的输出(一个注意力块或 MLP 块)。
问题在于缺少选择性访问能力。不同类型的层接收到的是同一个聚合状态,尽管不同层可能更适合采用不同的加权方式。由于这种递归是纯加法的,越靠后的层还必须学会生成越来越大的输出,才能影响不断累积的残差,这可能使训练变得不稳定。 AttnRes 不再平等对待所有层,而是给求和中的每一项乘上一个专门的权重,让模型可以根据上下文,为最有用的层赋予更高的重要性。
$$ h_l = \alpha_0 \cdot h_1 + \sum_{i=1}^{l-1} \alpha_i \cdot f_i(h_i) $$
每个权重 alpha_i 都通过查询与键的点积计算得到。查询是逐层学习的,而键和值来自更早的残差流状态。分数经过归一化,使其总和为 1,然后用于构造这些状态的加权组合。

因此,模型不必只以紧邻的前一层为条件。 AttnRes 让每一层都能有选择地访问更早的层输出,并允许学习得到的查询检索对当前计算最有用的表示。
下面的伪代码在块粒度上应用了相同思想。一个块是跨 12 个解码器层累积的注意力输出和 MLP 输出逐元素求和的结果,它被存储为单个深度表示,供之后的 AttnRes 混合使用。
如果每一层都应用残差注意力,训练和推理成本会增加太多。只在固定的块边界应用,则能以较低成本获得大部分收益。在 KimiK3 中,每经过 12 个解码器层就会形成一个边界。 23 个四层宏周期总共构成八个 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 检索、稀疏专家容量,以及对深度方向残差的选择性访问。最终得到的系统,会把新增容量投入具有明确功能作用的位置。
容量固定的关联记忆(维度固定)必须具备淘汰策略,因为纯加法的线性操作在容量用尽后,最终必然引入干扰。为此,系统需要门控、路由或衰减等学习式选择机制,而注意力则是最有效的选择性读取机制。