multimodalart's picture
multimodalart HF Staff
Gradio 6 takes theme/css on launch()
1889dd6 verified
Raw History Blame
46.5 kB
"""H3-World โ€” action-conditioned world model on MiniMax-H3.
``DANNY621/H3-World`` is a rank-32 LoRA (65.6M params, 0.198% of the 33.1B DiT) trained on 7,872
ABot-World-Explorer clips. It turns MiniMax-H3 into a *keyboard-driven* world model: a first frame plus
one WASD/IJKL key state **per latent video frame** rolls the world forward.
Two things beyond a plain ``load_lora_weights`` are needed, and this Space implements both.
1. **Actions are text.** The training run's manifest records ``action_mode: "text"`` and
``action_dim: 0`` โ€” no action tensors were trained at all (``action_tensors: 0``, ``lora_tensors:
208``). Each latent frame's key state is rendered into one short English sentence
(*"the man walks forward, camera pans left sharply"*) and appended to the scene prompt, so the
conditioning arrives through MiniMax-H3's own text channel.
2. **A directed attention mask.** MiniMax-H3 denoises one packed sequence
``[text | keyframe anchors | audio | video]`` under *full* self-attention, so by default every video
row sees every sentence and the per-frame binding is lost. The patch below restricts each video
row of latent frame ``f`` to sentence ``f`` โ€” and only among the sentences; the scene prompt, the
keyframe anchors, the audio rows and the whole video block stay fully visible. It is *directed*:
the sentences themselves are not restricted.
Implementing that as a dense ``[S, S]`` mask would force SDPA off its flash kernel. Instead the
attention is split by key region and recombined with an online-softmax (log-sum-exp) merge, which
is exact:
A = text rows before the sentences (flash, no mask)
B = everything after the text block (flash, no mask) <- 99% of the keys
C = the sentence rows (small masked matmul, |C| ~ 700)
so the expensive half keeps running on flash attention and only the ~700 sentence keys pay for the
mask.
MiniMax-H3 is 195.9 GiB in bfloat16, past a ZeroGPU worker's storage budget, so โ€” exactly like every
other MiniMax-H3 Space โ€” the 62.14 GiB Qwen3-VL conditioner lives in
``multimodalart/qwen3vl-conditioner`` and this half holds the 61.73 GiB transformer plus the two VAEs.
"""
from __future__ import annotations
import os
import re
import tempfile
import time
import traceback
from functools import cache
import spaces
import gradio as gr
import torch
MODEL_REPO = os.environ.get("H3_MODEL_REPO", "MiniMaxAI/MiniMax-H3")
LORA_REPO = os.environ.get("H3_LORA_REPO", "DANNY621/H3-World")
LORA_FILE = os.environ.get("H3_LORA_FILE", "step-10000.safetensors")
CONDITIONER_SPACE = os.environ.get("H3_CONDITIONER", "multimodalart/qwen3vl-conditioner")
PLACEMENT = os.environ.get("H3_PLACEMENT", "pack").lower()
ATTENTION = os.environ.get("H3_ATTENTION", "_native_cudnn").lower()
GPU_SIZE = os.environ.get("H3_GPU_SIZE", "xlarge")
FPS, FRAMES_PER_CHUNK, LATENTS_PER_CHUNK = 24, 17, 5
MIN_UI_DURATION, MAX_UI_DURATION = 2, 8
# The H3-World inference runs recorded in the checkpoint's own manifests: 124 frames, 480x832, 24 fps,
# 50 steps, seed 0. Keep those as the defaults so the Space reproduces the released results.
DEFAULT_DURATION = 5
DEFAULT_STEPS = 50
DEFAULT_SUBJECT = "the man"
# The conditioner offers canvases up to 1344x768, but H3-World was trained at 832x480 and the directed
# mask costs O(sequence x captions) โ€” at 1344x768 / 8s a request would need ~35 GPU-minutes, far past
# what any visitor's ZeroGPU quota can book. So the list is the cheap tier of each aspect ratio, which
# also keeps every canvas near the resolution the LoRA actually saw.
#
# The labels are a *wire contract*: the conditioner validates `canvas` against its own dropdown, so
# these strings must stay byte-identical to its choices.
CANVASES = {
"960x544 ยท 16:9 fast": (544, 960),
"1024x576 ยท 16:9 fast": (576, 1024),
"544x960 ยท 9:16 fast": (960, 544),
"544x544 ยท 1:1 fast": (544, 544),
"768x576 ยท 4:3 fast": (576, 768),
"576x768 ยท 3:4 fast": (768, 576),
"1152x512 ยท 21:9 fast": (512, 1152),
}
# H3-World was trained at 832x480 (1.733:1); 960x544 is the closest canvas the shared conditioner offers.
DEFAULT_CANVAS = "960x544 ยท 16:9 fast"
OUTPUT_DIR = os.path.join(tempfile.gettempdir(), "h3-world-out")
os.makedirs(OUTPUT_DIR, exist_ok=True)
PIPE = None
LOAD_ERROR: str | None = None
LOADED_IN: float | None = None
# โ”€โ”€ The action vocabulary โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
#
# Eight recorded key columns (`action_columns: ["W","A","S","D","I","J","K","L"]` in the checkpoint's
# manifests) plus `F`, the "sharp panning" bit the training pipeline synthesised from the per-frame
# camera delta. The named presets are the ones the author's own generalization runs used.
KEYS = "WASDIJKLF"
PRESETS: dict[str, str] = {
"still": "",
"forward": "W",
"back": "S",
"strafe-left": "A",
"strafe-right": "D",
"forward-left": "WA",
"forward-right": "WD",
"back-left": "SA",
"back-right": "SD",
"pan-left": "J",
"pan-right": "L",
"pan-left-fast": "JF",
"pan-right-fast": "LF",
"tilt-up": "K",
"tilt-down": "I",
# a couple of natural combinations of the above
"forward-pan-left": "WJ",
"forward-pan-right": "WL",
}
def caption_for(keys: str, subject: str = DEFAULT_SUBJECT) -> str:
"""Render one latent frame's key state as the sentence H3-World was trained to read.
The register is the model card's: ``"the man walks forward, camera pans left sharply"`` โ€” a
locomotion clause from ``W/A/S/D``, an optional camera clause from ``I/J/K/L``, and ``F`` turning
the camera adverb from ``slowly`` to ``sharply``.
"""
held = set(keys.upper())
motion = []
if "W" in held:
motion.append("walks forward")
elif "S" in held:
motion.append("walks backward")
if "A" in held:
motion.append("strafes left")
elif "D" in held:
motion.append("strafes right")
body = " and ".join(motion) if motion else "stands still"
camera = []
if "J" in held:
camera.append("camera pans left")
elif "L" in held:
camera.append("camera pans right")
if "K" in held:
camera.append("camera tilts up")
elif "I" in held:
camera.append("camera tilts down")
sentence = f"{subject.strip() or DEFAULT_SUBJECT} {body}"
if camera:
adverb = "sharply" if "F" in held else "slowly"
sentence += ", " + ", ".join(camera) + f" {adverb}"
return sentence
def parse_script(script: str, num_slots: int) -> list[str]:
"""Expand an action script into exactly ``num_slots`` per-latent-frame key states.
Syntax is one action per comma / newline, optionally ``*n`` for how many latent frames it holds:
``forward*12, pan-right-fast*10, still``
An action is either a preset name (``forward-left``) or a raw key combination (``WA``, ``LF``).
Items without a count share whatever slots are left over; a script that runs short is extended by
holding its last action.
"""
items: list[tuple[str, int | None]] = []
for chunk in re.split(r"[,\n;]+", script or ""):
chunk = chunk.strip()
if not chunk:
continue
match = re.match(r"^(.*?)(?:\s*[*xร—]\s*(\d+)\s*)?$", chunk)
name = (match.group(1) or "").strip()
count = int(match.group(2)) if match.group(2) else None
lowered = name.lower().replace(" ", "-").replace("_", "-")
if lowered in PRESETS:
keys = PRESETS[lowered]
else:
raw = name.upper().replace(" ", "").replace("+", "")
if raw and set(raw) <= set(KEYS):
keys = raw
elif lowered in {"none", "idle", "stop"}:
keys = ""
else:
raise gr.Error(
f"Unknown action {name!r}. Use a preset ({', '.join(sorted(PRESETS))}) "
f"or a raw key combination out of {KEYS}."
)
items.append((keys, count))
if not items:
items = [("", None)]
counts = [count or 0 for _, count in items]
free = [index for index, (_, count) in enumerate(items) if not count]
if free:
remaining = max(0, num_slots - sum(counts))
base, extra = divmod(remaining, len(free))
for position, index in enumerate(free):
counts[index] = base + (1 if position < extra else 0)
sequence: list[str] = []
for (keys, _), count in zip(items, counts):
sequence.extend([keys] * count)
if not sequence:
sequence = [items[0][0]]
while len(sequence) < num_slots:
sequence.append(sequence[-1])
return sequence[:num_slots]
def latent_frames_for(num_frames: int) -> int:
"""How many latent video frames โ€” i.e. how many action slots โ€” a frame count carries."""
return (num_frames - LATENTS_PER_CHUNK) // FRAMES_PER_CHUNK * LATENTS_PER_CHUNK + 2
def snap_frames(seconds: float) -> int:
"""The frame count MiniMax-H3's video VAE can decode: the next ``17 * n + 5`` at 24 fps."""
frames = max(1, round(float(seconds) * FPS))
while frames % FRAMES_PER_CHUNK != LATENTS_PER_CHUNK:
frames += 1
return frames
def summarize(sequence: list[str], subject: str = DEFAULT_SUBJECT) -> str:
"""Group a per-frame key sequence into a compact, human-readable timeline."""
if not sequence:
return "_no actions_"
rows, start = [], 0
for index in range(1, len(sequence) + 1):
if index == len(sequence) or sequence[index] != sequence[start]:
keys = sequence[start] or "โ€”"
span = f"{start}" if index - start == 1 else f"{start}โ€“{index - 1}"
rows.append(f"| `{span}` | `{keys}` | {caption_for(sequence[start], subject)} |")
start = index
return "| latent frames | keys | sentence |\n|---|---|---|\n" + "\n".join(rows)
# โ”€โ”€ Prompt assembly + per-sentence token spans โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
@cache
def tokenizer():
"""MiniMax-H3's own Qwen3-VL tokenizer โ€” a few MB, no model weights."""
from transformers import AutoTokenizer
return AutoTokenizer.from_pretrained(MODEL_REPO, subfolder="tokenizer")
def build_conditioning_text(scene_prompt: str, sequence: list[str], subject: str):
"""Return ``(prompt, num_prompt_tokens, cuts)`` for the scene prompt plus one sentence per frame.
``cuts[0]`` is the token index where the first sentence starts and ``cuts[i + 1]`` where sentence
``i`` ends, both counted inside the prompt alone. ``MiniMaxH3TextEncoderStep`` presents a request
as ``"<Picture 1>: " + <vision block> + prompt`` with **no chat template and no special tokens**,
so those offsets map onto the packed sequence by a single shift โ€” see ``directed_plan``.
"""
head = scene_prompt.strip() + "\n"
sentences = [caption_for(keys, subject) + "\n" for keys in sequence]
prompt = head + "".join(sentences)
def length(text: str) -> int:
return len(tokenizer()(text, add_special_tokens=False)["input_ids"])
cuts, running = [length(head)], head
for sentence in sentences:
running += sentence
cuts.append(length(running))
total = length(prompt)
if cuts[-1] != total:
# A BPE merge crossed a sentence boundary; the spans would be off, so refuse to mask rather
# than mask the wrong rows.
return prompt, total, None
return prompt, total, cuts
def directed_plan(num_text_tokens: int, num_prompt_tokens: int, cuts, num_latent_frames: int, rows_per_frame: int):
"""Turn prompt-local sentence cuts into absolute packed-sequence spans."""
if not cuts or len(cuts) != num_latent_frames + 1:
return None
offset = num_text_tokens - num_prompt_tokens
if offset < 0:
return None
spans = [(offset + cuts[i], offset + cuts[i + 1]) for i in range(num_latent_frames)]
if any(end <= start for start, end in spans):
return None
return {
"cap_start": offset + cuts[0],
"cap_end": offset + cuts[-1],
"spans": spans,
"num_latent_frames": num_latent_frames,
"rows_per_frame": rows_per_frame,
}
# โ”€โ”€ The directed-mask attention patch โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
DIRECTED: dict = {"plan": None}
def _lse_flash(q, k, v, scale):
out, lse = torch.ops.aten._scaled_dot_product_flash_attention(q, k, v, 0.0, False, False, scale=scale)[:2]
return out, lse
def _lse_efficient(q, k, v, scale):
out, lse = torch.ops.aten._scaled_dot_product_efficient_attention(
q, k, v, None, True, 0.0, False, scale=scale
)[:2]
return out, lse
# The two fused kernels that hand back a log-sum-exp, best first. The first one that runs is kept.
_LSE_BACKENDS = [("flash", _lse_flash), ("mem_efficient", _lse_efficient)]
def _flash_lse(query, key, value, scale):
"""Fused attention that also hands back its log-sum-exp. ``[B, S, H, D]`` in and out.
Degrades to the chunked fp32 path if neither fused kernel accepts these shapes โ€” the merge only
needs *some* partition of the softmax, so the fallback is still exact, just slower.
"""
while _LSE_BACKENDS:
name, backend = _LSE_BACKENDS[0]
q, k, v = (tensor.transpose(1, 2) for tensor in (query, key, value))
try:
out, lse = backend(q, k, v, scale)
return out.transpose(1, 2), lse[..., : query.shape[1]].float()
except Exception as error: # noqa: BLE001
detail = str(error).splitlines()[0][:200]
print(f"[gen] {name} log-sum-exp unavailable ({type(error).__name__}: {detail})", flush=True)
_LSE_BACKENDS.pop(0)
return _masked_lse(query, key, value, None, scale, chunk=512)
def _masked_lse(query, key, value, mask, scale, chunk: int = 2048):
"""Exact attention over a *small* key set, with a ``[Sq, Lk]`` boolean mask (True = visible).
The scores are accumulated in fp32 โ€” the key set is a few hundred rows, so this costs nothing and
keeps the merged softmax as precise as the flash half it is combined with.
"""
batch, length, heads, _ = query.shape
out = torch.empty_like(query)
lse = torch.empty(batch, heads, length, device=query.device, dtype=torch.float32)
k = key.transpose(1, 2).transpose(-1, -2).float() # [B, H, D, Lk]
v = value.transpose(1, 2) # [B, H, Lk, D]
for start in range(0, length, chunk):
stop = min(start + chunk, length)
q = query[:, start:stop].transpose(1, 2).float() # [B, H, q, D]
scores = torch.matmul(q, k) * scale
if mask is not None:
scores.masked_fill_(~mask[start:stop].unsqueeze(0).unsqueeze(0), float("-inf"))
lse[:, :, start:stop] = torch.logsumexp(scores, dim=-1)
probs = torch.softmax(scores, dim=-1).to(v.dtype)
out[:, start:stop] = torch.matmul(probs, v).transpose(1, 2).to(out.dtype)
del scores, probs, q
return out, lse
def _merge_lse(out_a, lse_a, out_b, lse_b, chunk: int = 4096):
"""Online-softmax merge of two partitions of the same softmax. Exact, not an approximation.
Written in place into ``out_a`` / ``lse_a`` and chunked over the sequence, so the fp32 working set
stays a few hundred MB instead of a few GB at MiniMax-H3's sequence lengths.
"""
for start in range(0, out_a.shape[1], chunk):
stop = min(start + chunk, out_a.shape[1])
a = lse_a[:, :, start:stop].transpose(1, 2).unsqueeze(-1)
b = lse_b[:, :, start:stop].transpose(1, 2).unsqueeze(-1)
peak = torch.maximum(a, b)
weight_a = torch.exp(a - peak)
weight_b = torch.exp(b - peak)
total = weight_a + weight_b
merged = (out_a[:, start:stop].float() * weight_a + out_b[:, start:stop].float() * weight_b) / total
out_a[:, start:stop] = merged.to(out_a.dtype)
lse_a[:, :, start:stop] = (peak + torch.log(total)).squeeze(-1).transpose(1, 2)
return out_a, lse_a
def directed_attention(query, key, value, plan):
"""Full self-attention with each video row's view of the *sentences* narrowed to its own.
Returns ``None`` when ``plan`` does not describe this call's sequence โ€” the token refiner runs the
same attention module over the text stream alone, and it must keep the unmasked path.
"""
length = query.shape[1]
cap_start, cap_end = plan["cap_start"], plan["cap_end"]
rows_per_frame = plan["rows_per_frame"]
num_latent_frames = plan["num_latent_frames"]
# The video block is the tail of the packed sequence, so its start follows from the geometry
# alone โ€” no need to re-derive the keyframe or audio row counts here.
video_start = length - num_latent_frames * rows_per_frame
if video_start <= cap_end:
return None
scale = query.shape[-1] ** -0.5
out, lse = _flash_lse(query, key[:, cap_end:], value[:, cap_end:], scale) # region B
if cap_start > 0: # region A
out, lse = _merge_lse(out, lse, *_flash_lse(query, key[:, :cap_start], value[:, :cap_start], scale))
# region C โ€” the sentences, the only place the mask bites.
visible = torch.ones(length, cap_end - cap_start, dtype=torch.bool, device=query.device)
for index, (start, stop) in enumerate(plan["spans"]):
rows = slice(video_start + index * rows_per_frame, video_start + (index + 1) * rows_per_frame)
visible[rows] = False
visible[rows, start - cap_start : stop - cap_start] = True
out_c, lse_c = _masked_lse(query, key[:, cap_start:cap_end], value[:, cap_start:cap_end], visible, scale)
out, _ = _merge_lse(out, lse, out_c, lse_c)
return out
def install_directed_processor() -> None:
"""Teach ``MiniMaxH3AttnProcessor`` the directed mask, in place.
Patching the class rather than swapping processor *instances* keeps whatever
``set_attention_backend`` configured on them, and covers every attention module the transformer
builds โ€” including ones created after this call.
"""
from diffusers.models.transformers import transformer_minimax_h3 as h3
if getattr(h3.MiniMaxH3AttnProcessor, "_h3world_directed", False):
return
base_call = h3.MiniMaxH3AttnProcessor.__call__
def __call__(self, attn, hidden_states, rotary_emb=None, attention_mask=None):
plan = DIRECTED.get("plan")
if plan is None or attention_mask is not None:
return base_call(self, attn, hidden_states, rotary_emb, attention_mask)
if attn.fused_projections:
query, key, value = attn.to_qkv(hidden_states).chunk(3, dim=-1)
else:
query, key, value = attn.to_q(hidden_states), attn.to_k(hidden_states), attn.to_v(hidden_states)
query = attn.norm_q(query.unflatten(-1, (attn.heads, -1)))
key = attn.norm_k(key.unflatten(-1, (attn.heads, -1)))
value = value.unflatten(-1, (attn.heads, -1))
if rotary_emb is not None:
query = h3._apply_rotary_emb(query, *rotary_emb)
key = h3._apply_rotary_emb(key, *rotary_emb)
attended = directed_attention(query, key, value, plan)
if attended is None: # not the packed sequence โ€” the token refiner's own stream
return base_call(self, attn, hidden_states, rotary_emb, attention_mask)
attended = attended.flatten(2, 3).type_as(query)
return attn.to_out[1](attn.to_out[0](attended))
h3.MiniMaxH3AttnProcessor.__call__ = __call__
h3.MiniMaxH3AttnProcessor._h3world_directed = True
# โ”€โ”€ LoRA loading (weight folding) โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
#
# H3-World ships 208 tensors in the *original* MiniMax-H3 naming
# (``blocks.N.attn.qkv_proj``, ``token_refiner.blocks.N.attn.out_proj``), so merging into the diffusers
# port means replaying the transforms `scripts/convert_minimax_h3_to_diffusers.py` applied to the base
# weights โ€” a delta is only valid in the layout of the weight it is added to. The subtle one is the
# fused QKV: the original checkpoint stores it **per-head interleaved**
# (``[head0: q k v, head1: q k v, ...]``), so the rows must be de-interleaved before being split into
# thirds. Splitting them into contiguous thirds directly scatters every head's q/k/v across all three
# projections and turns the adapter into structured noise on all 52 attention blocks.
def _reorder_interleaved_qkv(weight, num_heads: int, head_dim: int):
"""De-interleave per-head fused-QKV rows into ``[q_all; k_all; v_all]``."""
grouped = weight.reshape(num_heads, 3 * head_dim, *weight.shape[1:])
query, key, value = grouped.split(head_dim, dim=1)
return torch.cat(
[part.reshape(num_heads * head_dim, *weight.shape[1:]) for part in (query, key, value)], dim=0
)
def _lora_target_name(source_name: str) -> str:
if source_name.startswith("token_refiner.blocks."):
target = source_name.replace("token_refiner.blocks.", "token_refiner.refiner_blocks.", 1)
elif source_name.startswith("blocks."):
target = source_name.replace("blocks.", "transformer_blocks.", 1)
else:
target = source_name
return target.replace("final_layer.adaln_proj.linear", "norm_out.linear")
def _lora_targets(name: str, b_weight, num_heads: int, head_dim: int):
"""Yield ``(diffusers_param_key, row_transformed_B)`` for one original LoRA base name."""
target = _lora_target_name(name)
if target.endswith(".attn.qkv_proj"):
prefix = target.removesuffix("qkv_proj")
rows = _reorder_interleaved_qkv(b_weight, num_heads, head_dim)
for kind, part in zip(("q", "k", "v"), rows.split(num_heads * head_dim, dim=0)):
yield f"{prefix}to_{kind}.weight", part.contiguous()
elif target.endswith(".mlp.fc1"):
# SwiGLU gate/value swap: the checkpoint stores [gate, value]; diffusers stores [value, gate].
gate, value = b_weight.chunk(2, dim=0)
yield target.replace(".mlp.fc1", ".ff.net.0.proj") + ".weight", torch.cat([value, gate]).contiguous()
elif target.endswith(".mlp.fc2"):
yield target.replace(".mlp.fc2", ".ff.net.2") + ".weight", b_weight
elif target.endswith(".attn.out_proj"):
yield target.replace(".attn.out_proj", ".attn.to_out.0") + ".weight", b_weight
else:
yield target + ".weight", b_weight
def load_and_apply_lora(transformer) -> str:
"""Download H3-World, remap keys, and fold ``scale * (B @ A)`` into the base weights."""
from huggingface_hub import hf_hub_download
from safetensors.torch import load_file
lora = load_file(hf_hub_download(LORA_REPO, LORA_FILE))
suffix_a, suffix_b = ".lora_A.default.weight", ".lora_B.default.weight"
bases = sorted({key[: -len(suffix_a)] for key in lora if key.endswith(suffix_a)})
unexpected = [key for key in lora if not key.endswith((suffix_a, suffix_b))]
if unexpected:
raise ValueError(f"{LORA_FILE} holds {len(unexpected)} non-LoRA tensors, e.g. {unexpected[:5]}")
if not bases:
raise ValueError(f"No lora_A/lora_B pairs found in {LORA_FILE}")
ranks = set()
for name in bases:
if f"{name}{suffix_b}" not in lora:
raise ValueError(f"LoRA is missing the lora_B twin of {name}{suffix_a}")
ranks.add(lora[f"{name}{suffix_a}"].shape[0])
if len(ranks) != 1:
raise ValueError(f"LoRA mixes ranks {sorted(ranks)}")
rank = ranks.pop()
num_heads = transformer.config.num_attention_heads
head_dim = transformer.config.attention_head_dim
# alpha == rank in the training run, so the merge scale is 1.0.
scale = float(os.environ.get("H3_LORA_SCALE", "1.0"))
params = dict(transformer.named_parameters())
folded, missed = 0, []
with torch.no_grad():
for name in bases:
a = lora[f"{name}{suffix_a}"]
b = lora[f"{name}{suffix_b}"]
for key, b_part in _lora_targets(name, b, num_heads, head_dim):
param = params.get(key)
if param is None:
missed.append(key)
continue
delta = scale * (b_part.to(torch.float32) @ a.to(torch.float32))
if delta.shape != param.shape:
raise ValueError(
f"LoRA delta for `{key}` has shape {tuple(delta.shape)}, "
f"base weight is {tuple(param.shape)}"
)
param.data = (param.data.float() + delta.to(param.device)).to(param.dtype)
folded += 1
if missed:
raise ValueError(
f"{len(missed)} LoRA targets matched no transformer weight, e.g. {missed[:5]}. "
"The LoRA and the diffusers transformer disagree on module naming."
)
return f"H3-World merged ยท {len(bases)} targets, rank {rank}, scale {scale:g} -> {folded} weight deltas"
# โ”€โ”€ Model loading โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
def lower_duration_floor(seconds: float = MIN_UI_DURATION) -> None:
"""Let the pipeline generate below its 5 s floor."""
from diffusers.modular_pipelines.minimax_h3.modular_pipeline import MiniMaxH3ModularPipeline
MiniMaxH3ModularPipeline.min_duration = property(lambda self: float(seconds))
def load_models() -> str | None:
"""Load the denoising half at startup: transformer + VAEs, fold H3-World in, patch attention."""
global PIPE, LOAD_ERROR, LOADED_IN
if PIPE is not None or LOAD_ERROR is not None:
return LOAD_ERROR
started = time.time()
try:
from diffusers import ComponentsManager
from h3_split_blocks import MiniMaxH3GeneratorBlocks
lower_duration_floor()
install_directed_processor()
manager = ComponentsManager()
blocks = MiniMaxH3GeneratorBlocks()
print(f"[gen] loading {[c.name for c in blocks.expected_components]} from {MODEL_REPO} ...", flush=True)
pipe = blocks.init_pipeline(MODEL_REPO, components_manager=manager, collection="h3")
pipe.load_components(dtype=torch.bfloat16)
pipe.vae.set_attention_backend("native")
pipe.audio_vae.set_attention_backend("native")
pipe.transformer.set_attention_backend(ATTENTION)
print(f"[gen] {load_and_apply_lora(pipe.transformer)}", flush=True)
if PLACEMENT == "pack":
pipe.transformer.to("cuda")
PIPE = pipe
LOADED_IN = time.time() - started
print(f"[gen] ready in {LOADED_IN:.0f}s", flush=True)
except Exception as error:
traceback.print_exc()
LOAD_ERROR = f"**Loading failed** after {time.time() - started:.0f}s: `{type(error).__name__}: {error}`"
return LOAD_ERROR
# โ”€โ”€ Remote conditioning โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
@cache
def conditioner():
from gradio_client import Client
return Client(CONDITIONER_SPACE)
def conditioner_client(ip_token):
if not ip_token:
return conditioner()
from gradio_client import Client
return Client(CONDITIONER_SPACE, headers={"x-ip-token": ip_token})
def encode_remote(prompt, image_path, canvas, num_frames, ip_token=None):
"""Ask ``multimodalart/qwen3vl-conditioner`` for ``prompt_embeds`` + tags + resolved geometry."""
from gradio_client import handle_file
from safetensors import safe_open
def call():
return conditioner_client(ip_token).predict(
prompt=prompt,
image_path=handle_file(image_path) if image_path else None,
last_image_path=None,
canvas=canvas,
num_frames=num_frames,
rewrite_prompt=False,
api_name="/encode",
)
try:
path, plan = call()
except Exception as first:
print(f"[conditioner] retrying with a fresh client after: {first}", flush=True)
conditioner.cache_clear()
path, plan = call()
with safe_open(path, framework="pt") as handle:
return handle.get_tensor("prompt_embeds"), handle.get_tensor("text_token_tags"), handle.metadata(), plan
# โ”€โ”€ Geometry helpers โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
def _snap_canvas(aspect: float, current_canvas: str) -> str:
"""The supported canvas whose aspect ratio is closest to ``aspect`` (smallest at that ratio)."""
fastest: dict[float, tuple[str, tuple[int, int]]] = {}
for label, (h, w) in CANVASES.items():
ratio = w / h
if ratio not in fastest or w * h < fastest[ratio][1][0] * fastest[ratio][1][1]:
fastest[ratio] = (label, (h, w))
ratio = min(fastest, key=lambda r: abs(r - aspect))
cur_h, cur_w = CANVASES[current_canvas]
if abs(cur_w / cur_h - aspect) <= abs(ratio - aspect):
return current_canvas
return fastest[ratio][0]
def _cover_crop(image_path: str, canvas_label: str) -> str:
"""Center cover-crop ``image_path`` to ``canvas_label``'s aspect ratio, into a fresh temp file.
Never in place: the same helper runs on the bundled example assets, and a request must not rewrite
the repository's own files.
"""
from PIL import Image as _Image, ImageOps as _ImageOps
h, w = CANVASES[canvas_label]
target = w / h
img = _ImageOps.exif_transpose(_Image.open(image_path)).convert("RGB")
if abs(img.width / img.height - target) > 1e-3:
if img.width / img.height > target:
new_w = int(img.height * target)
left = (img.width - new_w) // 2
img = img.crop((left, 0, left + new_w, img.height))
else:
new_h = int(img.width / target)
top = (img.height - new_h) // 2
img = img.crop((0, top, img.width, top + new_h))
out = os.path.join(OUTPUT_DIR, f"kf-{int(time.time() * 1e6)}.png")
img.save(out)
return out
def _as_path(value):
if isinstance(value, dict):
value = value.get("path") or (value.get("url") or "").removeprefix("/gradio_api/file=")
return value or None
def _fit_keyframe(image_path, current_canvas):
from PIL import Image as _Image
with _Image.open(image_path) as img:
aspect = img.width / img.height
label = _snap_canvas(aspect, current_canvas)
return _cover_crop(image_path, label), label
def _fit_keyframe_ui(image_path, current_canvas):
"""``image.upload`` handler: snap the canvas to the frame, then crop the frame to the canvas."""
path = _as_path(image_path)
if not path:
return gr.update(), gr.update()
cropped, label = _fit_keyframe(path, current_canvas)
return gr.update(value=cropped), gr.update(value=label)
# โ”€โ”€ GPU duration estimation โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
# Fitted on this Space's own RTX PRO 6000 pool, over three measured requests at 960x544:
#
# 16 steps ยท 56 frames ยท 9,180 rows ยท directed 45 s 16 steps ยท 56 frames ยท undirected 32 s
# 50 steps ยท 124 frames ยท 19,380 rows ยท directed 303 s
#
# `_ATTN_*` is the unmasked block cost per step (linear + quadratic in the packed sequence) and
# `_MASK` is the directed mask's own term, which scales as sequence x caption-rows rather than
# sequence squared โ€” the whole reason the mask is affordable at all.
_ATTN_LINEAR, _ATTN_QUADRATIC = 5.393e-5, 1.690e-9
_MASK = 6.508e-7
_TOKENS_PER_CAPTION = 8
_DECODE_BASE, _DECODE_PER_DEFAULT_CANVAS, _DEFAULT_CANVAS_PIXELS = 15, 15, 960 * 544 * 124
# ~12% over the fit, plus one cold worker's weight placement.
_MARGIN, _PLACEMENT_ALLOWANCE = 1.12, 12
def get_duration(prompt_embeds, tags, plan, image, height, width, num_frames, steps, seed, *args, **kwargs):
"""Estimate GPU seconds for one ``_generate`` call, from measurements rather than guesswork."""
height, width, num_frames, steps = int(height), int(width), int(num_frames), int(steps)
patches = (height // 32) * (width // 32)
latent_frames = latent_frames_for(num_frames)
rows = latent_frames * patches + (1 if image is not None else 0) * patches
per_step = _ATTN_LINEAR * rows + _ATTN_QUADRATIC * rows**2
if plan is not None:
per_step += _MASK * rows * len(plan["spans"]) * _TOKENS_PER_CAPTION
decode = _DECODE_BASE + _DECODE_PER_DEFAULT_CANVAS * (height * width * num_frames) / _DEFAULT_CANVAS_PIXELS
return max(60, int((steps * per_step + decode) * _MARGIN) + _PLACEMENT_ALLOWANCE)
# โ”€โ”€ Inference โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
@spaces.GPU(duration=get_duration, size=GPU_SIZE)
def _generate(prompt_embeds, tags, plan, image, height, width, num_frames, steps, seed):
"""Denoise + decode on GPU time, with the directed mask active for this request's geometry."""
if PLACEMENT == "pack":
PIPE.vae.to("cuda")
PIPE.audio_vae.to("cuda")
elif PLACEMENT == "lazy":
PIPE.to("cuda")
DIRECTED["plan"] = plan
try:
state = PIPE(
prompt_embeds=prompt_embeds.to("cuda"),
text_token_tags=tags,
image=image,
height=height,
width=width,
num_frames=num_frames,
# `MiniMaxH3Scheduler.set_timesteps` counts *sigma grid points*, terminal 0.0 included, and
# runs `len(sigmas) - 1` model evaluations โ€” so ask for one more point than steps wanted.
num_inference_steps=int(steps) + 1,
generator=torch.Generator("cpu").manual_seed(int(seed)),
)
finally:
DIRECTED["plan"] = None
return state.get("videos")[0], state.get("audio")[0].cpu(), state.get("sampling_rate")
def _caller_ip_token() -> str | None:
from gradio.context import LocalContext
request = LocalContext.request.get()
return request.headers.get("x-ip-token") if request is not None else None
def generate(
prompt: str,
image_path: str | None = None,
script: str = "forward",
canvas: str = DEFAULT_CANVAS,
duration: float = DEFAULT_DURATION,
steps: int = DEFAULT_STEPS,
seed: int = 0,
subject: str = DEFAULT_SUBJECT,
directed: bool = True,
progress=gr.Progress(track_tqdm=True),
):
"""Roll a world forward from one frame under a keyboard action script.
Args:
prompt: A motion-free description of the scene, as in the ABot captions H3-World trained on.
image_path: The first frame to continue from. Optional, but this is a world model โ€” give it one.
script: Actions per latent frame, e.g. ``"forward*12, pan-right-fast*10, still"``. Presets are
still / forward / back / strafe-left / strafe-right / forward-left / forward-right /
back-left / back-right / pan-left / pan-right / pan-left-fast / pan-right-fast / tilt-up /
tilt-down, or a raw key combination out of W A S D I J K L F.
canvas: Output resolution; snapped to the first frame's aspect ratio when one is given.
duration: Seconds of video, rounded up to the next frame count the VAE can decode.
steps: Denoising steps. The released H3-World results use 50.
seed: Random seed.
subject: How the per-frame sentences refer to the character.
directed: Bind each sentence to its own latent frame with H3-World's directed attention mask.
Turning it off is the ablation: the sentences become one global prompt.
Returns:
The generated video (with MiniMax-H3's native soundtrack) and a markdown report holding the
per-frame action timeline.
"""
if LOAD_ERROR:
raise gr.Error(LOAD_ERROR.replace("**", "").replace("`", ""))
if PIPE is None:
raise gr.Error("The model is still loading โ€” watch the Space logs and retry shortly.")
if not prompt or not prompt.strip():
raise gr.Error("H3-World still needs a scene description alongside the actions.")
from PIL import Image, ImageOps
from diffusers.utils import encode_video
first = _as_path(image_path)
if first:
first, canvas = _fit_keyframe(first, canvas)
num_frames = snap_frames(duration)
num_latent_frames = latent_frames_for(num_frames)
sequence = parse_script(script, num_latent_frames)
text, num_prompt_tokens, cuts = build_conditioning_text(prompt, sequence, subject)
progress(0.0, desc=f"Encoding the scene and {num_latent_frames} action sentences ...")
conditioned = time.time()
prompt_embeds, tags, metadata, _ = encode_remote(
text, first, canvas, num_frames, ip_token=_caller_ip_token()
)
condition_seconds = time.time() - conditioned
height, width, num_frames = (int(metadata[key]) for key in ("height", "width", "num_frames"))
plan = None
if directed:
plan = directed_plan(
int(prompt_embeds.shape[1]),
num_prompt_tokens,
cuts,
num_latent_frames,
(height // 32) * (width // 32),
)
if plan is None:
print("[gen] WARNING: could not resolve sentence spans; running without the directed mask", flush=True)
keyframe = ImageOps.exif_transpose(Image.open(first)).convert("RGB") if first else None
progress(0.1, desc=f"Generating {num_frames / FPS:.1f}s at {width}x{height} in {int(steps)} steps ...")
started = time.time()
frames, audio, sampling_rate = _generate(
prompt_embeds, tags, plan, keyframe, height, width, num_frames, int(steps), int(seed)
)
generate_seconds = time.time() - started
path = os.path.join(OUTPUT_DIR, f"h3-world-{int(time.time() * 1000)}.mp4")
encode_video(frames, fps=FPS, output_path=path, audio=audio, audio_sample_rate=sampling_rate)
mask_note = (
f"directed mask on {num_latent_frames} sentence spans"
if plan is not None
else "directed mask **off** (sentences act as one global prompt)"
)
report = (
f"{width}x{height} ยท {num_frames} frames ({num_frames / FPS:.2f}s) ยท {int(steps)} steps ยท "
f"{mask_note} ยท conditioner {condition_seconds:.0f}s ยท denoise + decode {generate_seconds:.0f}s ยท "
f"seed {int(seed)}\n\n"
f"<details><summary>Action timeline</summary>\n\n{summarize(sequence, subject)}\n\n</details>"
)
print(f"[gen] {report.splitlines()[0]}", flush=True)
return path, report
def preview(script: str, duration: float, subject: str) -> str:
"""Render the per-latent-frame sentences an action script expands to, without generating."""
try:
slots = latent_frames_for(snap_frames(duration))
return f"**{slots} latent frames**\n\n" + summarize(parse_script(script, slots), subject)
except gr.Error as error:
return f"โš ๏ธ {error.message if hasattr(error, 'message') else error}"
except Exception as error: # noqa: BLE001
return f"โš ๏ธ {error}"
# โ”€โ”€ UI โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€โ”€
INTRO = """# H3-World โ€” a keyboard-driven world model
<div align="center">
<a href="https://huggingface.co/DANNY621/H3-World" target="_blank" rel="noopener"><strong>[ LoRA ]</strong></a> &nbsp;
<a href="https://huggingface.co/MiniMaxAI/MiniMax-H3" target="_blank" rel="noopener"><strong>[ base model ]</strong></a> &nbsp;
<a href="https://huggingface.co/datasets/acvlab/ABot-World-Explorer-500h" target="_blank" rel="noopener"><strong>[ training data ]</strong></a>
</div>
[**H3-World**](https://huggingface.co/DANNY621/H3-World) is a rank-32 LoRA on
[MiniMax-H3](https://huggingface.co/MiniMaxAI/MiniMax-H3)'s 33B DiT that turns it into an
**action-conditioned world model**: give it a first frame and a **WASD / IJKL key state per latent
video frame**, and it rolls the world forward under your input.
Each frame's key state becomes one short sentence โ€” *"the man walks forward, camera pans left
sharply"* โ€” and a **directed attention mask** binds that sentence to that frame inside MiniMax-H3's
packed sequence, which is what makes the control per-frame rather than a global prompt.
`W`/`A`/`S`/`D` move ยท `I`/`J`/`K`/`L` aim the camera ยท `F` makes the camera move sharp."""
CSS = """
.main.fillable { max-width: 1150px !important; }
"""
EXAMPLES = [
[
"The scene is an urban street intersection under bright daylight, featuring a mix of low-rise "
"commercial buildings with weathered facades, including a red-and-cream tiled structure with East "
"Asian architectural motifs and shuttered storefronts bearing faded signage; palm trees, concrete "
"sidewalks with red curbs, utility poles, and distant high-rises contribute to a sun-bleached, "
"slightly gritty Los Santos aesthetic, with asphalt roads marked by white lane lines and scattered "
"debris, all rendered in realistic textures of concrete, metal, and painted wood under clear skies.",
"assets/street_intersection.jpg",
"forward*20, forward-right*10, pan-right-fast*7",
],
[
"This is an expansive outdoor Western landscape featuring rolling grassy hills interspersed with "
"weathered sandstone buttes and scattered boulders, under a bright blue sky with soft cumulus "
"clouds; the terrain exhibits naturalistic textures of dry grass, cracked earth, and eroded rock "
"faces in muted ochres, sage greens, and pale grays, evoking a sun-drenched, arid yet verdant "
"high-plains environment with a rustic, late-19th-century frontier aesthetic.",
"assets/western_hills.jpg",
"forward*14, pan-left*12, forward-left*11",
],
[
"The scene is an expansive, multi-level indoor parking garage constructed of raw concrete, "
"featuring thick support pillars, a ribbed ceiling, and perforated block walls that allow dappled "
"daylight to filter through oval-shaped openings, casting soft, irregular patches of light across "
"the asphalt floor marked with faded yellow parking lines and directional arrows; a male character "
"stands centrally, wearing a short-sleeved yellow floral-patterned shirt over a white undershirt "
"and light-wash jeans, under diffuse, slightly cool ambient lighting that emphasizes the gritty "
"textures of weathered concrete and oil-stained pavement.",
"assets/parking_garage.jpg",
"still*6, forward*16, strafe-left*15",
],
]
with gr.Blocks(title="H3-World") as demo:
gr.Markdown(INTRO)
with gr.Row(equal_height=False):
with gr.Column():
image = gr.Image(label="First frame", type="filepath", height=260)
prompt = gr.Textbox(
label="Scene description",
lines=4,
placeholder="Describe the world โ€” the setting, materials, light. Leave the motion to the actions.",
)
script = gr.Textbox(
label="Action script",
value="forward*20, forward-right*10, pan-right-fast*7",
lines=2,
info=(
"One action per comma, optional ร—count of latent frames. Presets: "
+ ", ".join(sorted(PRESETS))
+ " โ€” or raw keys like WA, LF."
),
)
run = gr.Button("Roll the world forward", variant="primary", size="lg")
with gr.Accordion("Advanced", open=False):
duration_slider = gr.Slider(
label="Duration (s)",
minimum=MIN_UI_DURATION,
maximum=MAX_UI_DURATION,
step=1,
value=DEFAULT_DURATION,
)
steps_slider = gr.Slider(
label="Steps", minimum=16, maximum=50, step=1, value=DEFAULT_STEPS
)
canvas = gr.Dropdown(label="Canvas", choices=list(CANVASES), value=DEFAULT_CANVAS)
seed = gr.Number(label="Seed", value=0, precision=0)
subject = gr.Textbox(label="Subject", value=DEFAULT_SUBJECT)
directed = gr.Checkbox(
label="Directed attention mask",
value=True,
info="Off = ablation: every sentence becomes one global prompt again.",
)
with gr.Column():
result = gr.Video(label="Generated world", autoplay=True)
report = gr.Markdown()
with gr.Accordion("Per-frame sentences", open=False):
preview_md = gr.Markdown(preview("forward*20, forward-right*10, pan-right-fast*7",
DEFAULT_DURATION, DEFAULT_SUBJECT))
INPUTS = [prompt, image, script, canvas, duration_slider, steps_slider, seed, subject, directed]
OUTPUTS = [result, report]
image.upload(_fit_keyframe_ui, [image, canvas], [image, canvas])
for control in (script, duration_slider, subject):
control.change(preview, [script, duration_slider, subject], [preview_md])
gr.Examples(
examples=EXAMPLES,
inputs=[prompt, image, script],
fn=generate,
outputs=OUTPUTS,
cache_examples=True,
cache_mode="lazy",
)
run.click(generate, inputs=INPUTS, outputs=OUTPUTS, api_name="generate")
load_models()
if __name__ == "__main__":
# Gradio 6 takes `theme` / `css` on `launch()`, not on the `Blocks` constructor.
demo.launch(
theme=gr.themes.Citrus(),
css=CSS,
show_error=True,
mcp_server=True,
allowed_paths=[OUTPUT_DIR],
)