From d227a46406733f9aeb623edb1e4daef46ad4b55f Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 15 Jun 2026 22:04:43 +0300 Subject: [PATCH] SDNQ handle hf_quantizer and modules_to_not_use_matmul in loader --- modules/sdnq/loader.py | 73 +++++++++++++++++++++++---------------- modules/sdnq/quantizer.py | 4 +++ 2 files changed, 47 insertions(+), 30 deletions(-) diff --git a/modules/sdnq/loader.py b/modules/sdnq/loader.py index 43cf06166..1faff51d6 100644 --- a/modules/sdnq/loader.py +++ b/modules/sdnq/loader.py @@ -3,8 +3,8 @@ import json import torch from diffusers.models.modeling_utils import ModelMixin -from .common import dtype_dict, is_fp8_mm_supported, use_tensorwise_fp8_matmul, check_torch_compile, conv_types, linear_types -from .quantizer import QuantizationMethod, SDNQConfig, sdnq_post_load_quant +from .common import dtype_dict, is_fp8_mm_supported, use_tensorwise_fp8_matmul, check_torch_compile, linear_types +from .quantizer import QuantizationMethod, SDNQConfig, SDNQQuantizer, sdnq_post_load_quant from .quant_utils import prepare_weight_for_matmul, prepare_svd_for_matmul from .utils import get_quant_args_from_config, check_param_name_in from .forward import get_forward_func @@ -168,6 +168,10 @@ def load_sdnq_model(model_path: str, model_cls: ModelMixin | None = None, file_n model.config["quantization_config"] = quantization_config.to_dict() except Exception: pass + if hasattr(model, "hf_quantizer"): + model.hf_quantizer.quantization_config = quantization_config + else: + model.hf_quantizer = SDNQQuantizer(quantization_config) model = post_process_model(model) if (dtype is not None) or (dequantize_fp32 is not None) or (use_quantized_matmul is not None): @@ -209,25 +213,21 @@ def apply_sdnq_options_to_module(model, quantization_config: SDNQConfig, dtype: else: param_name = module_name if hasattr(module, "sdnq_dequantizer"): + param_name = param_name + ".weight" layer_class_name = module.original_class.__name__ current_use_quantized_matmul = use_quantized_matmul - if layer_class_name in conv_types: - current_use_quantized_matmul = None - elif check_param_name_in(param_name, quantization_config.modules_to_not_use_matmul) is not None: - current_use_quantized_matmul = None - if not is_fp8_mm_supported and module.sdnq_dequantizer.quantized_matmul_dtype in {"fp8", "float8_e4m3fn"}: - current_use_quantized_matmul = False - - if current_use_quantized_matmul: - if layer_class_name in conv_types: - output_channel_size, channel_size = module.sdnq_dequantizer.original_shape[:2] - elif layer_class_name in linear_types: - output_channel_size, channel_size = module.sdnq_dequantizer.original_shape - else: + if layer_class_name in linear_types: + if not is_fp8_mm_supported and module.sdnq_dequantizer.quantized_matmul_dtype in {"fp8", "float8_e4m3fn"}: current_use_quantized_matmul = False - current_use_quantized_matmul = current_use_quantized_matmul and channel_size >= 32 and output_channel_size >= 32 # pylint: disable=possibly-used-before-assignment - current_use_quantized_matmul = current_use_quantized_matmul and output_channel_size % 16 == 0 and channel_size % 16 == 0 # pylint: disable=possibly-used-before-assignment + elif check_param_name_in(param_name, quantization_config.modules_to_not_use_matmul) is not None: + current_use_quantized_matmul = None + if current_use_quantized_matmul: + output_channel_size, channel_size = module.sdnq_dequantizer.original_shape + current_use_quantized_matmul = current_use_quantized_matmul and channel_size >= 32 and output_channel_size >= 32 + current_use_quantized_matmul = current_use_quantized_matmul and output_channel_size % 16 == 0 and channel_size % 16 == 0 + else: + current_use_quantized_matmul = None if dtype is not None and module.sdnq_dequantizer.result_dtype not in {torch.float32, torch.float64}: module.sdnq_dequantizer.result_dtype = dtype @@ -259,19 +259,27 @@ def apply_sdnq_options_to_module(model, quantization_config: SDNQConfig, dtype: if module.zero_point is not None: module.zero_point.data = module.zero_point.to(dtype=scale_dtype) - if current_use_quantized_matmul is not None and current_use_quantized_matmul != module.sdnq_dequantizer.use_quantized_matmul: - if not module.sdnq_dequantizer.re_quantize_for_matmul and not dtype_dict[module.sdnq_dequantizer.weights_dtype]["is_packed"]: - module.scale.t_() - module.weight.t_() - if current_use_quantized_matmul: - module.weight.data = prepare_weight_for_matmul(module.weight) - else: - module.scale.data = module.scale.contiguous() - module.weight.data = module.weight.contiguous() - if module.svd_up is not None: - module.svd_up.data, module.svd_down.data = prepare_svd_for_matmul(module.svd_up.t_(), module.svd_down.t_(), current_use_quantized_matmul) - module.sdnq_dequantizer.use_quantized_matmul = current_use_quantized_matmul - module.forward_func = get_forward_func(module.original_class.__name__, module.sdnq_dequantizer.quantized_matmul_dtype, current_use_quantized_matmul) + if current_use_quantized_matmul is not None: + if current_use_quantized_matmul != module.sdnq_dequantizer.use_quantized_matmul: + if not module.sdnq_dequantizer.re_quantize_for_matmul and not dtype_dict[module.sdnq_dequantizer.weights_dtype]["is_packed"]: + module.scale.t_() + module.weight.t_() + if current_use_quantized_matmul: + module.weight.data = prepare_weight_for_matmul(module.weight) + else: + module.scale.data = module.scale.contiguous() + module.weight.data = module.weight.contiguous() + if module.svd_up is not None: + module.svd_up.data, module.svd_down.data = prepare_svd_for_matmul(module.svd_up.t_(), module.svd_down.t_(), current_use_quantized_matmul) + module.sdnq_dequantizer.use_quantized_matmul = current_use_quantized_matmul + module.forward_func = get_forward_func(module.original_class.__name__, module.sdnq_dequantizer.quantized_matmul_dtype, current_use_quantized_matmul) + if ( + not module.sdnq_dequantizer.use_quantized_matmul + and (use_quantized_matmul or (use_quantized_matmul is None and quantization_config.use_quantized_matmul)) + and check_param_name_in(param_name, quantization_config.modules_to_not_use_matmul) is None + ): + quantization_config.modules_to_not_use_matmul.append(param_name) + setattr(model, module_name, module) else: setattr(model, module_name, apply_sdnq_options_to_module(module, quantization_config, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul, full_param_name=param_name)) @@ -304,4 +312,9 @@ def apply_sdnq_options_to_model(model, dtype: torch.dtype | None = None, dequant model.config["quantization_config"].dequantize_fp32 = dequantize_fp32 except Exception: pass + if hasattr(model, "hf_quantizer"): + if use_quantized_matmul is not None: + model.hf_quantizer.quantization_config.use_quantized_matmul = use_quantized_matmul + if dequantize_fp32 is not None: + model.hf_quantizer.quantization_config.dequantize_fp32 = dequantize_fp32 return model diff --git a/modules/sdnq/quantizer.py b/modules/sdnq/quantizer.py index 70bda2a8d..28eb744a5 100644 --- a/modules/sdnq/quantizer.py +++ b/modules/sdnq/quantizer.py @@ -748,6 +748,10 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): def is_serializable(self, *args, **kwargs) -> bool: # pylint: disable=unused-argument, invalid-overridden-method return not self.quantization_config.is_training + @property + def supports_safetensors_serialization(self) -> bool: + return self.is_serializable() + @property def is_trainable(self): return self.quantization_config.is_training