Prhub

#24532 Cherry pick weight_checker fp8 dequant fix and non-persistent buffer skip from #21494

原始 PR 作者 fzyzcjy 合并时间 2026-05-06 21:13 文件变更 1 提交数 1 评论 1 代码增减 +9 / -5

执行摘要

修复权值检查器 FP8 反量化与非持久化缓冲区跳过

权值检查器在加载具有非持久化缓冲区(如 cos_sin_cache、freqs_cis)的模型时,会意外重置这些缓存,导致后续校验失败;同时 FP8 缩放张量 weight_scale_inv 未被跳过比较,造成重复比较;此外,inverse_transform_scale_ue8m0 无条件调用,在非 Blackwell 场景下会引发错误。PR #21494 针对这些问题进行了修复,本次将其 cherry-pick 到 main 分支。

可以合并,无需进一步审查。建议在后续开发中关注 weight_checker 的扩展性,例如将非持久化缓冲区名称配置化,便于维护。

讨论亮点

该 PR 没有 review 讨论。由于是简单的 cherry-pick 且改动较小,直接合并。

实现拆解

  1. 跳过非持久化缓冲区(_reset_tensors):在 python/sglang/srt/utils/weight_checker.py 的 _reset_tensors 方法中,遍历模型参数时添加条件判断,若名称包含 "cos_sin_cache" 或 "freqs_cis" 则跳过随机重置。这些缓冲区为非持久化,模型加载后会自动重新计算,因此不需要被重置。
  2. 扩展跳过比较名单(_postprocess_tensors):在 skip_compare_names 列表中不仅加入量化权重名,还加入对应的 weight_scale_inv 名称,避免 FP8 缩放张量被重复比较导致误报。
  3. 条件化 FP8 反量化(_postprocess_tensors):将对 inverse_transform_scale_ue8m0 的调用由无条件改为仅当 w_s.dtype == torch.int32 时执行(UE8M0 压缩格式,用于 Blackwell DeepGEMM),否则直接使用原始 w_s。这使得 block_quant_dequant 在不同硬件上都能正确工作。
文件 模块 状态 重要度
python/sglang/srt/utils/weight_checker.py 权值检查器 modified 6.02

关键符号

_reset_tensors _postprocess_tensors

关键源码片段

python/sglang/srt/utils/weight_checker.py core-logic

全部变更均在此文件中,包括跳过非持久化缓冲区、添加 weight_scale_inv 跳过、条件化 FP8 反量化。

# _reset_tensors:跳过非持久化缓冲区(如 cos_sin_cache、freqs_cis)
def _reset_tensors(self):
    for name, param in self._model_state():
        # 非持久化缓冲区在权重加载后会重新计算,无需重置
        if "cos_sin_cache" in name or "freqs_cis" in name:
            continue
        param.copy_(_random_like(param))
​
​
# _postprocess_tensors:改进 FP8 缩放张量处理和跳过逻辑
def _postprocess_tensors(
    raw: Dict[str, torch.Tensor],
) -> Iterable[Tuple[str, bool, torch.Tensor]]:
    from sglang.srt.debug_utils.dumper import get_tensor_info
​
    skip_compare_names = []
​
    # 找量化权重名称
    quant_names = [
        name
        for name in raw
        if name.endswith("weight") and name.replace("weight", "weight_scale_inv") in raw
    ]
    skip_compare_names += quant_names
    # 将 weight_scale_inv 也加入跳过列表,避免重复比较
    skip_compare_names += [
        name.replace("weight", "weight_scale_inv") for name in quant_names
    ]
    for name in quant_names:
        w_q = raw[name]
        w_s = raw[name.replace("weight", "weight_scale_inv")]
​
        try:
            # 仅在 UE8M0 压缩格式(Blackwell DeepGEMM)时进行逆变换
            if w_s.dtype == torch.int32:
                w_s = inverse_transform_scale_ue8m0(w_s, mn=w_q.shape[-2])
            # 其他情况直接使用原始 w_s
            w_dequant = block_quant_dequant(
                w_q,
                w_s,
                block_size=[128, 128],
                dtype=torch.bfloat16,
            )
            yield name, True, w_dequant
        except Exception as e:
            e.add_note(
                f"when handling {name=} {get_tensor_info(w_q)=} {get_tensor_info(w_s)=}"
            )
            raise
​
    for name in raw:
        should_compare = name not in skip_compare_names
        yield name, should_compare, raw[name]

评论区精华

没有提炼出高价值讨论线程

当前评论区没有形成足够清晰的争议点或结论,后续有更多讨论时会体现在这里。

风险与影响

风险极低。变更仅限于 weight_checker.py 文件,逻辑为功能修复和强化,不涉及核心推理路径。唯一需要注意的是一致性:若后续修改了 FP8 量化逻辑或新增非持久化缓冲区,需同步更新此处的名称过滤列表。

影响范围小,仅作用于权值检查器(调试工具)的执行路径。修复后,在含有 cos_sin_cache/freqs_cis 或 FP8 量化权重的模型上进行权值比较时,可以避免误报和错误,提高调试可靠性。对用户无可见功能变化。

仅调试工具变更

关联 Issue

未识别关联 Issue

当前没有检测到明确关联的 Issue 链接,后续同步到相关引用后会出现在这里。

完整报告

参与讨论