Support attention quantization for diffusers >= 0.35.0 (#608)

## What does this PR do?

**Type of change:**

new feature

**Overview:** ?

Attention mechanism has changed from diffusers 0.35.

Many model attentions are now subclass of a new Mixin class:
AttentionModuleMixin, which is not a sub class of Attention

To fix it, patch the mixin class by forcing to use native attention
impl so the existing function monkey patch still work.


## Testing

manual quant of Wan, Flux

---------

Signed-off-by: Shengliang Xu <shengliangx@nvidia.com>
This commit is contained in:
Shengliang Xu
2025-12-03 11:20:48 -08:00
committed by GitHub
parent 1524251cf1
commit 8844a2b2bc
3 changed files with 45 additions and 16 deletions
+5 -5
View File
@@ -760,8 +760,6 @@ class Quantizer:
self.logger.info("Disabling specific quantizers...")
mtq.disable_quantizer(backbone, model_filter_func)
mtq.print_quant_summary(backbone)
self.logger.info("Quantization completed successfully")
@@ -816,7 +814,6 @@ class ExportManager:
backbone: torch.nn.Module,
model_type: ModelType,
quant_format: QuantFormat,
quantize_mha: bool,
) -> None:
"""
Export model to ONNX format.
@@ -831,7 +828,6 @@ class ExportManager:
return
self.logger.info(f"Starting ONNX export to {self.config.onnx_dir}")
check_conv_and_mha(backbone, quant_format == QuantFormat.FP4, quantize_mha)
if quant_format == QuantFormat.FP8 and self._has_conv_layers(backbone):
self.logger.info(
@@ -1118,12 +1114,16 @@ def main() -> None:
export_manager.save_checkpoint(backbone)
check_conv_and_mha(
backbone, quant_config.format == QuantFormat.FP4, quant_config.quantize_mha
)
mtq.print_quant_summary(backbone)
export_manager.export_onnx(
pipe,
backbone,
model_config.model_type,
quant_config.format,
quantize_mha=quant_config.quantize_mha,
)
logger.info(
f"Quantization process completed successfully! Time taken = {time.time() - s} seconds"
+14 -10
View File
@@ -25,6 +25,7 @@ from diffusers.models.lora import LoRACompatibleConv, LoRACompatibleLinear
from diffusers.utils import load_image
import modelopt.torch.quantization as mtq
from modelopt.torch.quantization.plugins.diffusers import AttentionModuleMixin
USE_PEFT = True
try:
@@ -44,21 +45,24 @@ def filter_func_default(name: str) -> bool:
def check_conv_and_mha(backbone, if_fp4, quantize_mha):
for _, module in backbone.named_modules():
for name, module in backbone.named_modules():
if isinstance(module, (torch.nn.Conv1d, torch.nn.Conv2d, torch.nn.Conv3d)) and if_fp4:
module.weight_quantizer.disable()
module.input_quantizer.disable()
elif isinstance(module, Attention):
if not quantize_mha:
continue
print(f"Disabled NVFP4 Conv layer quantization for layer {name}")
elif isinstance(module, (Attention, AttentionModuleMixin)):
head_size = int(module.inner_dim / module.heads)
module.q_bmm_quantizer.disable()
module.k_bmm_quantizer.disable()
module.v_bmm_quantizer.disable()
module.softmax_quantizer.disable()
module.bmm2_output_quantizer.disable()
if head_size % 16 != 0:
if not quantize_mha or head_size % 16 != 0:
module.q_bmm_quantizer.disable()
module.k_bmm_quantizer.disable()
module.v_bmm_quantizer.disable()
module.softmax_quantizer.disable()
module.bmm2_output_quantizer.disable()
setattr(module, "_disable_fp8_mha", True)
print(f"Disabled Attention layer quantization for layer {name}")
else:
setattr(module, "_disable_fp8_mha", False)