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