Skip to content
Gains Summary
Main Navigation 首页 / Home
C++ 编程 / C++ Programming
系统与高性能 / Systems & Performance
Web 开发 / Web Development
人工智能 / Artificial Intelligence
工业软件 / Industrial Software
其他内容 / Other Topics
C++ 编程 / C++系统与性能 / SystemsWeb 开发 / Web人工智能 / AI工业软件 / Industrial

外观

Sidebar Navigation

← 人工智能 / Artificial Intelligence

AI 基础设施 / AI Infrastructure

集群基础设施 / Cluster Infrastructure

1. GPU 集群基础设施全景——训练框架之下、硬件之上的那一层 / GPU Cluster Infrastructure Between Training Frameworks and Hardware

2. GPU 集群硬件架构——从 NVLink 到 InfiniBand / GPU Cluster Hardware from NVLink to InfiniBand

3. 异构硬件生态——CPU/DPU/NPU 的集群角色 / Roles of CPUs, DPUs, and NPUs in Heterogeneous Clusters

4. GPU 虚拟化与资源隔离——一张卡多人用 / GPU Virtualization and Resource Isolation

5. 作业调度系统——Kubernetes 和 Slurm / Job Scheduling with Kubernetes and Slurm

6. 多作业与多租户管理——让集群被所有人高效使用 / Multi-Job and Multi-Tenant Cluster Management

7. 网络架构与 RDMA——让 GPU 之间的通信更快 / Network Architecture and RDMA for Faster GPU Communication

8. NCCL 集群组网——大规模集合通信调优 / NCCL Cluster Networking and Collective Communication Tuning

9. 分布式存储——让数据跑得比 GPU 快 / Distributed Storage That Keeps GPUs Fed with Data

10. 集群运营与故障处理——让万卡集群稳定运行 / Operations and Failure Recovery for Large GPU Clusters

训练系统 / Training Systems

1. AI Infra 训练侧全景——让千亿参数模型跑起来需要什么 / Training-Side AI Infrastructure for Hundred-Billion-Parameter Models

2. GPU 硬件基础——为什么 GPU 比 CPU 快,显存为什么总是不够 / GPU Hardware, Parallel Throughput, and Memory Capacity

3. 分布式训练——如何把大模型分到多张卡上 / Distributing Large-Model Training Across Multiple GPUs

4. 显存优化——让 70B 模型在有限显存中跑起来 / Memory Optimization for Running 70B Models

5. 混合精度与通信——BF16 为什么是 LLM 训练的主流选择 / Mixed Precision and Communication with BF16

6. 预训练——Scaling Laws、数据工程与训练稳定性 / Pretraining with Scaling Laws, Data Engineering, and Stability

7. 后训练 SFT——从预训练模型到助手模型 / Supervised Fine-Tuning from Pretrained Model to Assistant

8. 后训练 RLHF/DPO——从助手模型到对齐模型 / RLHF and DPO from Assistant Model to Aligned Model

9. 高效微调——LoRA 和 QLoRA 让大模型走进消费级 GPU / Efficient Fine-Tuning with LoRA and QLoRA on Consumer GPUs

10. 训练工程——千卡集群的管理与故障恢复 / Training Engineering for Thousand-GPU Cluster Operations and Recovery

本页目录

混合精度与通信——BF16 为什么是 LLM 训练的主流选择 / Mixed Precision and Communication with BF16 ​

📅 创建时间:2026-06-02 🏷️ 标签:#混合精度 #BF16 #FP8 #FP16 #NCCL #通信优化 📚 前置知识:[[01-gpu-hardware]](GPU 硬件基础) [[02-distributed-training]](分布式训练) 📚 相关知识:[[03-memory-optimization]](显存优化) [[05-pretraining]](预训练稳定性)


场景:FP16 训练 Loss 溢出,换 BF16 后稳定了 ​

┌─────────────────────────────────────────────────────────────────┐
│                                                             │
│  凌晨 3 点,你的训练脚本报错了。                            │
│                                                             │
│  第一次尝试:FP16 训练                                   │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  Loss: 3.45 → 12.78 → nan                        │  │
│  │  梯度爆炸了!Loss 直接溢出变成 NaN。                │  │
│  │  排查:learning rate、初始化、FP16 动态范围       │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  第二次尝试:换成 BF16 训练                              │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  Loss: 3.45 → 3.21 → 2.98 → 2.75                │  │
│  │  稳定下降,没有溢出。                               │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  为什么同一个模型,FP16 会溢出,BF16 就不会?              │
│  这就是本章要回答的问题。                                  │
│                                                             │
└─────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21

第1节:浮点数基础——为什么精度用 16 位就够了 ​

浮点数的结构 ​

┌─────────────────────────────────────────────────────────────┐
│                 浮点数的三部分构成                              │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  一个浮点数由三部分组成:                                   │
│                                                             │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  (-1)^sign  ×  (1 + fraction)  ×  2^exponent      │  │
│  │      1 bit        10/8/7 bits        5/8/7 bits     │  │
│  │                                                             │  │
│  │  sign     = 符号位(正 or 负)                        │  │
│  │  exponent = 指数位(表示数量级,10^x 中的 x)         │  │
│  │  fraction = 尾数位(表示精度,1.xxxx 中的 xxxx)      │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  动态范围(Dynamic Range)= 能表示的最大数和最小数之间的比值  │
│  精度(Precision)= 两个相邻可表示数之间的最小间隔            │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19

FP16 vs BF16 vs FP32——结构对比 ​

┌─────────────────────────────────────────────────────────────┐
│                 FP16 vs BF16 vs FP32 内存布局对比               │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  FP32(单精度,4 字节 = 32 bits):                        │
│  ┌────────┬──────────┬───────────────────┐             │
│  │ Sign   │ Exponent │ Fraction(尾数)   │             │
│  │  1 bit │  8 bits  │    23 bits        │             │
│  │  正/负  │  2^e     │  精度位            │             │
│  └────────┴──────────┴───────────────────┘             │
│  动态范围:~10^38  (和 BF16 一样大)                      │
│  精度:~7 位十进制小数                                     │
│                                                             │
│  FP16(半精度,2 字节 = 16 bits):                       │
│  ┌────────┬──────────┬───────────────────┐             │
│  │ Sign   │ Exponent │ Fraction(尾数)   │             │
│  │  1 bit │  5 bits  │    10 bits        │             │
│  └────────┴──────────┴───────────────────┘             │
│  动态范围:~65504  (最大约 6.5×10^4)                    │
│  精度:~4 位十进制小数                                     │
│                                                             │
│  BF16(Brain Float,2 字节 = 16 bits):                  │
│  ┌────────┬──────────┬───────────────────┐             │
│  │ Sign   │ Exponent │ Fraction(尾数)   │             │
│  │  1 bit │  8 bits  │    7 bits        │             │
│  └────────┴──────────┴───────────────────┘             │
│  动态范围:~10^38  (和 FP32 一样大!)                    │
│  精度:~3 位十进制小数                                     │
│                                                             │
│  关键差异:                                               │
│  FP16 把 16 bits 分给指数(5)和尾数(10)                │
│  BF16 把 16 bits 中 8 bits 给指数,7 bits 给尾数        │
│  → BF16 的指数和 FP32 一样多,动态范围相同               │
│  → BF16 只比 FP16 少了 3 bits 的精度,但动态范围大得多   │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36

为什么 FP16 会溢出 ​

┌─────────────────────────────────────────────────────────────┐
│                 FP16 溢出 vs BF16 稳定的根本原因                 │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  问题出在动态范围:                                        │
│                                                             │
│  FP16 的最大表示数:                                       │
│  (1 + 0.1111111111)₂ × 2^30 = 1.999 × 2^30 ≈ 65504   │
│  → 任何超过 65504 的数字,在 FP16 中就会变成 Infinity     │
│                                                             │
│  LLM 训练中的大数字来源:                                  │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  1. 损失值大:Cross-Entropy Loss 早期可能是 10-20   │  │
│  │  2. 梯度累积:多层链式乘法,梯度可能非常大            │  │
│  │  3. 权重数值:LayerNorm、Attention 等的中间值        │  │
│  │  4. Loss Scaling:动态损失缩放翻车时更容易溢出        │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  BF16 的最大表示数:                                       │
│  (1 + 0.1111111)₂ × 2^254 ≈ 3.4 × 10^38              │
│  → 和 FP32 一样大,训练中的任何中间值都不会溢出             │
│                                                             │
│  结论:                                                   │
│  BF16 = FP32 的动态范围 + FP16 的显存效率                 │
│  这就是为什么 2020 年后,所有 LLM 训练都迁移到了 BF16     │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27

第2节:混合精度训练原理——什么时候用 FP32 ​

什么是混合精度训练 ​

┌─────────────────────────────────────────────────────────────┐
│                 混合精度训练:不同步骤用不同精度                 │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  核心观察:                                                │
│  → Forward 和 Backward 用 BF16 足够(动态范围大,不溢出)  │
│  → Optimizer State 用 FP32 更好(精度高,更新更准确)      │
│                                                             │
│  混合精度训练的工作流程:                                    │
│                                                             │
│  Forward(BF16):                                         │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  Input (FP32)                                       │  │
│  │      ↓                                              │  │
│  │  模型权重转 BF16                                     │  │
│  │      ↓                                              │  │
│  │  Forward 计算(全部 BF16)→ Output                  │  │
│  │      ↓                                              │  │
│  │  Loss 计算(BF16)                                   │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  Backward(BF16):                                        │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  Loss (BF16)                                       │  │
│  │      ↓                                              │  │
│  │  反向传播(全部 BF16)→ 梯度(BF16)               │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  参数更新(FP32):                                         │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  梯度 BF16 → 转 FP32                                │  │
│  │      ↓                                              │  │
│  │  Optimizer State (FP32) + 梯度 (FP32) → 更新     │  │
│  │      ↓                                              │  │
│  │  更新后的权重(FP32)→ 转 BF16(用于下一步 Forward)│  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38

PyTorch 混合精度实现 ​

python
from torch.cuda.amp import autocast, GradScaler

# GradScaler:管理 loss scale,防止 BF16 下溢出
scaler = GradScaler()

for data, labels in dataloader:
    data = data.cuda()
    labels = labels.cuda()

    # Forward:用 autocast 自动把支持的操作转为 BF16/FP16
    with autocast(dtype=torch.bfloat16):
        outputs = model(data)
        loss = loss_fn(outputs, labels)

    # Backward:用 scaler.scale 防止梯度溢出
    scaler.scale(loss).backward()

    # 参数更新:用 scaler.step(内部做了 unscaling)
    scaler.step(optimizer)
    scaler.update()  # 动态调整 loss scale
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20

Loss Scaling——FP16 时代的补救方案 ​

┌─────────────────────────────────────────────────────────────┐
│                 Loss Scaling 的原理(FP16 时代的技术)            │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  背景:FP16 动态范围小,容易溢出                            │
│                                                             │
│  核心思想:                                                │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  梯度值太小(小于 FP16 最小精度)?                    │  │
│  │  → 乘以一个大的 scale(S),变成 FP16 可表示的数    │  │
│  │  → 正常更新                                         │  │
│  │  → 除以 S,恢复到原始尺度                            │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  动态 Loss Scaling:                                        │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  如果当前 batch 的梯度全是 0(溢出变 Inf):         │  │
│  │  → Loss scale 翻倍(2× S)                         │  │
│  │  → 重算                                             │  │
│  │  如果连续 N 个 batch 没有溢出:                       │  │
│  │  → Loss scale 减半(S/2)                         │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  BF16 不需要 Loss Scaling(动态范围够大)                   │
│  但 GradScaler 在 BF16 训练中仍然有用(极端情况保护)       │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27

第3节:FP8——H100 的下一代精度 ​

FP8 的两种格式 ​

┌─────────────────────────────────────────────────────────────┐
│                 FP8:H100 的 8 位浮点精度                      │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  H100 支持两种 FP8 格式(NVIDIA 发明):                    │
│                                                             │
│  E4M3(4 位指数 + 3 位尾数):                            │
│  ┌────────┬──────────┬───────────────────┐             │
│  │ Sign   │ Exponent │ Fraction         │             │
│  │  1 bit │  4 bits  │    3 bits       │             │
│  └────────┴──────────┴───────────────────┘             │
│  范围:±448(最大),精度约 2 位小数                       │
│  → 用于:Forward 计算(精度要求相对较低)                  │
│                                                             │
│  E5M2(5 位指数 + 2 位尾数):                            │
│  ┌────────┬──────────┬───────────────────┐             │
│  │ Sign   │ Exponent │ Fraction         │             │
│  │  1 bit │  5 bits  │    2 bits       │             │
│  └────────┴──────────┴───────────────────┘             │
│  范围:±57344(最大),精度约 1 位小数                    │
│  → 用于:Backward 计算(动态范围更重要)                  │
│                                                             │
│  为什么 Forward 和 Backward 用不同精度?                    │
│  Forward:需要高精度(模型权重对精度敏感)                 │
│  Backward:梯度动态范围大,精度要求较低                   │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27

FP8 在 LLM 训练中的应用 ​

┌─────────────────────────────────────────────────────────────┐
│                 FP8 训练的实际应用                             │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  FP8 的速度优势:                                          │
│  H100 FP8 算力:989 TFLOPS(稀疏后 1978 TFLOPS)        │
│  H100 BF16 算力:989 TFLOPS                               │
│  H100 FP16 算力:989 TFLOPS                               │
│  → FP8 峰值算力和 BF16 一样,但显存减半!                 │
│                                                             │
│  但 FP8 的挑战:                                           │
│  1. 精度损失:动态范围更小,需要精细的 scaling             │
│  2. 实现复杂度:Transformer Engine 是关键                   │
│  3. 不是所有算子都支持 FP8(Softmax、LayerNorm 等仍需 BF16)│
│                                                             │
│  典型配置(2024-2025 年):                                │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  Forward:FP8(矩阵乘法)+ BF16(Softmax/LayerNorm)│  │
│  │  Backward:FP8(梯度计算)+ BF16(通信)             │  │
│  │  Optimizer:FP32(保持高精度)                       │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  Transformer Engine(TE):                                 │
│  → 自动处理精度转换 + 动态 scaling                        │
│  → DeepSpeed 和 Megatron 都集成了 TE                      │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27

FP16 / BF16 / FP8 精度选择对比 ​

┌─────────────────────────────────────────────────────────────┐
│                 训练精度选择全景图                              │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  │ 精度    │ 显存占用 │ 动态范围 │ 训练稳定性 │ 速度   │
│  ├─────────┼──────────┼──────────┼────────────┼─────────┤
│  │ FP32    │  100%   │  最大   │  最稳定    │ 最慢   │
│  ├─────────┼──────────┼──────────┼────────────┼─────────┤
│  │ BF16    │  50%    │  最大   │  稳定      │  快    │
│  ├─────────┼──────────┼──────────┼────────────┼─────────┤
│  │ FP16    │  50%    │  小    │  不稳定(需LS)│ 快    │
│  ├─────────┼──────────┼──────────┼────────────┼─────────┤
│  │ FP8     │  25%    │  中    │  较稳定    │  最快  │
│                                                             │
│  2020-2023:BF16 是主流选择(GPT-4 之前)                 │
│  2024-2025:FP8 开始在 H100 上普及(训练速度优势)         │
│  FP16 在现代 LLM 训练中已基本被淘汰                        │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19

第4节:NCCL 通信原语——分布式训练的通信基础 ​

为什么分布式训练需要通信 ​

┌─────────────────────────────────────────────────────────────┐
│                 分布式训练中的通信需求                            │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  数据并行(DP):                                           │
│  每张卡独立计算梯度 → 需要同步 → AllReduce                  │
│                                                             │
│  张量并行(TP):                                          │
│  每张卡算一部分矩阵乘法 → 需要聚合 → AllReduce              │
│                                                             │
│  流水线并行(PP):                                        │
│  GPU 之间的激活值和梯度传递 → 点对点通信(P2P)            │
│                                                             │
│  FSDP/ZeRO-3:                                            │
│  参数分片 → 需要时广播 → AllGather                         │
│                                                             │
│  通信 = 分布式训练的隐形瓶颈                                │
│  通信慢 → GPU 等待 → MFU 下降                            │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20

三大通信原语 ​

┌─────────────────────────────────────────────────────────────┐
│                 AllReduce / AllGather / ReduceScatter              │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  AllReduce(最常用):                                      │
│  ┌─────────────────────────────────────────────────────┐  │
│  │                                                     │  │
│  │  GPU 0: [1, 2, 3, 4]      AllReduce → [10, 20, 30, 40]  │
│  │  GPU 1: [2, 4, 6, 8]      ────────────────→ 每张卡都一样  │
│  │  GPU 2: [3, 6, 9, 12]                            │  │
│  │  GPU 3: [4, 8, 12, 16]                           │  │
│  │                                                     │  │
│  │  操作:对应位置求和/平均/最大值等                  │  │
│  │  用途:数据并行梯度同步                             │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  AllGather(广播):                                       │
│  ┌─────────────────────────────────────────────────────┐  │
│  │                                                     │  │
│  │  GPU 0: [1, 2]                                   │  │
│  │  GPU 1: [3, 4]      AllGather → [1, 2, 3, 4]   │  │
│  │  GPU 2: [5, 6]       ──────────────→ 每张卡都一样  │  │
│  │  GPU 3: [7, 8]                                   │  │
│  │                                                     │  │
│  │  操作:收集所有 GPU 的数据,拼成完整结果             │  │
│  │  用途:ZeRO-3 参数分片广播                         │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  ReduceScatter(分片归约):                               │
│  ┌─────────────────────────────────────────────────────┐  │
│  │                                                     │  │
│  │  GPU 0: [1, 2, 3, 4]                             │  │
│  │  GPU 1: [2, 4, 6, 8]  ReduceScatter → GPU0:[3,6,9,12]   │
│  │  GPU 2: [3, 6, 9, 12]  ──────────────→ GPU1:[6,12,18,24] │
│  │  GPU 3: [4, 8, 12, 16]               → GPU2:[9,18,27,36]   │
│  │                                 → GPU3:[12,24,36,48]  │
│  │                                                     │  │
│  │  操作:先求和,然后分片发给各 GPU                  │  │
│  │  用途:ZeRO 梯度分片收集                           │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42

NCCL——GPU 集合通信的标准库 ​

┌─────────────────────────────────────────────────────────────┐
│                 NCCL:NVIDIA Collective Communications Library      │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  NCCL 是 NVIDIA 开发的 GPU 间通信库:                        │
│  → 所有主流分布式训练框架都在用(DeepSpeed / Megatron / FSDP)│
│  → 支持 GPUDirect RDMA(绕过 CPU 直接通信)                 │
│  → 自动选择最优通信路径(NVLink / IB / PCIe)              │
│                                                             │
│  NCCL Ring vs Tree:                                       │
│                                                             │
│  Ring AllReduce(链路环形):                               │
│  GPU0 → GPU1 → GPU2 → GPU3 → GPU0                       │
│  → 通信量:2×(N-1)/N × data(次优但稳定)                │
│  → 适合:跨节点通信(IB)                                  │
│                                                             │
│  Tree AllReduce(树形):                                  │
│  ┌─────────────────────────────────────────────────────┐  │
│  │         GPU0                                        │  │
│  │        /    \                                      │  │
│  │     GPU1    GPU2                                   │  │
│  │      |        |                                    │  │
│  │     GPU3    GPU4 ...                               │  │
│  └─────────────────────────────────────────────────────┘  │
│  → 通信量:(N-1)/N × data(最优)                        │
│  → 适合:节点内通信(NVLink)                             │
│                                                             │
│  NCCL 自动选择:节点内用 Tree,跨节点用 Ring              │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30

通信与计算的重叠——隐藏延迟 ​

┌─────────────────────────────────────────────────────────────┐
│                 通信与计算重叠:让 GPU 不等待                       │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  朴素做法(串行):                                        │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  Forward ──────────────────▶ Backward ──────────▶  │  │
│  │  全部完成后 → AllReduce ─────────────────────────▶  │  │
│  │                          ↑ GPU 在这里空等               │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  重叠做法(并行):                                        │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  Micro-batch 1: F1 ──▶ B1 ──▶ AllReduce           │  │
│  │  Micro-batch 2:       F2 ──▶ B2 (和上一步 AllReduce 并行) │  │
│  │  Micro-batch 3:             F3 ──▶ B3 (和上一步 AllReduce 并行) │  │
│  │                                                     │  │
│  │  通信时间被计算隐藏 → 总时间大幅减少                  │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  在 FSDP 中的应用:                                        │
│  ┌─────────────────────────────────────────────────────┐  │
│  │  当前 micro-batch: AllGather 参数                   │  │
│  │  上一个 micro-batch: Backward 计算(并行)           │  │
│  │  下一个 micro-batch: Forward 计算(并行)           │  │
│  └─────────────────────────────────────────────────────┘  │
│                                                             │
│  实现关键:CUDA Stream(见 [[01-gpu-hardware]])            │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30

升华:精度与效率的永恒博弈 ​

┌─────────────────────────────────────────────────────────────┐
│              从 FP16 到 BF16 到 FP8:精度选择的工程哲学            │
├─────────────────────────────────────────────────────────────┤
│                                                             │
│  1. 精度损失的接受度在提高                                 │
│     → FP32 → BF16:动态范围换取稳定性(可接受)           │
│     → BF16 → FP8:精度换取速度(可接受)                  │
│     → 趋势:能用低精度就不用高精度,能省就省               │
│                                                             │
│  2. 通信永远是瓶颈                                         │
│     → GPU 算力增长 > 显存带宽增长 > 互联带宽增长          │
│     → 越往后,通信越可能成为瓶颈                          │
│     → AllReduce 的优化(Ring vs Tree vs 拓扑感知)是工程重点│
│                                                             │
│  3. BF16 是当前训练的主流,FP8 是未来                     │
│     → H100 的 FP8 硬件支持已经成熟                        │
│     → 但框架支持(Transformer Engine)还在完善中           │
│     → 消费级显卡(4090)不支持 BF16 以外的新精度          │
│                                                             │
│  一句话总结:                                               │
│  混合精度训练的本质是在"精度够用"的边界上,尽可能节省显存和计算。│
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23

"AI 可查 vs 必须理解"清单 ​

AI 可查:
✅ Transformer Engine 的具体配置参数
✅ NCCL 通信的 benchmark 数据
✅ 不同精度在不同 GPU 型号上的支持情况

必须理解:
🔴 FP16 / BF16 / FP32 的内存布局差异,以及为什么 BF16 比 FP16 更适合 LLM 训练
🔴 混合精度训练的流程:Forward BF16 + Optimizer FP32
🔴 FP8 的 E4M3 / E5M2 两种格式及其适用场景
🔴 NCCL 三大通信原语:AllReduce / AllGather / ReduceScatter 的使用场景
🔴 为什么通信与计算重叠能提升 MFU
🔴 为什么张量并行需要 NVLink(带宽需求),而数据并行可以用 InfiniBand
1
2
3
4
5
6
7
8
9
10
11
12

学习状态:🟡 开始学习

最后更新于:

Pager
上一篇4. 显存优化——让 70B 模型在有限显存中跑起来 / Memory Optimization for Running 70B Models
下一篇6. 预训练——Scaling Laws、数据工程与训练稳定性 / Pretraining with Scaling Laws, Data Engineering, and Stability

持续记录,持续成长

Copyright © Tidenflow