add version to sdnq

This commit is contained in:
Disty0
2025-11-28 00:45:22 +03:00
parent c2ee7c0328
commit 55cf627ac6
4 changed files with 56 additions and 29 deletions
+1
View File
@@ -1,6 +1,7 @@
from .quantizer import QuantizationMethod, SDNQConfig, SDNQQuantizer, sdnq_post_load_quant, apply_sdnq_to_module, sdnq_quantize_layer
from .loader import save_sdnq_model, load_sdnq_model
__version__ = "0.1.0"
__all__ = [
"QuantizationMethod",
+1
View File
@@ -5,6 +5,7 @@ import torch
from modules import shared, devices
sdnq_version = "0.1.0"
dtype_dict = {
"int32": {"min": -2147483648, "max": 2147483647, "num_bits": 32, "sign": 1, "exponent": 0, "mantissa": 31, "target_dtype": torch.int32, "torch_dtype": torch.int32, "storage_dtype": torch.int32, "is_unsigned": False, "is_integer": True, "is_packed": False},
+31 -15
View File
@@ -4,7 +4,7 @@ import torch
from diffusers.models.modeling_utils import ModelMixin
from .common import dtype_dict, use_tensorwise_fp8_matmul
from .quantizer import SDNQConfig, sdnq_post_load_quant, prepare_weight_for_matmul, prepare_svd_for_matmul
from .quantizer import SDNQConfig, sdnq_post_load_quant, prepare_weight_for_matmul, prepare_svd_for_matmul, get_quant_args_from_config
from .forward import get_forward_func
from .file_loader import load_files
@@ -97,17 +97,6 @@ def load_sdnq_model(model_path: str, model_cls: ModelMixin = None, file_name: st
if model_cls is None:
raise ValueError(f"Cannot determine model class for {model_path}, please provide model_cls argument")
quantization_config.pop("is_integer", None)
quantization_config.pop("quant_method", None)
quantization_config.pop("quantization_device", None)
quantization_config.pop("return_device", None)
quantization_config.pop("non_blocking", None)
quantization_config.pop("add_skip_keys", None)
quantization_config.pop("use_static_quantization", None)
quantization_config.pop("use_stochastic_rounding", None)
quantization_config.pop("use_grad_ckpt", None)
quantization_config.pop("is_training", None)
if hasattr(model_cls, "load_config") and hasattr(model_cls, "from_config"):
config = model_cls.load_config(model_path)
model = model_cls.from_config(config)
@@ -117,7 +106,7 @@ def load_sdnq_model(model_path: str, model_cls: ModelMixin = None, file_name: st
else:
model = model_cls(**model_config)
model = sdnq_post_load_quant(model, torch_dtype=dtype, add_skip_keys=False, **quantization_config)
model = sdnq_post_load_quant(model, torch_dtype=dtype, add_skip_keys=False, **get_quant_args_from_config(quantization_config))
key_mapping = getattr(model, "_checkpoint_conversion_mapping", None)
files = []
@@ -167,7 +156,7 @@ def post_process_model(model):
return model
def apply_sdnq_options_to_model(model, dtype: torch.dtype = None, dequantize_fp32: bool = None, use_quantized_matmul: bool = None):
def apply_sdnq_options_to_module(model, dtype: torch.dtype = None, dequantize_fp32: bool = None, use_quantized_matmul: bool = None):
has_children = list(model.children())
if not has_children:
if dtype is not None and getattr(model, "dtype", torch.float32) != torch.float32:
@@ -212,5 +201,32 @@ def apply_sdnq_options_to_model(model, dtype: torch.dtype = None, dequantize_fp3
module.forward = module.forward.__get__(module, module.__class__)
setattr(model, module_name, module)
else:
setattr(model, module_name, apply_sdnq_options_to_model(module, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul))
setattr(model, module_name, apply_sdnq_options_to_module(module, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul))
return model
def apply_sdnq_options_to_model(model, dtype: torch.dtype = None, dequantize_fp32: bool = None, use_quantized_matmul: bool = None):
model = apply_sdnq_options_to_module(model, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul)
if hasattr(model, "quantization_config"):
if use_quantized_matmul is not None:
model.quantization_config.use_quantized_matmul = use_quantized_matmul
if dequantize_fp32 is not None:
model.quantization_config.dequantize_fp32 = dequantize_fp32
if hasattr(model, "config"):
try:
if hasattr(model.config, "quantization_config"):
if use_quantized_matmul is not None:
model.config.quantization_config.use_quantized_matmul = use_quantized_matmul
if dequantize_fp32 is not None:
model.config.quantization_config.dequantize_fp32 = dequantize_fp32
except Exception:
pass
try:
if hasattr(model.config, "get") and model.config.get("quantization_config", None) is not None:
if use_quantized_matmul is not None:
model.config["quantization_config"].use_quantized_matmul = use_quantized_matmul
if dequantize_fp32 is not None:
model.config["quantization_config"].dequantize_fp32 = dequantize_fp32
except Exception:
pass
return model
+23 -14
View File
@@ -16,7 +16,7 @@ from accelerate import init_empty_weights
from accelerate.utils import set_module_tensor_to_device
from modules import devices, shared
from .common import dtype_dict, common_skip_keys, module_skip_keys_dict, accepted_weight_dtypes, accepted_matmul_dtypes, allowed_types, linear_types, conv_types, conv_transpose_types, compile_func, use_tensorwise_fp8_matmul, use_contiguous_mm
from .common import sdnq_version, dtype_dict, common_skip_keys, module_skip_keys_dict, accepted_weight_dtypes, accepted_matmul_dtypes, allowed_types, linear_types, conv_types, conv_transpose_types, compile_func, use_tensorwise_fp8_matmul, use_contiguous_mm
from .dequantizer import SDNQDequantizer, dequantize_sdnq_model
from .packed_int import pack_int_symetric, pack_int_asymetric
from .forward import get_forward_func
@@ -130,6 +130,25 @@ def check_param_name_in(param_name: str, param_list: List[str]) -> bool:
return False
def get_quant_args_from_config(quantization_config: Union["SDNQConfig", dict]) -> dict:
if isinstance(quantization_config, SDNQConfig):
quantization_config_dict = quantization_config.to_dict()
else:
quantization_config_dict = quantization_config.copy()
quantization_config_dict.pop("is_integer", None)
quantization_config_dict.pop("quant_method", None)
quantization_config_dict.pop("quantization_device", None)
quantization_config_dict.pop("return_device", None)
quantization_config_dict.pop("non_blocking", None)
quantization_config_dict.pop("add_skip_keys", None)
quantization_config_dict.pop("use_static_quantization", None)
quantization_config_dict.pop("use_stochastic_rounding", None)
quantization_config_dict.pop("use_grad_ckpt", None)
quantization_config_dict.pop("is_training", None)
quantization_config_dict.pop("sdnq_version", None)
return quantization_config_dict
def get_minimum_dtype(weights_dtype: str, param_name: str, modules_dtype_dict: Dict[str, List[str]]):
if len(modules_dtype_dict.keys()) > 0:
for key, value in modules_dtype_dict.items():
@@ -719,19 +738,8 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
self.quantization_config.non_blocking = False
self.quantization_config.add_skip_keys = False
quantization_config_dict = self.quantization_config.to_dict()
quantization_config_dict.pop("is_integer", None)
quantization_config_dict.pop("quant_method", None)
quantization_config_dict.pop("quantization_device", None)
quantization_config_dict.pop("return_device", None)
quantization_config_dict.pop("non_blocking", None)
quantization_config_dict.pop("add_skip_keys", None)
quantization_config_dict.pop("use_static_quantization", None)
quantization_config_dict.pop("use_stochastic_rounding", None)
quantization_config_dict.pop("use_grad_ckpt", None)
quantization_config_dict.pop("is_training", None)
with init_empty_weights():
model = sdnq_post_load_quant(model, torch_dtype=self.torch_dtype, add_skip_keys=False, **quantization_config_dict)
model = sdnq_post_load_quant(model, torch_dtype=self.torch_dtype, add_skip_keys=False, **get_quant_args_from_config(self.quantization_config))
if self.quantization_config.add_skip_keys:
if keep_in_fp32_modules is not None:
@@ -890,8 +898,9 @@ class SDNQConfig(QuantizationConfigMixin):
self.return_device = return_device
self.modules_to_not_convert = modules_to_not_convert
self.modules_dtype_dict = modules_dtype_dict
self.post_init()
self.is_integer = dtype_dict[self.weights_dtype]["is_integer"]
self.sdnq_version = sdnq_version
self.post_init()
def post_init(self):
r"""