Prhub

#28974 [weight checker] refactor: add precision branch; allow ULP quant err; used chunked compare

原始 PR 作者 yueming-yuan 合并时间 2026-06-29 09:44 文件变更 7 提交数 21 评论 10 代码增减 +706 / -189

执行摘要

重构 weight checker 加入 ULP 量化误差容差与分块比较

Weight checker 的 compare/checksum 操作要求两个块量化权重在反量化后比特级精确,但同一源权重经过两次独立但正确的量化(如权重更新前后重新量化)会导致每个元素有合法差异。使用手工调节的平均误差阈值既太宽松(均值会掩盖大张量中的元素破坏),又对每个模型/精度来说是任意的。此 PR 用有原理的逐元素容差替代:每个块量化对是同一源权重的一次忠实量化,因此每个元素在反量化空间中可能偏离最多 own representation 的 1 ULP,即两者可能相差最多 expect_ulp + actual_ulp。通过 allow_quant_error 选择启用,默认保持精确相等(无行为变化)。

建议精读 weight_checker_comparator.py 中的 ComparableWeight 设计与 _quant_ulp 实现,该方案为量化比较提供了理论依据。整体设计优雅,值得在类似场景复用。

讨论亮点

reviewer fzyzcjy 提出了几点改进建议,均被采纳:

  • 统一使用 NamedTuple 代替裸元组提高可读性(CheckEntryCompareResult 等)。
  • 不应在 weight_checker.py 中导入 comparator 的私有符号(_CHUNK_NUMEL),改为公开常量。
  • 将 comparator 的测试独立到 test_weight_checker_comparator.py 文件。
  • reviewer fzyzcjy 还提出了一个随机思考:检查初始 HF 权重是否被篡改为 NaN 以检测 local attention buffer 问题,但未在此 PR 中实现。
  • reviewer hnyls2002 和 fzyzcjy 均给出了批准,总体方向认可。

实现拆解

  1. 提取独立比较器模块(新增 weight_checker_comparator.py):定义 ComparableWeight 基类和 Fp8BlockComparableRawComparable 子类,封装块量化权重与原始张量的反量化与分块迭代逻辑。核心方法 _quant_ulp 利用 torch.frexp 计算逐元素 ULP 间距,供后续容差比较使用。iter_chunks 分块返回 (反量化值, ULP) 元组,控制单次 GPU 内存峰值在 64M 元素以内。

  2. 量化方法路由select_comparable_weight):根据 module.quant_method 类型分发到对应 ComparableWeight 子类,对未支持的精度(如 nvfp4、int4)直接抛出 NotImplementedError,避免静默错误比较。

  3. 改进 WeightChecker 核心逻辑weight_checker.py):将原来的 _postprocess_tensors 替换为 _build_check_entries,后者利用 _build_quantized_set 扫描模型模块生成量化映射,然后将每个张量包装为 ComparableWeight 对象。_compare 方法新增 allow_quant_error 参数,传递给 compare_weights 函数。

  4. 扩展调用链路:修改 model_runner.check_weightsweight_updater.check_weights 以接受并传递 allow_quant_error 参数;CheckWeightsReqInput 新增 allow_quant_error: bool 字段。

  5. 新增大量单元测试test_weight_checker_comparator.pytest_weight_checker.py):覆盖 _quant_ulp 的暴力验证(对所有 256 种 fp8 位模式验证间距)、compare_weights 在相同/微小差异/位翻转/NaN 下的行为、_build_check_entries 对不同精度张量的包装正确性、分块与不分块结果的一致性。

文件 模块 状态 重要度
python/sglang/srt/utils/weight_checker_comparator.py 比较器 added 9.07
python/sglang/srt/utils/weight_checker.py 检查器 modified 8.85
test/registered/unit/utils/test_weight_checker_comparator.py 测试 added 8.05
test/registered/unit/utils/test_weight_checker.py 测试 modified 7.82
python/sglang/srt/model_executor/model_runner.py 模型执行 modified 6.28
python/sglang/srt/managers/scheduler_components/weight_updater.py 权重更新 modified 5.87
python/sglang/srt/managers/io_struct.py 入参结构 modified 4.58

关键符号

ComparableWeight._quant_ulp Fp8BlockComparable.iter_chunks Fp8BlockComparable.dequantize compare_weights select_comparable_weight WeightChecker.handle WeightChecker._compare WeightChecker._compute_checksum _build_quantized_set _build_check_entries ModelRunner.check_weights WeightUpdater.check_weights

关键源码片段

python/sglang/srt/utils/weight_checker_comparator.py dependency-wiring

新增核心比较器模块,定义 ComparableWeight 基类、ULP 计算、分块迭代、Fp8BlockComparable 实现等,是重构的核心。

class ComparableWeight:
    """量化比较基类;每个子类对应一种量化精度或原始张量。"""
​
    @staticmethod
    def _quant_ulp(w_q: torch.Tensor) -> torch.Tensor:
        """计算 w_q 中每个元素在其自身 dtype 下的 ULP 间距。"""
        finfo = torch.finfo(w_q.dtype)
        x = w_q.to(torch.float32).abs()
        # frexp: x = m * 2^e,其中 m in [0.5, 1),所以 2^(e-1) 是 x 的 binade 基数
        _, exponent = torch.frexp(x)
        binade = torch.exp2((exponent - 1).to(torch.float32))
        # 零和次正规数共享最小正规数的间距
        binade = binade.masked_fill(x < finfo.smallest_normal, finfo.smallest_normal)
        return binade * finfo.eps
​
    def iter_chunks(self) -> Iterable[Tuple[torch.Tensor, Optional[torch.Tensor]]]:
        raise NotImplementedError
​
    def dequantize(self, dtype: torch.dtype = torch.bfloat16) -> torch.Tensor:
        raise NotImplementedError
​
​
class Fp8BlockComparable(ComparableWeight):
    """Deepseek 风格 FP8 块量化权重。"""
​
    def __init__(self, w_q: torch.Tensor, w_s: torch.Tensor):
        self.w_q = w_q
        self.w_s = w_s
​
    def __repr__(self) -> str:
        return f"fp8_block(shape={tuple(self.w_q.shape)} dtype={self.w_q.dtype})"
​
    @staticmethod
    def _normalize_scale(w_q: torch.Tensor, w_s: torch.Tensor) -> torch.Tensor:
        # 若 scale 以 ue8m0 打包为 int32,需先解包
        if w_s.dtype == torch.int32:
            w_s = inverse_transform_scale_ue8m0(w_s, mn=w_q.shape[-2])
        return w_s.to(torch.float32)
​
    @staticmethod
    def _infer_block_size(w_q: torch.Tensor, w_s: torch.Tensor) -> list:
        k, s_k = w_q.shape[-1], w_s.shape[-1]
        assert k % s_k == 0, f"cannot infer block size from {w_q.shape=} {w_s.shape=}"
        block = k // s_k
        return [block, block]
​
    def _scale_and_block_size(self):
        s = self._normalize_scale(self.w_q, self.w_s)
        return s, self._infer_block_size(self.w_q, s)
​
    def iter_chunks(self):
        """分块迭代 (反量化张量, ULP 张量) 对,每块不超过 CHUNK_NUMEL 元素。"""
        s, block_size = self._scale_and_block_size()
        for q, s_chunk in self._iter_quant_chunks(self.w_q, s, block_size[0]):
            q, s_chunk = q.cuda(), s_chunk.cuda()
            yield (
                block_quant_dequant(q, s_chunk, block_size, dtype=torch.bfloat16),
                block_quant_dequant(
                    self._quant_ulp(q), s_chunk, block_size, dtype=torch.float32
                ),
            )

评论区精华

统一使用 NamedTuple 提高可读性 style

reviewer fzyzcjy 建议使用 dataclass 或 NamedTuple 代替元组以提升可读性。

结论:采纳,将 Entry 改为 CheckEntry(NamedTuple),CompareResult(NamedTuple) 等。 · 已解决

不应导入私有符号 style

reviewer fzyzcjy 指出不应在 weight_checker.py 中导入 comparator 的私有符号 _CHUNK_NUMEL。

结论:将 CHUNK_NUMEL 改为公开常量。 · 已解决

分离 comparator 测试 测试

reviewer fzyzcjy 建议将 comparator 的测试分离到单独文件。

结论:创建 test_weight_checker_comparator.py。 · 已解决

随机思考:检测 NaN 初始化权重 other

reviewer fzyzcjy 提出检查初始 HF 权重是否被篡改为 NaN 以检测 local attention buffer 问题。

结论:未在本 PR 中实现,可能作为未来改进。 · unresolved

风险与影响

  • 回归风险:重构涉及比较逻辑重写,尽管有大量测试但可能遗漏边界情况(如非常规块大小、未覆盖的量化格式)。但默认行为不变(精确相等),回归影响有限。
  • 性能风险:分块比较引入额外 GPU 内存操作,可能使比较速度变慢。但对单个小权重影响不大,大张量时 chunk 机制可避免 OOM。
  • 兼容性风险:新增 ComparableWeight 接口,未来扩展新精度需实现子类;未支持精度会抛 NotImplementedError,可能暴露之前静默无比较的问题。
  • 测试依赖 GPU:所有测试标记为 CUDA 门控,无法在 CPU 环境运行。
  • 用户:无行为变化(默认 allow_quant_error=False)。需要量化容差时传入参数即可。checksum 行为一致(仍反量化后哈希)。
  • 系统:比较逻辑更精确,避免因合法量化误差导致误报。分块比较可能减少 GPU OOM。
  • 团队:模块化设计便于添加新精度支持。ComparableWeight 子类封装了格式细节,降低维护成本。
重构核心路径 默认行为不变但需测试 新增依赖接口 GPU 内存分块依赖

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论