Prhub

#21431 [Diffusion] [AMD] Online MXFP4 and FP8 Quantization for Multimodal Generation

原始 PR 作者 ColinZ22 合并时间 2026-05-14 08:52 文件变更 10 提交数 20 评论 27 代码增减 +417 / -17

执行摘要

为多模态生成添加在线 MXFP4/FP8 量化支持

根据 PR body:'Adding Online MXFP4 (For AMD GPUs) and FP8 Quantization for multimodal (image and video) generation with models like Z-Image-Turbo and Wan 2.2.' 目标是在不依赖预量化检查点的情况下,通过在线量化减少模型推理时的显存占用并提升吞吐,特别针对 AMD ROCm 平台支持 MXFP4 格式。

该 PR 值得认真阅读,特别是对量化扩展有兴趣的开发者。Mxfp4ConfigMxfp4LinearMethod 的设计模式(继承 QuantizationConfigLinearMethodBase)可作为后续添加新量化方法的参考。Flash Attention 的 FA2 回退处理也是一个良好的兼容性示范。建议重点关注 quantization-ignored-layers 的传递正确性以及 FP8 后端的修复进展。

讨论亮点
  • _is_hip 本地变量:mickqian 建议避免使用模块级 _is_hip 本地变量,ColinZ22 解释这是代码库中多处使用的惯例(如 activation.pyfp8_utils.py),用于缓存 is_hip() 结果。该模式被保留。
  • 量化忽略层参数的可用性:avjves 发现 quantization-ignored-layers CLI 参数的值似乎没有被任何代码实际使用,Mxfp4Configignored_layers 仅在 from_config 中被设置为空列表。ColinZ22 未直接回复,但后续提交(72c846a)添加了 quantization-ignored-layers 的传递逻辑;不过在本次 PR 合并时是否完全修复仍存疑。
  • 帮助文本完善:mickqian 建议在 --quantization 的帮助信息中添加详细说明,包括自动检测逻辑和 MXFP4 的硬件要求,ColinZ22 采纳了建议。
  • Flash Attention 回退设计:在 flash_attention_v3.py 中将 raise NotImplementedError 替换为 FA2 回退,避免了非 CUDA 平台(如 ROCm)上扩散模型的 Attention 崩溃。
  • 文档补充:在 quantization.mdcli.md 中增加了新参数的用法说明和硬件要求。

实现拆解

  1. 新增 MXFP4 量化配置与方法:在 mxfp4.py 中实现 Mxfp4Config(继承 QuantizationConfig)和 Mxfp4LinearMethod(继承 LinearMethodBase)。Mxfp4Config 负责判断层是否应该跳过量化(基于 ignored_layers 和输出维度阈值 _MXFP4_MIN_OUTPUT_DIM),并为每个线性层分配合适的 LinearMethodBaseMxfp4LinearMethodprocess_weights_after_loading() 中调用 AITER 的 dynamic_mxfp4_quant 将 BF16/FP16 权重在线量化为 MXFP4,并在 apply() 中使用 gemm_a4w4shuffle_weight 执行量化 GEMM。
  2. 扩展 FP8 量化配置:在 fp8.pyFp8Config 中增加 packed_modules_mapping 参数,并在 get_quant_method() 中依赖 is_layer_skipped 处理融合层映射,确保 FP8 在线量化也能正确跳过指定层。
  3. 模型接入调整:修改 zimage.py 中的 FeedForward 类,使其构造函数接受 quant_configprefix,并将量化配置传递给内部的线性层。同时为 ZImageTransformer2DBlock 添加 packed_modules_mapping 类属性,用于描述融合参数(如 w13 对应 w1w3),供 is_layer_skipped 正确处理参数名模式。
  4. 加载器与配置提升:在 transformer_load_utils.pyresolve_transformer_quant_load_spec() 中,从模型类获取 packed_modules_mapping 并注入到量化配置对象。在 server_args.py 中添加 quantizationquantization_ignored_layers 两个 CLI 参数,它们的值会传递给量化后端。
  5. Flash Attention 兼容性:在 flash_attention_v3.py 中,当 sgl-kernel 的 FA3 不可用时(如 ROCm 或 sm<90),不再直接抛异常,而是尝试调用 flash_attn 包的 FA2 实现,保证多模态生成在其他平台上也能运行。
文件 模块 状态 重要度
python/sglang/multimodal_gen/runtime/layers/quantization/mxfp4.py 多模态生成 added 9.23
python/sglang/multimodal_gen/runtime/models/dits/zimage.py 多模态生成 modified 7.18
python/sglang/jit_kernel/flash_attention_v3.py JIT 内核 modified 6.74
python/sglang/multimodal_gen/runtime/server_args.py 多模态生成 modified 6.13
python/sglang/multimodal_gen/runtime/layers/quantization/fp8.py 多模态生成 modified 5.73
python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py 多模态生成 modified 5.59
docs/diffusion/quantization.md 文档 modified 3.59
docs/diffusion/api/cli.md 文档 modified 1.89
python/sglang/multimodal_gen/runtime/layers/quantization/__init__.py 多模态生成 modified 5.39
python/sglang/multimodal_gen/runtime/loader/fsdp_load.py 多模态生成 modified 5.29

关键符号

Mxfp4Config.__init__ Mxfp4Config.get_name Mxfp4Config.get_supported_act_dtypes Mxfp4Config.get_min_capability Mxfp4Config.get_config_filenames Mxfp4Config.from_config Mxfp4Config.get_quant_method Mxfp4LinearMethod.__init__ Mxfp4LinearMethod.create_weights Mxfp4LinearMethod.process_weights_after_loading Mxfp4LinearMethod.apply Fp8Config.__init__ Fp8Config.get_quant_method FeedForward.__init__ ZImageTransformer2DBlock.__init__ resolve_transformer_quant_load_spec flash_attn_varlen_func

关键源码片段

python/sglang/multimodal_gen/runtime/models/dits/zimage.py data-contract

模型接入的关键改动:FeedForward 类新增 quant_config 和 prefix 参数以支持量化,ZImageTransformer2DBlock 在创建 FeedForward 时传递量化配置并提供了 packed_modules_mapping。

class FeedForward(nn.Module):
    def __init__(
        self,
        dim: int,
        hidden_dim: int,
        quant_config: Optional[QuantizationConfig] = None,
        prefix: str = "",
    ):
        super().__init__()
        # 使用 MergedColumnParallelLinear 融合 gate 与 up 投影( fused )
        self.w13 = MergedColumnParallelLinear(
            dim,
            [hidden_dim, hidden_dim],
            bias=False,
            gather_output=False,
            quant_config=quant_config,
            prefix=f"{prefix}.w13",
        )
        self.w2 = RowParallelLinear(
            hidden_dim,
            dim,
            bias=False,
            input_is_parallel=True,
            quant_config=quant_config,
            prefix=f"{prefix}.w2",
        )
        self.act = SiluAndMul()
​
    def forward(self, x):
        x13, _ = self.w13(x)
        x = self.act(x13)
        out, _ = self.w2(x)
        return out# 在 ZImageTransformer2DBlock.__init__ 中创建 FeedForward 时传递 quant_config:
else:
    self.feed_forward = FeedForward(
        dim=dim,
        hidden_dim=hidden_dim,
        quant_config=quant_config,
        prefix=f"{prefix}.feed_forward",
    )

评论区精华

MXFP4 层跳过阈值与 ASM kernel 精度问题 设计

BowenBao 在 mxfp4.py 中评论 'ditto',指向 `_MXFP4_MIN_OUTPUT_DIM = 256` 的设计决策:gemm_a4w4 ASM kernel 在小输出维度时精度下降,因此对该类层跳过 MXFP4 量化。

结论:保持 256 阈值,并在文档和注释中说明原因。 · 已解决

`quantization-ignored-layers` CLI 参数可能未被使用 正确性

avjves 指出 `quantization-ignored-layers` 的值被保存到 `ServerArgs.quantization_ignored_layers`,但似乎没有被传递到任何量化配置的 `ignored_layers` 字段。ColinZ22 未直接回复。

结论:未完全确认修复状态,后续提交添加了部分传递路径,但 `Mxfp4Config` 的 `from_config` 仍设置为空列表。可能仍有缺陷。 · unresolved

`_is_hip` 模块级变量的使用 style

mickqian 询问能否避免使用这些本地变量,ColinZ22 解释这是代码库中的惯例,用于缓存 `is_hip()` 结果,并引用其他文件作为证据。

结论:接受使用 `_is_hip` 模块级变量,作为代码惯例保留。 · 已解决

FP8 后端在 main 上 broken 正确性

HaiShaw 报告 FP8 路径在 main 上出问题,并 @yichiche 和 @yctseng0211。yichiche 回复已有 PR #26261 修复。@ColinZ22 建议在 `Fp8LinearMethod.apply()` 上添加 `@torch.compiler.disable` 以解决 torch compile 冲突。

结论:FP8 后端启动异常需要独立 PR 修复;torch compile 禁用作为临时解决方案。 · 已解决

文档补充与新参数帮助文本 documentation

mickqian 要求更新 cli.md 和 quantization.md,并提供更详细的 `--quantization` 帮助文本,包括自动检测逻辑和 MXFP4 的硬件要求。ColinZ22 采纳了这些建议。

结论:文档已更新,帮助文本包含完整说明。 · 已解决

风险与影响

  1. 硬件依赖性:MXFP4 要求 AMD ROCm 且 GPU 为 MI350+(gfx95x)。在不支持的硬件上使用 --quantization mxfp4 会导致运行时错误或性能下降(目前没有优雅的 fallback)。
  2. FP8 后端稳定性:HaiShaw 在 review 中指出 FP8 路径已在 main 上出问题,依赖后续 PR(#26261)修复。若用户使用 FP8 在线量化可能遇到崩溃或精度问题。
  3. quantization-ignored-layers 可能失效:avjves 发现该参数未被实际传递到量化配置的 ignored_layers 字段。如果未修复,用户指定忽略某些层时不会生效,导致这些层被错误量化(可能影响精度)。
  4. Torch Compile 兼容性:MXFP4 和 FP8 的 GEMM 内核目前无法与 torch.compile 一起工作(Inductor 无法 lower aten._scaled_mm)。ColinZ22 建议添加 @torch.compiler.disable 强制图中断,但当前多模态生成默认关闭 torch compile,后续启用时需同步修复。
  5. 精度风险:对于输出维度小于 256 的线性层,Mxfp4Config 会强制回退到未量化,这是合理的,但可能遗漏某些边界情况。
  • 用户:AMD 用户(尤其是 MI350+)现在可以享受 MXFP4 量化的显存节省和速度提升;所有用户都能使用 FP8 在线量化(尽管后端可能需要修复);不再需要预先量化模型。新 CLI 参数提供了细粒度控制。
  • 系统:增加了对 AITER 加 MXFP4 后端的依赖;Flash Attention 在非 CUDA 平台上自动回退到 FA2,提升了平台兼容性。
  • 团队:本 PR 整合了之前分散的在线量化尝试(#20922、#23373),为后续量化扩展(如 W4A8、ModelSlim)奠定了统一的 QuantizationConfig / LinearMethodBase 架构。
AMD 硬件依赖 MXFP4 仅限 ROCm/MI350+ FP8 后端兼容性风险 CLI 参数可能未使用 Torch Compile 冲突

关联 Issue

未识别关联 Issue

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

完整报告

参与讨论