From 33416f6f82cc8eeedc4dc38120882677fae2816f Mon Sep 17 00:00:00 2001 From: Dity0 Date: Sun, 9 Aug 2026 15:58:01 +0300 Subject: [PATCH] Accept SDNQConfig in post load quant --- modules/sdnq/common.py | 2 +- modules/sdnq/quantizer.py | 63 +++++++++++++++++++++------------------ 2 files changed, 35 insertions(+), 30 deletions(-) diff --git a/modules/sdnq/common.py b/modules/sdnq/common.py index 94db1cbe9..92801ce25 100644 --- a/modules/sdnq/common.py +++ b/modules/sdnq/common.py @@ -6,7 +6,7 @@ import torch from modules import shared -sdnq_version = "0.2.4" +sdnq_version = "0.2.5" sdnq_keys = {"weight", "scale", "zero_point", "svd_up", "svd_down"} torch_version = torch.__version__[:4] diff --git a/modules/sdnq/quantizer.py b/modules/sdnq/quantizer.py index e93e8b76e..a0f216538 100644 --- a/modules/sdnq/quantizer.py +++ b/modules/sdnq/quantizer.py @@ -504,11 +504,15 @@ def sdnq_post_load_quant( quantization_device: torch.device | None = None, return_device: torch.device | None = None, torch_dtype: torch.dtype | None = None, + quantization_config: "SDNQConfig" = None, pre_quantized: bool = False, ) -> torch.nn.Module: if pre_quantized: add_skip_keys = False use_dynamic_quantization = False + if quantization_config is not None: + quantization_config.add_skip_keys = False + quantization_config.use_dynamic_quantization = False elif ( hasattr(model, "quantization_config") or (hasattr(model, "config") and hasattr(model.config, "quantization_config")) @@ -516,35 +520,36 @@ def sdnq_post_load_quant( ): raise RuntimeError("Quantizing a pre-quantized model is not supported!") - quantization_config = SDNQConfig( - weights_dtype=weights_dtype, - quantized_matmul_dtype=quantized_matmul_dtype, - hadamard_group_size=hadamard_group_size, - group_size=group_size, - svd_rank=svd_rank, - svd_steps=svd_steps, - dynamic_loss_threshold=dynamic_loss_threshold, - use_svd=use_svd, - use_hadamard=use_hadamard, - quant_conv=quant_conv, - quant_embedding=quant_embedding, - use_quantized_matmul=use_quantized_matmul, - use_quantized_matmul_conv=use_quantized_matmul_conv, - use_dynamic_quantization=use_dynamic_quantization, - use_stochastic_rounding=use_stochastic_rounding, - dequantize_fp32=dequantize_fp32, - non_blocking=non_blocking, - add_skip_keys=add_skip_keys, - minimum_allowed_numel=minimum_allowed_numel, - minimum_allowed_channel_size=minimum_allowed_channel_size, - modules_to_not_convert=modules_to_not_convert, - modules_to_not_use_matmul=modules_to_not_use_matmul, - modules_dtype_dict=modules_dtype_dict, - modules_quant_config=modules_quant_config, - quantization_device=quantization_device, - return_device=return_device, - ) - if add_skip_keys: + if quantization_config is None: + quantization_config = SDNQConfig( + weights_dtype=weights_dtype, + quantized_matmul_dtype=quantized_matmul_dtype, + hadamard_group_size=hadamard_group_size, + group_size=group_size, + svd_rank=svd_rank, + svd_steps=svd_steps, + dynamic_loss_threshold=dynamic_loss_threshold, + use_svd=use_svd, + use_hadamard=use_hadamard, + quant_conv=quant_conv, + quant_embedding=quant_embedding, + use_quantized_matmul=use_quantized_matmul, + use_quantized_matmul_conv=use_quantized_matmul_conv, + use_dynamic_quantization=use_dynamic_quantization, + use_stochastic_rounding=use_stochastic_rounding, + dequantize_fp32=dequantize_fp32, + non_blocking=non_blocking, + add_skip_keys=add_skip_keys, + minimum_allowed_numel=minimum_allowed_numel, + minimum_allowed_channel_size=minimum_allowed_channel_size, + modules_to_not_convert=modules_to_not_convert, + modules_to_not_use_matmul=modules_to_not_use_matmul, + modules_dtype_dict=modules_dtype_dict, + modules_quant_config=modules_quant_config, + quantization_device=quantization_device, + return_device=return_device, + ) + if quantization_config.add_skip_keys: model, quantization_config = add_module_skip_keys(model, quantization_config) model.eval()