vllm_patch.py — 这个文件为 vLLM 推理框架提供 NVFP4 量化权重的动态加载补丁¶
模块路径: verl.utils.qat.vllm_patch
文件概述¶
这个文件为 vLLM 推理框架提供 NVFP4 量化权重的动态加载补丁。verl 的训练-推理循环中需要反复将训练中更新的权重加载到 vLLM 推理引擎中。由于 vLLM 的原生权重后处理逻辑假设只加载一次,本文件通过 monkey-patch 方式使其支持多次重载。
支持三种量化方案:Dense W4A16、Dense W4A4、MoE NVFP4。
关键代码讲解¶
1. ParamMetaDict -- 参数元数据字典¶
class ParamMetaDict(dict):
"""支持参数重建和 tensor swap 的字典"""
def __init__(self, model, device=None):
self._layer_meta_cache = {} # 缓存每层的元数据
self._tensor_swap_layers = {} # 需要 tensor swap 的层
self._build_mappings()
def __getitem__(self, key):
if key in dict.keys(self):
return super().__getitem__(key)
# 如果参数被删除,尝试从元数据重建
param = self._try_rebuild(key)
if param is not None:
return param
raise KeyError(f"Parameter not found: {key}")
核心问题:vLLM 的 process_weights_after_loading 会删除 HF 格式的参数(如 weight_packed),转换为 Marlin 格式。当需要重新加载权重时,ParamMetaDict 能从保存的元数据重建这些参数。
2. W4A16 补丁¶
def patched_w4a16_process_weights_after_loading(self, layer):
is_first_call = _check_first_call(layer)
if is_first_call:
save_param_meta(layer, "weight_packed") # 保存元数据(首次)
save_param_meta(layer, "weight_global_scale")
save_param_meta(layer, "weight_scale")
# 转换为 Marlin 格式
marlin_weight = ops.gptq_marlin_repack(...)
weight_scale_permuted = marlin_permute_scales(...)
if is_first_call:
layer.weight = Parameter(marlin_weight, requires_grad=False)
layer._marlin_tensor_refs = {"weight_scale": layer.weight_scale.data}
else:
layer.weight.data.copy_(marlin_weight) # 原地更新,保持 CUDA Graph 地址稳定
marlin_scale_ref.copy_(marlin_weight_scale)
# 删除 HF 格式参数
delattr(layer, "weight_packed")
delattr(layer, "weight_global_scale")
关键设计:
- 首次调用创建新 Parameter,后续调用通过 data.copy_() 原地更新
- _marlin_tensor_refs 保存原始 tensor 引用,确保 CUDA Graph 不失效
3. 应用补丁¶
_PATCH_TARGETS = [
("vllm...CompressedTensorsW4A16Fp4.process_weights_after_loading", patched_w4a16),
("vllm...CompressedTensorsW4A4Fp4.process_weights_after_loading", patched_w4a4),
("vllm...CompressedTensorsW4A4Nvfp4MoEMethod.process_weights_after_loading", patched_moe),
]
def apply_qat_patches():
for target, replacement in _PATCH_TARGETS:
p = patch(target, replacement)
p.start()
使用 unittest.mock.patch 替换 vLLM 的方法。
4. 准备重载¶
def prepare_qat_for_load_weights(model, device=None):
"""在权重重载前调用:重建被删除的 HF 格式参数"""
param_meta = ParamMetaDict(inner_model, device=device)
param_meta.prepare_for_reload() # Tensor swap
# 重建所有被删除的参数
for layer_name, cache_entry in param_meta._layer_meta_cache.items():
for param_name, pm in cache_entry["meta"].items():
new_param = _create_param_from_meta(module, param_name, pm, device)
module.register_parameter(param_name, new_param)
核心类/函数列表¶
| 类/函数名 | 作用 |
|---|---|
ParamMetaDict |
支持参数重建的字典 |
apply_qat_patches |
应用所有 vLLM 补丁 |
prepare_qat_for_load_weights |
准备权重重载 |
manual_process_weights_after_loading |
手动触发权重后处理 |
save_param_meta |
保存参数元数据 |
patched_w4a16_process_weights_after_loading |
W4A16 补丁 |
patched_w4a4_process_weights_after_loading |
W4A4 补丁 |
patched_nvfp4_moe_process_weights_after_loading |
MoE 补丁 |
与其他模块的关系¶
- 被
qat/__init__.py导出 - 在 vLLM 推理引擎初始化前调用
apply_qat_patches - 配合
quantizer.py产出的量化权重使用
小结¶
这是 QAT 模块中最复杂的文件,解决了一个核心工程问题:如何让 vLLM 支持多次权重重载。关键技术包括参数元数据缓存、tensor 引用稳定性(for CUDA Graph)、以及对 Dense/MoE 多种层类型的适配。