diff --git a/modules/sdnq/__init__.py b/modules/sdnq/__init__.py index bca210c1b..e88bfc8d2 100644 --- a/modules/sdnq/__init__.py +++ b/modules/sdnq/__init__.py @@ -4,6 +4,7 @@ from typing import Any, Dict, List, Tuple, Optional, Union from dataclasses import dataclass from enum import Enum +import re import torch from diffusers.quantizers.base import DiffusersQuantizer from diffusers.quantizers.quantization_config import QuantizationConfigMixin @@ -270,7 +271,11 @@ class SDNQQuantizer(DiffusersQuantizer): ): if param_name.endswith(".weight"): split_param_name = param_name.split(".") - if param_name not in self.modules_to_not_convert and not any(param in split_param_name for param in self.modules_to_not_convert): + if ( + param_name not in self.modules_to_not_convert + and not any(param in split_param_name for param in self.modules_to_not_convert) + and not any("*" in param and re.match(param.replace("*", ".*"), param_name) for param in self.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: @@ -303,7 +308,11 @@ class SDNQQuantizer(DiffusersQuantizer): 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): + 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("*", ".*"), param_name) for param in value) + ): key = key.lower() if key in {"8bit", "8bits"}: if dtype_dict[weights_dtype]["num_bits"] != 8: diff --git a/pipelines/model_qwen.py b/pipelines/model_qwen.py index 9c7984dd5..b6e57c95f 100644 --- a/pipelines/model_qwen.py +++ b/pipelines/model_qwen.py @@ -30,7 +30,7 @@ def load_qwen(checkpoint_info, diffusers_load_config={}): # cls_name = nunchaku.pipeline.pipeline_qwenimage.NunchakuQwenImagePipeline # we dont need this if transformer is None: - transformer = generic.load_transformer(repo_id, cls_name=diffusers.QwenImageTransformer2DModel, load_config=diffusers_load_config, modules_dtype_dict={"minimum_6bit": ["img_mod", "pos_embed", "time_text_embed", "img_in", "txt_in", "norm_out"]}) + transformer = generic.load_transformer(repo_id, cls_name=diffusers.QwenImageTransformer2DModel, load_config=diffusers_load_config, modules_dtype_dict={"minimum_6bit": ["pos_embed", "time_text_embed", "img_in", "txt_in", "norm_out", "transformer_blocks.0.img_mod.1.weight", "transformer_blocks.59.img_mod.1.weight"]}) repo_te = 'Qwen/Qwen-Image' # if 'Qwen-Lightning' in repo_id or 'Qwen-Image-Edit' in repo_id else repo_id text_encoder = generic.load_text_encoder(repo_te, cls_name=transformers.Qwen2_5_VLForConditionalGeneration, load_config=diffusers_load_config)