大模型训练与推理原理
两个核心问题:训练时内存装了什么?推理时为什么快?
一、训练阶段:内存里到底装了什么?
短答案
是的,训练时几乎所有东西都必须装进 GPU 显存(VRAM),而且远不止"参数本身"——实际占用是参数大小的 4~20 倍。
训练内存的四大组成
| 组成部分 | 占用估算(FP32) | 说明 |
|---|---|---|
| 模型参数 | 4 bytes/参数 | 神经网络的权重矩阵 W |
| 梯度(Gradients) | 4 bytes/参数 | 反向传播时,每个参数对应一个梯度 |
| 优化器状态(Optimizer States) | 8 bytes/参数 | Adam 需存 m(一阶矩)和 v(二阶矩) |
| 中间激活值(Activations) | 随 batch size 线性增长 | 正向传播中每一层的输出,反向传播时要用 |
以 Adam 优化器、FP32 精度为例:
总显存 ≈ 参数量 × (4 + 4 + 8) bytes = 参数量 × 16 bytes
例:GPT-3 (175B 参数)
= 175×10⁹ × 16 bytes ≈ 2.8 TB
→ 需要 350 张 A100 (80GB) !
为什么梯度和优化器状态是必须的?
梯度 是反向传播(Backpropagation)的核心产物:
损失函数 L → 对每个参数 w 求偏导 → ∂L/∂w(梯度)→ 更新 w
每个参数都有对应的梯度,所以梯度和参数一样大。
Adam 优化器状态 记录了历史梯度信息,让参数更新更稳定:
m:梯度的指数移动平均("方向")v:梯度平方的指数移动平均("步长自适应")
这两个状态每个参数各占 4 bytes,合计 8 bytes/参数。
训练的三大省内存技术
1. 混合精度训练(Mixed Precision)
- 正向/反向传播用 FP16(2 bytes/参数)
- 优化器保留 FP32 主副本(防止数值精度丢失)
- 实际节省约 1.5×,但完全消除不了
2. 梯度检查点(Gradient Checkpointing)
- 不保存所有中间激活值
- 反向传播时重新计算需要的激活值
- 显存降低 √n 倍,但计算量增加 ~33%
3. 模型并行 + ZeRO(零冗余优化)
- 把参数、梯度、优化器状态切分到多张 GPU 上
- ZeRO-3:每张 GPU 只存全部参数的 1/N 份
- 允许在数百张 GPU 上训练万亿参数模型
训练时的计算流程(一个 step)
① 读取一批数据(mini-batch)
② 正向传播 → 计算预测输出 ŷ,保存各层激活值
③ 计算损失 L = Loss(ŷ, y)
④ 反向传播 → 从后往前逐层计算梯度 ∂L/∂w
⑤ 优化器更新 → w ← w - lr × Adam(梯度)
⑥ 清空梯度,准备下一 step
关键:每个 step 都要读写所有参数 + 梯度 + 优化器状态,所以训练极其耗显存,也是为什么训练比推理慢几十倍的根本原因。
二、推理阶段:为什么可以很快?
短答案
推理时确实需要加载全部参数到 GPU,但:
- 只需要正向传播,没有梯度和优化器状态
- KV Cache 让生成每个 token 不用重算历史
- 硬件针对矩阵乘法极度优化
推理时的内存
推理显存 ≈ 参数量 × 精度
FP16: 2 bytes/参数
INT8: 1 byte/参数(量化后)
INT4: 0.5 bytes/参数
例:Llama-3 70B,FP16
= 70×10⁹ × 2 bytes ≈ 140 GB
→ 约需 2 张 H100 (80GB)
与训练相比,没有梯度(0),没有优化器状态(0),没有存储激活值的必要,显存需求降低 8~20 倍。
推理时是怎么"逐字生成"的?
以 Transformer 架构为例,输入 "What is 1+1?" 生成回答:
Step 1: 把输入 token 转成向量(Embedding)
Step 2: 经过 N 层 Transformer Block
每层:
① Self-Attention:每个 token 和其他 token 算相关性
② FFN(前馈网络):做非线性变换
Step 3: 输出层 → Softmax → 概率分布 → 采样/贪心 → 下一个 token
Step 4: 把新 token 拼回输入,重复 Step 1~3
关键问题:每次生成新 token,都要重新计算所有历史 token 吗?
KV Cache:推理速度快的核心秘密
问题所在
Self-Attention 中,计算新 token 的注意力需要所有历史 token 的 Key 和 Value 矩阵。
如果没有缓存:
- 生成第 100 个 token → 要重算前 99 个 token 的 K、V
- 生成第 1000 个 token → 要重算前 999 个 token 的 K、V
- 时间复杂度 O(n²),越长越慢
KV Cache 的解法
把每个 token 算出的 K、V 矩阵缓存在内存里:
生成第 100 个 token:
- 只需计算第 100 个 token 自己的 Q、K、V
- 历史 K、V 直接从缓存读取
- 时间复杂度从 O(n²) → O(n),每步 O(1) 增量
代价:KV Cache 占显存(每个 token、每一层、每个注意力头都要存 K 和 V)。长上下文(如 128K tokens)时 KV Cache 可能占用几十 GB。
为什么第一个 token 比后续 token 慢?
这是实际使用时经常观察到的现象:
| 阶段 | 名称 | 特征 |
|---|---|---|
| 处理输入 prompt | Prefill 阶段 | 并行处理所有输入 token,计算并填充 KV Cache |
| 逐步生成输出 | Decode 阶段 | 每次只生成 1 个 token,速度均匀 |
- Prefill:计算量大(输入越长越慢),但只做一次
- Decode:计算量小(增量),但 GPU 利用率低(每次只算 1 个 token,算力浪费)
所以你感觉"第一个字出来很慢,之后很快"——那是 Prefill 在"预热"。
让推理更快的工程技术
1. 量化(Quantization)
- 把 FP16 参数压缩成 INT8 / INT4
- 显存减少 2~4 倍,速度提升 1.5~3 倍
- 精度略有损失,但高质量量化(GPTQ / AWQ)几乎无感
2. 批处理(Batching)
- 把多个用户请求合并成一个 batch 同时推理
- GPU 擅长大矩阵计算,batch 越大 GPU 利用率越高
- Continuous Batching:动态合并不同长度序列,让 GPU 永不空转
3. 推测解码(Speculative Decoding)
- 用小模型先"草稿"生成几个 token
- 大模型一次并行验证多个 token
- 等效加速 2~3 倍,输出质量不变
4. Flash Attention
- 重写 Attention 计算的内存访问顺序
- 减少 GPU HBM 读写次数,速度提升 2~4 倍
- 现已是所有主流推理框架的标配
5. 专用硬件
- A100/H100:专门优化矩阵乘法(Tensor Core)
- GPU HBM 带宽:A100 = 2 TB/s,H100 = 3.35 TB/s
- 推理的瓶颈往往是"把参数从显存搬到计算核心"(memory-bound),HBM 带宽直接决定速度上限
三、训练 vs 推理 对比总结
| 对比维度 | 训练 | 推理 |
|---|---|---|
| 是否需要加载全部参数 | ✅ 是 | ✅ 是 |
| 梯度 | ✅ 需要存储 | ❌ 不需要 |
| 优化器状态 | ✅ 需要存储 | ❌ 不需要 |
| 激活值 | ✅ 需要保留 | ❌ 逐层丢弃 |
| 显存占用(175B参数) | ~2.8 TB | ~350 GB(FP16) |
| 计算图方向 | 正向 + 反向 | 仅正向 |
| 关键优化 | 梯度检查点 / ZeRO | KV Cache / 量化 |
| 典型吞吐 | 数千 tokens/s(多机) | 数十~百 tokens/s(单机) |
四、直觉类比
- 参数 = 一本厚厚的百科全书(推理时要把整本书搬上工作台)
- KV Cache = 读书时划的重点笔记(不用每次重读,翻笔记即可)
- 量化 = 把百科全书印成缩印版(内容差不多,但轻很多)
- 批处理 = 一次给 10 个学生讲课,比一对一讲 10 次省时
- Prefill = 老师先通读学生的问题(慢)
- Decode = 逐字回答(快而均匀)
See Also
- wiki/ai/The-Future-is-for-Everyone-Zuckerberg.md — Zuckerberg 关于 AI 未来的演讲
Sources: 综合 Transformer 原文 (Vaswani 2017)、DeepSpeed ZeRO 论文、FlashAttention 论文、vLLM 技术报告
Updated: 2026-08-26