SDNQ add common keys

This commit is contained in:
Disty0
2025-10-31 00:21:54 +03:00
parent da3d183f96
commit 76d699dc09
2 changed files with 36 additions and 5 deletions
+29 -2
View File
@@ -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"]
+7 -3
View File
@@ -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))