"""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 ``": " + + 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"
Action timeline\n\n{summarize(sequence, subject)}\n\n
" ) 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
[ LoRA ]   [ base model ]   [ training data ]
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 = """""" # ── 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 = """""" def _pad_key(key: str) -> str: return (f'') 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'
{top}' f'{bottom}{label}
') BUILDER_HTML = ( KEYCAP_CSS + BUILDER_CSS + '
' + _pad("W", "ASD", "move") + _pad("I", "JKL", "aim camera") + _pad("", "F", "sharp") + '
' + '' + '' + '' + '' + 'slots' + '' + '' + '
' + '
' + '' + 'Record plays the take in real time and replaces the timeline' + " · Esc stops
" ) 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 '—'; let out = ""; for (const key of keys) out += '' + key + ""; 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 += '' + caps(run.k) + '' + run.n + ""; } const laid = total(); if (laid < slots) { html += ''; } 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], )