diff --git a/modules/lora/lora_apply.py b/modules/lora/lora_apply.py index 2221fbf3e..f18922040 100644 --- a/modules/lora/lora_apply.py +++ b/modules/lora/lora_apply.py @@ -196,6 +196,7 @@ def network_add_weights(self: Union[torch.nn.Conv2d, torch.nn.Linear, torch.nn.G torch_dtype=devices.dtype, group_size=shared.opts.sdnq_quantize_weights_group_size, svd_rank=shared.opts.sdnq_svd_rank, + svd_steps=shared.opts.sdnq_svd_steps, use_svd=shared.opts.sdnq_use_svd, quant_conv=shared.opts.sdnq_quantize_conv_layers, use_quantized_matmul=shared.opts.sdnq_use_quantized_matmul, diff --git a/modules/model_quant.py b/modules/model_quant.py index 5f88d3174..0bb72a75a 100644 --- a/modules/model_quant.py +++ b/modules/model_quant.py @@ -206,6 +206,7 @@ def create_sdnq_config(kwargs = None, allow: bool = True, module: str = 'Model', weights_dtype=weights_dtype, group_size=shared.opts.sdnq_quantize_weights_group_size, svd_rank=shared.opts.sdnq_svd_rank, + svd_steps=shared.opts.sdnq_svd_steps, use_svd=shared.opts.sdnq_use_svd, quant_conv=shared.opts.sdnq_quantize_conv_layers, use_quantized_matmul=shared.opts.sdnq_use_quantized_matmul, @@ -217,7 +218,7 @@ def create_sdnq_config(kwargs = None, allow: bool = True, module: str = 'Model', modules_to_not_convert=modules_to_not_convert, modules_dtype_dict=modules_dtype_dict.copy(), ) - log.debug(f'Quantization: module="{module}" type=sdnq mode=pre dtype={weights_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} group_size={shared.opts.sdnq_quantize_weights_group_size} svd_rank={shared.opts.sdnq_svd_rank} use_svd={shared.opts.sdnq_use_svd} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} dequantize_fp32={shared.opts.sdnq_dequantize_fp32} quantize_with_gpu={shared.opts.sdnq_quantize_with_gpu} quantization_device={quantization_device} return_device={return_device} device_map={shared.opts.device_map} offload_mode={shared.opts.diffusers_offload_mode} non_blocking={shared.opts.diffusers_offload_nonblocking} modules_to_not_convert={modules_to_not_convert} modules_dtype_dict={modules_dtype_dict}') + log.debug(f'Quantization: module="{module}" type=sdnq mode=pre dtype={weights_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} group_size={shared.opts.sdnq_quantize_weights_group_size} svd_rank={shared.opts.sdnq_svd_rank} svd_steps={shared.opts.sdnq_svd_steps} use_svd={shared.opts.sdnq_use_svd} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} dequantize_fp32={shared.opts.sdnq_dequantize_fp32} quantize_with_gpu={shared.opts.sdnq_quantize_with_gpu} quantization_device={quantization_device} return_device={return_device} device_map={shared.opts.device_map} offload_mode={shared.opts.diffusers_offload_mode} non_blocking={shared.opts.diffusers_offload_nonblocking} modules_to_not_convert={modules_to_not_convert} modules_dtype_dict={modules_dtype_dict}') if kwargs is None: return sdnq_config else: @@ -525,6 +526,7 @@ def sdnq_quantize_model(model, op=None, sd_model=None, do_gc: bool = True, weigh torch_dtype=devices.dtype, group_size=shared.opts.sdnq_quantize_weights_group_size, svd_rank=shared.opts.sdnq_svd_rank, + svd_steps=shared.opts.sdnq_svd_steps, use_svd=shared.opts.sdnq_use_svd, quant_conv=shared.opts.sdnq_quantize_conv_layers, use_quantized_matmul=shared.opts.sdnq_use_quantized_matmul, @@ -562,7 +564,7 @@ def sdnq_quantize_model(model, op=None, sd_model=None, do_gc: bool = True, weigh if do_gc: devices.torch_gc(force=True, reason='sdnq') - log.debug(f'Quantization: module="{op if op is not None else model.__class__}" type=sdnq mode=post dtype={weights_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} group_size={shared.opts.sdnq_quantize_weights_group_size} svd_rank={shared.opts.sdnq_svd_rank} use_svd={shared.opts.sdnq_use_svd} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} dequantize_fp32={shared.opts.sdnq_dequantize_fp32} quantize_with_gpu={shared.opts.sdnq_quantize_with_gpu} quantization_device={quantization_device} return_device={return_device} device_map={shared.opts.device_map} offload_mode={shared.opts.diffusers_offload_mode} non_blocking={shared.opts.diffusers_offload_nonblocking} modules_to_not_convert={modules_to_not_convert} modules_dtype_dict={modules_dtype_dict}') + log.debug(f'Quantization: module="{op if op is not None else model.__class__}" type=sdnq mode=post dtype={weights_dtype} matmul={shared.opts.sdnq_use_quantized_matmul} group_size={shared.opts.sdnq_quantize_weights_group_size} svd_rank={shared.opts.sdnq_svd_rank} svd_steps={shared.opts.sdnq_svd_steps} use_svd={shared.opts.sdnq_use_svd} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} dequantize_fp32={shared.opts.sdnq_dequantize_fp32} quantize_with_gpu={shared.opts.sdnq_quantize_with_gpu} quantization_device={quantization_device} return_device={return_device} device_map={shared.opts.device_map} offload_mode={shared.opts.diffusers_offload_mode} non_blocking={shared.opts.diffusers_offload_nonblocking} modules_to_not_convert={modules_to_not_convert} modules_dtype_dict={modules_dtype_dict}') return model @@ -570,7 +572,7 @@ def sdnq_quantize_weights(sd_model): try: t0 = time.time() from modules import shared, devices, sd_models - log.debug(f"Quantization: type=SDNQ modules={shared.opts.sdnq_quantize_weights} dtype={shared.opts.sdnq_quantize_weights_mode} dtype_te={shared.opts.sdnq_quantize_weights_mode_te} matmul={shared.opts.sdnq_use_quantized_matmul} svd_rank={shared.opts.sdnq_svd_rank} use_svd={shared.opts.sdnq_use_svd} group_size={shared.opts.sdnq_quantize_weights_group_size} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} quantize_with_gpu={shared.opts.sdnq_quantize_with_gpu} dequantize_fp32={shared.opts.sdnq_dequantize_fp32} pre_forward={shared.opts.diffusers_offload_pre}") + log.debug(f"Quantization: type=SDNQ modules={shared.opts.sdnq_quantize_weights} dtype={shared.opts.sdnq_quantize_weights_mode} dtype_te={shared.opts.sdnq_quantize_weights_mode_te} matmul={shared.opts.sdnq_use_quantized_matmul} svd_rank={shared.opts.sdnq_svd_rank} svd_steps={shared.opts.sdnq_svd_steps} use_svd={shared.opts.sdnq_use_svd} group_size={shared.opts.sdnq_quantize_weights_group_size} quant_conv={shared.opts.sdnq_quantize_conv_layers} matmul_conv={shared.opts.sdnq_use_quantized_matmul_conv} quantize_with_gpu={shared.opts.sdnq_quantize_with_gpu} dequantize_fp32={shared.opts.sdnq_dequantize_fp32} pre_forward={shared.opts.diffusers_offload_pre}") global quant_last_model_name, quant_last_model_device # pylint: disable=global-statement sd_model = sd_models.apply_function_to_model(sd_model, sdnq_quantize_model, shared.opts.sdnq_quantize_weights, op="sdnq") diff --git a/modules/sdnq/common.py b/modules/sdnq/common.py index 57a4f0aaf..698a35633 100644 --- a/modules/sdnq/common.py +++ b/modules/sdnq/common.py @@ -80,3 +80,19 @@ if use_torch_compile: else: def compile_func(fn, **kwargs): # pylint: disable=unused-argument return fn + + +module_skip_keys_dict = { + "FluxTransformer2DModel": [ + ["single_transformer_blocks.0.norm.linear.weight", ".time_text_embed", ".context_embedder", ".x_embedder", ".proj_out", ".norm_out", "pos_embed"], + {} + ], + "ChromaTransformer2DModel": [ + ["distilled_guidance_layer", ".time_text_embed", ".context_embedder", ".x_embedder", ".proj_out", ".norm_out", "pos_embed"], + {} + ], + "QwenImageTransformer2DModel": [ + ["transformer_blocks.0.img_mod.1.weight", ".time_text_embed", ".txt_in", ".img_in", ".proj_out", ".norm_out", "pos_embed"], + {} + ], +} diff --git a/modules/sdnq/loader.py b/modules/sdnq/loader.py index e2ecce092..30a23dc70 100644 --- a/modules/sdnq/loader.py +++ b/modules/sdnq/loader.py @@ -10,46 +10,39 @@ from .dequantizer import dequantize_symmetric, re_quantize_int8, re_quantize_fp8 def get_module_names(model: ModelMixin) -> list: modules_names = model._internal_dict.keys() # pylint: disable=protected-access - modules_names = [m for m in modules_names if not m.startswith('_')] + modules_names = [m for m in modules_names if not m.startswith("_")] modules_names = [m for m in modules_names if isinstance(getattr(model, m, None), torch.nn.Module)] modules_names = sorted(set(modules_names)) return modules_names +def unset_config_on_save(config: SDNQConfig) -> SDNQConfig: + config.quantization_config.quantization_device = None + config.quantization_config.return_device = None + config.quantization_config.non_blocking = False + config.quantization_config.add_skip_keys = False + return config + + def save_sdnq_model(model: ModelMixin, model_path: str, max_shard_size: str = "10GB", is_pipeline: bool = False, sdnq_config: SDNQConfig = None) -> None: if is_pipeline: for module_name in get_module_names(model): module = getattr(model, module_name, None) if hasattr(module, "config") and hasattr(module.config, "quantization_config") and isinstance(module.config.quantization_config, SDNQConfig): - module.config.quantization_config.quantization_device = None - module.config.quantization_config.return_device = None - module.config.quantization_config.non_blocking = False - module.config.quantization_config.add_skip_keys = False + module.config.quantization_config = unset_config_on_save(module.config.quantization_config) if hasattr(module, "quantization_config") and isinstance(module.quantization_config, SDNQConfig): - module.quantization_config.quantization_device = None - module.quantization_config.return_device = None - module.quantization_config.non_blocking = False - module.quantization_config.add_skip_keys = False + module.quantization_config = unset_config_on_save(module.quantization_config) else: if hasattr(model, "config") and hasattr(model.config, "quantization_config") and isinstance(model.config.quantization_config, SDNQConfig): - model.config.quantization_config.quantization_device = None - model.config.quantization_config.return_device = None - model.config.quantization_config.non_blocking = False - model.config.quantization_config.add_skip_keys = False + model.config.quantization_config = unset_config_on_save(model.config.quantization_config) if hasattr(model, "quantization_config") and isinstance(model.quantization_config, SDNQConfig): - model.quantization_config.quantization_device = None - model.quantization_config.return_device = None - model.quantization_config.non_blocking = False - model.quantization_config.add_skip_keys = False + model.quantization_config = unset_config_on_save(model.quantization_config) model.save_pretrained(model_path, max_shard_size=max_shard_size) # actual save quantization_config_path = os.path.join(model_path, "quantization_config.json") if sdnq_config is not None: # if provided, save global config - sdnq_config.quantization_device = None - sdnq_config.return_device = None - sdnq_config.non_blocking = False - sdnq_config.add_skip_keys = False + sdnq_config = unset_config_on_save(sdnq_config) sdnq_config.to_json_file(quantization_config_path) if is_pipeline: @@ -69,7 +62,7 @@ def save_sdnq_model(model: ModelMixin, model_path: str, max_shard_size: str = "1 model.config.quantization_config.to_json_file(quantization_config_path) -def load_sdnq_model(model_path: str, model_cls: ModelMixin = None, file_name: str = None, dtype: torch.dtype = None, device: torch.device = 'cpu', dequantize_fp32: bool = None, use_quantized_matmul: bool = None, model_config: dict = None, quantization_config: dict = None) -> ModelMixin: +def load_sdnq_model(model_path: str, model_cls: ModelMixin = None, file_name: str = None, dtype: torch.dtype = None, device: torch.device = "cpu", dequantize_fp32: bool = None, use_quantized_matmul: bool = None, model_config: dict = None, quantization_config: dict = None) -> ModelMixin: from accelerate import init_empty_weights from safetensors.torch import safe_open @@ -83,7 +76,7 @@ def load_sdnq_model(model_path: str, model_cls: ModelMixin = None, file_name: st if model_config is None: try: - with open(os.path.join(model_path, 'config.json'), "r", encoding="utf-8") as f: + with open(os.path.join(model_path, "config.json"), "r", encoding="utf-8") as f: model_config = json.load(f) except Exception: model_config = {} diff --git a/modules/sdnq/quantizer.py b/modules/sdnq/quantizer.py index 68d1c355a..b4e60a704 100644 --- a/modules/sdnq/quantizer.py +++ b/modules/sdnq/quantizer.py @@ -12,7 +12,7 @@ from diffusers.quantizers.quantization_config import QuantizationConfigMixin from diffusers.utils import get_module_from_name from modules import devices, shared -from .common import dtype_dict, accepted_weights, use_tensorwise_fp8_matmul, allowed_types, conv_types, conv_transpose_types, use_contiguous_mm +from .common import dtype_dict, module_skip_keys_dict, accepted_weights, use_tensorwise_fp8_matmul, allowed_types, conv_types, conv_transpose_types, use_contiguous_mm from .dequantizer import dequantizer_dict, dequantize_sdnq_model from .forward import get_forward_func @@ -49,13 +49,13 @@ def quantize_weight(weight: torch.FloatTensor, reduction_axes: Union[int, List[i return quantized_weight, scale, zero_point -def apply_svdquant(weight: torch.FloatTensor, rank: int = 32) -> Tuple[torch.FloatTensor, torch.FloatTensor, torch.FloatTensor]: +def apply_svdquant(weight: torch.FloatTensor, rank: int = 32, niter: int = 8) -> 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) - U, S, svd_down = torch.svd_lowrank(weight, q=rank) + 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_() weight = weight.sub_(torch.mm(svd_up, svd_down)) @@ -64,8 +64,79 @@ def apply_svdquant(weight: torch.FloatTensor, rank: int = 32) -> Tuple[torch.Flo return weight, svd_up, svd_down +def check_param_name_in(param_name: str, param_list: List[str]) -> bool: + split_param_name = param_name.split(".") + for param in param_list: + if param.startswith("."): + if param_name.startswith(param[1:]): + return True + else: + continue + if ( + param_name == param + or param in split_param_name + or ("*" in param and re.match(param.replace(".*", "\\.*").replace("*", ".*"), param_name)) + ): + return True + return False + + +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): + key = key.lower() + if key in {"8bit", "8bits"}: + if dtype_dict[weights_dtype]["num_bits"] != 8: + return "int8" + elif key.startswith("minimum_"): + minimum_bits_str = key.removeprefix("minimum_").removesuffix("bits").removesuffix("bit") + 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 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 = [] + 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 + + 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 + elif 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(layer, weights_dtype="int8", torch_dtype=None, group_size=0, svd_rank=32, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, dequantize_fp32=False, non_blocking=False, quantization_device=None, return_device=None, param_name=None): # pylint: disable=unused-argument +def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, dequantize_fp32=False, non_blocking=False, quantization_device=None, return_device=None, param_name=None): # pylint: disable=unused-argument layer_class_name = layer.__class__.__name__ if layer_class_name in allowed_types: num_of_groups = 1 @@ -128,7 +199,7 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz layer.weight.data = layer.weight.to(dtype=torch.float32) if use_svd: - layer.weight.data, svd_up, svd_down = apply_svdquant(layer.weight, rank=svd_rank) + layer.weight.data, svd_up, svd_down = apply_svdquant(layer.weight, rank=svd_rank, niter=svd_steps) if use_quantized_matmul: svd_up = svd_up.t_() svd_down = svd_down.t_() @@ -237,7 +308,7 @@ def sdnq_quantize_layer(layer, weights_dtype="int8", torch_dtype=None, group_siz return layer -def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_size=0, svd_rank=32, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, dequantize_fp32=False, non_blocking=False, quantization_device=None, return_device=None, modules_to_not_convert: List[str] = None, modules_dtype_dict: Dict[str, List[str]] = None, full_param_name="", op=None): # pylint: disable=unused-argument +def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_size=0, svd_rank=32, svd_steps=8, use_svd=False, quant_conv=False, use_quantized_matmul=False, use_quantized_matmul_conv=False, dequantize_fp32=False, non_blocking=False, quantization_device=None, return_device=None, modules_to_not_convert: List[str] = None, modules_dtype_dict: Dict[str, List[str]] = None, full_param_name="", op=None): # pylint: disable=unused-argument has_children = list(model.children()) if not has_children: return model @@ -252,12 +323,7 @@ def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_si param_name = full_param_name + "." + param_name if hasattr(module, "weight") and module.weight is not None: param_name = param_name + ".weight" - split_param_name = param_name.split(".") - if ( - param_name in modules_to_not_convert - or any(param in split_param_name for param in modules_to_not_convert) - or any("*" in param and re.match(param.replace(".*", "\\.*").replace("*", ".*"), param_name) for param in modules_to_not_convert) - ): + if check_param_name_in(param_name, modules_to_not_convert): continue layer_class_name = module.__class__.__name__ if layer_class_name in allowed_types: @@ -265,33 +331,15 @@ def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_si continue else: continue - if len(modules_dtype_dict.keys()) > 0: - for key, value in modules_dtype_dict.items(): - if ( - param_name in value - or any(param in split_param_name for param in value) - or any("*" in param and re.match(param.replace(".*", "\\.*").replace("*", ".*"), param_name) for param in value) - ): - key = key.lower() - if key in {"8bit", "8bits"}: - if dtype_dict[weights_dtype]["num_bits"] != 8: - weights_dtype = "int8" - elif key.startswith("minimum_"): - minimum_bits_str = key.removeprefix("minimum_").removesuffix("bits").removesuffix("bit") - minimum_bits = int(minimum_bits_str) - if dtype_dict[weights_dtype]["num_bits"] < minimum_bits: - weights_dtype = "int" + minimum_bits_str - if minimum_bits <= 4: - weights_dtype = "u" + weights_dtype - else: - weights_dtype = key + weights_dtype = get_minimum_dtype(weights_dtype, param_name, modules_dtype_dict) module = sdnq_quantize_layer( module, weights_dtype=weights_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, use_quantized_matmul=use_quantized_matmul, @@ -308,6 +356,7 @@ def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_si torch_dtype=torch_dtype, group_size=group_size, svd_rank=svd_rank, + svd_steps=svd_steps, use_svd=use_svd, quant_conv=quant_conv, use_quantized_matmul=use_quantized_matmul, @@ -324,32 +373,13 @@ def apply_sdnq_to_module(model, weights_dtype="int8", torch_dtype=None, group_si return model -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 = [] - 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, "_skip_layerwise_casting_patterns", None) is not None: - modules_to_not_convert.extend(model._skip_layerwise_casting_patterns) # pylint: disable=protected-access - if model.__class__.__name__ == "ChromaTransformer2DModel": - modules_to_not_convert.append("distilled_guidance_layer") - elif model.__class__.__name__ == "QwenImageTransformer2DModel": - modules_to_not_convert.extend(["transformer_blocks.0.img_mod.1.weight", "time_text_embed", "img_in", "txt_in", "proj_out", "norm_out", "pos_embed"]) - if "minimum_6bit" not in modules_dtype_dict.keys(): - modules_dtype_dict["minimum_6bit"] = ["img_mod"] - else: - modules_dtype_dict["minimum_6bit"].append("img_mod") - return model, modules_to_not_convert, modules_dtype_dict - - def sdnq_post_load_quant( model, weights_dtype="int8", torch_dtype: torch.dtype = None, group_size: int = 0, svd_rank: int = 32, + svd_steps: int = 8, use_svd: bool = False, quant_conv: bool = False, use_quantized_matmul: bool = False, @@ -373,6 +403,7 @@ def sdnq_post_load_quant( torch_dtype=torch_dtype, group_size=group_size, svd_rank=svd_rank, + svd_steps=svd_steps, use_svd=use_svd, quant_conv=quant_conv, use_quantized_matmul=use_quantized_matmul, @@ -389,6 +420,7 @@ def sdnq_post_load_quant( weights_dtype=weights_dtype, group_size=group_size, svd_rank=svd_rank, + svd_steps=svd_steps, use_svd=use_svd, quant_conv=quant_conv, use_quantized_matmul=use_quantized_matmul, @@ -433,12 +465,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): if hasattr(layer, "sdnq_dequantizer"): return True elif param_name.endswith(".weight"): - split_param_name = param_name.split(".") - if ( - param_name not in self.quantization_config.modules_to_not_convert - and not any(param in split_param_name for param in self.quantization_config.modules_to_not_convert) - and not any("*" in param and re.match(param.replace(".*", "\\.*").replace("*", ".*"), param_name) for param in self.quantization_config.modules_to_not_convert) - ): + if not check_param_name_in(param_name, self.quantization_config.modules_to_not_convert): layer_class_name = get_module_from_name(model, param_name)[0].__class__.__name__ if layer_class_name in allowed_types: if layer_class_name in conv_types or layer_class_name in conv_transpose_types: @@ -485,28 +512,8 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): setattr(layer, tensor_name, param_value) return - weights_dtype = self.quantization_config.weights_dtype - if len(self.quantization_config.modules_dtype_dict.keys()) > 0: - split_param_name = param_name.split(".") - for key, value in self.quantization_config.modules_dtype_dict.items(): - if ( - param_name in value - or any(param in split_param_name for param in value) - or any("*" in param and re.match(param.replace(".*", "\\.*").replace("*", ".*"), param_name) for param in value) - ): - key = key.lower() - if key in {"8bit", "8bits"}: - if dtype_dict[weights_dtype]["num_bits"] != 8: - weights_dtype = "int8" - elif key.startswith("minimum_"): - minimum_bits_str = key.removeprefix("minimum_").removesuffix("bits").removesuffix("bit") - minimum_bits = int(minimum_bits_str) - if dtype_dict[weights_dtype]["num_bits"] < minimum_bits: - weights_dtype = "int" + minimum_bits_str - if minimum_bits <= 4: - weights_dtype = "u" + weights_dtype - else: - weights_dtype = key + torch_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 @@ -516,7 +523,6 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): if self.quantization_config.quantization_device is not None: target_device = self.quantization_config.quantization_device - torch_dtype = param_value.dtype if self.torch_dtype is None else self.torch_dtype if param_value.dtype == torch.float32 and devices.same_device(param_value.device, target_device): param_value = param_value.clone() else: @@ -530,6 +536,7 @@ class SDNQQuantizer(DiffusersQuantizer, HfQuantizer): 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, use_svd=self.quantization_config.use_svd, quant_conv=self.quantization_config.quant_conv, use_quantized_matmul=self.quantization_config.use_quantized_matmul, @@ -638,6 +645,8 @@ class SDNQConfig(QuantizationConfigMixin): group_size = 0 will automatically select a group size based on weights_dtype. svd_rank (`int`, *optional*, defaults to `32`): The rank size used for the SVDQuant algorithm. + svd_steps (`int`, *optional*, defaults to `8`): + The number of iterations to use in svd lowrank estimation. use_svd (`bool`, *optional*, defaults to `False`): Enabling this option will use SVDQuant algorithm on top of SDNQ quantization. quant_conv (`bool`, *optional*, defaults to `False`): @@ -668,6 +677,7 @@ class SDNQConfig(QuantizationConfigMixin): weights_dtype: str = "int8", group_size: int = 0, svd_rank: int = 32, + svd_steps: int = 8, use_svd: bool = False, quant_conv: bool = False, use_quantized_matmul: bool = False, @@ -685,6 +695,7 @@ class SDNQConfig(QuantizationConfigMixin): self.quant_method = QuantizationMethod.SDNQ self.group_size = group_size self.svd_rank = svd_rank + self.svd_steps = svd_steps self.use_svd = use_svd self.quant_conv = quant_conv self.use_quantized_matmul = use_quantized_matmul diff --git a/modules/shared.py b/modules/shared.py index f987cdeb8..0c3315f26 100644 --- a/modules/shared.py +++ b/modules/shared.py @@ -202,6 +202,7 @@ options_templates.update(options_section(("quantization", "Model Quantization"), "sdnq_modules_dtype_dict": OptionInfo("{}", "Modules dtype dict"), "sdnq_quantize_weights_group_size": OptionInfo(0, "Group size", gr.Slider, {"minimum": -1, "maximum": 4096, "step": 1}), "sdnq_svd_rank": OptionInfo(32, "SVDQuant Rank size", gr.Slider, {"minimum": 1, "maximum": 512, "step": 1}), + "sdnq_svd_steps": OptionInfo(8, "SVDQuant quantization steps", gr.Slider, {"minimum": 1, "maximum": 128, "step": 1}), "sdnq_use_svd": OptionInfo(False, "Use SVDQuant quantization", gr.Checkbox), "sdnq_quantize_conv_layers": OptionInfo(False, "Quantize convolutional layers", gr.Checkbox), "sdnq_dequantize_compile": OptionInfo(devices.has_triton(), "Dequantize using torch.compile", gr.Checkbox), diff --git a/pipelines/model_qwen.py b/pipelines/model_qwen.py index d9a98d9be..08faf5624 100644 --- a/pipelines/model_qwen.py +++ b/pipelines/model_qwen.py @@ -49,8 +49,7 @@ def load_qwen(checkpoint_info, diffusers_load_config={}): subfolder=transformer_subfolder, cls_name=diffusers.QwenImageTransformer2DModel, load_config=diffusers_load_config, - modules_dtype_dict={"minimum_6bit": ["img_mod"]}, - modules_to_not_convert=["transformer_blocks.0.img_mod.1.weight", "time_text_embed", "img_in", "txt_in", "proj_out", "norm_out", "pos_embed"], + modules_to_not_convert=["transformer_blocks.0.img_mod.1.weight"], ) repo_te = 'Qwen/Qwen-Image'