Describe the bug
Some FLUX.1 Kontext LoRAs exported by fal store the block keys under base_model.model. but the global embedder keys (time_in, vector_in, txt_in, img_in, guidance_in) without that prefix. FluxPipeline.lora_state_dict routes them to _convert_fal_kontext_lora_to_diffusers (any key containing base_model), which does not map the embedder keys and then raises on them as leftovers.
Reproduction
CPU only, no checkpoint download. The state dict mirrors the fal Kontext layout with zero tensors:
import torch
from diffusers import FluxPipeline
prefix, rank, inner, mlp = "base_model.model.", 1, 3072, 12288
sd = {}
def add(name, out_features):
sd[f"{prefix}{name}.lora_A.weight"] = torch.zeros(rank, 1)
sd[f"{prefix}{name}.lora_B.weight"] = torch.zeros(out_features, rank)
for i in range(19):
for m in ["img_mod.lin", "txt_mod.lin", "img_mlp.0", "img_mlp.2", "txt_mlp.0", "txt_mlp.2", "img_attn.proj", "txt_attn.proj"]:
add(f"double_blocks.{i}.{m}", 1)
for m in ["img_attn.qkv", "txt_attn.qkv"]:
add(f"double_blocks.{i}.{m}", 3 * inner)
for i in range(38):
add(f"single_blocks.{i}.modulation.lin", 1)
add(f"single_blocks.{i}.linear1", 3 * inner + mlp)
add(f"single_blocks.{i}.linear2", 1)
add("final_layer.linear", 1)
# global embedders without the base_model.model. prefix
for name in ["time_in.in_layer", "time_in.out_layer", "vector_in.in_layer", "vector_in.out_layer", "txt_in", "img_in", "guidance_in.in_layer", "guidance_in.out_layer"]:
sd[f"{name}.lora_A.weight"] = torch.zeros(rank, 1)
sd[f"{name}.lora_B.weight"] = torch.zeros(1, rank)
FluxPipeline.lora_state_dict(sd)
Logs
ValueError: `original_state_dict` should be empty at this point but has original_state_dict.keys()=dict_keys(['time_in.in_layer.lora_A.weight', 'time_in.in_layer.lora_B.weight', ..., 'guidance_in.out_layer.lora_B.weight']).
System Info
diffusers main (031b279), torch 2.14.0 (CPU), transformers 5.17.0, peft 0.21.0, Python 3.12.
Who can help?
@sayakpaul
Describe the bug
Some FLUX.1 Kontext LoRAs exported by fal store the block keys under
base_model.model.but the global embedder keys (time_in,vector_in,txt_in,img_in,guidance_in) without that prefix.FluxPipeline.lora_state_dictroutes them to_convert_fal_kontext_lora_to_diffusers(any key containingbase_model), which does not map the embedder keys and then raises on them as leftovers.Reproduction
CPU only, no checkpoint download. The state dict mirrors the fal Kontext layout with zero tensors:
Logs
System Info
diffusers main (031b279), torch 2.14.0 (CPU), transformers 5.17.0, peft 0.21.0, Python 3.12.
Who can help?
@sayakpaul