跳转至

sglang_fp8_utils.py — 这个文件实现了 SGLang 推理引擎的 FP8 量化支持

模块路径: verl.utils.sglang.sglang_fp8_utils

文件概述

这个文件实现了 SGLang 推理引擎的 FP8 量化支持。与 vllm_fp8_utils.py 功能类似,但针对 SGLang 推理框架。由于 SGLang 不像 vLLM 那样有内置的 FP8 模型检测机制,这里使用基于参数名的规则来判断哪些参数需要量化。

关键代码讲解

1. 参数名判断规则

def should_quantize_param(param_name: str) -> bool:
    """基于参数名判断是否需要 FP8 量化"""
    if not param_name.endswith(".weight"):
        return False

    # 排除的层
    exclude_patterns = [
        "embed_tokens", "lm_head", "layernorm", "norm",
        "ln_", "embeddings", "mlp.gate.weight"
    ]
    for pattern in exclude_patterns:
        if pattern in param_name.lower():
            return False

    # 包含的层(Linear 相关)
    include_patterns = [
        "q_proj", "k_proj", "v_proj", "o_proj",
        "gate_proj", "up_proj", "down_proj",
        "fc1", "fc2", "mlp"
    ]
    for pattern in include_patterns:
        if pattern in param_name.lower():
            return True

    return False  # 默认不量化

白名单 + 黑名单策略:先排除不该量化的层(embedding、norm 等),再包含应该量化的层(各种投影层)。

2. 权重量化

def quant_weights_by_name(weights, quant_config, dtype=torch.bfloat16):
    """基于参数名的流式 FP8 量化"""
    for k, v in weights:
        if not should_quantize_param(k):
            yield (k, v)
            continue
        try:
            param_lp, param_scale = scaled_fp8_blockwise(
                v.to(dtype), weight_block_size=weight_block_size
            )
            yield (k, param_lp)
            yield (k + "_scale_inv", param_scale)
            del param_lp, param_scale
        except Exception as e:
            logger.error(f"Failed to quantize {k}: {e}")
            yield (k, v)  # 量化失败则使用原始权重

与 vLLM 版本的区别: - 使用参数名规则而非模块类型检测 - 有更完善的错误处理(量化失败时 fallback 到原始权重) - scale 参数统一命名为 _scale_inv

核心类/函数列表

函数名 作用
should_quantize_param 根据参数名判断是否量化
quant_weights_by_name 流式 FP8 量化

与其他模块的关系

  • 使用 verl.utils.kernel.fp8_kernel.scaled_fp8_blockwise 做量化
  • 与 vllm_fp8_utils.py 功能对应,但针对 SGLang
  • 在 SGLang 推理引擎的权重加载流程中使用

小结

SGLang 的 FP8 量化实现比 vLLM 版本更简洁,使用参数名规则代替模块类型检测。这种方法更通用但可能不够精确。容错机制(量化失败时 fallback)提高了鲁棒性。