mirror of
https://github.com/NVIDIA/Model-Optimizer.git
synced 2026-10-02 03:14:52 +08:00
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:
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user