From f8155ab47f381ba1a89ddd96499a925ab6fd3af3 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 7 Apr 2025 12:38:56 +0300 Subject: [PATCH] Allow loading Fun Reward LoRAs https://huggingface.co/alibaba-pai/Wan2.1-Fun-Reward-LoRAs/tree/main --- nodes.py | 117 ++++++++++++++++++++++++++++++++++++++++++++++++++++++- 1 file changed, 116 insertions(+), 1 deletion(-) diff --git a/nodes.py b/nodes.py index 9d6f087..435fbe7 100644 --- a/nodes.py +++ b/nodes.py @@ -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.')