FP8深度解析:深度学习的“压缩饼干“
·
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辅助下完成。
更多推荐



所有评论(0)