大模型推理全流程:Prefill + Decode 完整链路
以输入 "今天天气真" → 输出 "好" 为起点,覆盖 Prefill(预填充)与 Decode(逐词生成)两个阶段,含 KV Cache 加速原理。
总览
┌──────────┐ ┌──────────┐ ┌──────────┐ ┌──────────┐ ┌──────────┐ ┌──────────┐
│ 分词 │ → │ 嵌入 │ → │ 位置编码 │ → │ Transformer│ → │ 输出投影 │ → │ 采样解码 │
│Tokenizer │ │Embedding │ │Pos.Encode│ │ 层×N │ │ LM Head │ │ Sampling │
└──────────┘ └──────────┘ └──────────┘ └──────────┘ └──────────┘ └──────────┘
离散文本 稠密向量 注入位置 全局交互 词表打分 选出Token
第一步:分词(Tokenization)------ 文本拆块
大模型不认识中文字符或英文字母,它只认识整数。
过程
用户输入 "今天天气真",分词器按 BPE(Byte Pair Encoding)等算法将其切分为词表中存在的子词,并映射为整数 ID:
css
"今天天气真" → ["今天", "天气", "真"] → [2521, 1083, 29576, 17582]
关键细节
| 概念 | 说明 |
|---|---|
| 词表(Vocabulary) | 模型的"字典",通常 50,000 ~ 250,000 个条目,超出词表的字符会回退到更细粒度的子词或字节 |
| BPE 算法 | 从字节级出发,逐步合并高频相邻对,兼顾常见词和罕见词 |
| 特殊 Token | <bos>(句首)、<eos>(句尾)、<pad>(填充)等控制 token |
| 中文分词特殊性 | 中文无天然空格分隔,BPE 会在字符级或子词级切分,常见的如"今天"整体为一个 token,罕见的组合会拆得更细 |
输入输出
makefile
输入: "今天天气真"(字符串)
输出: [2521, 1083, 29576, 17582](4 个整数 ID)
第二步:嵌入(Embedding)------ 数字转语义向量
整数 ID 本身不含任何语义信息。"2521"和"2522"之间并不比"2521"和"9999"之间"更相似"。嵌入层的任务就是赋予每个 ID 一个语义化表示。
过程
ini
Token ID: [2521, 1083, 29576, 17582]
↓ ↓ ↓ ↓ 查表
↓ ↓ ↓ ↓
向量: [v₁] [v₂] [v₃] [v₄] 每个 v 是 d 维向量(如 4096 维)
最终: 4 × d 的矩阵(4 个 token,每个 4096 维)
嵌入矩阵
嵌入矩阵是一个巨大的可训练参数表:
ini
Embedding Matrix
┌─────────────────────────────────┐
│ ID=0 → [0.12, -0.34, ...] │ d 维
│ ID=1 → [0.87, 0.21, ...] │ d 维
│ ID=2 → [-0.56, 0.93, ...] │ d 维
│ ... │
│ ID=50000 → [0.03, -0.18, ...] │ d 维
└─────────────────────────────────┘
词表大小 V × 隐藏维度 d
(例如 50,000 × 4096 ≈ 2 亿参数)
查表操作本质是:取出 ID 对应的那一行向量。
语义空间的含义
在这个高维向量空间中:
- "猫"和"狗"的距离 < "猫"和"汽车"的距离
- 语义关系呈现为向量方向上的规律性(如
king - man + woman ≈ queen)
- 这些性质不是手动设计,而是通过海量语料训练自动涌现的
输入输出
ini
输入: [2521, 1083, 29576, 17582] (4 个整数 ID)
输出: [[0.12, -0.34, ...], (4 × 4096 矩阵)
[0.87, 0.21, ...],
[-0.56, 0.93, ...],
[0.03, -0.18, ...]]
第三步:位置编码(Positional Encoding)------ 注入顺序信息
Transformer 的 Attention 是并行计算的,所有 token 同时处理。如果不加位置信息,"今天天气真"和"天今真气天"在模型眼中将一模一样。
RoPE(旋转位置编码,主流方案)
LLaMA、Qwen 等现代模型使用 RoPE(Rotary Position Embedding):
- 不额外增加可学习参数
- 通过旋转 Q 和 K 向量来注入相对位置信息
- 核心性质:两个 token 之间的注意力分数取决于它们的相对距离,而非绝对位置
scss
RoPE(Q, pos=i) · RoPE(K, pos=j) → 只依赖于 (i - j)
这意味着即使训练时最大长度是 2048,推理时也可以外推到更长的上下文(配合 NTK 感知缩放等技巧)。
其他位置编码方案
| 方案 | 原理 | 代表模型 |
|---|---|---|
| 绝对位置编码(Sinusoidal) | 用正弦/余弦函数生成固定位置向量,与 Embedding 相加 | Transformer 原论文 |
| 可学习位置编码 | 位置向量作为可训练参数 | BERT、GPT-1 |
| RoPE | 通过旋转矩阵修改 Q、K,编码相对位置 | LLaMA、Qwen、ChatGLM |
| ALiBi | 在注意力分数上直接加一个与距离成正比的偏置 | BLOOM |
输入输出
less
输入: 4 × 4096 的 Embedding 矩阵
过程: 对 Q 和 K 向量按位置 i 施加旋转(角度 = i × θ_base^(d/2))
输出: 4 × 4096 的带位置信息的矩阵(形状不变,语义变了)
第四步:核心计算(Transformer 层)------ 全序列交互与深度加工
这是模型的主体,数据穿过 N 层相同的 Transformer Block(LLaMA-70B 有 80 层,小模型约 12~32 层)。每一层由两个子层组成。
4.1 单层内部结构(Pre-LN,现代 LLM 实际方案)
scss
输入 X (L × d)
│
├──────────────────────────────────────────────┐
│ │
▼ │
LayerNorm(X) → X_norm │
│ │
▼ │
┌─────────────────────────┐ │
│ Attention 子层 │ │
│ │ │
│ X_norm → Q, K, V │ │
│ Q × K^T → 注意力矩阵 │ ← 全序列交互 │
│ softmax(QK^T/√d) × V │ │
│ → Out_attn (L × d) │ │
└─────────────────────────┘ │
│ │
└────── X + Out_attn ────→ X' ─────────────────┤
│
▼ │
LayerNorm(X') → X'_norm │
│ │
▼ │
┌─────────────────────────┐ │
│ FFN 子层 │ │
│ │ │
│ W₁: d → 4d │ ← 逐 Token 独立处理 │
│ GELU / SwiGLU │ │
│ W₂: 4d → d │ │
│ → Out_ffn (L × d) │ │
└─────────────────────────┘ │
│ │
└────── X' + Out_ffn ───→ X_next ──────────────→ 下一层
4.2 自注意力机制(Self-Attention)------ 全序列交互
这是唯一发生 token 间信息交换的地方。
scss
对每个 token i:
1. 线性投影: x_i → W_Q·x_i (Query), W_K·x_i (Key), W_V·x_i (Value)
2. 计算相关性: score(i, j) = Q_i · K_j / √d_k (d_k 是每个头的维度)
3. 归一化: α(i, *) = softmax(score(i, *)) (对 j 维度归一化)
4. 加权聚合: Out_i = Σ_j α(i, j) · V_j (按注意力权重加权所有位置的 V)
最终输出: Out_attn ∈ L × d
多头的意义:Q、K、V 被拆成 h 个头各自独立计算,最后拼接。不同头关注不同模式(语法结构、长距离依赖、指代关系等)。
| 步骤 | 操作 | token 间交互? |
|---|---|---|
| 投影 Q,K,V | 线性变换 | 否(逐 token) |
| Q×K^T | 计算注意力分数 | 是(全序列交叉) |
| Softmax | 归一化为概率 | 是 |
| 加权 V | 按权重聚合 | 是 |
| 输出投影 | 拼接多头 + 线性变换 | 否(逐 token) |
关键理解:如果去掉 Q×K^T 这一步,Transformer 就退化成了对每个 token 独立操作的模型。这一步是模型"理解上下文"的数学本质。
4.3 FFN(前馈神经网络)------ 逐 Token 深度加工
scss
FFN(x) = W₂ · σ(W₁ · x + b₁) + b₂
其中:
W₁: d → 4d(升维,如 4096 → 16384)
σ: 激活函数(GELU 或 SwiGLU)
W₂: 4d → d(降维,如 16384 → 4096)
| 特性 | Attention | FFN |
|---|---|---|
| Token 间交互 | ✅ L × L 矩阵交叉 | ❌ 逐 token 独立 |
| 计算复杂度 | O(L² · d) | O(L · d²) |
| 功能类比 | 全体开会讨论 | 各自回工位思考 |
| 存储知识 | 关联与上下文 | 事实性知识(经验上多在 FFN 权重中) |
SwiGLU 变体(LLaMA 采用):
scss
FFN_SwiGLU(x) = (W₃ · x) ⊙ (SiLU(W₁ · x)) · W₂
其中 ⊙ 是逐元素乘法,SiLU = x · σ(x)
参数更多(多一个 W₃),但效果更好
4.4 逐层传递
css
第 1 层: X₀ (Embedding + 位置编码) → Attn → FFN → X₁
第 2 层: X₁ → Attn → FFN → X₂
...
第 N 层: X_{N-1} → Attn → FFN → X_N
每一层的输入输出形状都是 L × d。
残差连接确保梯度可以直通,深层梯度不消失。
4.5 Pre-LN vs Post-LN
| 维度 | Post-LN(原论文) | Pre-LN(现代 LLM) |
|---|---|---|
| LN 位置 | 残差之后 | 子层之前 |
| 残差路径 | 经过 LN | 直通,LN 在旁路 |
| 训练稳定性 | 深层难收敛 | 稳定,可堆上百层 |
| 代表模型 | Transformer 2017 | GPT-2/3、LLaMA、BERT |
输入输出
makefile
输入: L × d 矩阵(带位置编码的 Embedding)
过程: 穿过 N 层 Transformer Block,各层 Attention → FFN → 残差
输出: L × d 矩阵(每个 token 的最终隐藏状态都已融合了整段 Prompt 的全局信息)
第五步:输出投影(LM Head)------ 词表打分
所有 Token 穿过最后一层 Transformer 后,模型取最后一个位置的隐藏状态(因为它已经通过 Attention 融合了整个 Prompt 的上下文)。
过程
ini
X_N (L × d)
│
▼
取最后一个 Token: h_last ∈ 1 × d (例如 [0.23, -0.87, 0.45, ...])
│
▼
LM Head (线性层): h_last · W_lm (W_lm: d × V,通常与 Embedding 矩阵共享权重)
│
▼
Logits: ∈ 1 × V (例如 V=50000,每个候选词一个原始分数)
Logits 的特点
- 原始分数,未归一化
- 可能有正有负,数值范围较大(如 -50 ~ +100)
- 值越大的候选词,模型认为越"合适"
json
Logits 示例(简化,实际 V=50000+):
"好" : 8.2 "坏" : -2.1 "的" : 6.8
"啊" : 5.3 "吧" : 3.9 "呀" : 4.1
"天" : -1.2 "不" : 2.3 ...其余 49992 个分数
权重共享
许多模型(如 LLaMA)将 Embedding 矩阵和 LM Head 的权重共享(tied weights),减少参数量且隐含编码了输入输出的对称性。
输入输出
makefile
输入: 最后一个 Token 的隐藏状态 h_last (1 × d)
过程: h_last × W_lm (线性投影,d → V)
输出: Logits (1 × V),每个候选词一个原始分数
第六步:采样与解码(Sampling & Decode)------ 输出首个 Token
Logits 不能直接用于选择 token,需要经过概率化和采样。
6.1 Softmax → 概率分布
scss
P(token_i) = exp(logit_i) / Σ_j exp(logit_j)
将任意实数范围映射到 (0, 1) 且总和为 1。
makefile
Logits: 好=8.2 的=6.8 啊=5.3 吧=3.9 坏=-2.1 天=-1.2 ...
↓ ↓ ↓ ↓ ↓ ↓
Softmax: 0.52 0.15 0.06 0.02 0.0001 0.0002 ...
6.2 Temperature ------ 控制"创造性"
在 Softmax 之前,所有 Logits 除以 Temperature T:
scss
P(token_i) = exp(logit_i / T) / Σ_j exp(logit_j / T)
| T 值 | 效果 | 适用场景 |
|---|---|---|
| T → 0 | 几乎总是选最高分 token(确定性) | 代码生成、翻译 |
| T = 1 | 原始概率分布 | 通用对话 |
| T > 1 | 分布更平坦,更"随机" | 创意写作、头脑风暴 |
6.3 Top-k 采样
只保留概率最高的 k 个候选词,其余概率置零,重新归一化后再采样。
ini
Top-k=50: 只从"最可能"的 50 个词中选
6.4 Top-p(核采样)
按概率从高到低累加,当累积概率 ≥ p 时截断。比 Top-k 更自适应。
css
Top-p=0.9: 保留概率最高的词,直到累计概率达到 90%
例如:好(0.52) + 的(0.15) + 啊(0.06) + 吧(0.02) + ... 直到累计 0.9
6.5 组合采样(常见实践)
css
Temperature + Top-p + Top-k 三者组合使用:
1. Logits / T(调整"温度")
2. Top-k 预截断(去除极端低分噪声)
3. Top-p 终截断(自适应宽度)
4. 从剩余分布中随机采样
6.6 反向解码 ------ 输出文字
yaml
采样的 Token ID: 2648
↓
词表反查: 2648 → "好"
↓
输出: "好"
第七步:Decode 阶段 ------ 逐词生成与 KV Cache
首个 Token 输出后,模型进入 Decode(解码)阶段 。与前六个步骤不同,Decode 不是一次性计算,而是一个循环 :每生成一个新 Token,就把它拼到序列末尾,再重复整个过程,直到输出终止符 <eos>。
makefile
循环: 今天天气真好 → 今天天气真好, → 今天天气真好,适合 → ... → 今天天气真好,适合出去走走。
生成"好"后 生成","后 生成"适合"后 生成"走走"后
如果每次生成新 Token 都把整个序列重新算一遍(无 KV Cache),那么生成第 N 个 Token 时,Attention 的计算量是 O(N²)。生成一段 1000 token 的文本,总计算量将达到 O(N³) 级别------这在工程上完全不可接受。
KV Cache 正是解决这个问题的关键技术。
7.1 KV Cache 的核心思想:存起来,不必重算
在 Prefill 阶段,所有 Prompt Token 的 Key 和 Value 都已计算完毕。这些 K、V 向量不会改变------后续 token 的加入不会影响前面 token 的 K、V。
KV Cache 就是把这些 K、V 保存在显存里,后续 Decode 时直接读取,不再重新计算。
ini
Prefill 阶段结束时:
KV Cache = [K₀, K₁, K₂, K₃, ... K_{L-1}] (L 个 K,每个 d 维)
[V₀, V₁, V₂, V₃, ... V_{L-1}] (L 个 V,每个 d 维)
Decode 第 1 步(生成第 L+1 个 token):
只计算新 token 的 Q_new, K_new, V_new
Q_new 去 × 所有缓存的 K → Attention Weights
Attention Weights × 所有缓存的 V → 输出
把 K_new, V_new 追加到 KV Cache 中
7.2 一次逐步拆解:无 Cache vs 有 Cache
假设已生成 5 个 token,现在要生成第 6 个。
无 KV Cache(每步重算全部):
css
Step 6
┌────────────────────────────────────┐
│ 把前 5 个 token + 新 token │
│ 全部重新送入 Transformer │
│ │
│ Q₁ Q₂ Q₃ Q₄ Q₅ Q₆ (6 × d) │
│ K₁ K₂ K₃ K₄ K₅ K₆ (6 × d) ← 全重算 │
│ V₁ V₂ V₃ V₄ V₅ V₆ (6 × d) ← 全重算 │
│ │
│ Attention: 6 × 6 矩阵交叉 │
└────────────────────────────────────┘
计算量: O(N²) per step,N 步累计 O(N³)
有 KV Cache(增量计算):
scss
Step 6
┌────────────────────────────────────┐
│ 只计算新 token 的 Q,K,V │
│ │
│ Q_new (1 × d) ← 只算一个 │
│ K_new (1 × d) ← 只算一个 │
│ V_new (1 × d) ← 只算一个 │
│ │
│ 从 Cache 读取: │
│ K_cache (5 × d) ← 直接读 │
│ V_cache (5 × d) ← 直接读 │
│ │
│ Attention: Q_new × K_cache^T │
│ (1 × d)·(d × 5) = 1 × 5 │
└────────────────────────────────────┘
计算量: O(N) per step,N 步累计 O(N²)
7.3 复杂度变化的数学本质
| 操作 | 无 KV Cache | 有 KV Cache | 降幅 |
|---|---|---|---|
| 计算 Q、K、V | O(N) 每步算 N 个 | O(1) 每步只算 1 个 | N 倍 |
| Q × K^T 矩阵乘法 | O(N²) | O(N)(1×d × d×N) | N 倍 |
| 加权 V | O(N²) | O(N) | N 倍 |
| 每步总量 | O(N²) | O(N) | N 倍 |
| 生成 N 个 token 总量 | O(N³) | O(N²) | N 倍 |
核心理解:KV Cache 的本质是用空间换时间------把前面所有 token 的 K、V 缓存在显存里,避免每一步都重算。
7.4 逐步操作流程
ini
┌─────────────────────────────────────────────────────────────┐
│ Decode 第 t 步(当前序列长度 = L + t) │
│ │
│ 1. Embedding + RoPE │
│ 只计算新 token 的嵌入向量 + 位置编码 │
│ 输出: 1 × d │
│ │
│ 2. 穿过 Transformer 层 × N │
│ ┌──────────────────────────────────────┐ │
│ │ 每层的 Attention: │ │
│ │ Q_new (1×d) ← 新 token 的 Query │ │
│ │ K_cache (L+t-1 × d) ← 显存直接读 │ ← 核心加速点 │
│ │ V_cache (L+t-1 × d) ← 显存直接读 │ │
│ │ │ │
│ │ Scores = Q_new · K_cache^T │ (1 × d)·(d × N) │
│ │ Out = Softmax(Scores) · V_cache │ 1 × d │
│ └──────────────────────────────────────┘ │
│ │
│ 3. LM Head │
│ 隐藏状态 → Logits (1 × V) │
│ │
│ 4. 采样 → Token ID │
│ │
│ 5. 更新 KV Cache │
│ K_new, V_new 追加到 Cache 末尾 │
│ Cache 长度: L+t-1 → L+t │
└─────────────────────────────────────────────────────────────┘
7.5 Prefill vs Decode 对比总结
| 维度 | Prefill | Decode |
|---|---|---|
| 处理方式 | 一次并行处理所有 Prompt Token | 逐 Token 串行生成 |
| 输入规模 | L 个 Token(Prompt 长度) | 1 个 Token(新生成的那个) |
| Attention 计算 | Q、K、V 全部计算(L×L 矩阵) | 只算新 QKV,K、V 从 Cache 读 |
| 每步复杂度 | O(L²) | O(L) |
| 计算瓶颈 | 计算密集(Compute-bound) | 内存带宽密集(Memory-bound) |
| KV Cache | Prefill 结束后初始化 | 每步追加 1 个 K、V 对 |
| 类比 | 全班同学一次性填报所有信息 | 老师每次翻看花名册问一个问题 |
7.6 KV Cache 的显存代价
KV Cache 虽快,但吃显存。对于一个 L 层、h 头、d 维的模型:
yaml
每个 token 的 KV Cache 大小 = 2 × L × d × 2 (K + V)
= 4 × L × d (FP16 下为 2 字节/元素)
例: LLaMA-7B (32 层, 4096 维)
每个 token ≈ 4 × 32 × 4096 × 2 bytes ≈ 1 MB
2048 token 上下文 ≈ 2 GB 的 KV Cache
这还没算多头拆分、batch size 等额外开销。当生成上万 token 的长文时,KV Cache 会占用大量显存,且由于长度动态增长,容易产生显存碎片。
业界使用 PagedAttention(分页注意力) 技术解决此问题,将 KV Cache 按固定大小的"页"来管理,类同操作系统的虚拟内存分页机制,有效减少碎片并提升显存利用率。(该技术是 vLLM 推理框架的核心创新。)
深度分析:Prefill 阶段的真正计算瓶颈
一个常见的直觉误区是认为"Attention 计算(Q×K^T)是 Transformer 最耗时的部分"。但在实际工程中,尤其是日常对话场景(序列长度较短),线性投影(Q、K、V 的计算)才是 Prefill 阶段真正的耗时大头。
8.1 计算量对比:投影 vs Attention
| 操作 | FLOPs | 当 L=1000, d=4096 时 | 占比 |
|---|---|---|---|
| QKV 线性投影 | 3 · L · d² | 3×1000×16,777,216 ≈ 50.3 GFLOPs | ~95% |
| Attention QK^T | L² · d | 1,000,000×4,096 ≈ 4.1 GFLOPs | ~5% |
| Attention × V | L² · d | 同上 ≈ 4.1 GFLOPs | --- |
| FFN | 8 · L · d² | ~134 GFLOPs | --- |
当 L < d 时(日常对话几乎总是如此),3 · d² 远大于 L² · d。GPU 绝大部分时间花在做矩阵乘法(投影)上,而不是做注意力交叉计算。
这是反直觉的------大家总把 Attention 挂在嘴边,但数学上投影才是计算量主体。d²(4096² = 1600 万)这一项是天文数字级别的。
8.2 为什么 Attention 看起来"不耗时"?------ FlashAttention 的降维打击
传统的 Attention 需要把 L × L 的注意力分数矩阵写入显存,这会导致严重的内存读写瓶颈(Memory-bound)。
FlashAttention 通过分块计算(Tiling)技术解决了这个问题:
yaml
传统 Attention:
Q×K^T → 写入 L×L 矩阵到显存 → Softmax → 读取 → ×V
↑ 巨大的显存读写开销
FlashAttention:
分块加载 Q, K, V → 块内计算 Softmax → 块内加权求和 → 输出
↑ 从不需要把完整的 L×L 矩阵写入显存
结果:Attention 被优化成了极其高效的计算密集型操作,反而让线性投影成了新的瓶颈------投影是 Memory-bound 操作,需要反复读巨大的 W_q, W_k, W_v 权重矩阵。
8.3 业界优化方案一:算子融合(Operator Fusion)
现代推理框架不会让 GPU 分别计算 X·W_q、X·W_k、X·W_v。
less
优化前(三次独立 Kernel 调用):
读取 X → 读取 W_q → 计算 → 写入 Q
读取 X → 读取 W_k → 计算 → 写入 K ← 每次都要重新读 X
读取 X → 读取 W_v → 计算 → 写入 V
优化后(一次 Kernel 调用):
W_q, W_k, W_v 拼接为一个大矩阵 W_qkv
读取 X → 读取 W_qkv → 一次计算 → Q, K, V 一起输出
↑ 减少了对 X 和权重的重复读取,显存读写次数降为 1/3
vLLM、TensorRT-LLM 等框架将此作为基础优化。
8.4 业界优化方案二:GQA / MQA 架构
从模型架构层面削减 K、V 投影的计算量。
| 注意力类型 | 头数关系 | K,V 投影计算量 | 代表模型 |
|---|---|---|---|
| MHA(多头注意力) | Q头数 = K头数 = V头数 | 基准 (100%) | 原版 Transformer、GPT-3 |
| GQA(分组查询注意力) | Q头数 > K头数 = V头数(分组共享) | 大幅减少 | LLaMA 2 70B、LLaMA 3 |
| MQA(多查询注意力) | 所有 Q 头共享一组 K,V | 最少 | PaLM、Falcon |
bash
MHA: 8 个 Q 头 → 对应 8 组 K,V → W_k 的宽度 = 8 × (d/h)
GQA: 8 个 Q 头 → 对应 2 组 K,V → W_k 的宽度 = 2 × (d/h) ← 压缩了 4 倍
MQA: 8 个 Q 头 → 对应 1 组 K,V → W_k 的宽度 = 1 × (d/h) ← 压缩了 8 倍
GQA 是当前主流选择:在推理速度和模型质量之间取得了最佳平衡。LLaMA 3 全系列均采用 GQA。
8.5 Prefill 阶段瓶颈全景
scss
┌──────────────────────────────────────────────────┐
│ Prefill 阶段计算时间分布(典型场景) │
│ │
│ ████████████████████████████░░░░ QKV 投影 (~40%)│
│ ████████████████████████░░░░░░░░ FFN 投影 (~35%)│
│ ██████░░░░░░░░░░░░░░░░░░░░░░░░░░ Attention (~10%)│
│ ████░░░░░░░░░░░░░░░░░░░░░░░░░░░░ LM Head (~10%)│
│ ██░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░ 其他 (~5%) │
│ │
│ 核心瓶颈:大量 d×d 矩阵乘法(Memory-bound) │
└──────────────────────────────────────────────────┘
总结:Prefill 阶段表面是在做上下文关联(Attention),实际上 GPU 大部分算力消耗在把原始文本"翻译"成特征空间(线性投影)的过程中。 FlashAttention 已经把 Attention 优化得足够高效,瓶颈转移到了投影计算。业界通过算子融合和 GQA/MQA 架构双管齐下来应对。
完整流程总结
Prefill 阶段(一次并行处理 Prompt)
lua
"今天天气真"
│
▼
Tokenizer ──────→ [2521, 1083, 29576, 17582] (4 个整数 ID)
│
▼
Embedding ──────→ [[v₁], [v₂], [v₃], [v₄]] (4 × 4096 矩阵)
│
▼
RoPE ──────────→ Q,K 按位置旋转 (注入位置信息)
│
▼
Layer 1 ~ N ───→ Attention(全序列交互)+ FFN(逐词精炼) (L × d → L × d)
│ └→ 所有 K,V 存入 KV Cache (初始化 Cache)
▼
LM Head ───────→ h_last → 词表投影 (1 × 50000 Logits)
│
▼
Softmax + 采样 ─→ 选出最高概率 Token
│
▼
Detokenizer ───→ "好" (整数 → 文字)
Decode 阶段(逐 Token 循环生成)
markdown
本轮新 Token ──→ Embedding + RoPE ──→ 穿过 N 层 Transformers ──→ LM Head ──→ 采样输出
↑ │
│ 每层 Attention: │
│ Q_new · K_cache^T · V_cache │
│ (从显存读取,不重算) │
│ │
└─────────── 追加 K_new, V_new 到 KV Cache ─────────────────────────────────┘
循环直到输出 <eos> 或达到最大长度
完整阶段对比
| 阶段 | 输入 | 输出 | 计算特点 |
|---|---|---|---|
| 分词 | 字符串 | Token ID 序列 | 基于词表的规则匹配 |
| 嵌入 | Token ID 序列 | 稠密向量矩阵 L×d | 查表操作 |
| 位置编码 | 向量矩阵 | 带位置信息的矩阵 | 对 Q,K 施加旋转或加偏置 |
| Transformer 层 (Prefill) | 矩阵 L×d | 矩阵 L×d + KV Cache | Attention 全序列交互,FFN 逐 token |
| Transformer 层 (Decode) | 单 token + KV Cache | 单 token + 更新 Cache | Attention 仅 Q_new × K_cache,O(N) |
| 输出投影 | 最后一个 token 向量 | Logits (1×V) | 线性投影 |
| 采样解码 | Logits + 策略 | 文字 Token | 概率化 + 采样 + 反查 |
附录:关键维度参考
| 模型 | 层数 N | 隐藏维度 d | FFN 中间维度 | 词表大小 V | 头数 |
|---|---|---|---|---|---|
| LLaMA-7B | 32 | 4096 | 11008 | 32000 | 32 |
| LLaMA-13B | 40 | 5120 | 13824 | 32000 | 40 |
| LLaMA-70B | 80 | 8192 | 28672 | 32000 | 64 |
| Qwen-7B | 32 | 4096 | 11008 | 151936 | 32 |
| GPT-3 175B | 96 | 12288 | 49152 | 50257 | 96 |