mirror of
https://github.com/vladmandic/automatic
synced 2026-09-20 01:31:13 +02:00
update sdnq
This commit is contained in:
@@ -7,9 +7,9 @@ import torch
|
||||
from ...common import use_contiguous_mm # noqa: TID252
|
||||
|
||||
|
||||
def check_mats(input: torch.Tensor, weight: torch.Tensor) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
def check_mats(input: torch.Tensor, weight: torch.Tensor, allow_contiguous_mm: bool = True) -> Tuple[torch.Tensor, torch.Tensor]:
|
||||
input = input.contiguous()
|
||||
if use_contiguous_mm:
|
||||
if allow_contiguous_mm and use_contiguous_mm:
|
||||
weight = weight.contiguous()
|
||||
elif weight.is_contiguous():
|
||||
weight = weight.t().contiguous().t()
|
||||
|
||||
@@ -36,7 +36,7 @@ def fp8_matmul(
|
||||
input = input.flatten(0,-2)
|
||||
svd_bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
|
||||
input, input_scale = quantize_fp_mm_input(input)
|
||||
input, weight = check_mats(input, weight)
|
||||
input, weight = check_mats(input, weight, allow_contiguous_mm=False)
|
||||
if bias is not None and bias.dtype != torch.bfloat16:
|
||||
bias = bias.to(dtype=torch.bfloat16)
|
||||
result = torch._scaled_mm(input, weight, scale_a=input_scale, scale_b=scale, bias=bias, out_dtype=torch.bfloat16)
|
||||
|
||||
@@ -43,7 +43,7 @@ def fp8_matmul_tensorwise(
|
||||
bias = torch.mm(torch.mm(input.to(dtype=svd_down.dtype), svd_down), svd_up)
|
||||
dummy_input_scale = torch.ones(1, device=input.device, dtype=torch.float32)
|
||||
input, scale = quantize_fp_mm_input_tensorwise(input, scale)
|
||||
input, weight = check_mats(input, weight)
|
||||
input, weight = check_mats(input, weight, allow_contiguous_mm=False)
|
||||
if bias is not None:
|
||||
return dequantize_symmetric_with_bias(torch._scaled_mm(input, weight, scale_a=dummy_input_scale, scale_b=dummy_input_scale, bias=None, out_dtype=scale.dtype), scale, bias, dtype=return_dtype, result_shape=output_shape)
|
||||
else:
|
||||
|
||||
+118
-76
@@ -43,16 +43,21 @@ def get_scale_symmetric(weight: torch.FloatTensor, reduction_axes: Union[int, Li
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def quantize_weight(weight: torch.FloatTensor, reduction_axes: Union[int, List[int]], weights_dtype: str, use_stochastic_rounding: bool = False) -> Tuple[torch.Tensor, torch.FloatTensor, torch.FloatTensor]:
|
||||
def quantize_weight(weight: torch.FloatTensor, reduction_axes: Union[int, List[int]], weights_dtype: str, dtype: torch.dtype = None, use_stochastic_rounding: bool = False) -> Tuple[torch.Tensor, torch.FloatTensor, torch.FloatTensor]:
|
||||
weight = weight.to(dtype=torch.float32)
|
||||
|
||||
if dtype_dict[weights_dtype]["is_unsigned"]:
|
||||
scale, zero_point = get_scale_asymmetric(weight, reduction_axes, weights_dtype)
|
||||
if dtype is not None:
|
||||
scale = scale.to(dtype=dtype)
|
||||
zero_point = zero_point.to(dtype=dtype)
|
||||
quantized_weight = torch.sub(weight, zero_point).div_(scale)
|
||||
else:
|
||||
scale = get_scale_symmetric(weight, reduction_axes, weights_dtype)
|
||||
quantized_weight = torch.div(weight, scale)
|
||||
zero_point = None
|
||||
if dtype is not None:
|
||||
scale = scale.to(dtype=dtype)
|
||||
quantized_weight = torch.div(weight, scale)
|
||||
|
||||
if dtype_dict[weights_dtype]["is_integer"]:
|
||||
if use_stochastic_rounding:
|
||||
@@ -68,7 +73,7 @@ def quantize_weight(weight: torch.FloatTensor, reduction_axes: Union[int, List[i
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def apply_svdquant(weight: torch.FloatTensor, rank: int = 32, niter: int = 8) -> Tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]:
|
||||
def apply_svdquant(weight: torch.FloatTensor, rank: int = 32, niter: int = 8, dtype: torch.dtype = None) -> Tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]:
|
||||
reshape_weight = False
|
||||
if weight.ndim > 2: # convs
|
||||
reshape_weight = True
|
||||
@@ -78,6 +83,9 @@ def apply_svdquant(weight: torch.FloatTensor, rank: int = 32, niter: int = 8) ->
|
||||
U, S, svd_down = torch.svd_lowrank(weight, q=rank, niter=niter)
|
||||
svd_up = torch.mul(U, S.unsqueeze(0))
|
||||
svd_down = svd_down.t_()
|
||||
if dtype is not None:
|
||||
svd_up = svd_up.to(dtype=dtype)
|
||||
svd_down = svd_down.to(dtype=dtype)
|
||||
weight = weight.sub(torch.mm(svd_up, svd_down))
|
||||
if reshape_weight:
|
||||
weight = weight.unflatten(-1, (*weight_shape[1:],)) # pylint: disable=possibly-used-before-assignment
|
||||
@@ -139,6 +147,9 @@ def get_quant_args_from_config(quantization_config: Union["SDNQConfig", dict]) -
|
||||
quantization_config_dict.pop("use_grad_ckpt", None)
|
||||
quantization_config_dict.pop("is_training", None)
|
||||
quantization_config_dict.pop("sdnq_version", None)
|
||||
if quantization_config_dict.get("modules_quant_config", None) is not None:
|
||||
for key in quantization_config_dict["modules_quant_config"].keys():
|
||||
quantization_config_dict["modules_quant_config"][key] = get_quant_args_from_config(quantization_config_dict["modules_quant_config"][key])
|
||||
return quantization_config_dict
|
||||
|
||||
|
||||
@@ -169,6 +180,14 @@ def get_minimum_dtype(weights_dtype: str, param_name: str, modules_dtype_dict: D
|
||||
return weights_dtype
|
||||
|
||||
|
||||
def get_quant_kwargs(quant_kwargs: dict, modules_quant_config: Dict[str, dict]) -> dict:
|
||||
if check_param_name_in(quant_kwargs["param_name"], modules_quant_config.keys()):
|
||||
for key, value in modules_quant_config.items():
|
||||
quant_kwargs[key] = value
|
||||
quant_kwargs["weights_dtype"] = get_minimum_dtype(quant_kwargs["weights_dtype"], quant_kwargs["param_name"], quant_kwargs["modules_dtype_dict"])
|
||||
return quant_kwargs
|
||||
|
||||
|
||||
def add_module_skip_keys(model, modules_to_not_convert: List[str] = None, modules_dtype_dict: Dict[str, List[str]] = None):
|
||||
if modules_to_not_convert is None:
|
||||
modules_to_not_convert = []
|
||||
@@ -211,6 +230,8 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
|
||||
is_conv_transpose_type = False
|
||||
is_linear_type = False
|
||||
result_shape = None
|
||||
scale_dtype = None
|
||||
|
||||
original_shape = weight.shape
|
||||
original_stride = weight.stride()
|
||||
weight = weight.detach()
|
||||
@@ -274,9 +295,20 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
|
||||
reduction_axes = -1
|
||||
use_quantized_matmul = False
|
||||
|
||||
if (
|
||||
not dequantize_fp32
|
||||
and dtype_dict[weights_dtype]["num_bits"] <= 8
|
||||
and not (
|
||||
use_quantized_matmul
|
||||
and not dtype_dict[quantized_matmul_dtype]["is_integer"]
|
||||
and (not use_tensorwise_fp8_matmul or dtype_dict[quantized_matmul_dtype]["num_bits"] == 16)
|
||||
)
|
||||
):
|
||||
scale_dtype = torch_dtype
|
||||
|
||||
if use_svd:
|
||||
try:
|
||||
weight, svd_up, svd_down = apply_svdquant(weight, rank=svd_rank, niter=svd_steps)
|
||||
weight, svd_up, svd_down = apply_svdquant(weight, rank=svd_rank, niter=svd_steps, dtype=scale_dtype)
|
||||
if use_quantized_matmul:
|
||||
svd_up = svd_up.t_()
|
||||
svd_down = svd_down.t_()
|
||||
@@ -335,30 +367,21 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
|
||||
else:
|
||||
group_size = -1
|
||||
|
||||
weight, scale, zero_point = quantize_weight(weight, reduction_axes, weights_dtype, use_stochastic_rounding=(use_stochastic_rounding and not skip_sr))
|
||||
if (
|
||||
not dequantize_fp32
|
||||
and dtype_dict[weights_dtype]["num_bits"] <= 8
|
||||
and not (
|
||||
use_quantized_matmul
|
||||
and not dtype_dict[quantized_matmul_dtype]["is_integer"]
|
||||
and (not use_tensorwise_fp8_matmul or dtype_dict[quantized_matmul_dtype]["num_bits"] == 16)
|
||||
)
|
||||
):
|
||||
scale = scale.to(dtype=torch_dtype)
|
||||
if zero_point is not None:
|
||||
zero_point = zero_point.to(dtype=torch_dtype)
|
||||
if svd_up is not None:
|
||||
svd_up = svd_up.to(dtype=torch_dtype)
|
||||
svd_down = svd_down.to(dtype=torch_dtype)
|
||||
|
||||
cast_scale = True
|
||||
transpose_weights = False
|
||||
re_quantize_for_matmul = re_quantize_for_matmul or num_of_groups > 1
|
||||
if use_quantized_matmul and not re_quantize_for_matmul and not dtype_dict[weights_dtype]["is_packed"]:
|
||||
transpose_weights = True
|
||||
if not use_tensorwise_fp8_matmul and not dtype_dict[quantized_matmul_dtype]["is_integer"]:
|
||||
cast_scale = False
|
||||
|
||||
weight, scale, zero_point = quantize_weight(weight, reduction_axes, weights_dtype, dtype=(scale_dtype if cast_scale else None), use_stochastic_rounding=(use_stochastic_rounding and not skip_sr))
|
||||
|
||||
if transpose_weights:
|
||||
scale.t_()
|
||||
weight.t_()
|
||||
weight = prepare_weight_for_matmul(weight)
|
||||
if not use_tensorwise_fp8_matmul and not dtype_dict[quantized_matmul_dtype]["is_integer"]:
|
||||
scale = scale.to(dtype=torch.float32)
|
||||
|
||||
sdnq_dequantizer = SDNQDequantizer(
|
||||
result_dtype=torch_dtype,
|
||||
@@ -528,7 +551,7 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", quantized_matmul_dtype=None
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, dynamic_loss_threshold=1e-2, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_dynamic_quantization=False, use_stochastic_rounding=False, dequantize_fp32=False, non_blocking=False, modules_to_not_convert: List[str] = None, modules_dtype_dict: Dict[str, List[str]] = None, quantization_device=None, return_device=None, full_param_name=""): # pylint: disable=unused-argument
|
||||
def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, dynamic_loss_threshold=1e-2, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_dynamic_quantization=False, use_stochastic_rounding=False, dequantize_fp32=False, non_blocking=False, modules_to_not_convert: List[str] = None, modules_dtype_dict: Dict[str, List[str]] = None, modules_quant_config: Dict[str, dict] = None, quantization_device=None, return_device=None, full_param_name=""): # pylint: disable=unused-argument
|
||||
has_children = list(model.children())
|
||||
if not has_children:
|
||||
return model, modules_to_not_convert, modules_dtype_dict
|
||||
@@ -536,6 +559,8 @@ def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=Non
|
||||
modules_to_not_convert = []
|
||||
if modules_dtype_dict is None:
|
||||
modules_dtype_dict = {}
|
||||
if modules_quant_config is None:
|
||||
modules_quant_config = {}
|
||||
for module_name, module in model.named_children():
|
||||
if full_param_name:
|
||||
param_name = full_param_name + "." + module_name
|
||||
@@ -549,29 +574,30 @@ def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=Non
|
||||
if layer_class_name in allowed_types and module.weight.dtype in {torch.float32, torch.float16, torch.bfloat16}:
|
||||
if (layer_class_name in conv_types or layer_class_name in conv_transpose_types) and not quant_conv:
|
||||
continue
|
||||
module, modules_to_not_convert, modules_dtype_dict = sdnq_quantize_layer(
|
||||
module,
|
||||
weights_dtype=get_minimum_dtype(weights_dtype, param_name, modules_dtype_dict),
|
||||
quantized_matmul_dtype=quantized_matmul_dtype,
|
||||
torch_dtype=torch_dtype,
|
||||
group_size=group_size,
|
||||
svd_rank=svd_rank,
|
||||
svd_steps=svd_steps,
|
||||
dynamic_loss_threshold=dynamic_loss_threshold,
|
||||
use_svd=use_svd,
|
||||
quant_conv=quant_conv,
|
||||
use_quantized_matmul=use_quantized_matmul,
|
||||
use_quantized_matmul_conv=use_quantized_matmul_conv,
|
||||
use_dynamic_quantization=use_dynamic_quantization,
|
||||
use_stochastic_rounding=use_stochastic_rounding,
|
||||
dequantize_fp32=dequantize_fp32,
|
||||
non_blocking=non_blocking,
|
||||
quantization_device=quantization_device,
|
||||
return_device=return_device,
|
||||
modules_to_not_convert=modules_to_not_convert,
|
||||
modules_dtype_dict=modules_dtype_dict,
|
||||
param_name=param_name,
|
||||
)
|
||||
quant_kwargs = {
|
||||
"weights_dtype": weights_dtype,
|
||||
"quantized_matmul_dtype": quantized_matmul_dtype,
|
||||
"torch_dtype": torch_dtype,
|
||||
"group_size": group_size,
|
||||
"svd_rank": svd_rank,
|
||||
"svd_steps": svd_steps,
|
||||
"dynamic_loss_threshold": dynamic_loss_threshold,
|
||||
"use_svd": use_svd,
|
||||
"quant_conv": quant_conv,
|
||||
"use_quantized_matmul": use_quantized_matmul,
|
||||
"use_quantized_matmul_conv": use_quantized_matmul_conv,
|
||||
"use_dynamic_quantization": use_dynamic_quantization,
|
||||
"use_stochastic_rounding": use_stochastic_rounding,
|
||||
"dequantize_fp32": dequantize_fp32,
|
||||
"non_blocking": non_blocking,
|
||||
"quantization_device": quantization_device,
|
||||
"return_device": return_device,
|
||||
"modules_to_not_convert": modules_to_not_convert,
|
||||
"modules_dtype_dict": modules_dtype_dict,
|
||||
"param_name": param_name,
|
||||
}
|
||||
quant_kwargs = get_quant_kwargs(quant_kwargs, modules_quant_config)
|
||||
module, modules_to_not_convert, modules_dtype_dict = sdnq_quantize_layer(module, **quant_kwargs)
|
||||
setattr(model, module_name, module)
|
||||
|
||||
module, modules_to_not_convert, modules_dtype_dict = apply_sdnq_to_module(
|
||||
@@ -595,6 +621,7 @@ def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=Non
|
||||
return_device=return_device,
|
||||
modules_to_not_convert=modules_to_not_convert,
|
||||
modules_dtype_dict=modules_dtype_dict,
|
||||
modules_quant_config=modules_quant_config,
|
||||
full_param_name=param_name,
|
||||
)
|
||||
setattr(model, module_name, module)
|
||||
@@ -620,18 +647,22 @@ def sdnq_post_load_quant(
|
||||
dequantize_fp32: bool = False,
|
||||
non_blocking: bool = False,
|
||||
add_skip_keys:bool = True,
|
||||
modules_to_not_convert: List[str] = None,
|
||||
modules_dtype_dict: Dict[str, List[str]] = None,
|
||||
quantization_device: Optional[torch.device] = None,
|
||||
return_device: Optional[torch.device] = None,
|
||||
modules_to_not_convert: Optional[List[str]] = None,
|
||||
modules_dtype_dict: Optional[Dict[str, List[str]]] = None,
|
||||
modules_quant_config: Optional[Dict[str, dict]] = None,
|
||||
):
|
||||
if modules_to_not_convert is None:
|
||||
modules_to_not_convert = []
|
||||
if modules_dtype_dict is None:
|
||||
modules_dtype_dict = {}
|
||||
if modules_quant_config is None:
|
||||
modules_quant_config = {}
|
||||
|
||||
modules_to_not_convert = modules_to_not_convert.copy()
|
||||
modules_dtype_dict = modules_dtype_dict.copy()
|
||||
modules_quant_config = modules_quant_config.copy()
|
||||
if add_skip_keys:
|
||||
model, modules_to_not_convert, modules_dtype_dict = add_module_skip_keys(model, modules_to_not_convert, modules_dtype_dict)
|
||||
|
||||
@@ -652,6 +683,7 @@ def sdnq_post_load_quant(
|
||||
add_skip_keys=add_skip_keys,
|
||||
modules_to_not_convert=modules_to_not_convert,
|
||||
modules_dtype_dict=modules_dtype_dict,
|
||||
modules_quant_config=modules_quant_config,
|
||||
quantization_device=quantization_device,
|
||||
return_device=return_device,
|
||||
)
|
||||
@@ -676,12 +708,14 @@ def sdnq_post_load_quant(
|
||||
non_blocking=non_blocking,
|
||||
modules_to_not_convert=modules_to_not_convert,
|
||||
modules_dtype_dict=modules_dtype_dict,
|
||||
modules_quant_config=modules_quant_config,
|
||||
quantization_device=quantization_device,
|
||||
return_device=return_device,
|
||||
)
|
||||
|
||||
quantization_config.modules_to_not_convert = modules_to_not_convert
|
||||
quantization_config.modules_dtype_dict = modules_dtype_dict
|
||||
quantization_config.modules_quant_config = modules_quant_config
|
||||
|
||||
model.quantization_config = quantization_config
|
||||
if hasattr(model, "config"):
|
||||
@@ -798,16 +832,37 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
|
||||
return
|
||||
|
||||
torch_dtype = kwargs.get("dtype", param_value.dtype if self.torch_dtype is None else self.torch_dtype)
|
||||
weights_dtype = get_minimum_dtype(self.quantization_config.weights_dtype, param_name, self.quantization_config.modules_dtype_dict)
|
||||
|
||||
if self.quantization_config.return_device is not None:
|
||||
return_device = self.quantization_config.return_device
|
||||
else:
|
||||
return_device = target_device
|
||||
|
||||
if self.quantization_config.quantization_device is not None:
|
||||
target_device = self.quantization_config.quantization_device
|
||||
|
||||
quant_kwargs = {
|
||||
"weights_dtype": self.quantization_config.weights_dtype,
|
||||
"quantized_matmul_dtype": self.quantization_config.quantized_matmul_dtype,
|
||||
"torch_dtype": torch_dtype,
|
||||
"group_size": self.quantization_config.group_size,
|
||||
"svd_rank": self.quantization_config.svd_rank,
|
||||
"svd_steps": self.quantization_config.svd_steps,
|
||||
"dynamic_loss_threshold": self.quantization_config.dynamic_loss_threshold,
|
||||
"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,
|
||||
"use_dynamic_quantization": self.quantization_config.use_dynamic_quantization,
|
||||
"use_stochastic_rounding": self.quantization_config.use_stochastic_rounding,
|
||||
"dequantize_fp32": self.quantization_config.dequantize_fp32,
|
||||
"non_blocking": self.quantization_config.non_blocking,
|
||||
"modules_to_not_convert": self.quantization_config.modules_to_not_convert,
|
||||
"modules_dtype_dict": self.quantization_config.modules_dtype_dict,
|
||||
"quantization_device": None,
|
||||
"return_device": return_device,
|
||||
"param_name": param_name,
|
||||
}
|
||||
quant_kwargs = get_quant_kwargs(quant_kwargs, self.quantization_config.modules_quant_config)
|
||||
|
||||
if param_value.dtype == torch.float32 and devices.same_device(param_value.device, target_device):
|
||||
param_value = param_value.clone()
|
||||
else:
|
||||
@@ -815,29 +870,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
|
||||
|
||||
layer, tensor_name = get_module_from_name(model, param_name)
|
||||
layer.weight = torch.nn.Parameter(param_value, requires_grad=False)
|
||||
layer, self.quantization_config.modules_to_not_convert, self.quantization_config.modules_dtype_dict = sdnq_quantize_layer(
|
||||
layer,
|
||||
weights_dtype=weights_dtype,
|
||||
quantized_matmul_dtype=self.quantization_config.quantized_matmul_dtype,
|
||||
torch_dtype=torch_dtype,
|
||||
group_size=self.quantization_config.group_size,
|
||||
svd_rank=self.quantization_config.svd_rank,
|
||||
svd_steps=self.quantization_config.svd_steps,
|
||||
dynamic_loss_threshold=self.quantization_config.dynamic_loss_threshold,
|
||||
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,
|
||||
use_dynamic_quantization=self.quantization_config.use_dynamic_quantization,
|
||||
use_stochastic_rounding=self.quantization_config.use_stochastic_rounding,
|
||||
dequantize_fp32=self.quantization_config.dequantize_fp32,
|
||||
non_blocking=self.quantization_config.non_blocking,
|
||||
modules_to_not_convert=self.quantization_config.modules_to_not_convert,
|
||||
modules_dtype_dict=self.quantization_config.modules_dtype_dict,
|
||||
quantization_device=None,
|
||||
return_device=return_device,
|
||||
param_name=param_name,
|
||||
)
|
||||
layer, self.quantization_config.modules_to_not_convert, self.quantization_config.modules_dtype_dict = sdnq_quantize_layer(layer, **quant_kwargs)
|
||||
|
||||
layer.weight._is_hf_initialized = True # pylint: disable=protected-access
|
||||
if hasattr(layer, "scale"):
|
||||
@@ -1005,10 +1038,13 @@ class SDNQConfig(QuantizationConfigMixin):
|
||||
return_device (`torch.device`, *optional*, defaults to `None`):
|
||||
Used to set which device will the quantized weights be sent back to.
|
||||
modules_to_not_convert (`list`, *optional*, default to `None`):
|
||||
The list of modules to not quantize, useful for quantizing models that explicitly require to have some
|
||||
The list of modules to not quantize. Useful for quantizing models that explicitly require to have some
|
||||
modules left in their original precision (e.g. Whisper encoder, Llava encoder, Mixtral gate layers).
|
||||
modules_dtype_dict (`dict`, *optional*, default to `None`):
|
||||
The dict of dtypes and list of modules, useful for quantizing some modules with a different dtype.
|
||||
The dict of dtypes and list of modules. Useful for quantizing some modules with a different dtype.
|
||||
modules_quant_config (`dict`, *optional*, default to `None`):
|
||||
The dict of modules and a dict of quantization kwargs to use for that module.
|
||||
Useful for quantizing some modules with a different quantization config.
|
||||
"""
|
||||
|
||||
def __init__( # pylint: disable=super-init-not-called
|
||||
@@ -1034,6 +1070,7 @@ class SDNQConfig(QuantizationConfigMixin):
|
||||
return_device: Optional[torch.device] = None,
|
||||
modules_to_not_convert: Optional[List[str]] = None,
|
||||
modules_dtype_dict: Optional[Dict[str, List[str]]] = None,
|
||||
modules_quant_config: Optional[Dict[str, dict]] = None,
|
||||
is_training: bool = False,
|
||||
**kwargs, # pylint: disable=unused-argument
|
||||
):
|
||||
@@ -1063,6 +1100,7 @@ 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.modules_quant_config = modules_quant_config
|
||||
self.is_integer = dtype_dict[self.weights_dtype]["is_integer"]
|
||||
self.sdnq_version = sdnq_version
|
||||
self.post_init()
|
||||
@@ -1103,8 +1141,12 @@ class SDNQConfig(QuantizationConfigMixin):
|
||||
if not isinstance(key, str) or not isinstance(value, list):
|
||||
raise ValueError(f"modules_dtype_dict must be a dictionary of strings and lists but got {type(key)} and {type(value)}")
|
||||
|
||||
if self.modules_quant_config is None:
|
||||
self.modules_quant_config = {}
|
||||
|
||||
self.modules_to_not_convert = self.modules_to_not_convert.copy()
|
||||
self.modules_dtype_dict = self.modules_dtype_dict.copy()
|
||||
self.modules_quant_config = self.modules_quant_config.copy()
|
||||
|
||||
def to_dict(self):
|
||||
quantization_config_dict = self.__dict__.copy() # make serializable
|
||||
|
||||
Reference in New Issue
Block a user