multimodalart's picture
multimodalart HF Staff
fix: dark-mode text color override for Citrus theme
e5fe8ce verified
Raw
History Blame Contribute Delete
19.6 kB
"""MiniWorld — camera-controlled video world model simulator (ZeroGPU).
Mirrors the authors' reference inference path
python -m miniworld.sample --dataset re10k --custom_camera_trajectory ...
one-to-one: a single init image is Wan2.2-VAE-encoded into the clean seed
latent, a procedural camera path is turned into ray-encoding conditioning, and
the AR-diffusion denoiser rolls the world forward chunk-by-chunk with a
position-bounded streaming KV cache and streaming VAE decode.
"""
from __future__ import annotations
import math
import os
import tempfile
import time
os.environ.setdefault("PYTORCH_CUDA_ALLOC_CONF", "expandable_segments:True")
import spaces # noqa: E402 (must precede any torch / CUDA work)
import gradio as gr # noqa: E402
import numpy as np # noqa: E402
import torch # noqa: E402
from einops import rearrange # noqa: E402
from huggingface_hub import hf_hub_download # noqa: E402
from PIL import Image # noqa: E402
# MiniWorld checkpoints are plain `torch.save` dicts that carry a `meta` blob of
# plain-python objects next to the tensors, so they need the full unpickler.
_ORIG_TORCH_LOAD = torch.load
def _torch_load(*args, **kwargs):
kwargs.setdefault("weights_only", False)
return _ORIG_TORCH_LOAD(*args, **kwargs)
torch.load = _torch_load
from miniworld.conditioning.actions import ( # noqa: E402
ConditioningConfig,
build_cond_seq_for_batch,
)
from miniworld.conditioning.trajectories import build_custom_trajectory # noqa: E402
from miniworld.denoiser import DenoiserConfig, build_denoiser_from_mode # noqa: E402
from miniworld.vae.codec import StreamingVAEDecoder, vae_encode # noqa: E402
from miniworld.vae.wan22_vae import Wan2_2_VAE # noqa: E402
# --------------------------------------------------------------------------- #
# Constants (match scripts/sample_re10k.sh) #
# --------------------------------------------------------------------------- #
MINIWORLD_REPO = "zhaoyian01/MiniWorld"
MINIWORLD_CKPT = "MiniWorld_1b_re10k.pt"
VAE_REPO = "Wan-AI/Wan2.2-TI2V-5B"
VAE_FILE = "Wan2.2_VAE.pth"
RESIZE_H, RESIZE_W = 240, 320
SPATIAL_DOWNSAMPLE = 16
LATENT_CHANNELS = 48
POSE_ENC_FREQ = 15
DF_CHUNK_SIZE = 2
DF_ARDIFF_STEP = 5
STREAM_INFLIGHT_CHUNKS = 8
STREAM_MAX_CACHE_CHUNKS = 24
STREAM_SINK_SIZE = 1
SAMPLE_HISTORY_LEN = 1
SAVE_FPS = 8
MAX_SEED = np.iinfo(np.int32).max
H_LAT, W_LAT = RESIZE_H // SPATIAL_DOWNSAMPLE, RESIZE_W // SPATIAL_DOWNSAMPLE
TRAJECTORIES = [
"orbit_right",
"orbit_left",
"pan_right",
"pan_left",
"forward",
"backward",
"tilt_up",
"tilt_down",
"spiral",
"zoom_in",
"zoom_out",
"static",
]
# --------------------------------------------------------------------------- #
# Model construction #
# --------------------------------------------------------------------------- #
print("Fetching Wan2.2 VAE ...", flush=True)
vae_path = hf_hub_download(VAE_REPO, VAE_FILE)
print("Fetching MiniWorld-1B (RealEstate10K) ...", flush=True)
ckpt_path = hf_hub_download(MINIWORLD_REPO, MINIWORLD_CKPT)
_ckpt = torch.load(ckpt_path, map_location="cpu")
_meta: dict = {}
_weights = None
if isinstance(_ckpt, dict):
# `miniworld/sample.py` expects a training checkpoint wrapper; the *released*
# weights are a bare state dict of `net.*` tensors, so support both.
for _key in ("ema_model", "model", "ema", "state_dict", "module"):
cand = _ckpt.get(_key)
if isinstance(cand, dict) and cand:
_weights = cand
_meta = _ckpt.get("meta") or {}
print(f"[Checkpoint] using wrapped weights under {_key!r}", flush=True)
break
if _weights is None and any(
isinstance(k, str) and k.startswith("net.") for k in _ckpt
):
_weights = _ckpt
print("[Checkpoint] bare state dict (no training wrapper)", flush=True)
if _weights is None:
raise RuntimeError(
"Unrecognised MiniWorld checkpoint layout; top-level keys: "
f"{list(_ckpt)[:8] if isinstance(_ckpt, dict) else type(_ckpt)}"
)
def _resolve_latent_frames() -> int:
for key in ("latent_frames", "trained_num_frames"):
val = int(_meta.get(key, 0) or 0)
if val > 0:
return val
freqs = _weights.get("net.feat_rope.freqs_cos")
tokens_per_frame = H_LAT * W_LAT
if freqs is not None and freqs.shape[0] % tokens_per_frame == 0:
return int(freqs.shape[0] // tokens_per_frame)
raise RuntimeError("Cannot determine the checkpoint's latent frame count")
LATENT_FRAMES = _resolve_latent_frames()
WM_MODEL = str(_meta.get("wm_model") or "1B")
TRAINED_NUM_FRAMES = int(_meta.get("trained_num_frames", 0) or 0) or LATENT_FRAMES
MAX_TOTAL_LEN = TRAINED_NUM_FRAMES
print(
f"[Checkpoint] wm_model={WM_MODEL} latent_frames={LATENT_FRAMES} "
f"trained_num_frames={TRAINED_NUM_FRAMES}",
flush=True,
)
denoiser = build_denoiser_from_mode(
DenoiserConfig(
wm_model=WM_MODEL,
latent_size=(H_LAT, W_LAT),
latent_channels=LATENT_CHANNELS,
latent_frames=LATENT_FRAMES,
wm_mlp_ratio=4.0,
wm_use_qknorm=True,
wm_use_checkpoint=False,
cond_dim=4 * 6 * 2 * POSE_ENC_FREQ,
cond_per_token=True,
adaln_mode="adaln_lora",
cond_dropout_prob=0.1,
timestep_baseshift=2.667,
timestep_shift=-1.0,
num_sampling_steps=100,
cfg_scale=2.0,
cfg_interval_min=0.2,
cfg_interval_max=1.0,
df_chunk_size=DF_CHUNK_SIZE,
df_ardiff_step=DF_ARDIFF_STEP,
)
).eval()
_missing, _unexpected = denoiser.load_state_dict(_weights, strict=False)
if _missing or _unexpected:
raise RuntimeError(
f"MiniWorld checkpoint does not match the built model.\n"
f"missing ({len(_missing)}): {_missing[:12]}\n"
f"unexpected ({len(_unexpected)}): {_unexpected[:12]}"
)
denoiser.trained_num_frames = TRAINED_NUM_FRAMES
print(f"[Checkpoint] loaded: all {len(_weights)} keys matched", flush=True)
del _ckpt, _weights
denoiser = denoiser.to("cuda")
vae = Wan2_2_VAE(vae_pth=vae_path, device="cuda")
vae.model.requires_grad_(False)
vae.model.eval()
_COND_CFG = ConditioningConfig(
use_pose_cond=True, use_action_cond=False, pose_enc_freq=POSE_ENC_FREQ
)
# --------------------------------------------------------------------------- #
# Helpers #
# --------------------------------------------------------------------------- #
def _prepare_init_frame(image) -> torch.Tensor:
"""PIL / ndarray -> ``(H, W, C)`` float32 in [-1, 1] (== `load_init_image`)."""
if image is None:
raise gr.Error("Please provide an initial frame.")
if isinstance(image, np.ndarray):
image = Image.fromarray(image)
arr = np.asarray(image.convert("RGB"), dtype=np.float32) / 255.0
img = torch.from_numpy(arr).permute(2, 0, 1).unsqueeze(0)
if tuple(img.shape[-2:]) != (RESIZE_H, RESIZE_W):
img = torch.nn.functional.interpolate(
img, size=(RESIZE_H, RESIZE_W), mode="bilinear", align_corners=False
)
return img.squeeze(0).permute(1, 2, 0).contiguous() * 2.0 - 1.0
def _write_mp4(frames: np.ndarray, fps: int) -> str:
import imageio.v2 as imageio
path = os.path.join(tempfile.mkdtemp(), "miniworld.mp4")
writer = imageio.get_writer(
path,
fps=fps,
codec="libx264",
quality=8,
macro_block_size=1,
ffmpeg_params=["-pix_fmt", "yuv420p"],
)
try:
for frame in frames:
writer.append_data(frame)
finally:
writer.close()
return path
def _rollout_tflops(total_len: int, steps: int) -> float:
"""Replay the streaming schedule to cost a rollout in DiT TFLOPs.
The AR-diffusion schedule is not linear in ``total_len`` (short rollouts
that fit inside the in-flight window run the *full* sampler), so the ZeroGPU
reservation is derived from the same bookkeeping the sampler does.
"""
chunk = DF_CHUNK_SIZE
ar = DF_ARDIFF_STEP
inflight = STREAM_INFLIGHT_CHUNKS
max_cache_frames = STREAM_MAX_CACHE_CHUNKS * chunk
total_chunks = (total_len + chunk - 1) // chunk
eff = steps if total_chunks <= inflight else min(steps, inflight * ar)
prev = [0] * total_chunks
masks = []
n_rows = 0
while any(p != eff for p in prev):
row = [0] * total_chunks
for i in range(total_chunks):
row[i] = prev[i] + 1 if (i == 0 or prev[i - 1] == eff) else row[i - 1] - ar
row[i] = max(0, min(eff, row[i]))
masks.append([row[i] != prev[i] for i in range(total_chunks)])
prev = row
n_rows += 1
if n_rows > 4000: # safety valve
break
terminal = min(inflight, total_chunks)
committed = 0
cache_frames = 0
tflops = 0.0
# per-forward TFLOPs for a 1B DiT: 0.6 per query frame (linear layers) plus
# 0.01548 per (query frame x key frame) (attention), at 300 tokens/frame.
for step in range(n_rows):
if terminal < total_chunks and masks[step][terminal]:
terminal += 1
win_sc = max(0, terminal - inflight)
while committed < win_sc:
frames = min((committed + 1) * chunk, total_len) - committed * chunk
tflops += 2 * frames * (0.6 + 0.01548 * (cache_frames + frames))
cache_frames = min(cache_frames + frames, max_cache_frames)
committed += 1
if terminal <= win_sc:
continue
q_frames = min(terminal * chunk, total_len) - win_sc * chunk
tflops += 2 * q_frames * (0.6 + 0.01548 * (cache_frames + q_frames))
return tflops
# Calibrated on this Space's ZeroGPU H200 slice: measured 46.0s / 72.3s / 157.2s
# at total_len 20 / 32 / 64 against 2395 / 3599 / 7703 modelled TFLOP, i.e. a
# very clean 47.8 TFLOP/s (streaming VAE decode overlaps the denoiser, so it
# needs no separate term).
_TFLOPS_PER_SEC = 47.8
_VAE_SEC_PER_LATENT_FRAME = 0.0
_FIXED_OVERHEAD_SEC = 2.0
def _duration(*args, **kwargs) -> int:
total_len, steps = 32, 100
if len(args) >= 4:
total_len = int(args[3])
if len(args) >= 8:
steps = int(args[7])
total_len = int(kwargs.get("total_len", total_len))
steps = int(kwargs.get("num_sampling_steps", steps))
total_len = max(4, min(total_len, 64))
est = (
_FIXED_OVERHEAD_SEC
+ _rollout_tflops(total_len, steps) / _TFLOPS_PER_SEC
+ _VAE_SEC_PER_LATENT_FRAME * total_len
)
return int(min(400, math.ceil(est * 1.15)))
# --------------------------------------------------------------------------- #
# Inference #
# --------------------------------------------------------------------------- #
@spaces.GPU(duration=_duration)
@torch.no_grad()
def simulate(
image,
trajectory: str = "orbit_right",
magnitude: float = 3.0,
total_len: int = 32,
seed: int = 0,
randomize_seed: bool = True,
cfg_scale: float = 2.0,
num_sampling_steps: int = 100,
focal_norm: float = 0.5,
progress=gr.Progress(track_tqdm=True),
):
init_frame = _prepare_init_frame(image)
total_len = max(4, min(int(total_len), MAX_TOTAL_LEN))
if randomize_seed:
seed = int(np.random.randint(0, MAX_SEED))
seed = int(seed) % (MAX_SEED + 1)
device = torch.device("cuda")
denoiser.cfg_scale = float(cfg_scale)
denoiser.steps = int(num_sampling_steps)
# `build_custom_trajectory` spreads the whole path evenly over the rollout,
# so a fixed magnitude means *faster* per-frame motion in a shorter clip.
# The authors' guidance is to scale it linearly with length to keep the
# apparent speed constant (3.0 @ total_len 64 -> 4.5 @ 96), so the slider is
# exposed as a speed in "magnitude at 64 latent frames" units.
magnitude_eff = float(magnitude) * total_len / 64.0
videos = init_frame.unsqueeze(0).unsqueeze(0).to(device) # (1, 1, H, W, C)
poses = (
build_custom_trajectory(
trajectory,
num_frames=4 * (total_len - 1) + 1,
focal_norm=float(focal_norm),
magnitude=magnitude_eff,
)
.unsqueeze(0)
.to(device)
)
generator = torch.Generator(device="cpu").manual_seed(seed)
noise = torch.randn(
1, LATENT_CHANNELS, total_len, H_LAT, W_LAT,
generator=generator, dtype=torch.float32,
).to(device)
start = time.perf_counter()
with torch.autocast(device_type="cuda", dtype=torch.bfloat16, enabled=True):
latents = vae_encode(
vae, rearrange(videos, "b t h w c -> b c t h w").contiguous()
)
_, c_latent, _, h_lat, w_lat = latents.shape
full_latents = latents.new_zeros(1, c_latent, total_len, h_lat, w_lat)
full_latents[:, :, :1] = latents[:, :, :1]
cond_seq = build_cond_seq_for_batch(
cfg=_COND_CFG,
poses=poses,
actions=None,
t_latent=total_len,
h_lat=h_lat,
w_lat=w_lat,
)
_, pred_rgb = denoiser.generate_eval_latents_streaming(
full_latents,
cond_seq,
total_len=total_len,
history_len=SAMPLE_HISTORY_LEN,
max_cache_chunks=STREAM_MAX_CACHE_CHUNKS,
inflight_chunks=STREAM_INFLIGHT_CHUNKS,
sink_frames=STREAM_SINK_SIZE,
stream_decoder=StreamingVAEDecoder(vae),
noise=noise.to(full_latents.dtype),
)
elapsed = time.perf_counter() - start
video = ((pred_rgb[0].permute(1, 2, 3, 0).clamp(-1, 1) + 1.0) * 127.5).to(
torch.uint8
)
frames = video.cpu().numpy()
path = _write_mp4(frames, SAVE_FPS)
n = int(frames.shape[0])
print(f"[Timing] total_len={total_len} steps={num_sampling_steps}: "
f"{elapsed:.2f}s (reserved {_duration(None, trajectory, magnitude, total_len, seed, False, cfg_scale, num_sampling_steps)}s)",
flush=True)
return (
path,
seed,
f"**{n} frames** @ {SAVE_FPS} fps ({n / SAVE_FPS:.1f}s) · "
f"{total_len} latent frames · `{trajectory}` · speed {magnitude:g} "
f"(magnitude {magnitude_eff:.2f}) · "
f"seed `{seed}` · {elapsed:.1f}s of GPU time",
)
# --------------------------------------------------------------------------- #
# UI #
# --------------------------------------------------------------------------- #
CSS = "#col-container { max-width: 1060px; margin: 0 auto; } .dark .gradio-container { color: var(--body-text-color); }"
with gr.Blocks() as demo:
with gr.Column(elem_id="col-container"):
gr.Markdown(
f"""
# 🌍 MiniWorld · camera-controlled world model
Hand MiniWorld-1B **one frame and a camera path** and it rolls the
world forward autoregressively — no text prompt, no reference video,
no ground-truth poses. A position-bounded streaming KV cache plus
causal Wan2.2 VAE decoding keep the horizon open, so a rollout can
run to {4 * (MAX_TOTAL_LEN - 1) + 1} frames from a
{TRAINED_NUM_FRAMES}-latent-frame checkpoint.
Model: [`zhaoyian01/MiniWorld`](https://huggingface.co/zhaoyian01/MiniWorld)
(RealEstate10K, {WM_MODEL}) · Paper:
[MiniWorld: Democratizing the Training of Video World Models from Scratch](https://huggingface.co/papers/2608.01127)
· Code: [zhao-yian/MiniWorld](https://github.com/zhao-yian/MiniWorld)
"""
)
with gr.Row():
with gr.Column():
image = gr.Image(
label="Initial frame",
type="pil",
height=270,
sources=["upload", "clipboard"],
)
trajectory = gr.Dropdown(
label="Camera trajectory",
choices=TRAJECTORIES,
value="orbit_right",
)
magnitude = gr.Slider(
label="Camera speed",
minimum=0.5,
maximum=8.0,
step=0.5,
value=3.0,
info="3.0 is the paper's default: clear, stable motion. "
"1.0 is nearly static, 8.0 breaks down late. Scaled "
"internally with rollout length so the apparent speed "
"stays constant.",
)
total_len = gr.Slider(
label="Rollout length (latent frames)",
minimum=20,
maximum=MAX_TOTAL_LEN,
step=4,
value=32,
info=f"Each latent frame decodes to 4 RGB frames at {SAVE_FPS} fps; "
f"{MAX_TOTAL_LEN}{4 * (MAX_TOTAL_LEN - 1) + 1} frames.",
)
run_button = gr.Button("Simulate", variant="primary")
with gr.Column():
result = gr.Video(
label="Rollout", autoplay=True, loop=True, height=270
)
info = gr.Markdown()
with gr.Accordion("Advanced settings", open=False):
with gr.Row():
seed = gr.Slider(
label="Seed", minimum=0, maximum=MAX_SEED, step=1, value=0
)
randomize_seed = gr.Checkbox(label="Randomize seed", value=True)
with gr.Row():
cfg_scale = gr.Slider(
label="Guidance scale (CFG)",
minimum=1.0,
maximum=5.0,
step=0.1,
value=2.0,
)
num_sampling_steps = gr.Slider(
label="Sampling steps",
minimum=20,
maximum=100,
step=10,
value=100,
info="Effective steps per chunk are capped by the streaming "
"schedule at in-flight chunks × AR step = 40.",
)
focal_norm = gr.Slider(
label="Normalized focal length",
minimum=0.3,
maximum=1.2,
step=0.05,
value=0.5,
info="0.5 matches typical RealEstate10K intrinsics; smaller = wider FOV.",
)
gr.Examples(
examples=[
["examples/kitchen.png", "orbit_right", 3.0, 32],
["examples/deck.png", "forward", 3.0, 32],
["examples/garden.png", "pan_left", 3.0, 32],
],
inputs=[image, trajectory, magnitude, total_len],
outputs=[result, seed, info],
fn=simulate,
cache_examples=True,
cache_mode="lazy",
)
gr.on(
triggers=[run_button.click],
fn=simulate,
inputs=[
image,
trajectory,
magnitude,
total_len,
seed,
randomize_seed,
cfg_scale,
num_sampling_steps,
focal_norm,
],
outputs=[result, seed, info],
)
demo.queue().launch(theme=gr.themes.Citrus(), css=CSS, mcp_server=True)