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
44 changes: 44 additions & 0 deletions src/diffusers/loaders/lora_conversion_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"),
Expand Down Expand Up @@ -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.",
Expand Down
139 changes: 138 additions & 1 deletion tests/pipelines/qwenimage21/test_qwenimage21.py
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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
Expand Down Expand Up @@ -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."""
Loading