Spaces:
Running on Zero
Running on Zero
multimodalart HF Staff
Drop the static keycap legend, now that the builder replaces it
f984226 verified Download app.py from hugging-apps/h3-world-action-demo: direct link, hf CLI and curl.
- Browser
- Download file 73.6 kB
-
https://huggingface.co/spaces/hugging-apps/h3-world-action-demo/resolve/main/app.py
- Command line
-
hf download hf://spaces/hugging-apps/h3-world-action-demo/app.py
-
curl -L -o app.py https://huggingface.co/spaces/hugging-apps/h3-world-action-demo/resolve/main/app.py
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 ─────────────────────────────── | |
| 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 ────────────────────────────────────────────────────── | |
| 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 ──────────────────────────────────────────────────────────────── | |
| 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> | |
| <a href="https://huggingface.co/MiniMaxAI/MiniMax-H3" target="_blank" rel="noopener"><strong>[ base model ]</strong></a> | |
| <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' | |
| + " · 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">—</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], | |
| ) | |