执行摘要
- 一句话:为多模态生成添加在线 MXFP4/FP8 量化支持
- 推荐动作:该 PR 值得认真阅读,特别是对量化扩展有兴趣的开发者。
Mxfp4Config 和 Mxfp4LinearMethod 的设计模式(继承 QuantizationConfig、LinearMethodBase)可作为后续添加新量化方法的参考。Flash Attention 的 FA2 回退处理也是一个良好的兼容性示范。建议重点关注 quantization-ignored-layers 的传递正确性以及 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 格式。
实现拆解
- 新增 MXFP4 量化配置与方法:在
mxfp4.py 中实现 Mxfp4Config(继承 QuantizationConfig)和 Mxfp4LinearMethod(继承 LinearMethodBase)。Mxfp4Config 负责判断层是否应该跳过量化(基于 ignored_layers 和输出维度阈值 _MXFP4_MIN_OUTPUT_DIM),并为每个线性层分配合适的 LinearMethodBase。Mxfp4LinearMethod 在 process_weights_after_loading() 中调用 AITER 的 dynamic_mxfp4_quant 将 BF16/FP16 权重在线量化为 MXFP4,并在 apply() 中使用 gemm_a4w4 或 shuffle_weight 执行量化 GEMM。
- 扩展 FP8 量化配置:在
fp8.py 的 Fp8Config 中增加 packed_modules_mapping 参数,并在 get_quant_method() 中依赖 is_layer_skipped 处理融合层映射,确保 FP8 在线量化也能正确跳过指定层。
- 模型接入调整:修改
zimage.py 中的 FeedForward 类,使其构造函数接受 quant_config 和 prefix,并将量化配置传递给内部的线性层。同时为 ZImageTransformer2DBlock 添加 packed_modules_mapping 类属性,用于描述融合参数(如 w13 对应 w1 和 w3),供 is_layer_skipped 正确处理参数名模式。
- 加载器与配置提升:在
transformer_load_utils.py 的 resolve_transformer_quant_load_spec() 中,从模型类获取 packed_modules_mapping 并注入到量化配置对象。在 server_args.py 中添加 quantization 和 quantization_ignored_layers 两个 CLI 参数,它们的值会传递给量化后端。
- Flash Attention 兼容性:在
flash_attention_v3.py 中,当 sgl-kernel 的 FA3 不可用时(如 ROCm 或 sm<90),不再直接抛异常,而是尝试调用 flash_attn 包的 FA2 实现,保证多模态生成在其他平台上也能运行。
关键文件:
python/sglang/multimodal_gen/runtime/layers/quantization/mxfp4.py(模块 多模态生成;类别 source;类型 dependency-wiring;符号 Mxfp4Config, init, get_name, get_supported_act_dtypes): 核心新增文件,实现了 MXFP4 在线量化的配置类(Mxfp4Config)和 LinearMethod(Mxfp4LinearMethod),包含了动态量化 logic、硬件检测、层跳过逻辑以及与 AITER 内核的集成。
python/sglang/multimodal_gen/runtime/models/dits/zimage.py(模块 多模态生成;类别 source;类型 data-contract;符号 init): 模型接入的关键改动:FeedForward 类新增 quant_config 和 prefix 参数以支持量化,ZImageTransformer2DBlock 在创建 FeedForward 时传递量化配置并提供了 packed_modules_mapping。
python/sglang/jit_kernel/flash_attention_v3.py(模块 JIT 内核;类别 source;类型 dependency-wiring): Flash Attention V3 函数的兼容性增强:当 sgl-kernel FA3 不支持时(如 ROCm 或 SM<90),自动回退到 flash_attn 包的 FA2 实现,避免多模态生成在非 CUDA 平台上崩溃。
python/sglang/multimodal_gen/runtime/server_args.py(模块 多模态生成;类别 source;类型 core-logic): 新增 --quantization 和 --quantization-ignored-layers 两个 CLI 参数的定义和帮助文本,用户通过它们控制在线量化的启用与层排除。
python/sglang/multimodal_gen/runtime/layers/quantization/fp8.py(模块 多模态生成;类别 source;类型 core-logic): 扩展 Fp8Config 以支持 packed_modules_mapping,并修改 is_layer_skipped 使用 fused_mapping,使得 FP8 在线量化能正确处理融合层参数名。
python/sglang/multimodal_gen/runtime/loader/transformer_load_utils.py(模块 多模态生成;类别 source;类型 core-logic): 在 resolve_transformer_quant_load_spec 中从模型类获取 packed_modules_mapping 并注入量化配置,确保量化后端能正确识别融合层。
docs/diffusion/quantization.md(模块 文档;类别 docs;类型 documentation): 文档更新,增加了 MXFP4 在线量化、新 CLI 参数的用法说明和硬件要求。
docs/diffusion/api/cli.md(模块 文档;类别 docs;类型 documentation): 文档更新,提及了新服务器参数。
python/sglang/multimodal_gen/runtime/layers/quantization/__init__.py(模块 多模态生成;类别 source;类型 dependency-wiring): 导入了 mxfp4 模块,使量化配置系统能发现 Mxfp4Config。
python/sglang/multimodal_gen/runtime/loader/fsdp_load.py(模块 多模态生成;类别 source;类型 core-logic): 添加了 'weight_scale' 和 'input_scale' 到忽略的 FSDP key 列表,避免 FP8/MXFP4 量化 scale 参数被错误地视为状态 dict 的一部分。
关键符号: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
模型接入的关键改动: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 要求 AMD ROCm 且 GPU 为 MI350+(gfx95x)。在不支持的硬件上使用
--quantization mxfp4 会导致运行时错误或性能下降(目前没有优雅的 fallback)。
- FP8 后端稳定性:HaiShaw 在 review 中指出 FP8 路径已在 main 上出问题,依赖后续 PR(#26261)修复。若用户使用 FP8 在线量化可能遇到崩溃或精度问题。
quantization-ignored-layers 可能失效:avjves 发现该参数未被实际传递到量化配置的 ignored_layers 字段。如果未修复,用户指定忽略某些层时不会生效,导致这些层被错误量化(可能影响精度)。
- Torch Compile 兼容性:MXFP4 和 FP8 的 GEMM 内核目前无法与
torch.compile 一起工作(Inductor 无法 lower aten._scaled_mm)。ColinZ22 建议添加 @torch.compiler.disable 强制图中断,但当前多模态生成默认关闭 torch compile,后续启用时需同步修复。
- 精度风险:对于输出维度小于 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 冲突
关联脉络
- PR #20922 [Diffusion] Online quantization support (FP8): 为扩散模型添加了初始的在线量化支持(FP8),但缺少 MXFP4。本 PR 在此基础上升级并覆盖了 MXFP4 功能。
- PR #23373 [AMD] MXFP4 quantization for diffusion: 另一个独立的 MXFP4 量化 PR,BowenBao 评论说本 PR 包含了其超集,建议合并本 PR 取代 #23373。
- PR #26261 [AMD] Fix aiter FP8 backend: 修复 FP8 量化后端在 main 上 broken 的问题,与本次 PR 中的 FP8 路径直接相关。
参与讨论