架构演进 · 仅解码器

LLaMA

LLaMA 3 8B · GPT-2 之后改了什么

对比
LLaMA 3 8B vs GPT-2 small
层
32 GPT-2 12
dmodel
4,096 GPT-2 768
头
32 个 q · 8 个 kv GPT-2 12
dff
14,336 GPT-2 3,072
词表
128,256 GPT-2 50,257
上下文
8,192 GPT-2 1,024

全部步骤

  1. GPT-2 vs LLaMA

    LLaMA 保留了 GPT-2 的块:前置归一化、残差相加、因果注意力。下排标出的四处有改动;点击其中一处跳过去。另外还有不同:任何地方都没有偏置项,输出矩阵不与嵌入共享(不是 WEᵀ),最后的归一化是 RMSNorm,词表有 128K 个词元。

    4 处改动 · 同样的块

  2. 旋转位置编码

    GPT-2 只在输入处加一次学习得到的位置向量。LLaMA 则在每个注意力层内,把查询和键的每一对数旋转一个随位置增大的角度。用下方的 q、k 和平移控件来验证。

    θⱼ = base^(−2j / dhead)

  3. RMSNorm

    LayerNorm 把每个词元中心化并缩放到单位离散度。RMSNorm 只按均方根重新缩放:更简单,略快,而且同样稳定。

    x / √(mean(x²) + ε)

  4. SwiGLU MLP

    GPT-2 的 MLP 先变宽、施加 GELU、再变窄。LLaMA 的 MLP 并排跑两个投影,让一个给另一个做门控,这种结构叫 SwiGLU。

    (SiLU(x·Wgate) ⊙ x·Wup) · Wdown

  5. 分组查询注意力

    生成时,过去词元的键和值保存在 KV 缓存中,它随每个词元、每一层增长。GPT-2 给每个头各自的键和值;LLaMA 3 让每个键/值头被 4 个查询头共享,所以缓存小了 4 倍。

    32 个 q 头 · 8 个 kv 头

代码

# RMSNorm
x = x * torch.rsqrt(x.pow(2).mean(-1, keepdim=True) + eps) * self.weight
# RoPE, inside attention: rotate each pair of q and k by position × θ
q, k = apply_rotary_pos_emb(q, k, cos, sin)
# GQA: 8 key/value heads serve 32 query heads
k, v = repeat_kv(k, 4), repeat_kv(v, 4)
# SwiGLU MLP
y = self.down_proj(F.silu(self.gate_proj(x)) * self.up_proj(x))
# the block, shaped like GPT-2's
h = x + self.self_attn(self.input_layernorm(x))
out = h + self.mlp(self.post_attention_layernorm(h))

延伸阅读