mirror of
https://github.com/vladmandic/automatic
synced 2026-09-10 14:58:44 +02:00
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:
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user