Transformer 前向传播
现代 decoder block 只有七个操作。用 NumPy 老老实实写一遍,之后所有 kernel 优化就都有了对照的正确答案。
- RMSNorm
- RoPE
- 分组查询注意力
- SwiGLU
- 残差流
为什么重要
问题
接下来你要花十七章让 model.forward() 变快。在那之前,先得有一个毫无歧义地正确的版本:融合 kernel 产出「看着像模像样、其实是垃圾」的结果时,你需要有个东西能拿来对拍。
现代 decoder block 只有七个操作,整个 Llama 系列架构就是这个 block 的重复。没有 encoder、没有 cross-attention、没有可学习的位置嵌入、没有 bias。下面的全部内容用 250 行 NumPy 就装得下。
核心思路
解法
这就是完整的 block。先读一遍;本章余下部分都是注解。
def block(x, w, pos, cache=None):
# ---- attention sublayer -------------------------------------
h = rms_norm(x, w.attn_norm)
q = h @ w.wq # [T, n_heads * head_dim]
k = h @ w.wk # [T, n_kv_heads * head_dim]
v = h @ w.wv
q, k = rope(q, pos), rope(k, pos) # position enters here, only here
a = grouped_query_attention(q, k, v, causal=True)
x = x + a @ w.wo # residual
# ---- feed-forward sublayer ----------------------------------
h = rms_norm(x, w.ffn_norm)
x = x + (silu(h @ w.w1) * (h @ w.w3)) @ w.w2 # SwiGLU
return x其中有四个设计选择值得理解,因为每一个都改变了引擎必须做的事:
- pre-norm,而不是 post-norm。 归一化发生在进入每个子层的路上,残差流本身从不被归一化。这让深层堆叠可训练。对我们而言,它意味着残差流是这个 block 唯一的状态,其余一切都是它的纯函数。
- RMSNorm,而不是 LayerNorm。 不减均值、不加 bias。一次归约,一次乘法。
- RoPE,而不是可学习位置。 位置信息是通过旋转 q 和 k 中成对的维度注入的。残差流里什么都没有被加进去,注意力看到的只有相对位置,YaRN 这类上下文扩展技巧正是建立在这一点上。单靠 RoPE 在超过训练长度后仍会退化,这些技巧就是为此而生。
- SwiGLU,而不是 GELU-MLP。 三次 matmul 而非两次:一个门控、一个升维、一个降维。
RoPE 只作用于 q 和 k,绝不作用于 v
工作原理
工作原理
分组查询注意力,以及它为什么存在
完整的多头注意力给每个 query 头配一套独立的 key 和 value 头。开销恰好落在最不该落的地方:KV 缓存的大小与 KV 头数成正比,而缓存容量决定了你能同时跑多少请求。
GQA 让一个 KV 头被一组 query 头共享。Llama-3-70B 用 64 个 query 头配 8 个 KV 头,缓存小 8 倍,质量损失小到所有人现在都这么做。多查询注意力(MQA)是极端情形:总共只有一个 KV 头。
def grouped_query_attention(q, k, v, n_heads, n_kv_heads, causal=True):
head_dim = q.shape[-1] // n_heads
q = q.reshape(-1, n_heads, head_dim)
k = k.reshape(-1, n_kv_heads, head_dim)
v = v.reshape(-1, n_kv_heads, head_dim)
reps = n_heads // n_kv_heads # e.g. 4 query heads per kv head
k = np.repeat(k, reps, axis=1) # broadcast, do not copy in a real kernel
v = np.repeat(v, reps, axis=1)
scores = np.einsum("qhd,khd->hqk", q, k) / np.sqrt(head_dim)
if causal:
scores += causal_mask(q.shape[0], k.shape[0])
return np.einsum("hqk,khd->qhd", softmax(scores), v).reshape(-1, n_heads * head_dim)np.repeat 是教学道具,不是 kernel
FLOPs 和字节花在了哪里
两个数字就能描述 block 里的每个操作:它做多少算术,以及为此必须搬运多少字节。二者之比就是算术强度,它决定了 GPU 是在干活还是在等待。
prefill 阶段序列很长,权重被数百个 token 复用,算术强度很高。decode 阶段序列长度是 1:你读完整个模型只为产出一个 token,算术强度塌缩到约每字节 1 次浮点运算,而硬件想要的是 300。下面的模拟器会把这一点变得具体。
动手观察
动手试试
- 参数量 / 层
- 45.1M
- 参数量 / 模型
- 992.0M
- head 维度
- 64
- GQA 比例
- 4:1
柱子表示当前序列长度下每次前向的乘加次数。每层合计 90.2M,整个模型每 token 合计 2.0G。
输出形状 [1, 2048]
参数量 2.0k · 浮点运算数 4.1k
不减均值,也不加 bias:RMSNorm 就是把 LayerNorm 里无关紧要的部分删掉。它很便宜,但要把残差流完整读一遍,所以受带宽限制。
每种色块是一个被 4 个 query 头共享的 KV 头。KV 缓存:每层每 token 2048 B,完整多头则要 8192 B;省下 4.0×,质量几乎无损。
- 每读取一字节权重对应的浮点运算次数
H100 大约需要每字节 300 次浮点运算才能把算力吃满,而序列长度为 1 时只有 1 左右。权重被全速读入,乘法器却在空转。这道落差,正是批处理(S09)和投机解码(S14)成为全课程杠杆率最高的两项优化的原因。
有两个实验值得做。第一,把 kv 头数 从 32 拖到 1:每 token 的缓存开销降低 32 倍,参数量几乎不动,这就是 GQA 的价值所在。第二,把 序列长度 从 1 拖到 512,看算术强度爬升。本课程后面所有关于批处理的章节,都是在设法换来这段爬升,同时又不让用户等待。
亲手实现
实现
RoPE 值得单独看一眼,因为你第一次写出来的实现,几乎从来不是权重所期待的那一个。
def rope(x, positions, theta=10000.0):
"""Rotate pairs of dimensions by an angle proportional to position."""
T, D = x.shape
half = D // 2
# frequency per dimension pair: low dims rotate fast, high dims slow
inv_freq = 1.0 / (theta ** (np.arange(0, half) * 2.0 / D))
angles = positions[:, None] * inv_freq[None, :] # [T, half]
cos, sin = np.cos(angles), np.sin(angles)
# NeoX layout: first half pairs with second half.
x1, x2 = x[:, :half], x[:, half:]
out = np.empty_like(x)
out[:, :half] = x1 * cos - x2 * sin
out[:, half:] = x2 * cos + x1 * sin
return outRoPE 有两种布局,而且互不兼容
i 与 i + D/2 配对(如上)。GPT-J / 交错布局则把 2i 与 2i+1 配对。两者都叫「RoPE」。选错了,模型照样能跑,照样输出语法正确的英文,只是隐隐地不连贯。这是本课程中最令人泄气的一个 bug。HF 格式的 Llama 和 Qwen 检查点用 NeoX(转换器会对 wq/wk 做置换来实现);Meta 原版和 GGUF 的 Llama 权重则是交错布局。去查 config 和文件格式,不要猜。另一个陷阱是 theta 基数。长上下文模型会把它调大(Llama-3 用 500000 而不是 10000),以减慢低频维度的旋转、拉长它们的波长,让远距离的位置仍可区分。最高频的那一对维度根本不依赖 theta。用默认基数加载一个长上下文模型,几千 token 之后质量会断崖式下跌。
$ python code/s03_transformer.py只需要 NumPy — 查看环境准备.
生产实践
在生产环境中
- picoLM 用约 340 行 C 写出这个 block,并预先计算好 RoPE 的正余弦表。三角函数被提到内层循环之外,因为热循环里的
sinf比它周围的 matmul 还慢。 - vLLM 把这里的每一行都换成了融合 kernel:RMSNorm 与残差加法融合、QKV 打包成一次 matmul、RoPE 融进注意力的前导部分。
- quant.cpp 在加载时自动探测架构变体(Qwen3 的 QK-norm、NeoX 还是交错 RoPE、双 FFN),因为同一个 GGUF 加载器要服务七种架构。
练习
- 1把
q_proj、k_proj、v_proj融合成对拼接权重的一次 matmul,再把结果切开。确认输出逐位一致,并测量加速比。 - 2在 NeoX 之外再实现一份交错布局的 RoPE,写个测试展示相同权重下两者产生不同的注意力分数。日后真遇上这个故障,你就能认出来。
- 3在 config 开关后面加上 QK-norm(像 Qwen3 那样,在 RoPE 之前对 q 和 k 各做一次 RMSNorm)。验证开关关闭时输出不变——然后在 norm 权重全为 1 的情况下打开它,观察输出仍然会变,因为 RMSNorm 会按每个向量自身的 RMS 对它做缩放。
继续学习
接下来
block 的终点是 logits。S04 用 temperature、top-k、top-p、min-p 和重复惩罚把它们变成一个 token,并说明采样器虽然只花微秒级时间,却是大多数「感知到的质量问题」真正的所在。
自测
习题
先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。
为什么 RoPE 作用于 q 和 k,却从不作用于 v?