diff --git a/modules/sdnq/common.py b/modules/sdnq/common.py index 483e1dd9a..aaee862ee 100644 --- a/modules/sdnq/common.py +++ b/modules/sdnq/common.py @@ -337,7 +337,7 @@ weights_dtype_order = [ use_torch_compile = shared.opts.sdnq_dequantize_compile # this setting requires a full restart of the webui to apply -def check_torch_compile(): # dynamo can be disabled after startup +def check_torch_compile() -> bool: # dynamo can be disabled after startup return use_torch_compile and not torch._dynamo.config.disable # pylint: disable=protected-access diff --git a/modules/sdnq/dequantizer.py b/modules/sdnq/dequantizer.py index afaf2781c..dd22c6133 100644 --- a/modules/sdnq/dequantizer.py +++ b/modules/sdnq/dequantizer.py @@ -13,7 +13,17 @@ from .layers import SDNQLayer @devices.inference_context() -def dequantize_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False) -> torch.FloatTensor: +def dequantize_asymmetric( + weight: torch.Tensor, + scale: torch.FloatTensor, + zero_point: torch.FloatTensor, + svd_up: torch.FloatTensor | None = None, + svd_down: torch.FloatTensor | None = None, + hadamard: torch.FloatTensor | None = None, + dtype: torch.dtype = None, + result_shape: torch.Size = None, + skip_quantized_matmul: bool = False, +) -> torch.FloatTensor: result = torch.addcmul(zero_point, weight.to(dtype=scale.dtype), scale) if result_shape is not None: result = result.view(result_shape) @@ -37,7 +47,17 @@ def dequantize_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, ze @devices.inference_context() -def dequantize_symmetric(weight: torch.CharTensor, scale: torch.FloatTensor, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False, re_quantize_for_matmul: bool = False) -> torch.FloatTensor: +def dequantize_symmetric( + weight: torch.Tensor, + scale: torch.FloatTensor, + svd_up: torch.FloatTensor | None = None, + svd_down: torch.FloatTensor | None = None, + hadamard: torch.FloatTensor | None = None, + dtype: torch.dtype = None, + result_shape: torch.Size = None, + skip_quantized_matmul: bool = False, + re_quantize_for_matmul: bool = False, +) -> torch.FloatTensor: result = weight.to(dtype=scale.dtype).mul_(scale) if skip_quantized_matmul and not re_quantize_for_matmul: result.t_() @@ -63,7 +83,7 @@ def dequantize_symmetric(weight: torch.CharTensor, scale: torch.FloatTensor, svd @devices.inference_context() -def dequantize_symmetric_with_bias(weight: torch.CharTensor, scale: torch.FloatTensor, bias: torch.FloatTensor, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None) -> torch.FloatTensor: +def dequantize_symmetric_with_bias(weight: torch.Tensor, scale: torch.FloatTensor, bias: torch.FloatTensor, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None) -> torch.FloatTensor: if hadamard is not None: result = rotate_hadamard(weight.to(dtype=scale.dtype).mul_(scale), hadamard=hadamard).add_(bias) else: @@ -76,22 +96,22 @@ def dequantize_symmetric_with_bias(weight: torch.CharTensor, scale: torch.FloatT @devices.inference_context() -def dequantize_packed_int_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False) -> torch.FloatTensor: +def dequantize_packed_int_asymmetric(weight: torch.Tensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False) -> torch.FloatTensor: return dequantize_asymmetric(unpack_int(weight, weights_dtype, shape), scale, zero_point, svd_up=svd_up, svd_down=svd_down, hadamard=hadamard, dtype=dtype, result_shape=result_shape, skip_quantized_matmul=skip_quantized_matmul) @devices.inference_context() -def dequantize_packed_int_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False, re_quantize_for_matmul: bool = False) -> torch.FloatTensor: +def dequantize_packed_int_symmetric(weight: torch.Tensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False, re_quantize_for_matmul: bool = False) -> torch.FloatTensor: return dequantize_symmetric(unpack_int(weight, weights_dtype, shape, dtype=scale.dtype), scale, svd_up=svd_up, svd_down=svd_down, hadamard=hadamard, dtype=dtype, result_shape=result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul) @devices.inference_context() -def dequantize_packed_float_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False) -> torch.FloatTensor: +def dequantize_packed_float_asymmetric(weight: torch.Tensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False) -> torch.FloatTensor: return dequantize_asymmetric(unpack_float(weight, weights_dtype, shape), scale, zero_point, svd_up=svd_up, svd_down=svd_down, hadamard=hadamard, dtype=dtype, result_shape=result_shape, skip_quantized_matmul=skip_quantized_matmul) @devices.inference_context() -def dequantize_packed_float_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False, re_quantize_for_matmul: bool = False) -> torch.FloatTensor: +def dequantize_packed_float_symmetric(weight: torch.Tensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None, dtype: torch.dtype = None, result_shape: torch.Size = None, skip_quantized_matmul: bool = False, re_quantize_for_matmul: bool = False) -> torch.FloatTensor: return dequantize_symmetric(unpack_float(weight, weights_dtype, shape), scale, svd_up=svd_up, svd_down=svd_down, hadamard=hadamard, dtype=dtype, result_shape=result_shape, skip_quantized_matmul=skip_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul) @@ -119,7 +139,7 @@ def re_quantize_fp_mm(weight: torch.FloatTensor, matmul_dtype: str = "float8_e4m @devices.inference_context() -def re_quantize_matmul_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, matmul_dtype: str, result_shape: torch.Size = None, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: +def re_quantize_matmul_asymmetric(weight: torch.Tensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, matmul_dtype: str, result_shape: torch.Size = None, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: weight = dequantize_asymmetric(weight, scale, zero_point, svd_up=svd_up, svd_down=svd_down, hadamard=hadamard, dtype=scale.dtype, result_shape=result_shape) if dtype_dict[matmul_dtype]["is_integer"]: return re_quantize_int_mm(weight) @@ -128,7 +148,7 @@ def re_quantize_matmul_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTe @devices.inference_context() -def re_quantize_matmul_symmetric(weight: torch.CharTensor, scale: torch.FloatTensor, matmul_dtype: str, result_shape: torch.Size = None, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: +def re_quantize_matmul_symmetric(weight: torch.Tensor, scale: torch.FloatTensor, matmul_dtype: str, result_shape: torch.Size = None, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: weight = dequantize_symmetric(weight, scale, svd_up=svd_up, svd_down=svd_down, hadamard=hadamard, dtype=scale.dtype, result_shape=result_shape) if dtype_dict[matmul_dtype]["is_integer"]: return re_quantize_int_mm(weight) @@ -137,27 +157,27 @@ def re_quantize_matmul_symmetric(weight: torch.CharTensor, scale: torch.FloatTen @devices.inference_context() -def re_quantize_matmul_packed_int_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: +def re_quantize_matmul_packed_int_asymmetric(weight: torch.Tensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: return re_quantize_matmul_asymmetric(unpack_int(weight, weights_dtype, shape), scale, zero_point, matmul_dtype, svd_up=svd_up, svd_down=svd_down, hadamard=hadamard, result_shape=result_shape) @devices.inference_context() -def re_quantize_matmul_packed_int_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size = None, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: +def re_quantize_matmul_packed_int_symmetric(weight: torch.Tensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size = None, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: return re_quantize_matmul_symmetric(unpack_int(weight, weights_dtype, shape, dtype=scale.dtype), scale, matmul_dtype, svd_up=svd_up, svd_down=svd_down, hadamard=hadamard, result_shape=result_shape) @devices.inference_context() -def re_quantize_matmul_packed_float_asymmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: +def re_quantize_matmul_packed_float_asymmetric(weight: torch.Tensor, scale: torch.FloatTensor, zero_point: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: return re_quantize_matmul_asymmetric(unpack_float(weight, weights_dtype, shape), scale, zero_point, matmul_dtype, svd_up=svd_up, svd_down=svd_down, hadamard=hadamard, result_shape=result_shape) @devices.inference_context() -def re_quantize_matmul_packed_float_symmetric(weight: torch.ByteTensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size = None, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: +def re_quantize_matmul_packed_float_symmetric(weight: torch.Tensor, scale: torch.FloatTensor, shape: torch.Size, weights_dtype: str, matmul_dtype: str, result_shape: torch.Size = None, svd_up: torch.FloatTensor | None = None, svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None) -> tuple[torch.Tensor, torch.FloatTensor]: return re_quantize_matmul_symmetric(unpack_float(weight, weights_dtype, shape), scale, matmul_dtype, svd_up=svd_up, svd_down=svd_down, hadamard=hadamard, result_shape=result_shape) @devices.inference_context() -def dequantize_sdnq_module(model: torch.nn.Module): +def dequantize_sdnq_module(model: torch.nn.Module) -> torch.nn.Module: if isinstance(model, SDNQLayer): model = model.dequantize() has_children = list(model.children()) @@ -172,7 +192,7 @@ def dequantize_sdnq_module(model: torch.nn.Module): @devices.inference_context() -def dequantize_sdnq_model(model: torch.nn.Module): +def dequantize_sdnq_model(model: torch.nn.Module) -> torch.nn.Module: model = dequantize_sdnq_module(model) if hasattr(model, "quantization_method"): del model.quantization_method @@ -265,7 +285,7 @@ class SDNQDequantizer: svd_down: torch.FloatTensor | None = None, hadamard: torch.FloatTensor | None = None, non_hadamard: bool = True, - ): # pylint: disable=unused-argument + ) -> tuple[torch.Tensor, torch.FloatTensor]: # pylint: disable=unused-argument if hadamard is None and self.use_hadamard and not non_hadamard: hadamard = get_hadamard(self.hadamard_group_size, dtype=self.result_dtype, device=weight.device) if self.is_packed: @@ -298,7 +318,7 @@ class SDNQDequantizer: non_hadamard: bool = False, skip_compile: bool = False, dtype: torch.dtype = None, - ): # pylint: disable=unused-argument + ) -> torch.FloatTensor: # pylint: disable=unused-argument if dtype is None: dtype = self.result_dtype if hadamard is None and self.use_hadamard and not non_hadamard: diff --git a/modules/sdnq/kernels/triton_atten.py b/modules/sdnq/kernels/triton_atten.py index b54b1253e..8adec1acf 100644 --- a/modules/sdnq/kernels/triton_atten.py +++ b/modules/sdnq/kernels/triton_atten.py @@ -46,7 +46,7 @@ def sdnq_attn_kernel( q_dtype: tl.constexpr, v_dtype: tl.constexpr, out_dtype: tl.constexpr, mask_dtype: tl.constexpr, BLOCK_SIZE_M: tl.constexpr, BLOCK_SIZE_N: tl.constexpr, -): # pylint: disable=unused-argument +) -> None: # pylint: disable=unused-argument start_m = tl.program_id(0) off_h = tl.program_id(1) off_z = tl.program_id(2) @@ -190,7 +190,7 @@ def quantize_attn( hadamard_group_size: int = 256, matmul_dtype: str = "int8", pv_matmul_dtype: str | None = None, -): +) -> tuple[torch.Tensor]: if matmul_dtype == "auto": matmul_dtype = "int8" if scale is None: @@ -242,7 +242,7 @@ def get_attn_inputs( pv_matmul_dtype: str | None = None, do_quantize: bool = True, out_dtype: torch.dtype | None = None, -) -> torch.FloatTensor: +) -> tuple[torch.Tensor, float, torch.dtype]: QZ, QH, QN, QHD = query.shape _, _, KN, KHD = key.shape _, _, _, VHD = value.shape diff --git a/modules/sdnq/kernels/triton_mm.py b/modules/sdnq/kernels/triton_mm.py index eb73e96fe..bbdd7a692 100644 --- a/modules/sdnq/kernels/triton_mm.py +++ b/modules/sdnq/kernels/triton_mm.py @@ -46,7 +46,7 @@ def triton_mm_kernel( BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, -): # pylint: disable=unused-argument +) -> None: # pylint: disable=unused-argument pid = tl.program_id(axis=0) num_pid_m: tl.constexpr = tl.cdiv(M, BLOCK_SIZE_M) num_pid_n: tl.constexpr = tl.cdiv(N, BLOCK_SIZE_N) @@ -113,7 +113,7 @@ def triton_mm_td_kernel( BLOCK_SIZE_N: tl.constexpr, BLOCK_SIZE_K: tl.constexpr, GROUP_SIZE_M: tl.constexpr, -): # pylint: disable=unused-argument +) -> None: # pylint: disable=unused-argument pid = tl.program_id(axis=0) num_pid_m: tl.constexpr = tl.cdiv(M, BLOCK_SIZE_M) num_pid_n: tl.constexpr = tl.cdiv(N, BLOCK_SIZE_N) diff --git a/modules/sdnq/layers/__init__.py b/modules/sdnq/layers/__init__.py index d1c7adcf0..6835d059e 100644 --- a/modules/sdnq/layers/__init__.py +++ b/modules/sdnq/layers/__init__.py @@ -1,8 +1,10 @@ +from collections.abc import Callable + import torch class SDNQLayer(torch.nn.Module): - def __init__(self, original_layer, forward_func): + def __init__(self, original_layer: torch.nn.Module, forward_func: Callable): torch.nn.Module.__init__(self) for key, value in original_layer.__dict__.items(): if key not in {"forward", "forward_func", "original_class", "state_dict", "load_state_dict"}: @@ -11,7 +13,7 @@ class SDNQLayer(torch.nn.Module): self.forward_func = forward_func @property - def dtype(self: torch.nn.Module): + def dtype(self: torch.nn.Module) -> torch.dtype: return self.sdnq_dequantizer.result_dtype if hasattr(self, "sdnq_dequantizer") else self.weight.dtype def dequantize(self: torch.nn.Module): @@ -27,7 +29,7 @@ class SDNQLayer(torch.nn.Module): def forward(self, *args, **kwargs) -> torch.Tensor: return self.forward_func(self, *args, **kwargs) - def __repr__(self): + def __repr__(self) -> str: return f"{self.__class__.__name__}(original_class={self.original_class} forward_func={self.forward_func} sdnq_dequantizer={repr(getattr(self, 'sdnq_dequantizer', None))})" @@ -67,7 +69,7 @@ torch.serialization.add_safe_globals([SDNQConvTranspose2d]) torch.serialization.add_safe_globals([SDNQConvTranspose3d]) -def get_sdnq_wrapper_class(original_layer, forward_func): +def get_sdnq_wrapper_class(original_layer: torch.nn.Module, forward_func: Callable) -> SDNQLayer: match original_layer.__class__.__name__: case "Linear": return SDNQLinear(original_layer, forward_func) diff --git a/modules/sdnq/layers/embedding/forward.py b/modules/sdnq/layers/embedding/forward.py index 482178b02..50370ec7e 100644 --- a/modules/sdnq/layers/embedding/forward.py +++ b/modules/sdnq/layers/embedding/forward.py @@ -60,7 +60,7 @@ def quantized_embedding( def quantized_embedding_forward(self: torch.nn.Module, input: torch.Tensor) -> torch.FloatTensor: if self.sdnq_dequantizer.use_hadamard: - hadamard = get_hadamard(self.sdnq_dequantizer.hadamard_group_size, dtype=input.dtype, device=input.device) + hadamard = get_hadamard(self.sdnq_dequantizer.hadamard_group_size, dtype=self.sdnq_dequantizer.result_dtype, device=input.device) else: hadamard = None diff --git a/modules/sdnq/loader.py b/modules/sdnq/loader.py index 5ac4ca864..489e8945b 100644 --- a/modules/sdnq/loader.py +++ b/modules/sdnq/loader.py @@ -1,7 +1,6 @@ import os import json import torch -from diffusers.models.modeling_utils import ModelMixin from modules import shared from .common import dtype_dict, is_fp8_mm_supported, use_tensorwise_fp8_matmul, check_torch_compile, linear_types @@ -12,7 +11,7 @@ from .forward import get_forward_func from .file_loader import load_files -def get_module_names(model: ModelMixin) -> list: +def get_module_names(model: torch.nn.Module) -> list: modules_names = model._internal_dict.keys() # pylint: disable=protected-access modules_names = [m for m in modules_names if not m.startswith("_")] modules_names = [m for m in modules_names if isinstance(getattr(model, m, None), torch.nn.Module)] @@ -20,7 +19,7 @@ def get_module_names(model: ModelMixin) -> list: return modules_names -def normalize_tied_weights_keys_for_save(model: ModelMixin, is_pipeline: bool = False) -> list[tuple[torch.nn.Module, object]]: +def normalize_tied_weights_keys_for_save(model: torch.nn.Module, is_pipeline: bool = False) -> list[tuple[torch.nn.Module, object]]: normalized_modules = [] modules_to_walk = [] if is_pipeline: @@ -45,7 +44,13 @@ def restore_tied_weights_keys_after_save(normalized_modules: list[tuple[torch.nn submodule._tied_weights_keys = tied_weights_keys # pylint: disable=protected-access -def save_sdnq_model(model: ModelMixin, model_path: str, max_shard_size: str = "5GB", is_pipeline: bool = False, sdnq_config: SDNQConfig | None = None) -> None: +def save_sdnq_model( + model: torch.nn.Module, + model_path: str, + max_shard_size: str = "5GB", + is_pipeline: bool = False, + sdnq_config: SDNQConfig | None = None, +) -> None: normalized_modules = normalize_tied_weights_keys_for_save(model, is_pipeline=is_pipeline) try: model.save_pretrained(model_path, max_shard_size=max_shard_size) # actual save @@ -73,7 +78,18 @@ def save_sdnq_model(model: ModelMixin, model_path: str, max_shard_size: str = "5 model.config.quantization_config.to_json_file(quantization_config_path) -def load_sdnq_model(model_path: str, model_cls: ModelMixin | None = None, file_name: str | None = None, dtype: torch.dtype | None = None, device: torch.device = "cpu", dequantize_fp32: bool | None = None, use_quantized_matmul: bool | None = None, model_config: dict | None = None, quantization_config: dict | None = None, load_method: str = "safetensors") -> ModelMixin: +def load_sdnq_model( + model_path: str, + model_cls: torch.nn.Module | None = None, + file_name: str | None = None, + dtype: torch.dtype | None = None, + device: torch.device = "cpu", + dequantize_fp32: bool | None = None, + use_quantized_matmul: bool | None = None, + model_config: dict | None = None, + quantization_config: dict | None = None, + load_method: str = "safetensors", +) -> torch.nn.Module: from accelerate import init_empty_weights with init_empty_weights(): @@ -180,7 +196,7 @@ def load_sdnq_model(model_path: str, model_cls: ModelMixin | None = None, file_n return model -def post_process_model(model): +def post_process_model(model: torch.nn.Module) -> torch.nn.Module: has_children = list(model.children()) if not has_children: return model @@ -202,7 +218,14 @@ def post_process_model(model): return model -def apply_sdnq_options_to_module(model, quantization_config: SDNQConfig, dtype: torch.dtype | None = None, dequantize_fp32: bool | None = None, use_quantized_matmul: bool | None = None, full_param_name: str = ""): +def apply_sdnq_options_to_module( + model: torch.nn.Module, + quantization_config: SDNQConfig, + dtype: torch.dtype | None = None, + dequantize_fp32: bool | None = None, + use_quantized_matmul: bool | None = None, + full_param_name: str = "", +) -> torch.nn.Module: has_children = list(model.children()) if not has_children: if dtype is not None and getattr(model, "dtype", torch.float32) not in {torch.float32, torch.float64}: @@ -287,7 +310,7 @@ def apply_sdnq_options_to_module(model, quantization_config: SDNQConfig, dtype: return model -def apply_sdnq_options_to_model(model, dtype: torch.dtype | None = None, dequantize_fp32: bool | None = None, use_quantized_matmul: bool | None = None): +def apply_sdnq_options_to_model(model: torch.nn.Module, dtype: torch.dtype | None = None, dequantize_fp32: bool | None = None, use_quantized_matmul: bool | None = None) -> torch.nn.Module: if use_quantized_matmul and not check_torch_compile(): shared.log.warning("SDNQ: Quantized MatMul requires a working Triton install for best performance.") model = apply_sdnq_options_to_module(model, model.quantization_config, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul) diff --git a/modules/sdnq/quantizer.py b/modules/sdnq/quantizer.py index 3fbe4cb67..e33ed41bc 100644 --- a/modules/sdnq/quantizer.py +++ b/modules/sdnq/quantizer.py @@ -37,7 +37,7 @@ from .forward import get_forward_func from .layers import get_sdnq_wrapper_class from .quant_utils import quantize_weight, apply_svdquant, apply_hadamard, prepare_weight_for_matmul, prepare_svd_for_matmul -from .utils import check_param_name_in, get_quant_args_from_config, get_quant_kwargs, add_module_skip_keys +from .utils import check_param_name_in, get_quant_args_from_config, get_quant_kwargs, get_quantized_matmul_dtype, add_module_skip_keys class QuantizationMethod(str, Enum): @@ -66,7 +66,7 @@ def sdnq_quantize_layer_weight( skip_sr: bool = False, param_name: str | None = None, # pylint: disable=unused-argument torch_dtype: torch.dtype | None = None, -): +) -> tuple[SDNQDequantizer, dict[str, torch.Tensor]]: num_of_groups = 1 is_conv_type = False is_conv_transpose_type = False @@ -80,13 +80,7 @@ def sdnq_quantize_layer_weight( if torch_dtype is None: torch_dtype = weight.dtype - if quantized_matmul_dtype is None: - if dtype_dict[weights_dtype]["is_integer"]: - quantized_matmul_dtype = "int8" - elif dtype_dict[weights_dtype]["num_bits"] < 16: - quantized_matmul_dtype = "float8_e4m3fn" - else: - quantized_matmul_dtype = "float16" + quantized_matmul_dtype = get_quantized_matmul_dtype(weights_dtype, quantized_matmul_dtype) re_quantize_for_matmul = bool( dtype_dict[weights_dtype]["is_unsigned"] @@ -244,7 +238,7 @@ def sdnq_quantize_layer_weight( layer_class_name=layer_class_name, ) - return sdnq_dequantizer, {"weight": weight, "scale": scale, "zero_point": zero_point, "svd_up": svd_up, "svd_down": svd_down} + return (sdnq_dequantizer, {"weight": weight, "scale": scale, "zero_point": zero_point, "svd_up": svd_up, "svd_down": svd_down}) @devices.inference_context() @@ -266,7 +260,8 @@ def sdnq_quantize_layer_weight_dynamic( hadamard: torch.FloatTensor | None = None, param_name: str | None = None, torch_dtype: torch.dtype | None = None, -): + quantization_config: "SDNQConfig" = None, +) -> None | tuple[SDNQDequantizer, dict[str, torch.Tensor]]: if torch_dtype is None: torch_dtype = weight.dtype if dynamic_loss_threshold is None or dynamic_loss_threshold < 0: @@ -299,17 +294,34 @@ def sdnq_quantize_layer_weight_dynamic( quantization_loss = None for i in range(weights_dtype_order.index(weights_dtype), len(weights_dtype_order)): + add_param_to_not_use_matmul = False current_weights_dtype = weights_dtype_order[i] - if quantized_matmul_dtype is None and not is_fp8_mm_supported and not dtype_dict[current_weights_dtype]["is_integer"] and dtype_dict[current_weights_dtype]["num_bits"] < 16: + current_quantized_matmul_dtype = get_quantized_matmul_dtype(current_weights_dtype, quantized_matmul_dtype) + if quantized_matmul_dtype is None and not is_fp8_mm_supported and current_quantized_matmul_dtype in {"fp8", "float8_e4m3fn", "float8_e5m2"}: current_use_quantized_matmul = False + add_param_to_not_use_matmul = use_quantized_matmul else: current_use_quantized_matmul = use_quantized_matmul + if ( + (dtype_dict[current_quantized_matmul_dtype]["is_integer"] and not dtype_dict[current_weights_dtype]["is_integer"]) + or ( + dtype_dict[current_weights_dtype]["is_unsigned"] and not dtype_dict[current_quantized_matmul_dtype]["is_unsigned"] + and dtype_dict[current_quantized_matmul_dtype]["num_bits"] == dtype_dict[current_weights_dtype]["num_bits"] + ) + or ( + dtype_dict[weights_dtype]["num_bits"] <= dtype_dict[current_quantized_matmul_dtype]["num_bits"] + and dtype_dict[current_quantized_matmul_dtype]["num_bits"] < dtype_dict[current_weights_dtype]["num_bits"] + ) + ): + current_use_quantized_matmul = False + add_param_to_not_use_matmul = True + sdnq_dequantizer, weight_data = sdnq_quantize_layer_weight( weight, layer_class_name=layer_class_name, weights_dtype=current_weights_dtype, - quantized_matmul_dtype=quantized_matmul_dtype, + quantized_matmul_dtype=current_quantized_matmul_dtype, torch_dtype=torch_dtype, hadamard_group_size=hadamard_group_size, group_size=group_size, @@ -317,11 +329,11 @@ def sdnq_quantize_layer_weight_dynamic( svd_steps=svd_steps, use_svd=False, use_hadamard=False, - using_pre_calculated_svd=use_svd, - using_pre_rotated_hadamard=use_hadamard, use_quantized_matmul=current_use_quantized_matmul, use_stochastic_rounding=use_stochastic_rounding, dequantize_fp32=dequantize_fp32, + using_pre_calculated_svd=use_svd, + using_pre_rotated_hadamard=use_hadamard, param_name=param_name, ) @@ -347,13 +359,27 @@ def sdnq_quantize_layer_weight_dynamic( ).div_(weight_std) if quantization_loss <= dynamic_loss_threshold: del original_weight_fp32 - return sdnq_dequantizer, weight_data + if quantization_config is not None: + if sdnq_dequantizer.weights_dtype not in quantization_config.modules_dtype_dict.keys(): + quantization_config.modules_dtype_dict[sdnq_dequantizer.weights_dtype] = [param_name] + else: + quantization_config.modules_dtype_dict[sdnq_dequantizer.weights_dtype].append(param_name) + if add_param_to_not_use_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) + return (sdnq_dequantizer, weight_data), quantization_config + else: + return (sdnq_dequantizer, weight_data) + del original_weight_fp32 - return None + if quantization_config is not None: + quantization_config.modules_to_not_convert.append(param_name) + return None, quantization_config + else: + return None @devices.inference_context() -def sdnq_quantize_layer(layer, quantization_config: "SDNQConfig", torch_dtype: torch.dtype | None = None, param_name: str = "", quant_kwargs: dict | None = None): # pylint: disable=unused-argument +def sdnq_quantize_layer(layer: torch.nn.Module, quantization_config: "SDNQConfig", torch_dtype: torch.dtype | None = None, param_name: str = "", quant_kwargs: dict | None = None) -> tuple[torch.nn.Module, "SDNQConfig"]: # pylint: disable=unused-argument if torch_dtype is None: torch_dtype = layer.weight.dtype if quant_kwargs is None: @@ -379,7 +405,7 @@ def sdnq_quantize_layer(layer, quantization_config: "SDNQConfig", torch_dtype: t layer.weight.data = layer.weight.to(quantization_device, non_blocking=non_blocking, copy=False) if use_dynamic_quantization: - weight_data = sdnq_quantize_layer_weight_dynamic(layer.weight, **quant_kwargs) + weight_data, quantization_config = sdnq_quantize_layer_weight_dynamic(layer.weight, quantization_config=quantization_config, **quant_kwargs) else: weight_data = sdnq_quantize_layer_weight(layer.weight, **quant_kwargs) @@ -395,24 +421,19 @@ def sdnq_quantize_layer(layer, quantization_config: "SDNQConfig", torch_dtype: t setattr(layer, key, value) del weight_data - if use_dynamic_quantization: - if layer.sdnq_dequantizer.weights_dtype not in quantization_config.modules_dtype_dict.keys(): - quantization_config.modules_dtype_dict[layer.sdnq_dequantizer.weights_dtype] = [param_name] - else: - quantization_config.modules_dtype_dict[layer.sdnq_dequantizer.weights_dtype].append(param_name) - - if quant_kwargs["use_quantized_matmul"] and not layer.sdnq_dequantizer.use_quantized_matmul: + if ( + quant_kwargs["use_quantized_matmul"] and not layer.sdnq_dequantizer.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) else: layer.weight = torch.nn.Parameter(layer.weight.to(return_device, dtype=torch_dtype, non_blocking=non_blocking, copy=False), requires_grad=False) - if use_dynamic_quantization: - quantization_config.modules_to_not_convert.append(param_name) return layer, quantization_config @devices.inference_context() -def apply_sdnq_to_module(model, quantization_config: "SDNQConfig", torch_dtype: torch.dtype | None = None, full_param_name: str = ""): # pylint: disable=unused-argument +def apply_sdnq_to_module(model: torch.nn.Module, quantization_config: "SDNQConfig", torch_dtype: torch.dtype | None = None, full_param_name: str = "") -> tuple[torch.nn.Module, "SDNQConfig"]: # pylint: disable=unused-argument if not list(model.children()): return model, quantization_config for module_name, module in model.named_children(): @@ -470,7 +491,7 @@ def sdnq_post_load_quant( return_device: torch.device | None = None, torch_dtype: torch.dtype | None = None, pre_quantized: bool = False, -): +) -> torch.nn.Module: if pre_quantized: add_skip_keys = False use_dynamic_quantization = False @@ -529,7 +550,7 @@ def sdnq_post_load_quant( class SDNQQuantize: - def __init__(self, hf_quantizer): + def __init__(self, hf_quantizer: "SDNQQuantizer"): self.hf_quantizer = hf_quantizer def convert( @@ -539,7 +560,7 @@ class SDNQQuantize: full_layer_name: str | None = None, missing_keys: list[str] | None = None, # pylint: disable=unused-argument **kwargs, # pylint: disable=unused-argument - ) -> dict[str, torch.FloatTensor]: + ) -> dict[str, torch.Tensor]: _module_name, value = tuple(input_dict.items())[0] value = value[0] self.hf_quantizer.create_quantized_param(model, value, full_layer_name, value.device) @@ -563,7 +584,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): required_packages = None torch_dtype = None - def __str__(self): + def __str__(self) -> str: return f"SDNQQuantizer(torch_dtype={self.torch_dtype}, requires_parameters_quantization={self.requires_parameters_quantization}, use_keep_in_fp32_modules={self.use_keep_in_fp32_modules}, requires_calibration={self.requires_calibration}, required_packages={self.required_packages})" def check_if_quantized_param( @@ -572,7 +593,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): param_value: torch.Tensor, # pylint: disable=unused-argument param_name: str, *args, **kwargs, # pylint: disable=unused-argument - ): + ) -> bool: if self.pre_quantized: layer, _tensor_name = get_module_from_name(model, param_name) if hasattr(layer, "sdnq_dequantizer") and param_name.rsplit(".", maxsplit=1)[-1] in sdnq_keys: @@ -600,7 +621,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): param_name: str, target_device: torch.device, *args, **kwargs, # pylint: disable=unused-argument - ): + ) -> None: layer, tensor_name = get_module_from_name(model, param_name) if self.pre_quantized: @@ -653,7 +674,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): device_map, # pylint: disable=unused-argument keep_in_fp32_modules: list[str] | None = None, **kwargs, # pylint: disable=unused-argument - ): + ) -> None: if self.pre_quantized: self.quantization_config.quantization_device = None self.quantization_config.return_device = None @@ -671,7 +692,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): model, self.quantization_config = add_module_skip_keys(model, self.quantization_config) # pylint: disable=attribute-defined-outside-init - def _process_model_after_weight_loading(self, model, **kwargs): # pylint: disable=unused-argument + def _process_model_after_weight_loading(self, model: torch.nn.Module, **kwargs) -> torch.nn.Module: # pylint: disable=unused-argument model.quantization_config = self.quantization_config model.quantization_method = QuantizationMethod.SDNQ if hasattr(model, "config"): @@ -707,7 +728,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): devices.torch_gc(force=True, reason="sdnq") return model - def get_quantize_ops(self): + def get_quantize_ops(self) -> SDNQQuantize: return SDNQQuantize(self) def adjust_max_memory(self, max_memory: dict[str, int | str]) -> dict[str, int | str]: @@ -728,10 +749,10 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): # diffusers return state_dict, {} - def get_accelerator_warm_up_factor(self): + def get_accelerator_warm_up_factor(self) -> int: return 32 // dtype_dict[self.quantization_config.weights_dtype]["num_bits"] - def _dequantize(self, model): + def _dequantize(self, model: torch.nn.Module) -> torch.nn.Module: return dequantize_sdnq_model(model) def is_serializable(self, *args, **kwargs) -> bool: # pylint: disable=unused-argument, invalid-overridden-method @@ -742,7 +763,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): return self.is_serializable() @property - def is_trainable(self): + def is_trainable(self) -> bool: return self.quantization_config.is_training @property @@ -750,7 +771,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): return self.is_trainable() @property - def is_compileable(self): + def is_compileable(self) -> bool: return True def check_quantized_param(self, *args, **kwargs) -> bool: @@ -759,13 +780,13 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): """ return self.check_if_quantized_param(*args, **kwargs) - def param_needs_quantization(self, model, param_name: str, *args, **kwargs) -> bool: + def param_needs_quantization(self, model: torch.nn.Module, param_name: str, *args, **kwargs) -> bool: """ needed for transformers compatibility, returns self.check_if_quantized_param """ return self.check_if_quantized_param(model, None, param_name, *args, **kwargs) - def get_cuda_warm_up_factor(self): + def get_cuda_warm_up_factor(self) -> int: """ needed for transformers compatibility, returns self.get_accelerator_warm_up_factor """ @@ -928,7 +949,7 @@ class SDNQConfig(QuantizationConfigMixin): self.sdnq_version = sdnq_version self.post_init() - def post_init(self): + def post_init(self) -> None: r""" Safety checker that arguments are correct """ @@ -987,13 +1008,13 @@ class SDNQConfig(QuantizationConfigMixin): for key, value in self.modules_dtype_dict.items(): self.modules_dtype_dict[key] = list(set(value)) - def to_dict(self): + def to_dict(self) -> dict: quantization_config_dict = self.__dict__.copy() # make serializable quantization_config_dict["quantization_device"] = str(quantization_config_dict["quantization_device"]) if quantization_config_dict["quantization_device"] is not None else None quantization_config_dict["return_device"] = str(quantization_config_dict["return_device"]) if quantization_config_dict["return_device"] is not None else None return quantization_config_dict - def __str__(self): + def __str__(self) -> str: return f"SDNQConfig(weights_dtype={self.weights_dtype} quantization_device={self.quantization_device} return_device={self.return_device} group_size={self.group_size} use_quantized_matmul={self.use_quantized_matmul} quantized_matmul_dtype={self.quantized_matmul_dtype} quant_conv={self.quant_conv} quant_embedding={self.quant_embedding} use_quantized_matmul_conv={self.use_quantized_matmul_conv} use_static_quantization={self.use_static_quantization} use_dynamic_quantization={self.use_dynamic_quantization} dynamic_loss_threshold={self.dynamic_loss_threshold} use_stochastic_rounding={self.use_stochastic_rounding} use_hadamard={self.use_hadamard} hadamard_group_size={self.hadamard_group_size} use_svd={self.use_svd} svd_rank={self.svd_rank} svd_steps={self.svd_steps} dequantize_fp32={self.dequantize_fp32} non_blocking={self.non_blocking} add_skip_keys={self.add_skip_keys} modules_to_not_convert={self.modules_to_not_convert} modules_to_not_use_matmul={self.modules_to_not_use_matmul} modules_dtype_dict={self.modules_dtype_dict} modules_quant_config={self.modules_quant_config} )" diff --git a/modules/sdnq/utils.py b/modules/sdnq/utils.py index 3038ca6d2..672984bbc 100644 --- a/modules/sdnq/utils.py +++ b/modules/sdnq/utils.py @@ -121,6 +121,17 @@ def get_quant_kwargs(layer: torch.nn.Module, quantization_config, torch_dtype: t return quant_kwargs +def get_quantized_matmul_dtype(weights_dtype: str, quantized_matmul_dtype: str | None = None) -> str: + if quantized_matmul_dtype is None: + if dtype_dict[weights_dtype]["is_integer"]: + quantized_matmul_dtype = "int8" + elif dtype_dict[weights_dtype]["num_bits"] < 16: + quantized_matmul_dtype = "float8_e4m3fn" + else: + quantized_matmul_dtype = "float16" + return quantized_matmul_dtype + + def add_module_skip_keys(model: torch.nn.Module, quantization_config): if getattr(model, "_keep_in_fp32_modules", None) is not None: quantization_config.modules_to_not_convert.extend(model._keep_in_fp32_modules) # pylint: disable=protected-access @@ -140,15 +151,7 @@ def add_module_skip_keys(model: torch.nn.Module, quantization_config): else: quantization_config.modules_dtype_dict[key] = value - if quantization_config.quantized_matmul_dtype is None: - if dtype_dict[quantization_config.weights_dtype]["is_integer"]: - quantized_matmul_dtype = "int8" - elif dtype_dict[quantization_config.weights_dtype]["num_bits"] < 16: - quantized_matmul_dtype = "float8_e4m3fn" - else: - quantized_matmul_dtype = "float16" - else: - quantized_matmul_dtype = quantization_config.quantized_matmul_dtype + quantized_matmul_dtype = get_quantized_matmul_dtype(quantization_config.weights_dtype, quantization_config.quantized_matmul_dtype) quantization_config.modules_to_not_use_matmul.extend(skip_key_list[2].get(quantized_matmul_dtype, [])) else: quantization_config.modules_to_not_convert.extend(common_skip_keys)