diff --git a/src/diffusers/loaders/lora_conversion_utils.py b/src/diffusers/loaders/lora_conversion_utils.py index cef2e88454a8..cd6b796e917e 100644 --- a/src/diffusers/loaders/lora_conversion_utils.py +++ b/src/diffusers/loaders/lora_conversion_utils.py @@ -2275,6 +2275,9 @@ def convert_key(key: str) -> str: ("img", "mlp"), ("txt", "mod"), ("img", "mod"), + # Qwen-Image-2.1 SwiGLU: `img_mlp.gate_layer`, and the fused `img_mlp.gate_up` + ("gate", "layer"), + ("gate", "up"), # triplets ("add", "q", "proj"), ("add", "k", "proj"), @@ -2359,10 +2362,51 @@ def get_alpha_scales(down_weight, alpha_key): if len(state_dict) > 0: raise ValueError(f"`state_dict` should be empty at this point but has {state_dict.keys()=}") + converted_state_dict = _split_qwen_image21_fused_gate_up_lora(converted_state_dict) + converted_state_dict = {f"transformer.{k}": v for k, v in converted_state_dict.items()} return converted_state_dict +def _split_qwen_image21_fused_gate_up_lora(state_dict): + """ + Split LoRAs that target the fused SwiGLU `img_mlp.gate_up` projection of Qwen-Image-2.1 into its two diffusers + layers. + + ComfyUI (and trainers built on its layout, such as ai-toolkit) fuse `img_mlp.gate_layer` and `img_mlp.proj` into + one `img_mlp.gate_up` Linear whose output rows are `[gate; up]`, the same order + `convert_qwen_image21_transformer_checkpoint_to_diffusers` splits a full checkpoint in. Both halves read the same + input, so the fused `lora_A` is shared and `lora_B` is split along its rows: `B @ A` of the fused layer is exactly + `cat([B_gate @ A, B_up @ A])`, so the split is lossless. Any alpha has to be folded into the weights beforehand. + + Only Qwen-Image-2.1 has an `img_mlp.gate_up` module. The original Qwen-Image's `img_mlp` is a `FeedForward` + (`net.0.proj` / `net.2`), so its keys pass through unchanged. + """ + fused = ".img_mlp.gate_up." + converted_state_dict = {} + for key, value in state_dict.items(): + if fused not in key: + converted_state_dict[key] = value + continue + + gate_key = key.replace(fused, ".img_mlp.gate_layer.") + proj_key = key.replace(fused, ".img_mlp.proj.") + if ".lora_A." in key: + converted_state_dict[gate_key] = value + converted_state_dict[proj_key] = value.clone() + else: + if value.shape[0] % 2 != 0: + raise ValueError( + f"Expected the fused `img_mlp.gate_up` LoRA weight {key} to have an even number of output rows, " + f"but got shape {tuple(value.shape)}." + ) + gate, up = value.chunk(2, dim=0) + converted_state_dict[gate_key] = gate.clone() + converted_state_dict[proj_key] = up.clone() + + return converted_state_dict + + def _convert_non_diffusers_anima_lora_to_diffusers(state_dict): rename_dict = { "blocks.": "transformer_blocks.", diff --git a/tests/pipelines/qwenimage21/test_qwenimage21.py b/tests/pipelines/qwenimage21/test_qwenimage21.py index fd653da72ad4..12109fdc0329 100644 --- a/tests/pipelines/qwenimage21/test_qwenimage21.py +++ b/tests/pipelines/qwenimage21/test_qwenimage21.py @@ -32,8 +32,10 @@ QwenImage21Pipeline, QwenImage21Transformer2DModel, ) +from diffusers.utils import logging +from diffusers.utils.import_utils import is_peft_available -from ...testing_utils import assert_tensors_close +from ...testing_utils import CaptureLogger, assert_tensors_close, torch_device from ..testing_utils import ( BasePipelineTesterConfig, LoraMemoryTesterMixin, @@ -43,6 +45,10 @@ ) +if is_peft_available(): + from peft.tuners.tuners_utils import BaseTunerLayer + + # The pipeline hardcodes `vae_scale_factor = 16` and rounds height/width down to a multiple of 32, so 32 is the # smallest resolution that survives: a 2x2 latent, which is exactly one vision slot's worth of target tokens. IMAGE_SIZE = 32 @@ -276,6 +282,137 @@ class TestQwenImage21PipelineLoRA(QwenImage21PipelineTesterConfig, LoraTesterMix def test_simple_inference_with_text_denoiser_lora_and_scale(self, base_pipe_output): pass + # Modules a typical ai-toolkit / ComfyUI Qwen-Image-2.1 LoRA trains: attention, the fused SwiGLU + # `img_mlp.gate_up` (ComfyUI's `[gate_layer; proj]` row-concatenation) and `img_mlp.out`. + FUSED_LORA_MODULES = ("attn.to_q", "attn.to_out.0", "img_mlp.gate_up", "img_mlp.out") + + def get_fused_gate_up_lora(self, transformer, rank=4, alpha=None): + """ + Build a LoRA in the fused `img_mlp.gate_up` layout, plus the equivalent LoRA in the diffusers layout: + `gate_layer` and `proj` share the fused `lora_A` and each take their half of the fused `lora_B` rows, with + `alpha / rank` folded into `lora_A`. + """ + generator = torch.Generator("cpu").manual_seed(0) + scale = 1.0 if alpha is None else alpha / rank + fused, split = {}, {} + for i, block in enumerate(transformer.transformer_blocks): + for module_name in self.FUSED_LORA_MODULES: + if module_name == "img_mlp.gate_up": + in_features = block.img_mlp.gate_layer.in_features + out_features = block.img_mlp.gate_layer.out_features + block.img_mlp.proj.out_features + else: + module = block.get_submodule(module_name) + in_features, out_features = module.in_features, module.out_features + + lora_A = torch.randn(rank, in_features, generator=generator) + lora_B = torch.randn(out_features, rank, generator=generator) + fused[f"transformer_blocks.{i}.{module_name}"] = (lora_A, lora_B) + + prefix = f"transformer.transformer_blocks.{i}" + if module_name == "img_mlp.gate_up": + gate_B, proj_B = lora_B.chunk(2, dim=0) + for split_name, split_B in (("img_mlp.gate_layer", gate_B), ("img_mlp.proj", proj_B)): + split[f"{prefix}.{split_name}.lora_A.weight"] = lora_A * scale + split[f"{prefix}.{split_name}.lora_B.weight"] = split_B + else: + split[f"{prefix}.{module_name}.lora_A.weight"] = lora_A * scale + split[f"{prefix}.{module_name}.lora_B.weight"] = lora_B + return fused, split + + @staticmethod + def to_checkpoint_format(fused, layout, alpha=None): + """Serialize `{module: (lora_A, lora_B)}` the way each supported non-diffusers trainer does.""" + state_dict = {} + for module, (lora_A, lora_B) in fused.items(): + if layout == "ai-toolkit": + key, down, up = f"diffusion_model.{module}", "lora_A.weight", "lora_B.weight" + elif layout == "kohya": + key, down, up = f"lora_unet_{module.replace('.', '_')}", "lora_down.weight", "lora_up.weight" + elif layout == "diffsynth": + key, down, up = module, "lora_A.default.weight", "lora_B.default.weight" + else: + raise ValueError(layout) + state_dict[f"{key}.{down}"] = lora_A + state_dict[f"{key}.{up}"] = lora_B + if alpha is not None: + state_dict[f"{key}.alpha"] = torch.tensor(float(alpha)) + return state_dict + + @pytest.mark.parametrize("layout, alpha", [("ai-toolkit", None), ("kohya", 8), ("diffsynth", None)]) + def test_load_lora_weights_splits_fused_gate_up(self, base_pipe_output, layout, alpha): + """ + A LoRA on ComfyUI's fused `img_mlp.gate_up` must load into `img_mlp.gate_layer` and `img_mlp.proj` instead of + being dropped as unexpected keys, and produce the same output as the equivalent pre-split diffusers LoRA. + """ + pipe = self.get_pipeline().to(torch_device) + fused, split = self.get_fused_gate_up_lora(pipe.transformer, alpha=alpha) + state_dict = self.to_checkpoint_format(fused, layout, alpha=alpha) + + logger = logging.get_logger("diffusers.utils.peft_utils") + logger.setLevel(logging.WARNING) + with CaptureLogger(logger) as cap_logger: + pipe.load_lora_weights(state_dict, adapter_name="fused") + assert "unexpected keys" not in cap_logger.out + assert "missing keys" not in cap_logger.out + + for block in pipe.transformer.transformer_blocks: + for module in (block.img_mlp.gate_layer, block.img_mlp.proj, block.img_mlp.out, block.attn.to_q): + assert isinstance(module, BaseTunerLayer) + assert "fused" in module.lora_A + output_fused = self.run_pipe(pipe) + assert not torch.allclose(output_fused, base_pipe_output, atol=1e-3, rtol=1e-3) + + reference_pipe = self.get_pipeline().to(torch_device) + reference_pipe.load_lora_weights(split, adapter_name="split") + output_split = self.run_pipe(reference_pipe) + + assert_tensors_close(output_fused, output_split, atol=1e-5, rtol=1e-5) + + def test_lora_state_dict_fused_gate_up_split_is_lossless(self): + """`B @ A` of the fused layer is the row-concatenation of the `gate_layer` and `proj` deltas.""" + transformer = self.get_dummy_components()["transformer"] + fused, _ = self.get_fused_gate_up_lora(transformer) + converted = self.pipeline_class.lora_state_dict(self.to_checkpoint_format(fused, "ai-toolkit")) + + assert not any("gate_up" in k for k in converted) + for i in range(len(transformer.transformer_blocks)): + lora_A, lora_B = fused[f"transformer_blocks.{i}.img_mlp.gate_up"] + prefix = f"transformer.transformer_blocks.{i}.img_mlp" + split_delta = torch.cat( + [ + converted[f"{prefix}.{name}.lora_B.weight"] @ converted[f"{prefix}.{name}.lora_A.weight"] + for name in ("gate_layer", "proj") + ] + ) + assert_tensors_close(split_delta, lora_B @ lora_A, atol=1e-6, rtol=1e-6) + # The halves are copies, so the converted state dict can be serialized as is. + gate_A = converted[f"{prefix}.gate_layer.lora_A.weight"] + proj_A = converted[f"{prefix}.proj.lora_A.weight"] + assert gate_A.data_ptr() != proj_A.data_ptr() + + def test_fused_gate_up_lora_matches_fused_swiglu(self): + """ + Pins the row order independently of the converter: the loaded MLP must match ComfyUI's fused SwiGLU, which + computes `out(silu(gate) * up)` with `gate, up = gate_up(x).chunk(2, dim=-1)`, on the LoRA-updated fused weight. + """ + pipe = self.get_pipeline() + fused, _ = self.get_fused_gate_up_lora(pipe.transformer) + mlp = pipe.transformer.transformer_blocks[0].img_mlp + fused_weight = torch.cat([mlp.gate_layer.weight, mlp.proj.weight]).detach().clone() + out_weight = mlp.out.weight.detach().clone() + + pipe.load_lora_weights(self.to_checkpoint_format(fused, "ai-toolkit"), adapter_name="fused") + + lora_A, lora_B = fused["transformer_blocks.0.img_mlp.gate_up"] + out_A, out_B = fused["transformer_blocks.0.img_mlp.out"] + hidden_states = torch.randn(2, 3, fused_weight.shape[1], generator=torch.Generator("cpu").manual_seed(0)) + with torch.no_grad(): + gate, up = (hidden_states @ (fused_weight + lora_B @ lora_A).T).chunk(2, dim=-1) + expected = (torch.nn.functional.silu(gate) * up) @ (out_weight + out_B @ out_A).T + actual = mlp(hidden_states) + + assert_tensors_close(actual, expected, atol=1e-4, rtol=1e-4) + class TestQwenImage21PipelineLoRAMemory(QwenImage21PipelineTesterConfig, LoraMemoryTesterMixin): """LoRA x memory-optimization tests for the Qwen-Image 2.1 pipeline."""