SDNQ post process pre-quants after load

This commit is contained in:
Disty0
2025-12-08 01:08:53 +03:00
parent 0835ca6f66
commit 6e05a12a49
2 changed files with 7 additions and 0 deletions
+4
View File
@@ -146,9 +146,13 @@ def post_process_model(model):
return model
for module_name, module in model.named_children():
if hasattr(module, "sdnq_dequantizer"):
module.weight.requires_grad_(False)
if module.sdnq_dequantizer.use_quantized_matmul and not module.sdnq_dequantizer.re_quantize_for_matmul:
module.weight.data = prepare_weight_for_matmul(module.weight)
if module.zero_point is not None:
module.zero_point.requires_grad_(False)
if module.svd_up is not None:
module.svd_up, module.svd_down = module.svd_up.requires_grad_(False), module.svd_down.requires_grad_(False)
module.svd_up.data, module.svd_down.data = prepare_svd_for_matmul(module.svd_up, module.svd_down, module.sdnq_dequantizer.use_quantized_matmul)
setattr(model, module_name, module)
else:
+3
View File
@@ -793,6 +793,9 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
model.quantization_method = QuantizationMethod.SDNQ
def _process_model_after_weight_loading(self, model, **kwargs): # pylint: disable=unused-argument
if self.pre_quantized:
from .loader import post_process_model
model = post_process_model(model)
if self.quantization_config.is_training:
from .training import convert_sdnq_model_to_training
model = convert_sdnq_model_to_training(