diff --git a/modules/sdnq/__init__.py b/modules/sdnq/__init__.py index c67bf0db7..e53746b3b 100644 --- a/modules/sdnq/__init__.py +++ b/modules/sdnq/__init__.py @@ -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", diff --git a/modules/sdnq/common.py b/modules/sdnq/common.py index 49db90bf4..6f7d6ba50 100644 --- a/modules/sdnq/common.py +++ b/modules/sdnq/common.py @@ -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}, diff --git a/modules/sdnq/loader.py b/modules/sdnq/loader.py index 49efdedc3..0821e62ad 100644 --- a/modules/sdnq/loader.py +++ b/modules/sdnq/loader.py @@ -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 diff --git a/modules/sdnq/quantizer.py b/modules/sdnq/quantizer.py index 156ac23e3..1e383f0bd 100644 --- a/modules/sdnq/quantizer.py +++ b/modules/sdnq/quantizer.py @@ -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"""