vllm_fp8_utils.py — 这个文件实现了 vLLM 推理引擎的 FP8 量化支持¶
模块路径: verl.utils.vllm.vllm_fp8_utils
文件概述¶
这个文件实现了 vLLM 推理引擎的 FP8 量化支持。在 verl 的训练-推理循环中,训练使用 BF16/FP16 精度,但推理时可以使用 FP8 量化来提升速度。本文件提供了权重的 FP8 量化、加载、以及 vLLM 权重后处理的补丁。
关键代码讲解¶
1. FP8 状态管理¶
@dataclass
class FP8State:
seen_params: set = field(default_factory=lambda: set()) # 已检查的参数
fp8_param_names: set = field(default_factory=lambda: set()) # 确认是 FP8 的参数
vllm_patches: list = field(default_factory=lambda: [])
fp8_state = FP8State() # 全局单例
2. 判断是否为 FP8 权重¶
def is_fp8_weight(name, model):
if name not in fp8_state.seen_params:
fp8_state.seen_params.add(name)
if name.endswith("weight"):
module = get_module_from_param_name(model, name)
if isinstance(module, LinearBase) and module.weight.dtype == torch.float8_e4m3fn:
fp8_state.fp8_param_names.add(name)
elif isinstance(module, FusedMoE) and module.w13_weight.dtype == torch.float8_e4m3fn:
fp8_state.fp8_param_names.add(name)
return name in fp8_state.fp8_param_names
通过检查模块类型和现有权重的 dtype 来判断。
3. 权重量化(流式 Generator)¶
def quant_weights(weights, model, quant_config, dtype=torch.bfloat16):
"""将 BF16 权重量化为 FP8,以 generator 方式流式输出"""
for k, v in weights:
if not is_fp8_weight(k, model):
yield (k, v)
continue
# FP8 blockwise 量化
param_lp, param_scale = scaled_fp8_blockwise(v.to(dtype), weight_block_size=quant_config.weight_block_size)
yield (k, param_lp) # 量化后的权重
yield (k + "_scale_inv", param_scale) # scale 参数
del v, param_lp, param_scale # 主动释放内存
Generator 模式的好处:不需要将所有量化结果存入内存,逐个产出。
4. vLLM 权重后处理补丁¶
def process_weights_after_loading_for_vllm11(self, layer):
"""替换 vLLM 的权重后处理,避免创建新 Parameter(保留 weight_loader)"""
def _create_param_from_subclass_attributes(custom_param):
param = Parameter(custom_param.data, requires_grad=False)
# 复制自定义属性(如 weight_loader)
for attr in custom_attributes:
setattr(param, attr, getattr(custom_param, attr))
param.subclass_type = type(custom_param) # 保存原始类型
return param
# ... 处理权重和 scale
核心问题:vLLM 原生的 process_weights_after_loading 会创建新的 Parameter 对象,丢失 weight_loader 属性。补丁版本保留了这个属性,使得权重可以被多次重载。
5. 应用补丁¶
def apply_vllm_fp8_patches():
# 根据 vLLM 版本选择对应的补丁
patcher1 = patch(
"vllm...Fp8LinearMethod.process_weights_after_loading",
process_weights_after_loading_for_vllm11 if vllm >= "0.11.0"
else process_weights_after_loading_for_vllm10
)
patcher1.start()
# 类似地补丁 MoE 层
核心类/函数列表¶
| 类/函数名 | 作用 |
|---|---|
FP8State |
FP8 参数名的缓存 |
is_fp8_weight |
判断参数是否为 FP8 权重 |
quant_weights |
流式 FP8 量化 |
load_quanted_weights |
加载量化权重到 vLLM |
apply_vllm_fp8_patches |
应用 vLLM 补丁 |
process_weights_after_loading_for_vllm10/11 |
不同版本的补丁 |
与其他模块的关系¶
- 使用
verl.utils.kernel.fp8_kernel.scaled_fp8_blockwise做量化 - 需要 vLLM 的 FP8 量化配置
- 支持 vLLM v0.10 和 v0.11 两个大版本
小结¶
FP8 推理是提升大模型推理速度的重要技术。这个文件解决了 verl 训练循环中的 FP8 量化和加载问题,同时通过补丁机制确保 vLLM 的权重后处理不会破坏 verl 需要的 weight_loader 属性。