SDNQ add native pre-quant loader support to from_pretrained

This commit is contained in:
Disty0
2025-10-11 16:19:11 +03:00
parent 6bc83bc296
commit f7286c90d5
+129 -42
View File
@@ -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