SDNQ: Don't re-quantize for matmul with dyn quant, Fix nn.Embedding and Typing

This commit is contained in:
Disty0
2026-07-04 02:54:19 +03:00
parent 2d5903e91b
commit d62a0690a2
9 changed files with 161 additions and 92 deletions
+1 -1
View File
@@ -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
+37 -17
View File
@@ -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:
+3 -3
View File
@@ -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
+2 -2
View File
@@ -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)
+6 -4
View File
@@ -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)
+1 -1
View File
@@ -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
+31 -8
View File
@@ -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)
+68 -47
View File
@@ -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} )"
+12 -9
View File
@@ -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)