FlashAttention-4 的下一步不只是把某个矩阵乘法换成 FP8,而是把低精度覆盖到注意力计算的完整路径:前向、反向,以及贯穿其中的块级缩放。该方案扩展了 MXFP8 支持,在 LLM 形状上报告了最高 2.85 PFLOPS 的前向性能和 2 PFLOPS 的反向性能;内部测试形状下,FA4 MX8 达到 2.54 PFLOPS。
这些数字属于特定硬件、内核实现和工作负载组合,不能直接视为所有模型都能获得的加速。但它清楚地说明了一件事:在 Blackwell 这类面向低精度计算优化的 GPU 上,注意力性能的瓶颈已经从“是否使用 FP8”转向“如何管理 FP8 的数值范围、缩放粒度和数据流”。
为什么块级缩放比单一张量缩放更重要
低精度格式的有效动态范围有限。对整个张量使用一个 scale 很简单,却容易出现两种浪费:
- 少数极大值决定 scale,导致大部分元素只能使用很少的有效精度;
- 不同 token、head 或局部矩阵的分布差异被一个全局 scale 掩盖。
块级缩放把张量切成固定大小的 block,每个 block 保存自己的缩放因子。量化时,可以近似写成:
q_block = cast(x_block / scale_block, FP8)
恢复值 = cast(q_block, FP32) * scale_block
MXFP8 的关键不只是 FP8 数据本身,还包括这些 scale 如何存储、加载、应用,以及如何与矩阵乘、归约和 softmax 融合。对注意力来说,Q、K、V、累加器和中间结果的数值角色不同,不能简单地把所有张量都改成同一种精度。
端到端低精度注意力的难点
1. QKᵀ、softmax 和 PV 的误差性质不同
QKᵀ 通常可以利用 FP8 输入与更高精度累加器完成,而 softmax 对数值范围更敏感。实际内核往往需要在局部 tile 内进行稳定化处理,例如减去当前行最大值,再执行指数和归一化。
PV 的误差路径又不同:V 的量化误差会直接影响输出,过于激进的缩放可能在长序列或稀疏分布下放大尾部误差。因此,端到端设计要同时考虑:
- 输入 Q、K、V 的块边界;
- scale 的布局和带宽;
- tile 内的高精度累加;
- softmax 的最大值与归一化和;
- 反向传播中 dQ、dK、dV 的误差累积。
2. 反向传播更容易暴露误差
训练场景不能只看前向输出的平均误差。反向计算会重复使用注意力概率、激活和梯度,局部量化误差可能在长序列、多层堆叠后变得明显。
因此,MXFP8 backward 的价值不只是减少显存读写,还在于它证明了低精度可以覆盖训练关键路径。工程上仍应保留高精度累加、必要的统计量,以及对梯度溢出和异常值的监控。
一个可运行的块缩放实验
下面的 Python 示例用 PyTorch 演示块级、二次幂 scale 的基本机制。它不是 FlashAttention-4 或 Blackwell 内核的实现,而是一个可以直接运行的数值实验:先按 block 量化,再反量化,并测量误差。需要安装 PyTorch;如果当前版本支持 torch.float8_e4m3fn,代码会尝试使用该类型,否则退化为截断后的 FP32 模拟。
import torch
def block_quantize_mxfp8(x: torch.Tensor, block_size: int = 32):
"""教育用途的块缩放示例:scale 使用二次幂,数据使用 E4M3 或 FP32 模拟。"""
if x.ndim != 2:
raise ValueError("demo expects a 2-D tensor")
rows, cols = x.shape
if cols % block_size != 0:
raise ValueError("the last dimension must be divisible by block_size")
blocks = x.reshape(rows, cols // block_size, block_size)
max_abs = blocks.abs().amax(dim=-1, keepdim=True).clamp_min(1e-12)
# 用二次幂近似 MX 风格的共享 scale,便于硬件快速处理。
scale = torch.pow(2.0, torch.ceil(torch.log2(max_abs / 448.0)))
normalized = blocks / scale
normalized = normalized.clamp(-448.0, 448.0)
fp8_type = getattr(torch, "float8_e4m3fn", None)
if fp8_type is not None:
q = normalized.to(fp8_type)
restored = q.to(torch.float32) * scale
else:
# 没有 FP8 类型时,用 FP32 截断模拟量化误差。
q = normalized.round()
restored = q * scale
return q, scale, restored.reshape_as(x)
if __name__ == "__main__":
torch.manual_seed(7)
x = torch.randn(4, 128, dtype=torch.float32) * 3.0
q, scales, x_hat = block_quantize_mxfp8(x, block_size=32)
abs_error = (x - x_hat).abs()
rel_error = abs_error.mean() / x.abs().mean().clamp_min(1e-12)
print("quantized shape:", tuple(q.shape))
print("scale shape:", tuple(scales.shape))
print("mean absolute error:", abs_error.mean().item())
print("mean relative error:", rel_error.item())
这个实验适合用来比较 block size、激活分布和异常值处理策略。要把它改造成真实内核,需要进一步明确硬件支持的 FP8 编码、scale 元数据布局、tile 尺寸和累加精度,不能把这段 Python 代码当作生产级 FlashAttention 实现。
如何验证性能,而不是只看理论峰值
PFLOPS 是有用的汇总指标,但注意力内核通常受形状和带宽影响很大。建议至少覆盖以下维度:
- batch size、序列长度和 head dimension;
- prefill 与 decode;
- 训练前向、训练反向和推理;
- causal 与非 causal mask;
- 不同 block size 和 scale 布局;
- FP16/BF16、FP8 和 MXFP8 的输出误差与吞吐。
可以先用项目提供的 benchmark 入口记录基线。下面是一个可改造的命令模板,参数名需要按具体仓库调整:
# 示例:把路径、GPU 数量和 shape 改成实际 benchmark 的参数
CUDA_VISIBLE_DEVICES=0 python benchmark_attention.py \
--dtype mxfp8 \
--batch-size 8 \
--seqlen 4096 \
--num-heads 32 \
--head-dim 128 \
--backward \
--warmup 100 \
--iters 500
性能分析时不要只记录 kernel 的峰值时间。还应检查 scale 的加载是否成为额外开销、反向是否产生更多中间张量、以及量化后是否需要回退到高精度路径。一个“理论吞吐更高但频繁回退”的实现,端到端训练速度可能并不占优。
采用建议与边界
如果目标平台是 Blackwell,并且工作负载主要是大规模 Transformer,MXFP8 attention 值得作为独立实验路径评估。比较稳妥的落地顺序是:
- 用 BF16 或 FP16 结果建立准确性基线;
- 只替换前向,检查 logits、attention 输出和端到端模型指标;
- 再启用反向,观察梯度范数、loss 曲线和溢出情况;
- 按真实 LLM 形状测量吞吐、显存带宽和 step time;
- 为异常分布、短序列和不支持的形状保留高精度回退。
FA4 MXFP8 的报告性能展示了块缩放注意力的潜力,但峰值数字不是唯一目标。真正可用的方案应在数值稳定性、内核覆盖率、scale 元数据成本和训练收敛之间取得平衡。对开发者来说,最值得关注的变化是:低精度注意力正在从“单个 GEMM 的优化”演变为“覆盖整个注意力数据流的系统设计”。