Skip to content
Closed
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
35 changes: 32 additions & 3 deletions src/diffusers/loaders/lora_conversion_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -3124,7 +3124,8 @@ def _convert_non_diffusers_minimax_h3_lora_to_diffusers(state_dict):

Every known producer trains against the original checkpoint's module names — ai-toolkit under a `diffusion_model.`
prefix, the reference `generate.py` / ComfyUI checkpoints under no prefix at all, musubi-tuner under a flattened
`lora_unet_` one — so the prefix is optional and the module names are what identifies the format. Handles:
`lora_unet_` one, DiffSynth-Studio under no prefix but with peft's `.default.` infix — so the prefix is optional
and the module names are what identifies the format. Handles:

- `diffusion_model.` prefix removal, and bare `blocks.` / `token_refiner.` / `final_layer.` keys
- musubi-tuner's flattened `lora_unet_blocks_0_attn_qkv_proj` -> `blocks.0.attn.qkv_proj`
Expand All @@ -3139,6 +3140,16 @@ def _convert_non_diffusers_minimax_h3_lora_to_diffusers(state_dict):
"""
state_dict = {k.removeprefix("diffusion_model."): v for k, v in state_dict.items()}

# DiffSynth-Studio runs the raw checkpoint's per-head-interleaved fused QKV verbatim, where every other producer
# reorders it to `[q_all; k_all; v_all]` at load time as the reference implementation does. The two are
# indistinguishable by shape; what identifies DiffSynth is that it writes peft's `.default.` infix over these names.
is_diffsynth = any(".lora_A.default.weight" in k or ".lora_B.default.weight" in k for k in state_dict)
if is_diffsynth:
state_dict = {
k.replace(".lora_A.default.", ".lora_A.").replace(".lora_B.default.", ".lora_B."): v
for k, v in state_dict.items()
}

# musubi-tuner (kohya sd-scripts) writes one flat module name per key under a `lora_unet_` prefix, with every `.`
# collapsed to `_`. H3's own module names contain underscores (`qkv_proj`, `token_refiner`, `final_layer`,
# `video_out`), so the dot path is recovered by matching the whole flattened name against the module vocabulary
Expand Down Expand Up @@ -3230,8 +3241,8 @@ def pull(base):
source = f"{source_prefix}.{i}"
target = f"{target_prefix}.{i}"

# Fused qkv -> split to_q / to_k / to_v (shared down/lora_A, chunk up/lora_B in thirds). Both producers
# consume the fused rows as `[q_all; k_all; v_all]`, so no per-head de-interleave is involved.
# Fused qkv -> split to_q / to_k / to_v (shared down/lora_A, chunk up/lora_B in thirds), after the
# per-head de-interleave a DiffSynth file needs.
qkv = pull(f"{source}.attn.qkv_proj")
if qkv is not None:
down, up = qkv
Expand All @@ -3240,6 +3251,24 @@ def pull(base):
f"`{source}.attn.qkv_proj` has {up.shape[0]} output rows, which is not divisible by 3. "
"This is not a fused MiniMax-H3 QKV projection."
)
if is_diffsynth:
# `[head0: q k v, head1: q k v, ...]` -> `[q_all; k_all; v_all]`. Every released MiniMax-H3
# partition has a 128-wide head, which is what lets the head count be read off the row count.
head_dim = 128
if up.shape[0] % (3 * head_dim) != 0:
raise ValueError(
f"`{source}.attn.qkv_proj` has {up.shape[0]} output rows, which is not a multiple of "
f"3 * {head_dim}. This is not a per-head-interleaved MiniMax-H3 QKV projection."
)
num_heads = up.shape[0] // (3 * head_dim)
grouped = up.reshape(num_heads, 3 * head_dim, up.shape[1])
up = torch.cat(
[
head_slice.reshape(num_heads * head_dim, up.shape[1])
for head_slice in grouped.split(head_dim, dim=1)
],
dim=0,
)
up_q, up_k, up_v = torch.chunk(up, 3, dim=0)
for proj, up_proj in (("to_q", up_q), ("to_k", up_k), ("to_v", up_v)):
converted_state_dict[f"{target}.attn.{proj}.lora_A.weight"] = down.clone()
Expand Down
20 changes: 14 additions & 6 deletions src/diffusers/loaders/lora_pipeline.py
Original file line number Diff line number Diff line change
Expand Up @@ -6748,6 +6748,12 @@ class MiniMaxH3LoraLoaderMixin(LoraBaseMixin):
and degrades output), so routing is explicit: converted state dicts target `transformer.`; reach `transformer_ref`
with a `transformer_ref.`-prefixed file or `load_into_transformer_ref=True`.

DiffSynth-Studio LoRAs (e.g.
[DiffSynth-Studio/MiniMax-H3-LoRA-LineartAnime](https://huggingface.co/DiffSynth-Studio/MiniMax-H3-LoRA-LineartAnime))
are trained against the raw checkpoint's per-head-interleaved fused QKV and are de-interleaved on conversion. Their
fp32 factors make the unfused LoRA path compute in fp32; `.to(torch.bfloat16)` on the model after loading restores
the bf16 memory budget.

LoRAs trained against a pruned checkpoint (the `*_pruned_*` files in
[Comfy-Org/MiniMax-H3](https://huggingface.co/Comfy-Org/MiniMax-H3);
[joyfox/MiniMax-H3-Turbo](https://huggingface.co/joyfox/MiniMax-H3-Turbo) is one) fail with a size mismatch.
Expand Down Expand Up @@ -6815,20 +6821,22 @@ def lora_state_dict(
logger.warning(warn_msg)
state_dict = {k: v for k, v in state_dict.items() if "dora_scale" not in k}

# The original checkpoint's module names, none of which a diffusers name shares. Checked before the peft dump
# below, which DiffSynth-Studio's files would otherwise match: they carry `.default.` over these same names.
is_non_diffusers_format = any(
k.startswith(("diffusion_model.", "blocks.", "token_refiner.blocks.", "final_layer.", "lora_unet_"))
for k in state_dict
)
# A peft dump: diffusers module names carrying peft's `.default.` infix and no component prefix. The missing
# prefix is what keeps this from shadowing a file diffusers itself wrote.
is_unprefixed_diffusers_format = any(".default.weight" in k for k in state_dict) and not any(
k.startswith((f"{cls.transformer_name}.", f"{cls.transformer_ref_name}.")) for k in state_dict
)
if is_unprefixed_diffusers_format:
state_dict = {f"{cls.transformer_name}.{k.replace('.default.', '.')}": v for k, v in state_dict.items()}

is_non_diffusers_format = any(
k.startswith(("diffusion_model.", "blocks.", "token_refiner.", "final_layer.", "lora_unet_"))
for k in state_dict
)
if is_non_diffusers_format:
state_dict = _convert_non_diffusers_minimax_h3_lora_to_diffusers(state_dict)
elif is_unprefixed_diffusers_format:
state_dict = {f"{cls.transformer_name}.{k.replace('.default.', '.')}": v for k, v in state_dict.items()}

# Published H3 LoRAs are alpha-less and mixed-rank (64 on attention and FFN, 16 on the AdaLN projections), and
# `get_peft_kwargs` would scale one of the two rank groups by `alpha / r`, so `alpha == rank` is pinned below.
Expand Down
Loading