Kimi K 3竟是GPT-2的22580倍,博主「肝」48小时发现:七年进化大模型不只是参数暴涨
最近,月之暗面 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) 级别线性增长。