SDNQ add "*" support and upcast only the first and last layer's img_mod to 6 bit with Qwen Image

This commit is contained in:
Disty0
2025-08-20 03:24:19 +03:00
parent a78d51ea07
commit 47ff01fd3b
2 changed files with 12 additions and 3 deletions
+11 -2
View File
@@ -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:
+1 -1
View File
@@ -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)