机器之心编辑部
最近,月之暗面 kimi 正式开源 Kimi K3 完整模型权重,Kimi K3 是一款总参数量达 2.8 万亿、上下文窗口达 100 万 token 的 MoE 大模型,更是全球首个落地的近 3 万亿参数级开源大模型,引起业界热议。
其中一个博主 ali@waterloo_intern 意识到,其实从 2019 年 OpenAI 发布的参数量约 1.24 亿的 GPT-2,到 2026 年 2.8 万亿参数量的 Kimi K3,只有短短七年的时间,但两个模型规模相差 22580 倍!
简单换算,相当于把大约 22580 个 GPT-2 Small 装进一个 Kimi K3。
这引起了他的好奇:「但这一切,真的只是规模变大了吗?」
对此,ali 称自己花了约 48 小时阅读 Kimi K3 的建模代码和 8 篇论文,最终理清了从 2019 年 GPT-2 到 Kimi K3 的完整技术谱系。「我将带你回顾我们是如何一步步走到今天,以及从 GPT-2 到 Kimi K3,模型究竟发生了多少变化?又有哪些东西其实始终没有改变?我们会沿着这条技术演进路线,梳理最终通向 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 嵌入和位置嵌入:
![]()
把每一个 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,便可以选择下一个 token。
这也是仅解码器生成方式的一处低效之处:模型会为输入序列中的每一个位置计算表示,但在每一步解码时,真正会被用到的只有最后一个位置的 logits。如果没有缓存机制,在生成下一个 token 时,大量计算都需要重新执行。
![]()
KV Cache 源于一个非常直接的观察:当新生成的 token 被追加到输入序列后,模型原本需要重新计算此前所有 token 的投影。将这些 token 对应的 Key 和 Value 向量保存下来,就可以避免这部分重复计算。
这些被保存的数据,就是 KV Cache。它会保留前面 N-1 个 token 的向量,规模可能变得非常庞大,甚至形成内存带宽瓶颈。
总体来看,在词表规模约为 5 万、包含 12 个 Transformer 模块、12 个注意力头、嵌入维度为 768 的情况下,这个基线模型大约拥有 1.24 亿个参数。
线性注意力
Softmax 注意力是在 q・k 乘积完成之后再施加非线性变换,因此每一个 Query 都会与每一个 Key 相互耦合。而线性注意力则会分别对 q 和 k 应用特征映射,例如 ELU+1。这样一来,矩阵乘法就可以重新结合,持续增长的 K、V 向量也能够被压缩进一个固定大小的 D×D 状态中。
作者表示,论文中关于 O (N²) 的描述一度让他感到困惑。严格来说,「Transformer 每个时间步的计算成本会随当前序列长度的平方增长」并不准确。FlashAttention 解决的正是这个问题…… 随后他才发现,这篇论文发表于 2020 年。
当时,训练通常会显式构建完整的 N×N 注意力矩阵,FlashAttention 还没有出现,而许多参考级的自回归实现也没有使用 KV Cache,需要反复计算此前所有 token 的历史状态。
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 Cache 的大小会随着序列长度以 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 注意力则使用指数运算;
- 除以所有分数之和,完成归一化;
- 根据归一化后的权重,对 Value 进行加权平均。
线性注意力保留了注意力机制的基本计算形式,但为了使 QK 分数非负,它采用了一种表达能力相对较弱的特征映射。
DeltaNet(快速权重编程器)
有限容量的缓存,必然需要覆盖已有信息,或者将新信息与旧信息合并。来自第 i-1 个 token 的状态并不会获得一个独立的存储槽位,而是被写入同一个 D 到 D 矩阵中。因此,新的 Query 无法再从中取回每个历史 token 彼此完全隔离的表示。
这种累加写入,正是效率提升的来源。通过加法更新缓存,而非不断拼接新的内容,缓存规模便不会随着序列长度以 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,随后使用同一个 Key 读取,就会得到 k @ (k.T @ v),即 (k @ k.T) v,实际上等于 k 的范数平方乘以 v。因此,读取结果会被 Key 的范数平方缩放。将 k 归一化为单位长度,或者直接用结果除以其范数,就可以精确恢复出 v。
Q 同样可以看作一个学习得到的指针。Wq 和 Wk 都从同一条残差流中读取信息,而当模型查询某个事实时,对应的 Query 会指向这个事实最初写入时所对应的 Key 方向。
在更新状态时,模型首先会检查当前 Key 能够从缓存中读取出什么信息。接着,它用希望存储的新 Value 减去当前已经读出的旧信息,再将这个差值与 Key 相乘,并把结果加回状态矩阵。
这样一来,旧信息会被移除,新信息则会写入原来的位置。
DeltaNet(通过 Delta Rule 并行化线性 Transformer)
这是本文最难理解的一部分,博主表示,为了真正搞懂它,大约花了七个小时,因此接下来会从具体实现出发,一步步展开解释。
简单来说,DeltaNet 实现了一种一阶线性递归,并使用广义 Householder 变换矩阵作为状态转移矩阵,从而支持按分块并行执行前向传播,实现更适合硬件的线性时间训练。
它会把输入和输出划分为若干个大小为 C 的分块,并根据前一个分块的最终状态,以及当前分块中的 Query、Key、Value 矩阵,计算该分块的输出。
实际需要解决的问题是 prefill,也就是上下文预填充阶段。
如果直接在长度为 T) 的序列上实现 Delta Rule,计算过程大致如下:
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 Rule,直接实现线性注意力的 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)
采用分块形式,可以得到一种效率更高的实现方式。通过一个例子,更容易理解其中的计算机制:
![]()
当 C=N 时,这种方法会退化为标准的 O (N²) 注意力;当 C=1 时,它对应普通的线性注意力。
在二者之间选择不同的 C,相当于在计算量和硬件利用率之间进行权衡:分块内部会增加一部分计算,但能够更充分地利用硬件。
在实践中,C 通常会被设置为 64 或 128,因为 Tensor Core 的指令能够在这类粒度上高效运行,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,也就是递归式的计算顺序:先构建状态,再用 Query 从状态中读取信息。标准注意力的计算量会以 O (N²) 增长,而这种方法不会。
具体来说,在每个分块内部,仍然执行真正的注意力计算,也就是带掩码的 QKᵀ与 V 相乘;而在分块之间,会把所有历史信息压缩进状态,再通过一次矩阵乘法将其读取出来。
因此,整体计算成本可以拆分成两部分:
- 第一部分是固定开销 2Ld²:这部分来自状态矩阵的计算,与分块大小 C 无关;
- 第二部分会随着 C 增长 2LCd:它对应分布在矩阵对角线上的分块内注意力分数矩阵。
完整注意力只是 C=L 的特殊情况,此时第二项会变成 2L²d,计算复杂度也就重新变成了平方级。
因此,从 FLOPs 的角度看,C 越小,需要执行的计算越少。
当 C=1 时,理论 FLOPs 最低,但它未必能带来最短的实际运行时间。只要计算任务能够高效映射到 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 Transformer 与 DeltaNet Transformer。
![]()
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 Rule 可以单独更新某一条事实,却无法让其余事实自然衰减。
因此,Gated Delta Rule 将 Mamba 的门控更新规则与 Delta Rule 结合起来。它引入参数 alpha:当 alpha=1 时,更新退化为纯 Delta Rule;当 alpha=0 时,记忆会被完全清空。
这里的难点,是如何继续使用前文介绍的分块并行方法来实现这一机制。
具体实现仍然采用上一节介绍的 DeltaNet 重参数化方法。整体数学形式几乎相同,只增加了一项:一个由数据动态决定、取值范围在 0 到 1 之间的标量,用来控制旧状态的衰减程度。
这样一来,模型便同时具备了有效学习键值关联的能力,以及自适应管理记忆的能力。
相应的代码改动如下:
![]()
其中,γʳ/γⁱ项用于计算累计衰减。
假设某个 token 在时间步 x 被写入,并在 x+t 时被读取,那么它所经历的累计缩放为:αₓαₓ₊₁αₓ₊₂…αₓ₊ₜ。
这可以看作前缀和计算在乘法形式下的对应版本。
最终得到的架构如下:
![]()
KDA / Kimi Linear
发展到这一步,研究人员开始尝试混合架构:在同一个模型中组合多种注意力机制,例如将 Gated DeltaNet 与 Mamba 结合起来。
Kimi Linear 之所以受到关注,核心在于它提出了一项重要结论:在控制变量的对比实验中,Kimi Linear 的表现超过了全注意力架构。
论文作者将其描述为一种可以直接替换传统注意力的架构方案,不仅模型效果更好,解码吞吐量最高还能提升至原来的 6 倍。
Kimi Linear 对 Gated DeltaNet 的主要改进,是引入了更加细粒度的门控机制。
此前,模型只使用一个标量控制整体衰减;Kimi Linear 则为每一个通道分别学习一个衰减值。
![]()
KDA 的更新规则依然相似,但对应代码变成了下面这样:
![]()
其中,alpha.reshape (nb, C, d) 体现了这篇论文最重要的贡献:对记忆衰减进行细粒度控制。
与 DeltaNet Transformer 相比,Kimi Linear 架构主要引入了三项变化:
- 采用混合架构,在模型中交替插入多头潜在注意力(Multi-head Latent Attention,MLA)层;
- 使用混合专家(MoE)层替代传统 MLP;
- 通过 alpha 投影,为 DeltaNet 增加额外容量。
![]()
需要理解的重点是,这并非单纯、盲目地扩大模型规模。新增的容量有着明确的数学用途:逐通道缩放机制,使模型能够更加精细地控制记忆衰减。
Scaling Law 依然成立,但模型容量必须被增加在正确的位置,并采用系统真正能够利用的形式。在这条架构演进路径上,每一种新架构增加容量,都是为了解决上一代系统中某个具体的局限。
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 中,一个学习得到的路由器会根据点积相似度,将每个 token 分配给一组专家网络。
Kimi K3 总共拥有 898 个专家,其中两个是共享专家,会处理所有 token;剩余的 896 个专家中,路由器会为每个 token 选择 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)
模型还会先将输入降维投影到共享专家空间,并在汇总共享专家的输出后,再将其升维投影回去:
![]()
这揭示了模型推理中一个反复出现的难题:如果没有融合算子,新激活函数的运行速度几乎比原始计算路径慢 3 倍。
一个能够抵消部分开销的优化是,让专家网络在压缩后的潜在空间中运行,这样可以大幅加快专家网络的前向传播,并将 FLOPs 几乎减少一半。
剩余的改动包括 MLA Query LoRA、输出门控,以及每隔 12 层加入一次分块式注意力残差(AttnRes )。AttnRes 会使推理延迟增加约 2%,但它能带来两项重要收益:
- 有选择地检索早期表示,从而缓解残差流中的信息稀释和隐藏状态幅度不断增长的问题;
- 获得约 1.25 倍的计算优势。
AttnRes 和 MLA 从不同方向解决了同一个底层局限。
KDA 层使用固定大小的状态,因此不可避免地需要丢弃一部分信息。MLA 从 token 上下文中检索信息,而 AttnRes 则从网络深度方向上更早的表示中进行检索。
AttnRes(注意力残差)
在每一次前向传播中,输入都会经过一系列堆叠的网络层。这里,每一层都由一个注意力模块(KDA 或 MLA)和一个 MLP 或 MoE 模块组成。
通常情况下,某一层的输入,是原始嵌入与此前所有层输出之和,并且所有部分的权重都相同:
![]()
![]()
这种方式的问题在于缺少选择性访问能力。
不同类型的网络层接收到的都是同一个聚合状态,尽管它们可能更适合使用不同的权重组合。
此外,由于这种递归完全依赖加法,越靠后的层必须学会产生幅度越来越大的输出,才能对不断累积的残差产生足够影响,这可能导致训练过程不稳定。
AttnRes 不再平等对待所有层,而是为求和表达式中的每一项乘上一个专门计算的权重,使模型能够根据当前上下文,更加重视其中最有用的网络层:
![]()
每一个权重 alpha_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 检索、稀专家容量,以及对深度方向残差表示的选择性访问结合在一起。
最终得到的,是一个会将额外容量投入到明确功能位置上的系统。
归根结底,一个容量固定的关联记忆系统,也就是维度保持不变的记忆系统,必须具备某种淘汰策略。
因为当记忆达到容量上限后,纯累加式的线性操作必然会导致不同信息相互干扰。
因此,系统必须引入门控、路由或衰减等学习得到的选择机制;而注意力机制,仍然是目前最有效的选择性读取方式。
更多信息,可以查看完整文章了解!
https://x.com/waterloo_intern/status/2081762991532560503
https://x.com/waterloo_intern/status/2081762065392541951
特别声明:以上内容(如有图片或视频亦包括在内)为自媒体平台“网易号”用户上传并发布,本平台仅提供信息存储服务。
Notice: The content above (including the pictures and videos if any) is uploaded and posted by a user of NetEase Hao, which is a social media platform and only provides information storage services.