MiniMax-H3-Turbo-Loras / lora_stack.py
Alexander
Fold the LoRA rack once it settles, and book GPU time from measurements
005fad1
Raw
History Blame Contribute Delete
17.2 kB
"""The LoRA stack: a catalog, arbitrary URLs, and the adapters they turn into on the transformer.
diffusers carries the MiniMax-H3 LoRA loader itself (`MiniMaxH3LoraLoaderMixin`), including the conversion of the
ComfyUI / musubi layouts onto `MiniMaxH3Transformer3DModel`'s module names, so a slot is one `load_lora_weights`
call and the whole stack is one `set_adapters`. Alpha-less files load at `alpha == rank`; a `__metadata__` alpha is
honored instead. Weights are per slot and change for free - nothing is folded into the checkpoint.
A slot resolves in the main process (`LoraSpec.fetch`, which downloads) and is applied inside the GPU worker
(`apply`), so no file is ever pulled off the network on booked GPU time.
"""
from __future__ import annotations
import os
import re
import tempfile
import urllib.request
from dataclasses import dataclass, replace
NONE_CHOICE = "None"
CUSTOM_CHOICE = "Custom URL"
CACHE_DIR = os.path.join(tempfile.gettempdir(), "h3-lora-cache")
# How many adapters may stay resident on the card before the ones outside the current stack are dropped.
MAX_RESIDENT = int(os.getenv("H3_MAX_ADAPTERS", "8"))
# Identical requests in a row before the rack is folded into the weights. 1 folds immediately, 0 never folds — the
# fold makes every later request with that rack free, and makes the next change to it cost one unfold plus one fold.
FUSE_AFTER = int(os.getenv("H3_FUSE_AFTER", "2"))
# Every entry was checked against the released `transformer/` partition. LoRAs trained on a *pruned* checkpoint
# (joyfox/MiniMax-H3-Turbo is the known one) fail with a size mismatch and are deliberately not listed; the same
# repositories' `*_full*` files are, where they exist.
CATALOG: dict[str, dict] = {
"Turbo 4-step · larryvrh": {
"repo": "larryvrh/MiniMax-H3-Turbo-Lora",
"weights": "minimax_h3_turbo_4step_ema_ckpt850.safetensors",
"adapter": "turbo_larry",
"default_weight": 1.0,
"suggested_steps": 6,
"url": "https://huggingface.co/larryvrh/MiniMax-H3-Turbo-Lora",
"about": "4-step distillation of the joint video+audio path. The usual starting point.",
"hint": "Drop steps to 4-6. Above ~1.2 it sharpens into artifacts, below ~0.8 it goes soft.",
},
"Turbo 4-step · lightx2v": {
"repo": "lightx2v/Minimax-h3-Turbo",
"weights": "minimax_h3_fl2v_turbo_4step_v1.0_768p_bf16.safetensors",
"adapter": "turbo_lightx_4",
"default_weight": 1.0,
"suggested_steps": 5,
"url": "https://huggingface.co/lightx2v/Minimax-h3-Turbo",
"about": "Alternative 4-step distillation, trained at 768p on the first/last-frame branch.",
"hint": "Trained for keyframed requests; holds up on text-only ones too.",
},
"Turbo 8-step · lightx2v": {
"repo": "lightx2v/Minimax-h3-Turbo",
"weights": "minimax_h3_fl2v_turbo_8step_v1.0_bf16.safetensors",
"adapter": "turbo_lightx_8",
"default_weight": 1.0,
"suggested_steps": 8,
"url": "https://huggingface.co/lightx2v/Minimax-h3-Turbo",
"about": "8-step sibling: slower than the 4-step files, steadier on motion.",
"hint": "Give it 8 steps - it is not distilled down to 4.",
},
"Turbo 4-step · InstantX": {
"repo": "InstantX/MiniMax-H3-Turbo-Lora-Diffusers",
"weights": "minimax_h3_turbo_4step_ema_ckpt500_diffusers.safetensors",
"adapter": "turbo_instantx",
"default_weight": 1.0,
"suggested_steps": 5,
"url": "https://huggingface.co/InstantX/MiniMax-H3-Turbo-Lora-Diffusers",
"about": "The ckpt500 EMA turbo LoRA, already converted to diffusers module names.",
"hint": "An earlier checkpoint of the larryvrh round - useful for A/B against it.",
},
"Realism · People": {
"repo": "fal/MiniMax-H3-Realism-People-LoRA",
"weights": "h3-realism-people-t2v-i2v-r2v.safetensors",
"adapter": "realism_people",
"default_weight": 0.8,
"url": "https://huggingface.co/fal/MiniMax-H3-Realism-People-LoRA",
"about": "Photographic skin, hair and lighting on human subjects.",
"hint": "Stacks well under a turbo LoRA. Back off toward 0.6 if faces start looking waxy.",
},
"Motion Adapter": {
"repo": "MATLOWAI/MiniMax-H3-Motion-Adapter",
"weights": "minimax_h3_motion_adapter_pilot_r16.safetensors",
"adapter": "motion_adapter",
"default_weight": 0.8,
"url": "https://huggingface.co/MATLOWAI/MiniMax-H3-Motion-Adapter",
"about": "Pushes the amount of movement in the shot up.",
"hint": "Rank 16 pilot - subtle by design. Raise it before you blame the prompt.",
},
"Spatial Physics": {
"repo": "Jojocodex/minimax-h3-spatial-physics-lora",
"weights": "wushu_spatial_physics_v2_full.safetensors",
"adapter": "spatial_physics",
"default_weight": 0.8,
"url": "https://huggingface.co/Jojocodex/minimax-h3-spatial-physics-lora",
"about": "Steadier object motion, contact and weight.",
"hint": "Pairs with Wushu Action for fight choreography that keeps its footing.",
},
"Wushu Action": {
"repo": "Jojocodex/minimax-h3-wushu-action-lora",
"weights": "wushu_action_h3_lora_v4_2000_full.safetensors",
"adapter": "wushu_action",
"default_weight": 0.8,
"url": "https://huggingface.co/Jojocodex/minimax-h3-wushu-action-lora",
"about": "Martial-arts choreography: strikes, spins, weapon work.",
"hint": "Name the move in the prompt ('spinning hook kick'), not just the style.",
},
"Lineart Anime": {
"repo": "DiffSynth-Studio/MiniMax-H3-LoRA-LineartAnime",
"weights": "model.safetensors",
"adapter": "lineart_anime",
"default_weight": 1.0,
"url": "https://huggingface.co/DiffSynth-Studio/MiniMax-H3-LoRA-LineartAnime",
"about": "Flat-shaded lineart anime, trained against the raw fused-QKV checkpoint.",
"hint": "Its factors are fp32; the transformer is cast back to bf16 after loading.",
},
}
CHOICES = [NONE_CHOICE, *CATALOG, CUSTOM_CHOICE]
def default_weight(key: str) -> float:
return CATALOG.get(key, {}).get("default_weight", 1.0)
def suggested_steps(keys) -> int | None:
"""The step count the fastest turbo LoRA in a stack was distilled for, if any."""
steps = [CATALOG[key]["suggested_steps"] for key in keys if key in CATALOG and "suggested_steps" in CATALOG[key]]
return min(steps) if steps else None
def about(key: str) -> tuple[str, str, str]:
"""`(about, hint, url)` of a catalog entry, empty strings for anything else."""
entry = CATALOG.get(key, {})
return entry.get("about", ""), entry.get("hint", ""), entry.get("url", "")
# ---------------------------------------------------------------------------
# Sources
# ---------------------------------------------------------------------------
def parse_source(raw: str) -> tuple[str, str, str | None]:
"""Any LoRA reference into `(kind, source, weight_name)`, `kind` being hf / civitai / direct."""
raw = raw.strip()
match = re.match(r"https?://huggingface\.co/([^/]+/[^/?#]+)(?:/(?:blob|resolve)/[^/]+/([^?#]+))?", raw)
if match:
return "hf", match.group(1), match.group(2)
if "civitai.com" in raw:
return "civitai", raw, None
if raw.startswith("http"):
return "direct", raw, None
# `user/repo` or `user/repo/inside/file.safetensors`
if re.match(r"^[A-Za-z0-9_.-]+/[A-Za-z0-9_.-]", raw):
owner, name, *rest = raw.split("/", 2)
return "hf", f"{owner}/{name}", rest[0] if rest else None
raise ValueError(f"Not a LoRA source: {raw!r}")
def _civitai_download_url(url: str) -> str:
"""Any CivitAI URL into a direct download URL."""
if "/api/download/models/" in url:
return url
match = re.search(r"[?&]modelVersionId=(\d+)", url)
if not match:
raise ValueError(
"A CivitAI URL needs its modelVersionId, e.g. https://civitai.com/models/123?modelVersionId=456"
)
token = os.getenv("CIVITAI_API_KEY", "")
url = f"https://civitai.com/api/download/models/{match.group(1)}"
return f"{url}?token={token}" if token else url
def _download(url: str) -> str:
"""Fetch a LoRA file to the local cache once, and return its path."""
os.makedirs(CACHE_DIR, exist_ok=True)
name = url.split("?")[0].rstrip("/").split("/")[-1]
if not name.endswith((".safetensors", ".sft")):
name = re.sub(r"[^A-Za-z0-9]", "_", url)[-64:] + ".safetensors"
path = os.path.join(CACHE_DIR, name)
if not os.path.exists(path):
request = urllib.request.Request(url, headers={"User-Agent": "minimax-h3-turbo-loras"})
with urllib.request.urlopen(request, timeout=300) as response, open(path + ".part", "wb") as handle:
while chunk := response.read(1 << 22):
handle.write(chunk)
os.replace(path + ".part", path)
return path
@dataclass(frozen=True)
class LoraSpec:
"""One filled slot, on its way to the GPU worker. Picklable on purpose - it crosses a process boundary."""
label: str
weight: float
adapter: str
source: str # a hub id, or a URL
weights_file: str | None = None
path: str | None = None # filled by `fetch`
def fetch(self) -> "LoraSpec":
"""Put the file on local disk. Runs in the main process, off booked GPU time."""
if self.path:
return self
kind, source, weight_name = parse_source(self.source)
if kind == "hf":
from huggingface_hub import hf_hub_download
name = self.weights_file or weight_name
if name is None:
name = _only_safetensors(source)
return replace(self, path=hf_hub_download(source, name), weights_file=name)
url = _civitai_download_url(source) if kind == "civitai" else source
return replace(self, path=_download(url))
def _only_safetensors(repo: str) -> str:
"""The single `.safetensors` in a repo, when the user did not name one."""
from huggingface_hub import list_repo_files
files = [f for f in list_repo_files(repo) if f.endswith(".safetensors")]
if not files:
raise ValueError(f"`{repo}` holds no .safetensors file.")
if len(files) > 1:
raise ValueError(f"`{repo}` holds {len(files)} .safetensors files - name one, e.g. {repo}/{files[0]}")
return files[0]
def slot_spec(key: str, url: str, weight: float) -> LoraSpec | None:
"""A UI slot into a spec, or None when the slot is empty."""
if not key or key == NONE_CHOICE:
return None
if key == CUSTOM_CHOICE:
url = (url or "").strip()
if not url:
return None
parse_source(url) # fail here, where the slot is still nameable
return LoraSpec(
label=url,
weight=float(weight),
adapter="custom_" + re.sub(r"[^A-Za-z0-9]", "_", url)[-40:],
source=url,
)
entry = CATALOG[key]
return LoraSpec(
label=key,
weight=float(weight),
adapter=entry["adapter"],
source=entry["repo"],
weights_file=entry["weights"],
)
def build_stack(slots) -> list[LoraSpec]:
"""`[(key, url, weight), ...]` into the specs to apply, first occurrence of an adapter winning."""
stack: list[LoraSpec] = []
seen: set[str] = set()
for key, url, weight in slots:
spec = slot_spec(key, url, weight)
if spec is None or spec.adapter in seen:
continue
seen.add(spec.adapter)
stack.append(spec)
return stack
# ---------------------------------------------------------------------------
# Application - inside the GPU worker
# ---------------------------------------------------------------------------
def _resident(pipe) -> set[str]:
"""Adapters already injected into the transformer. The model is the source of truth: a fresh worker has none."""
return set(getattr(pipe.transformer, "peft_config", {}) or {})
def stack_key(specs) -> tuple:
"""What makes two racks the same rack: the adapters and the weights they run at."""
return tuple(sorted((spec.adapter, round(float(spec.weight), 4)) for spec in specs))
def _state(pipe) -> dict:
"""The fold bookkeeping, kept on the transformer itself.
A recycled ZeroGPU worker re-forks the transformer from the startup copy, whose weights carry no fold and whose
`peft_config` is empty — and with the bookkeeping on that same object, the state re-derives instead of lying.
"""
transformer = pipe.transformer
if not hasattr(transformer, "_stack_state"):
transformer._stack_state = {"fused": None, "key": None, "repeats": 0}
return transformer._stack_state
def apply(pipe, specs) -> str:
"""Load whatever is missing, activate the rack with its weights, fold it once it settles, and report.
Unfused adapters cost an extra pair of matmuls on every linear layer of every step, which is the price of a
weight slider that moves for free. So the rack is left unfused while it is being tuned, and folded into the
bfloat16 weights once the same rack comes back `FUSE_AFTER` requests running: the fold is `weight * (B @ A)`
added to the weights once, and every later request with that same rack pays nothing at all for its LoRAs.
A slot that fails to load is reported and skipped - one unusable URL does not cost the whole request.
"""
import torch
active: list[str] = []
weights: list[float] = []
notes: list[str] = []
state = _state(pipe)
requested = stack_key(specs)
# A fold that no longer matches the request comes out of the weights before anything else touches them - in
# particular before a new adapter is injected into layers that still carry a merged one.
if state["fused"] is not None and state["fused"] != requested:
print(f"[lora] unfolding the previous rack ({len(state['fused'])} adapter(s))", flush=True)
pipe.unfuse_lora()
state["fused"] = None
for spec in specs:
try:
if spec.path is None:
raise ValueError("this slot was never fetched")
if spec.adapter not in _resident(pipe):
directory, filename = os.path.split(spec.path)
pipe.load_lora_weights(directory, weight_name=filename, adapter_name=spec.adapter)
# DiffSynth-Studio's fp32 factors would otherwise drag the unfused LoRA path into fp32.
pipe.transformer.to(torch.bfloat16)
print(f"[lora] loaded {spec.label} as `{spec.adapter}`", flush=True)
active.append(spec.adapter)
weights.append(float(spec.weight))
except Exception as error:
print(f"[lora] {spec.label} failed: {type(error).__name__}: {error}", flush=True)
notes.append(f"{spec.label} failed ({type(error).__name__})")
# The marker names what actually made it onto the card, which is the requested rack unless a slot failed.
key = stack_key(spec for spec in specs if spec.adapter in active)
if state["fused"] is not None and state["fused"] != key:
pipe.unfuse_lora()
state["fused"] = None
state["repeats"] = state["repeats"] + 1 if key == state["key"] else 1
state["key"] = key
if not active:
pipe.disable_lora()
return "; ".join(["none", *notes]) if notes else "none"
if state["fused"] == key:
# The weights already carry this rack. Nothing to activate, nothing to pay.
applied = ", ".join(f"{spec.label} @ {spec.weight:g}" for spec in specs if spec.adapter in active)
return "; ".join([f"{applied} (folded)", *notes])
_evict(pipe, keep=set(active), state=state)
pipe.set_adapters(active, adapter_weights=weights)
pipe.enable_lora()
folded = ""
if FUSE_AFTER and state["repeats"] >= FUSE_AFTER:
# `set_adapters` above put each slot's weight on its adapter, and PEFT merges with exactly that scaling —
# so the pipeline-level scale stays at 1.0. Merged adapters are skipped by the LoRA path, not applied twice.
print(f"[lora] folding {len(active)} adapter(s) into the weights", flush=True)
pipe.fuse_lora(adapter_names=active, lora_scale=1.0)
state["fused"] = key
folded = " (folded)"
applied = ", ".join(f"{spec.label} @ {spec.weight:g}" for spec in specs if spec.adapter in active)
return "; ".join([applied + folded, *notes])
def _evict(pipe, keep: set[str], state: dict) -> None:
"""Drop resident adapters outside the current rack once too many have piled up on the card.
Never while a fold is standing: deleting an adapter whose delta sits in the weights would strand it there.
"""
if state["fused"] is not None:
return
resident = _resident(pipe)
if len(resident) <= MAX_RESIDENT:
return
stale = sorted(resident - keep)
if stale:
print(f"[lora] evicting {len(stale)} adapter(s): {', '.join(stale)}", flush=True)
pipe.delete_adapters(stale)