Allow loading Fun Reward LoRAs

https://huggingface.co/alibaba-pai/Wan2.1-Fun-Reward-LoRAs/tree/main
This commit is contained in:
kijai
2025-04-07 12:38:56 +03:00
parent dce705f5d7
commit f8155ab47f
+116 -1
View File
@@ -205,8 +205,123 @@ def standardize_lora_key_format(lora_sd):
# Diffusers format
if k.startswith('transformer.'):
k = k.replace('transformer.', 'diffusion_model.')
# Fun LoRA format
if k.startswith('lora_unet__'):
# Split into main path and weight type parts
parts = k.split('.')
main_part = parts[0] # e.g. lora_unet__blocks_0_cross_attn_k
weight_type = '.'.join(parts[1:]) if len(parts) > 1 else None # e.g. lora_down.weight
# from finetrainer format
# Process the main part - convert from underscore to dot format
if 'blocks_' in main_part:
# Extract components
components = main_part[len('lora_unet__'):].split('_')
# Start with diffusion_model
new_key = "diffusion_model"
# Add blocks.N
if components[0] == 'blocks':
new_key += f".blocks.{components[1]}"
# Handle different module types
idx = 2
if idx < len(components):
if components[idx] == 'self' and idx+1 < len(components) and components[idx+1] == 'attn':
new_key += ".self_attn"
idx += 2
elif components[idx] == 'cross' and idx+1 < len(components) and components[idx+1] == 'attn':
new_key += ".cross_attn"
idx += 2
elif components[idx] == 'ffn':
new_key += ".ffn"
idx += 1
# Add the component (k, q, v, o) and handle img suffix
if idx < len(components):
component = components[idx]
idx += 1
# Check for img suffix
if idx < len(components) and components[idx] == 'img':
component += '_img'
idx += 1
new_key += f".{component}"
# Handle weight type - this is the critical fix
if weight_type:
if weight_type == 'alpha':
new_key += '.alpha'
elif weight_type == 'lora_down.weight' or weight_type == 'lora_down':
new_key += '.lora_A.weight'
elif weight_type == 'lora_up.weight' or weight_type == 'lora_up':
new_key += '.lora_B.weight'
else:
# Keep original weight type if not matching our patterns
new_key += f'.{weight_type}'
# Add .weight suffix if missing
if not new_key.endswith('.weight'):
new_key += '.weight'
k = new_key
else:
# For other lora_unet__ formats (head, embeddings, etc.)
new_key = main_part.replace('lora_unet__', 'diffusion_model.')
# Fix specific component naming patterns
new_key = new_key.replace('_self_attn', '.self_attn')
new_key = new_key.replace('_cross_attn', '.cross_attn')
new_key = new_key.replace('_ffn', '.ffn')
new_key = new_key.replace('blocks_', 'blocks.')
new_key = new_key.replace('head_head', 'head.head')
new_key = new_key.replace('img_emb', 'img_emb')
new_key = new_key.replace('text_embedding', 'text.embedding')
new_key = new_key.replace('time_embedding', 'time.embedding')
new_key = new_key.replace('time_projection', 'time.projection')
# Replace remaining underscores with dots, carefully
parts = new_key.split('.')
final_parts = []
for part in parts:
if part in ['img_emb', 'self_attn', 'cross_attn']:
final_parts.append(part) # Keep these intact
else:
final_parts.append(part.replace('_', '.'))
new_key = '.'.join(final_parts)
# Handle weight type
if weight_type:
if weight_type == 'alpha':
new_key += '.alpha'
elif weight_type == 'lora_down.weight' or weight_type == 'lora_down':
new_key += '.lora_A.weight'
elif weight_type == 'lora_up.weight' or weight_type == 'lora_up':
new_key += '.lora_B.weight'
else:
new_key += f'.{weight_type}'
if not new_key.endswith('.weight'):
new_key += '.weight'
k = new_key
# Handle special embedded components
special_components = {
'time.projection': 'time_projection',
'img.emb': 'img_emb',
'text.emb': 'text_emb',
'time.emb': 'time_emb',
}
for old, new in special_components.items():
if old in k:
k = k.replace(old, new)
# Fix diffusion.model -> diffusion_model
if k.startswith('diffusion.model.'):
k = k.replace('diffusion.model.', 'diffusion_model.')
# Finetrainer format
if '.attn1.' in k:
k = k.replace('.attn1.', '.cross_attn.')
k = k.replace('.to_k.', '.k.')