跳转至

混合精度训练:BF16 → FP8,哪些层能量化

更新日期:2026-04-15


一、精度演进

flowchart LR
    fp32["FP32<br/>(2017 baseline)"]
    fp16["FP16<br/>(2018 V100)"]
    bf16["BF16<br/>(2020 A100)"]
    fp8["FP8<br/>(2022 H100)"]
    fp4["FP4<br/>(2024 B200)"]

    fp32 --> fp16 --> bf16 --> fp8 --> fp4

    classDef stage fill:#fff,stroke:#cc785c,color:#1a1a1a;
    class fp32,fp16,bf16,fp8,fp4 stage

每一代精度降级 2× → 算力 / 显存 / 带宽都翻倍。但精度风险也累积,需要相应的稳定性技巧(loss scaling / per-tile scaling / FP32 master copy)。


二、各精度对比

BF16 vs FP16 的关键区别:动态范围。FP16 容易在训练中出现 inf/nan,BF16 范围和 FP32 相同,训练稳定得多。


三、BF16 训练(当前主流)

3.1 混合精度策略

# 关键: 不是所有东西都 BF16!
# 某些需要高精度的部分保留 FP32

class MixedPrecisionTrainer:
    def __init__(self, model, optimizer):
        # FP32 master weights (优化器状态)
        self.master_weights = {p: p.data.float() for p in model.parameters()}

        # 模型内的 BF16 权重 (前向/反向用)
        # 由 autocast 自动处理
        self.model = model

    def step(self, batch):
        # 前向: BF16
        with torch.autocast('cuda', dtype=torch.bfloat16):
            loss = self.model(batch)

        # 反向: 梯度默认在 BF16 (可以配置为 FP32)
        loss.backward()

        # 优化器更新: 在 FP32 上
        for p in self.model.parameters():
            master = self.master_weights[p]
            # 梯度 cast 到 FP32
            grad_fp32 = p.grad.float()
            # FP32 更新
            master -= lr * (m_fp32 / (sqrt(v_fp32) + eps))
            # 更新后的权重 cast 回 BF16
            p.data = master.bfloat16()

3.2 精度配置表


四、FP8 训练(DeepSeek 创新)

DeepSeek-V3 是首个在 671B 规模成功使用 FP8 训练的模型,节省 50% 计算量。参考 DeepSeek-V3 FP8 部分

4.1 哪些层能 FP8

可以使用 FP8 的层:

  • 线性层矩阵乘法(主要 FLOPs 来源)

  • FFN 的两个 Linear

  • MoE 的专家 FFN

必须保留 BF16/FP32 的层:

  • Embedding — 需要高精度查表

  • LM Head / 输出投影 — 对精度敏感

  • Attention 的 softmax — exp 易溢出

  • LayerNorm — 统计量需要高精度

  • RoPE — 旋转精度要求

  • Optimizer 状态 — 累积更新需要高精度

  • 梯度累加 — 防止累积误差

4.2 Tile-wise 量化

DeepSeek-V3 的关键技巧:每个小块 (tile) 独立量化,而非整个矩阵用一个 scale。

def tile_wise_fp8_quant(W, tile_size=128):
    """
    将权重矩阵分为 128x128 的 tile, 每个 tile 独立量化
    减少 outlier 对整体 scale 的影响
    """
    M, N = W.shape
    W_fp8 = torch.zeros_like(W, dtype=torch.float8_e4m3fn)
    scales = torch.zeros(M // tile_size, N // tile_size)

    for i in range(0, M, tile_size):
        for j in range(0, N, tile_size):
            tile = W[i:i+tile_size, j:j+tile_size]
            scale = tile.abs().max() / 448.0  # E4M3 max
            W_fp8[i:i+tile_size, j:j+tile_size] = (tile / scale).to(torch.float8_e4m3fn)
            scales[i//tile_size, j//tile_size] = scale

    return W_fp8, scales

# 推理时:
# W_dequant = W_fp8 * scales_broadcasted
# 但矩阵乘法可以直接在 FP8 上做, 最后才 scale

4.3 FP8 训练的稳定性技巧

4.4 FP8 的收益


五、FP4 / NVFP4 探索

Blackwell GPU 原生支持 FP4 (NVFP4)。推理时 FP4 质量已经可用,训练仍在研究。 NVFP4 是 4-bit 浮点格式,目前仅用于推理(2026 主流)。训练使用 FP4 需要更复杂的量化策略,仍在研究阶段。

Blackwell FP4 性能: B200 FP4 约 9000 TFLOPS(对比 BF16 的 2250 TFLOPS),理论加速达 4x


六、不同场景的选择

flowchart TB
    start["要训 LLM"]
    q1{"硬件"}

    start --> q1

    q1 -->|"V100/T4"| fp16["FP16 + loss scaling<br/>(老硬件唯一选择)"]
    q1 -->|"A100"| bf16["BF16 mixed precision<br/>(标准选择)"]
    q1 -->|"H100/H200"| h["BF16 起步 → FP8 (E4M3 + per-tile)"]
    q1 -->|"B100/B200"| b["FP8 默认 / FP4 探索"]

    classDef stage fill:#fff,stroke:#cc785c,color:#1a1a1a;
    classDef decision fill:#f5f3eb,stroke:#bdb9ab,color:#1a1a1a;
    class start,fp16,bf16,h,b stage
    class q1 decision
场景 推荐 备注
单机 / 小模型 BF16 简单稳定
H100 1B+ FP8 (DeepSeek-V3 路线) 训练 2× 加速
推理 FP8 (E4M3) / INT4 (AWQ) 内存减半
边缘 / 端侧 INT4 GGUF 极致压缩

七、调试混合精度训练

7.1 常见问题

7.2 Loss Scaling(FP16 时代遗留)

# FP16 时代的技巧, BF16/FP8 已不需要
# 但了解历史背景有助于理解

# FP16 的梯度容易 underflow (<6e-5 会变 0)
# 解决: 将 loss 乘以大数 (loss scale)
# 反向后梯度放大, 除以相同 scale 恢复

loss_scale = 65536
with autocast(dtype=torch.float16):
    loss = model(batch) * loss_scale
loss.backward()
# 检查梯度是否 overflow
if grad_has_inf_or_nan():
    loss_scale /= 2  # 减小 scale
    optimizer.zero_grad()  # 跳过这一步
else:
    # 梯度除以 scale
    for p in params:
        p.grad /= loss_scale
    optimizer.step()

参考文献


上级 · D. 预训练 Recipe