跳转至

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 内存节省和速度提升。