mirror of
https://github.com/vladmandic/automatic
synced 2026-09-18 16:54:33 +02:00
SDNQ add common keys
This commit is contained in:
+29
-2
@@ -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"]
|
||||
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user