Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
154 changes: 154 additions & 0 deletions src/diffusers/loaders/lora_conversion_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -2736,6 +2736,160 @@ def _convert_to_ai_toolkit_cat(sds_sd, ait_sd, sds_key, ait_keys, dims=None):
return ait_sd


def _bake_lokr_alpha_(state_dict):
"""
Consume `.alpha` keys by baking the LyCORIS `alpha / rank` scaling into the left Kronecker factor. The scaling only
applies when a factor is rank-decomposed (`lokr_w1_a/b` or `lokr_w2_a/b`); when both factors are stored as full
matrices, LoKr applies no alpha scaling and the alpha key is simply dropped.
"""
for alpha_key in [k for k in state_dict if k.endswith(".alpha")]:
alpha = state_dict.pop(alpha_key).item()
module = alpha_key.removesuffix(".alpha")
w1_b = state_dict.get(f"{module}.lokr_w1_b")
w2_b = state_dict.get(f"{module}.lokr_w2_b")
rank = w2_b.shape[0] if w2_b is not None else w1_b.shape[0] if w1_b is not None else None
if rank is None:
continue
w1_key = f"{module}.lokr_w1" if f"{module}.lokr_w1" in state_dict else f"{module}.lokr_w1_a"
state_dict[w1_key] = state_dict[w1_key] * (alpha / rank)


def _convert_non_diffusers_lokr_to_diffusers(state_dict):
"""
Convert a non-diffusers LoKr state dict whose module paths already match the diffusers model (e.g. ai-toolkit
Z-Image checkpoints with keys like `diffusion_model.layers.0.attention.to_q.lokr_w1`) to the peft-loadable format:
the `diffusion_model.` prefix is replaced with `transformer.` and the `.alpha` keys are consumed.
"""
state_dict = {k.removeprefix("diffusion_model."): v for k, v in state_dict.items()}
_bake_lokr_alpha_(state_dict)

non_lokr_keys = [k for k in state_dict if ".lokr_" not in k]
if non_lokr_keys:
raise ValueError(f"`state_dict` contains unexpected non-LoKr keys: {non_lokr_keys}.")

return {f"transformer.{k}": v for k, v in state_dict.items()}


_LOKR_SUFFIXES = ("lokr_w1", "lokr_w1_a", "lokr_w1_b", "lokr_w2", "lokr_w2_a", "lokr_w2_b")


def _convert_non_diffusers_flux2_lokr_to_diffusers(state_dict):
"""
Convert a BFL-format Flux2 LoKr state dict (e.g. trained with ai-toolkit, keys like
`diffusion_model.double_blocks.0.img_attn.qkv.lokr_w1`) to the peft-loadable diffusers format.

BFL checkpoints apply LoKr to the fused QKV projections of the double blocks. Unlike a LoRA delta, a Kronecker
product delta over the fused projection cannot be split exactly into separate Q/K/V factors, so these are mapped to
the model's fused `to_qkv`/`to_added_qkv` projections instead; `Flux2LoraLoaderMixin.load_lora_weights` fuses the
model's projections before injecting such an adapter.
"""
original_state_dict = {k.removeprefix("diffusion_model."): v for k, v in state_dict.items()}
_bake_lokr_alpha_(original_state_dict)

converted_state_dict = {}

num_double_layers = 0
num_single_layers = 0
for key in original_state_dict.keys():
if key.startswith("single_blocks."):
num_single_layers = max(num_single_layers, int(key.split(".")[1]) + 1)
elif key.startswith("double_blocks."):
num_double_layers = max(num_double_layers, int(key.split(".")[1]) + 1)

def _remap(bfl_path, diffusers_path):
for suffix in _LOKR_SUFFIXES:
weight = original_state_dict.pop(f"{bfl_path}.{suffix}", None)
if weight is not None:
converted_state_dict[f"{diffusers_path}.{suffix}"] = weight

for sl in range(num_single_layers):
_remap(f"single_blocks.{sl}.linear1", f"single_transformer_blocks.{sl}.attn.to_qkv_mlp_proj")
_remap(f"single_blocks.{sl}.linear2", f"single_transformer_blocks.{sl}.attn.to_out")

for dl in range(num_double_layers):
tb = f"transformer_blocks.{dl}"
db = f"double_blocks.{dl}"

_remap(f"{db}.img_attn.qkv", f"{tb}.attn.to_qkv")
_remap(f"{db}.txt_attn.qkv", f"{tb}.attn.to_added_qkv")

_remap(f"{db}.img_attn.proj", f"{tb}.attn.to_out.0")
_remap(f"{db}.txt_attn.proj", f"{tb}.attn.to_add_out")

_remap(f"{db}.img_mlp.0", f"{tb}.ff.linear_in")
_remap(f"{db}.img_mlp.2", f"{tb}.ff.linear_out")
_remap(f"{db}.txt_mlp.0", f"{tb}.ff_context.linear_in")
_remap(f"{db}.txt_mlp.2", f"{tb}.ff_context.linear_out")

extra_mappings = {
"img_in": "x_embedder",
"txt_in": "context_embedder",
"time_in.in_layer": "time_guidance_embed.timestep_embedder.linear_1",
"time_in.out_layer": "time_guidance_embed.timestep_embedder.linear_2",
"guidance_in.in_layer": "time_guidance_embed.guidance_embedder.linear_1",
"guidance_in.out_layer": "time_guidance_embed.guidance_embedder.linear_2",
"final_layer.linear": "proj_out",
"final_layer.adaLN_modulation.1": "norm_out.linear",
"single_stream_modulation.lin": "single_stream_modulation.linear",
"double_stream_modulation_img.lin": "double_stream_modulation_img.linear",
"double_stream_modulation_txt.lin": "double_stream_modulation_txt.linear",
}
for bfl_key, diffusers_key in extra_mappings.items():
_remap(bfl_key, diffusers_key)

if len(original_state_dict) > 0:
raise ValueError(f"`original_state_dict` should be empty at this point but has {original_state_dict.keys()=}.")

return {f"transformer.{k}": v for k, v in converted_state_dict.items()}
Comment thread
sayakpaul marked this conversation as resolved.


# Mapping from LyCORIS underscore-encoded sub-paths to dotted Flux2 module paths.
_LYCORIS_FLUX2_SUBPATH_MAP = {
"attn_to_q": "attn.to_q",
"attn_to_k": "attn.to_k",
"attn_to_v": "attn.to_v",
"attn_to_out_0": "attn.to_out.0",
"attn_to_add_out": "attn.to_add_out",
"attn_add_q_proj": "attn.add_q_proj",
"attn_add_k_proj": "attn.add_k_proj",
"attn_add_v_proj": "attn.add_v_proj",
"attn_to_qkv_mlp_proj": "attn.to_qkv_mlp_proj",
"attn_to_out": "attn.to_out",
"ff_linear_in": "ff.linear_in",
"ff_linear_out": "ff.linear_out",
"ff_context_linear_in": "ff_context.linear_in",
"ff_context_linear_out": "ff_context.linear_out",
}


def _convert_lycoris_flux2_lokr_to_diffusers(state_dict):
"""
Convert a LyCORIS-format Flux2 LoKr state dict (keys like `lycoris_transformer_blocks_0_attn_to_q.lokr_w1`) to the
peft-loadable diffusers format. LyCORIS wraps the diffusers model directly and encodes each module path with
underscores, which are decoded through a lookup of the known block sub-paths.
"""
state_dict = dict(state_dict)
_bake_lokr_alpha_(state_dict)

lycoris_key_pattern = re.compile(r"^lycoris_((?:single_)?transformer_blocks)_(\d+)_(.+)\.(.+)$")

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We expect the stat_dict to always contain lycoris prefix? If so, do we want to check against it and raise if needed?

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Yes, the converter is only selected when keys start with lycoris_. Unrecognized keys already ended in an error, but it was the generic "should be empty" one. It now raises upfront and names them: These keys are not LyCORIS Flux2 LoKr keys: [...]. There's a test for it.


converted_state_dict = {}
unrecognized_keys = []
for key, value in state_dict.items():
match = lycoris_key_pattern.match(key)
diffusers_sub_path = _LYCORIS_FLUX2_SUBPATH_MAP.get(match.group(3)) if match is not None else None
if diffusers_sub_path is None:
unrecognized_keys.append(key)
continue
container, block_idx, _, suffix = match.groups()
converted_state_dict[f"transformer.{container}.{block_idx}.{diffusers_sub_path}.{suffix}"] = value

if unrecognized_keys:
raise ValueError(f"These keys are not LyCORIS Flux2 LoKr keys: {unrecognized_keys}.")

return converted_state_dict


def _convert_non_diffusers_z_image_lora_to_diffusers(state_dict):
"""
Convert non-diffusers ZImage LoRA state dict to diffusers format.
Expand Down
31 changes: 23 additions & 8 deletions src/diffusers/loaders/lora_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -45,13 +45,16 @@
_convert_hunyuan_video_lora_to_diffusers,
_convert_kohya_flux2_lora_to_diffusers,
_convert_kohya_flux_lora_to_diffusers,
_convert_lycoris_flux2_lokr_to_diffusers,
_convert_musubi_wan_lora_to_diffusers,
_convert_non_diffusers_ace_step_lora_to_diffusers,
_convert_non_diffusers_anima_lora_to_diffusers,
_convert_non_diffusers_flux2_lokr_to_diffusers,
_convert_non_diffusers_flux2_lora_to_diffusers,
_convert_non_diffusers_hidream_lora_to_diffusers,
_convert_non_diffusers_ideogram4_lora_to_diffusers,
_convert_non_diffusers_krea2_lora_to_diffusers,
_convert_non_diffusers_lokr_to_diffusers,
_convert_non_diffusers_lora_to_diffusers,
_convert_non_diffusers_ltx2_lora_to_diffusers,
_convert_non_diffusers_ltxv_lora_to_diffusers,
Expand Down Expand Up @@ -5443,14 +5446,18 @@ def lora_state_dict(
has_lora_unet = any(k.startswith("lora_unet_") for k in state_dict)
has_diffusion_model = any(k.startswith("diffusion_model.") for k in state_dict)
has_default = any("default." in k for k in state_dict)
if has_alphas_in_sd or has_lora_unet or has_diffusion_model or has_default:
is_lokr = any(".lokr_" in k for k in state_dict)
if is_lokr:
# ai-toolkit Z-Image LoKr checkpoints store module paths that already match the diffusers model.
if has_diffusion_model or has_alphas_in_sd:
state_dict = _convert_non_diffusers_lokr_to_diffusers(state_dict)
elif has_alphas_in_sd or has_lora_unet or has_diffusion_model or has_default:
state_dict = _convert_non_diffusers_z_image_lora_to_diffusers(state_dict)

out = (state_dict, metadata) if return_lora_metadata else state_dict
return out

@require_peft_backend

Copy link
Copy Markdown
Member

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This now becomes exactly the same as Flux2. So, we could change the "# Copied from ..." statement accordingly?

# Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights
def load_lora_weights(
self,
pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor],
Expand All @@ -5471,9 +5478,9 @@ def load_lora_weights(
kwargs["return_lora_metadata"] = True
state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs)

is_correct_format = all("lora" in key for key in state_dict.keys())
is_correct_format = all("lora" in key or "lokr" in key for key in state_dict.keys())
if not is_correct_format:
raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.")
raise ValueError("Invalid adapter checkpoint. We currently support LoRA and LoKr.")

self.load_lora_into_transformer(
state_dict,
Expand Down Expand Up @@ -5831,15 +5838,23 @@ def lora_state_dict(
if is_peft_format:
state_dict = {k.replace("base_model.model.", "diffusion_model."): v for k, v in state_dict.items()}

is_lokr = any(".lokr_" in k for k in state_dict)
is_ai_toolkit = any(k.startswith("diffusion_model.") for k in state_dict)
if is_ai_toolkit:
if is_lokr:
if any(k.startswith("lycoris_") for k in state_dict):
state_dict = _convert_lycoris_flux2_lokr_to_diffusers(state_dict)
elif is_ai_toolkit:
state_dict = _convert_non_diffusers_flux2_lokr_to_diffusers(state_dict)
elif not any(k.startswith("transformer.") for k in state_dict):
# Bare dotted diffusers module paths (e.g. SimpleTuner exports), possibly with alpha keys.
state_dict = _convert_non_diffusers_lokr_to_diffusers(state_dict)
elif is_ai_toolkit:
state_dict = _convert_non_diffusers_flux2_lora_to_diffusers(state_dict)

out = (state_dict, metadata) if return_lora_metadata else state_dict
return out

@require_peft_backend
# Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights
def load_lora_weights(
self,
pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor],
Expand All @@ -5860,9 +5875,9 @@ def load_lora_weights(
kwargs["return_lora_metadata"] = True
state_dict, metadata = self.lora_state_dict(pretrained_model_name_or_path_or_dict, **kwargs)

is_correct_format = all("lora" in key for key in state_dict.keys())
is_correct_format = all("lora" in key or "lokr" in key for key in state_dict.keys())
if not is_correct_format:
raise ValueError("Invalid LoRA checkpoint. Make sure all LoRA param names contain `'lora'` substring.")
raise ValueError("Invalid adapter checkpoint. We currently support LoRA and LoKr.")

self.load_lora_into_transformer(
state_dict,
Expand Down
Loading
Loading