@@ -2741,6 +2741,166 @@ def _convert_to_ai_toolkit_cat(sds_sd, ait_sd, sds_key, ait_keys, dims=None):
27412741 return ait_sd
27422742
27432743
2744+ def _bake_lokr_alpha (state_dict ):
2745+ """
2746+ Consume `.alpha` keys by baking the LyCORIS `alpha / rank` scaling into the left Kronecker factor. The scaling only
2747+ applies when a factor is rank-decomposed (`lokr_w1_a/b` or `lokr_w2_a/b`); when both factors are stored as full
2748+ matrices, LoKr applies no alpha scaling and the alpha key is simply dropped.
2749+ """
2750+ for alpha_key in [k for k in state_dict if k .endswith (".alpha" )]:
2751+ alpha = state_dict .pop (alpha_key ).item ()
2752+ module = alpha_key .removesuffix (".alpha" )
2753+ w1_b = state_dict .get (f"{ module } .lokr_w1_b" )
2754+ w2_b = state_dict .get (f"{ module } .lokr_w2_b" )
2755+ rank = w2_b .shape [0 ] if w2_b is not None else w1_b .shape [0 ] if w1_b is not None else None
2756+ if rank is None :
2757+ continue
2758+ w1_key = f"{ module } .lokr_w1" if f"{ module } .lokr_w1" in state_dict else f"{ module } .lokr_w1_a"
2759+ state_dict [w1_key ] = state_dict [w1_key ] * (alpha / rank )
2760+
2761+
2762+ def _convert_non_diffusers_lokr_to_diffusers (state_dict ):
2763+ """
2764+ Convert a non-diffusers LoKr state dict whose module paths already match the diffusers model (e.g. ai-toolkit
2765+ Z-Image checkpoints with keys like `diffusion_model.layers.0.attention.to_q.lokr_w1`) to the peft-loadable format:
2766+ the `diffusion_model.` prefix is replaced with `transformer.` and the `.alpha` keys are consumed.
2767+ """
2768+ state_dict = {k .removeprefix ("diffusion_model." ): v for k , v in state_dict .items ()}
2769+ _bake_lokr_alpha (state_dict )
2770+
2771+ non_lokr_keys = [k for k in state_dict if ".lokr_" not in k ]
2772+ if non_lokr_keys :
2773+ raise ValueError (f"`state_dict` contains unexpected non-LoKr keys: { non_lokr_keys } ." )
2774+
2775+ return {f"transformer.{ k } " : v for k , v in state_dict .items ()}
2776+
2777+
2778+ _LOKR_SUFFIXES = ("lokr_w1" , "lokr_w1_a" , "lokr_w1_b" , "lokr_w2" , "lokr_w2_a" , "lokr_w2_b" )
2779+
2780+
2781+ def _convert_non_diffusers_flux2_lokr_to_diffusers (state_dict ):
2782+ """
2783+ Convert a BFL-format Flux2 LoKr state dict (e.g. trained with ai-toolkit, keys like
2784+ `diffusion_model.double_blocks.0.img_attn.qkv.lokr_w1`) to the peft-loadable diffusers format.
2785+
2786+ BFL checkpoints apply LoKr to the fused QKV projections of the double blocks. Unlike a LoRA delta, a Kronecker
2787+ product delta over the fused projection cannot be split exactly into separate Q/K/V factors, so these are mapped to
2788+ the model's fused `to_qkv`/`to_added_qkv` projections instead; `Flux2LoraLoaderMixin.load_lora_weights` fuses the
2789+ model's projections before injecting such an adapter.
2790+ """
2791+ original_state_dict = {k .removeprefix ("diffusion_model." ): v for k , v in state_dict .items ()}
2792+ _bake_lokr_alpha (original_state_dict )
2793+
2794+ converted_state_dict = {}
2795+
2796+ # Some Flux2 LoKr checkpoints already store expanded diffusers block names; accept those as-is.
2797+ for key in list (original_state_dict .keys ()):
2798+ if key .startswith (("single_transformer_blocks." , "transformer_blocks." )):
2799+ converted_state_dict [key ] = original_state_dict .pop (key )
2800+
2801+ num_double_layers = 0
2802+ num_single_layers = 0
2803+ for key in original_state_dict .keys ():
2804+ if key .startswith ("single_blocks." ):
2805+ num_single_layers = max (num_single_layers , int (key .split ("." )[1 ]) + 1 )
2806+ elif key .startswith ("double_blocks." ):
2807+ num_double_layers = max (num_double_layers , int (key .split ("." )[1 ]) + 1 )
2808+
2809+ def _remap (bfl_path , diffusers_path ):
2810+ for suffix in _LOKR_SUFFIXES :
2811+ weight = original_state_dict .pop (f"{ bfl_path } .{ suffix } " , None )
2812+ if weight is not None :
2813+ converted_state_dict [f"{ diffusers_path } .{ suffix } " ] = weight
2814+
2815+ for sl in range (num_single_layers ):
2816+ _remap (f"single_blocks.{ sl } .linear1" , f"single_transformer_blocks.{ sl } .attn.to_qkv_mlp_proj" )
2817+ _remap (f"single_blocks.{ sl } .linear2" , f"single_transformer_blocks.{ sl } .attn.to_out" )
2818+
2819+ for dl in range (num_double_layers ):
2820+ tb = f"transformer_blocks.{ dl } "
2821+ db = f"double_blocks.{ dl } "
2822+
2823+ _remap (f"{ db } .img_attn.qkv" , f"{ tb } .attn.to_qkv" )
2824+ _remap (f"{ db } .txt_attn.qkv" , f"{ tb } .attn.to_added_qkv" )
2825+
2826+ _remap (f"{ db } .img_attn.proj" , f"{ tb } .attn.to_out.0" )
2827+ _remap (f"{ db } .txt_attn.proj" , f"{ tb } .attn.to_add_out" )
2828+
2829+ _remap (f"{ db } .img_mlp.0" , f"{ tb } .ff.linear_in" )
2830+ _remap (f"{ db } .img_mlp.2" , f"{ tb } .ff.linear_out" )
2831+ _remap (f"{ db } .txt_mlp.0" , f"{ tb } .ff_context.linear_in" )
2832+ _remap (f"{ db } .txt_mlp.2" , f"{ tb } .ff_context.linear_out" )
2833+
2834+ extra_mappings = {
2835+ "img_in" : "x_embedder" ,
2836+ "txt_in" : "context_embedder" ,
2837+ "time_in.in_layer" : "time_guidance_embed.timestep_embedder.linear_1" ,
2838+ "time_in.out_layer" : "time_guidance_embed.timestep_embedder.linear_2" ,
2839+ "guidance_in.in_layer" : "time_guidance_embed.guidance_embedder.linear_1" ,
2840+ "guidance_in.out_layer" : "time_guidance_embed.guidance_embedder.linear_2" ,
2841+ "final_layer.linear" : "proj_out" ,
2842+ "final_layer.adaLN_modulation.1" : "norm_out.linear" ,
2843+ "single_stream_modulation.lin" : "single_stream_modulation.linear" ,
2844+ "double_stream_modulation_img.lin" : "double_stream_modulation_img.linear" ,
2845+ "double_stream_modulation_txt.lin" : "double_stream_modulation_txt.linear" ,
2846+ }
2847+ for bfl_key , diffusers_key in extra_mappings .items ():
2848+ _remap (bfl_key , diffusers_key )
2849+
2850+ if len (original_state_dict ) > 0 :
2851+ raise ValueError (f"`original_state_dict` should be empty at this point but has { original_state_dict .keys ()= } ." )
2852+
2853+ return {f"transformer.{ k } " : v for k , v in converted_state_dict .items ()}
2854+
2855+
2856+ # Mapping from LyCORIS underscore-encoded sub-paths to dotted Flux2 module paths.
2857+ _LYCORIS_FLUX2_SUBPATH_MAP = {
2858+ "attn_to_q" : "attn.to_q" ,
2859+ "attn_to_k" : "attn.to_k" ,
2860+ "attn_to_v" : "attn.to_v" ,
2861+ "attn_to_out_0" : "attn.to_out.0" ,
2862+ "attn_to_add_out" : "attn.to_add_out" ,
2863+ "attn_add_q_proj" : "attn.add_q_proj" ,
2864+ "attn_add_k_proj" : "attn.add_k_proj" ,
2865+ "attn_add_v_proj" : "attn.add_v_proj" ,
2866+ "attn_to_qkv_mlp_proj" : "attn.to_qkv_mlp_proj" ,
2867+ "attn_to_out" : "attn.to_out" ,
2868+ "ff_linear_in" : "ff.linear_in" ,
2869+ "ff_linear_out" : "ff.linear_out" ,
2870+ "ff_context_linear_in" : "ff_context.linear_in" ,
2871+ "ff_context_linear_out" : "ff_context.linear_out" ,
2872+ }
2873+
2874+
2875+ def _convert_lycoris_flux2_lokr_to_diffusers (state_dict ):
2876+ """
2877+ Convert a LyCORIS-format Flux2 LoKr state dict (keys like `lycoris_transformer_blocks_0_attn_to_q.lokr_w1`) to the
2878+ peft-loadable diffusers format. LyCORIS wraps the diffusers model directly and encodes each module path with
2879+ underscores, which are decoded through a lookup of the known block sub-paths.
2880+ """
2881+ state_dict = dict (state_dict )
2882+ _bake_lokr_alpha (state_dict )
2883+
2884+ lycoris_key_pattern = re .compile (r"^lycoris_((?:single_)?transformer_blocks)_(\d+)_(.+)\.(.+)$" )
2885+
2886+ converted_state_dict = {}
2887+ for key in list (state_dict .keys ()):
2888+ match = lycoris_key_pattern .match (key )
2889+ if match is None :
2890+ continue
2891+ container , block_idx , sub_path , suffix = match .groups ()
2892+ diffusers_sub_path = _LYCORIS_FLUX2_SUBPATH_MAP .get (sub_path )
2893+ if diffusers_sub_path is None :
2894+ continue
2895+ diffusers_key = f"transformer.{ container } .{ block_idx } .{ diffusers_sub_path } .{ suffix } "
2896+ converted_state_dict [diffusers_key ] = state_dict .pop (key )
2897+
2898+ if len (state_dict ) > 0 :
2899+ raise ValueError (f"`state_dict` should be empty at this point but has { state_dict .keys ()= } ." )
2900+
2901+ return converted_state_dict
2902+
2903+
27442904def _convert_non_diffusers_z_image_lora_to_diffusers (state_dict ):
27452905 """
27462906 Convert non-diffusers ZImage LoRA state dict to diffusers format.
0 commit comments