跳转至

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 多种层类型的适配。