mirror of
https://github.com/vladmandic/automatic
synced 2026-09-19 17:24:32 +02:00
SDNQ add minimum_allowed_channel_size
This commit is contained in:
@@ -496,6 +496,7 @@ def sdnq_post_load_quant(
|
||||
non_blocking: bool = False,
|
||||
add_skip_keys:bool = True,
|
||||
minimum_allowed_numel: int = 16384,
|
||||
minimum_allowed_channel_size: int = 32,
|
||||
modules_to_not_convert: list[str] | None = None,
|
||||
modules_to_not_use_matmul: list[str] | None = None,
|
||||
modules_dtype_dict: dict[str, list[str]] | None = None,
|
||||
@@ -535,6 +536,7 @@ def sdnq_post_load_quant(
|
||||
non_blocking=non_blocking,
|
||||
add_skip_keys=add_skip_keys,
|
||||
minimum_allowed_numel=minimum_allowed_numel,
|
||||
minimum_allowed_channel_size=minimum_allowed_channel_size,
|
||||
modules_to_not_convert=modules_to_not_convert,
|
||||
modules_to_not_use_matmul=modules_to_not_use_matmul,
|
||||
modules_dtype_dict=modules_dtype_dict,
|
||||
@@ -873,6 +875,8 @@ class SDNQConfig(QuantizationConfigMixin):
|
||||
Disabling this option won't add model specific keys to modules_to_not_convert, modules_to_not_use_matmul and modules_dtype_dict.
|
||||
minimum_allowed_numel (`int`, *optional*, defaults to `16384`):
|
||||
Layers that have less than `minimum_allowed_numel` elements in them will be skipped and added to `modules_to_not_convert`.
|
||||
minimum_allowed_channel_size (`int`, *optional*, defaults to `32`):
|
||||
Layers that have less than `minimum_allowed_channel_size` channels in them will be skipped and added to `modules_to_not_convert`.
|
||||
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
|
||||
modules left in their original precision (e.g. Whisper encoder, Llava encoder, Mixtral gate layers).
|
||||
@@ -917,6 +921,7 @@ class SDNQConfig(QuantizationConfigMixin):
|
||||
non_blocking: bool = False,
|
||||
add_skip_keys: bool = True,
|
||||
minimum_allowed_numel: int = 16384,
|
||||
minimum_allowed_channel_size: int = 32,
|
||||
modules_to_not_convert: list[str] | None = None,
|
||||
modules_to_not_use_matmul: list[str] | None = None,
|
||||
modules_dtype_dict: dict[str, list[str]] | None = None,
|
||||
@@ -949,6 +954,7 @@ class SDNQConfig(QuantizationConfigMixin):
|
||||
self.non_blocking = non_blocking
|
||||
self.add_skip_keys = add_skip_keys
|
||||
self.minimum_allowed_numel = minimum_allowed_numel
|
||||
self.minimum_allowed_channel_size = minimum_allowed_channel_size
|
||||
self.modules_to_not_convert = modules_to_not_convert
|
||||
self.modules_to_not_use_matmul = modules_to_not_use_matmul
|
||||
self.modules_dtype_dict = modules_dtype_dict
|
||||
|
||||
+13
-3
@@ -44,13 +44,23 @@ def check_param_name_in(param_name: str, param_list: list[str]) -> str:
|
||||
|
||||
|
||||
def check_quant_is_allowed(layer_class_name: str, weight: torch.Tensor, quantization_config, pre_quantized: bool = False) -> bool:
|
||||
return bool(
|
||||
if (
|
||||
layer_class_name in allowed_types
|
||||
and weight.dtype in {torch.float64, torch.float32, torch.float16, torch.bfloat16}
|
||||
and not (layer_class_name in embedding_types and not quantization_config.quant_embedding)
|
||||
and not ((layer_class_name in conv_types or layer_class_name in conv_transpose_types) and not quantization_config.quant_conv)
|
||||
and (pre_quantized or weight.numel() >= quantization_config.minimum_allowed_numel)
|
||||
)
|
||||
):
|
||||
if pre_quantized:
|
||||
return True
|
||||
if layer_class_name in conv_types:
|
||||
channel_size = weight.shape[1]
|
||||
elif layer_class_name in conv_transpose_types:
|
||||
channel_size = weight.shape[0]
|
||||
else:
|
||||
channel_size = weight.shape[-1]
|
||||
if channel_size >= quantization_config.minimum_allowed_channel_size and weight.numel() >= quantization_config.minimum_allowed_numel:
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def check_quantized_matmul_is_allowed(use_quantized_matmul: bool, output_channel_size: int, channel_size: int) -> bool:
|
||||
|
||||
Reference in New Issue
Block a user