From 23f2deaa584e7e5a494936766129045be18e8d5e Mon Sep 17 00:00:00 2001 From: Disty0 Date: Mon, 6 Oct 2025 02:04:28 +0300 Subject: [PATCH] fix enable_quantized_mamtul --- modules/sdnq/dequantizer.py | 11 ++++++++++- modules/sdnq/loader.py | 18 +++++++++--------- 2 files changed, 19 insertions(+), 10 deletions(-) diff --git a/modules/sdnq/dequantizer.py b/modules/sdnq/dequantizer.py index 72a058555..38f09f0a1 100644 --- a/modules/sdnq/dequantizer.py +++ b/modules/sdnq/dequantizer.py @@ -4,7 +4,7 @@ from typing import Tuple, Optional import torch -from .common import dtype_dict, compile_func +from .common import dtype_dict, compile_func, use_tensorwise_fp8_matmul from .packed_int import pack_int_symetric, unpack_int_symetric, pack_int_asymetric, unpack_int_asymetric @@ -67,6 +67,15 @@ def re_quantize_int8(weight: torch.FloatTensor) -> Tuple[torch.CharTensor, torch return weight, scale +def re_quantize_fp8(weight: torch.FloatTensor) -> Tuple[torch.CharTensor, torch.FloatTensor]: + if weight.ndim > 2: # convs + weight = weight.flatten(1,-1) + weight, scale = quantize_fp8(weight.t(), dim=0) + if not use_tensorwise_fp8_matmul: + scale = scale.to(dtype=torch.float32) + return weight, scale + + def re_quantize_matmul_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, result_shape: torch.Size, svd_up: Optional[torch.FloatTensor] = None, svd_down: Optional[torch.FloatTensor] = None) -> Tuple[torch.CharTensor, torch.FloatTensor]: return re_quantize_int8(dequantize_asymmetric(weight, scale, zero_point, scale.dtype, result_shape, svd_up=svd_up, svd_down=svd_down)) diff --git a/modules/sdnq/loader.py b/modules/sdnq/loader.py index 6a625ba17..aabd2496c 100644 --- a/modules/sdnq/loader.py +++ b/modules/sdnq/loader.py @@ -4,9 +4,9 @@ import torch from safetensors import safe_open from diffusers.models.modeling_utils import ModelMixin +from .common import use_contiguous_mm from .quantizer import SDNQConfig, apply_sdnq_to_module -from .common import use_contiguous_mm, use_tensorwise_fp8_matmul -from .dequantizer import dequantize_symmetric_compiled, quantize_fp8 +from .dequantizer import dequantize_symmetric_compiled, re_quantize_int8, re_quantize_fp8 def save_sdnq_model(model: ModelMixin, sdnq_config: SDNQConfig, model_path: str, max_shard_size: str = "10GB") -> None: @@ -45,14 +45,14 @@ def enable_quantized_mamtul(model): for module in model.children(): if hasattr(module, "sdnq_dequantizer"): if not module.sdnq_dequantizer.use_quantized_matmul: - if module.sdnq_dequantizer.weights_dtype in {"int8", "float8_e4m3fn"} and module.sdnq_dequantizer.result_shape != module.weight.shape: - if module.sdnq_dequantizer.weights_dtype == "int8": - module.weight.data, module.scale.data = module.sdnq_dequantizer.re_quantize_matmul(module.weight, module.scale, module.zero_point, None, None) - elif module.sdnq_dequantizer.weights_dtype == "float8_e4m3fn": - module.weight.data, module.scale.data = quantize_fp8(dequantize_symmetric_compiled(module.weight, module.scale, module.sdnq_dequantizer.result_dtype, module.sdnq_dequantizer.result_shape)) + if module.sdnq_dequantizer.weights_dtype in {"int8", "float8_e4m3fn"}: + if module.sdnq_dequantizer.re_quantize_for_matmul: + if module.sdnq_dequantizer.weights_dtype == "int8": + module.weight.data, module.scale.data = re_quantize_int8(dequantize_symmetric_compiled(module.weight, module.scale, torch.float32, module.sdnq_dequantizer.result_shape)) + else: + module.weight.data, module.scale.data = re_quantize_fp8(dequantize_symmetric_compiled(module.weight, module.scale, torch.float32, module.sdnq_dequantizer.result_shape)) + else: module.weight.data, module.scale.data = module.weight.t_(), module.scale.t_() - if not use_tensorwise_fp8_matmul: - module.scale.data = module.scale.to(dtype=torch.float32) if use_contiguous_mm: module.weight.data = module.weight.contiguous() elif module.weight.is_contiguous():