mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
SDNQ add native pre-quant loader support to from_pretrained
This commit is contained in:
+129
-42
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user