大模型训练与推理原理

两个核心问题:训练时内存装了什么?推理时为什么快?


一、训练阶段:内存里到底装了什么?

短答案

是的,训练时几乎所有东西都必须装进 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 优化器状态 记录了历史梯度信息,让参数更新更稳定:

这两个状态每个参数各占 4 bytes,合计 8 bytes/参数。


训练的三大省内存技术

1. 混合精度训练(Mixed Precision)

2. 梯度检查点(Gradient Checkpointing)

3. 模型并行 + ZeRO(零冗余优化)


训练时的计算流程(一个 step)

① 读取一批数据(mini-batch)
② 正向传播 → 计算预测输出 ŷ,保存各层激活值
③ 计算损失 L = Loss(ŷ, y)
④ 反向传播 → 从后往前逐层计算梯度 ∂L/∂w
⑤ 优化器更新 → w ← w - lr × Adam(梯度)
⑥ 清空梯度,准备下一 step

关键:每个 step 都要读写所有参数 + 梯度 + 优化器状态,所以训练极其耗显存,也是为什么训练比推理慢几十倍的根本原因。


二、推理阶段:为什么可以很快?

短答案

推理时确实需要加载全部参数到 GPU,但:

  1. 只需要正向传播,没有梯度和优化器状态
  2. KV Cache 让生成每个 token 不用重算历史
  3. 硬件针对矩阵乘法极度优化

推理时的内存

推理显存 ≈ 参数量 × 精度
  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 的 KeyValue 矩阵。

如果没有缓存:

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 在"预热"。


让推理更快的工程技术

1. 量化(Quantization)

2. 批处理(Batching)

3. 推测解码(Speculative Decoding)

4. Flash Attention

5. 专用硬件


三、训练 vs 推理 对比总结

对比维度 训练 推理
是否需要加载全部参数 ✅ 是 ✅ 是
梯度 ✅ 需要存储 ❌ 不需要
优化器状态 ✅ 需要存储 ❌ 不需要
激活值 ✅ 需要保留 ❌ 逐层丢弃
显存占用(175B参数) ~2.8 TB ~350 GB(FP16)
计算图方向 正向 + 反向 仅正向
关键优化 梯度检查点 / ZeRO KV Cache / 量化
典型吞吐 数千 tokens/s(多机) 数十~百 tokens/s(单机)

四、直觉类比


See Also


Sources: 综合 Transformer 原文 (Vaswani 2017)、DeepSpeed ZeRO 论文、FlashAttention 论文、vLLM 技术报告
Updated: 2026-08-26