multimodalart's picture
multimodalart HF Staff
Drop the static keycap legend, now that the builder replaces it
f984226 verified
Raw History Blame Contribute Delete
73.6 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 cuts the sentences'
outgoing edges, as the reference's ``mask_mod`` does: sentence ``f`` is readable only by itself and
by the video rows of latent frame ``f``, and as a query it reads frame ``f``'s video rows only.
The token refiner applies the same rule as one segment per sentence.
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 json
import os
import re
import tempfile
import time
import traceback
from functools import cache
import spaces
import gradio as gr
import torch
import h3_turbo_lora
def _bill_examples_to_the_clicker() -> None:
import inspect
from gradio.context import LocalContext
from gradio.helpers import Examples
if "request=None" not in inspect.getsource(Examples.cache):
raise RuntimeError("example-quota patch: Examples.cache no longer passes "
"request=None; re-read gradio/helpers.py before bumping "
"sdk_version, this patch may now be a no-op.")
original = gr.Blocks.process_api
async def process_api(self, *args, **kwargs):
if kwargs.get("request") is None and not kwargs.get("explicit_call"):
kwargs["request"] = LocalContext.request.get(None)
kwargs["event_id"] = kwargs.get("event_id") or LocalContext.event_id.get(None)
return await original(self, *args, **kwargs)
gr.Blocks.process_api = process_api
_bill_examples_to_the_clicker()
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. 50 is still reachable on the Steps slider (it is its maximum).
DEFAULT_DURATION = 5
DEFAULT_STEPS = 28
DEFAULT_SUBJECT = "the man"
# The two sampling modes the quality comparison is between: MiniMax-H3's own 28-step default with
# H3-World alone, and `larryvrh/MiniMax-H3-Turbo-Lora`'s few-step LoRA folded in on top of it.
MODE_BASE = "28 steps · no turbo LoRA"
MODE_TURBO = "8 steps · turbo LoRA"
MODES: dict[str, tuple[int, bool]] = {
MODE_BASE: (DEFAULT_STEPS, False),
MODE_TURBO: (h3_turbo_lora.TURBO_STEPS, True),
}
DEFAULT_MODE = MODE_BASE
# 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
TURBO_ERROR: str | 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",
}
MOTION = {"W": "walks forward", "S": "walks backward",
"A": "strafes left", "D": "strafes right"}
MOTION_ORDER = ("W", "S", "A", "D")
MOTION_IDLE = "stands still"
PAN_KEY = {"J": "left", "L": "right"}
TILT_KEY = {"I": "tilts down", "K": "tilts up"}
CAMERA_IDLE = "holds steady"
CAMERA_FOLLOW = "follows him"
FAST_KEY = "F"
def _purify(held: dict[str, bool], pairs) -> None:
for a, b in pairs:
if held[a] and held[b]:
held[a] = held[b] = False
def _motion_clause(held: dict[str, bool]) -> str:
_purify(held, (("W", "S"), ("A", "D")))
words = [MOTION[name] for name in MOTION_ORDER if held[name]]
return " and ".join(words) if words else MOTION_IDLE
def _camera_clause(held: dict[str, bool], moving: bool) -> str:
_purify(held, (("J", "L"), ("I", "K")))
parts = []
for key, side in PAN_KEY.items():
if held[key]:
parts.append(f"pans {side} {'sharply' if held[FAST_KEY] else 'slowly'}")
for key, word in TILT_KEY.items():
if held[key]:
parts.append(word)
if parts:
return " and ".join(parts)
return CAMERA_FOLLOW if moving else CAMERA_IDLE
def caption_for(keys: str, subject: str = DEFAULT_SUBJECT) -> str:
upper = keys.upper()
held = {name: (name in upper) for name in KEYS}
motion = _motion_clause(dict(held))
camera = _camera_clause(dict(held), motion != MOTION_IDLE)
return f"{subject.strip() or DEFAULT_SUBJECT} {motion}, camera {camera}"
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}
# Bumped whenever the patched processor's behaviour changes; see `install_directed_processor`.
DIRECTED_PATCH_VERSION = "leak_out+leak_in+refiner_segments"
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):
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()
v = value.transpose(1, 2)
for start in range(0, length, chunk):
stop = min(start + chunk, length)
q = query[:, start:stop].transpose(1, 2).float()
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.nan_to_num(torch.softmax(scores, dim=-1), nan=0.0).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 _refiner_segmented(query, key, value, plan, scale):
cap_start = plan["cap_start"]
segments = ([(0, cap_start)] if cap_start > 0 else []) + list(plan["spans"])
out = torch.zeros_like(query)
for start, stop in segments:
if stop <= start:
continue
q = query[:, start:stop]
piece, _ = _flash_lse(q, key[:, start:stop], value[:, start:stop], scale)
out[:, start:stop] = piece.to(out.dtype)
return out
def directed_attention(query, key, value, plan):
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"]
scale = query.shape[-1] ** -0.5
if length == cap_end:
return _refiner_segmented(query, key, value, plan, scale)
video_start = length - num_latent_frames * rows_per_frame
if video_start <= cap_end:
return None
spans = plan["spans"]
def video_rows(index):
return slice(video_start + index * rows_per_frame, video_start + (index + 1) * rows_per_frame)
out, lse = _flash_lse(query, key[:, cap_end:], value[:, cap_end:], scale)
if cap_start > 0:
out, lse = _merge_lse(out, lse, *_flash_lse(query, key[:, :cap_start], value[:, :cap_start], scale))
visible = torch.zeros(length, cap_end - cap_start, dtype=torch.bool, device=query.device)
for index, (start, stop) in enumerate(spans):
column = slice(start - cap_start, stop - cap_start)
visible[start:stop, column] = True
visible[video_rows(index), column] = 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)
q_ann = query[:, cap_start:cap_end]
out_ann, lse_ann = _flash_lse(
q_ann, key[:, cap_end:video_start], value[:, cap_end:video_start], scale)
if cap_start > 0:
out_ann, lse_ann = _merge_lse(
out_ann, lse_ann, *_flash_lse(q_ann, key[:, :cap_start], value[:, :cap_start], scale))
own = torch.zeros(cap_end - cap_start, cap_end - cap_start, dtype=torch.bool, device=query.device)
for start, stop in spans:
local = slice(start - cap_start, stop - cap_start)
own[local, local] = True
out_ann, lse_ann = _merge_lse(
out_ann, lse_ann,
*_masked_lse(q_ann, key[:, cap_start:cap_end], value[:, cap_start:cap_end], own, scale))
for index, (start, stop) in enumerate(spans):
local = slice(start - cap_start, stop - cap_start)
rows = video_rows(index)
piece, piece_lse = _flash_lse(q_ann[:, local], key[:, rows], value[:, rows], scale)
merged, merged_lse = _merge_lse(
out_ann[:, local].clone(), lse_ann[:, :, local].clone(), piece, piece_lse)
out_ann[:, local] = merged
lse_ann[:, :, local] = merged_lse
out[:, cap_start:cap_end] = out_ann.to(out.dtype)
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
# Version the guard, not just its presence: a reload re-executes this module but the
# patched class object survives it, so a bare boolean would keep the PREVIOUS closure
# — still pointing at the previous `directed_attention` — installed forever.
if getattr(h3.MiniMaxH3AttnProcessor, "_h3world_directed", None) == DIRECTED_PATCH_VERSION:
return
base_call = getattr(h3.MiniMaxH3AttnProcessor, "_h3world_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._h3world_base_call = base_call
h3.MiniMaxH3AttnProcessor.__call__ = __call__
h3.MiniMaxH3AttnProcessor._h3world_directed = DIRECTED_PATCH_VERSION
# ── 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, TURBO_ERROR
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)
# The turbo LoRA is only *prepared* here — its factors stay resident and are folded in (and
# back out) per request, so both sampling modes are one click apart.
try:
print(f"[gen] {h3_turbo_lora.prepare(pipe.transformer)}", flush=True)
except Exception as error: # noqa: BLE001
traceback.print_exc()
TURBO_ERROR = f"{type(error).__name__}: {error}"
print(f"[gen] WARNING: turbo LoRA unavailable ({TURBO_ERROR})", 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
# Folding (or unfolding) the turbo LoRA rewrites 312 weights through a rank-128 matmul each — a few
# seconds of card, but budget generously since a worker may have to unfold the other state first.
_TURBO_FOLD_ALLOWANCE = 60
def get_duration(
prompt_embeds, tags, plan, image, height, width, num_frames, steps, seed, turbo=False, *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
# Either direction can cost a fold: a turbo request folds it in, and the next base request on a
# warm worker folds it back out. So the allowance rides along whenever the LoRA is loaded at all.
fold = _TURBO_FOLD_ALLOWANCE if TURBO_ERROR is None else 0
return max(60, int((steps * per_step + decode) * _MARGIN) + _PLACEMENT_ALLOWANCE) + fold
# ── Inference ────────────────────────────────────────────────────────────────
@spaces.GPU(duration=get_duration, size=GPU_SIZE)
def _generate(prompt_embeds, tags, plan, image, height, width, num_frames, steps, seed, turbo=False):
"""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")
# In or out, in place, on the card the weights already sit on.
folded = h3_turbo_lora.set_active(PIPE.transformer, turbo)
print(f"[gen] turbo LoRA {'folded in' if folded else 'off'}", flush=True)
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 = 0,
seed: int = 0,
subject: str = DEFAULT_SUBJECT,
directed: bool = True,
mode: str = DEFAULT_MODE,
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, or ``0`` for whatever ``mode`` asks for (28 without the turbo LoRA,
8 with it). 28 is MiniMax-H3's default; 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.
mode: ``"28 steps · no turbo LoRA"`` or ``"8 steps · turbo LoRA"`` — the second folds
``larryvrh/MiniMax-H3-Turbo-Lora``'s few-step LoRA in on top of H3-World and drops the
step count to 8, for a quality-versus-speed comparison at the same seed.
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.")
mode_steps, turbo = MODES.get(str(mode), MODES[DEFAULT_MODE])
if turbo and TURBO_ERROR:
raise gr.Error(f"The turbo LoRA failed to load, so only {MODE_BASE} is available: {TURBO_ERROR}")
steps = int(steps) or mode_steps
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), turbo
)
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)"
)
turbo_note = "turbo LoRA **on** (8-step distillation)" if turbo else "no turbo LoRA"
report = (
f"{width}x{height} · {num_frames} frames ({num_frames / FPS:.2f}s) · {int(steps)} steps · "
f"{turbo_note} · "
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 is a rank-32 LoRA on 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.
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; }
"""
# ── The action-script legend, drawn as keycaps ───────────────────────────────
#
# The script box used to carry its syntax as help text — a comma-separated dump of every preset
# name — which is exactly the thing a picture does better. This renders the same vocabulary as the
# keys each action actually presses. Nothing about the input or `parse_script` changes: the same
# preset names and the same raw key combinations are still what gets typed.
KEY_ROLE = {"W": "move", "A": "move", "S": "move", "D": "move",
"I": "cam", "J": "cam", "K": "cam", "L": "cam", "F": "sharp"}
KEYCAP_CSS = """<style>
.h3k { font-size: 12px; line-height: 1.4; }
.h3k-legend { display: flex; gap: 20px; align-items: flex-end; flex-wrap: wrap; margin: 0 0 12px; }
.h3k-pad { display: flex; flex-direction: column; align-items: center; gap: 3px; }
.h3k-padrow { display: flex; gap: 3px; min-height: 22px; }
.h3k-lab { font-size: 10px; letter-spacing: .06em; text-transform: uppercase; opacity: .6; }
.h3k-key { display: inline-flex; align-items: center; justify-content: center;
min-width: 22px; height: 22px; padding: 0 4px; border-radius: 5px;
font: 600 11px/1 ui-monospace, SFMono-Regular, Menlo, monospace;
border: 1px solid rgba(128,128,128,.40); border-bottom-width: 2px;
background: rgba(128,128,128,.12); }
.h3k-move { background: rgba(99,102,241,.20); border-color: rgba(99,102,241,.55); }
.h3k-cam { background: rgba(16,185,129,.20); border-color: rgba(16,185,129,.55); }
.h3k-sharp { background: rgba(245,158,11,.22); border-color: rgba(245,158,11,.60); }
.h3k-idle { opacity: .5; }
.h3k-grid { display: flex; flex-wrap: wrap; gap: 6px; }
.h3k-chip { display: inline-flex; align-items: center; gap: 5px;
padding: 3px 7px; border-radius: 8px;
border: 1px solid rgba(128,128,128,.28); background: rgba(128,128,128,.06); }
.h3k-name { font-size: 11px; opacity: .8; }
.h3k-code { font: 600 11px/1 ui-monospace, SFMono-Regular, Menlo, monospace; opacity: .85; }
.h3k-sep { opacity: .4; }
</style>"""
# ── The keypad builder ───────────────────────────────────────────────────────
#
# `forward*20, forward-right*10, pan-right-fast*7` is a *performance* written down. Typing it means
# converting an intention ("walk forward for most of it, then swing the camera right") into slot
# arithmetic that has to land on the 37 latent frames a 5s clip carries — the one part of this Space
# that asks the visitor to do the model's bookkeeping. The builder below lets them play the take
# instead: a live WASD/IJKL pad whose held state is sampled once per latent frame, in real time,
# for exactly as long as the clip lasts.
#
# It is strictly an *input method*. Everything downstream still reads the `script` textbox, so typing,
# `gr.Examples` and the preview are untouched; the builder only writes the same string a hand would.
# It also reads the box back (`parse` in the JS below), so an example click or a hand edit redraws
# the timeline rather than leaving it showing a take that is no longer what will be generated.
def _canonical_keys(keys: str) -> str:
"""A key combination in `KEYS` order, so a set of held keys has exactly one spelling."""
return "".join(key for key in KEYS if key in keys)
PRESET_BY_KEYS = {_canonical_keys(keys): name for name, keys in PRESETS.items()}
BUILDER_CONFIG = {
"keys": KEYS,
"role": KEY_ROLE,
"presets": PRESETS,
"names": PRESET_BY_KEYS,
"fps": FPS,
"framesPerChunk": FRAMES_PER_CHUNK,
"latentsPerChunk": LATENTS_PER_CHUNK,
"defaultDuration": DEFAULT_DURATION,
}
BUILDER_CSS = """<style>
.h3b { font-size: 12px; margin: -6px 0 12px; padding: 10px 12px; display: flex;
flex-direction: column; gap: 10px; border-radius: 10px;
border: 1px solid rgba(128,128,128,.28); background: rgba(128,128,128,.05); }
.h3b-top { display: flex; align-items: flex-end; gap: 22px; flex-wrap: wrap; }
.h3b-pads { display: flex; gap: 16px; align-items: flex-end; }
.h3b button.h3b-key { appearance: none; cursor: pointer; color: inherit;
padding: 0 4px; margin: 0; box-shadow: none; }
.h3b button.h3b-key:hover { border-color: rgba(128,128,128,.75); }
.h3b button.h3b-on { border-bottom-width: 1px; transform: translateY(1px); }
.h3b button.h3b-on.h3k-move { background: rgba(99,102,241,.55); }
.h3b button.h3b-on.h3k-cam { background: rgba(16,185,129,.55); }
.h3b button.h3b-on.h3k-sharp { background: rgba(245,158,11,.60); }
.h3b-controls { display: flex; align-items: center; gap: 8px; flex-wrap: wrap; }
.h3b button.h3b-btn { appearance: none; cursor: pointer; color: inherit; margin: 0;
padding: 6px 10px; border-radius: 8px; box-shadow: none;
font-family: inherit; font-size: 11px; font-weight: 600; line-height: 1;
border: 1px solid rgba(128,128,128,.35); background: rgba(128,128,128,.10); }
.h3b button.h3b-btn:hover { background: rgba(128,128,128,.20); }
.h3b button.h3b-ghost { background: transparent; font-weight: 400; opacity: .75; }
.h3b button.h3b-rec { border-color: rgba(239,68,68,.55); background: rgba(239,68,68,.14); }
.h3b-dot { display: inline-block; width: 7px; height: 7px; margin-right: 6px;
border-radius: 50%; background: rgb(239,68,68); vertical-align: 0; }
.h3b-recording button.h3b-rec { background: rgba(239,68,68,.40); }
.h3b-recording .h3b-dot { animation: h3b-blink 1s steps(2, start) infinite; }
@keyframes h3b-blink { 50% { opacity: .15; } }
.h3b-add { display: inline-flex; align-items: center; gap: 6px; }
.h3b input.h3b-num { width: 46px; margin: 0; padding: 5px 4px; text-align: center;
color: inherit; background: transparent; box-shadow: none;
border: 1px solid rgba(128,128,128,.35); border-radius: 6px;
font: 600 11px/1 ui-monospace, SFMono-Regular, Menlo, monospace; }
.h3b-check { display: inline-flex; align-items: center; gap: 4px; opacity: .75; cursor: pointer; }
.h3b-check input { margin: 0; }
/* the tick grid is one slot wide, so the track reads as N boxes however the runs fall */
.h3b-track { display: flex; height: 30px; border-radius: 8px; overflow: hidden;
border: 1px solid rgba(128,128,128,.30); background-color: rgba(128,128,128,.06);
background-image: repeating-linear-gradient(to right, rgba(128,128,128,.22) 0 1px,
transparent 1px calc(100% / var(--slots, 37))); }
.h3b-recording .h3b-track { box-shadow: 0 0 0 2px rgba(239,68,68,.35); }
.h3b-seg { flex-basis: 0; min-width: 0; overflow: hidden; display: flex; gap: 2px;
align-items: center; justify-content: center; border-right: 1px solid rgba(128,128,128,.35); }
.h3b-seg:last-child { border-right: none; }
.h3b-move { background: rgba(99,102,241,.16); }
.h3b-cam { background: rgba(16,185,129,.16); }
.h3b-free { background: repeating-linear-gradient(45deg,
rgba(128,128,128,.13) 0 5px, transparent 5px 10px); }
.h3b-track .h3k-key { min-width: 15px; height: 16px; padding: 0 2px;
font-size: 9px; border-bottom-width: 1px; }
.h3b-n { font: 600 10px/1 ui-monospace, SFMono-Regular, Menlo, monospace; opacity: .7; }
.h3b-foot { display: flex; align-items: baseline; gap: 8px; flex-wrap: wrap; }
.h3b-count { font: 600 11px/1 ui-monospace, SFMono-Regular, Menlo, monospace; }
.h3b-note { font-size: 11px; opacity: .7; }
.h3b-spacer { flex: 1 1 auto; }
.h3b-hint { font-size: 11px; opacity: .55; }
</style>"""
def _pad_key(key: str) -> str:
return (f'<button type="button" class="h3k-key h3k-{KEY_ROLE[key]} h3b-key" '
f'data-key="{key}">{key}</button>')
def _pad(keys_top: str, keys_bottom: str, label: str) -> str:
top = "".join(_pad_key(key) for key in keys_top)
bottom = "".join(_pad_key(key) for key in keys_bottom)
return (f'<div class="h3k-pad"><span class="h3k-padrow">{top}</span>'
f'<span class="h3k-padrow">{bottom}</span><span class="h3k-lab">{label}</span></div>')
BUILDER_HTML = (
KEYCAP_CSS
+ BUILDER_CSS
+ '<div class="h3b"><div class="h3b-top"><div class="h3b-pads">'
+ _pad("W", "ASD", "move")
+ _pad("I", "JKL", "aim camera")
+ _pad("", "F", "sharp")
+ '</div><div class="h3b-controls">'
+ '<button type="button" class="h3b-btn h3b-rec" data-act="record">'
+ '<span class="h3b-dot"></span><span data-label>Record</span></button>'
+ '<label class="h3b-check"><input type="checkbox" data-slow> half speed</label>'
+ '<span class="h3b-add"><button type="button" class="h3b-btn" data-act="add">+ hold for</button>'
+ '<input class="h3b-num" type="number" min="1" max="99" step="1" value="6" data-num>'
+ '<span class="h3b-hint">slots</span></span>'
+ '<button type="button" class="h3b-btn h3b-ghost" data-act="undo">Undo</button>'
+ '<button type="button" class="h3b-btn h3b-ghost" data-act="clear">Clear</button>'
+ '</div></div><div class="h3b-track" data-track></div>'
+ '<div class="h3b-foot"><span class="h3b-count" data-count></span>'
+ '<span class="h3b-note" data-note></span><span class="h3b-spacer"></span>'
+ '<span class="h3b-hint">Record plays the take in real time and replaces the timeline'
+ " &middot; Esc stops</span></div></div>"
)
BUILDER_JS = (
"const CFG = " + json.dumps(BUILDER_CONFIG, ensure_ascii=False) + ";\n"
+ r"""
"use strict";
const root = element.querySelector(".h3b");
if (!root) return;
if (window.__h3bTeardown) { try { window.__h3bTeardown(); } catch (err) {} }
const KEYS = CFG.keys, ROLE = CFG.role, PRESETS = CFG.presets, NAMES = CFG.names;
const track = root.querySelector("[data-track]");
const countEl = root.querySelector("[data-count]");
const noteEl = root.querySelector("[data-note]");
const labelEl = root.querySelector("[data-label]");
const numEl = root.querySelector("[data-num]");
const slowEl = root.querySelector("[data-slow]");
// The timeline, run-length encoded exactly the way the script string is: [{k: "W", n: 20}, ...].
// `held` is what the pad currently shows down — latched by clicking a cap, momentary from the
// keyboard. A take samples `held` once per latent frame; "+ hold for n" writes it n times at once.
let runs = [];
let held = new Set();
let take = null;
let seen = null; // the last textbox value we have accounted for, ours or the user's
let unknown = false; // the textbox holds something we could not lay out
function canon(keys) {
let out = "";
for (const key of KEYS) if (keys.has(key)) out += key;
return out;
}
function nameOf(keys) {
return NAMES[keys] !== undefined ? NAMES[keys] : keys;
}
function total() {
let sum = 0;
for (const run of runs) sum += run.n;
return sum;
}
function push(keys, n) {
if (n <= 0) return;
const last = runs[runs.length - 1];
if (last && last.k === keys) last.n += n; else runs.push({ k: keys, n: n });
}
// Read the live duration off the slider rather than caching it: it is the thing that decides how
// many slots a take has to fill, and it sits in an accordion the visitor can open mid-build.
function duration() {
const input = document.querySelector("#h3-duration input[type=range]")
|| document.querySelector("#h3-duration input");
const value = input ? parseFloat(input.value) : NaN;
return (isFinite(value) && value > 0) ? value : CFG.defaultDuration;
}
function slotCount() {
let frames = Math.max(1, Math.round(duration() * CFG.fps));
while (frames % CFG.framesPerChunk !== CFG.latentsPerChunk) frames += 1;
return Math.floor((frames - CFG.latentsPerChunk) / CFG.framesPerChunk) * CFG.latentsPerChunk + 2;
}
function caps(keys) {
if (!keys) return '<span class="h3k-key h3k-idle">&mdash;</span>';
let out = "";
for (const key of keys) out += '<span class="h3k-key h3k-' + ROLE[key] + '">' + key + "</span>";
return out;
}
function emit() {
return runs.map(function (run) { return nameOf(run.k) + "*" + run.n; }).join(", ");
}
function render() {
const slots = slotCount();
track.style.setProperty("--slots", slots);
let html = "";
for (const run of runs) {
const tint = /[WASD]/.test(run.k) ? "move" : (run.k ? "cam" : "idle");
html += '<span class="h3b-seg h3b-' + tint + '" style="flex-grow:' + run.n + '" title="'
+ nameOf(run.k) + " × " + run.n + '">' + caps(run.k)
+ '<span class="h3b-n">' + run.n + "</span></span>";
}
const laid = total();
if (laid < slots) {
html += '<span class="h3b-seg h3b-free" style="flex-grow:' + (slots - laid) + '"></span>';
}
track.innerHTML = html;
countEl.textContent = laid + " / " + slots + " slots";
// Say what `parse_script` will actually do with a timeline that does not land on the slot count,
// rather than calling it an error: both short and long scripts are legal, they just get padded
// by holding the last action, or cut at the end of the clip.
noteEl.textContent = laid === 0
? (unknown ? "the script above is not one the builder can lay out" : "nothing laid down yet")
: laid < slots ? "the last action holds for the remaining " + (slots - laid)
: laid > slots ? "the last " + (laid - slots) + " run past the end and get dropped"
: "an exact take";
for (const button of root.querySelectorAll("[data-key]")) {
button.classList.toggle("h3b-on", held.has(button.dataset.key));
}
}
function box() {
return document.querySelector("#h3-script textarea") || document.querySelector("#h3-script input");
}
// Gradio's textbox is a Svelte `bind:value`, which listens for `input` — so setting `.value` and
// dispatching one is what makes the Python side (and the `.change` preview) see the new script.
function write() {
const target = box();
if (!target) return;
const text = emit();
seen = text;
unknown = false;
target.value = text;
target.dispatchEvent(new Event("input", { bubbles: true }));
}
// The mirror image of `parse_script`, so the timeline shows what the *textbox* means — including
// an example click or a hand edit. Returns null for anything it cannot name, which leaves the
// error reporting to `preview()` where it already lives.
function parse(text, slots) {
const items = [];
for (let chunk of String(text == null ? "" : text).split(/[,\n;]+/)) {
chunk = chunk.trim();
if (!chunk) continue;
const match = /^(.*?)(?:\s*[*x×]\s*(\d+)\s*)?$/.exec(chunk);
if (!match) return null;
const name = (match[1] || "").trim();
const count = match[2] ? parseInt(match[2], 10) : null;
const lowered = name.toLowerCase().replace(/[ _]/g, "-");
let keys;
// through `canon` either way: a combination has one spelling in here, so `nameOf` can find it
// again. `parse_script` is order-blind, but `PRESETS` is not written in key order (`back-left`
// is `SA`), and an uncanonicalised hit would come back out of `emit` as raw keys.
if (Object.prototype.hasOwnProperty.call(PRESETS, lowered)) {
keys = canon(new Set(PRESETS[lowered].split("")));
} else {
const raw = name.toUpperCase().replace(/[ +]/g, "");
const letters = raw.split("");
if (raw && letters.every(function (key) { return KEYS.indexOf(key) >= 0; })) {
keys = canon(new Set(letters));
} else if (["none", "idle", "stop"].indexOf(lowered) >= 0) {
keys = "";
} else {
return null;
}
}
items.push([keys, count]);
}
if (!items.length) items.push(["", null]);
const counts = items.map(function (item) { return item[1] || 0; });
const free = [];
items.forEach(function (item, index) { if (!item[1]) free.push(index); });
if (free.length) {
let assigned = 0;
counts.forEach(function (count) { assigned += count; });
const remaining = Math.max(0, slots - assigned);
const base = Math.floor(remaining / free.length), extra = remaining % free.length;
free.forEach(function (index, position) { counts[index] = base + (position < extra ? 1 : 0); });
}
const out = [];
items.forEach(function (item, index) {
if (counts[index] <= 0) return;
const last = out[out.length - 1];
if (last && last.k === item[0]) last.n += counts[index];
else out.push({ k: item[0], n: counts[index] });
});
let sum = 0;
out.forEach(function (run) { sum += run.n; });
if (!out.length) out.push({ k: items[0][0], n: slots });
else if (sum < slots) out[out.length - 1].n += slots - sum;
else while (sum > slots) {
const last = out[out.length - 1];
const cut = Math.min(sum - slots, last.n);
last.n -= cut;
sum -= cut;
if (!last.n) out.pop();
}
return out;
}
// ── the take ───────────────────────────────────────────────────────────────
// One slot per latent frame, played out over the clip's own duration: a 5s take is 37 slots in 5
// seconds, so holding W for a second is worth about seven of them and the timeline you record is
// the timing you will watch back. Driven off `performance.now()` rather than a tick count, so a
// dropped frame moves the playhead instead of stretching the take.
function tick() {
if (!take) return;
const slots = slotCount();
const per = (duration() * 1000 / slots) * (slowEl && slowEl.checked ? 2 : 1);
const filled = Math.min(slots, Math.floor((performance.now() - take.t0) / per));
if (filled > take.filled) {
push(canon(held), filled - take.filled);
take.filled = filled;
render();
}
if (take.filled >= slots) { stop(); return; }
take.raf = requestAnimationFrame(tick);
}
function start() {
runs = [];
unknown = false;
const active = document.activeElement;
if (active && active !== document.body && active.blur) active.blur();
take = { t0: performance.now(), filled: 0, raf: 0 };
root.classList.add("h3b-recording");
if (labelEl) labelEl.textContent = "Stop";
render();
take.raf = requestAnimationFrame(tick);
}
function stop() {
if (!take) return;
cancelAnimationFrame(take.raf);
take = null;
root.classList.remove("h3b-recording");
if (labelEl) labelEl.textContent = "Record";
render();
write();
}
function onClick(event) {
const target = event.target instanceof Element ? event.target : null;
if (!target) return;
const cap = target.closest("[data-key]");
if (cap) {
const key = cap.dataset.key;
if (held.has(key)) held.delete(key); else held.add(key);
render();
return;
}
const button = target.closest("[data-act]");
if (!button) return;
const action = button.dataset.act;
if (action === "record") { if (take) stop(); else start(); return; }
if (action === "add") push(canon(held), Math.max(1, parseInt(numEl.value, 10) || 1));
else if (action === "undo") runs.pop();
else if (action === "clear") { runs = []; held.clear(); }
render();
write();
}
root.addEventListener("click", onClick);
// Keys are only intercepted while a take is running, or while the focus is inside the builder —
// otherwise `W` would stop reaching the prompt box, which is a textarea a visitor spends far more
// time in than this pad.
function keyed(event, down) {
if (event.metaKey || event.ctrlKey || event.altKey) return;
if (!take) {
const active = document.activeElement;
const inside = active && root.contains(active)
&& active.tagName !== "INPUT" && active.tagName !== "TEXTAREA";
if (!inside) return;
}
if (event.key === "Escape") { if (down) stop(); return; }
const key = (event.key || "").toUpperCase();
if (key.length !== 1 || KEYS.indexOf(key) < 0) return;
event.preventDefault();
if (down) held.add(key); else held.delete(key);
render();
}
function onKeyDown(event) { if (!event.repeat) keyed(event, true); }
function onKeyUp(event) { keyed(event, false); }
document.addEventListener("keydown", onKeyDown, true);
document.addEventListener("keyup", onKeyUp, true);
// Gradio sets a textbox's value straight on the DOM node, which fires no event and mutates no
// attribute — neither a listener nor a MutationObserver would see `gr.Examples` fill the script in.
// A string compare a few times a second is the one thing that does, and it keeps the timeline
// honest about what is actually queued to generate.
let lastSlots = slotCount();
const poll = setInterval(function () {
if (!root.isConnected) { teardown(); return; }
const target = box();
const slots = slotCount();
if (target && target.value !== seen) {
seen = target.value;
if (!take) {
const parsed = parse(seen, slots);
unknown = parsed === null;
runs = parsed || [];
render();
}
} else if (slots !== lastSlots) {
render();
}
lastSlots = slots;
}, 400);
function teardown() {
clearInterval(poll);
root.removeEventListener("click", onClick);
document.removeEventListener("keydown", onKeyDown, true);
document.removeEventListener("keyup", onKeyUp, true);
if (take) { cancelAnimationFrame(take.raf); take = null; }
if (window.__h3bTeardown === teardown) window.__h3bTeardown = null;
}
window.__h3bTeardown = teardown;
const initial = box();
seen = initial ? initial.value : null;
if (seen !== null) {
const parsed = parse(seen, slotCount());
unknown = parsed === null;
runs = parsed || [];
}
render();
"""
)
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,
elem_id="h3-script",
)
gr.HTML(
BUILDER_HTML,
js_on_load=BUILDER_JS,
elem_id="h3-builder",
apply_default_css=False,
)
mode = gr.Radio(
label="Sampling",
choices=list(MODES),
value=DEFAULT_MODE,
)
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,
# the builder reads this back to know how many slots a take has to fill
elem_id="h3-duration",
)
steps_slider = gr.Slider(
label="Steps",
minimum=4,
maximum=50,
step=1,
value=DEFAULT_STEPS,
info="Follows the sampling mode; 50 is the released H3-World configuration.",
)
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, mode]
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])
# Picking a sampling mode moves the Steps slider to that mode's own count; the slider stays an
# override, so 50 (the released configuration) is still reachable with the LoRA either way.
mode.change(lambda choice: gr.update(value=MODES.get(choice, MODES[DEFAULT_MODE])[0]), [mode], [steps_slider])
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],
)