Support additional Qwen LoRA naming conventions (#270)

This commit is contained in:
Filip Strand 2025-10-11 16:00:25 +02:00 committed by GitHub
parent 98e03a46f2
commit 1278bafc53
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 75 additions and 13 deletions

View File

@ -1,3 +1,4 @@
import re
from pathlib import Path from pathlib import Path
from typing import Dict, Tuple from typing import Dict, Tuple
@ -83,19 +84,20 @@ class LoRALoader:
# Pattern matching logic # Pattern matching logic
for pattern, mapping_info in lora_mappings.items(): for pattern, mapping_info in lora_mappings.items():
if "{block}" in pattern: if "{block}" in pattern:
# Extract block number from the weight key # Extract block number from the weight key - try both . and _ separators
parts = weight_key.split(".") # This handles both standard LoRA formats (dot-separated) and other formats (underscore-separated)
for i, part in enumerate(parts): # Find all numbers in the weight key
if part.isdigit(): numbers_in_key = re.findall(r'\d+', weight_key)
try: for num_str in numbers_in_key:
test_block_idx = int(part) try:
concrete_pattern = pattern.format(block=test_block_idx) test_block_idx = int(num_str)
if weight_key == concrete_pattern: concrete_pattern = pattern.format(block=test_block_idx)
found_mapping = mapping_info if weight_key == concrete_pattern:
block_idx = test_block_idx found_mapping = mapping_info
break block_idx = test_block_idx
except (ValueError, KeyError): break
continue except (ValueError, KeyError):
continue
if found_mapping: if found_mapping:
break break
else: else:

View File

@ -13,13 +13,18 @@ class QwenLoRAMapping(LoRAMapping):
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.attn.to_q.lora_up.weight", "transformer_blocks.{block}.attn.to_q.lora_up.weight",
"transformer.transformer_blocks.{block}.attn.to_q.lora.up.weight", "transformer.transformer_blocks.{block}.attn.to_q.lora.up.weight",
"diffusion_model.transformer_blocks.{block}.attn.to_q.lora_B.weight",
"lora_unet_transformer_blocks_{block}_attn_to_q.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.attn.to_q.lora_down.weight", "transformer_blocks.{block}.attn.to_q.lora_down.weight",
"transformer.transformer_blocks.{block}.attn.to_q.lora.down.weight", "transformer.transformer_blocks.{block}.attn.to_q.lora.down.weight",
"diffusion_model.transformer_blocks.{block}.attn.to_q.lora_A.weight",
"lora_unet_transformer_blocks_{block}_attn_to_q.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.attn.to_q.alpha", "transformer_blocks.{block}.attn.to_q.alpha",
"lora_unet_transformer_blocks_{block}_attn_to_q.alpha",
] ]
), ),
LoRATarget( LoRATarget(
@ -27,13 +32,18 @@ class QwenLoRAMapping(LoRAMapping):
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.attn.to_k.lora_up.weight", "transformer_blocks.{block}.attn.to_k.lora_up.weight",
"transformer.transformer_blocks.{block}.attn.to_k.lora.up.weight", "transformer.transformer_blocks.{block}.attn.to_k.lora.up.weight",
"diffusion_model.transformer_blocks.{block}.attn.to_k.lora_B.weight",
"lora_unet_transformer_blocks_{block}_attn_to_k.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.attn.to_k.lora_down.weight", "transformer_blocks.{block}.attn.to_k.lora_down.weight",
"transformer.transformer_blocks.{block}.attn.to_k.lora.down.weight", "transformer.transformer_blocks.{block}.attn.to_k.lora.down.weight",
"diffusion_model.transformer_blocks.{block}.attn.to_k.lora_A.weight",
"lora_unet_transformer_blocks_{block}_attn_to_k.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.attn.to_k.alpha", "transformer_blocks.{block}.attn.to_k.alpha",
"lora_unet_transformer_blocks_{block}_attn_to_k.alpha",
] ]
), ),
LoRATarget( LoRATarget(
@ -41,13 +51,18 @@ class QwenLoRAMapping(LoRAMapping):
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.attn.to_v.lora_up.weight", "transformer_blocks.{block}.attn.to_v.lora_up.weight",
"transformer.transformer_blocks.{block}.attn.to_v.lora.up.weight", "transformer.transformer_blocks.{block}.attn.to_v.lora.up.weight",
"diffusion_model.transformer_blocks.{block}.attn.to_v.lora_B.weight",
"lora_unet_transformer_blocks_{block}_attn_to_v.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.attn.to_v.lora_down.weight", "transformer_blocks.{block}.attn.to_v.lora_down.weight",
"transformer.transformer_blocks.{block}.attn.to_v.lora.down.weight", "transformer.transformer_blocks.{block}.attn.to_v.lora.down.weight",
"diffusion_model.transformer_blocks.{block}.attn.to_v.lora_A.weight",
"lora_unet_transformer_blocks_{block}_attn_to_v.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.attn.to_v.alpha", "transformer_blocks.{block}.attn.to_v.alpha",
"lora_unet_transformer_blocks_{block}_attn_to_v.alpha",
] ]
), ),
LoRATarget( LoRATarget(
@ -55,109 +70,154 @@ class QwenLoRAMapping(LoRAMapping):
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.attn.to_out.0.lora_up.weight", "transformer_blocks.{block}.attn.to_out.0.lora_up.weight",
"transformer.transformer_blocks.{block}.attn.to_out.0.lora.up.weight", "transformer.transformer_blocks.{block}.attn.to_out.0.lora.up.weight",
"diffusion_model.transformer_blocks.{block}.attn.to_out.0.lora_B.weight",
"lora_unet_transformer_blocks_{block}_attn_to_out_0.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.attn.to_out.0.lora_down.weight", "transformer_blocks.{block}.attn.to_out.0.lora_down.weight",
"transformer.transformer_blocks.{block}.attn.to_out.0.lora.down.weight", "transformer.transformer_blocks.{block}.attn.to_out.0.lora.down.weight",
"diffusion_model.transformer_blocks.{block}.attn.to_out.0.lora_A.weight",
"lora_unet_transformer_blocks_{block}_attn_to_out_0.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.attn.to_out.0.alpha", "transformer_blocks.{block}.attn.to_out.0.alpha",
"lora_unet_transformer_blocks_{block}_attn_to_out_0.alpha",
] ]
), ),
LoRATarget( LoRATarget(
model_path="transformer_blocks.{block}.attn.add_q_proj", model_path="transformer_blocks.{block}.attn.add_q_proj",
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.attn.add_q_proj.lora_up.weight", "transformer_blocks.{block}.attn.add_q_proj.lora_up.weight",
"diffusion_model.transformer_blocks.{block}.attn.add_q_proj.lora_B.weight",
"lora_unet_transformer_blocks_{block}_attn_add_q_proj.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.attn.add_q_proj.lora_down.weight", "transformer_blocks.{block}.attn.add_q_proj.lora_down.weight",
"diffusion_model.transformer_blocks.{block}.attn.add_q_proj.lora_A.weight",
"lora_unet_transformer_blocks_{block}_attn_add_q_proj.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.attn.add_q_proj.alpha", "transformer_blocks.{block}.attn.add_q_proj.alpha",
"lora_unet_transformer_blocks_{block}_attn_add_q_proj.alpha",
] ]
), ),
LoRATarget( LoRATarget(
model_path="transformer_blocks.{block}.attn.add_k_proj", model_path="transformer_blocks.{block}.attn.add_k_proj",
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.attn.add_k_proj.lora_up.weight", "transformer_blocks.{block}.attn.add_k_proj.lora_up.weight",
"diffusion_model.transformer_blocks.{block}.attn.add_k_proj.lora_B.weight",
"lora_unet_transformer_blocks_{block}_attn_add_k_proj.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.attn.add_k_proj.lora_down.weight", "transformer_blocks.{block}.attn.add_k_proj.lora_down.weight",
"diffusion_model.transformer_blocks.{block}.attn.add_k_proj.lora_A.weight",
"lora_unet_transformer_blocks_{block}_attn_add_k_proj.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.attn.add_k_proj.alpha", "transformer_blocks.{block}.attn.add_k_proj.alpha",
"lora_unet_transformer_blocks_{block}_attn_add_k_proj.alpha",
] ]
), ),
LoRATarget( LoRATarget(
model_path="transformer_blocks.{block}.attn.add_v_proj", model_path="transformer_blocks.{block}.attn.add_v_proj",
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.attn.add_v_proj.lora_up.weight", "transformer_blocks.{block}.attn.add_v_proj.lora_up.weight",
"diffusion_model.transformer_blocks.{block}.attn.add_v_proj.lora_B.weight",
"lora_unet_transformer_blocks_{block}_attn_add_v_proj.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.attn.add_v_proj.lora_down.weight", "transformer_blocks.{block}.attn.add_v_proj.lora_down.weight",
"diffusion_model.transformer_blocks.{block}.attn.add_v_proj.lora_A.weight",
"lora_unet_transformer_blocks_{block}_attn_add_v_proj.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.attn.add_v_proj.alpha", "transformer_blocks.{block}.attn.add_v_proj.alpha",
"lora_unet_transformer_blocks_{block}_attn_add_v_proj.alpha",
] ]
), ),
LoRATarget( LoRATarget(
model_path="transformer_blocks.{block}.attn.to_add_out", model_path="transformer_blocks.{block}.attn.to_add_out",
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.attn.to_add_out.lora_up.weight", "transformer_blocks.{block}.attn.to_add_out.lora_up.weight",
"diffusion_model.transformer_blocks.{block}.attn.to_add_out.lora_B.weight",
"lora_unet_transformer_blocks_{block}_attn_to_add_out.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.attn.to_add_out.lora_down.weight", "transformer_blocks.{block}.attn.to_add_out.lora_down.weight",
"diffusion_model.transformer_blocks.{block}.attn.to_add_out.lora_A.weight",
"lora_unet_transformer_blocks_{block}_attn_to_add_out.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.attn.to_add_out.alpha", "transformer_blocks.{block}.attn.to_add_out.alpha",
"lora_unet_transformer_blocks_{block}_attn_to_add_out.alpha",
] ]
), ),
LoRATarget( LoRATarget(
model_path="transformer_blocks.{block}.img_ff.mlp_in", model_path="transformer_blocks.{block}.img_ff.mlp_in",
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.img_mlp.net.0.proj.lora_up.weight", "transformer_blocks.{block}.img_mlp.net.0.proj.lora_up.weight",
"diffusion_model.transformer_blocks.{block}.img_mlp.net.0.proj.lora_B.weight",
"lora_unet_transformer_blocks_{block}_img_mlp_net_0_proj.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.img_mlp.net.0.proj.lora_down.weight", "transformer_blocks.{block}.img_mlp.net.0.proj.lora_down.weight",
"diffusion_model.transformer_blocks.{block}.img_mlp.net.0.proj.lora_A.weight",
"lora_unet_transformer_blocks_{block}_img_mlp_net_0_proj.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.img_mlp.net.0.proj.alpha", "transformer_blocks.{block}.img_mlp.net.0.proj.alpha",
"lora_unet_transformer_blocks_{block}_img_mlp_net_0_proj.alpha",
] ]
), ),
LoRATarget( LoRATarget(
model_path="transformer_blocks.{block}.img_ff.mlp_out", model_path="transformer_blocks.{block}.img_ff.mlp_out",
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.img_mlp.net.2.lora_up.weight", "transformer_blocks.{block}.img_mlp.net.2.lora_up.weight",
"diffusion_model.transformer_blocks.{block}.img_mlp.net.2.lora_B.weight",
"lora_unet_transformer_blocks_{block}_img_mlp_net_2.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.img_mlp.net.2.lora_down.weight", "transformer_blocks.{block}.img_mlp.net.2.lora_down.weight",
"diffusion_model.transformer_blocks.{block}.img_mlp.net.2.lora_A.weight",
"lora_unet_transformer_blocks_{block}_img_mlp_net_2.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.img_mlp.net.2.alpha", "transformer_blocks.{block}.img_mlp.net.2.alpha",
"lora_unet_transformer_blocks_{block}_img_mlp_net_2.alpha",
] ]
), ),
LoRATarget( LoRATarget(
model_path="transformer_blocks.{block}.txt_ff.mlp_in", model_path="transformer_blocks.{block}.txt_ff.mlp_in",
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.txt_mlp.net.0.proj.lora_up.weight", "transformer_blocks.{block}.txt_mlp.net.0.proj.lora_up.weight",
"diffusion_model.transformer_blocks.{block}.txt_mlp.net.0.proj.lora_B.weight",
"lora_unet_transformer_blocks_{block}_txt_mlp_net_0_proj.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.txt_mlp.net.0.proj.lora_down.weight", "transformer_blocks.{block}.txt_mlp.net.0.proj.lora_down.weight",
"diffusion_model.transformer_blocks.{block}.txt_mlp.net.0.proj.lora_A.weight",
"lora_unet_transformer_blocks_{block}_txt_mlp_net_0_proj.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.txt_mlp.net.0.proj.alpha", "transformer_blocks.{block}.txt_mlp.net.0.proj.alpha",
"lora_unet_transformer_blocks_{block}_txt_mlp_net_0_proj.alpha",
] ]
), ),
LoRATarget( LoRATarget(
model_path="transformer_blocks.{block}.txt_ff.mlp_out", model_path="transformer_blocks.{block}.txt_ff.mlp_out",
possible_up_patterns=[ possible_up_patterns=[
"transformer_blocks.{block}.txt_mlp.net.2.lora_up.weight", "transformer_blocks.{block}.txt_mlp.net.2.lora_up.weight",
"diffusion_model.transformer_blocks.{block}.txt_mlp.net.2.lora_B.weight",
"lora_unet_transformer_blocks_{block}_txt_mlp_net_2.lora_up.weight",
], ],
possible_down_patterns=[ possible_down_patterns=[
"transformer_blocks.{block}.txt_mlp.net.2.lora_down.weight", "transformer_blocks.{block}.txt_mlp.net.2.lora_down.weight",
"diffusion_model.transformer_blocks.{block}.txt_mlp.net.2.lora_A.weight",
"lora_unet_transformer_blocks_{block}_txt_mlp_net_2.lora_down.weight",
], ],
possible_alpha_patterns=[ possible_alpha_patterns=[
"transformer_blocks.{block}.txt_mlp.net.2.alpha", "transformer_blocks.{block}.txt_mlp.net.2.alpha",
"lora_unet_transformer_blocks_{block}_txt_mlp_net_2.alpha",
] ]
), ),
] ]