FP8深度解析:深度学习的"压缩饼干"

一、为什么需要FP8?从比特的角度看成本

1.1 深度学习的"数据饥渴症"

# 一个真实的例子:训练GPT-3

model_params = 175_000_000_000  # 1750亿参数

# 如果用FP32(传统精度)
fp32_size = model_params * 4  # 每个参数4字节
print(f"FP32模型大小: {fp32_size / 1e9:.0f} GB")
# 输出: 700 GB

# 如果用FP16(半精度)
fp16_size = model_params * 2  # 每个参数2字节
print(f"FP16模型大小: {fp16_size / 1e9:.0f} GB")
# 输出: 350 GB

# 如果用FP8(超低精度)
fp8_size = model_params * 1  # 每个参数1字节
print(f"FP8模型大小: {fp8_size / 1e9:.0f} GB")
# 输出: 175 GB

# 节省一半存储!

问题来了:把数据压缩一半,精度会不会崩掉?

答案是:只要设计得当,影响很小!


二、浮点数的秘密:符号、指数、尾数

2.1 浮点数的"三件套"

所有浮点数(FP32/FP16/FP8)都由三部分组成:

一个浮点数 = 符号位 × 2^指数 × (1 + 尾数)

┌────────────────────────────────────┐
│  -3.75 如何表示?                   │
├────────────────────────────────────┤
│  符号位(S): 1 (负数)                │
│  指数(E):   1 (表示2^1 = 2)         │
│  尾数(M):   0.875 (表示1.875)       │
│                                    │
│  计算: -1 × 2^1 × 1.875 = -3.75   │
└────────────────────────────────────┘

类比理解:浮点数就像科学计数法

-3750000 = -3.75 × 10^6
           ↑      ↑
         尾数    指数

浮点数:
-3.75 = -1.875 × 2^1
        ↑       ↑
       尾数    指数

2.2 指数和尾数的"跷跷板效应"

总共的比特是固定的,你要在范围精度之间做选择:

┌─────────────────────────────────────────────┐
│  场景1:拍摄风景照                           │
├─────────────────────────────────────────────┤
│  需求:范围大(从近景到远山)                │
│  不需要:极致细节                            │
│  选择:多给指数位 → E5M2                     │
└─────────────────────────────────────────────┘

┌─────────────────────────────────────────────┐
│  场景2:拍摄微距照片                         │
├─────────────────────────────────────────────┤
│  需求:精度高(捕捉花蕊纹理)                │
│  不需要:太大范围                            │
│  选择:多给尾数位 → E4M3                     │
└─────────────────────────────────────────────┘

三、FP8的两兄弟:E4M3 vs E5M2

3.1 E4M3:精度优先

格式:1位符号 + 4位指数 + 3位尾数

┌───┬───────────┬─────────┐
│ S │    E (4)  │  M (3)  │
└───┴───────────┴─────────┘
 1位    4位        3位

特点:
✓ 尾数多(3位)→ 精度高
✓ 可以表示更细腻的数值
✗ 指数少(4位)→ 范围小

数值范围:-448 到 +448
精度:约 ±0.1 的误差(在1.0附近)

适用场景:
- 神经网络的权重(通常在[-10, 10]范围内)
- 激活值(ReLU后都是正数,范围不大)
- Attention scores(softmax后在[0,1])

通俗比喻

E4M3就像一把精密刻度尺

  • 只能量0-50厘米(范围小)
  • 但刻度精确到0.1毫米(精度高)
  • 适合测量手机、书本这种小物件

3.2 E5M2:范围优先

格式:1位符号 + 5位指数 + 2位尾数

┌───┬─────────────┬───────┐
│ S │    E (5)    │ M (2) │
└───┴─────────────┴───────┘
 1位     5位        2位

特点:
✓ 指数多(5位)→ 范围大
✓ 可以表示超大/超小的数
✗ 尾数少(2位)→ 精度低

数值范围:-57344 到 +57344
精度:约 ±1 的误差(在1.0附近)

适用场景:
- 梯度(可能非常大或非常小)
- 损失值(初期可能爆炸)
- 中间激活值(某些层可能值很大)

通俗比喻

E5M2就像一把卷尺

  • 可以量0-100米(范围大)
  • 但刻度只精确到1厘米(精度低)
  • 适合测量房间、操场这种大物体

3.3 直观对比

# 用生活中的温度计类比

# E4M3 = 医用体温计
e4m3_thermometer = {
    "range": "35°C - 42°C",      # 范围小
    "precision": "0.1°C",         # 精度高
    "use_case": "测量人体体温",
}

# E5M2 = 工业温度计
e5m2_thermometer = {
    "range": "-200°C - 1000°C",  # 范围大
    "precision": "5°C",           # 精度低
    "use_case": "测量炼钢炉温度",
}

# 选哪个?
if 你要测人体体温:
    用E4M3  # 范围够用,精度重要
elif 你要测工业设备:
    用E5M2  # 范围重要,精度无所谓

四、Scaling:FP8的"变焦镜头"

4.1 问题:FP8范围太小怎么办?

# 真实场景:某层激活值范围

activation_values = [0.001, 0.005, 0.008, 0.012, ...]  # 都很小

# 如果直接用FP8存储
fp8_max = 448  # E4M3的最大值
fp8_min = 1.0 / 512  # E4M3的最小正数 ≈ 0.002

# 问题:0.001 < 0.002,存不下!会被截断为0
# 结果:丢失信息

解决方案:Scaling(缩放)

核心思想:把数值范围"变焦"到FP8能表示的范围

原始数据:[0.001, 0.005, 0.008]
            ↓ 放大1000倍
缩放后:  [1.0, 5.0, 8.0]  ← FP8能表示了!
            ↓ 存储为FP8
使用时:    ↓ 缩小1000倍
恢复:    [0.001, 0.005, 0.008]

4.2 Scaling的工作原理

# === Per-Tensor Scaling(逐张量缩放)===

def per_tensor_scaling(tensor):
    """
    整个tensor用同一个scale
    """
    # 步骤1:找到tensor的最大绝对值
    max_val = tensor.abs().max()
  
    # 步骤2:计算scale,使最大值映射到FP8的最大值
    fp8_max = 448  # E4M3的最大值
    scale = fp8_max / max_val
  
    # 步骤3:缩放并转换为FP8
    scaled_tensor = tensor * scale
    fp8_tensor = scaled_tensor.to(torch.float8_e4m3fn)
  
    # 步骤4:记住scale(后面恢复用)
    return fp8_tensor, scale


def dequantize(fp8_tensor, scale):
    """
    恢复原始精度
    """
    # 转回FP16/FP32
    fp16_tensor = fp8_tensor.to(torch.float16)
  
    # 除以scale恢复原始范围
    original_tensor = fp16_tensor / scale
  
    return original_tensor


# === 使用示例 ===
import torch

# 原始数据:范围[0, 0.01]
original = torch.tensor([0.001, 0.005, 0.008, 0.010])

print("原始值:", original)
# 输出: tensor([0.0010, 0.0050, 0.0080, 0.0100])

# 转换为FP8
fp8, scale = per_tensor_scaling(original)
print(f"Scale: {scale:.1f}")
# 输出: Scale: 44800.0

print("FP8存储:", fp8)
# 内部存储为: [44.8, 224.0, 358.4, 448.0] (理论值)

# 恢复
recovered = dequantize(fp8, scale)
print("恢复值:", recovered)
# 输出: tensor([0.0010, 0.0050, 0.0080, 0.0100])

# 误差
error = (recovered - original).abs().max()
print(f"最大误差: {error:.6f}")
# 输出: 最大误差: 0.000012 (非常小!)

4.3 两种Scaling策略

Per-Tensor Scaling(逐张量)
┌──────────────────────────────────────┐
│  整个tensor用一个scale                │
├──────────────────────────────────────┤
│  Tensor: [0.001, 0.5, 100.0, 0.002]  │
│           ↓                          │
│  Max = 100.0                         │
│  Scale = 448 / 100 = 4.48            │
│           ↓                          │
│  Scaled: [0.0048, 2.24, 448, 0.009]  │
└──────────────────────────────────────┘

优点:
✓ 简单,只需存1个scale值
✓ 计算快

缺点:
✗ 如果数据分布不均(有极端值),小的值精度很差
  例如:0.001被放大后还是很小,量化误差大
Per-Channel Scaling(逐通道)
┌────────────────────────────────────────┐
│  每个通道(channel)用独立的scale        │
├────────────────────────────────────────┤
│  Channel 0: [0.001, 0.002, 0.003]      │
│  → Scale0 = 448/0.003 = 149333         │
│                                        │
│  Channel 1: [10.0, 20.0, 30.0]         │
│  → Scale1 = 448/30 = 14.9              │
│                                        │
│  Channel 2: [100, 200, 300]            │
│  → Scale2 = 448/300 = 1.49             │
└────────────────────────────────────────┘

优点:
✓ 每个通道根据自己的范围调整
✓ 精度更高,尤其是不同通道范围差异大时

缺点:
✗ 需要存储多个scale(例如512个通道 = 512个scale)
✗ 计算稍慢

选哪个?

# 决策树

if 权重矩阵:
    # 权重通常每个通道分布差异大
    use_per_channel_scaling()
  
elif 激活值:
    # 激活值通常分布较均匀
    if 内存不紧张:
        use_per_channel_scaling()  # 更好的精度
    else:
        use_per_tensor_scaling()   # 更省内存
      
elif 梯度:
    # 梯度范围变化大,用per-tensor就够
    use_per_tensor_scaling()

五、FP8在神经网络中的实战

5.1 混合精度策略

神经网络的不同部分,对精度的需求不同:

┌─────────────────────────────────────────┐
│  FP32 (最高精度)                         │
├─────────────────────────────────────────┤
│  • 损失计算 (loss)                       │
│  • 优化器状态 (Adam的momentum)           │
│  • BatchNorm的running stats             │
│                                         │
│  为什么?误差累积会影响训练稳定性         │
└─────────────────────────────────────────┘

┌─────────────────────────────────────────┐
│  FP16 (中等精度)                         │
├─────────────────────────────────────────┤
│  • 主要计算 (矩阵乘法)                   │
│  • 梯度(部分)                          │
│                                         │
│  为什么?平衡速度和精度,成熟方案         │
└─────────────────────────────────────────┘

┌─────────────────────────────────────────┐
│  FP8-E4M3 (低精度-高精度)                │
├─────────────────────────────────────────┤
│  • 权重 (存储和前向传播)                 │
│  • 激活值 (层之间传递)                   │
│                                         │
│  为什么?范围固定,精度要求高             │
└─────────────────────────────────────────┘

┌─────────────────────────────────────────┐
│  FP8-E5M2 (低精度-大范围)                │
├─────────────────────────────────────────┤
│  • 梯度 (反向传播)                       │
│  • 某些激活函数输出 (如GELU)             │
│                                         │
│  为什么?梯度可能爆炸/消失,需要大范围     │
└─────────────────────────────────────────┘

5.2 一个完整的训练迭代

训练一个batch的数据流和精度转换:

T0: 输入数据
    [FP16] Input Tensor
      ↓
T1: 前向传播 - Layer 1
    [FP16] → Linear(weight_fp8_e4m3) → [FP8-E4M3] → ReLU → [FP16]
                    ↑ dequantize         ↑ quantize
                  
T2: 前向传播 - Layer 2
    [FP16] → Linear(weight_fp8_e4m3) → [FP8-E4M3] → ...
  
T3: 计算损失
    [FP16] prediction → [FP32] loss.item()
                        ↑ 高精度避免下溢
                      
T4: 反向传播
    [FP32] loss → [FP8-E5M2] gradients → [FP16] weight.grad
                  ↑ 大范围适应梯度爆炸
                
T5: 优化器更新
    [FP32] optimizer_state + [FP16] grad → [FP16] new_weight
    ↑ 高精度累积小更新

关键技巧在精度损失最小的地方用FP8

# 原则:
# 1. 瓶颈在哪里?→ 优化哪里
# 2. 能容忍多少误差?→ 选择精度

bottleneck_analysis = {
    "内存瓶颈": {
        "位置": "存储大模型权重",
        "方案": "权重用FP8-E4M3存储,计算时dequantize",
        "收益": "模型大小减半,显存占用减半"
    },
  
    "带宽瓶颈": {
        "位置": "GPU间通信(多卡训练)",
        "方案": "梯度用FP8-E5M2传输",
        "收益": "通信量减半,训练速度提升30%+"
    },
  
    "计算瓶颈": {
        "位置": "Transformer的FFN层",
        "方案": "激活值用FP8-E4M3传递",
        "收益": "激活值存储减半,大batch size可用"
    }
}

六、FP8的"副作用"与应对

6.1 量化噪声

问题:FP8的精度有限,引入噪声

┌─────────────────────────────────────┐
│  原始值 (FP16): 3.14159             │
│  FP8表示:       3.125 ≈ 3.14159    │
│  误差:          0.01659              │
└─────────────────────────────────────┘

累积效应:
Layer 1: 误差 0.016
Layer 2: 误差 0.016 → 累积 0.032
Layer 3: 误差 0.016 → 累积 0.048
...
Layer 96: 累积误差 1.5 (Transformer-96层)

后果:
• 训练不稳定(loss震荡)
• 模型精度下降(准确率降低1-2%)

解决方案:Residual Connection + LayerNorm

# Transformer的标准做法

class TransformerBlock:
    def forward(self, x):
        # FP8计算,有误差
        attention_out = self.attention(x)  # 可能有噪声
      
        # 残差连接:原始信号保留
        x = x + attention_out  # ← 关键!高精度的x"修正"低精度的输出
      
        # LayerNorm:归一化,减少误差累积
        x = self.layer_norm(x)
      
        return x

# 为什么有效?
# 残差连接相当于"高通滤波器",过滤掉累积的直流误差
# LayerNorm把分布拉回正常范围,防止误差放大

6.2 训练不稳定

症状:使用FP8后,训练loss曲线抖动剧烈

原因分析:

1. 梯度下溢
   小梯度 (1e-5) 在FP8中被截断为0
   → 某些参数无法更新
 
2. 梯度爆炸
   大梯度 (1e5) 超出FP8范围
   → 被clip成最大值,方向错误
 
3. Scale不匹配
   不同层的数值范围差异大,统一scale不合适

解决方案组合拳

# 1. 动态Scale调整
class DynamicScaler:
    def __init__(self):
        self.scale_history = []
  
    def adjust_scale(self, grad):
        # 每N步重新计算scale
        if step % 100 == 0:
            max_grad = grad.abs().max()
            self.scale = 448 / max_grad
          
        return grad * self.scale

# 2. 梯度裁剪
max_grad_norm = 1.0
torch.nn.utils.clip_grad_norm_(model.parameters(), max_grad_norm)

# 3. Warmup阶段用高精度
if epoch < warmup_epochs:
    use_fp16()  # 前几轮用FP16稳定训练
else:
    use_fp8()   # 后续用FP8加速

# 4. 混合:关键层用FP16,其他用FP8
model.attention.use_fp16 = True   # Attention精度敏感
model.ffn.use_fp8 = True          # FFN可以容忍低精度

七、FP8的硬件支持

7.1 为什么H100特别适合FP8?

传统GPU (A100):
┌──────────────────────────────────────┐
│  Tensor Core (FP16/BF16)             │
│  • 原生支持FP16矩阵乘法               │
│  • FP8需要软件模拟 (慢10倍)           │
└──────────────────────────────────────┘

Hopper GPU (H100):
┌──────────────────────────────────────┐
│  Transformer Engine                  │
│  ├─ FP8 Tensor Core (硬件原生)       │
│  │  • E4M3矩阵乘法                   │
│  │  • E5M2矩阵乘法                   │
│  ├─ 自动缩放单元                      │
│  │  • 硬件自动计算scale               │
│  │  • 无需CPU干预                    │
│  └─ 混合精度调度器                    │
│     • 自动选择最优精度                │
└──────────────────────────────────────┘

性能对比:
A100 (FP16): 312 TFLOPS
H100 (FP8):  3958 TFLOPS (12倍算力!)

7.2 Transformer Engine:FP8的"自动驾驶"

# 传统做法:手动管理精度和scaling
def manual_fp8_training(model, data):
    # 程序员要手动做很多事
    x_fp8, scale1 = quantize_to_fp8(data)
  
    weight_fp8, scale2 = quantize_to_fp8(model.weight)
  
    output = matmul_fp8(x_fp8, weight_fp8)
  
    output = dequantize(output, scale1 * scale2)
  
    # 痛苦:每一层都要这样写!

# Transformer Engine:全自动
import transformer_engine.pytorch as te

def auto_fp8_training(model, data):
    # 只需要把层替换为TE版本
    model.linear = te.Linear(in_features, out_features, ...)
  
    # 前向传播,TE自动处理一切
    output = model(data)
  
    # TE自动:
    # 1. 监控数值范围
    # 2. 动态调整scale
    # 3. 选择E4M3还是E5M2
    # 4. 在合适的地方转换精度
  
    # 程序员:无感知,就像用FP16一样简单!

八、实战建议

8.1 什么时候用FP8?

# 决策流程图

if GPU == "H100" or GPU == "H200":
    # 有硬件支持,强烈推荐
  
    if 任务 == "训练大模型" and 参数量 > 10B:
        use_fp8()  # ✓ 显存和速度都受益
      
    elif 任务 == "推理" and QPS要求高:
        use_fp8()  # ✓ 吞吐量提升巨大
      
    elif 任务 == "多卡训练" and 卡数 > 8:
        use_fp8_for_gradients()  # ✓ 通信瓶颈消除
      
    else:
        evaluate_first()  # ⚠ 先测试精度损失
      
else:  # A100, V100, etc.
    if 显存非常紧张:
        use_fp8_for_storage_only()  # ⚠ 仅存储用FP8,计算用FP16
    else:
        stick_to_fp16()  # ✗ 软件模拟FP8太慢,不值得

8.2 FP8使用的"黄金法则"

法则1:渐进式引入
❌ 错误:一次性全部换成FP8
✅ 正确:先换FFN层 → 再换Attention → 最后尝试梯度

法则2:先测精度,再上生产
❌ 错误:看到加速就直接部署
✅ 正确:在验证集上对比FP16和FP8的指标差异

法则3:Per-Channel > Per-Tensor
对于权重矩阵,多花点内存存scale,值得

法则4:关键路径用高精度
Loss计算、Softmax、LayerNorm:坚持用FP32/FP16

法则5:监控数值健康
定期检查:
• 激活值的范围(是否超出FP8范围)
• 梯度的范围(是否下溢/上溢)
• Scale的变化(是否稳定)

8.3 常见陷阱

# 陷阱1:忘记dequantize
x_fp8 = quantize(x, scale)
y_fp8 = model(x_fp8)
loss = criterion(y_fp8, target)  # ❌ y_fp8还没还原!

# 正确:
y_fp16 = dequantize(y_fp8, scale)
loss = criterion(y_fp16, target)  # ✓

# ────────────────────────────────────

# 陷阱2:Scale过期
scale = compute_scale(epoch=0)
for epoch in range(100):
    train(model, scale)  # ❌ scale应该动态更新!

# 正确:
for epoch in range(100):
    scale = compute_scale(epoch)  # ✓ 每轮重新计算
    train(model, scale)

# ────────────────────────────────────

# 陷阱3:混用不同scale
x_fp8 = quantize(x, scale_x)
w_fp8 = quantize(w, scale_w)
y_fp8 = matmul(x_fp8, w_fp8)
y = dequantize(y_fp8, scale_x)  # ❌ 应该是scale_x * scale_w!

# 正确:
y = dequantize(y_fp8, scale_x * scale_w)  # ✓

九、未来展望

FP8只是开始

当前 (2024):
FP8 (8-bit) - H100原生支持

近期 (2025-2026):
FP6 (6-bit) - 研究中
FP4 (4-bit) - 推理已可用 (GPTQ, AWQ)

未来 (2027+):
• 动态精度:根据层的重要性自动选择2-16位
• 混合格式:同一层内不同参数用不同精度
• 学习式量化:AI自动学习最优的量化策略

终极目标:
1-bit神经网络 (Binary Neural Networks)
理论上,1750亿参数只需要 21 GB!

对行业的影响

成本维度:
• 训练一个GPT-4级别模型:
  FP16: 需要25000张H100,3个月
  FP8:  需要12000张H100,1.5个月
  节省: 约5000万美元

能耗维度:
• 数据中心电费:
  FP8减少40%计算量 → 减少40%电费
  GPT-4规模训练:节省约1000万美元电费

环境维度:
• 碳排放:
  FP8减少数百吨CO2排放(单次训练)

十、总结

核心要点速记

fp8_summary = {
    "是什么": "8位浮点数,符号1位+指数4-5位+尾数2-3位",
  
    "为什么": "深度学习对精度不敏感,用更少比特可以大幅节省资源",
  
    "两种格式": {
        "E4M3": "精度高,范围小,适合权重和激活",
        "E5M2": "范围大,精度低,适合梯度"
    },
  
    "Scaling": "把数据缩放到FP8能表示的范围,用完再缩回来",
  
    "适用场景": "大模型训练、高吞吐推理、多卡通信",
  
    "硬件要求": "H100最佳,A100勉强(软件模拟),V100不推荐",
  
    "精度损失": "通常<1%,关键是设计得当(混合精度+动态scaling)",
  
    "性能收益": "内存减半、带宽减半、H100上算力12倍",
}

一句话总结

FP8就像深度学习的"压缩饼干":体积减半(8 bit),营养保留(精度损失<1%),只是需要"泡水还原"(scaling),但在H100这个"热水壶"里,这个过程快到无感知。


延伸阅读

  • NVIDIA Transformer Engine文档
  • FP8 Training论文 (Micikevicius et al., 2022)
  • H100 Whitepaper

Scale的通俗解释:像调望远镜一样

一、核心比喻

Scale就像望远镜的变焦

场景:你要用FP8相机(只能拍0-448米的物体)拍摄一只1毫米的蚂蚁

问题:蚂蚁太小了,FP8拍不清楚!

解决:
1. 用放大镜把蚂蚁放大1000倍 → 变成1米 (scaling)
2. FP8相机拍摄放大后的蚂蚁 (计算)
3. 照片缩小1000倍还原真实大小 (descaling)

二、什么时候用Scale?

规则很简单:数据超出FP8范围时就要scale

# 场景1:数据太小
activations = [0.0001, 0.0005, 0.0008]  # 太小了!
fp8_min = 0.002  # FP8能表示的最小正数

# 不scale会怎样?
# → 全部变成0,信息丢失!

# 用scale:
scale = 1000  # 放大1000倍
scaled = [0.1, 0.5, 0.8]  # 现在FP8能表示了!
# 场景2:数据太大
gradients = [500, 1000, 2000]  # 太大了!
fp8_max = 448  # E4M3的最大值

# 不scale会怎样?
# → 超出部分被截断成448,梯度错误!

# 用scale:
scale = 0.2  # 缩小5倍
scaled = [100, 200, 400]  # 现在FP8能表示了!

简单判断

if abs(data) < 0.01 or abs(data) > 100:
    需要scale  # ✓
else:
    可能不需要  # 大部分情况下FP8能直接表示

三、Scale的完整流程

就三步,记住口诀:缩放 → 计算 → 还原

# === 步骤1:缩放(进入FP8世界)===
x = [0.001, 0.005, 0.010]  # 原始数据,太小

scale_x = 448 / max(x)  # 算一个scale
# scale_x = 448 / 0.01 = 44800

x_fp8 = (x * scale_x).to_fp8()  # 放大并转FP8
# x_fp8 = [44.8, 224, 448] (FP8表示)

# ✓ 存下scale_x,后面要用!


# === 步骤2:计算(FP8世界里算)===
# 所有计算都在FP8进行,快且省内存

w = [0.5, 1.0, 1.5]
scale_w = 448 / max(w)
# scale_w = 448 / 1.5 = 298.67

w_fp8 = (w * scale_w).to_fp8()

# 矩阵乘法(FP8计算)
y_fp8 = x_fp8 @ w_fp8  # ← 核心计算在这里!
# 注意:y_fp8的真实值 = (x*scale_x) @ (w*scale_w)
#               = x @ w * (scale_x * scale_w)


# === 步骤3:还原(回到真实世界)===
y = y_fp8.to_fp16() / (scale_x * scale_w)  # ← 除以两个scale
#   ↑ 转回高精度    ↑ 把放大的倍数除回去

# 现在y就是真实的计算结果了!

四、关键问题解答

Q1: Scale之后再计算,还是计算之后再Scale?

答案:Scale之后再计算!

# ❌ 错误顺序:先算再scale
y = x @ w  # FP16计算,慢
y_fp8 = quantize(y, scale)  # 仅仅是存储用FP8,没加速

# ✓ 正确顺序:先scale再算
x_fp8 = quantize(x, scale_x)
w_fp8 = quantize(w, scale_w)
y_fp8 = x_fp8 @ w_fp8  # ← FP8计算,快!
y = dequantize(y_fp8, scale_x * scale_w)

原因:只有在FP8里计算,才能用上H100的FP8加速器(12倍算力)


Q2: 每次计算都要Scale吗?

不是!看情况:

# 场景A:权重(每层scale一次)
class LinearLayer:
    def __init__(self, weight):
        # 初始化时scale一次
        self.scale = compute_scale(weight)
        self.weight_fp8 = quantize(weight, self.scale)
        # ✓ 权重不变,scale也不变
  
    def forward(self, x):
        # 前向传播时直接用
        y = x @ self.weight_fp8
        return dequantize(y, self.scale)


# 场景B:激活值(每次都要scale)
def attention_layer(q, k, v):
    # 每次输入都不同,需要重新算scale
    scale_q = compute_scale(q)  # ← 每次都算
    q_fp8 = quantize(q, scale_q)
  
    scale_k = compute_scale(k)
    k_fp8 = quantize(k, scale_k)
  
    # 计算
    scores_fp8 = q_fp8 @ k_fp8.T
  
    # 还原
    scores = dequantize(scores_fp8, scale_q * scale_k)
    return scores

规则

  • 权重:初始化时scale,训练中固定 → 一次性
  • 激活值:每次输入不同 → 每次都scale
  • 梯度:每轮不同 → 每次都scale

Q3: Scale会影响计算结果吗?

理论上不影响,实际上有微小误差:

# 理论:
x = 0.001
scale = 1000
x_scaled = 0.001 * 1000 = 1.0

# FP8存储:1.0 → FP8表示还是1.0(精确)

# 还原:1.0 / 1000 = 0.001  # ✓ 完全一致


# ──────────────────────────────────

# 实际:
x = 0.00123456  # 很多小数位
scale = 1000
x_scaled = 1.23456

# FP8存储:1.23456 → 1.234(FP8精度有限)

# 还原:1.234 / 1000 = 0.001234  # ⚠️ 有误差
# 误差:0.001234 - 0.00123456 = 0.00000044  # 很小!

结论:误差主要来自FP8的精度限制,不是Scale本身的问题。


五、生活化类比总结

FP8 = 一把只能量0-50厘米的尺子

场景1:量蚂蚁(2毫米)
问题:尺子最小刻度是1毫米,量不准
方案:用放大镜放大50倍 → 变成10厘米 → 量完再除以50
这就是Scale!

场景2:量房间(8米)
问题:尺子只有50厘米,超出范围
方案:缩小到模型尺寸(1:20) → 变成40厘米 → 量完再乘以20
这也是Scale!

场景3:量书本(25厘米)
问题:刚好在范围内
方案:直接量,不需要Scale

六、记忆口诀

太大太小用Scale,
缩放计算再还原。
权重Scale一次搞定,
激活梯度次次算。

核心公式

# 万能公式
真实结果 = FP8计算结果 / (scale_1 * scale_2 * ... * scale_n)
         ↑                  ↑
      descale          所有参与计算的数据的scale乘积

就这么简单! 🎯

后记

2026年8月17日于上海,在claude opus 4.8辅助下完成。

Logo

有“AI”的1024 = 2048,欢迎大家加入2048 AI社区

更多推荐