mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
add version to sdnq
This commit is contained in:
@@ -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",
|
||||
|
||||
@@ -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
@@ -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
@@ -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"""
|
||||
|
||||
Reference in New Issue
Block a user