multimodalart's picture
multimodalart HF Staff
MiniMax-H3 Wushu Action LoRA demo
da11654 verified
Raw
History Blame Contribute Delete
7.69 kB
"""Wushu Action LoRA support for the diffusers MiniMax-H3 transformer.
The Jojocodex wushu-action LoRA (`Jojocodex/minimax-h3-wushu-action-lora`) targets the ComfyUI reference module
tree — `diffusion_model.blocks.N.attn.qkv_proj`, `diffusion_model.blocks.N.mlp.fc1`, etc. — with rank 16 and no
alpha metadata (so scale = 1.0, matching the convention where alpha == rank). The shipped `_pruned` file carries
416 keys over 208 base modules (52 transformer blocks x {attn.qkv_proj, attn.out_proj, mlp.fc1, mlp.fc2}, the two
`token_refiner` blocks included) and has its `adaln_proj` rows removed (`__metadata__: {"adaln_pruned": "true"}`),
which is what makes it stackable with the Turbo acceleration LoRA.
The LoRA is applied by folding `scale * (lora_B @ lora_A)` into the bf16 weights rather than as runtime wrappers:
folding costs nothing per request and keeps the transformer a plain `nn.Module` for the ZeroGPU startup packing.
Deltas are computed in float32 and round once on the way back into bf16.
The diffusers conversion transforms:
- strip the `diffusion_model.` prefix the ai-toolkit export uses
- fused-QKV row thirds onto `attn.to_q/k/v`
- `SwiGLU` gate/value swap onto `ff.net.0.proj`
- `fc2` -> `ff.net.2`, `attn.out_proj` -> `attn.to_out.0`
- `blocks.` -> `transformer_blocks.`
- `token_refiner.blocks.` -> `token_refiner.refiner_blocks.`
- `final_layer.adaln_proj.linear` -> `norm_out.linear` (not present in the pruned file)
"""
from __future__ import annotations
import os
import torch
# --- Wushu Action LoRA config ---
# The model card tells ComfyUI users to take the `_pruned` file; it is also the Turbo-compatible one.
WUSHU_LORA_REPO = os.environ.get("WUSHU_LORA_REPO", "Jojocodex/minimax-h3-wushu-action-lora")
WUSHU_LORA_FILE = os.environ.get("WUSHU_LORA_FILE", "wushu_action_h3_lora_v4_2000_pruned.safetensors")
# The card recommends strength 0.8~1.0 for the ComfyUI LoraLoader; folded at the top of that range.
WUSHU_LORA_STRENGTH = float(os.environ.get("WUSHU_LORA_STRENGTH", "1.0"))
# --- Turbo LoRA config (optional acceleration) ---
TURBO_LORA_REPO = os.environ.get("TURBO_LORA_REPO", "Comfy-Org/MiniMax-H3")
TURBO_LORA_FILE = os.environ.get(
"TURBO_LORA_FILE", "loras/minimax_h3_fl2v_turbo_4step_v1.0_768p_comfyui_bf16.safetensors"
)
TURBO_LORA_STRENGTH = float(os.environ.get("TURBO_LORA_STRENGTH", "1.0"))
USE_TURBO = os.environ.get("USE_TURBO", "1").lower() not in ("0", "off", "none", "false")
def _targets(name: str, b: torch.Tensor, inner_dim: int) -> list[tuple[str, torch.Tensor]]:
"""Map one reference-tree base name and its `lora_B` onto diffusers parameter key + row-transformed B."""
name = name.removeprefix("diffusion_model.")
if name.startswith("token_refiner.blocks."):
target = name.replace("token_refiner.blocks.", "token_refiner.refiner_blocks.", 1)
elif name.startswith("blocks."):
target = name.replace("blocks.", "transformer_blocks.", 1)
else:
target = name
target = target.replace("final_layer.adaln_proj.linear", "norm_out.linear")
if target.endswith(".attn.qkv_proj"):
prefix = target.removesuffix("qkv_proj")
return [
(f"{prefix}to_{kind}.weight", part.contiguous())
for kind, part in zip(("q", "k", "v"), b.split(inner_dim, dim=0))
]
if target.endswith(".mlp.fc1"):
gate, value = b.chunk(2, dim=0)
return [(target.replace(".mlp.fc1", ".ff.net.0.proj") + ".weight", torch.cat([value, gate]).contiguous())]
if target.endswith(".mlp.fc2"):
return [(target.replace(".mlp.fc2", ".ff.net.2") + ".weight", b)]
if target.endswith(".attn.out_proj"):
return [(target.replace(".attn.out_proj", ".attn.to_out.0") + ".weight", b)]
# `adaln_proj.linear` (block-level and the final `norm_out.linear`): identical row layout on both sides.
return [(target + ".weight", b)]
def _load_wushu_lora(inner_dim: int) -> dict:
"""Load the Jojocodex wushu action LoRA from the Hub."""
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
lora = load_file(hf_hub_download(WUSHU_LORA_REPO, WUSHU_LORA_FILE))
bases = sorted({key.rsplit(".lora_", 1)[0] for key in lora})
entries = []
for name in bases:
a = lora[f"{name}.lora_A.weight"]
b = lora[f"{name}.lora_B.weight"]
entries.extend((key, a, b_part) for key, b_part in _targets(name, b, inner_dim))
return {
"label": f"{WUSHU_LORA_REPO}/{WUSHU_LORA_FILE}",
"scale": WUSHU_LORA_STRENGTH, # alpha == rank, so the base scale is 1
"entries": entries,
}
def _load_turbo_lora(inner_dim: int) -> dict:
"""Load the MiniMax-H3 Turbo LoRA from Comfy-Org for 4-step accelerated inference.
The Comfy-Org Turbo LoRA uses the kohya format with explicit `.alpha` keys per LoRA layer, so each entry's
scale is `alpha / rank`, applied to `B` before the fused-QKV split.
"""
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
lora = load_file(hf_hub_download(TURBO_LORA_REPO, TURBO_LORA_FILE))
all_keys = list(lora.keys())
bases = sorted({key.rsplit(".lora_", 1)[0] if ".lora_" in key else key.rsplit(".alpha", 1)[0] for key in all_keys})
entries = []
for name in bases:
a = lora[f"{name}.lora_A.weight"]
b = lora[f"{name}.lora_B.weight"]
alpha_key = f"{name}.alpha"
if alpha_key in lora:
scale = float(lora[alpha_key]) / a.shape[0]
else:
scale = 1.0
b_scaled = b * (scale * TURBO_LORA_STRENGTH)
entries.extend((key, a, b_part) for key, b_part in _targets(name, b_scaled, inner_dim))
return {
"label": f"{TURBO_LORA_REPO}/{TURBO_LORA_FILE}",
"scale": 1.0, # scale already applied per-entry above
"entries": entries,
}
def _apply(entries, params, sign: float) -> None:
"""Fold (sign * scale * (B @ A)) into each target parameter in place."""
for key, a, b in entries:
param = params.get(key)
if param is None:
raise KeyError(f"LoRA target `{key}` not found in the transformer")
delta = sign * (b.to(torch.float32) @ a.to(torch.float32))
param.data = (param.data.float() + delta.to(param.device)).to(param.dtype)
def apply_lora(transformer) -> str | None:
"""Fold the wushu action LoRA — and, when enabled, the Turbo LoRA — into the transformer weights.
Returns a status line, or `None` when nothing could be folded.
"""
inner_dim = transformer.config.num_attention_heads * transformer.config.attention_head_dim
params = dict(transformer.named_parameters())
loaded = []
try:
wushu = _load_wushu_lora(inner_dim)
_apply(wushu["entries"], params, wushu["scale"])
loaded.append(f"wushu-action ({wushu['label']}, {len(wushu['entries'])} weights, scale={wushu['scale']})")
print(f"[lora] wushu action LoRA folded: {len(wushu['entries'])} weights", flush=True)
except Exception as error:
print(f"[lora] WARNING: failed to load the wushu action LoRA: {error}", flush=True)
if USE_TURBO:
try:
turbo = _load_turbo_lora(inner_dim)
_apply(turbo["entries"], params, turbo["scale"])
loaded.append(f"turbo ({turbo['label']}, {len(turbo['entries'])} weights)")
print(f"[lora] turbo LoRA folded: {len(turbo['entries'])} weights", flush=True)
except Exception as error:
print(f"[lora] WARNING: failed to load the turbo LoRA: {error}", flush=True)
if not loaded:
return None
return "LoRAs folded: " + " + ".join(loaded)