训练基础设施 - 从单卡到千卡集群 / Training Infrastructure from One GPU to Thousand-GPU Clusters
📅 创建时间:2026-05-31 🏷️ 标签:#分布式训练 #DataParallel #TensorParallel #PipelineParallel #ZeRO #DeepSpeed 📚 前置知识:[[11 - 训练基础扫盲]] [[12 - Post-Training Pipeline]] 🎯 文档定位:科普深入 × 专业浅显 — 讲清楚分布式训练的"为什么"和"怎么组合",不陷入 CUDA 编程细节
📋 本章目标
阅读完本文档后,你将能够:
- [ ] 理解 ML Infra 的三分框架(Pre-Training / SFT / RL Infra) ← 新增
- [ ] 解释为什么单卡训练不了大模型(显存/计算/通信瓶颈)
- [ ] 描述三种分布式训练策略的核心思想(DP / TP / PP)
- [ ] 理解 ZeRO 显存优化技术的三个阶段
- [ ] 掌握业界常见的并行策略组合(TP+PP+DP+ZeRO)
- [ ] 了解训练 Infra 的核心组件(GPU 集群拓扑、通信库、框架)
- [ ] 解释常见训练术语(BF16/FP8、NCCL、Checkpoint 等)
- [ ] 理解完整的 Post-Training Pipeline 的资源消耗
第0部分:ML Infra 的三分天下
在深入技术细节之前,先建立一个全局视野。ML Infra 可以分为三大块,每一块的挑战和技术栈都不同:
┌─────────────────────────────────────────────────────────────────┐
│ ML Infra 三分天下 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ 1. Pre-Training Infra(预训练 Infra) │
│ ───────────────────────────────────────────────────────────── │
│ 职责:支撑万亿 token 的预训练 │
│ 核心挑战: │
│ ├─ 千卡并行训练的稳定性(数千卡跑几个月不能挂) │
│ ├─ 极致的数据吞吐(PB 级数据怎么喂进去) │
│ └─ 超大规模 Checkpoint 管理 │
│ 代表工作:Meta 的 Megatron-LM、Google 的 TPU 基础设施 │
│ │
│ 2. Post-Training SFT Infra(SFT Infra) │
│ ───────────────────────────────────────────────────────────── │
│ 职责:支撑 SFT 阶段的监督微调 │
│ 核心挑战: │
│ ├─ 高质量标注数据的 pipeline(采集、清洗、存储) │
│ ├─ 多任务数据的配比(代码/对话/知识各占多少) │
│ └─ SFT 的分布式训练(比预训练简单,但仍需多卡) │
│ 资源量级:比预训练少 100-1000 倍 │
│ │
│ 3. Post-Training RL Infra(RL Infra) ← 翁家翌做的事情 │
│ ───────────────────────────────────────────────────────────── │
│ 职责:支撑 RLHF/DPO 的强化学习微调 │
│ 核心挑战: │
│ ├─ 多模型并发:LLM + Reward Model + Reference Model │
│ ├─ 实时数据流:生成 → 打分 → 更新 的循环管道 │
│ ├─ 训练稳定性:PPO 极易崩溃,监控和容错是关键 │
│ └─ 偏好数据管理:人类/AI 标注的收集和版本管理 │
│ 最复杂:RL Infra 需要同时运行多个模型,通信和数据同步最麻烦 │
│ │
└─────────────────────────────────────────────────────────────────┘回到翁家翌的例子:他说自己在 OpenAI 搭建了 Post-Training RL Infra,指的就是第三块——让 RLHF 从"能 work 的代码"变成"能 scale 到 175B 的生产系统"。这不是预训练,也不是 SFT,而是最复杂的 RL 训练工程。
Pre-Training vs RL Infra 的区别:
- 预训练 Infra:1 个模型,1 个数据流(数据 → 模型 → 更新)
- RL Infra:3 个模型,2 个数据流(生成 + 打分 + 更新,且互相依赖)
第1部分:为什么单卡训练不了大模型
1.1 三个瓶颈
训练大模型有三大瓶颈,显存瓶颈是最先卡死的:
┌─────────────────────────────────────────────────────────────────┐
│ 训练大模型的三大瓶颈 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ 瓶颈1:显存瓶颈(最先卡死) │
│ ───────────────────────────────────────────────────────────── │
│ │
│ 训练需要存储: │
│ ├─ 模型参数(FP16 下每个参数 2 字节) │
│ ├─ 梯度(和参数同量级) │
│ ├─ 优化器状态(Adam 需存动量等,12 字节/参数!) │
│ ├─ 激活值(Forward 时保存,反向时用) │
│ └─ 临时缓冲(算子融合等) │
│ │
│ 7B 模型在 FP16 下的显存占用: │
│ ├─ 参数:7B × 2B = 14 GB │
│ ├─ 梯度:7B × 2B = 14 GB │
│ ├─ 优化器状态:7B × 12B = 84 GB │
│ └─ 总计:> 100 GB │
│ │
│ 单卡 A100 80GB → 不够! │
│ │
│ 瓶颈2:计算瓶颈(太慢) │
│ ───────────────────────────────────────────────────────────── │
│ GPT-3(175B)预训练需要约 3640 PF-days │
│ 单卡 A100 算力 ~ 312 TFLOPS │
│ → 需要约 3000+ 年! │
│ │
│ 瓶颈3:通信瓶颈(数据传输太慢) │
│ ───────────────────────────────────────────────────────────── │
│ 分布式训练中 GPU 之间需要传输梯度 │
│ PCIe 带宽 32 GB/s vs NVLink 900 GB/s │
│ → 用 PCIe 通信会严重拖慢训练 │
│ │
└─────────────────────────────────────────────────────────────────┘1.2 7B/70B 模型的显存账
┌─────────────────────────────────────────────────────────────────┐
│ 7B / 70B 模型在 FP16 下的显存占用 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ 以 7B 参数模型为例(FP16 = 2 字节/参数): │
│ │
│ ├─ 模型参数:7B × 2 = 14 GB │
│ ├─ 梯度: 7B × 2 = 14 GB (必须存,用于反向传播) │
│ ├─ 优化器状态:7B × 12 = 84 GB (Adam 需要 m, v 两个状态) │
│ ├─ 激活值: ~7B × 2 = ~14 GB(取决于序列长度和 batch size) │
│ └─ 其他开销: ~5 GB │
│ │
│ 总计:~131 GB │
│ 单卡 A100 80GB → 不够! │
│ │
│ ───────────────────────────────────────────────────────────── │
│ │
│ 以 70B 参数模型为例(FP16 = 2 字节/参数): │
│ │
│ ├─ 模型参数:70B × 2 = 140 GB │
│ ├─ 梯度: 70B × 2 = 140 GB │
│ ├─ 优化器状态:70B × 12 = 840 GB │
│ ├─ 激活值: ~70B × 2 = ~140 GB │
│ └─ 其他开销: ~50 GB │
│ │
│ 总计:> 1.3 TB │
│ 需要约 16+ 张 A100 80GB 才能放下! │
│ │
└─────────────────────────────────────────────────────────────────┘第2部分:分布式训练三大策略
2.1 Data Parallelism(数据并行)
核心思想:每张卡有完整的模型,处理不同的数据 batch,最后同步梯度。
┌─────────────────────────────────────────────────────────────────┐
│ Data Parallelism(数据并行) │
├─────────────────────────────────────────────────────────────────┤
│ │
│ 场景:4 张 GPU,模型 70B(单卡放不下,但可以放) │
│ │
│ 数据集(1000 条) │
│ ↓ │
│ ┌──────────┬──────────┬──────────┬──────────┐ │
│ ↓ ↓ ↓ ↓ │ │
│ GPU 0 GPU 1 GPU 2 GPU 3 │ │
│ ┌────────┐ ┌────────┐ ┌────────┐ ┌────────┐ │ │
│ │ 完整 │ │ 完整 │ │ 完整 │ │ 完整 │ │ │
│ │ 模型 │ │ 模型 │ │ 模型 │ │ 模型 │ │ │
│ │(70B) │ │(70B) │ │(70B) │ │(70B) │ │ │
│ └───┬────┘ └───┬────┘ └───┬────┘ └───┬────┘ │ │
│ 处理250条 处理250条 处理250条 处理250条 │ │
│ ↓ ↓ ↓ ↓ │ │
│ 梯度A 梯度B 梯度C 梯度D │ │
│ ↓ ↓ ↓ ↓ │ │
│ └──────────┴──────────┴──────────┘ │ │
│ ↓ │ │
│ 梯度 AllReduce │ │
│ (求平均) │ │
│ ↓ │ │
│ 每张卡用平均梯度更新参数 │ │
│ │
│ 优点: │
│ ├─ 实现简单,几乎所有框架都支持 │
│ ├─ 加速比高(理想情况下 N 张卡加速 N 倍) │
│ └─ 通信量小(只传梯度,不传模型参数) │
│ │
│ 缺点: │
│ ├─ 每张卡都要存完整的模型 + 优化器状态(显存瓶颈依然存在) │
│ └─ 大到单卡放不下的模型,DP 本身不够 │
│ │
└─────────────────────────────────────────────────────────────────┘2.2 Tensor Parallelism(张量并行)
核心思想:把模型的一层横向切分到多张卡上,每张卡只算一部分。
┌─────────────────────────────────────────────────────────────────┐
│ Tensor Parallelism(张量并行) │
├─────────────────────────────────────────────────────────────────┤
│ │
│ 以矩阵乘法 Y = X × W 为例(Y: seq×d_out, W: d_in×d_out): │
│ │
│ 原始(单卡): │
│ │
│ X (seq×d_in) │
│ × │
│ W (d_in×d_out) │
│ = │
│ Y (seq×d_out) │
│ │
│ 张量并行(2 卡,横向切 W): │
│ │
│ X (seq×d_in) │
│ × │
│ ┌────────┬────────┐ │
│ │ W₁ │ W₂ │ ← W 被横向切分成 W₁, W₂ │
│ │(d_in× │(d_in× │ │
│ │ d_out/2)│ d_out/2)│ │
│ └────┬───┴───┬────┘ │
│ ↓ ↓ │
│ Y₁ (seq× Y₂ (seq× ← 两部分结果 │
│ d_out/2) d_out/2) │
│ ↓ ↓ │
│ ┌────┴────────┴────┐ │
│ │ AllReduce │ ← 需要通信把两部分加起来 │
│ └────────┬─────────┘ │
│ ↓ │
│ Y (seq×d_out) │
│ │
│ 在 Transformer 中的应用: │
│ Self-Attention 和 MLP 的矩阵乘法都被切分 │
│ 8 卡张量并行 → 每卡只存 1/8 的权重 │
│ │
│ 优点: │
│ ├─ 突破单卡显存限制 │
│ └─ 通信隐藏在计算中(AllReduce 和矩阵乘法重叠) │
│ │
│ 缺点: │
│ ├─ 需要修改模型代码(不兼容所有模型) │
│ ├─ GPU 间通信量大(每层都要通信) │
│ └─ 需要 NVLink 等高速互联(PCIe 会严重拖慢) │
│ │
└─────────────────────────────────────────────────────────────────┘2.3 Pipeline Parallelism(流水线并行)
核心思想:把模型的不同层放到不同卡上。
┌─────────────────────────────────────────────────────────────────┐
│ Pipeline Parallelism(流水线并行) │
├─────────────────────────────────────────────────────────────────┤
│ │
│ 以 70B 模型为例(80 层),4 卡流水线并行: │
│ │
│ 输入 │
│ ↓ │
│ ┌────────┐ │
│ │ GPU 0 │ ← 层 1-20 (约 17.5B 参数) │
│ └────┬───┘ │
│ ↓ GPU 0 → GPU 1 激活值传递 │
│ ┌────────┐ │
│ │ GPU 1 │ ← 层 21-40 (约 17.5B 参数) │
│ └────┬───┘ │
│ ↓ GPU 1 → GPU 2 │
│ ┌────────┐ │
│ │ GPU 2 │ ← 层 41-60 (约 17.5B 参数) │
│ └────┬───┘ │
│ ↓ GPU 2 → GPU 3 │
│ ┌────────┐ │
│ │ GPU 3 │ ← 层 61-80 (约 17.5B 参数) │
│ └────┬───┘ │
│ ↓ │
│ 输出 │
│ │
│ 流水线并行的问题: │
│ "Bubble"(气泡)= GPU 空闲等待 │
│ │
│ 理想情况: │
│ GPU0: [F0][F1][F2][F3][F4][F5][F6][F7] │
│ GPU1: [F0][F1][F2][F3][F4][F5][F6][F7] │
│ GPU2: [F0][F1][F2][F3][F4][F5][F6][F7] │
│ GPU3: [F0][F1][F2][F3][F4][F5][F6][F7] │
│ ↑____气泡区域(GPU空闲)____↑ │
│ │
│ 解决:Micro-Batch(将 batch 再细分成更小的 micro-batch) │
│ 7B/70B 模型通常用 1-16 个 micro-batch │
│ │
│ 优点: │
│ ├─ 通信量小(只传激活值,相邻层之间) │
│ └─ 每张卡显存压力小(只存部分层) │
│ │
│ 缺点: │
│ ├─ 流水线气泡(GPU 利用率不满) │
│ └─ 实现复杂,需要精细调度 │
│ │
└─────────────────────────────────────────────────────────────────┘2.4 三种并行策略对比
| Data Parallel | Tensor Parallel | Pipeline Parallel | |
|---|---|---|---|
| 切分维度 | 数据 | 层内权重 | 层 |
| 通信对象 | 所有 GPU(AllReduce) | 邻居 GPU(AllReduce) | 邻居 GPU(P2P) |
| 通信量 | 中(每步传梯度) | 高(每层传激活值) | 低(只传激活值) |
| 通信位置 | 反向传播时 | 每层计算后 | 相邻层传递 |
| 显存节省 | 无(每卡存全模型) | 每卡只存 1/N | 每卡只存 1/N 层 |
| GPU 利用率 | 高 | 高 | 中(有气泡) |
| 实现难度 | 低 | 高 | 中 |
第3部分:ZeRO — 显存优化技术
3.1 ZeRO 的核心思想
ZeRO(Zero Redundancy Optimizer) 是 DeepSpeed 提出的显存优化技术。
核心观察:
Data Parallel 中每张卡都存了"重复"的数据:
├─ 完整模型参数(所有卡都有)
├─ 完整梯度(所有卡都有)
└─ 完整优化器状态(所有卡都有)
→ 显存浪费严重!只有一张卡的数据是"有用的",其他都是"冗余的"
ZeRO 的解决思路:
不存"完整"的,而是每张卡只存"一部分",用通信换显存3.2 ZeRO Stage 1/2/3
┌─────────────────────────────────────────────────────────────────┐
│ ZeRO 三个阶段 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ ZeRO-1(优化器状态分片): │
│ ───────────────────────────────────────────────────────────── │
│ 每张卡只存 1/N 的优化器状态 │
│ │
│ 单卡显存节省约 4 倍 │
│ 模型参数和梯度仍全量存储 │
│ │
│ 例如 70B 模型 + 8 卡 ZeRO-1: │
│ 优化器状态:840 GB → 105 GB/卡 │
│ │
│ ───────────────────────────────────────────────────────────── │
│ ZeRO-2(梯度分片): │
│ ───────────────────────────────────────────────────────────── │
│ 每张卡只存 1/N 的优化器状态 + 1/N 的梯度 │
│ │
│ 单卡显存节省约 8 倍 │
│ 模型参数仍全量存储 │
│ │
│ 例如 70B 模型 + 8 卡 ZeRO-2: │
│ 优化器状态:840 GB → 105 GB/卡 │
│ 梯度: 140 GB → 17.5 GB/卡 │
│ │
│ ───────────────────────────────────────────────────────────── │
│ ZeRO-3(参数分片): │
│ ───────────────────────────────────────────────────────────── │
│ 每张卡只存 1/N 的优化器状态 + 梯度 + 参数 │
│ │
│ 单卡显存节省约 N 倍(N=GPU数量) │
│ 通信量增加(需要时广播参数) │
│ │
│ 例如 70B 模型 + 8 卡 ZeRO-3: │
│ 模型参数: 140 GB → 17.5 GB/卡 │
│ 梯度: 140 GB → 17.5 GB/卡 │
│ 优化器状态:840 GB → 105 GB/卡 │
│ │
└─────────────────────────────────────────────────────────────────┘3.3 ZeRO 与并行的组合
ZeRO 不是替代 DP/TP/PP,而是互补的:
常见组合策略:
TP + PP + DP(无 ZeRO):
├─ 大模型训练经典组合
├─ Megatron-Deepspeed 方案
└─ 需要高速互联(NVLink + InfiniBand)
ZeRO-3 + PP(DeepSpeed 方案):
├─ ZeRO-3 分担参数
├─ PP 减少流水线气泡
└─ 适合网络带宽一般的集群
ZeRO-3 + TP(更常见):
├─ TP 已经是层内切分
├─ ZeRO-3 主要分摊优化器状态
└─ 常见于 70B+ 模型的训练
FSDP(Fully Sharded Data Parallel)= ZeRO-3 + DataParallel
├─ 在 PyTorch 原生支持(FSDP API)
└─ 本质上是 ZeRO-3 的另一种封装第4部分:训练基础设施核心组件
4.1 GPU 集群拓扑
┌─────────────────────────────────────────────────────────────────┐
│ GPU 集群硬件拓扑 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ 单机 8 卡 A100(DGX A100): │
│ │
│ ┌───────────────────────────────────────┐ │
│ │ DGX A100(8 卡) │ │
│ │ ┌────┐┌────┐┌────┐┌────┐ │ │
│ │ │GPU0││GPU1││GPU2││GPU3│ │ │
│ │ └─┬──┘└─┬──┘└─┬──┘└─┬──┘ │ │
│ │ └──────┴──────┴──────┘ │ │
│ │ NVSwitch(全互联) │ │
│ │ ┌──────┬──────┬──────┬──────┐ │ │
│ │ ┌─┴──┐┌─┴──┐┌─┴──┐┌─┴──┐ │ │ │
│ │ │GPU4││GPU5││GPU6││GPU7│ │ │ │
│ │ └────┘└────┘└────┘└────┘ │ │ │
│ └───────────────────────────────────────┘ │
│ ↓ │
│ NVLink:900 GB/s(双向) │
│ NVSwitch:每个 GPU 和其他所有 GPU 全互联 │
│ │
│ ───────────────────────────────────────────────────────────── │
│ │
│ 多机集群(通过 InfiniBand 互联): │
│ │
│ ┌────────────┐ ┌────────────┐ ┌────────────┐ │
│ │ DGX A100 │ │ DGX A100 │ │ DGX A100 │ │
│ │ Node 0 │ IB │ Node 1 │ IB │ Node 2 │ │
│ │ 8×A100 │ ──→ │ 8×A100 │ ──→ │ 8×A100 │ │
│ └────────────┘ └────────────┘ └────────────┘ │
│ ↓ │
│ InfiniBand HDR:400 Gb/s(约 50 GB/s) │
│ NVLink-Network(NVLink 跨节点扩展) │
│ │
└─────────────────────────────────────────────────────────────────┘4.2 关键硬件术语
| 术语 | 全称 | 作用 | |
|---|---|---|---|
| NVLink | NVIDIA Link | GPU 间高速互联(900 GB/s),比 PCIe 快 10 倍 | |
| NVSwitch | NVIDIA Switch | DGX 机身内的全互联交换机 | |
| InfiniBand | IB | 跨节点高速网络(400-800 Gb/s) | |
| NCCL | NVIDIA Collective Communications | 英伟达集合通信库(AllReduce、Broadcast 等) | |
| RoCE | RDMA over Converged Ethernet | InfiniBand over 以太网的替代方案 |
4.3 分布式训练框架
┌─────────────────────────────────────────────────────────────────┐
│ 主流分布式训练框架 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ DeepSpeed(微软): │
│ ├─ ZeRO(1/2/3)显存优化 │
│ ├─ 3D 并行(DP + TP + PP) │
│ ├─ 混合精度训练(FP16/BF16) │
│ ├─ 训练checkpoint 压缩 │
│ └─ 大量 LLM 训练使用(Llama、Qwen 等开源模型) │
│ │
│ Megatron-LM(英伟达): │
│ ├─ Tensor Parallelism 实现 │
│ ├─ 高效的张量并行 Attention + MLP │
│ ├─ 流水线并行支持 │
│ └─ 通常和 DeepSpeed 组合使用(Megatron-Deepspeed) │
│ │
│ ColossalAI(潞晨科技): │
│ ├─ 统一的并行策略抽象(auto_parallel) │
│ ├─ 异构训练(CPU-GPU 协同) │
│ └─ PyTorch 兼容性较好 │
│ │
│ PyTorch FSDP: │
│ ├─ PyTorch 原生的 Fully Sharded Data Parallel │
│ ├─ 本质上是 ZeRO-3 的 PyTorch 实现 │
│ └─ API 简单,适合快速实验 │
│ │
└─────────────────────────────────────────────────────────────────┘4.4 混合精度训练(BF16 vs FP16 vs FP32)
┌─────────────────────────────────────────────────────────────────┐
│ 训练精度选择 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ FP32(单精度浮点): │
│ ├─ 4 字节/参数 │
│ ├─ 精度最高,但显存和速度都差 │
│ └─ 主optimizer states 用这个 │
│ │
│ FP16(半精度): │
│ ├─ 2 字节/参数 │
│ ├─ 速度快,显存省一半 │
│ ├─ 但动态范围窄(65535 → 65504),训练可能不稳定 │
│ └─ 2017-2019 年的主流 │
│ │
│ BF16(Brain Float 16,谷歌提出): │
│ ├─ 2 字节/参数 │
│ ├─ 动态范围和 FP32 一样大(3.4×10³⁸) │
│ ├─ 精度比 FP16 低,但训练更稳定 │
│ └─ 2020 年后成为 LLM 训练的主流选择 │
│ │
│ FP8(8 位浮点,NVIDIA H100 支持): │
│ ├─ 1 字节/参数 │
│ ├─ 显存进一步节省 │
│ ├─ 精度挑战大,需要细致校准 │
│ └─ 2024-2025 年开始流行 │
│ │
│ 训练常用配置: │
│ ├─ Forward/Backward:BF16(速度和显存效率) │
│ ├─ Optimizer States:FP32(保证精度) │
│ └─ 这就是"混合精度训练" │
│ │
└─────────────────────────────────────────────────────────────────┘第5部分:完整 Pipeline 与资源消耗
5.1 Post-Training Pipeline 完整流程图
┌─────────────────────────────────────────────────────────────────┐
│ LLM Post-Training Pipeline │
├─────────────────────────────────────────────────────────────────┤
│ │
│ Pretrained Model │
│ │ │
│ ├─ 参数量:7B / 13B / 70B / 405B │
│ ├─ 精度:FP16 / BF16 │
│ └─ 来源:预训练结束后的 checkpoint │
│ │ │
│ ↓ │
│ ┌─────────────────────────────────────────────────────────┐ │
│ │ 阶段1:数据收集与清洗 │ │
│ │ ├─ 任务分布设计(代码/对话/知识/推理各占多少) │ │
│ │ ├─ 数据来源(开源 + 人工标注) │ │
│ │ ├─ 质量过滤(去重、有毒内容过滤、语言识别) │ │
│ │ └─ 规模:SFT 10K-1M条,偏好数据 10K-100K对 │ │
│ └─────────────────────────────────────────────────────────┘ │
│ │ │
│ ↓ │
│ ┌─────────────────────────────────────────────────────────┐ │
│ │ 阶段2:SFT(Supervised Fine-Tuning) │ │
│ │ ├─ 数据:Prompt-Response 对 │ │
│ │ ├─ 规模:10K - 1M 条数据 │ │
│ │ ├─ 资源:8-64 张 A100/H100 │ │
│ │ ├─ 耗时:数天到数周 │ │
│ │ └─ 产出:SFT Model(能遵循指令,但可能不够对齐) │ │
│ └─────────────────────────────────────────────────────────┘ │
│ │ │
│ ↓ │
│ ┌─────────────────────────────────────────────────────────┐ │
│ │ 阶段3:Reward Model 训练(RLHF 路线) │ │
│ │ ├─ 数据:人类偏好对比数据(10K-100K 对比) │ │
│ │ ├─ 资源:8-64 张卡 │ │
│ │ └─ 耗时:数天 │ │
│ │ │ │
│ │ 或跳过此阶段(DPO 路线) │ │
│ └─────────────────────────────────────────────────────────┘ │
│ │ │
│ ↓ │
│ ┌─────────────────────────────────────────────────────────┐ │
│ │ 阶段4:对齐微调(RLHF / DPO / CAI) │ │
│ │ ├─ RLHF:PPO 微调,需 RM + Ref Model │ │
│ │ ├─ DPO:直接偏好优化,只需 Ref Model │ │
│ │ ├─ 资源:8-64 张卡 │ │
│ │ ├─ 耗时:数天到数周 │ │
│ │ └─ 产出:Alignment 后的模型 │ │
│ └─────────────────────────────────────────────────────────┘ │
│ │ │
│ ↓ │
│ ┌─────────────────────────────────────────────────────────┐ │
│ │ 阶段5:Safety Red Teaming │ │
│ │ ├─ 专门测试有害内容、越狱 jailbreak 等 │ │
│ │ ├─ 方式:人工红队 + 自动红队 │ │
│ │ └─ 耗时:数周到数月(持续迭代) │ │
│ └─────────────────────────────────────────────────────────┘ │
│ │ │
│ ↓ │
│ ┌─────────────────────────────────────────────────────────┐ │
│ │ 阶段6:Benchmark 评测 │ │
│ │ ├─ MMLU:多任务知识理解(57 个学科) │ │
│ │ ├─ HumanEval:代码生成(164 道编程题) │ │
│ │ ├─ GSM8K:数学推理(中学数学) │ │
│ │ ├─ MT-Bench:多轮对话 │ │
│ │ └─ 耗时:数天 │ │
│ └─────────────────────────────────────────────────────────┘ │
│ │ │
│ ↓ │
│ ┌─────────────────────────────────────────────────────────┐ │
│ │ 模型发布 │ │
│ │ ├─ Base 版本:预训练模型(供研究) │ │
│ │ ├─ Instruct 版本:SFT 后(能遵循指令) │ │
│ │ └─ Chat 版本:完整对齐后(安全 + 有用) │ │
│ └─────────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────┘5.2 资源消耗一览
┌─────────────────────────────────────────────────────────────────┐
│ 不同规模模型的训练资源参考 │
├─────────────────────────────────────────────────────────────────┤
│ │
│ 7B 模型(70 亿参数): │
│ ├─ 预训练:512-1024 张 A100,2-3 个月 │
│ ├─ SFT:8-64 张卡,2-7 天 │
│ ├─ DPO:8-64 张卡,1-3 天 │
│ └─ 单卡可微调(用 LoRA/QLoRA) │
│ │
│ 13B 模型: │
│ ├─ 预训练:1024-2048 张 A100,2-4 个月 │
│ ├─ SFT:16-64 张卡,1-2 周 │
│ └─ 单卡微调困难(需要 LoRA/QLoRA) │
│ │
│ 70B 模型: │
│ ├─ 预训练:2048-4096 张 H100/A100,3-6 个月 │
│ ├─ SFT:64-128 张卡,2-4 周 │
│ ├─ DPO:64-128 张卡,1-2 周 │
│ └─ 必须多卡训练 │
│ │
│ 405B 模型(Llama 3.1): │
│ ├─ 预训练:16384 张 H100,~54 天(官方数据) │
│ ├─ SFT:128-512 张卡,数周 │
│ └─ 需要 TP+PP+DP+ZeRO 全套并行 │
│ │
│ Post-Training 全流程(7B): │
│ SFT + RM + DPO ≈ 100-500 A100-GPU-days │
│ 预训练 ≈ 10000+ A100-GPU-days │
│ Post-Training 只占 < 5% 的总算力 │
│ │
└─────────────────────────────────────────────────────────────────┘第6部分:常见术语扫盲
6.1 训练策略术语
| 术语 | 全称 | 解释 | |
|---|---|---|---|
| FT | Fine-tuning | 全量微调,更新所有参数 | |
| PEFT | Parameter-Efficient FT | 参数高效微调,只改部分参数 | |
| LoRA | Low-Rank Adaptation | 低秩适配,最流行的 PEFT 方法 | |
| QLoRA | Quantized LoRA | LoRA + 4-bit 量化,单卡可训大模型 | |
| Adapter | Adapter Tuning | 插入小型适配层,不改原模型 | |
| RLHF | RL from Human Feedback | 人类反馈强化学习 | |
| RLAIF | RL from AI Feedback | AI 反馈强化学习 | |
| DPO | Direct Preference Optimization | 直接偏好优化,RLHF 简化版 | |
| ORPO | Odds Ratio Preference Optimization | 一种新的对齐方法 |
6.2 并行策略术语
| 术语 | 全称 | 解释 | |
|---|---|---|---|
| DP | Data Parallelism | 数据并行,每卡完整模型 | |
| TP | Tensor Parallelism | 张量并行,层内切分 | |
| PP | Pipeline Parallelism | 流水线并行,层间切分 | |
| ZeRO | Zero Redundancy Optimizer | 显存优化,三阶段 | |
| FSDP | Fully Sharded Data Parallel | PyTorch 原生 ZeRO-3 | |
| EP | Expert Parallelism | MoE 模型专用,专家路由并行 |
6.3 硬件与 Infra 术语
| 术语 | 全称 | 解释 | |
|---|---|---|---|
| NCCL | NVIDIA Collective Comms | 英伟达 GPU 集合通信库 | |
| NVLink | NVIDIA Link | GPU 间高速互联(900 GB/s) | |
| IB | InfiniBand | 跨节点高速网络 | |
| BF16 | Brain Float 16 | LLM 训练主流精度 | |
| FP16 | Float 16 | 半精度,2019 年前主流 | |
| FP8 | Float 8 | 8 位精度,H100 开始支持 | |
| CKPT | Checkpoint | 模型训练快照,断点恢复用 | |
| DS | DeepSpeed | 微软分布式训练框架 |
核心总结
总结1:分布式训练的三大策略
Data Parallel:每卡完整模型,不同数据
Tensor Parallel:每卡一层的一部分(层内切分)
Pipeline Parallel:每卡不同的层(层间切分)
实际用哪个?
├─ 单卡能放下的模型:DP 足够
├─ 大模型(单卡放不下):TP + PP + DP
└─ 显存不够:ZeRO 来凑总结2:ZeRO 的三个阶段
ZeRO-1:优化器状态分片(省 4× 显存)
ZeRO-2:梯度 + 优化器状态分片(省 8× 显存)
ZeRO-3:参数 + 梯度 + 优化器状态全分片(省 N× 显存)
原则:用通信换显存总结3:Post-Training Pipeline
Pretrained Model → 数据清洗 → SFT → RM → DPO/RLHF → 红队 → 评测 → 发布
整个 Post-Training 算力 < 预训练的 5%
但对齐质量决定了模型是否"好用"总结4:硬件选择
H100 > A100 > 3090(性价比)
互联:NVLink + IB > PCIe
精度:BF16 是 LLM 训练主流章节测试
测试0:ML Infra 三分
ML Infra 可以分为哪三类?请分别描述它们的核心挑战和资源量级。
测试1:显存瓶颈
计算 13B 参数模型在 BF16 精度下的显存占用(参数 + 梯度 + 优化器状态),并判断单卡 A100 80GB 是否足够。
测试2:并行策略选择
如果要训练一个单卡放不下(200GB+)但层数不太深的模型,应该优先考虑哪两种并行策略?为什么?
测试3:ZeRO vs TP
ZeRO-3 和 Tensor Parallelism 都能突破单卡显存限制,它们的主要区别是什么?
测试4:Pipeline 气泡
解释流水线并行中"气泡"(bubble)是怎么产生的,以及 Micro-Batch 如何缓解它。
测试5:完整 Pipeline
描述一个 70B 模型从 Base Model 到发布的完整 Post-Training Pipeline,并估算各环节的资源消耗。
测试6:术语选择
某创业公司想在 4 张 3090(24GB)上微调一个 7B 模型,你会推荐哪些技术和术语?
参考答案
测试0答案
答案:
1. Pre-Training Infra(预训练 Infra)
├─ 核心挑战:千卡并行稳定性、PB 级数据吞吐、超大规模 Checkpoint
└─ 资源量级:千卡 × 数月
2. Post-Training SFT Infra(SFT Infra)
├─ 核心挑战:高质量标注数据的 pipeline、多任务数据配比
└─ 资源量级:8-128 卡 × 数天到数周
3. Post-Training RL Infra(RL Infra)
├─ 核心挑战:多模型并发(LLM+RM+Ref)、实时数据流管道、PPO 训练稳定性
└─ 资源量级:8-128 卡 × 数天到数月测试1答案
答案:
13B 参数模型 BF16 精度(2 字节/参数):
├─ 模型参数:13B × 2B = 26 GB
├─ 梯度: 13B × 2B = 26 GB
├─ 优化器状态:13B × 12B = 156 GB(Adam 在 FP32)
└─ 总计:> 208 GB
单卡 A100 80GB:不够!(差 128GB+)
解决方案:
├─ ZeRO-3 + 8 卡:每卡 208/8 ≈ 26 GB ← 刚好够
├─ QLoRA(4-bit 量化):单卡可训
└─ DeepSpeed ZeRO-3:减少优化器状态的精度测试2答案
答案:优先考虑 Tensor Parallelism(TP) + Pipeline Parallelism(PP)。
解析:
"单卡放不下但层数不太深"的特点:
├─ 参数量大(需要切分)
└─ 但不是特别深(不需要太多 PP stage)
优先选择:
1. Tensor Parallelism:
- 层内横向切分,最适合参数量大的模型
- 可以把单层的大矩阵乘法分散到多卡
- 每卡显存压力大幅减少
2. Pipeline Parallelism(配合 TP):
- 如果单靠 TP 还不够,加上 PP
- 把不同层放到不同卡上
- 额外减少每卡的显存压力
不优先选择 Data Parallel:
- DP 每卡都要存完整模型
- 无法解决"单卡放不下"的问题测试3答案
答案:
核心区别:切分维度不同
ZeRO-3:
├─ 切分的是"同一个东西的不同副本"
│ (参数、梯度、优化器状态被分片到不同卡)
├─ 每张卡在需要时才获取完整的参数(通信换取显存)
└─ 通信:AllGather(获取参数)+ ReduceScatter(同步梯度)
Tensor Parallelism:
├─ 切分的是"计算本身"
│ (一个矩阵乘法被横向或纵向切分)
├─ 每张卡始终只有自己的部分,不需要广播完整参数
└─ 通信:每层计算后需要 AllReduce
类比:
ZeRO-3 = 把一本书复印 N 份,每人一页(用时再借)
TP = 把一本书撕成 N 份,每人几页(永远只有自己的)
实际用法:TP + ZeRO-3 组合(Megatron + DeepSpeed)测试4答案
答案:
气泡(bubble)产生的原因:
流水线并行中,GPU 之间需要等待:
- GPU 0 算完第 1-20 层,才能把激活值传给 GPU 1
- GPU 1 必须等收到激活值才能开始算
时序图(4 卡,无 micro-batch):
GPU0: [F][F][F][F] [B][B][B][B]
GPU1: [wait][F][F][F][F] [wait][B][B][B][B]
GPU2: [wait][F][F][F][F] [wait][B][B][B][B]
GPU3: [wait][F][F][F][F][wait][B][B][B][B]
↑_____________巨大的气泡区域_____________↑
每个 GPU 在等待上游数据时都在空闲
Micro-Batch 的缓解方法:
将一个 batch 分成多个 micro-batch:
batch = 32 → micro_batch_size = 4(8 个 micro-batch)
GPU0: [F0][F1][F2][F3][F4][F5][F6][F7] [B7][B6][B5][B4][B3][B2][B1][B0]
GPU1: [F0][F1][F2][F3][F4][F5][F6][F7] [wait][B7][B6][B5][B4][B3][B2][B1][B0]
GPU2: [F0][F1][F2][F3][F4][F5][F6][F7] [B7][B6][B5][B4][B3][B2][B1][B0]
GPU3: [F0][F1][F2][F3][F4][F5][F6][F7] [B7][B6][B5][B4][B3][B2][B1][B0]
气泡大大减少!
GPU0 的反向算完时,GPU1 的正向刚好传完数据测试5答案
答案:
70B 模型完整 Post-Training Pipeline:
1. 数据收集与清洗
├─ 资源:主要是人力 + 少量 GPU
└─ 耗时:4-8 周
2. SFT
├─ 数据量:10 万-100 万条
├─ 资源:64-128 张 A100/H100
└─ 耗时:2-4 周
3. Reward Model(RLHF 路线)
├─ 数据量:10-30 万对偏好数据
├─ 资源:8-32 张卡
└─ 耗时:1-2 周
4. DPO 或 RLHF 对齐
├─ 资源:64-128 张卡(DPO 稍少)
└─ 耗时:2-4 周
5. Safety Red Teaming
├─ 资源:主要是人工
└─ 耗时:4-12 周(持续迭代)
6. Benchmark 评测
├─ 资源:少量 GPU(做推理评测)
└─ 耗时:1-2 周
总耗时:3-6 个月(整个 Post-Training)
对比预训练(估算):数千 GPU × 数月
Post-Training 约占总训练成本的 5% 以下测试6答案
答案:
推荐技术组合:
1. QLoRA(核心)
├─ 4-bit 量化模型主体
├─ LoRA 适配器只训练 1-2% 参数
└─ 4 张 3090(24GB)可以运行 7B 模型
2. DeepSpeed ZeRO-2 或 ZeRO-3
├─ ZeRO-2:梯度 + 优化器状态分片
└─ 节省显存,配合 QLoRA 使用
3. 梯度检查点(Gradient Checkpointing)
├─ 用计算换显存
└─ 减少激活值的显存占用
4. BF16 或 FP16 混合精度
├─ Forward BF16,Optimizer FP32
└─ 平衡速度和精度
不推荐:
├─ 全量微调 FT(3090 显存不够)
├─ 张量并行 TP(3090 不支持 NVLink,通信太慢)
└─ 标准 RLHF(太复杂,3090 跑不动)相关笔记
- [[11 - 训练基础扫盲]] - 训练循环、Loss、优化器基础
- [[12 - Post-Training Pipeline]] - SFT / RLHF / DPO 对齐方法
下一步学习
- [ ] 回到 00 - LLM 学习路线总览 重新规划学习路径
学习状态:🟡 待学习