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)