diff --git a/src/diffusers/loaders/lora_conversion_utils.py b/src/diffusers/loaders/lora_conversion_utils.py index cef2e88454a8..bfaec3b0b5a7 100644 --- a/src/diffusers/loaders/lora_conversion_utils.py +++ b/src/diffusers/loaders/lora_conversion_utils.py @@ -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+)_(.+)\.(.+)$") + + 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. diff --git a/src/diffusers/loaders/lora_pipeline.py b/src/diffusers/loaders/lora_pipeline.py index 1003aa57c420..46867b2257f7 100644 --- a/src/diffusers/loaders/lora_pipeline.py +++ b/src/diffusers/loaders/lora_pipeline.py @@ -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 - # 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, diff --git a/src/diffusers/loaders/peft.py b/src/diffusers/loaders/peft.py index 5ab34f4f0a3d..3f0652c979c5 100644 --- a/src/diffusers/loaders/peft.py +++ b/src/diffusers/loaders/peft.py @@ -37,7 +37,12 @@ set_adapter_layers, set_weights_and_activate_adapters, ) -from ..utils.peft_utils import _create_lora_config, _maybe_warn_for_unhandled_keys +from ..utils.peft_utils import ( + _create_lokr_config, + _create_lora_config, + _maybe_fuse_qkv_projections_for_lokr, + _maybe_warn_for_unhandled_keys, +) from .lora_base import _fetch_state_dict, _func_optionally_disable_offloading from .unet_loader_utils import _maybe_expand_lora_scales @@ -217,56 +222,65 @@ def load_lora_adapter( "Please choose an existing adapter name or set `hotswap=False` to prevent hotswapping." ) - # check with first key if is not in peft format - first_key = next(iter(state_dict.keys())) - if "lora_A" not in first_key: - state_dict = convert_unet_state_dict_to_peft(state_dict) - - # Control LoRA from SAI is different from BFL Control LoRA - # https://huggingface.co/stabilityai/control-lora - # https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors - is_sai_sd_control_lora = "lora_controlnet" in state_dict - if is_sai_sd_control_lora: - state_dict = convert_sai_sd_control_lora_state_dict_to_peft(state_dict) - - rank = {} - for key, val in state_dict.items(): - # Cannot figure out rank from lora layers that don't have at least 2 dimensions. - # Bias layers in LoRA only have a single dimension - if "lora_B" in key and val.ndim > 1: - # Check out https://github.com/huggingface/peft/pull/2419 for the `^` symbol. - # We may run into some ambiguous configuration values when a model has module - # names, sharing a common prefix (`proj_out.weight` and `blocks.transformer.proj_out.weight`, - # for example) and they have different LoRA ranks. - rank[f"^{key}"] = val.shape[1] - - if network_alphas is not None and len(network_alphas) >= 1: - alpha_keys = [k for k in network_alphas.keys() if k.startswith(f"{prefix}.")] - network_alphas = { - k.removeprefix(f"{prefix}."): v for k, v in network_alphas.items() if k in alpha_keys - } - - # adapter_name if adapter_name is None: adapter_name = get_adapter_name(self) - # create LoraConfig - lora_config = _create_lora_config( - state_dict, - network_alphas, - metadata, - rank, - model_state_dict=self.state_dict(), - adapter_name=adapter_name, - ) + # LoKr adapters (Kronecker product factors) use a different peft config and state dict layout + # (`{module}.lokr_w1` etc.) than LoRA; detect them before any LoRA-specific key handling. + is_lokr = any(".lokr_" in k for k in state_dict) + + if is_lokr: + if hotswap: + raise ValueError("Hotswapping LoKr adapters is not supported.") + _maybe_fuse_qkv_projections_for_lokr(self, state_dict) + adapter_config = _create_lokr_config(state_dict, metadata) + else: + # check with first key if is not in peft format + first_key = next(iter(state_dict.keys())) + if "lora_A" not in first_key: + state_dict = convert_unet_state_dict_to_peft(state_dict) + + # Control LoRA from SAI is different from BFL Control LoRA + # https://huggingface.co/stabilityai/control-lora + # https://huggingface.co/comfyanonymous/ControlNet-v1-1_fp16_safetensors + is_sai_sd_control_lora = "lora_controlnet" in state_dict + if is_sai_sd_control_lora: + state_dict = convert_sai_sd_control_lora_state_dict_to_peft(state_dict) + + rank = {} + for key, val in state_dict.items(): + # Cannot figure out rank from lora layers that don't have at least 2 dimensions. + # Bias layers in LoRA only have a single dimension + if "lora_B" in key and val.ndim > 1: + # Check out https://github.com/huggingface/peft/pull/2419 for the `^` symbol. + # We may run into some ambiguous configuration values when a model has module + # names, sharing a common prefix (`proj_out.weight` and `blocks.transformer.proj_out.weight`, + # for example) and they have different LoRA ranks. + rank[f"^{key}"] = val.shape[1] + + if network_alphas is not None and len(network_alphas) >= 1: + alpha_keys = [k for k in network_alphas.keys() if k.startswith(f"{prefix}.")] + network_alphas = { + k.removeprefix(f"{prefix}."): v for k, v in network_alphas.items() if k in alpha_keys + } + + # create LoraConfig + adapter_config = _create_lora_config( + state_dict, + network_alphas, + metadata, + rank, + model_state_dict=self.state_dict(), + adapter_name=adapter_name, + ) - # Adjust LoRA config for Control LoRA - if is_sai_sd_control_lora: - lora_config.lora_alpha = lora_config.r - lora_config.alpha_pattern = lora_config.rank_pattern - lora_config.bias = "all" - lora_config.modules_to_save = lora_config.exclude_modules - lora_config.exclude_modules = None + # Adjust LoRA config for Control LoRA + if is_sai_sd_control_lora: + adapter_config.lora_alpha = adapter_config.r + adapter_config.alpha_pattern = adapter_config.rank_pattern + adapter_config.bias = "all" + adapter_config.modules_to_save = adapter_config.exclude_modules + adapter_config.exclude_modules = None # None: ) +def _maybe_fuse_qkv_projections_for_lokr(model, state_dict) -> None: + """ + Fuse the model's QKV projections when a peft-format LoKr state dict targets fused `to_qkv` / `to_added_qkv` + projections that the model does not have yet. + + BFL-format Flux2 LoKr checkpoints apply LoKr to the fused QKV projections. Unlike a LoRA delta, a Kronecker product + delta over the fused projection cannot be split exactly into separate Q/K/V factors, so the model's projections are + fused instead and the adapter maps 1:1. + """ + fused_targets = { + module + for module in (k.rpartition(".lokr_")[0] for k in state_dict if ".lokr_" in k) + if module.rsplit(".", 1)[-1] in ("to_qkv", "to_added_qkv") + } + named_modules = dict(model.named_modules()) + if all(module in named_modules for module in fused_targets) or not hasattr(model, "fuse_qkv_projections"): + return + + if getattr(model, "is_quantized", False): + raise ValueError( + "This LoKr checkpoint targets fused QKV projections. Fusing concatenates the Q/K/V weights into a new " + "`nn.Linear`, which is not possible with quantized weights. Please load the transformer without " + "quantization." + ) + + # Fusing replaces to_q/to_k/to_v (and the add_*_proj) with a single projection, which would orphan any adapter + # already injected on the unfused ones. + from peft.tuners.tuners_utils import BaseTunerLayer + + unfused_projections = {"to_q", "to_k", "to_v", "add_q_proj", "add_k_proj", "add_v_proj"} + adapted = [ + name + for name, module in named_modules.items() + if isinstance(module, BaseTunerLayer) and name.rsplit(".", 1)[-1] in unfused_projections + ] + if adapted: + raise ValueError( + "This LoKr checkpoint targets fused QKV projections, but an adapter is already loaded on the unfused " + f"projections (e.g. `{adapted[0]}`). Unload it with `unload_lora_weights()` before loading this checkpoint." + ) + + logger.info("The LoKr checkpoint targets fused QKV projections; calling `fuse_qkv_projections()` on the model.") + model.fuse_qkv_projections() + + +def _create_lokr_config(state_dict, metadata): + """ + Create a `LoKrConfig` from a peft-format LoKr state dict (keys like `{module}.lokr_w1`). + + Without metadata, the config is inferred from the tensor shapes. The checkpoint alpha is expected to be already + baked into the weights by the state dict conversion, so `alpha` is set equal to the rank (runtime scaling 1.0). + peft re-derives each module's Kronecker factorization from `decompose_factor` and only creates rank-decomposed + factors when the rank is small compared to the factorized dimensions, so for modules whose checkpoint factors are + full matrices the rank is set to `max(lokr_w2.shape)` to make peft create full matrices as well. + """ + from peft import LoKrConfig + from peft.tuners.lokr.layer import factorization + + if metadata is not None: + try: + return LoKrConfig(**metadata) + except TypeError as e: + raise TypeError("`LoKrConfig` class could not be instantiated.") from e + + modules = sorted({k.rpartition(".lokr_")[0] for k in state_dict if ".lokr_" in k}) + + # Reconstruct each module's factorized dimensions, (out_l, out_k) x (in_m, in_n), from the checkpoint. A factor + # is either stored as a full matrix (`lokr_w1`) or rank-decomposed into `lokr_w1_a @ lokr_w1_b` (same for w2). + factorizations = {} + rank_dict = {} + for module in modules: + w1, w1_a, w1_b = (state_dict.get(f"{module}.lokr_w1{s}") for s in ("", "_a", "_b")) + w2, w2_a, w2_b = (state_dict.get(f"{module}.lokr_w2{s}") for s in ("", "_a", "_b")) + out_l, in_m = w1.shape if w1 is not None else (w1_a.shape[0], w1_b.shape[1]) + out_k, in_n = w2.shape if w2 is not None else (w2_a.shape[0], w2_b.shape[1]) + factorizations[module] = ((out_l, out_k), (in_m, in_n)) + if w2_a is not None: + rank_dict[module] = w2_a.shape[1] + elif w1_a is not None: + rank_dict[module] = w1_a.shape[1] + else: + rank_dict[module] = max(w2.shape) + + # Find the `decompose_factor` under which peft reproduces the checkpoint factorizations. A fixed factor shows up + # as the left dimension of the modules it divides (modules it does not divide fall back to a near-square + # factorization, like with factor -1), so every observed left dimension is a candidate. + left_dims = {dims[0][0] for dims in factorizations.values()} + decompose_factor = None + for candidate in sorted(left_dims) + [-1]: + if all( + factorization(out_l * out_k, candidate) == (out_l, out_k) + and factorization(in_m * in_n, candidate) == (in_m, in_n) + for (out_l, out_k), (in_m, in_n) in factorizations.values() + ): + decompose_factor = candidate + break + if decompose_factor is None: + raise ValueError( + "Could not infer a `decompose_factor` that reproduces the Kronecker factorizations of this LoKr " + "state dict. Please open an issue: https://github.com/huggingface/diffusers/issues/new" + ) + + r = collections.Counter(rank_dict.values()).most_common(1)[0][0] + rank_pattern = {k: v for k, v in rank_dict.items() if v != r} + + lokr_config_kwargs = { + "r": r, + "alpha": r, + "rank_pattern": rank_pattern, + "alpha_pattern": dict(rank_pattern), + "target_modules": modules, + "decompose_both": any(".lokr_w1_a" in k for k in state_dict), + "decompose_factor": decompose_factor, + } + try: + return LoKrConfig(**lokr_config_kwargs) + except TypeError as e: + raise TypeError("`LoKrConfig` class could not be instantiated.") from e + + def _create_lora_config( state_dict, network_alphas, metadata, rank_pattern_dict, is_unet=True, model_state_dict=None, adapter_name=None ): @@ -397,7 +517,7 @@ def _maybe_warn_for_unhandled_keys(incompatible_keys, adapter_name: str) -> None # Check only for unexpected keys. unexpected_keys = getattr(incompatible_keys, "unexpected_keys", None) if unexpected_keys: - lora_unexpected_keys = [k for k in unexpected_keys if ".lora_" in k] + lora_unexpected_keys = [k for k in unexpected_keys if ".lora_" in k or "lokr_" in k] if lora_unexpected_keys: warn_msg = ( f"Loading adapter weights from state_dict led to unexpected keys found in the model:" @@ -407,7 +527,7 @@ def _maybe_warn_for_unhandled_keys(incompatible_keys, adapter_name: str) -> None # Filter missing keys specific to the current adapter. missing_keys = getattr(incompatible_keys, "missing_keys", None) if missing_keys: - lora_missing_keys = [k for k in missing_keys if ".lora_" in k and adapter_name in k] + lora_missing_keys = [k for k in missing_keys if (".lora_" in k or "lokr_" in k) and adapter_name in k] if lora_missing_keys: warn_msg += ( f"Loading adapter weights from state_dict led to missing keys in the model:" diff --git a/tests/models/testing_utils/__init__.py b/tests/models/testing_utils/__init__.py index 2d7d5ae23257..8bab73f67ca4 100644 --- a/tests/models/testing_utils/__init__.py +++ b/tests/models/testing_utils/__init__.py @@ -17,6 +17,7 @@ from .common import BaseModelTesterConfig, ModelTesterMixin from .compile import TorchCompileTesterMixin from .ip_adapter import IPAdapterTesterMixin +from .lokr import LoKrTesterMixin from .lora import LoraHotSwappingForModelTesterMixin, LoraTesterMixin from .memory import CPUOffloadTesterMixin, GroupOffloadTesterMixin, LayerwiseCastingTesterMixin, MemoryTesterMixin from .parallelism import ( @@ -80,6 +81,7 @@ "GroupOffloadTesterMixin", "IPAdapterTesterMixin", "LayerwiseCastingTesterMixin", + "LoKrTesterMixin", "LoraHotSwappingForModelTesterMixin", "LoraTesterMixin", "MemoryTesterMixin", diff --git a/tests/models/testing_utils/lokr.py b/tests/models/testing_utils/lokr.py new file mode 100644 index 000000000000..d1589d28fadd --- /dev/null +++ b/tests/models/testing_utils/lokr.py @@ -0,0 +1,139 @@ +# coding=utf-8 +# Copyright 2026 HuggingFace Inc. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import pytest +import torch +import torch.nn as nn + +from diffusers.utils.import_utils import is_peft_available + +from ...testing_utils import assert_tensors_close, is_lora, require_peft_backend, torch_device +from .common import BaseModelOutputMixin + + +if is_peft_available(): + from peft.tuners.lokr.layer import LoKrLayer, factorization + + from diffusers.loaders.peft import PeftAdapterMixin + + +def make_lokr_factors(out_features, in_features, factor=4, rank=None): + """ + Random LoKr factors for an `out_features x in_features` layer, factorized like peft does with `decompose_factor`. + With `rank`, the right factor is stored rank-decomposed (`lokr_w2_a @ lokr_w2_b`). + + Returns the factors keyed by their state dict suffix, and the delta weight they encode. + """ + out_l, out_k = factorization(out_features, factor) + in_m, in_n = factorization(in_features, factor) + w1 = torch.randn(out_l, in_m) + if rank is None: + w2 = torch.randn(out_k, in_n) + return {"lokr_w1": w1, "lokr_w2": w2}, torch.kron(w1, w2) + w2_a, w2_b = torch.randn(out_k, rank), torch.randn(rank, in_n) + return {"lokr_w1": w1, "lokr_w2_a": w2_a, "lokr_w2_b": w2_b}, torch.kron(w1, w2_a @ w2_b) + + +def check_lokr_deltas(model, expected_deltas, adapter_name="default", atol=1e-5): + """Check that exactly the expected modules carry a LoKr adapter, each with the expected delta weight.""" + named_modules = dict(model.named_modules()) + adapted = {name for name, module in named_modules.items() if isinstance(module, LoKrLayer)} + assert adapted == set(expected_deltas) + for name, expected in expected_deltas.items(): + delta = named_modules[name].get_delta_weight(adapter_name).cpu() + assert_tensors_close(delta, expected, atol=atol, rtol=0, msg=f"Wrong LoKr delta on {name}") + + +@is_lora +@require_peft_backend +class LoKrTesterMixin(BaseModelOutputMixin): + """ + Mixin class for testing loading LoKr (LyCORIS Kronecker product) adapters with `load_lora_adapter`. + + Expected from config mixin: + - model_class: The model class to test + + Expected methods from config mixin: + - get_init_dict(): Returns dict of arguments to initialize the model + - get_dummy_inputs(): Returns dict of inputs to pass to the model forward pass + + Pytest mark: lora + Use `pytest -m "not lora"` to skip these tests + """ + + def setup_method(self): + if not issubclass(self.model_class, PeftAdapterMixin): + pytest.skip(f"PEFT is not supported for this model ({self.model_class.__name__}).") + + def _flatten_output(self, output): + # Some models (e.g. Z-Image) return a list of per-sample tensors. + if isinstance(output, (list, tuple)): + return torch.cat([t.flatten() for t in output]) + return output + + def _model_output(self, model, inputs_dict): + return self._flatten_output(model(**inputs_dict, return_dict=False)[0]) + + def get_lokr_state_dict(self, model, rank=None): + """ + A peft-format LoKr state dict on every attention `to_q` and `to_v`, with the expected delta of each. With + `rank`, the `to_v` factors are rank-decomposed, so the config inference has to mix both kinds. + """ + state_dict, expected_deltas = {}, {} + for name, module in model.named_modules(): + projection = name.rsplit(".", 1)[-1] + if isinstance(module, nn.Linear) and projection in ("to_q", "to_v"): + factors, delta = make_lokr_factors( + module.out_features, module.in_features, rank=rank if projection == "to_v" else None + ) + state_dict.update({f"{name}.{suffix}": weight for suffix, weight in factors.items()}) + expected_deltas[name] = delta + return state_dict, expected_deltas + + @pytest.mark.parametrize("rank", [None, 1], ids=["full_factors", "rank_decomposed_factors"]) + @torch.no_grad() + def test_lokr_adapter_loads_exact_kronecker_deltas(self, base_model_output, rank): + torch.manual_seed(0) + model = self.model_class(**self.get_init_dict()).eval().to(torch_device) + state_dict, expected_deltas = self.get_lokr_state_dict(model, rank=rank) + + model.load_lora_adapter(state_dict, prefix=None, adapter_name="default") + + check_lokr_deltas(model, expected_deltas) + output = self._model_output(model, self.get_dummy_inputs()) + base_output = self._flatten_output(base_model_output) + assert not torch.allclose(output, base_output, atol=1e-4, rtol=1e-4), "Output should differ with LoKr" + + @torch.no_grad() + def test_lokr_unload_restores_base_output(self, base_model_output): + torch.manual_seed(0) + model = self.model_class(**self.get_init_dict()).eval().to(torch_device) + state_dict, _ = self.get_lokr_state_dict(model) + + model.load_lora_adapter(state_dict, prefix=None, adapter_name="default") + model.unload_lora() + + assert not any(isinstance(module, LoKrLayer) for module in model.modules()) + output = self._model_output(model, self.get_dummy_inputs()) + assert_tensors_close(output, self._flatten_output(base_model_output), atol=1e-4, rtol=1e-4) + + def test_lokr_hotswap_raises(self): + torch.manual_seed(0) + model = self.model_class(**self.get_init_dict()).eval().to(torch_device) + state_dict, _ = self.get_lokr_state_dict(model) + model.load_lora_adapter(state_dict, prefix=None, adapter_name="default") + + with pytest.raises(ValueError, match="Hotswapping LoKr adapters is not supported"): + model.load_lora_adapter(state_dict, prefix=None, adapter_name="default", hotswap=True) diff --git a/tests/models/transformers/test_models_transformer_flux2.py b/tests/models/transformers/test_models_transformer_flux2.py index 3263ce68202c..cc4f5cad966a 100644 --- a/tests/models/transformers/test_models_transformer_flux2.py +++ b/tests/models/transformers/test_models_transformer_flux2.py @@ -17,9 +17,11 @@ import subprocess import sys +import pytest import torch from diffusers import Flux2Transformer2DModel +from diffusers.loaders.lora_pipeline import Flux2LoraLoaderMixin from diffusers.models.transformers.transformer_flux2 import ( Flux2KVAttnProcessor, Flux2KVCache, @@ -36,6 +38,7 @@ ContextParallelTesterMixin, GGUFCompileTesterMixin, GGUFTesterMixin, + LoKrTesterMixin, LoraHotSwappingForModelTesterMixin, LoraTesterMixin, MemoryTesterMixin, @@ -47,6 +50,7 @@ TorchCompileTesterMixin, TrainingTesterMixin, ) +from ..testing_utils.lokr import check_lokr_deltas, make_lokr_factors enable_full_determinism() @@ -202,6 +206,116 @@ class TestFlux2TransformerLoRA(Flux2TransformerTesterConfig, LoraTesterMixin): """LoRA adapter tests for Flux2 Transformer.""" +class TestFlux2TransformerLoKr(Flux2TransformerTesterConfig, LoKrTesterMixin): + """LoKr adapter tests for Flux2 Transformer, including the Flux2 LoKr checkpoint formats.""" + + # ai-toolkit stores a placeholder alpha for full-matrix factors, where LoKr applies no scaling. + placeholder_alpha = torch.tensor(9999220736.0) + + def get_bfl_qkv_state_dict(self, model): + """A BFL-format LoKr state dict on the fused QKV projections of the first double block.""" + to_q = model.transformer_blocks[0].attn.to_q + state_dict, expected_deltas = {}, {} + for bfl_path, diffusers_path in [ + ("double_blocks.0.img_attn.qkv", "transformer_blocks.0.attn.to_qkv"), + ("double_blocks.0.txt_attn.qkv", "transformer_blocks.0.attn.to_added_qkv"), + ]: + factors, expected_deltas[diffusers_path] = make_lokr_factors(3 * to_q.out_features, to_q.in_features) + state_dict.update({f"diffusion_model.{bfl_path}.{k}": v for k, v in factors.items()}) + state_dict[f"diffusion_model.{bfl_path}.alpha"] = self.placeholder_alpha + return state_dict, expected_deltas + + @torch.no_grad() + def test_lokr_bfl_checkpoint(self): + # BFL checkpoints (e.g. ai-toolkit) apply LoKr to the fused QKV projections. A Kronecker product delta cannot + # be split exactly into Q/K/V, so loading fuses the model's projections and maps the adapter 1:1. + torch.manual_seed(0) + model = self.model_class(**self.get_init_dict()).eval().to(torch_device) + state_dict, expected_deltas = self.get_bfl_qkv_state_dict(model) + for bfl_path, diffusers_path in [ + ("single_blocks.0.linear1", "single_transformer_blocks.0.attn.to_qkv_mlp_proj"), + ("double_blocks.0.img_attn.proj", "transformer_blocks.0.attn.to_out.0"), + ("double_blocks.0.img_mlp.0", "transformer_blocks.0.ff.linear_in"), + ]: + linear = model.get_submodule(diffusers_path) + factors, expected_deltas[diffusers_path] = make_lokr_factors(linear.out_features, linear.in_features) + state_dict.update({f"diffusion_model.{bfl_path}.{k}": v for k, v in factors.items()}) + state_dict[f"diffusion_model.{bfl_path}.alpha"] = self.placeholder_alpha + + converted = Flux2LoraLoaderMixin.lora_state_dict(state_dict) + model.load_lora_adapter(converted, prefix="transformer", adapter_name="default") + + assert model.transformer_blocks[0].attn.fused_projections + check_lokr_deltas(model, expected_deltas) + + def test_lokr_fused_qkv_checkpoint_refuses_when_unfused_projections_are_adapted(self): + # Fusing would replace to_q and orphan the adapter already injected there. + from peft import LoraConfig + + model = self.model_class(**self.get_init_dict()).eval().to(torch_device) + model.add_adapter(LoraConfig(r=2, target_modules=["to_q"]), adapter_name="lora") + state_dict, _ = self.get_bfl_qkv_state_dict(model) + converted = Flux2LoraLoaderMixin.lora_state_dict(state_dict) + + with pytest.raises(ValueError, match="already loaded on the unfused projections"): + model.load_lora_adapter(converted, prefix="transformer", adapter_name="lokr") + assert not model.transformer_blocks[0].attn.fused_projections + + @torch.no_grad() + def test_lokr_lycoris_checkpoint(self): + # LyCORIS wraps the diffusers model and encodes module paths with underscores under a `lycoris_` prefix. + torch.manual_seed(0) + model = self.model_class(**self.get_init_dict()).eval().to(torch_device) + state_dict, expected_deltas = {}, {} + for diffusers_path in [ + "single_transformer_blocks.0.attn.to_qkv_mlp_proj", + "transformer_blocks.0.attn.to_q", + "transformer_blocks.0.attn.to_out.0", + "transformer_blocks.0.ff.linear_in", + ]: + linear = model.get_submodule(diffusers_path) + factors, expected_deltas[diffusers_path] = make_lokr_factors(linear.out_features, linear.in_features) + lycoris_path = "lycoris_" + diffusers_path.replace(".", "_") + state_dict.update({f"{lycoris_path}.{k}": v for k, v in factors.items()}) + state_dict[f"{lycoris_path}.alpha"] = torch.tensor(16.0) + + converted = Flux2LoraLoaderMixin.lora_state_dict(state_dict) + model.load_lora_adapter(converted, prefix="transformer", adapter_name="default") + + check_lokr_deltas(model, expected_deltas) + + def test_lokr_lycoris_checkpoint_with_unknown_keys_raises(self): + state_dict = { + "lycoris_transformer_blocks_0_attn_to_q.lokr_w1": torch.randn(4, 4), + "lycoris_transformer_blocks_0_attn_norm_q.lokr_w1": torch.randn(4, 4), + } + with pytest.raises(ValueError, match="lycoris_transformer_blocks_0_attn_norm_q.lokr_w1"): + Flux2LoraLoaderMixin.lora_state_dict(state_dict) + + @torch.no_grad() + def test_lokr_diffusers_names_checkpoint(self): + # Checkpoints that store the diffusers module paths directly, with alpha keys and no prefix (e.g. SimpleTuner, + # `bghira/flux2-klein-9b-distillation-lokr`). Alpha scales the rank-decomposed factors only. + torch.manual_seed(0) + model = self.model_class(**self.get_init_dict()).eval().to(torch_device) + rank, alpha = 2, 1.0 + state_dict, expected_deltas = {}, {} + for diffusers_path, factor_rank in [ + ("single_transformer_blocks.0.attn.to_out", None), + ("transformer_blocks.0.attn.to_k", rank), + ]: + linear = model.get_submodule(diffusers_path) + factors, delta = make_lokr_factors(linear.out_features, linear.in_features, rank=factor_rank) + state_dict.update({f"{diffusers_path}.{k}": v for k, v in factors.items()}) + state_dict[f"{diffusers_path}.alpha"] = torch.tensor(alpha) + expected_deltas[diffusers_path] = delta if factor_rank is None else (alpha / rank) * delta + + converted = Flux2LoraLoaderMixin.lora_state_dict(state_dict) + model.load_lora_adapter(converted, prefix="transformer", adapter_name="default") + + check_lokr_deltas(model, expected_deltas) + + class TestFlux2TransformerLoRAHotSwap(Flux2TransformerTesterConfig, LoraHotSwappingForModelTesterMixin): """LoRA hot-swapping tests for Flux2 Transformer.""" diff --git a/tests/models/transformers/test_models_transformer_z_image.py b/tests/models/transformers/test_models_transformer_z_image.py index 35bafa5702ae..1a6881c13126 100644 --- a/tests/models/transformers/test_models_transformer_z_image.py +++ b/tests/models/transformers/test_models_transformer_z_image.py @@ -18,6 +18,7 @@ import torch from diffusers import ZImageTransformer2DModel +from diffusers.loaders.lora_pipeline import ZImageLoraLoaderMixin from diffusers.utils.torch_utils import randn_tensor from ...testing_utils import assert_tensors_close, torch_device @@ -25,6 +26,7 @@ AutoRoundCompileTesterMixin, AutoRoundTesterMixin, BaseModelTesterConfig, + LoKrTesterMixin, LoraTesterMixin, MemoryTesterMixin, ModelTesterMixin, @@ -32,6 +34,7 @@ TorchCompileTesterMixin, TrainingTesterMixin, ) +from ..testing_utils.lokr import check_lokr_deltas, make_lokr_factors # Z-Image requires torch.use_deterministic_algorithms(False) due to complex64 RoPE operations @@ -183,6 +186,36 @@ class TestZImageTransformerLoRA(ZImageTransformerTesterConfig, LoraTesterMixin): """LoRA adapter tests for Z-Image Transformer.""" +class TestZImageTransformerLoKr(ZImageTransformerTesterConfig, LoKrTesterMixin): + """LoKr adapter tests for Z-Image Transformer, including the ai-toolkit Z-Image LoKr checkpoint format.""" + + @torch.no_grad() + def test_lokr_ai_toolkit_checkpoint(self): + # ai-toolkit stores the diffusers module paths under a `diffusion_model.` prefix. Full-matrix factors come with + # a placeholder alpha, where LoKr applies no scaling; rank-decomposed factors are scaled by `alpha / rank`. + torch.manual_seed(0) + model = self.model_class(**self.get_init_dict()).eval().to(torch_device) + rank = 1 + state_dict, expected_deltas = {}, {} + for module, factor_rank, alpha in [ + ("layers.0.attention.to_q", None, 9999220736.0), + ("layers.0.feed_forward.w1", None, 9999220736.0), + ("layers.0.adaLN_modulation.0", None, 9999220736.0), + ("layers.0.attention.to_v", rank, 0.5), + ]: + linear = model.get_submodule(module) + factors, delta = make_lokr_factors(linear.out_features, linear.in_features, rank=factor_rank) + state_dict.update({f"diffusion_model.{module}.{k}": v for k, v in factors.items()}) + state_dict[f"diffusion_model.{module}.alpha"] = torch.tensor(alpha) + expected_deltas[module] = delta if factor_rank is None else (alpha / rank) * delta + + converted = ZImageLoraLoaderMixin.lora_state_dict(state_dict) + assert all(k.startswith("transformer.") and ".lokr_" in k for k in converted) + model.load_lora_adapter(converted, prefix="transformer", adapter_name="default") + + check_lokr_deltas(model, expected_deltas) + + # TODO: Add pretrained_model_name_or_path once a tiny Z-Image model is available on the Hub # class TestZImageTransformerBitsAndBytes(ZImageTransformerTesterConfig, BitsAndBytesTesterMixin): # """BitsAndBytes quantization tests for Z-Image Transformer."""