diff --git a/modules/sdnq/quantizer.py b/modules/sdnq/quantizer.py index 58083079a..501f70c62 100644 --- a/modules/sdnq/quantizer.py +++ b/modules/sdnq/quantizer.py @@ -1,6 +1,6 @@ # pylint: disable=redefined-builtin,no-member,protected-access -from typing import Dict, List, Tuple, Optional, Union +from typing import Any, Dict, List, Tuple, Optional, Union from dataclasses import dataclass from enum import Enum @@ -230,9 +230,7 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz use_quantized_matmul=use_quantized_matmul, re_quantize_for_matmul=re_quantize_for_matmul, ) - layer.weight.data = layer.sdnq_dequantizer.pack_weight(layer.weight).to(return_device, non_blocking=non_blocking) - layer.sdnq_dequantizer = layer.sdnq_dequantizer.to(return_device, non_blocking=non_blocking) layer.forward = get_forward_func(layer_class_name, use_quantized_matmul, dtype_dict[weights_dtype]["is_integer"], use_tensorwise_fp8_matmul) layer.forward = layer.forward.__get__(layer, layer.__class__) @@ -425,6 +423,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): def __init__(self, quantization_config, **kwargs): super().__init__(quantization_config, **kwargs) self.modules_to_not_convert = [] + self.updated_expected_keys = False def check_if_quantized_param( self, @@ -433,6 +432,8 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): param_name: str, *args, **kwargs, # pylint: disable=unused-argument ): + if self.pre_quantized and self.updated_expected_keys and (param_name.endswith(".scale") or param_name.endswith(".zero_point") or param_name.endswith(".svd_up") or param_name.endswith(".svd_down")): + return True if param_name.endswith(".weight"): split_param_name = param_name.split(".") if ( @@ -471,11 +472,22 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): param_value: torch.FloatTensor, param_name: str, target_device: torch.device, + state_dict: Dict[str, Any], *args, **kwargs, # pylint: disable=unused-argument ): - weights_dtype = self.quantization_config.weights_dtype - torch_dtype = param_value.dtype if self.torch_dtype is None else self.torch_dtype + if self.pre_quantized and self.updated_expected_keys and (param_name.endswith(".scale") or param_name.endswith(".zero_point") or param_name.endswith(".svd_up") or param_name.endswith(".svd_down")): + layer, tensor_name = get_module_from_name(model, param_name) + return_dtype = torch.float32 if self.quantization_config.dequantize_fp32 else self.torch_dtype if self.torch_dtype is not None else param_value.dtype + if param_value is not None: + if param_value.dtype == return_dtype and devices.same_device(param_value.device, target_device): + param_value = param_value.clone() + else: + param_value = param_value.to(target_device, dtype=return_dtype) + param_value = torch.nn.Parameter(param_value, requires_grad=False) + setattr(layer, tensor_name, param_value) + return + weights_dtype = self.quantization_config.weights_dtype if len(self.quantization_config.modules_dtype_dict.keys()) > 0: split_param_name = param_name.split(".") for key, value in self.quantization_config.modules_dtype_dict.items(): @@ -506,29 +518,88 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): if self.quantization_config.quantization_device is not None: target_device = self.quantization_config.quantization_device - if param_value.dtype == torch.float32 and devices.same_device(param_value.device, target_device): - param_value = param_value.clone() - else: - param_value = param_value.to(target_device, non_blocking=self.quantization_config.non_blocking).to(dtype=torch.float32) + if not self.pre_quantized: + torch_dtype = param_value.dtype if self.torch_dtype is None else self.torch_dtype + if param_value.dtype == torch.float32 and devices.same_device(param_value.device, target_device): + param_value = param_value.clone() + else: + param_value = param_value.to(target_device, non_blocking=self.quantization_config.non_blocking).to(dtype=torch.float32) - layer, _ = get_module_from_name(model, param_name) - layer.weight = torch.nn.Parameter(param_value, requires_grad=False) - layer = sdnq_quantize_layer( - layer, - weights_dtype=weights_dtype, - torch_dtype=torch_dtype, - group_size=self.quantization_config.group_size, - svd_rank=self.quantization_config.svd_rank, - use_svd=self.quantization_config.use_svd, - quant_conv=self.quantization_config.quant_conv, - use_quantized_matmul=self.quantization_config.use_quantized_matmul, - use_quantized_matmul_conv=self.quantization_config.use_quantized_matmul_conv, - dequantize_fp32=self.quantization_config.dequantize_fp32, - non_blocking=self.quantization_config.non_blocking, - quantization_device=None, - return_device=return_device, - param_name=param_name, - ) + layer, _ = get_module_from_name(model, param_name) + layer.weight = torch.nn.Parameter(param_value, requires_grad=False) + layer = sdnq_quantize_layer( + layer, + weights_dtype=weights_dtype, + torch_dtype=torch_dtype, + group_size=self.quantization_config.group_size, + svd_rank=self.quantization_config.svd_rank, + use_svd=self.quantization_config.use_svd, + quant_conv=self.quantization_config.quant_conv, + use_quantized_matmul=self.quantization_config.use_quantized_matmul, + use_quantized_matmul_conv=self.quantization_config.use_quantized_matmul_conv, + dequantize_fp32=self.quantization_config.dequantize_fp32, + non_blocking=self.quantization_config.non_blocking, + quantization_device=None, + return_device=return_device, + param_name=param_name, + ) + else: + from accelerate import init_empty_weights + layer, _ = get_module_from_name(model, param_name) + torch_dtype = layer.weight.dtype if self.torch_dtype is None else self.torch_dtype + + if self.updated_expected_keys: # prevent overwrite to meta + scale, zero_point, svd_up, svd_down = None, None, None, None + if hasattr(layer, "scale") and layer.scale.device.type != "meta": + scale = layer.scale + del layer.scale + if hasattr(layer, "zero_point") and layer.zero_point.device.type != "meta": + zero_point = layer.zero_point + del layer.zero_point + if hasattr(layer, "svd_up") and layer.svd_up.device.type != "meta": + svd_up = layer.svd_up + del layer.svd_up + if hasattr(layer, "svd_down") and layer.svd_down.device.type != "meta": + svd_down = layer.svd_down + del layer.svd_down + + with init_empty_weights(): + # add sdnq_dequantizer + layer = sdnq_quantize_layer( + layer, + weights_dtype=weights_dtype, + torch_dtype=torch_dtype, + group_size=self.quantization_config.group_size, + svd_rank=self.quantization_config.svd_rank, + use_svd=self.quantization_config.use_svd, + quant_conv=self.quantization_config.quant_conv, + use_quantized_matmul=self.quantization_config.use_quantized_matmul, + use_quantized_matmul_conv=self.quantization_config.use_quantized_matmul_conv, + dequantize_fp32=self.quantization_config.dequantize_fp32, + non_blocking=self.quantization_config.non_blocking, + quantization_device="meta", + return_device="meta", + param_name=param_name, + ) + + layer.weight = torch.nn.Parameter(param_value.clone().to(target_device), requires_grad=False) + if self.updated_expected_keys: # Transformers + if scale is not None: + layer.scale = torch.nn.Parameter(scale, requires_grad=False) + if zero_point is not None: + layer.zero_point = torch.nn.Parameter(zero_point, requires_grad=False) + if svd_up is not None: + layer.svd_up = torch.nn.Parameter(svd_up, requires_grad=False) + if svd_down is not None: + layer.svd_down = torch.nn.Parameter(svd_down, requires_grad=False) + else: # Diffusers doesn't have the API for updating expected keys + layer_key = param_name.removesuffix(".weight") + layer.scale = torch.nn.Parameter(state_dict[layer_key + ".scale"].clone().to(target_device), requires_grad=False) + if layer.zero_point is not None: + layer.zero_point = torch.nn.Parameter(state_dict[layer_key + ".zero_point"].clone().to(target_device), requires_grad=False) + if layer.svd_up is not None: + layer.svd_up = torch.nn.Parameter(state_dict[layer_key + ".svd_up"].clone().to(target_device), requires_grad=False) + layer.svd_down = torch.nn.Parameter(state_dict[layer_key + ".svd_down"].clone().to(target_device), requires_grad=False) def adjust_max_memory(self, max_memory: Dict[str, Union[int, str]]) -> Dict[str, Union[int, str]]: max_memory = {key: val * 0.80 for key, val in max_memory.items()} @@ -560,6 +631,19 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): model.config.quantization_config = self.quantization_config model.quantization_config = self.quantization_config + if self.pre_quantized and hasattr(model, "get_parameter_or_buffer"): + from functools import wraps + @wraps(model.get_parameter_or_buffer) + def get_parameter_or_buffer(self, target: str): + try: + return self.original_get_parameter_or_buffer(target) + except Exception as e: + if target.endswith(".scale") or target.endswith(".zero_point") or target.endswith(".svd_up") or target.endswith(".svd_down"): + return None + raise e + model.original_get_parameter_or_buffer = model.get_parameter_or_buffer + model.get_parameter_or_buffer = get_parameter_or_buffer.__get__(model, model.__class__) + def _process_model_after_weight_loading(self, model, **kwargs): # pylint: disable=unused-argument if shared.opts.diffusers_offload_mode != "none": model = model.to(devices.cpu) @@ -575,6 +659,24 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): """ return self.get_accelerator_warm_up_factor() + def update_unexpected_keys(self, model, unexpected_keys: List[str], *args, **kwargs) -> List[str]: # pylint: disable=unused-argument + if not self.pre_quantized: + return unexpected_keys + new_unexpected_keys = [] + for key in unexpected_keys: + if not (key.endswith(".scale") or key.endswith(".zero_point") or key.endswith(".svd_up") or key.endswith(".svd_down")): + new_unexpected_keys.append(key) + return new_unexpected_keys + + def update_expected_keys(self, model, expected_keys: List[str], loaded_keys: list[str], *args, **kwargs) -> List[str]: # pylint: disable=unused-argument + if not self.pre_quantized: + return expected_keys + self.updated_expected_keys = True + for key in loaded_keys: + if key.endswith(".scale") or key.endswith(".zero_point") or key.endswith(".svd_up") or key.endswith(".svd_down"): + expected_keys.append(key) + return expected_keys + def update_tp_plan(self, config, *args, **kwargs): # pylint: disable=unused-argument """ needed for transformers compatibilty, no-op function @@ -587,12 +689,6 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): """ return config - def update_unexpected_keys(self, model, unexpected_keys: List[str], *args, **kwargs) -> List[str]: # pylint: disable=unused-argument - """ - needed for transformers compatibilty, no-op function - """ - return unexpected_keys - def update_missing_keys_after_loading(self, model, missing_keys: List[str], *args, **kwargs) -> List[str]: # pylint: disable=unused-argument """ needed for transformers compatibilty, no-op function @@ -605,12 +701,6 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): """ return state_dict - def update_expected_keys(self, model, expected_keys: List[str], *args, **kwargs) -> List[str]: # pylint: disable=unused-argument - """ - needed for transformers compatibilty, no-op function - """ - return expected_keys - def update_param_name(self, param_name: str, *args, **kwargs) -> str: # pylint: disable=unused-argument """ needed for transformers compatibilty, no-op function @@ -624,9 +714,6 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): return dtype def is_serializable(self, *args, **kwargs) -> bool: # pylint: disable=unused-argument, invalid-overridden-method - """ - needed for transformers compatibilty, returns True - """ return True @property