Repository navigation
[LoRA] add LoKr adapter support (Z-Image, Flux2/Klein) #14163
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
base: main
Are you sure you want to change the base?
Changes from all commits
557d685
67ae266
7330b4e
470b0da
90bd8aa
b9c5c58
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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()} | ||
|
|
||
|
|
||
| # 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+)_(.+)\.(.+)$") | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. We expect the
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Yes, the converter is only selected when keys start with |
||
|
|
||
| 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. | ||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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, | ||
|
|
@@ -5443,14 +5446,19 @@ 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 | ||
|
Member
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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?
Contributor
Author
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. Done in b9c5c58: it now carries |
||
| # Copied from diffusers.loaders.lora_pipeline.CogVideoXLoraLoaderMixin.load_lora_weights | ||
| # Copied from diffusers.loaders.lora_pipeline.Flux2LoraLoaderMixin.load_lora_weights | ||
| def load_lora_weights( | ||
| self, | ||
| pretrained_model_name_or_path_or_dict: str | dict[str, torch.Tensor], | ||
|
|
@@ -5471,9 +5479,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, | ||
|
|
@@ -5831,15 +5839,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], | ||
|
|
@@ -5860,9 +5876,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, | ||
|
|
||
Uh oh!
There was an error while loading. Please reload this page.