From 76d699dc09c4f9996d50491ccba8e9a3ef110d1e Mon Sep 17 00:00:00 2001 From: Disty0 Date: Fri, 31 Oct 2025 00:21:54 +0300 Subject: [PATCH] SDNQ add common keys --- modules/sdnq/common.py | 31 +++++++++++++++++++++++++++++-- modules/sdnq/quantizer.py | 10 +++++++--- 2 files changed, 36 insertions(+), 5 deletions(-) diff --git a/modules/sdnq/common.py b/modules/sdnq/common.py index 2f8904fa0..e83af542b 100644 --- a/modules/sdnq/common.py +++ b/modules/sdnq/common.py @@ -82,6 +82,28 @@ else: return fn +common_skip_keys = ( + ".lm_head", + ".patch_embedding", + ".time_embed", + ".time_text_embed", + ".context_embedder", + ".condition_embedder", + ".x_embedder", + ".emb_in", + ".txt_in", + ".img_in", + ".vid_in", + ".proj_out", + ".norm_out", + ".emb_out", + ".txt_out", + ".img_out", + ".vid_out", + ".final_layer", + "pos_embed", +) + 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"], @@ -99,12 +121,16 @@ module_skip_keys_dict = { ["scale_shift_table", ".rope", ".patch_embedding", ".condition_embedder", ".proj_out", ".norm_out", "pos_embed"], {} ], + "Gemma3nForCausalLM": [ + [".lm_head", "correction_coefs", "prediction_coefs", "embedding_projection"], + {} + ], "HunyuanImage3ForCausalMM": [ - [".patch_embed", ".time_embed", ".time_embed_2", ".final_layer", ".model.wte", ".model.ln_f", ".timestep_emb", ".vae", ".vision_aligner", ".vision_model.head", ".vision_model.post_layernorm", ".vision_model.embeddings", ".lm_head"], + [".lm_head", ".patch_embed", ".time_embed", ".time_embed_2", ".final_layer", ".model.wte", ".model.ln_f", ".timestep_emb", ".vae", ".vision_aligner", ".vision_model.head", ".vision_model.post_layernorm", ".vision_model.embeddings"], {} ], "Emu3ForCausalLM": [ - [".lm_head", ".vq_model", ".tokenizer", ".model.embed_tokens", ".model.norm"], + [".lm_head", ".vq_model", ".tokenizer"], {} ], "NaDiT": [ @@ -114,4 +140,5 @@ module_skip_keys_dict = { } module_skip_keys_dict["ChronoEditTransformer3DModel"] = module_skip_keys_dict["WanTransformer3DModel"] +module_skip_keys_dict["Gemma3nForConditionalGeneration"] = module_skip_keys_dict["Gemma3nForCausalLM"] module_skip_keys_dict["NaDiTUpscaler"] = module_skip_keys_dict["NaDiT"] diff --git a/modules/sdnq/quantizer.py b/modules/sdnq/quantizer.py index dec88e3a9..14ecfcc92 100644 --- a/modules/sdnq/quantizer.py +++ b/modules/sdnq/quantizer.py @@ -16,7 +16,7 @@ from accelerate import init_empty_weights from accelerate.utils import set_module_tensor_to_device from modules import devices, shared -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 .common import dtype_dict, common_skip_keys, 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 @@ -119,6 +119,8 @@ def add_module_skip_keys(model, modules_to_not_convert: List[str] = None, module 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: + 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: @@ -128,8 +130,10 @@ def add_module_skip_keys(model, modules_to_not_convert: List[str] = None, module 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 + 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))