fp8_kernel.py — FP8 量化内核¶
文件路径: verl/utils/kernel/fp8_kernel.py
文件概述¶
实现 FP8 (8-bit floating point) 块级量化内核,用于在推理引擎(vLLM/SGLang)中使用 FP8 权重以加速推理。提供 Triton 和 PyTorch 两种实现。
背景知识¶
FP8 是一种 8 位浮点格式(float8_e4m3fn),相比 FP16/BF16 节省一半内存,在 H100 等新 GPU 上有专门的硬件加速。块级量化 将权重矩阵分成小块,每块使用独立的缩放因子,比全局量化更精确。
核心函数¶
1. 统一入口¶
def scaled_fp8_blockwise(data_hp, weight_block_size):
"""
FP8 块级量化的统一入口。
自动选择 Triton(更快)或 PyTorch fallback。
Args:
data_hp: [M, N] 高精度输入
weight_block_size: [BLOCK_M, BLOCK_N] 块大小
Returns:
(fp8_data, descale): 量化结果和反量化缩放因子
"""
if _TRITON_AVAILABLE and not _DISABLE_TRITON_FP8:
return scaled_fp8_blockwise_triton(data_hp, weight_block_size)
return _scaled_fp8_blockwise_pytorch(data_hp, weight_block_size)
2. Triton 实现¶
@triton.jit
def _blockwise_cast_to_fp8_kernel(X, Y, S, ..., BLOCK_M, BLOCK_N):
"""单 kernel 完成:加载 → 求绝对值最大 → 计算缩放 → 量化 → 存储"""
pid_m = tl.program_id(0)
pid_n = tl.program_id(1)
# 加载一个 block 的数据
x = tl.load(X + ...)
# 计算 block 内绝对值最大
_absmax = tl.maximum(tl.max(tl.abs(x)), eps)
# 计算缩放因子
x_s = _absmax / fp8_max
# 量化
y_q = tl.clamp(x * (1.0 / x_s), fp8_min, fp8_max)
# 存储
tl.store(Y + ..., y_q)
tl.store(S + ..., x_s)
单个 Triton kernel 完成所有操作,没有中间张量分配。
3. PyTorch Fallback¶
def _scaled_fp8_blockwise_pytorch(data_hp, weight_block_size):
"""内存优化的 PyTorch 实现"""
# reshape 为 block 结构
data_hp = data_hp.reshape(blk_m, block_size0, blk_n, block_size1)
# 计算 per-block max
max_abs = data_hp.abs().amax(dim=-1, keepdim=True)
# 量化
data_hp.mul_(scale_fp)
data_hp.clamp_(-max_dtype, max_dtype)
fp_data = data_hp.to(FP8_DTYPE)
使用 in-place 操作和显式 del 释放中间张量,最小化内存峰值。
核心函数列表¶
| 函数 | 说明 |
|---|---|
scaled_fp8_blockwise() |
统一入口(自动选择实现) |
scaled_fp8_blockwise_triton() |
Triton 实现 |
_scaled_fp8_blockwise_pytorch() |
PyTorch fallback |
blockwise_cast_to_fp8_triton() |
Triton kernel 封装 |
is_triton_available() |
检查 Triton 可用性 |
环境变量¶
VERL_DISABLE_TRITON_FP8=1: 强制使用 PyTorch fallback
与其他模块的关系¶
- 被 FSDP 模型权重同步到推理引擎时使用
- 依赖
device.py检查 GPU 算力
小结¶
fp8_kernel.py 提供高效的 FP8 块级量化,通过 Triton kernel 在单次 GPU 计算中完成量化,为大模型推理提供 2x 内存节省和速度提升。