Published on

关于 Prefill 一些模糊点的理解

Authors
  • avatar
    Name
    Vegetog
    Twitter

虽然学过 prefill,但我对它的并行方式一直比较模糊,因此记录一下。

本文以标准 decoder-only Transformer 的因果注意力为例。

一、关于 Transformer 每一层的并行

1. 同一层内,token 之间的依赖

计算这一层 C 的输出,需要的是 A、B 在“上一层的输出”,而不是 A、B 在“这一层的输出”。

假设现在计算第 2 层。

第 1 层已经完成,所以第 2 层拿到三个已知向量:

a₁:A 在第 1 层的输出
b₁:B 在第 1 层的输出
c₁:C 在第 1 层的输出

首先,第 2 层分别用它们计算自己的 Q、K、V。省略 Norm 等细节:

a₁ → q_A、k_A、v_A
b₁ → q_B、k_B、v_B
c₁ → q_C、k_C、v_C

这三行没有互相依赖。 例如 k_B = b₁ @ W_K,只需要已经知道的 b₁,不需要先计算 A 在第 2 层的输出。

接下来,计算 attention:

A 的 attention 输出 = 用 q_A 查询 [k_A],加权 [v_A]

B 的 attention 输出 = 用 q_B 查询 [k_A, k_B],加权 [v_A, v_B]

C 的 attention 输出 = 用 q_C 查询 [k_A, k_B, k_C],加权 [v_A, v_B, v_C]

到这一步,所有需要的 Q/K/V 都已经算好了,所以三行也可以并行。

C 确实需要 A、B 的 K/V,但不需要等待 A、B 的 attention 输出。

一种错误的理解是:

算出本层 a₂ → 才能算本层 b₂ → 才能算本层 c₂

实际依赖是:

上一层输出                  本层输出

a₁ ───────────────────────→ a₂

a₁、b₁ ───────────────────→ b₂

a₁、b₁、c₁ ───────────────→ c₂

右侧的 a₂、b₂、c₂ 都依赖左侧已经存在的数据,右侧之间没有这条串行依赖链。

用单头 attention 的简化伪代码表达就是(省略 Norm、RoPE 等细节):

# 上一层已经完成,三行都已知
X = stack([a1, b1, c1])       # [3, D]

# 一起算出三个位置的 Q/K/V
Q = X @ Wq
K = X @ Wk
V = X @ Wv

# 一起计算三行 attention
scores = (Q @ K.T) / (Q.shape[-1] ** 0.5)
scores = causal_mask(scores)  # 禁止 A 看 B/C,禁止 B 看 C
O = softmax(scores, dim=-1) @ V

# O 的三行分别是 A、B、C 的 attention 输出

这里 causal mask 限制的是“能读哪些位置”,不是“哪个位置必须先执行完”。

同一层内部当然仍有步骤依赖:先产生 Q/K/V,再用它们计算 attention。但这不等于 A、B、C 三个位置必须依次完成这一层


2. 一层完整的计算过程

这一层的输入:       [a₁, b₁, c₁]
RMSNorm             三个 token 各自归一化,可以并行
QKV 投影            三个 token 一起做矩阵乘
RoPE                各位置的 Q/K 独立旋转,可以并行
Attention           各 query 位置可以并行,按因果规则读取 K/V
输出投影 Wo          三个 token 一起做矩阵乘
残差相加             三个 token 各自相加,可以并行
RMSNorm             三个 token 各自归一化,可以并行
FFN / SwiGLU        三个 token 各自加工,可以并行
残差相加             三个 token 各自相加,可以并行
这一层的输出:       [a₂, b₂, c₂]

3. 不同 head 之间的并行

没有相互数据依赖的计算具有并行机会;是否真正同时执行,还取决于 kernel 实现和硬件资源。

不同 head 的 attention 可以独立计算,最后将输出拼接,再经过输出投影。GQA 中多个 query head 即使共享 K/V,也可以并行读取它们。

二、关于瓶颈分析

1. 单 token decode:权重复用较少

先把 attention 暂时放在一边,只看 Transformer 里一个 Linear:

Y=XWY=XW

假设:

W: [4096, 4096],FP16

它包含约 1678 万个元素,占 32 MiB

单请求 decode 时:

X: [1, 4096]
W: [4096, 4096]
Y: [1, 4096]

这次矩阵乘法约需要:

F=2×4096×4096=33,554,432 FLOPF=2\times4096\times4096=33{,}554{,}432\ \mathrm{FLOP}

如果权重需要从显存读一遍,仅权重就是约 3355 万字节。

于是忽略其他流量,算术强度大约是:

I=计算量搬运字节数1 FLOP/ByteI=\frac{\text{计算量}}{\text{搬运字节数}}\approx1\ \mathrm{FLOP/Byte}

搬进来一个权重,给一个 token 用一下,这次算子就结束了。

2. 矩阵化 prefill:多个 token 共享权重

假设一次有 512 个 token:

X: [512, 4096]
W: [4096, 4096]
Y: [512, 4096]

计算量变为前者的 512 倍,但同一份权重可以服务 512 行输入。

这里将 512 个 token 的向量按行叠成 X,每行对应一个 token。同一层的同一个投影(例如 W_Q)对所有 token 使用相同参数,因此可以一起计算 Y = X @ W;这个 Linear 不混合不同 token 的行。

仅按权重流量估算:

I512 FLOP/ByteI\approx512\ \mathrm{FLOP/Byte}

实际还要算输入、输出、分块重复加载等,因此这只是展示复用机会的简化估计。

GPU 的矩阵乘 kernel 会把权重块和输入块加载到片上存储,进行多次乘加,再处理下一块。并行处理多个 token,让同一块权重有更多被重复使用的机会。

所以常见现象是:

小 batch decode:
搬很多权重,只服务少量 token
→ 容易受显存带宽限制

较长 prompt 的矩阵化 prefill:
同一份权重服务大量 token
→ 更容易充分使用计算单元

3. 耗时下界:计算与搬运取较大者

这里 t 表示所分析任务(例如一次矩阵乘法)的耗时。用一个粗略下界理解:

tmax(FPeff,MBeff)t\gtrsim\max\left(\frac{F}{P_{\mathrm{eff}}},\frac{M}{B_{\mathrm{eff}}}\right)
  • F:浮点运算总量,单位 FLOP。

  • P_eff:有效计算吞吐,单位 FLOP/s。

  • M:搬运数据量,单位 Byte。

  • B_eff:有效内存带宽,单位 Byte/s。

    两个分数分别表示计算所需时间和数据搬运所需时间。不同数据块的搬运与计算可以重叠,因此取较大值作为粗略下界,而不是直接相加。

    例如,计算需要 10 ms、搬运需要 20 ms,则耗时至少约为 20 ms,容易受带宽限制。即使计算吞吐翻倍,搬运时间不变,下界仍是 20 ms。

    实际还存在重叠不充分、kernel 启动、同步和调度等开销,因此真实耗时可能更长。有效吞吐和带宽不一定达到硬件标称峰值。