跳到正文
LLM 推理
S03模型本身·251

Transformer 前向传播

现代 decoder block 只有七个操作。用 NumPy 老老实实写一遍,之后所有 kernel 优化就都有了对照的正确答案。

  • RMSNorm
  • RoPE
  • 分组查询注意力
  • SwiGLU
  • 残差流

为什么重要

问题

接下来你要花十七章让 model.forward() 变快。在那之前,先得有一个毫无歧义地正确的版本:融合 kernel 产出「看着像模像样、其实是垃圾」的结果时,你需要有个东西能拿来对拍。

现代 decoder block 只有七个操作,整个 Llama 系列架构就是这个 block 的重复。没有 encoder、没有 cross-attention、没有可学习的位置嵌入、没有 bias。下面的全部内容用 250 行 NumPy 就装得下。

核心思路

解法

这就是完整的 block。先读一遍;本章余下部分都是注解。

一个 decoder blockpython
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

KV 缓存之所以成立,靠的就是这一点。注意力分数取决于 q 与 k 的相对位置,而旋转恰好编码了这一点。value 不携带任何位置信息,缓存下来的 v 永远有效;把旋转后的 K 缓存起来,位置也就再不用重算。

工作原理

工作原理

示意图decoder block 解剖
一个 DECODER BLOCK —— 重复 L 次残差流 [T, d_model]注意力子层rms_normq_projk_projv_projRoPE只作用于 q、k注意力softmax(qkᵀ/√d)vo_proj+ 残差KV cache唯一保留的状态k 和 v 是在 RoPE 之后缓存的。q 每一步都被丢弃。前馈子层rms_normW1 gateW3 upsilu×W2 down+ 残差Pre-norm:每个子层读取一份归一化后的副本,再把未归一化的修正量写回残差流。FFN 占了全部参数的约 2/3;decode 时,它也就占你每个 token 搬运字节的约 2/3。DECODE 时的形状x[1, d]q[1, H·dh]k,v[1, Hkv·dh]out[1, d]
S05 到 S08 四章讲的都是右侧那条虚线引出的分支:k 和 v 是唯一值得在多次前向之间保留的张量。

分组查询注意力,以及它为什么存在

完整的多头注意力给每个 query 头配一套独立的 key 和 value 头。开销恰好落在最不该落的地方:KV 缓存的大小与 KV 头数成正比,而缓存容量决定了你能同时跑多少请求。

GQA 让一个 KV 头被一组 query 头共享。Llama-3-70B 用 64 个 query 头配 8 个 KV 头,缓存小 8 倍,质量损失小到所有人现在都这么做。多查询注意力(MQA)是极端情形:总共只有一个 KV 头。

GQA 是复制,不是重写python
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

把重复后的 K 和 V 真的物化出来,等于把 GQA 省下的显存又还回去。生产 kernel 会在注意力内层循环里直接索引那个共享的 KV 头:在 FlashAttention 里是一次 stride 计算,在 PagedAttention 里是一次 block table 查表。

FLOPs 和字节花在了哪里

两个数字就能描述 block 里的每个操作:它做多少算术,以及为此必须搬运多少字节。二者之比就是算术强度,它决定了 GPU 是在干活还是在等待。

prefill 阶段序列很长,权重被数百个 token 复用,算术强度很高。decode 阶段序列长度是 1:你读完整个模型只为产出一个 token,算术强度塌缩到约每字节 1 次浮点运算,而硬件想要的是 300。下面的模拟器会把这一点变得具体。

动手观察

动手试试

模拟器block 解剖浏览器
参数量 / 层
45.1M
参数量 / 模型
992.0M
head 维度
64
GQA 比例
4:1
一个 decoder block · 点击某个算子

柱子表示当前序列长度下每次前向的乘加次数。每层合计 90.2M,整个模型每 token 合计 2.0G。

已选中 · rms_norm

输出形状 [1, 2048]

参数量 2.0k · 浮点运算数 4.1k

不减均值,也不加 bias:RMSNorm 就是把 LayerNorm 里无关紧要的部分删掉。它很便宜,但要把残差流完整读一遍,所以受带宽限制。

分组查询注意力

每种色块是一个被 4 个 query 头共享的 KV 头。KV 缓存:每层每 token 2048 B,完整多头则要 8192 B;省下 4.0×,质量几乎无损。

算术强度 —— decode 为什么受带宽限制
1.0seq 1
1.0seq 8
1.0seq 64
1.0seq 512
  • 每读取一字节权重对应的浮点运算次数

H100 大约需要每字节 300 次浮点运算才能把算力吃满,而序列长度为 1 时只有 1 左右。权重被全速读入,乘法器却在空转。这道落差,正是批处理(S09)和投机解码(S14)成为全课程杠杆率最高的两项优化的原因。

有两个实验值得做。第一,把 kv 头数 从 32 拖到 1:每 token 的缓存开销降低 32 倍,参数量几乎不动,这就是 GQA 的价值所在。第二,把 序列长度 从 1 拖到 512,看算术强度爬升。本课程后面所有关于批处理的章节,都是在设法换来这段爬升,同时又不让用户等待。

亲手实现

实现

RoPE 值得单独看一眼,因为你第一次写出来的实现,几乎从来不是权重所期待的那一个。

code/s03_transformer.py(节选)python
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 out

RoPE 有两种布局,而且互不兼容

NeoX(对半切分)布局把维度 ii + D/2 配对(如上)。GPT-J / 交错布局则把 2i2i+1 配对。两者都叫「RoPE」。选错了,模型照样能跑,照样输出语法正确的英文,只是隐隐地不连贯。这是本课程中最令人泄气的一个 bug。HF 格式的 Llama 和 Qwen 检查点用 NeoX(转换器会对 wq/wk 做置换来实现);Meta 原版和 GGUF 的 Llama 权重则是交错布局。去查 config 和文件格式,不要猜。

另一个陷阱是 theta 基数。长上下文模型会把它调大(Llama-3 用 500000 而不是 10000),以减慢低频维度的旋转、拉长它们的波长,让远距离的位置仍可区分。最高频的那一对维度根本不依赖 theta。用默认基数加载一个长上下文模型,几千 token 之后质量会断崖式下跌。

在本地运行
构建一个随机初始化的小型 Llama 风格模型,跑一次前向,并打印每个子层之后残差流的形状、均值和标准差。还包含一项数值检查:当 n_kv_heads == n_heads 时,GQA 与普通多头注意力完全一致。
$ python code/s03_transformer.py
预期输出: 一张逐层激活统计表(数值稳定)、GQA 等价性检查通过的断言,以及一份参数量拆解,显示 FFN 占绝对主导。

只需要 NumPy — 查看环境准备.

生产实践

在生产环境中

  • picoLM 用约 340 行 C 写出这个 block,并预先计算好 RoPE 的正余弦表。三角函数被提到内层循环之外,因为热循环里的 sinf 比它周围的 matmul 还慢。
  • vLLM 把这里的每一行都换成了融合 kernel:RMSNorm 与残差加法融合、QKV 打包成一次 matmul、RoPE 融进注意力的前导部分。
  • quant.cpp 在加载时自动探测架构变体(Qwen3 的 QK-norm、NeoX 还是交错 RoPE、双 FFN),因为同一个 GGUF 加载器要服务七种架构。

练习

  1. 1
    q_projk_projv_proj 融合成对拼接权重的一次 matmul,再把结果切开。确认输出逐位一致,并测量加速比。
  2. 2
    在 NeoX 之外再实现一份交错布局的 RoPE,写个测试展示相同权重下两者产生不同的注意力分数。日后真遇上这个故障,你就能认出来。
  3. 3
    在 config 开关后面加上 QK-norm(像 Qwen3 那样,在 RoPE 之前对 q 和 k 各做一次 RMSNorm)。验证开关关闭时输出不变——然后在 norm 权重全为 1 的情况下打开它,观察输出仍然会变,因为 RMSNorm 会按每个向量自身的 RMS 对它做缩放。

继续学习

接下来

block 的终点是 logits。S04 用 temperature、top-k、top-p、min-p 和重复惩罚把它们变成一个 token,并说明采样器虽然只花微秒级时间,却是大多数「感知到的质量问题」真正的所在。

自测

习题

先作答,再看解析。答错比答对更有价值,因为解析会指出你该回头重读哪一部分。

自测 1 题 / 共 6

为什么 RoPE 作用于 q 和 k,却从不作用于 v?

得分 0/6