mirror of
https://github.com/vladmandic/automatic
synced 2026-09-17 16:24:33 +02:00
Refactor SDNQ quantizer handling and add modules_to_not_use_matmul
This commit is contained in:
@@ -1,4 +1,4 @@
|
||||
from .quantizer import QuantizationMethod, SDNQConfig, SDNQQuantizer, sdnq_post_load_quant, apply_sdnq_to_module, sdnq_quantize_layer
|
||||
from .quantizer import QuantizationMethod, SDNQConfig, SDNQQuantizer, apply_sdnq_to_module, sdnq_post_load_quant, sdnq_quantize_layer
|
||||
from .loader import save_sdnq_model, load_sdnq_model
|
||||
from .common import sdnq_version
|
||||
|
||||
|
||||
+44
-20
@@ -6,7 +6,7 @@ import torch
|
||||
|
||||
from modules import shared, devices
|
||||
|
||||
sdnq_version = "0.1.8"
|
||||
sdnq_version = "0.1.9"
|
||||
|
||||
dtype_dict = {
|
||||
### Integers
|
||||
@@ -437,85 +437,109 @@ common_skip_keys = (
|
||||
"wte",
|
||||
)
|
||||
|
||||
|
||||
# modules_to_not_convert: ["x_embedder", "y_embedder"]
|
||||
# modules_to_not_use_matmul: {"int8": ["x_embedder", "y_embedder"], "float8_e4m3fn": ["x_embedder", "y_embedder"]}
|
||||
# modules_dtype_dict: {"minimum_6bit": ["x_embedder", "y_embedder"]}
|
||||
|
||||
module_skip_keys_dict = {
|
||||
"FluxTransformer2DModel": [
|
||||
["single_transformer_blocks.0.norm.linear.weight", "time_text_embed", "time_embed", "context_embedder", "x_embedder", ".proj_out", "norm_out"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"Flux2Transformer2DModel": [
|
||||
["double_stream_modulation_img", "double_stream_modulation_txt", "single_stream_modulation", "time_guidance_embed", "context_embedder", "x_embedder", ".proj_out", "norm_out"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"ChromaTransformer2DModel": [
|
||||
["distilled_guidance_layer", "time_text_embed", "context_embedder", "x_embedder", ".proj_out", "norm_out"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"QwenImageTransformer2DModel": [
|
||||
["transformer_blocks.0.img_mod.1.weight", "time_text_embed", "txt_in", "img_in", "proj_out", "norm_out"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"WanTransformer3DModel": [
|
||||
["scale_shift_table", "patch_embedding", "condition_embedder", "proj_out", "norm_out"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"LongCatVideoTransformer3DModel": [
|
||||
["blocks.0.adaLN_modulation.1.weight", "x_embedder", "t_embedder", "y_embedder", "final_layer"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"LTX2VideoTransformer3DModel": [
|
||||
[
|
||||
"audio_time_embed", "time_embed", "audio_caption_projection", "caption_projection", "proj_in", "audio_proj_in", "proj_out", "audio_proj_out",
|
||||
"av_cross_attn_audio_scale_shift", "av_cross_attn_audio_v2a_gate", "av_cross_attn_video_a2v_gate", "av_cross_attn_video_scale_shift",
|
||||
],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"Lumina2Transformer2DModel": [
|
||||
["layers.0.norm1.linear.weight", "time_caption_embed", "x_embedder", "norm_out"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"ZImageTransformer2DModel": [
|
||||
["layers.0.adaLN_modulation.0.weight", "t_embedder", "cap_embedder", "siglip_embedder", "all_x_embedder", "all_final_layer"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"CosmosTransformer3DModel": [
|
||||
["transformer_blocks.0.norm*", "patch_embed", "time_embed", "norm_out", "proj_out", "crossattn_proj"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"GlmImageTransformer2DModel": [
|
||||
["transformer_blocks.0.norm1.linear.weight", "image_projector", "glyph_projector", "prior_projector", "time_condition_embed", "norm_out", "proj_out"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"GlmImageForConditionalGeneration": [
|
||||
["lm_head", "patch_embed", "embeddings", "embed_tokens", "vqmodel"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"HunyuanImage3ForCausalMM": [
|
||||
["lm_head", "patch_embed", "time_embed", "time_embed_2", "final_layer", "wte", "ln_f", "timestep_emb", "vae", "vision_aligner", "head", "post_layernorm", "embeddings"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"Emu3ForCausalLM": [
|
||||
["lm_head", "vq_model", "tokenizer"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"Gemma3nForCausalLM": [
|
||||
["lm_head", "correction_coefs", "prediction_coefs", "embedding_projection"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"Gemma4ForConditionalGeneration": [
|
||||
["lm_head", "embed_audio", "embed_vision", "patch_embedder", "embed_tokens", "subsample_conv_projection", "output_proj"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"MoondreamModel": [
|
||||
["lm_head", "region", "wte", "post_ln", "proj_mlp", "patch_emb", "pos_emb"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"NaDiT": [
|
||||
[".emb_in", ".txt_in", ".vid_in", ".emb_scale", ".vid_out", ".vid_out_norm", ".vid_out_ada"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
"HiDreamO1Qwen3VLTransformer": [
|
||||
["lm_head", "embed_tokens", "x_embedder", "t_embedder1", "final_layer2", "patch_embed", "pos_embed"],
|
||||
{}
|
||||
{},
|
||||
{},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
@@ -50,17 +50,7 @@ def int8_matmul(
|
||||
|
||||
def quantized_linear_forward_int8_matmul(self, input: torch.FloatTensor) -> torch.FloatTensor:
|
||||
if torch.numel(input) / input.shape[-1] < 32:
|
||||
dequantized_weight = self.sdnq_dequantizer(
|
||||
self.weight,
|
||||
self.scale,
|
||||
self.zero_point,
|
||||
self.svd_up,
|
||||
self.svd_down,
|
||||
skip_quantized_matmul=True,
|
||||
)
|
||||
if input.dtype != dequantized_weight.dtype:
|
||||
input = input.to(dtype=dequantized_weight.dtype)
|
||||
return torch.nn.functional.linear(input, dequantized_weight, self.bias)
|
||||
return torch.nn.functional.linear(input, self.sdnq_dequantizer(self.weight, self.scale, self.zero_point, self.svd_up, self.svd_down, skip_quantized_matmul=True), self.bias)
|
||||
if self.sdnq_dequantizer.re_quantize_for_matmul:
|
||||
weight, scale = self.sdnq_dequantizer.re_quantize_matmul(self.weight, self.scale, self.zero_point, None, None)
|
||||
quantized_weight_shape = None
|
||||
|
||||
+14
-4
@@ -4,7 +4,9 @@ import torch
|
||||
from diffusers.models.modeling_utils import ModelMixin
|
||||
|
||||
from .common import dtype_dict, is_fp8_mm_supported, use_tensorwise_fp8_matmul, check_torch_compile, conv_types, linear_types
|
||||
from .quantizer import SDNQConfig, sdnq_post_load_quant, prepare_weight_for_matmul, prepare_svd_for_matmul, get_quant_args_from_config
|
||||
from .quantizer import SDNQConfig, sdnq_post_load_quant
|
||||
from .quant_utils import prepare_weight_for_matmul, prepare_svd_for_matmul
|
||||
from .utils import get_quant_args_from_config, check_param_name_in
|
||||
from .forward import get_forward_func
|
||||
from .file_loader import load_files
|
||||
|
||||
@@ -191,16 +193,24 @@ def post_process_model(model):
|
||||
return model
|
||||
|
||||
|
||||
def apply_sdnq_options_to_module(model, dtype: torch.dtype | None = None, dequantize_fp32: bool | None = None, use_quantized_matmul: bool | None = None):
|
||||
def apply_sdnq_options_to_module(model, quantization_config: SDNQConfig, dtype: torch.dtype | None = None, dequantize_fp32: bool | None = None, use_quantized_matmul: bool | None = None, full_param_name: str = ""):
|
||||
has_children = list(model.children())
|
||||
if not has_children:
|
||||
if dtype is not None and getattr(model, "dtype", torch.float32) not in {torch.float32, torch.float64}:
|
||||
model = model.to(dtype=dtype)
|
||||
return model
|
||||
for module_name, module in model.named_children():
|
||||
if full_param_name:
|
||||
param_name = full_param_name + "." + module_name
|
||||
else:
|
||||
param_name = module_name
|
||||
if hasattr(module, "sdnq_dequantizer"):
|
||||
layer_class_name = module.original_class.__name__
|
||||
current_use_quantized_matmul = use_quantized_matmul
|
||||
if layer_class_name in conv_types:
|
||||
current_use_quantized_matmul = None
|
||||
elif check_param_name_in(param_name, quantization_config.modules_to_not_use_matmul) is not None:
|
||||
current_use_quantized_matmul = None
|
||||
|
||||
if not is_fp8_mm_supported and module.sdnq_dequantizer.quantized_matmul_dtype in {"fp8", "float8_e4m3fn"}:
|
||||
current_use_quantized_matmul = False
|
||||
@@ -260,14 +270,14 @@ def apply_sdnq_options_to_module(model, dtype: torch.dtype | None = None, dequan
|
||||
module.forward_func = get_forward_func(module.original_class.__name__, module.sdnq_dequantizer.quantized_matmul_dtype, current_use_quantized_matmul)
|
||||
setattr(model, module_name, module)
|
||||
else:
|
||||
setattr(model, module_name, apply_sdnq_options_to_module(module, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul))
|
||||
setattr(model, module_name, apply_sdnq_options_to_module(module, quantization_config, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul, full_param_name=param_name))
|
||||
return model
|
||||
|
||||
|
||||
def apply_sdnq_options_to_model(model, dtype: torch.dtype | None = None, dequantize_fp32: bool | None = None, use_quantized_matmul: bool | None = None):
|
||||
if use_quantized_matmul and not check_torch_compile():
|
||||
raise RuntimeError("SDNQ Quantized MatMul requires a working Triton install.")
|
||||
model = apply_sdnq_options_to_module(model, dtype=dtype, dequantize_fp32=dequantize_fp32, use_quantized_matmul=use_quantized_matmul)
|
||||
model = apply_sdnq_options_to_module(model, model.quantization_config, 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
|
||||
|
||||
@@ -0,0 +1,92 @@
|
||||
import torch
|
||||
|
||||
from modules import devices
|
||||
from .common import dtype_dict, use_contiguous_mm
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def get_scale_asymmetric(weight: torch.FloatTensor, reduction_axes: int | list[int], weights_dtype: str) -> tuple[torch.FloatTensor, torch.FloatTensor]:
|
||||
zero_point = torch.amin(weight, dim=reduction_axes, keepdims=True)
|
||||
scale = torch.amax(weight, dim=reduction_axes, keepdims=True).sub_(zero_point).div_(dtype_dict[weights_dtype]["max"] - dtype_dict[weights_dtype]["min"])
|
||||
if dtype_dict[weights_dtype]["min"] != 0:
|
||||
zero_point.sub_(torch.mul(scale, dtype_dict[weights_dtype]["min"]))
|
||||
return scale, zero_point
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def get_scale_symmetric(weight: torch.FloatTensor, reduction_axes: int | list[int], weights_dtype: str) -> torch.FloatTensor:
|
||||
return torch.amax(weight.abs(), dim=reduction_axes, keepdims=True).div_(dtype_dict[weights_dtype]["max"])
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def quantize_weight(weight: torch.FloatTensor, reduction_axes: int | list[int], weights_dtype: str, dtype: torch.dtype = None, use_stochastic_rounding: bool = False) -> tuple[torch.Tensor, torch.FloatTensor, torch.FloatTensor]:
|
||||
if weight.dtype != torch.float64:
|
||||
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)
|
||||
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:
|
||||
quantized_weight.add_(torch.randn_like(quantized_weight), alpha=0.1)
|
||||
quantized_weight.round_()
|
||||
else:
|
||||
if use_stochastic_rounding:
|
||||
mantissa_difference = 1 << (23 - dtype_dict[weights_dtype]["mantissa"])
|
||||
quantized_weight = quantized_weight.to(dtype=torch.float32).view(dtype=torch.int32)
|
||||
quantized_weight = quantized_weight.add_(torch.randint_like(quantized_weight, low=0, high=mantissa_difference, dtype=torch.int32)).bitwise_and_(-mantissa_difference).view(dtype=torch.float32)
|
||||
quantized_weight.nan_to_num_()
|
||||
quantized_weight = quantized_weight.clamp_(dtype_dict[weights_dtype]["min"], dtype_dict[weights_dtype]["max"]).to(dtype_dict[weights_dtype]["torch_dtype"])
|
||||
return quantized_weight, scale, zero_point
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
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
|
||||
weight_shape = weight.shape
|
||||
weight = weight.flatten(1,-1)
|
||||
if weight.dtype != torch.float64:
|
||||
weight = weight.to(dtype=torch.float32)
|
||||
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
|
||||
return weight, svd_up, svd_down
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def prepare_weight_for_matmul(weight: torch.Tensor) -> torch.Tensor:
|
||||
if use_contiguous_mm:
|
||||
weight = weight.contiguous()
|
||||
elif weight.is_contiguous():
|
||||
weight = weight.t_().contiguous().t_()
|
||||
return weight
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def prepare_svd_for_matmul(svd_up: torch.FloatTensor, svd_down: torch.FloatTensor, use_quantized_matmul: bool) -> tuple[torch.FloatTensor, torch.FloatTensor]:
|
||||
if svd_up is not None:
|
||||
if use_quantized_matmul:
|
||||
svd_up = prepare_weight_for_matmul(svd_up)
|
||||
else:
|
||||
svd_up = svd_up.contiguous()
|
||||
if svd_down is not None:
|
||||
svd_down = prepare_weight_for_matmul(svd_down)
|
||||
return svd_up, svd_down
|
||||
+108
-467
@@ -1,10 +1,8 @@
|
||||
# pylint: disable=redefined-builtin,no-member,protected-access
|
||||
|
||||
from typing import Union
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
|
||||
import re
|
||||
import torch
|
||||
|
||||
from transformers.quantizers import HfQuantizer
|
||||
@@ -15,235 +13,22 @@ from diffusers.utils import get_module_from_name
|
||||
from accelerate import init_empty_weights
|
||||
|
||||
from modules import devices, shared
|
||||
from .common import sdnq_version, dtype_dict, common_skip_keys, module_skip_keys_dict, accepted_weight_dtypes, accepted_matmul_dtypes, weights_dtype_order, allowed_types, linear_types, embedding_types, conv_types, conv_transpose_types, compile_func, is_fp8_mm_supported, use_tensorwise_fp8_matmul, use_contiguous_mm, check_torch_compile
|
||||
from .common import sdnq_version, dtype_dict, accepted_weight_dtypes, accepted_matmul_dtypes, weights_dtype_order, allowed_types, linear_types, embedding_types, conv_types, conv_transpose_types, compile_func, is_fp8_mm_supported, use_tensorwise_fp8_matmul, check_torch_compile
|
||||
from .dequantizer import SDNQDequantizer, dequantize_sdnq_model
|
||||
from .packed_int import pack_int
|
||||
from .packed_float import pack_float
|
||||
from .forward import get_forward_func
|
||||
from .layers import get_sdnq_wrapper_class
|
||||
|
||||
from .quant_utils import quantize_weight, apply_svdquant, prepare_weight_for_matmul, prepare_svd_for_matmul
|
||||
from .utils import check_param_name_in, get_quant_args_from_config, get_quant_kwargs, add_module_skip_keys
|
||||
|
||||
|
||||
class QuantizationMethod(str, Enum):
|
||||
SDNQ = "sdnq"
|
||||
SDNQ_TRAINING = "sdnq_training"
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def get_scale_asymmetric(weight: torch.FloatTensor, reduction_axes: int | list[int], weights_dtype: str) -> tuple[torch.FloatTensor, torch.FloatTensor]:
|
||||
zero_point = torch.amin(weight, dim=reduction_axes, keepdims=True)
|
||||
scale = torch.amax(weight, dim=reduction_axes, keepdims=True).sub_(zero_point).div_(dtype_dict[weights_dtype]["max"] - dtype_dict[weights_dtype]["min"])
|
||||
if dtype_dict[weights_dtype]["min"] != 0:
|
||||
zero_point.sub_(torch.mul(scale, dtype_dict[weights_dtype]["min"]))
|
||||
return scale, zero_point
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def get_scale_symmetric(weight: torch.FloatTensor, reduction_axes: int | list[int], weights_dtype: str) -> torch.FloatTensor:
|
||||
return torch.amax(weight.abs(), dim=reduction_axes, keepdims=True).div_(dtype_dict[weights_dtype]["max"])
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def quantize_weight(weight: torch.FloatTensor, reduction_axes: int | list[int], weights_dtype: str, dtype: torch.dtype = None, use_stochastic_rounding: bool = False) -> tuple[torch.Tensor, torch.FloatTensor, torch.FloatTensor]:
|
||||
if weight.dtype != torch.float64:
|
||||
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)
|
||||
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:
|
||||
quantized_weight.add_(torch.randn_like(quantized_weight), alpha=0.1)
|
||||
quantized_weight.round_()
|
||||
else:
|
||||
if use_stochastic_rounding:
|
||||
mantissa_difference = 1 << (23 - dtype_dict[weights_dtype]["mantissa"])
|
||||
quantized_weight = quantized_weight.to(dtype=torch.float32).view(dtype=torch.int32)
|
||||
quantized_weight = quantized_weight.add_(torch.randint_like(quantized_weight, low=0, high=mantissa_difference, dtype=torch.int32)).bitwise_and_(-mantissa_difference).view(dtype=torch.float32)
|
||||
quantized_weight.nan_to_num_()
|
||||
quantized_weight = quantized_weight.clamp_(dtype_dict[weights_dtype]["min"], dtype_dict[weights_dtype]["max"]).to(dtype_dict[weights_dtype]["torch_dtype"])
|
||||
return quantized_weight, scale, zero_point
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
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
|
||||
weight_shape = weight.shape
|
||||
weight = weight.flatten(1,-1)
|
||||
if weight.dtype != torch.float64:
|
||||
weight = weight.to(dtype=torch.float32)
|
||||
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
|
||||
return weight, svd_up, svd_down
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def prepare_weight_for_matmul(weight: torch.Tensor) -> torch.Tensor:
|
||||
if use_contiguous_mm:
|
||||
weight = weight.contiguous()
|
||||
elif weight.is_contiguous():
|
||||
weight = weight.t_().contiguous().t_()
|
||||
return weight
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def prepare_svd_for_matmul(svd_up: torch.FloatTensor, svd_down: torch.FloatTensor, use_quantized_matmul: bool) -> tuple[torch.FloatTensor, torch.FloatTensor]:
|
||||
if svd_up is not None:
|
||||
if use_quantized_matmul:
|
||||
svd_up = prepare_weight_for_matmul(svd_up)
|
||||
else:
|
||||
svd_up = svd_up.contiguous()
|
||||
if svd_down is not None:
|
||||
svd_down = prepare_weight_for_matmul(svd_down)
|
||||
return svd_up, svd_down
|
||||
|
||||
|
||||
def check_param_name_in(param_name: str, param_list: list[str]) -> str:
|
||||
split_param_name = param_name.split(".")
|
||||
for param in param_list:
|
||||
if param.startswith("."):
|
||||
if param_name.startswith(param[1:]):
|
||||
return param
|
||||
else:
|
||||
continue
|
||||
if (
|
||||
param_name == param
|
||||
or param in split_param_name
|
||||
or ("*" in param and re.match(param.replace(".*", "\\.*").replace("*", ".*"), param_name))
|
||||
):
|
||||
return param
|
||||
return None
|
||||
|
||||
|
||||
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_dynamic_quantization", 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)
|
||||
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
|
||||
|
||||
|
||||
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():
|
||||
if check_param_name_in(param_name, value) is not None:
|
||||
key = key.lower()
|
||||
if key.startswith("minimum") or key.endswith("bit") or key.endswith("bits"):
|
||||
minimum_bits_str = key.removeprefix("minimum").removeprefix("-").removeprefix("_").removesuffix("bits").removesuffix("bit").removesuffix("-").removesuffix("_")
|
||||
if minimum_bits_str.startswith("uint"):
|
||||
is_unsigned = True
|
||||
minimum_bits_str = minimum_bits_str.removeprefix("uint")
|
||||
else:
|
||||
is_unsigned = False
|
||||
minimum_bits_str = minimum_bits_str.removeprefix("int")
|
||||
minimum_bits = int(minimum_bits_str)
|
||||
if dtype_dict[weights_dtype]["num_bits"] < minimum_bits:
|
||||
if is_unsigned or minimum_bits <= 4:
|
||||
return "uint" + minimum_bits_str
|
||||
else:
|
||||
return "int" + minimum_bits_str
|
||||
else:
|
||||
return key
|
||||
return weights_dtype
|
||||
|
||||
|
||||
def get_quant_kwargs(quant_kwargs: dict, modules_quant_config: dict[str, dict]) -> dict:
|
||||
param_key = check_param_name_in(quant_kwargs["param_name"], modules_quant_config.keys())
|
||||
if param_key is not None:
|
||||
for key, value in modules_quant_config[param_key].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 update_modules_quant_config(quant_kwargs: dict, modules_quant_config: dict[str, dict], layer: torch.nn.Module) -> dict[str, dict]:
|
||||
layer_class_name = layer.__class__.__name__
|
||||
if layer_class_name in conv_types:
|
||||
use_quantized_matmul_key = "use_quantized_matmul_conv"
|
||||
else:
|
||||
use_quantized_matmul_key = "use_quantized_matmul"
|
||||
if (
|
||||
hasattr(layer, "sdnq_dequantizer")
|
||||
and (layer_class_name in linear_types or layer_class_name in conv_types)
|
||||
and quant_kwargs["use_dynamic_quantization"] and quant_kwargs[use_quantized_matmul_key]
|
||||
and quant_kwargs["quantized_matmul_dtype"] is None and not is_fp8_mm_supported
|
||||
and not dtype_dict[layer.sdnq_dequantizer.weights_dtype]["is_integer"] and dtype_dict[layer.sdnq_dequantizer.weights_dtype]["num_bits"] < 16
|
||||
and not layer.sdnq_dequantizer.use_quantized_matmul
|
||||
):
|
||||
if quant_kwargs["param_name"] not in modules_quant_config.keys():
|
||||
modules_quant_config[quant_kwargs["param_name"]] = {}
|
||||
modules_quant_config[quant_kwargs["param_name"]][use_quantized_matmul_key] = False
|
||||
return modules_quant_config
|
||||
|
||||
|
||||
def add_module_skip_keys(model, modules_to_not_convert: list[str] | None = None, modules_dtype_dict: dict[str, list[str]] | None = None):
|
||||
if modules_to_not_convert is None:
|
||||
modules_to_not_convert = []
|
||||
if modules_dtype_dict is None:
|
||||
modules_dtype_dict = {}
|
||||
if getattr(model, "_keep_in_fp32_modules", None) is not None:
|
||||
modules_to_not_convert.extend(model._keep_in_fp32_modules) # pylint: disable=protected-access
|
||||
if getattr(model, "_tied_weights_keys", None) is not None:
|
||||
if isinstance(model._tied_weights_keys, dict): # pylint: disable=protected-access
|
||||
modules_to_not_convert.extend(model._tied_weights_keys.keys()) # pylint: disable=protected-access
|
||||
modules_to_not_convert.extend(model._tied_weights_keys.values()) # pylint: disable=protected-access
|
||||
else:
|
||||
modules_to_not_convert.extend(model._tied_weights_keys) # pylint: disable=protected-access
|
||||
|
||||
skip_key_list = module_skip_keys_dict.get(model.__class__.__name__, None)
|
||||
if skip_key_list is not None:
|
||||
modules_to_not_convert.extend(skip_key_list[0])
|
||||
for key, value in skip_key_list[1].items():
|
||||
if key in modules_dtype_dict.keys():
|
||||
modules_dtype_dict[key].extend(value)
|
||||
else:
|
||||
modules_dtype_dict[key] = value
|
||||
else:
|
||||
modules_to_not_convert.extend(common_skip_keys)
|
||||
if getattr(model, "_skip_layerwise_casting_patterns", None) is not None:
|
||||
modules_to_not_convert.extend(model._skip_layerwise_casting_patterns) # pylint: disable=protected-access
|
||||
|
||||
# dedupe
|
||||
modules_to_not_convert = list(set(modules_to_not_convert))
|
||||
for key, value in modules_dtype_dict.items():
|
||||
modules_dtype_dict[key] = list(set(value))
|
||||
|
||||
return model, modules_to_not_convert, modules_dtype_dict
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, use_quantized_matmul=False, use_stochastic_rounding=False, dequantize_fp32=True, using_pre_calculated_svd=False, skip_sr=False, param_name=None): # pylint: disable=unused-argument
|
||||
num_of_groups = 1
|
||||
@@ -286,9 +71,7 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
|
||||
is_conv_type = True
|
||||
reduction_axes = 1
|
||||
output_channel_size, channel_size = weight.shape[:2]
|
||||
if use_quantized_matmul:
|
||||
use_quantized_matmul = channel_size >= 32 and output_channel_size >= 32
|
||||
use_quantized_matmul = use_quantized_matmul and output_channel_size % 16 == 0 and channel_size % 16 == 0
|
||||
use_quantized_matmul = use_quantized_matmul and channel_size >= 32 and output_channel_size >= 32 and output_channel_size % 16 == 0 and channel_size % 16 == 0
|
||||
if use_quantized_matmul and not re_quantize_for_matmul and not dtype_dict[weights_dtype]["is_packed"]:
|
||||
result_shape = weight.shape
|
||||
weight = weight.flatten(1,-1)
|
||||
@@ -305,9 +88,7 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
|
||||
output_channel_size, channel_size = weight.shape
|
||||
except Exception as e:
|
||||
raise ValueError(f"SDNQ: param_name={param_name} layer_class_name={layer_class_name} weight_shape={weight.shape} weights_dtype={weights_dtype} quantized_matmul_dtype={quantized_matmul_dtype} unsupported") from e
|
||||
if use_quantized_matmul:
|
||||
use_quantized_matmul = channel_size >= 32 and output_channel_size >= 32
|
||||
use_quantized_matmul = use_quantized_matmul and output_channel_size % 16 == 0 and channel_size % 16 == 0
|
||||
use_quantized_matmul = use_quantized_matmul and channel_size >= 32 and output_channel_size >= 32 and output_channel_size % 16 == 0 and channel_size % 16 == 0
|
||||
else:
|
||||
if weight.ndim > 1:
|
||||
output_channel_size, channel_size = weight.shape[-2:]
|
||||
@@ -404,12 +185,21 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
|
||||
weight.t_()
|
||||
weight = prepare_weight_for_matmul(weight)
|
||||
|
||||
quantized_weight_shape = weight.shape
|
||||
if dtype_dict[weights_dtype]["is_packed"]:
|
||||
if dtype_dict[weights_dtype]["is_integer"]:
|
||||
weight = pack_int(weight, weights_dtype)
|
||||
else:
|
||||
weight = pack_float(weight, weights_dtype)
|
||||
else:
|
||||
weight = weight.to(dtype=dtype_dict[weights_dtype]["torch_dtype"])
|
||||
|
||||
sdnq_dequantizer = SDNQDequantizer(
|
||||
result_dtype=torch_dtype,
|
||||
result_shape=result_shape,
|
||||
original_shape=original_shape,
|
||||
original_stride=original_stride,
|
||||
quantized_weight_shape=weight.shape,
|
||||
quantized_weight_shape=quantized_weight_shape,
|
||||
weights_dtype=weights_dtype,
|
||||
quantized_matmul_dtype=quantized_matmul_dtype,
|
||||
group_size=group_size,
|
||||
@@ -421,19 +211,11 @@ def sdnq_quantize_layer_weight(weight, layer_class_name=None, weights_dtype="int
|
||||
layer_class_name=layer_class_name,
|
||||
)
|
||||
|
||||
if dtype_dict[weights_dtype]["is_packed"]:
|
||||
if dtype_dict[weights_dtype]["is_integer"]:
|
||||
weight = pack_int(weight, weights_dtype)
|
||||
else:
|
||||
weight = pack_float(weight, weights_dtype)
|
||||
else:
|
||||
weight = weight.to(dtype=dtype_dict[weights_dtype]["torch_dtype"])
|
||||
|
||||
return weight, scale, zero_point, svd_up, svd_down, sdnq_dequantizer
|
||||
return sdnq_dequantizer, {"weight": weight, "scale": scale, "zero_point": zero_point, "svd_up": svd_up, "svd_down": svd_down}
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def sdnq_quantize_layer_weight_dynamic(weight, layer_class_name=None, weights_dtype="uint4", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, dynamic_loss_threshold=None, use_svd=False, use_quantized_matmul=False, use_dynamic_quantization=False, use_stochastic_rounding=False, dequantize_fp32=True, param_name=None): # pylint: disable=unused-argument
|
||||
def sdnq_quantize_layer_weight_dynamic(weight, layer_class_name=None, weights_dtype="uint4", quantized_matmul_dtype=None, group_size=0, svd_rank=32, svd_steps=8, dynamic_loss_threshold=None, use_svd=False, use_quantized_matmul=False, use_stochastic_rounding=False, dequantize_fp32=True, quantization_config=None, torch_dtype=None, param_name=None): # pylint: disable=unused-argument
|
||||
if torch_dtype is None:
|
||||
torch_dtype = weight.dtype
|
||||
if dynamic_loss_threshold is None or dynamic_loss_threshold < 0:
|
||||
@@ -463,7 +245,7 @@ def sdnq_quantize_layer_weight_dynamic(weight, layer_class_name=None, weights_dt
|
||||
else:
|
||||
current_use_quantized_matmul = use_quantized_matmul
|
||||
|
||||
quantized_weight, scale, zero_point, _, _, sdnq_dequantizer = sdnq_quantize_layer_weight(
|
||||
sdnq_dequantizer, weight_data = sdnq_quantize_layer_weight(
|
||||
svd_weight,
|
||||
layer_class_name=layer_class_name,
|
||||
weights_dtype=current_weights_dtype,
|
||||
@@ -485,23 +267,32 @@ def sdnq_quantize_layer_weight_dynamic(weight, layer_class_name=None, weights_dt
|
||||
svd_down = svd_down.t_()
|
||||
svd_is_transposed = True
|
||||
|
||||
quantization_loss = torch.nn.functional.mse_loss(weight, sdnq_dequantizer(quantized_weight, scale, zero_point, svd_up, svd_down, skip_quantized_matmul=sdnq_dequantizer.use_quantized_matmul, dtype=weight.dtype, skip_compile=True)).div_(weight_std)
|
||||
weight_data["svd_up"] = svd_up
|
||||
weight_data["svd_down"] = svd_down
|
||||
|
||||
quantization_loss = torch.nn.functional.mse_loss(weight, sdnq_dequantizer(**weight_data, skip_quantized_matmul=sdnq_dequantizer.use_quantized_matmul, dtype=weight.dtype, skip_compile=True)).div_(weight_std)
|
||||
if quantization_loss <= dynamic_loss_threshold:
|
||||
return (quantized_weight, scale, zero_point, svd_up, svd_down, sdnq_dequantizer)
|
||||
return sdnq_dequantizer, weight_data
|
||||
return None
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
def sdnq_quantize_layer(layer, weights_dtype="int8", quantized_matmul_dtype=None, torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, dynamic_loss_threshold=None, use_svd=False, quant_conv=False, quant_embedding=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_dynamic_quantization=False, use_stochastic_rounding=False, dequantize_fp32=True, non_blocking=False, modules_to_not_convert=None, modules_dtype_dict=None, quantization_device=None, return_device=None, param_name=None): # pylint: disable=unused-argument
|
||||
def sdnq_quantize_layer(layer, quantization_config: "SDNQConfig", torch_dtype: torch.dtype | None = None, param_name: str = "", quant_kwargs: dict | None = None): # pylint: disable=unused-argument
|
||||
if quant_kwargs is None:
|
||||
quant_kwargs = get_quant_kwargs(layer, quantization_config, torch_dtype=torch_dtype, param_name=param_name)
|
||||
|
||||
layer_class_name = layer.__class__.__name__
|
||||
if layer_class_name in embedding_types:
|
||||
if not quant_embedding:
|
||||
return layer, modules_to_not_convert, modules_dtype_dict
|
||||
use_quantized_matmul = False
|
||||
elif layer_class_name in conv_transpose_types or layer_class_name in conv_types:
|
||||
if not quant_conv:
|
||||
return layer, modules_to_not_convert, modules_dtype_dict
|
||||
use_quantized_matmul = use_quantized_matmul_conv
|
||||
if (
|
||||
(layer_class_name in embedding_types and not quantization_config.quant_embedding)
|
||||
or ((layer_class_name in conv_transpose_types or layer_class_name in conv_types) and not quantization_config.quant_conv)
|
||||
):
|
||||
quantization_config.modules_to_not_convert.append(param_name)
|
||||
return layer, quantization_config
|
||||
|
||||
return_device = quant_kwargs.pop("return_device")
|
||||
quantization_device = quant_kwargs.pop("quantization_device")
|
||||
non_blocking = quant_kwargs.pop("non_blocking")
|
||||
use_dynamic_quantization = quant_kwargs.pop("use_dynamic_quantization")
|
||||
|
||||
layer.weight.requires_grad_(False)
|
||||
if return_device is None:
|
||||
@@ -510,85 +301,42 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", quantized_matmul_dtype=None
|
||||
layer.weight.data = layer.weight.to(quantization_device, non_blocking=non_blocking)
|
||||
|
||||
if use_dynamic_quantization:
|
||||
weight_data = sdnq_quantize_layer_weight_dynamic(
|
||||
layer.weight,
|
||||
layer_class_name=layer_class_name,
|
||||
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,
|
||||
use_quantized_matmul=use_quantized_matmul,
|
||||
use_stochastic_rounding=use_stochastic_rounding,
|
||||
dequantize_fp32=dequantize_fp32,
|
||||
param_name=param_name,
|
||||
)
|
||||
weight_data = sdnq_quantize_layer_weight_dynamic(layer.weight, **quant_kwargs)
|
||||
else:
|
||||
weight_data = sdnq_quantize_layer_weight(
|
||||
layer.weight,
|
||||
layer_class_name=layer_class_name,
|
||||
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,
|
||||
use_svd=use_svd,
|
||||
use_quantized_matmul=use_quantized_matmul,
|
||||
use_stochastic_rounding=use_stochastic_rounding,
|
||||
dequantize_fp32=dequantize_fp32,
|
||||
param_name=param_name,
|
||||
)
|
||||
weight_data = sdnq_quantize_layer_weight(layer.weight, **quant_kwargs)
|
||||
|
||||
if weight_data is not None:
|
||||
(
|
||||
layer.weight.data,
|
||||
layer.scale, layer.zero_point,
|
||||
layer.svd_up, layer.svd_down,
|
||||
layer.sdnq_dequantizer,
|
||||
) = weight_data
|
||||
layer.sdnq_dequantizer, weight_data = weight_data
|
||||
layer = get_sdnq_wrapper_class(layer, get_forward_func(layer_class_name, layer.sdnq_dequantizer.quantized_matmul_dtype, layer.sdnq_dequantizer.use_quantized_matmul))
|
||||
|
||||
for key, value in weight_data.items():
|
||||
if isinstance(value, (torch.Tensor, torch.nn.Parameter)):
|
||||
setattr(layer, key, torch.nn.Parameter(value.to(return_device, non_blocking=non_blocking), requires_grad=False))
|
||||
setattr(getattr(layer, key), "_is_hf_initialized", True)
|
||||
else:
|
||||
setattr(layer, key, value)
|
||||
del weight_data
|
||||
|
||||
layer = get_sdnq_wrapper_class(layer, get_forward_func(layer_class_name, layer.sdnq_dequantizer.quantized_matmul_dtype, layer.sdnq_dequantizer.use_quantized_matmul))
|
||||
layer.weight = torch.nn.Parameter(layer.weight.to(return_device, non_blocking=non_blocking), requires_grad=False)
|
||||
layer.scale = torch.nn.Parameter(layer.scale.to(return_device, non_blocking=non_blocking), requires_grad=False)
|
||||
if layer.zero_point is not None:
|
||||
layer.zero_point = torch.nn.Parameter(layer.zero_point.to(return_device, non_blocking=non_blocking), requires_grad=False)
|
||||
if layer.svd_up is not None:
|
||||
layer.svd_up = torch.nn.Parameter(layer.svd_up.to(return_device, non_blocking=non_blocking), requires_grad=False)
|
||||
layer.svd_down = torch.nn.Parameter(layer.svd_down.to(return_device, non_blocking=non_blocking), requires_grad=False)
|
||||
|
||||
if use_dynamic_quantization:
|
||||
if modules_dtype_dict is None:
|
||||
modules_dtype_dict = {}
|
||||
if layer.sdnq_dequantizer.weights_dtype not in modules_dtype_dict.keys():
|
||||
modules_dtype_dict[layer.sdnq_dequantizer.weights_dtype] = [param_name]
|
||||
if layer.sdnq_dequantizer.weights_dtype not in quantization_config.modules_dtype_dict.keys():
|
||||
quantization_config.modules_dtype_dict[layer.sdnq_dequantizer.weights_dtype] = [param_name]
|
||||
else:
|
||||
modules_dtype_dict[layer.sdnq_dequantizer.weights_dtype].append(param_name)
|
||||
quantization_config.modules_dtype_dict[layer.sdnq_dequantizer.weights_dtype].append(param_name)
|
||||
|
||||
if quant_kwargs["use_quantized_matmul"] and not layer.sdnq_dequantizer.use_quantized_matmul:
|
||||
quantization_config.modules_to_not_use_matmul.append(param_name)
|
||||
else:
|
||||
layer.weight = layer.weight.to(return_device, dtype=torch_dtype, non_blocking=non_blocking)
|
||||
if use_dynamic_quantization:
|
||||
if modules_to_not_convert is None:
|
||||
modules_to_not_convert = []
|
||||
modules_to_not_convert.append(param_name)
|
||||
quantization_config.modules_to_not_convert.append(param_name)
|
||||
|
||||
return layer, modules_to_not_convert, modules_dtype_dict
|
||||
return layer, quantization_config
|
||||
|
||||
|
||||
@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=None, use_svd=False, quant_conv=False, quant_embedding=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, use_dynamic_quantization=False, use_stochastic_rounding=False, dequantize_fp32=True, non_blocking=False, modules_to_not_convert: list[str] | None = None, modules_dtype_dict: dict[str, list[str]] | None = None, modules_quant_config: dict[str, dict] | None = 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, modules_quant_config)
|
||||
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 = {}
|
||||
def apply_sdnq_to_module(model, quantization_config: "SDNQConfig", torch_dtype: torch.dtype | None = None, full_param_name: str = ""): # pylint: disable=unused-argument
|
||||
if not list(model.children()):
|
||||
return model
|
||||
for module_name, module in model.named_children():
|
||||
if full_param_name:
|
||||
param_name = full_param_name + "." + module_name
|
||||
@@ -596,69 +344,22 @@ def apply_sdnq_to_module(model, weights_dtype="int8", quantized_matmul_dtype=Non
|
||||
param_name = module_name
|
||||
if hasattr(module, "weight") and module.weight is not None:
|
||||
param_name = param_name + ".weight"
|
||||
if check_param_name_in(param_name, modules_to_not_convert) is not None:
|
||||
continue
|
||||
layer_class_name = module.__class__.__name__
|
||||
if layer_class_name in allowed_types and module.weight.dtype in {torch.float64, torch.float32, torch.float16, torch.bfloat16}:
|
||||
if layer_class_name in embedding_types and not quant_embedding:
|
||||
continue
|
||||
if (layer_class_name in conv_types or layer_class_name in conv_transpose_types) and not quant_conv:
|
||||
continue
|
||||
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,
|
||||
"quant_embedding": quant_embedding,
|
||||
"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)
|
||||
modules_quant_config = update_modules_quant_config(quant_kwargs, modules_quant_config, module)
|
||||
if (
|
||||
layer_class_name in allowed_types
|
||||
and module.weight.dtype in {torch.float64, torch.float32, torch.float16, torch.bfloat16}
|
||||
and check_param_name_in(param_name, quantization_config.modules_to_not_convert) is None
|
||||
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)
|
||||
):
|
||||
module, quantization_config = sdnq_quantize_layer(module, quantization_config, torch_dtype=torch_dtype, param_name=param_name)
|
||||
setattr(model, module_name, module)
|
||||
else:
|
||||
quantization_config.modules_to_not_convert.append(param_name)
|
||||
|
||||
module, (modules_to_not_convert, modules_dtype_dict, modules_quant_config) = apply_sdnq_to_module(
|
||||
module,
|
||||
dynamic_loss_threshold=dynamic_loss_threshold,
|
||||
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,
|
||||
use_svd=use_svd,
|
||||
quant_conv=quant_conv,
|
||||
quant_embedding=quant_embedding,
|
||||
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,
|
||||
modules_quant_config=modules_quant_config,
|
||||
full_param_name=param_name,
|
||||
)
|
||||
module, quantization_config = apply_sdnq_to_module(module, quantization_config, torch_dtype=torch_dtype, full_param_name=param_name)
|
||||
setattr(model, module_name, module)
|
||||
return model, (modules_to_not_convert, modules_dtype_dict, modules_quant_config)
|
||||
return model, quantization_config
|
||||
|
||||
|
||||
@devices.inference_context()
|
||||
@@ -684,22 +385,10 @@ def sdnq_post_load_quant(
|
||||
quantization_device: torch.device | None = None,
|
||||
return_device: torch.device | None = None,
|
||||
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,
|
||||
modules_quant_config: dict[str, dict] | None = 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)
|
||||
|
||||
quantization_config = SDNQConfig(
|
||||
weights_dtype=weights_dtype,
|
||||
group_size=group_size,
|
||||
@@ -717,41 +406,17 @@ def sdnq_post_load_quant(
|
||||
non_blocking=non_blocking,
|
||||
add_skip_keys=add_skip_keys,
|
||||
modules_to_not_convert=modules_to_not_convert,
|
||||
modules_to_not_use_matmul=modules_to_not_use_matmul,
|
||||
modules_dtype_dict=modules_dtype_dict,
|
||||
modules_quant_config=modules_quant_config,
|
||||
quantization_device=quantization_device,
|
||||
return_device=return_device,
|
||||
)
|
||||
if add_skip_keys:
|
||||
model, quantization_config = add_module_skip_keys(model, quantization_config)
|
||||
|
||||
model.eval()
|
||||
model, (modules_to_not_convert, modules_dtype_dict, modules_quant_config) = apply_sdnq_to_module(
|
||||
model,
|
||||
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,
|
||||
quant_embedding=quant_embedding,
|
||||
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,
|
||||
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 = apply_sdnq_to_module(model, quantization_config, torch_dtype=torch_dtype)
|
||||
|
||||
model.quantization_config = quantization_config
|
||||
if hasattr(model, "config"):
|
||||
@@ -826,6 +491,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
|
||||
return True
|
||||
else:
|
||||
return True
|
||||
self.quantization_config.modules_to_not_convert.append(param_name)
|
||||
return False
|
||||
|
||||
def check_quantized_param(self, *args, **kwargs) -> bool:
|
||||
@@ -849,8 +515,9 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
|
||||
target_device: torch.device,
|
||||
*args, **kwargs, # pylint: disable=unused-argument
|
||||
):
|
||||
layer, tensor_name = get_module_from_name(model, param_name)
|
||||
|
||||
if self.pre_quantized:
|
||||
layer, tensor_name = get_module_from_name(model, param_name)
|
||||
if param_value is not None:
|
||||
if tensor_name == "weight":
|
||||
return_dtype = param_value.dtype
|
||||
@@ -880,32 +547,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
|
||||
return
|
||||
|
||||
torch_dtype = kwargs.get("dtype", param_value.dtype if self.torch_dtype is None else self.torch_dtype)
|
||||
|
||||
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,
|
||||
"quant_embedding": self.quantization_config.quant_embedding,
|
||||
"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": self.quantization_config.quantization_device,
|
||||
"return_device": self.quantization_config.return_device,
|
||||
"param_name": param_name,
|
||||
}
|
||||
quant_kwargs = get_quant_kwargs(quant_kwargs, self.quantization_config.modules_quant_config)
|
||||
|
||||
quant_kwargs = get_quant_kwargs(layer, self.quantization_config, torch_dtype=torch_dtype, param_name=param_name)
|
||||
if quant_kwargs["return_device"] is None:
|
||||
quant_kwargs["return_device"] = target_device
|
||||
if quant_kwargs["quantization_device"] is not None:
|
||||
@@ -917,19 +559,9 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
|
||||
else:
|
||||
param_value = param_value.to(target_device, non_blocking=self.quantization_config.non_blocking).to(dtype=torch.float32 if param_value.dtype != torch.float64 else torch.float64)
|
||||
|
||||
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, **quant_kwargs)
|
||||
self.quantization_config.modules_quant_config = update_modules_quant_config(quant_kwargs, self.quantization_config.modules_quant_config, layer)
|
||||
layer, self.quantization_config = sdnq_quantize_layer(layer, self.quantization_config, torch_dtype=torch_dtype, param_name=param_name, quant_kwargs=quant_kwargs)
|
||||
|
||||
layer.weight._is_hf_initialized = True # pylint: disable=protected-access
|
||||
if hasattr(layer, "scale"):
|
||||
layer.scale._is_hf_initialized = True # pylint: disable=protected-access
|
||||
if layer.zero_point is not None:
|
||||
layer.zero_point._is_hf_initialized = True # pylint: disable=protected-access
|
||||
if layer.svd_up is not None:
|
||||
layer.svd_up._is_hf_initialized = True # pylint: disable=protected-access
|
||||
layer.svd_down._is_hf_initialized = True # pylint: disable=protected-access
|
||||
parent_module, tensor_name = get_module_from_name(model, param_name.removesuffix(tensor_name).removesuffix("."))
|
||||
setattr(parent_module, tensor_name, layer)
|
||||
|
||||
@@ -943,16 +575,6 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
|
||||
def adjust_target_dtype(self, target_dtype: torch.dtype) -> torch.dtype: # pylint: disable=unused-argument,arguments-renamed
|
||||
return dtype_dict[self.quantization_config.weights_dtype]["target_dtype"]
|
||||
|
||||
def update_torch_dtype(self, torch_dtype: torch.dtype | None = None) -> torch.dtype:
|
||||
self.torch_dtype = torch_dtype
|
||||
return torch_dtype
|
||||
|
||||
def update_dtype(self, dtype: torch.dtype | None = None) -> torch.dtype:
|
||||
"""
|
||||
needed for transformers compatibilty, returns self.update_torch_dtype
|
||||
"""
|
||||
return self.update_torch_dtype(dtype)
|
||||
|
||||
def _process_model_before_weight_loading( # pylint: disable=arguments-differ
|
||||
self,
|
||||
model,
|
||||
@@ -974,9 +596,12 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
|
||||
self.quantization_config.modules_to_not_convert.extend(keep_in_fp32_modules)
|
||||
if hasattr(self, "get_modules_to_not_convert") and hasattr(model, "tie_weights"):
|
||||
self.quantization_config.modules_to_not_convert.extend(self.get_modules_to_not_convert(model, add_default_skips=True))
|
||||
model, self.quantization_config.modules_to_not_convert, self.quantization_config.modules_dtype_dict = add_module_skip_keys(
|
||||
model, self.quantization_config.modules_to_not_convert, self.quantization_config.modules_dtype_dict
|
||||
)
|
||||
model, self.quantization_config = add_module_skip_keys(model, self.quantization_config)
|
||||
|
||||
|
||||
def _process_model_after_weight_loading(self, model, **kwargs): # pylint: disable=unused-argument
|
||||
model.quantization_config = self.quantization_config
|
||||
model.quantization_method = QuantizationMethod.SDNQ
|
||||
if hasattr(model, "config"):
|
||||
try:
|
||||
model.config.quantization_config = self.quantization_config
|
||||
@@ -986,10 +611,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
|
||||
model.config["quantization_config"] = self.quantization_config.to_dict()
|
||||
except Exception:
|
||||
pass
|
||||
model.quantization_config = self.quantization_config
|
||||
model.quantization_method = QuantizationMethod.SDNQ
|
||||
|
||||
def _process_model_after_weight_loading(self, model, **kwargs): # pylint: disable=unused-argument
|
||||
if self.pre_quantized:
|
||||
from .loader import post_process_model
|
||||
model = post_process_model(model)
|
||||
@@ -1004,6 +626,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer):
|
||||
use_stochastic_rounding=self.quantization_config.use_stochastic_rounding,
|
||||
dequantize_fp32=self.quantization_config.dequantize_fp32,
|
||||
)
|
||||
|
||||
if shared.opts.diffusers_offload_mode != "none":
|
||||
try:
|
||||
model = model.to(device=devices.cpu)
|
||||
@@ -1123,6 +746,7 @@ class SDNQConfig(QuantizationConfigMixin):
|
||||
quantization_device: torch.device | None = None,
|
||||
return_device: torch.device | None = None,
|
||||
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,
|
||||
modules_quant_config: dict[str, dict] | None = None,
|
||||
is_training: bool = False,
|
||||
@@ -1154,6 +778,7 @@ class SDNQConfig(QuantizationConfigMixin):
|
||||
self.quantization_device = quantization_device
|
||||
self.return_device = return_device
|
||||
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
|
||||
self.modules_quant_config = modules_quant_config
|
||||
self.is_integer = dtype_dict[self.weights_dtype]["is_integer"]
|
||||
@@ -1180,6 +805,15 @@ class SDNQConfig(QuantizationConfigMixin):
|
||||
elif not isinstance(self.modules_to_not_convert, list):
|
||||
raise ValueError(f"modules_to_not_convert must be a list but got {type(self.modules_to_not_convert)}")
|
||||
|
||||
if self.modules_to_not_use_matmul is None:
|
||||
self.modules_to_not_use_matmul = []
|
||||
elif isinstance(self.modules_to_not_use_matmul, str):
|
||||
self.modules_to_not_use_matmul = [self.modules_to_not_use_matmul]
|
||||
elif isinstance(self.modules_to_not_use_matmul, tuple):
|
||||
self.modules_to_not_use_matmul = list(self.modules_to_not_use_matmul)
|
||||
elif not isinstance(self.modules_to_not_use_matmul, list):
|
||||
raise ValueError(f"modules_to_not_use_matmul must be a list but got {type(self.modules_to_not_use_matmul)}")
|
||||
|
||||
if self.modules_dtype_dict is None:
|
||||
self.modules_dtype_dict = {}
|
||||
elif not isinstance(self.modules_dtype_dict, dict):
|
||||
@@ -1200,9 +834,16 @@ class SDNQConfig(QuantizationConfigMixin):
|
||||
self.modules_quant_config = {}
|
||||
|
||||
self.modules_to_not_convert = self.modules_to_not_convert.copy()
|
||||
self.modules_to_not_use_matmul = self.modules_to_not_use_matmul.copy()
|
||||
self.modules_dtype_dict = self.modules_dtype_dict.copy()
|
||||
self.modules_quant_config = self.modules_quant_config.copy()
|
||||
|
||||
# dedupe
|
||||
self.modules_to_not_convert = list(set(self.modules_to_not_convert))
|
||||
self.modules_to_not_use_matmul = list(set(self.modules_to_not_use_matmul))
|
||||
for key, value in self.modules_dtype_dict.items():
|
||||
self.modules_dtype_dict[key] = list(set(value))
|
||||
|
||||
def to_dict(self):
|
||||
quantization_config_dict = self.__dict__.copy() # make serializable
|
||||
quantization_config_dict["quantization_device"] = str(quantization_config_dict["quantization_device"]) if quantization_config_dict["quantization_device"] is not None else None
|
||||
|
||||
@@ -0,0 +1,162 @@
|
||||
import re
|
||||
import torch
|
||||
|
||||
from .common import dtype_dict, common_skip_keys, module_skip_keys_dict, conv_types, conv_transpose_types
|
||||
|
||||
|
||||
def check_param_name_in(param_name: str, param_list: list[str]) -> str:
|
||||
split_param_name = param_name.split(".")
|
||||
for param in param_list:
|
||||
if param.startswith("."):
|
||||
if param_name.startswith(param[1:]):
|
||||
return param
|
||||
else:
|
||||
continue
|
||||
if (
|
||||
param_name == param
|
||||
or param in split_param_name
|
||||
or ("*" in param and re.match(param.replace(".*", "\\.*").replace("*", ".*"), param_name))
|
||||
):
|
||||
return param
|
||||
return None
|
||||
|
||||
|
||||
def get_quant_args_from_config(quantization_config: dict) -> dict:
|
||||
from .quantizer import SDNQConfig
|
||||
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_dynamic_quantization", 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)
|
||||
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
|
||||
|
||||
|
||||
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():
|
||||
if check_param_name_in(param_name, value) is not None:
|
||||
key = key.lower()
|
||||
if key.startswith("minimum") or key.endswith("bit") or key.endswith("bits"):
|
||||
minimum_bits_str = key.removeprefix("minimum").removeprefix("-").removeprefix("_").removesuffix("bits").removesuffix("bit").removesuffix("-").removesuffix("_")
|
||||
if minimum_bits_str.startswith("uint"):
|
||||
is_unsigned = True
|
||||
minimum_bits_str = minimum_bits_str.removeprefix("uint")
|
||||
else:
|
||||
is_unsigned = False
|
||||
minimum_bits_str = minimum_bits_str.removeprefix("int")
|
||||
minimum_bits = int(minimum_bits_str)
|
||||
if dtype_dict[weights_dtype]["num_bits"] < minimum_bits:
|
||||
if is_unsigned or minimum_bits <= 4:
|
||||
return "uint" + minimum_bits_str
|
||||
else:
|
||||
return "int" + minimum_bits_str
|
||||
else:
|
||||
return key
|
||||
return weights_dtype
|
||||
|
||||
|
||||
def get_quant_kwargs(layer: torch.nn.Module, quantization_config, torch_dtype: torch.dtype | None = None, param_name: str = "", **kwargs) -> dict:
|
||||
from .quantizer import SDNQConfig
|
||||
if not isinstance(quantization_config, SDNQConfig):
|
||||
quantization_config = SDNQConfig(**quantization_config)
|
||||
layer_class_name = layer.__class__.__name__
|
||||
|
||||
quant_kwargs = {
|
||||
"weights_dtype": quantization_config.weights_dtype,
|
||||
"quantized_matmul_dtype": quantization_config.quantized_matmul_dtype,
|
||||
"group_size": quantization_config.group_size,
|
||||
"svd_rank": quantization_config.svd_rank,
|
||||
"svd_steps": quantization_config.svd_steps,
|
||||
"dynamic_loss_threshold": quantization_config.dynamic_loss_threshold,
|
||||
"use_svd": quantization_config.use_svd,
|
||||
"use_quantized_matmul": quantization_config.use_quantized_matmul,
|
||||
"use_quantized_matmul_conv": quantization_config.use_quantized_matmul_conv,
|
||||
"use_dynamic_quantization": quantization_config.use_dynamic_quantization,
|
||||
"use_stochastic_rounding": quantization_config.use_stochastic_rounding,
|
||||
"dequantize_fp32": quantization_config.dequantize_fp32,
|
||||
"non_blocking": quantization_config.non_blocking,
|
||||
"quantization_device": quantization_config.quantization_device,
|
||||
"return_device": quantization_config.return_device,
|
||||
"layer_class_name": layer_class_name,
|
||||
"torch_dtype": torch_dtype,
|
||||
"param_name": param_name,
|
||||
}
|
||||
|
||||
for key, value in kwargs.items():
|
||||
quant_kwargs[key] = value
|
||||
|
||||
param_key = check_param_name_in(quant_kwargs["param_name"], quantization_config.modules_quant_config.keys())
|
||||
if param_key is not None:
|
||||
for key, value in quantization_config.modules_quant_config[param_key].items():
|
||||
quant_kwargs[key] = value
|
||||
|
||||
if layer_class_name in conv_transpose_types or layer_class_name in conv_types:
|
||||
quant_kwargs["use_quantized_matmul"] = quant_kwargs.pop("use_quantized_matmul_conv")
|
||||
else:
|
||||
quant_kwargs.pop("use_quantized_matmul_conv")
|
||||
|
||||
if not quant_kwargs["use_dynamic_quantization"]:
|
||||
quant_kwargs.pop("dynamic_loss_threshold")
|
||||
|
||||
quant_kwargs["weights_dtype"] = get_minimum_dtype(quant_kwargs["weights_dtype"], quant_kwargs["param_name"], quantization_config.modules_dtype_dict)
|
||||
if check_param_name_in(quant_kwargs["param_name"], quantization_config.modules_to_not_use_matmul) is not None:
|
||||
quant_kwargs["use_quantized_matmul"] = False
|
||||
|
||||
return quant_kwargs
|
||||
|
||||
|
||||
def add_module_skip_keys(model: torch.nn.Module, quantization_config):
|
||||
if getattr(model, "_keep_in_fp32_modules", None) is not None:
|
||||
quantization_config.modules_to_not_convert.extend(model._keep_in_fp32_modules) # pylint: disable=protected-access
|
||||
if getattr(model, "_tied_weights_keys", None) is not None:
|
||||
if isinstance(model._tied_weights_keys, dict): # pylint: disable=protected-access
|
||||
quantization_config.modules_to_not_convert.extend(model._tied_weights_keys.keys()) # pylint: disable=protected-access
|
||||
quantization_config.modules_to_not_convert.extend(model._tied_weights_keys.values()) # pylint: disable=protected-access
|
||||
else:
|
||||
quantization_config.modules_to_not_convert.extend(model._tied_weights_keys) # pylint: disable=protected-access
|
||||
|
||||
skip_key_list = module_skip_keys_dict.get(model.__class__.__name__, None)
|
||||
if skip_key_list is not None:
|
||||
quantization_config.modules_to_not_convert.extend(skip_key_list[0])
|
||||
for key, value in skip_key_list[1].items():
|
||||
if key in quantization_config.modules_dtype_dict.keys():
|
||||
quantization_config.modules_dtype_dict[key].extend(value)
|
||||
else:
|
||||
quantization_config.modules_dtype_dict[key] = value
|
||||
|
||||
if quantization_config.quantized_matmul_dtype is None:
|
||||
if dtype_dict[quantization_config.weights_dtype]["is_integer"]:
|
||||
quantized_matmul_dtype = "int8"
|
||||
elif dtype_dict[quantization_config.weights_dtype]["num_bits"] < 16:
|
||||
quantized_matmul_dtype = "float8_e4m3fn"
|
||||
else:
|
||||
quantized_matmul_dtype = "float16"
|
||||
else:
|
||||
quantized_matmul_dtype = quantization_config.quantized_matmul_dtype
|
||||
quantization_config.modules_to_not_use_matmul.extend(skip_key_list[2].get(quantized_matmul_dtype, []))
|
||||
else:
|
||||
quantization_config.modules_to_not_convert.extend(common_skip_keys)
|
||||
if getattr(model, "_skip_layerwise_casting_patterns", None) is not None:
|
||||
quantization_config.modules_to_not_convert.extend(model._skip_layerwise_casting_patterns) # pylint: disable=protected-access
|
||||
|
||||
# dedupe
|
||||
quantization_config.modules_to_not_convert = list(set(quantization_config.modules_to_not_convert))
|
||||
quantization_config.modules_to_not_use_matmul = list(set(quantization_config.modules_to_not_use_matmul))
|
||||
for key, value in quantization_config.modules_dtype_dict.items():
|
||||
quantization_config.modules_dtype_dict[key] = list(set(value))
|
||||
|
||||
return model, quantization_config
|
||||
Reference in New Issue
Block a user