Skip to content

Commit e6a1235

Browse files
[LoRA] add LoKr adapter support (Z-Image, Flux2/Klein)
Adds loading of LoKr (LyCORIS Kronecker product) adapters: - `load_lora_adapter` detects `lokr_` keys and injects a peft `LoKrConfig`, inferred from the tensor shapes via `_create_lokr_config` (decompose factor, per-module rank/alpha patterns). - State dict conversions for the formats in the wild: ai-toolkit Z-Image (dotted diffusers paths under `diffusion_model.`), ai-toolkit BFL Flux2 (fused qkv), LyCORIS underscore format, and bare dotted diffusers paths. - BFL fused-QKV LoKr cannot be split exactly into separate Q/K/V Kronecker factors, so `Flux2LoraLoaderMixin.load_lora_weights` fuses the model's QKV projections and maps the adapter 1:1 (exact). - Alpha follows the LyCORIS convention: scaling applies only to rank-decomposed factors and is baked into the weights at conversion. Fixes #13221
1 parent 284419b commit e6a1235

6 files changed

Lines changed: 624 additions & 108 deletions

File tree

‎src/diffusers/loaders/lora_conversion_utils.py‎

Lines changed: 160 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -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+
27442904
def _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

Comments
 (0)