Prhub

#46236 [Bugfix][Quant] Raise actionable error instead of bare assert for group-size/TP mismatch (#46230)

原始 PR 作者 ArsalanShakil 合并时间 2026-06-30 22:57 文件变更 6 提交数 35 评论 15 代码增减 +37 / -12

执行摘要

将 group-size/TP 不匹配的 bare assert 替换为可操作的 ValueError

Issue #46230 报告:加载 W4A16/W8A16 等分组量化 checkpoint 时遇到 AssertionError: assert input_size_per_partition % group_size == 0,用户无法得知如何解决。PR 旨在给出清晰错误信息和建议措施,避免用户困惑。

建议合并。该 PR 解决了用户实际遇到的模糊错误,是典型的开发者体验改进。设计上将 TP 相关的量化检查从 Marlin 特定代码中抽离到 distributed 层,体现了良好的模块化思考,未来类似校验可复用该 helper。

讨论亮点
  1. vadiklyutiy 指出新代码可能和已有的 check_marlin_supports_shape 重复;作者解释已将共享逻辑抽取到 helper,并保持 Marlin 特定约束检查独立,不会误拒合法形状。
  2. hmellor 建议将函数作为 TP 工具置于 distributed/utils.py,作者采纳。
  3. hmellor 认为 docstring 过于冗长,要求简化;作者压缩为一行。
  4. hmellor 认为不需要专门测试 x % y 的测试文件,作者删除测试文件。

实现拆解

  1. vllm/distributed/utils.py 新增 verify_group_size_divides_partition 函数,替代分散在各处的 bare assert,该函数验证 TP 分片是否包含整数个量化组,否则 raise ValueError 并附带诊断信息。
  2. marlin_utils.pyverify_marlin_supports_shape 中,将 group-size 整除检查委托给新函数,同时保留 Marlin 特定的 min_thread_n / min_thread_k 检查。
  3. 在四个 compressed-tensors scheme 文件(compressed_tensors_w4a8_int.pycompressed_tensors_wNa16.pycompressed_tensors_wNa8o8.pycompressed_tensors_w4a8_fp8.py)的 create_weights_register_weight 方法中,将原 assert 替换为对新 helper 的调用,部分调用点传入 layer_name 以提升错误定位精度。
  4. 根据 review 意见,将函数从 marlin_utils 移至 distributed/utils,并精简 docstring 和错误消息格式。
文件 模块 状态 重要度
vllm/distributed/utils.py 分布式工具 modified 6.79
vllm/model_executor/layers/quantization/utils/marlin_utils.py 量化工具 modified 6.35
vllm/model_executor/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a8_int.py 量化方案 modified 5.4
vllm/model_executor/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa16.py 量化方案 modified 5.34
vllm/model_executor/layers/quantization/compressed_tensors/schemes/compressed_tensors_wNa8o8.py 量化方案 modified 5.34
vllm/model_executor/layers/quantization/compressed_tensors/schemes/compressed_tensors_w4a8_fp8.py 量化方案 modified 5.2

关键符号

verify_group_size_divides_partition verify_marlin_supports_shape CompressedTensorsW4A8Int.create_weights CompressedTensorsWNA16.create_weights CompressedTensorsWNA8O8._register_weight CompressedTensorsW4A8FP8.create_weights

关键源码片段

vllm/distributed/utils.py core-logic

新增核心验证函数,所有调用点统一依赖此函数,是变更的中心。

def verify_group_size_divides_partition(
    input_size_per_partition: int,
    group_size: int,
    layer_name: str | None = None,
    extra_suggestion: str = "",
) -> None:
    """Validate that a TP-sharded layer holds a whole number of quant groups."""
    # 如果整除,直接返回
    if input_size_per_partition % group_size == 0:
        return
    # 构造包含 `layer_name`(如有)和 `extra_suggestion` 的错误消息
    location = f" for layer '{layer_name}'" if layer_name else ""
    raise ValueError(
        f"Weight {input_size_per_partition=}{location} is not divisible by "
        f"{group_size=}. This happens when tensor_parallel_size splits the layer input "
        "into shards that are not a whole number of quant groups. Consider reducing "
        f"tensor_parallel_size{extra_suggestion}."
    )
vllm/model_executor/layers/quantization/utils/marlin_utils.py data-contract

将 Marlin 形状校验中的 group-size 整除检查委托给共享 helper,保持 Marlin 特有检查独立。

def verify_marlin_supports_shape(
    output_size_per_partition: int,
    input_size_per_partition: int,
    input_size: int,
    group_size: int,
) -> None:
    # Validate output_size_per_partition
    if output_size_per_partition % GPTQ_MARLIN_MIN_THREAD_N != 0:
        raise ValueError(
            f"Weight output_size_per_partition = "
            f"{output_size_per_partition} is not divisible by "
            f" min_thread_n = {GPTQ_MARLIN_MIN_THREAD_N}. "
            "Consider reducing tensor_parallel_size or running "
            "with --quantization gptq."
        )
    # Validate input_size_per_partition
    if input_size_per_partition % GPTQ_MARLIN_MIN_THREAD_K != 0:
        raise ValueError(
            f"Weight input_size_per_partition = "
            f"{input_size_per_partition} is not divisible "
            f"by min_thread_k = {GPTQ_MARLIN_MIN_THREAD_K}. "
            "Consider reducing tensor_parallel_size or running "
            "with --quantization gptq."
        )
    # 当 group_size 小于 input_size 时,检查 group-size 整除
    if group_size < input_size:
        # 委托给共享验证函数,并补充 Marlin 专用建议
        verify_group_size_divides_partition(
            input_size_per_partition,
            group_size,
            extra_suggestion=" or running with --quantization gptq",
        )

评论区精华

与 check_marlin_supports_shape 的重复问题 设计

vadiklyutiy 指出新代码可能和已有的 `check_marlin_supports_shape` 重复。作者解释已将共享逻辑抽取到 helper,并保持 Marlin 特定约束检查独立,不会拒绝对 Marlin 无效但其他内核有效的形状。

结论:作者保持 helper 独立,`verify_marlin_supports_shape` 的 group-size 检查委托给 helper。没有重复。 · 已解决

函数放置位置:marlin_utils vs distributed/utils 设计

hmellor 建议将函数从 `marlin_utils` 移到 `distributed.utils`,认为这是 TP 层面的工具而非 Marlin 特有。作者采纳。

结论:函数移至 `vllm/distributed/utils.py`。 · 已解决

docstring 和错误消息长度 style

hmellor 指出 docstring 过于冗长,建议缩短。作者从命,将多行 docstring 压缩为一行,并精简错误消息中的措辞。

结论:docstring 缩短为一行,错误消息格式精简。 · 已解决

删除测试文件 测试

hmellor 认为不需要专门为 `x % y` 编写测试文件。作者同意并删除了 `tests/quantization/test_group_partition_divisibility.py`。

结论:测试文件被移除,功能验证依赖现有测试覆盖。 · 已解决

风险与影响

风险极低:本质是将 assert 替换为显式 ValueError,逻辑行为一致(不满足整除条件时停止加载)。但需要注意:原 assert 在 Python 最优模式下(-O)会被跳过,而 ValueError 始终生效,这实际上是行为改进(错误不会被静默忽略)。缺少测试覆盖是主要风险点(测试文件被删除),但该函数逻辑简单,未来变更可能引入注入。

影响范围限定于 compressed-tensors 分组量化模型加载(WNA16、WNA8A8、W4A8-FP8 等)以及 Marlin 内核校验路径。对正常整除通过的分片无任何性能影响,错误路径从无消息 abort 变为提供可操作指引,显著提升开发者体验。

缺少测试覆盖 低风险变更

关联 Issue

#46230 [Bug]: W4A16 or W8A16 Qwen3.5 9B meet AssertionError

完整报告

参与讨论