Spaces:
Running on Zero
Running on Zero
File size: 7,686 Bytes
da11654 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 | """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)
|