Spaces:
Running on Zero
Running on Zero
Fold larryvrh/MiniMax-H3-Turbo-Lora v4-600-EMA instead of the lightx2v 8-step file
Browse files- h3_turbo_lora.py +90 -46
h3_turbo_lora.py
CHANGED
|
@@ -1,21 +1,31 @@
|
|
| 1 |
-
"""The
|
| 2 |
-
|
| 3 |
-
|
| 4 |
-
|
| 5 |
-
|
| 6 |
-
|
| 7 |
-
|
| 8 |
-
|
| 9 |
-
|
| 10 |
-
|
| 11 |
-
|
| 12 |
-
|
| 13 |
-
|
| 14 |
-
|
| 15 |
-
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 19 |
|
| 20 |
The adapter is *folded* into the weights rather than wrapped as a runtime module, for the same
|
| 21 |
reason H3-World is: this Space patches `MiniMaxH3AttnProcessor` and drives the transformer's live
|
|
@@ -29,29 +39,59 @@ import os
|
|
| 29 |
|
| 30 |
import torch
|
| 31 |
|
| 32 |
-
TURBO_REPO = os.environ.get("H3_TURBO_REPO", "
|
| 33 |
-
TURBO_FILE = os.environ.get("H3_TURBO_FILE", "
|
| 34 |
-
# The
|
| 35 |
TURBO_STEPS = int(os.environ.get("H3_TURBO_STEPS", "8"))
|
| 36 |
-
#
|
| 37 |
-
TURBO_ALPHA = float(os.environ.get("H3_TURBO_ALPHA", "0"))
|
| 38 |
TURBO_STRENGTH = float(os.environ.get("H3_TURBO_STRENGTH", "1.0"))
|
| 39 |
|
| 40 |
-
SUFFIX_A, SUFFIX_B = ".lora_A.
|
| 41 |
|
| 42 |
|
| 43 |
-
def
|
| 44 |
-
"""
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 45 |
from huggingface_hub import hf_hub_download
|
| 46 |
-
from safetensors import safe_open
|
| 47 |
from safetensors.torch import load_file
|
| 48 |
|
| 49 |
-
|
| 50 |
-
with safe_open(path, framework="pt") as handle:
|
| 51 |
-
metadata = handle.metadata() or {}
|
| 52 |
-
alpha = TURBO_ALPHA or float(metadata.get("alpha", 8))
|
| 53 |
-
|
| 54 |
-
lora = load_file(path)
|
| 55 |
bases = sorted({key[: -len(SUFFIX_A)] for key in lora if key.endswith(SUFFIX_A)})
|
| 56 |
if not bases:
|
| 57 |
raise ValueError(f"No lora_A/lora_B pairs found in {TURBO_FILE}")
|
|
@@ -62,25 +102,29 @@ def load() -> dict:
|
|
| 62 |
if f"{name}{SUFFIX_B}" not in lora:
|
| 63 |
raise ValueError(f"LoRA is missing the lora_B twin of {name}{SUFFIX_A}")
|
| 64 |
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 69 |
|
| 70 |
-
entries = [(f"{name}.weight", lora[f"{name}{SUFFIX_A}"], lora[f"{name}{SUFFIX_B}"]) for name in bases]
|
| 71 |
return {
|
| 72 |
"label": f"{TURBO_REPO}/{TURBO_FILE}",
|
| 73 |
-
"scale": alpha
|
| 74 |
"entries": entries,
|
| 75 |
-
"
|
| 76 |
-
"
|
| 77 |
}
|
| 78 |
|
| 79 |
|
| 80 |
def _fold(transformer, spec: dict, sign: float) -> None:
|
| 81 |
"""Add ``sign * scale * (B @ A)`` to every target weight, in place.
|
| 82 |
|
| 83 |
-
The factors ride to the weight's own device before the matmul — the deltas are
|
| 84 |
total, which is seconds on the card and minutes on a Space's two vCPUs.
|
| 85 |
"""
|
| 86 |
params = dict(transformer.named_parameters())
|
|
@@ -100,7 +144,7 @@ def _fold(transformer, spec: dict, sign: float) -> None:
|
|
| 100 |
|
| 101 |
def prepare(transformer) -> str:
|
| 102 |
"""Load the adapter and validate it against ``transformer``, without folding it yet."""
|
| 103 |
-
spec = load()
|
| 104 |
params = dict(transformer.named_parameters())
|
| 105 |
missed = [key for key, _, _ in spec["entries"] if key not in params]
|
| 106 |
if missed:
|
|
@@ -116,9 +160,9 @@ def prepare(transformer) -> str:
|
|
| 116 |
)
|
| 117 |
transformer._turbo_state = {"active": False, "spec": spec}
|
| 118 |
return (
|
| 119 |
-
f"Turbo LoRA ready · {spec['label']} · {
|
| 120 |
-
f"
|
| 121 |
-
f"{TURBO_STEPS} steps when active"
|
| 122 |
)
|
| 123 |
|
| 124 |
|
|
|
|
| 1 |
+
"""The `larryvrh/MiniMax-H3-Turbo-Lora` few-step LoRA, folded into the same bf16 weights H3-World is merged into.
|
| 2 |
+
|
| 3 |
+
`minimax_h3_turbo_v4_step600_ema.safetensors` is the repo's recommended checkpoint (its `v4`
|
| 4 |
+
line, step 600, EMA). Everything below comes out of the file itself and the model card:
|
| 5 |
+
|
| 6 |
+
* It is a **ComfyUI-side checkpoint in the original MiniMax-H3 module naming** —
|
| 7 |
+
`blocks.N.attn.qkv_proj.lora_A.weight`, `token_refiner.blocks.N.mlp.fc1`,
|
| 8 |
+
`blocks.N.adaln_proj.linear`, `final_layer.adaln_proj.linear` — so, exactly like H3-World, the
|
| 9 |
+
keys have to be replayed through `convert_minimax_h3_to_diffusers.py`'s renames before the
|
| 10 |
+
deltas mean anything against the diffusers port: 259 LoRA pairs become 363 weight deltas.
|
| 11 |
+
|
| 12 |
+
One transform differs from `app.py`'s H3-World merge and it matters: the LoRA was trained
|
| 13 |
+
against `comfy.ldm.minimax.model`, whose `Attention.forward` reads
|
| 14 |
+
`self.qkv_proj(x).split(heads * head_dim, dim=-1)` — i.e. `Comfy-Org/MiniMax-H3` stores the
|
| 15 |
+
fused QKV as contiguous `[q_all; k_all; v_all]`, the reference's *in-memory* layout, not the raw
|
| 16 |
+
release shards' per-head interleave. So these rows are split into contiguous thirds and **not**
|
| 17 |
+
de-interleaved. The `mlp.fc1` `[gate; value]` -> `[value; gate]` swap diffusers' `SwiGLU` needs
|
| 18 |
+
still applies (comfy's swiglu also reads `[gate; value]`), and `adaln_proj.linear` /
|
| 19 |
+
`final_layer.adaln_proj.linear` are pure renames — the conversion script permutes no rows there.
|
| 20 |
+
* The file's own safetensors metadata records `application: "W_eff = W + lora_B @ lora_A"` and the
|
| 21 |
+
card states *"applied as a plain low-rank update, alpha = rank, so no extra scaling"* — hence a
|
| 22 |
+
fold scale of exactly **1.0**, with `H3_TURBO_STRENGTH` left as the card's strength dial (nudge
|
| 23 |
+
to ~1.05-1.2 for smear, ~0.8-0.95 for over-sharp grain). Ranks are mixed on purpose: 64 for the
|
| 24 |
+
attention/FFN projections, 16 for the AdaLN projections.
|
| 25 |
+
* The card's useful step range is **4-8**, with 6-8 recommended and no benefit past 8, so the step
|
| 26 |
+
count is overridden to **8 NFE** when the adapter is active. No scheduler swap and no CFG
|
| 27 |
+
change: MiniMax-H3 is already guidance-distilled and its own `MiniMaxH3Scheduler` is what this
|
| 28 |
+
Space keeps using with the turbo LoRA folded.
|
| 29 |
|
| 30 |
The adapter is *folded* into the weights rather than wrapped as a runtime module, for the same
|
| 31 |
reason H3-World is: this Space patches `MiniMaxH3AttnProcessor` and drives the transformer's live
|
|
|
|
| 39 |
|
| 40 |
import torch
|
| 41 |
|
| 42 |
+
TURBO_REPO = os.environ.get("H3_TURBO_REPO", "larryvrh/MiniMax-H3-Turbo-Lora")
|
| 43 |
+
TURBO_FILE = os.environ.get("H3_TURBO_FILE", "minimax_h3_turbo_v4_step600_ema.safetensors")
|
| 44 |
+
# The card's step range is 4-8 (6-8 recommended, nothing gained past 8).
|
| 45 |
TURBO_STEPS = int(os.environ.get("H3_TURBO_STEPS", "8"))
|
| 46 |
+
# alpha == rank in this checkpoint, so the update is applied as-is; this is the card's strength dial.
|
|
|
|
| 47 |
TURBO_STRENGTH = float(os.environ.get("H3_TURBO_STRENGTH", "1.0"))
|
| 48 |
|
| 49 |
+
SUFFIX_A, SUFFIX_B = ".lora_A.weight", ".lora_B.weight"
|
| 50 |
|
| 51 |
|
| 52 |
+
def _target_name(source_name: str) -> str:
|
| 53 |
+
"""Rename one original-layout module path onto its diffusers module path."""
|
| 54 |
+
if source_name.startswith("token_refiner.blocks."):
|
| 55 |
+
target = source_name.replace("token_refiner.blocks.", "token_refiner.refiner_blocks.", 1)
|
| 56 |
+
elif source_name.startswith("blocks."):
|
| 57 |
+
target = source_name.replace("blocks.", "transformer_blocks.", 1)
|
| 58 |
+
else:
|
| 59 |
+
target = source_name
|
| 60 |
+
return target.replace("final_layer.adaln_proj.linear", "norm_out.linear")
|
| 61 |
+
|
| 62 |
+
|
| 63 |
+
def _targets(name: str, b_weight, inner_dim: int):
|
| 64 |
+
"""Yield ``(diffusers_param_key, row_transformed_B)`` for one original LoRA base name."""
|
| 65 |
+
target = _target_name(name)
|
| 66 |
+
|
| 67 |
+
if target.endswith(".attn.qkv_proj"):
|
| 68 |
+
prefix = target.removesuffix("qkv_proj")
|
| 69 |
+
if b_weight.shape[0] != 3 * inner_dim:
|
| 70 |
+
raise ValueError(
|
| 71 |
+
f"{name} lora_B has {b_weight.shape[0]} rows, expected 3 x {inner_dim} for a fused QKV"
|
| 72 |
+
)
|
| 73 |
+
# Contiguous thirds — the comfy layout this adapter was trained against already holds
|
| 74 |
+
# `[q_all; k_all; v_all]`, so there is no per-head de-interleave here (see the module docstring).
|
| 75 |
+
for kind, part in zip(("q", "k", "v"), b_weight.split(inner_dim, dim=0)):
|
| 76 |
+
yield f"{prefix}to_{kind}.weight", part.contiguous()
|
| 77 |
+
elif target.endswith(".mlp.fc1"):
|
| 78 |
+
# SwiGLU gate/value swap: the checkpoint stores [gate, value]; diffusers stores [value, gate].
|
| 79 |
+
gate, value = b_weight.chunk(2, dim=0)
|
| 80 |
+
yield target.replace(".mlp.fc1", ".ff.net.0.proj") + ".weight", torch.cat([value, gate]).contiguous()
|
| 81 |
+
elif target.endswith(".mlp.fc2"):
|
| 82 |
+
yield target.replace(".mlp.fc2", ".ff.net.2") + ".weight", b_weight
|
| 83 |
+
elif target.endswith(".attn.out_proj"):
|
| 84 |
+
yield target.replace(".attn.out_proj", ".attn.to_out.0") + ".weight", b_weight
|
| 85 |
+
else: # `adaln_proj.linear`, `norm_out.linear` — pure renames
|
| 86 |
+
yield target + ".weight", b_weight
|
| 87 |
+
|
| 88 |
+
|
| 89 |
+
def load(inner_dim: int) -> dict:
|
| 90 |
+
"""Download the adapter and return ``{"label", "scale", "entries", ...}``; entries are ``(param_key, A, B)``."""
|
| 91 |
from huggingface_hub import hf_hub_download
|
|
|
|
| 92 |
from safetensors.torch import load_file
|
| 93 |
|
| 94 |
+
lora = load_file(hf_hub_download(TURBO_REPO, TURBO_FILE))
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 95 |
bases = sorted({key[: -len(SUFFIX_A)] for key in lora if key.endswith(SUFFIX_A)})
|
| 96 |
if not bases:
|
| 97 |
raise ValueError(f"No lora_A/lora_B pairs found in {TURBO_FILE}")
|
|
|
|
| 102 |
if f"{name}{SUFFIX_B}" not in lora:
|
| 103 |
raise ValueError(f"LoRA is missing the lora_B twin of {name}{SUFFIX_A}")
|
| 104 |
|
| 105 |
+
# Ranks are mixed by design (64 on attention/FFN, 16 on AdaLN) and alpha == rank for every one of
|
| 106 |
+
# them, so the fold scale is rank-independent.
|
| 107 |
+
ranks = sorted({lora[f"{name}{SUFFIX_A}"].shape[0] for name in bases})
|
| 108 |
+
|
| 109 |
+
entries = []
|
| 110 |
+
for name in bases:
|
| 111 |
+
a = lora[f"{name}{SUFFIX_A}"]
|
| 112 |
+
for key, b_part in _targets(name, lora[f"{name}{SUFFIX_B}"], inner_dim):
|
| 113 |
+
entries.append((key, a, b_part))
|
| 114 |
|
|
|
|
| 115 |
return {
|
| 116 |
"label": f"{TURBO_REPO}/{TURBO_FILE}",
|
| 117 |
+
"scale": TURBO_STRENGTH, # alpha == rank -> the update is applied as-is
|
| 118 |
"entries": entries,
|
| 119 |
+
"targets": len(bases),
|
| 120 |
+
"ranks": ranks,
|
| 121 |
}
|
| 122 |
|
| 123 |
|
| 124 |
def _fold(transformer, spec: dict, sign: float) -> None:
|
| 125 |
"""Add ``sign * scale * (B @ A)`` to every target weight, in place.
|
| 126 |
|
| 127 |
+
The factors ride to the weight's own device before the matmul — the deltas are a few TFLOP in
|
| 128 |
total, which is seconds on the card and minutes on a Space's two vCPUs.
|
| 129 |
"""
|
| 130 |
params = dict(transformer.named_parameters())
|
|
|
|
| 144 |
|
| 145 |
def prepare(transformer) -> str:
|
| 146 |
"""Load the adapter and validate it against ``transformer``, without folding it yet."""
|
| 147 |
+
spec = load(transformer.config.num_attention_heads * transformer.config.attention_head_dim)
|
| 148 |
params = dict(transformer.named_parameters())
|
| 149 |
missed = [key for key, _, _ in spec["entries"] if key not in params]
|
| 150 |
if missed:
|
|
|
|
| 160 |
)
|
| 161 |
transformer._turbo_state = {"active": False, "spec": spec}
|
| 162 |
return (
|
| 163 |
+
f"Turbo LoRA ready · {spec['label']} · {spec['targets']} targets -> "
|
| 164 |
+
f"{len(spec['entries'])} weight deltas, rank {'/'.join(str(r) for r in spec['ranks'])}, "
|
| 165 |
+
f"scale {spec['scale']:.4g} · {TURBO_STEPS} steps when active"
|
| 166 |
)
|
| 167 |
|
| 168 |
|