Flax FP8 量化实战指南:从fp8_ops低层 API 到Fp8DotGeneral/Fp8Einsum的完整实现解析
【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax
导读:本文以 Flax 官方量化指南 docs/guides/quantization/fp8_basics.md 为核心骨架,结合仓库源码 flax/linen/fp8_ops.py 与测试 tests/linen/linen_test.py 深入展开,系统讲解 JAX/Flax 中 FP8 量化的完整技术栈:FP8 数据类型与量化/反量化原理、当前缩放与延迟缩放两种缩放策略、Flax 低层
fp8_ops函数式 API、高层Fp8DotGeneral/Fp8Einsum即插即用模块,以及 FP8 专用参数(_overwrite_with_gradient)的更新、累积与调试方法。读完本文,你将掌握如何用 Flax 在 NVIDIA Hopper 及更新 GPU 上为 Dense 层、Einsum 运算(如 MoE 层)接入 FP8 训练,并理解其底层实现原理。
1. FP8 量化基础:为什么需要 Q/DQ
JAX 原生支持多种 FP8 格式,Flax 指南中主要涉及两类:
| 格式 | JAX dtype | 指数位/尾数位 | 最大可表示值 | 典型用途 |
|---|---|---|---|---|
| E4M3 | jnp.float8_e4m3fn | 4/3 | 448 | 前向传播中的激活值与权重 |
| E5M2 | jnp.float8_e5m2 | 5/2 | 57344 | 反向传播中的梯度 |
FP8 数据类型的取值范围非常有限,无法直接容纳完整精度的浮点张量。因此,在使用前必须将高精度数据缩放(scale)到 FP8 可表示范围内,这一过程称为量化(Quantization,Q);反之,将 FP8 数据还原回原始类型的过程称为反量化(De-Quantization,DQ)。核心关系为:
量化: x_fp8 ≈ clip(x_fp32 / scale) 反量化: x_fp32 ≈ x_fp8 × scale其中scale是缩放因子,通常按scale = amax(x) / MAX计算,amax表示取张量绝对值的最大值,MAX是目标 dtype(如 E4M3)的最大可表示值。
需要特别强调的是,本文所述内容依赖 XLA-FP8 特性,该特性仅在 NVIDIA Hopper 架构(compute capability 9.0)及更新的 GPU 上受支持。指南中的环境检测代码如下:
import flax import jax import re import pprint from jax import random from jax import numpy as jnp from jax._src import test_util as jtu from flax import linen as nn from flax.linen import fp8_ops e4m3 = jnp.float8_e4m3fn f32 = jnp.float32 E4M3_MAX = jnp.finfo(e4m3).max.astype(f32) # 断言当前 CUDA 计算能力不低于 9.0(Hopper),否则后续运行无意义 assert jtu.is_cuda_compute_capability_at_least("9.0")为了验证 FP8 路径是否真正生效,指南还定义了一个辅助函数,通过编译后的 HLO 文本中是否出现custom-call(f8e4m3fn..., f8e4m3fn...)来判断是否真正调用了 FP8 矩阵乘核函数:
def check_fp8_call(lowered): hlo = lowered.compile() if re.search(r"custom-call\(f8e4m3fn.*, f8e4m3fn.*", hlo.as_text()): print("Fp8 call detected!") else: print("No Fp8 call!")这个函数贯穿全文所有示例,是验证"FP8 是否真正被 XLA 采用"的最直接手段。
2. Flax 低层 API:手写 FP8 矩阵乘法
2.1 直接使用jnp.dot的两个局限
jnp.dot本身就支持 FP8 dtype 输入,因此下面的调用是合法的:
k0, k1 = random.split(random.key(0), 2) a = random.uniform(k0, (16, 32)) b = random.uniform(k1, (32, 64)) @jax.jit def dot_fp8(a, b): return jnp.dot(a.astype(e4m3), b.astype(e4m3), preferred_element_type=f32) check_fp8_call(dot_fp8.lower(a, b))但指南明确指出这种直接用法存在两个关键局限:
- 不支持自定义缩放因子:
jnp.dot对操作数默认使用 scale=1.0,无法针对每个张量的动态范围做缩放; - 自动微分不智能:
jnp.dot的反向传播不会自动按照推荐实践——即梯度使用 E5M2、激活/权重使用 E4M3——进行量化。
要克服这些局限、实现"正确"的 FP8 矩阵乘法,需要使用 Flax 提供的flax.linen.fp8_ops模块(源码)。
2.2 当前缩放(Current Scaling)
缩放因子按scale = amax(x) / MAX直接从当前操作数张量推导。代码示例如下:
@jax.jit def dot_fp8(a, b): a_scale = jnp.max(jnp.abs(a)) / E4M3_MAX b_scale = jnp.max(jnp.abs(b)) / E4M3_MAX a = fp8_ops.quantize(a, e4m3, a_scale, f32) b = fp8_ops.quantize(b, e4m3, b_scale, f32) c = jnp.dot(a, b, preferred_element_type=f32) c = fp8_ops.dequantize(c, f32, a_scale * b_scale) return c c = dot_fp8(a, b) check_fp8_call(dot_fp8.lower(a, b))工作流程分为三步:
- 量化:对两个操作数分别计算缩放因子,再通过
fp8_ops.quantize量化为 E4M3; - 高精度累加:
jnp.dot以preferred_element_type=f32完成矩阵乘,累加过程保持 FP32 精度; - 反量化:将结果乘以
a_scale * b_scale还原到原始数值范围(fp8_ops.dequantize)。
对照源码 fp8_ops.py 可看到quantize/dequantize的底层实现:
def quantize(x, q_dtype, scale, compute_dtype): dtype_max = get_fp8_max(q_dtype, compute_dtype) scaled_x = x / jnp.broadcast_to(scale.astype(compute_dtype), x.shape) clipped_x = jnp.clip(scaled_x, -dtype_max, dtype_max) # 先缩放,再截断到 FP8 范围 return clipped_x.astype(q_dtype) def dequantize(x, dq_dtype, scale): return x.astype(dq_dtype) * jnp.broadcast_to(scale.astype(dq_dtype), x.shape)可见quantize内部包含"除以缩放因子 → 截断到 ±MAX → 转换为 FP8"三个步骤,其中get_fp8_max支持 E4M3、E5M2 及对应的 FN 变体(见 fp8_ops.py)。指南也指出:虽然示例中两个输入都用 E4M3,但完全可以使用不同的 FP8 dtype(如一侧 E4M3、一侧 E5M2),量化方法和缩放因子的计算也可以按应用需求定制。
当前缩放的主要缺陷是计算a_scale和b_scale需要额外加载操作数张量,带来明显的性能开销。因此,指南推荐使用延迟缩放。
2.3 延迟缩放(Delayed Scaling)
延迟缩放的核心思想:缩放因子不再从当前张量实时计算,而是与一个 amax 历史(amax history)绑定。amax history 是一个存储最近若干步(例如 1024 步)amax 值的列表;缩放因子和 amax history 都从上一步计算而来,并作为模型参数维护(在 Flax 中存放于专门的变量集合,见第 5 节)。
延迟缩放的量化/反量化操作由fp8_ops.in_q和fp8_ops.out_dq提供:
fp8_ops.in_q:负责输入量化,同时更新 amax history 和缩放因子;fp8_ops.out_dq:负责输出反量化。
a_scale = jnp.array(1.0) b_scale = jnp.array(1.0) a_amax_hist = jnp.zeros((1024,)) b_amax_hist = jnp.zeros((1024,)) @jax.jit def dot_fp8(a, a_scale, a_amax_hist, b, b_scale, b_amax_hist): a, a_scale = fp8_ops.in_q(f32, e4m3, a, a_scale, a_amax_hist) b, b_scale = fp8_ops.in_q(f32, e4m3, b, b_scale, b_amax_hist) c = jnp.dot(a, b, preferred_element_type=f32) c = fp8_ops.out_dq(f32, a_scale, b_scale, c) return c c = dot_fp8(a, a_scale, a_amax_hist, b, b_scale, b_amax_hist) check_fp8_call(dot_fp8.lower(a, a_scale, a_amax_hist, b, b_scale, b_amax_hist))示例中先准备多组缩放因子与 amax history,视作上一步的计算结果;随后对jnp.dot的两个输入应用in_q,对输出应用out_dq。
结合源码可以更精确地理解延迟缩放的元数据更新机制(fp8_ops.py):
def compute_amax_history(x, amax_history): amax_update = jnp.max(jnp.abs(x)).astype(amax_history.dtype) new_history = jnp.roll(amax_history, shift=-1, axis=0).at[0].set(amax_update) return new_history def update_fp8_meta(x, q_dtype, scale, amax_history): ... amax_from_history = jnp.max(amax_history, axis=0) # 取历史最大值 new_scale = compute_scale(amax_from_history, scale, dtype_max) new_history = compute_amax_history(x, amax_history) # 滚动更新历史 return new_scale, new_history- 新缩放因子来自历史 amax 的最大值,
compute_scale的算法参考了 NVIDIA Transformer Engine 的update_fp8_metas(fp8_ops.py 注释中注明来源),且用jnp.where处理了 amax 为 0 或非有限值的边界情况; - 新历史通过
jnp.roll把最新 amax 写入队首、其余后移,实现固定长度(1024)的滑动窗口。
延迟缩放避免了每次前向都重新扫描整个操作数张量来计算 scale,因而显著降低了额外访存开销,是高性能 FP8 训练的实际选择。
3. Flax 高层 API:即插即用的 FP8 层
手动维护延迟缩放的 amax history 和缩放因子相当繁琐。为此,Flax 提供了两个高层模块,作为现有层的直接替代品:
| 高层 API | 替代对象 | 适用场景 |
|---|---|---|
fp8_ops.Fp8DotGeneral | lax.dot_general | nn.Dense/nn.DenseGeneral等点积层 |
fp8_ops.Fp8Einsum | jnp.einsum | Einsum 运算,如 Mixture of Experts(MoE)层 |
这两个模块自动处理所有 FP8 相关功能:量化/反量化、缩放因子更新、前向与反向传播的 FP8 dtype 选择(前向 E4M3、反向梯度 E5M2),用户无需手写任何 Q/DQ 代码。
3.1 为nn.Dense注入 FP8 点积
只需一行参数即可将默认的lax.dot_general替换为 FP8 版本:
model = nn.Dense(features=64, dot_general_cls=fp8_ops.Fp8DotGeneral) params = model.init(k0, A) @jax.jit def train_step(var, a): c = model.apply(var, a) return jnp.sum(c) check_fp8_call(train_step.lower(params, A))从源码看,nn.Dense与nn.DenseGeneral均声明了dot_general_cls: Any = None属性,并在构造点积函数时优先使用它(见 flax/linen/linear.py、linear.py)。也就是说,dot_general_cls是 Linen 线性层预留的标准扩展点,FP8 只是其应用之一。
模型的使用方式与普通 Dense 完全一致,但params中会额外包含 FP8 量化专用参数(缩放因子与 amax history,详见第 5 节)。
3.2 为 Einsum(如 MoE 层)启用 FP8
对于使用jnp.einsum的模型(如 MoE 层中的专家计算),可以将其替换为fp8_ops.Fp8Einsum:
from typing import Any class FooModule(nn.Module): einsum: Any = None @nn.compact def __call__(self, a, b): if self.einsum is not None: einsum_fn = self.einsum() elif self.einsum is None: einsum_fn = jnp.einsum c = einsum_fn("mk,kn->mn", a, b) return c model = FooModule(einsum=fp8_ops.Fp8Einsum) params = model.init(k0, a, b) @jax.jit def train_step(var, a, b): c = model.apply(var, a, b) return jnp.sum(c) check_fp8_call(train_step.lower(params, a, b))从源码实现看,Fp8Einsum.__call__(fp8_ops.py)假定rhs是权重且其 dtype 即实际计算 dtype,然后通过jnp.einsum(..., _dot_general=dot_general_fn)将内部点积替换为fp8_scaled_dot_general,从而复用同一套 FP8 量化逻辑。
3.3 高层 API 的底层调用链
无论走Fp8DotGeneral(源码中Fp8DotGeneral = Fp8DirectDotGeneralOp,见 fp8_ops.py)还是Fp8Einsum,最终都汇聚到fp8_scaled_dot_general(fp8_ops.py),其完整流程为:
- 对左操作数(激活)和右操作数(内核)分别调用
in_q,量化为 E4M3 并得到新缩放因子; - 调用
quantized_dot(基于lax.dot_general的自定义 VJP)完成 FP8 点积,以preferred_element_type高精度累加; - 对输出调用
out_dq反量化回原始类型; - 反向传播时(
quantized_dot_bwd,fp8_ops.py):梯度先经update_fp8_meta量化为E5M2,再通过转置点积计算左/右梯度并反量化——这正是指南所说"梯度用 E5M2、激活/权重用 E4M3"的推荐实践,且对前向/反向同时生效。
注:
Fp8DotGeneralBase同时提供了amax_history_length: int = 1024、e4m3_dtype、e5m2_dtype等可配置属性(fp8_ops.py),默认 amax 历史长度为 1024。仓库还提供面向 FP8 变体格式(FN 系列)的NANOOFp8DotGeneralOp(fp8_ops.py),相关行为在测试 linen_test.py 中与Fp8DotGeneral一并验证。
4. 验证与精度保障:测试如何佐证 FP8 实现
Fp8DotGeneral/Fp8Einsum的正确性并非空谈。仓库测试 tests/linen/linen_test.py 提供了强有力的实现证据:
test_fp8_einsum(linen_test.py):对多组形状与 einsum 方程(mk,kn->mn、...k,nk->...n、...k,kn->...n)分别用Fp8Einsum与纯 FP32jnp.einsum计算前向输出和梯度,断言二者误差在atol=1e-02, rtol=1e-02以内;test_fp8_dot_general_injection(linen_test.py):将dot_general_cls注入nn.DenseGeneral,对比启用/禁用 FP8 的输出与梯度,并校验注入后多出的_overwrite_with_gradient参数结构(三个(1024,)形状的 amax history + 三个(1,)形状的 scale);test_fp8_train_state(linen_test.py):手动复算 5 个训练步中期望的 amax history 与缩放因子(compute_scale+jnp.roll更新),与TrainState实际维护的 FP8 元数据逐项比对,验证延迟缩放的更新逻辑。
这些测试说明:FP8 量化带来的数值差异被严格控制在可接受误差带内,同时元数据的演化规律与手算一致——这也是读者在自己项目中接入 FP8 时可以复用的验证思路。
5. 操纵 FP8 参数:_overwrite_with_gradient解析
5.1 参数结构与三对 scale/amax_history
高层 API 内部由Fp8DotGeneralBase.setup()创建六项 FP8 元参数(fp8_ops.py),它们不属于普通的params集合,而是存放在名为_overwrite_with_gradient的独立变量集合中。指南通过树形结构展示(值已打码):
params_structure = flax.core.unfreeze(params).copy() params_structure = flax.traverse_util.flatten_dict(params_structure, sep='/') for key, value in params_structure.items(): params_structure[key] = '*' params_structure = flax.traverse_util.unflatten_dict(params_structure, sep='/') pprint.pprint(params_structure)输出:
{'_overwrite_with_gradient': {'Fp8Einsum_0': {'input_amax_history': '*', 'input_scale': '*', 'kernel_amax_history': '*', 'kernel_scale': '*', 'output_grad_amax_history': '*', 'output_grad_scale': '*'}}}除常规params外,_overwrite_with_gradient集合包含三对 amax_history 与 scale,分别服务于:
| 参数对 | 作用对象 | 使用的 FP8 格式 |
|---|---|---|
input_amax_history/input_scale | 激活(点积左操作数) | E4M3 |
kernel_amax_history/kernel_scale | 内核/权重(点积右操作数) | E4M3 |
output_grad_amax_history/output_grad_scale | 点积输出梯度 | E5M2 |
其初始化方式为:scale 用ones_init()初始化为 1.0,amax history 用zeros_init()初始化为长度 1024 的零向量(对应amax_history_length=1024)。源码中通过self.variable(OVERWRITE_WITH_GRADIENT, 'input_amax_history', ...)这类调用把变量挂入该集合,其中OVERWRITE_WITH_GRADIENT = '_overwrite_with_gradient'是模块级常量(fp8_ops.py)。
5.2 更新 FP8 参数:梯度直接覆盖
指南给出一个完整的手动更新示例。先执行一步训练得到梯度:
step_fn = jax.jit(jax.grad(train_step, (0, 1))) grads = step_fn(params, A) params = flax.core.unfreeze(params) params = flax.traverse_util.flatten_dict(params, sep='/') grads = flax.traverse_util.flatten_dict(grads[0], sep='/') for key, value in params.items(): if key.startswith('params'): params[key] = value + 0.01 * grads[key] if key.startswith('_overwrite_with_gradient'): params[key] = grads[key] params = flax.traverse_util.unflatten_dict(params, sep='/') params = flax.core.freeze(params)两类参数的更新策略截然不同:
params(普通模型参数):采用梯度下降new_param = old_param + lr * grad,0.01即学习率(实际项目可替换为 optax 优化器);_overwrite_with_gradient(FP8 元参数):直接使用梯度覆盖旧值——因为 scale 与 amax history 的"梯度"本身就是新的元数据(如新的缩放因子、更新后的历史),累加优化反而没有意义。
这种"用梯度覆盖"的语义正是该集合名称_overwrite_with_gradient的由来。
好消息是:flax.training.train_state.TrainState原生支持_overwrite_with_gradient集合(flax/training/train_state.py)。其apply_gradients中检测到该集合时,会对其单独执行"梯度直接覆盖",而对params走 optax 优化器;create时也会把该集合从 opt_state 初始化中排除(因为无需动量等状态)。因此,使用默认TrainState的用户无需修改任何训练脚本,只有使用自定义 TrainState 时才需要自己实现上述覆盖逻辑。
5.3 累积梯度:fp32_max_grad与分支场景
当同一参数被分支使用(如流水线并行中一个 minibatch 的多个 microbatch 共享同一组参数)时,autograd 会对来自各分支的梯度做加法累积。但对于_overwrite_with_gradient参数,加法累积没有意义——新的 scale/history 应当取各分支更新的最大值而非求和。
为此,Flax 引入自定义 dtypefp8_ops.fp32_max_grad。基本用法:
fmax32 = fp8_ops.fp32_max_grad def reuse_fp8_param(x, y, scale, amax_history): scale = scale.astype(fmax32) amax_history = amax_history.astype(fmax32) x = fp8_ops.in_qdq(f32, e4m3, x, scale, amax_history) y = fp8_ops.in_qdq(f32, e4m3, y, scale, amax_history) return x + y reuse_fp8_param_fn = jax.grad(reuse_fp8_param, (0, 1, 2, 3)) reuse_fp8_param_fn = jax.jit(reuse_fp8_param_fn) _, _, new_ah, new_sf = reuse_fp8_param_fn(2.0, 3.0, a_scale, a_amax_hist) print(new_ah, new_sf)将 scale 与 amax_history 转为fp32_max_grad后,同一对 scale/history 被fp8_ops.in_qdq使用两次,autograd 对两个分支的梯度取最大值,得到正确结果:
1.0 [3. 0. 0. ... 0. 0. 0.]而不做类型转换时,两个分支的梯度被相加:
2.0 [5. 0. 0. ... 0. 0. 0.]这一机制在源码中的落点是Fp8MetaTyRules.add(fp8_ops.py):它定义了 FP8 元 dtype 的加法语义为lax.max(先转回 float32 取最大值再转回元 dtype),zero则以负无穷为加法单位元。测试 linen_test.py 在jax.lax.scan循环中复用了同一对 fmax32 参数,验证了 amax history 按最大而非求和方式累积。
使用高层 API(
Fp8DotGeneral/Fp8Einsum)时,这一 dtype 转换已被内置(元参数默认即为fp32_max_grad),无需用户处理——只有使用低层in_qdq等 API 手动实现分支复用时才需要显式.astype(fmax32)。
6. 已弃用 API 与迁移指南
指南明确列出了历史上存在过的两代已弃用 API,并解释了弃用原因与迁移路径。
弃用原因:早期 API(fp8_ops.quantize_dequantize、fp8_ops.[in|out]_qdq、Fp8DotGeneralOp)依赖 XLA-FP8 的模式匹配能力,即由 XLA 将 QDQ→dot 序列自动匹配为 Q→fp8_cublas_gemm。但这种模式匹配方案很脆弱:QDQ 序列很容易被其他 XLA 优化破坏,导致 FP8 路径失效。
迁移对照表(旧写法 → 新写法):
| 旧 API | 新 API |
|---|---|
fp8_ops.quantize_dequantize → jnp.dot | fp8_ops.quantize → jnp.dot → fp8_ops.dequantize |
fp8_ops.in_qdq → jnp.dot → fp8_ops.out_qdq | fp8_ops.in_q → jnp.dot → fp8_ops.out_dq |
fp8_ops.Fp8DotGeneralOp | fp8_ops.Fp8DotGeneral(另提供 einsum 变体fp8_ops.Fp8Einsum) |
源码中同样保留了兼容痕迹:Fp8DotGeneralOp构造时会触发DeprecationWarning,提示改用Fp8DirectDotGeneralOp或Fp8Einsum(fp8_ops.py);而Fp8DotGeneral本身就是Fp8DirectDotGeneralOp的别名(fp8_ops.py),直接基于 FP8 点积 + 显式 Q/DQ 实现,不再依赖脆弱的模式匹配。
7. 完整工作流总结
综合全文,一个基于 Flax 高层 API 的 FP8 训练工作流可以归纳为五个步骤:
- 硬件检查:确认运行在 NVIDIA Hopper(compute capability 9.0+)GPU 上,并断言环境;
- 替换算子:Dense 类层设置
dot_general_cls=fp8_ops.Fp8DotGeneral,Einsum 场景(如 MoE)替换为fp8_ops.Fp8Einsum; - 初始化与验证:
model.init后自动获得含_overwrite_with_gradient集合的变量;可通过编译 HLO 中的custom-call(f8e4m3fn...)确认 FP8 核函数真正生效; - 训练更新:使用默认
flax.training.train_state.TrainState——它对params走 optax 优化,对_overwrite_with_gradient自动执行梯度直接覆盖; - 进阶调优:涉及参数分支复用(如流水线并行)时,借助
fp8_ops.fp32_max_grad的取最大值累积语义保证 FP8 元数据正确性。
如需进一步参考,可阅读 docs/guides/quantization/index.rst 了解量化指南目录结构;FP8 相关的完整示例与验证代码位于 docs/guides/quantization/fp8_basics.ipynb(与本文所讲文档同源的 notebook 版),底层实现见 flax/linen/fp8_ops.py,测试覆盖见 tests/linen/linen_test.py。
【免费下载链接】flaxFlax is a neural network library for JAX that is designed for flexibility.项目地址: https://gitcode.com/GitHub_Trending/fl/flax
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考