multimodalart HF Staff commited on
Commit
d53b4b0
·
verified ·
1 Parent(s): 52aec0b

Fold larryvrh/MiniMax-H3-Turbo-Lora v4-600-EMA instead of the lightx2v 8-step file

Browse files
Files changed (1) hide show
  1. h3_turbo_lora.py +90 -46
h3_turbo_lora.py CHANGED
@@ -1,21 +1,31 @@
1
- """The lightx2v 8-step Turbo LoRA, folded into the same bf16 weights H3-World is merged into.
2
-
3
- This mirrors the `h3_lora.py` of the Spaces that already run this adapter on MiniMax-H3 —
4
- [`MiniMaxAI/MiniMax-H3-Turbo-Lora`](https://huggingface.co/spaces/MiniMaxAI/MiniMax-H3-Turbo-Lora)
5
- (its `lightx` set) and
6
- [`hugging-apps/minimax-h3-turbo-sla-demo`](https://huggingface.co/spaces/hugging-apps/minimax-h3-turbo-sla-demo):
7
-
8
- * `lightx2v/Minimax-h3-Turbo` is a **PEFT checkpoint against the diffusers module tree itself** —
9
- `transformer_blocks.N.attn.to_q.lora_A.default.weight` and friends — so unlike H3-World it needs
10
- no key conversion at all: 312 targets (50 transformer blocks + 2 token-refiner blocks x
11
- `to_q/to_k/to_v/to_out.0/ff.net.0.proj/ff.net.2`) map name-for-name.
12
- * Rank is 128 and the file's own safetensors metadata records `alpha: 8`, so the fold scale is
13
- `alpha / rank = 0.0625` — exactly what `set_adapters(weights=1.0)` applies in lightx2v's
14
- reference inference script, and what the Spaces above use for this repo.
15
- * `minimax_h3_fl2v_turbo_8step_v1.0_bf16.safetensors` is an 8-NFE distillation, so the step count
16
- is overridden to **8** when it is active. No scheduler swap and no CFG change: MiniMax-H3 is
17
- already guidance-distilled and its own `MiniMaxH3Scheduler` is what every one of those Spaces
18
- keeps using with the turbo LoRA folded.
 
 
 
 
 
 
 
 
 
 
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", "lightx2v/Minimax-h3-Turbo")
33
- TURBO_FILE = os.environ.get("H3_TURBO_FILE", "minimax_h3_fl2v_turbo_8step_v1.0_bf16.safetensors")
34
- # The distillation's own NFE count; the card offers 4 as the faster, softer alternative.
35
  TURBO_STEPS = int(os.environ.get("H3_TURBO_STEPS", "8"))
36
- # 0 = read `alpha` out of the file's safetensors metadata (it is `8` for every lightx2v file).
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.default.weight", ".lora_B.default.weight"
41
 
42
 
43
- def load() -> dict:
44
- """Download the adapter and return ``{"label", "scale", "entries"}``; entries are ``(param_key, A, B)``."""
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
45
  from huggingface_hub import hf_hub_download
46
- from safetensors import safe_open
47
  from safetensors.torch import load_file
48
 
49
- path = hf_hub_download(TURBO_REPO, TURBO_FILE)
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
- ranks = {lora[f"{name}{SUFFIX_A}"].shape[0] for name in bases}
66
- if len(ranks) != 1:
67
- raise ValueError(f"LoRA mixes ranks {sorted(ranks)}")
68
- rank = ranks.pop()
 
 
 
 
 
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 / rank * TURBO_STRENGTH,
74
  "entries": entries,
75
- "rank": rank,
76
- "alpha": alpha,
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 5 TFLOP in
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']} · {len(spec['entries'])} targets, "
120
- f"rank {spec['rank']}, alpha {spec['alpha']:g} -> scale {spec['scale']:.4f} · "
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