Spaces:
Sleeping
Sleeping
File size: 64,840 Bytes
a6cc5f0 3bd48a2 a6cc5f0 3bd48a2 41c4616 3bd48a2 a6cc5f0 71ca7ab a6cc5f0 3bd48a2 a6cc5f0 3bd48a2 82eed55 a6cc5f0 3bd48a2 a6cc5f0 3bd48a2 a6cc5f0 3bd48a2 71ca7ab a6cc5f0 71ca7ab a6cc5f0 3bd48a2 a6cc5f0 71ca7ab a6cc5f0 4d63178 a6cc5f0 4d63178 a6cc5f0 11ac3b5 a6cc5f0 3bd48a2 a6cc5f0 3bd48a2 86d8d8b 3bd48a2 57616f7 3bd48a2 86d8d8b a6cc5f0 86d8d8b 56da3b2 380643c 56da3b2 a6cc5f0 277d2b5 a6cc5f0 9be5666 90a900b efa506a 9be5666 a6cc5f0 9be5666 c5a4e9d a6cc5f0 09f1883 a6cc5f0 86d8d8b a6cc5f0 56da3b2 459951f a6cc5f0 e018e6e a6cc5f0 3bd48a2 a6cc5f0 3bd48a2 1ea2d8c 3bd48a2 1ea2d8c 380643c 90a900b 380643c 3bd48a2 9b6581c 3bd48a2 9b6581c 3bd48a2 86d8d8b 3bd48a2 380643c 3bd48a2 c12a4ea 380643c e5afd16 380643c e5afd16 380643c e5afd16 c12a4ea e5afd16 380643c 3bd48a2 a6cc5f0 c5a4e9d a6cc5f0 3bd48a2 a6cc5f0 c20b770 86d8d8b 3bd48a2 a6cc5f0 3bd48a2 2ced563 a6cc5f0 e018e6e 459951f a6cc5f0 86d8d8b a6cc5f0 2ced563 99f1eda a6cc5f0 99f1eda 3bd48a2 99f1eda 3bd48a2 56da3b2 a6cc5f0 3bd48a2 99f1eda 3bd48a2 71ca7ab ca201a9 71ca7ab 81891de 37bc768 c9ff8ab 2ced563 c9ff8ab 2ced563 c9ff8ab 3bd48a2 d4ddf36 3bd48a2 d4ddf36 3bd48a2 d4ddf36 3bd48a2 c9ff8ab 3bd48a2 a6cc5f0 2ced563 a6cc5f0 81891de a6cc5f0 3bd48a2 75fea2a b343634 277d2b5 c5a4e9d 277d2b5 566c4e9 09f1883 566c4e9 cb8d9a7 09f1883 a7c8cf8 277d2b5 75fea2a 277d2b5 9be5666 a6cc5f0 b6e0e5d 1e9c586 a6cc5f0 b6e0e5d c5a4e9d 3bd48a2 b6e0e5d 9be5666 86d8d8b aaf1546 c12a4ea 86d8d8b 3bd48a2 a6cc5f0 1e9c586 a6cc5f0 86d8d8b a6cc5f0 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 832 833 834 835 836 837 838 839 840 841 842 843 844 845 846 847 848 849 850 851 852 853 854 855 856 857 858 859 860 861 862 863 864 865 866 867 868 869 870 871 872 873 874 875 876 877 878 879 880 881 882 883 884 885 886 887 888 889 890 891 892 893 894 895 896 897 898 899 900 901 902 903 904 905 906 907 908 909 910 911 912 913 914 915 916 917 918 919 920 921 922 923 924 925 926 927 928 929 930 931 932 933 934 935 936 937 938 939 940 941 942 943 944 945 946 947 948 949 950 951 952 953 954 955 956 957 958 959 960 961 962 963 964 965 966 967 968 969 970 971 972 973 974 975 976 977 978 979 980 981 982 983 984 985 986 987 988 989 990 991 992 993 994 995 996 997 998 999 1000 1001 1002 1003 1004 1005 1006 1007 1008 1009 1010 1011 1012 1013 1014 1015 1016 1017 1018 1019 1020 1021 1022 1023 1024 1025 1026 1027 1028 1029 1030 1031 1032 1033 1034 1035 1036 1037 1038 1039 1040 1041 1042 1043 1044 1045 1046 1047 1048 1049 1050 1051 1052 1053 1054 1055 1056 1057 1058 1059 1060 1061 1062 1063 1064 1065 1066 1067 1068 1069 1070 1071 1072 1073 1074 1075 1076 1077 1078 1079 1080 1081 1082 1083 1084 1085 1086 1087 1088 1089 1090 1091 1092 1093 1094 1095 1096 1097 1098 1099 1100 1101 1102 1103 1104 1105 1106 1107 1108 1109 1110 1111 1112 1113 1114 1115 1116 1117 1118 1119 1120 1121 1122 1123 1124 1125 1126 1127 1128 1129 1130 1131 1132 1133 1134 1135 1136 1137 1138 1139 1140 1141 1142 1143 1144 1145 1146 1147 1148 1149 1150 1151 1152 1153 1154 1155 1156 1157 1158 1159 1160 1161 1162 1163 1164 1165 1166 1167 1168 1169 1170 1171 1172 1173 1174 1175 1176 1177 1178 1179 1180 1181 1182 1183 1184 1185 1186 1187 1188 1189 1190 1191 1192 1193 1194 1195 1196 1197 1198 1199 1200 1201 1202 1203 1204 1205 1206 1207 1208 1209 1210 1211 1212 1213 1214 1215 1216 1217 1218 1219 1220 1221 1222 1223 1224 1225 1226 1227 1228 1229 1230 1231 1232 1233 1234 1235 1236 1237 1238 1239 1240 1241 1242 1243 1244 1245 1246 1247 1248 1249 1250 1251 1252 1253 1254 1255 1256 1257 1258 1259 1260 1261 1262 1263 1264 1265 1266 1267 1268 1269 1270 1271 1272 1273 1274 1275 1276 1277 1278 1279 1280 1281 1282 1283 1284 1285 1286 1287 1288 1289 1290 1291 1292 1293 1294 1295 1296 1297 1298 1299 1300 1301 1302 1303 1304 1305 1306 1307 1308 1309 1310 1311 1312 1313 1314 1315 1316 1317 1318 1319 1320 1321 1322 1323 1324 1325 1326 1327 1328 1329 1330 | """
Orbis 2 World Model — interactive rollout demo for Hugging Face Spaces.
This is the Orbis 2 counterpart to app.py (which drives the single-level Orbis 1
model). Orbis 2 is a hierarchical L1-L2 world model whose rollout script takes an
mp4 directly, so the flow differs from app.py:
1. User uploads a short driving clip.
2. We re-encode it to a constant CONTEXT_FPS (= the L1 frame rate) so the
rollout logic's exact integer-multiple frame-rate check passes.
3. We call into orbis2_app_engine.RolloutEngine (in-process, not a subprocess —
see that module's docstring), which samples the L1 (high-rate) and L2
(low-rate, further back) context windows from the tail of that video.
4. The context is tiled NUM_VIDEOS times so one minibatch rolls out that many
independently sampled futures under a single seed.
5. Generated frames (fake_images/sequence_XXXX/*.jpg) are encoded back into mp4s.
Checkpoint / config resolution (see resolve_exp_dir / download_checkpoints):
* If ORBIS2_EXP_DIR is set (and contains the config + checkpoint), it is used
as-is — this is the path that works today against a local training run.
* Otherwise, if ORBIS2_HF_REPO is set, the config + checkpoint are pulled from
that Hub repo into a local exp dir. Orbis 2 is not on the Hub yet; set these
env vars once the checkpoint is uploaded.
"""
import argparse
import hashlib
import json
import os
import random
import re
import subprocess
import tempfile
import traceback
import uuid
from pathlib import Path
import cv2
import gradio as gr
import numpy as np
import spaces
from huggingface_hub import snapshot_download
from orbis2_app_engine import get_engine
import torch
print(f"torch version: {torch.__version__}, cuda available: {torch.cuda.is_available()}, cuda version: {torch.version.cuda}")
# ----------------------------------------------------------------------------
# Workaround for gradio 5.9.1 / gradio_client bug (gradio-app/gradio#11722):
# get_api_info() walks the app's JSON schema and crashes with
# "TypeError: argument of type 'bool' is not iterable"
# when a schema node is a boolean (e.g. "additionalProperties": true), because
# get_type()/_json_schema_to_python_type() assume every schema is a dict. The
# main page route calls api_info() on every load, so without this the endpoint
# 500s continuously. Make the walker tolerate boolean schemas.
# ----------------------------------------------------------------------------
import gradio_client.utils as _gc_utils
_orig_get_type = _gc_utils.get_type
_orig_json_to_py = _gc_utils._json_schema_to_python_type
def _safe_get_type(schema):
if isinstance(schema, bool):
return "bool"
return _orig_get_type(schema)
def _safe_json_to_py(schema, defs=None):
if isinstance(schema, bool):
return "Any"
return _orig_json_to_py(schema, defs)
_gc_utils.get_type = _safe_get_type
_gc_utils._json_schema_to_python_type = _safe_json_to_py
# ----------------------------------------------------------------------------
# CLI args
# ----------------------------------------------------------------------------
def _parse_args():
parser = argparse.ArgumentParser(description="Orbis 2 hierarchical world model demo")
parser.add_argument(
"--ui-only", action="store_true",
help="Launch only the Gradio UI, without downloading or loading any "
"model/checkpoint. For UI/UX testing on a machine without a GPU: "
"'Generate rollouts' will show a placeholder error instead of "
"actually running inference.",
)
return parser.parse_args()
ARGS = _parse_args()
UI_ONLY = ARGS.ui_only
# ----------------------------------------------------------------------------
# Configuration — adjust these to your setup
# ----------------------------------------------------------------------------
# Persistent-storage location. Prefer the Space's persistent /data volume so large
# downloads/caches survive restarts; fall back to the equivalent path under the home
# dir if /data isn't writable (e.g. persistent storage not enabled).
def _pick_cache_dir(subpath: str = ".cache/huggingface") -> str:
for base in ("/data", os.path.expanduser("~")):
candidate = os.path.join(base, subpath)
try:
probe = Path(candidate)
probe.mkdir(parents=True, exist_ok=True)
test = probe / ".write_probe"
test.touch()
test.unlink()
return candidate
except OSError:
continue
return os.path.join(os.path.expanduser("~"), subpath)
HF_CACHE_DIR = os.path.abspath(os.environ.get("HF_HOME") or _pick_cache_dir())
os.environ["HF_HOME"] = HF_CACHE_DIR
HF_HUB_CACHE_DIR = str(Path(HF_CACHE_DIR) / "hub")
print(f"[startup] HF cache dir: {HF_CACHE_DIR}"
f"{'' if HF_CACHE_DIR.startswith('/data') else ' (ephemeral — /data not writable)'}")
# torch.compile artifacts (Inductor/Triton-compiled kernels) are tied to the exact
# torch/CUDA/cuDNN build and GPU model they were compiled on -- loading a cache built
# on a different combination can silently miscompile or just fail to load. So instead
# of one fixed filename, compile_cache_path() below fingerprints the current
# environment into the filename, and this dir (persistent across restarts, same as
# HF_CACHE_DIR above) is where those per-environment files accumulate.
COMPILE_CACHE_DIR = Path(os.path.abspath(_pick_cache_dir(".cache/orbis2_compile_cache")))
print(f"[startup] Compile cache dir: {COMPILE_CACHE_DIR}"
f"{'' if str(COMPILE_CACHE_DIR).startswith('/data') else ' (ephemeral — /data not writable)'}")
def compile_cache_path() -> Path:
"""Path to this exact environment's torch.compile cache file. Must be called with
CUDA already attached to the process (i.e. from inside a @spaces.GPU-decorated
call, not at module import time) -- HF Spaces ZeroGPU only attaches a GPU to that
call, not to the main process, so torch.cuda.* would report nothing useful before
then."""
parts = [f"torch{torch.__version__}", f"cu{torch.version.cuda or 'none'}"]
if torch.cuda.is_available():
parts.append(f"cudnn{torch.backends.cudnn.version()}")
parts.append(torch.cuda.get_device_name(0))
parts.append("sm%d%d" % torch.cuda.get_device_capability(0))
fingerprint = re.sub(r"[^A-Za-z0-9]+", "-", "_".join(parts)).strip("-")
return COMPILE_CACHE_DIR / f"compile_cache_{fingerprint}.pkl"
# --- Orbis 2 models (configs + checkpoints) ----------------------------------
# Three models are needed, laid out as <models>/{L1,L2,tok}/{config,checkpoints/last.ckpt}:
# L1 — detail predictor (the world model rollout_demo_v2.py loads directly)
# L2 — abstract predictor (frozen; referenced by L1's condition_preprocessor)
# tok — tokenizer (referenced by BOTH L1's and L2's tokenizer_config)
#
# L1's and L2's configs locate their dependencies via $ORBIS2_MODELS_DIR, which we
# export below so OmegaConf's expandvars resolves them wherever the tree lands.
CONFIG_NAME = "config_distill.yaml" # L1's config, relative to EXP_DIR
CKPT_NAME = "checkpoints/last.ckpt" # every model stores its weights here
#: (sub-dir, config filename) for each of the required models.
MODEL_SPECS = [
("L1", CONFIG_NAME),
("L2", "config.yaml"),
# L1's config_distill.yaml conditions on this frozen, distilled L2 predictor
# (a separate Hub folder from "L2") -- must be downloaded too, or L1 fails to
# build its condition_preprocessor with a FileNotFoundError on config.yaml.
("L2_distilled", "config.yaml"),
("tok", "config.yaml"),
]
# Hub repo holding the three models (root contains L1/, L2/, tok/).
HF_CKPT_REPO = os.environ.get("ORBIS2_HF_REPO", "sud0301/orbis2")
def _missing_models(root: Path) -> list[str]:
"""Return the models whose config or checkpoint is absent under `root`."""
return [sub for sub, cfg in MODEL_SPECS
if not (root / sub / cfg).exists() or not (root / sub / CKPT_NAME).exists()]
def resolve_models_dir() -> Path:
"""Return a directory containing L1/, L2/ and tok/, downloading it if needed.
A local tree wins if it's complete (a dev checkout, or ORBIS2_MODELS_DIR pointing
at one). Otherwise we snapshot the Hub repo and use the *cache* path directly
rather than copying into the Space — `local_dir=` would duplicate ~20 GB on disk.
The cache lives on /data when persistent storage is enabled, so the download
survives restarts.
"""
local = Path(os.environ.get("ORBIS2_MODELS_DIR") or (Path(__file__).parent / "models")).resolve()
if not _missing_models(local):
print(f"[startup] Using local Orbis 2 models: {local}")
return local
print(f"[startup] Fetching Orbis 2 models from {HF_CKPT_REPO} (~20 GB, first run only)…")
snapshot = Path(snapshot_download(
repo_id=HF_CKPT_REPO,
allow_patterns=[f"{sub}/**" for sub, _ in MODEL_SPECS],
cache_dir=HF_HUB_CACHE_DIR,
))
still_missing = _missing_models(snapshot)
if still_missing:
raise RuntimeError(
f"Downloaded {HF_CKPT_REPO} but these models are incomplete: "
f"{', '.join(still_missing)}. Each of L1/, L2/, tok/ must contain its "
f"config and '{CKPT_NAME}'."
)
print(f"[startup] Orbis 2 models ready: {snapshot}")
return snapshot
FRAME_H, FRAME_W = 288, 512
L1_FRAME_RATE = 10 # Hz to sample L1 context/rollout frames (and output fps)
CONTEXT_FPS = L1_FRAME_RATE # we re-encode uploads to this exact fps for the rollout script
if UI_ONLY:
print("[startup] --ui-only: skipping model/checkpoint download and load (UI/UX test mode).")
MODELS_DIR = None
EXP_DIR = None
ENGINE = None
# No model to introspect for the real lookback requirement -- 0 lets the player
# auto-seek stay a no-op instead of erroring.
MIN_CONTEXT_END_FRAME = 0
MIN_CONTEXT_TIME_S = 0.0
# No model to read model.num_pred_frames from -- 1 is just a placeholder so the
# UI-only slider below has *some* seconds-per-step to show; it's never used for
# an actual rollout (run_rollout raises before reaching the model in this mode).
FRAMES_PER_GEN_STEP = 1
else:
MODELS_DIR = resolve_models_dir()
os.environ["ORBIS2_MODELS_DIR"] = str(MODELS_DIR) # L1/L2 configs expandvars this
EXP_DIR = MODELS_DIR / "L1" # rollout_demo_v2.py --exp_dir
# Load the checkpoint once, at Space startup, on CPU (no GPU is attached outside a
# @spaces.GPU call). run_rollout() below only moves it to CUDA / compiles it, and
# only does that once too -- see orbis2_app_engine.py for why this avoids paying the
# multi-minute load+compile cost on every single generation request.
ENGINE = get_engine(EXP_DIR, CONFIG_NAME, CKPT_NAME)
# SPACES_ZERO_GPU is set by the platform only on ZeroGPU Spaces, where no GPU is
# attached to this (the main) process -- only to a @spaces.GPU-decorated call. On
# dedicated hardware CUDA is attached for the process's whole lifetime, so move the
# model there and compile it once now instead of paying that cost on the first
# request (run_rollout's own ENGINE.ensure_ready() call becomes a no-op once this
# has already run -- see RolloutEngine.ensure_ready's idempotency).
ON_ZERO_GPU = bool(os.environ.get("SPACES_ZERO_GPU"))
if not ON_ZERO_GPU:
print("[startup] not running on ZeroGPU (SPACES_ZERO_GPU unset) -- "
"moving model to CUDA and compiling now…")
ENGINE.ensure_ready(device="cuda", compile=True,
compile_artifacts=str(compile_cache_path()))
# Earliest point (in a CONTEXT_FPS-encoded upload) that has enough L1+L2 lookback to
# serve as the end of a valid context window -- the player auto-seeks here on load
# instead of sitting at time 0 (never valid: there's nothing before it). CPU-only,
# so safe to compute once here, right after the CPU-loaded ENGINE above.
MIN_CONTEXT_END_FRAME = ENGINE.min_context_end_frame(L1_FRAME_RATE)
MIN_CONTEXT_TIME_S = MIN_CONTEXT_END_FRAME / CONTEXT_FPS
FRAMES_PER_GEN_STEP = ENGINE.frames_per_rollout_step
# Seconds of output video one rollout step yields, at L1_FRAME_RATE fps -- lets the
# UI expose rollout length in seconds while the model itself still steps in whole
# rollout steps (see GEN_SECONDS_STEP / run_rollout's seconds -> steps conversion below).
SECONDS_PER_GEN_STEP = FRAMES_PER_GEN_STEP / L1_FRAME_RATE
DEFAULT_GEN_STEPS = 10 # rollout steps; each yields model.num_pred_frames frames
GEN_STEPS_MIN, GEN_STEPS_MAX, GEN_STEPS_STEP = 4, 30, 2 # bounds on the underlying step count
# The slider itself is in seconds (see build_demo()); these are its bounds, derived
# from the step-count bounds above via SECONDS_PER_GEN_STEP.
GEN_SECONDS_MIN = GEN_STEPS_MIN * SECONDS_PER_GEN_STEP
GEN_SECONDS_MAX = GEN_STEPS_MAX * SECONDS_PER_GEN_STEP
GEN_SECONDS_STEP = GEN_STEPS_STEP * SECONDS_PER_GEN_STEP
DEFAULT_GEN_SECONDS = DEFAULT_GEN_STEPS * SECONDS_PER_GEN_STEP
# L1 and L2 are sampled at different step counts. L1 (detail) needs more steps for
# image fidelity; L2 (abstract, consistency-distilled) is designed for very few.
DEFAULT_L1_STEPS = 4 # L1 sampler steps (NFE)
DEFAULT_L2_STEPS = 4 # L2 sampler steps (NFE) — matches the config's l2_pred_NFE
OUTPUT_FPS = L1_FRAME_RATE
NUM_VIDEOS = 3 # futures rolled out per run, batched into one minibatch
DEFAULT_ETA = 0.0 # 0 = deterministic ODE; >0 injects noise each solver step
# Very light cornflower-blue background + one darker cornflower-blue accent used
# consistently for the title/subtitle and interactive widgets (see build_demo()'s css=).
THEME_BG = "#eef3fb"
THEME_ACCENT = "#2f56a3"
TITLE_HTML = f"""
<div style="text-align:center; margin-bottom:8px;">
<div style="font-size:2.75rem; font-weight:800; color:{THEME_ACCENT}; line-height:1.1;">
Orbis 2
</div>
<div style="font-size:1.15rem; font-weight:500; color:{THEME_ACCENT}; margin-top:4px;">
A Hierarchical World Model for Driving
</div>
</div>
"""
AUTHOR_HTML = """
<div style="text-align:center; margin-bottom:4px;">
<a href="https://lmb.informatik.uni-freiburg.de/people/mittal/">Sudhanshu Mittal</a>*,
<a href="https://lmb.informatik.uni-freiburg.de/people/mousakha/">Arian Mousakhan</a>*,
<a href="https://lmb.informatik.uni-freiburg.de/people/galessos/">Silvio Galesso</a>*,
<a href="https://lmb.informatik.uni-freiburg.de/people/faridk/">Karim Farid</a>,
<a href="https://lmb.informatik.uni-freiburg.de/people/dienertj/">Johannes Dienert</a>,
<a href="https://lmb.informatik.uni-freiburg.de/people/sahayr/">Rajat Sahay</a>,
<a href="https://lmb.informatik.uni-freiburg.de/people/brox/index.html">Thomas Brox</a> — *main contributors<br>
University of Freiburg, Germany
</div>
<div style="text-align:center; margin-bottom:8px;">
<a href="https://lmb-freiburg.github.io/orbis2.github.io/">Project page</a> ·
<a href="https://github.com/lmb-freiburg/orbis">Code</a> ·
<a href="https://lmb-freiburg.github.io/orbis.github.io/">Orbis 1</a>
</div>
"""
DESCRIPTION = f"""
Upload a short driving clip (a few seconds is enough), or select one of the example
clips. The hierarchical model takes
the tail of the clip as L1 (high-rate) and L2 (low-rate, further back) context and
autoregressively predicts future frames.
Each run samples **{NUM_VIDEOS} rollouts** from the same context in a single
minibatch, so you see several plausible futures diverge from identical starting
conditions.
"""
PAPER_MD = """
### Abstract
> Current world models typically operate at a single abstraction level, favoring
> perceptual fidelity but lacking the spatial and semantic reasoning needed for
> downstream driving tasks. We propose a hierarchical driving world model that
> separates prediction into a high-level long-horizon scene forecaster and a
> low-level detail generator conditioned on it. This design improves both visual
> fidelity and spatial-semantic representation quality. We also introduce a
> two-stage training strategy: diffusion-forcing pretraining for richer
> representations, followed by teacher-forcing fine-tuning for stable
> autoregressive rollouts. Our method achieves state-of-the-art results on
> standard driving world model benchmarks, including long-horizon fidelity,
> counterfactual steering responsiveness, and internal representation quality.
### Method
Orbis 2 splits future prediction across two levels of abstraction:
- **Abstract predictor (L2)** — operates over a long temporal context to forecast a
future state in latent space, capturing abstract scene dynamics over long
horizons and enabling steering control.
- **Detail predictor (L1)** — conditioned on that abstract prediction, generates
fine-grained short-horizon frames, so high-fidelity local prediction stays
grounded in long-range temporal context.
Training runs in two stages: **diffusion-forcing pretraining** for richer
representations, followed by **teacher-forcing fine-tuning** for stable
autoregressive rollouts.
### Links & citation
```bibtex
@article{orbis2_2026,
author = {Mittal, Sudhanshu and Mousakhan, Arian and Galesso, Silvio and
Farid, Karim and Dienert, Johannes and Sahay, Rajat and Brox, Thomas},
title = {Orbis 2: A Hierarchical World Model for Driving},
journal = {arXiv preprint arXiv:2607.15898},
year = {2026},
}
```
"""
DEMO_MD = f"""
**Pipeline.** Your clip is re-encoded to a constant {CONTEXT_FPS} fps. The rollout
script samples two context windows from the tail of that video: an **L1** window at
{L1_FRAME_RATE} Hz (fine, recent frames) and an **L2** window sampled further back at
the frozen L2 predictor's own trained rate (coarse, long-horizon context). Each frame
is resized to {FRAME_H}×{FRAME_W}. The **L2 abstract predictor** forecasts a latent
future, and the **L1 detail predictor**, conditioned on it, generates the next frames
autoregressively.
**Rollout length.** *Generated video length (s)* sets how much video to generate; under
the hood this is rounded to a whole number of autoregressive steps, each of which emits
{FRAMES_PER_GEN_STEP} image frame{"s" if FRAMES_PER_GEN_STEP != 1 else ""} ({SECONDS_PER_GEN_STEP:.2g}s of video per step).
**Sampling.** The two levels are sampled independently, with their own step counts.
*L1 sampler steps* controls the detail predictor, which integrates the flow-matching
ODE and benefits from more steps (sharper frames, slower). *L2 sampler steps* controls
the abstract predictor, which is **consistency-distilled** and so is designed to run in
very few steps — raising it buys little.
**Randomness.** The {NUM_VIDEOS} rollouts share one context and one seed, but the
solver draws independent initial noise per batch element, so the futures diverge.
Leave *Seed* at -1 for a fresh draw each run, or set it explicitly to reproduce a run
exactly.
**Steering (optional).** Expand "Draw a steering trajectory" to freehand a path;
it's resampled and smoothed, then used to condition the rollout instead of running
unconditionally. Leave it untouched, or hit its Clear button, to roll out without
steering, as before.
**Note.** The clip must be long enough for the L2 look-back window; very short clips
are rejected with an explicit error. Frame and step counts are kept modest to fit the
ZeroGPU time budget.
"""
# ----------------------------------------------------------------------------
# Trajectory canvas (freehand steering input) — copied from app_traj.py, whose
# header documents these pieces as copy-paste-ready: the pure Python functions
# have no Gradio dependency, and CANVAS_HTML_JS / HEAD_SCRIPT are self-contained.
#
# Canvas coordinate convention: x = horizontal/lateral (X_MIN_M..X_MAX_M), y =
# vertical/forward (Y_MIN_M..Y_MAX_M), so points come back as [lateral, forward].
# rollout_demo_v2.py's --trajectory_file expects the opposite, [forward, lateral]
# (see trajectory_to_speed_yawrate in orbis2/evaluate/rollout_demo_v2.py) — the
# columns are swapped where the file is written, in run_rollout below.
# ----------------------------------------------------------------------------
# Both axes span the same number of meters and the canvas is square with the same
# px/meter on both axes, so a shape drawn on screen matches its real-world proportions
# instead of being stretched. X is centered on 0 (lateral: left/right of start); Y
# starts at 0 (forward: trajectory starts at the bottom, drawn upward).
TRAJ_EXTENT_M = 8 # meters spanned by each axis
TRAJ_X_MIN_M, TRAJ_X_MAX_M = -TRAJ_EXTENT_M / 2, TRAJ_EXTENT_M / 2
TRAJ_Y_MIN_M, TRAJ_Y_MAX_M = 0, TRAJ_EXTENT_M
TRAJ_PX_PER_M = 35
# Shrinks the on-screen canvas to 2/3 of its nominal (TRAJ_EXTENT_M * TRAJ_PX_PER_M)
# size. Safe to change independently of TRAJ_PX_PER_M: the JS below (metersToPixel /
# pixelToMeters) always normalizes by canvas.width/height, never a hardcoded pixel
# density, so it scales cleanly with whatever this canvas's actual size ends up being.
TRAJ_CANVAS_SCALE = 1.05 #2 / 3
TRAJ_CANVAS_SIZE_PX = int(TRAJ_EXTENT_M * TRAJ_PX_PER_M * TRAJ_CANVAS_SCALE)
# Canvas drawing units (comfortable -5..5m / 0..20m) are not calibrated to this
# model's dt=1/L1_FRAME_RATE=0.1s; scale positions up before the model sees them
# so a drawn stroke implies a realistic speed. Derived empirically: a near-full-
# height canvas stroke (~11.5m arc length over 200 samples) implied ~2.9 m/s at
# this model's dt, vs. ~24.2 m/s for a known-good reference trajectory (both
# figures computed with the same, since-corrected dt assumption — the 8.4x
# ratio between them is dt-invariant) — ratio ~8.4x.
TRAJ_REAL_WORLD_SCALE = 8.4
TRAJ_N_SAMPLES = 200 # resolution of the authoritative curve sent to the model
TRAJ_SMOOTHING_PASSES = 4 # more passes = smoother; 1 pass is a gentle nudge
TRAJ_SMOOTHING_WINDOW_FRAC = 0.15 # smoothing window as a fraction of TRAJ_N_SAMPLES
# TRAJ_N_SAMPLES points spread evenly over the calibrated dt=1/L1_FRAME_RATE pace
# span TRAJ_N_SAMPLES/L1_FRAME_RATE seconds of implied travel time, so point index
# i represents i/L1_FRAME_RATE seconds -- used to mark one point per second on the
# canvas overlay instead of all TRAJ_N_SAMPLES (too dense to read).
TRAJ_SAMPLES_PER_SECOND = L1_FRAME_RATE
def resample_by_arclength(points: np.ndarray, n_samples: int) -> np.ndarray:
"""Resample a raw (possibly unevenly-spaced) path to `n_samples` points
evenly spaced by arc length. Freehand strokes have point density that
depends on how fast the user moved the mouse, so this normalizes that
out before smoothing -- otherwise slow/fast sections would get smoothed
by different effective amounts.
"""
pts = np.asarray(points, dtype=float)
if len(pts) < 2:
return np.repeat(pts, max(n_samples, 1), axis=0)[:n_samples]
deltas = np.diff(pts, axis=0)
seg_lengths = np.sqrt((deltas ** 2).sum(axis=1))
cum_len = np.concatenate([[0.0], np.cumsum(seg_lengths)])
total_len = cum_len[-1]
if total_len == 0:
# User clicked without dragging: no real path, just repeat the point.
return np.repeat(pts[:1], n_samples, axis=0)
targets = np.linspace(0, total_len, n_samples)
x = np.interp(targets, cum_len, pts[:, 0])
y = np.interp(targets, cum_len, pts[:, 1])
return np.stack([x, y], axis=1)
def smooth_trajectory(
points: np.ndarray,
n_samples: int = TRAJ_N_SAMPLES,
passes: int = TRAJ_SMOOTHING_PASSES,
window_frac: float = TRAJ_SMOOTHING_WINDOW_FRAC,
) -> np.ndarray:
"""Turn a raw freehand stroke into a smooth model-facing trajectory:
resample to even spacing, then repeatedly box-filter both coordinate
channels (more passes ~ approximates a Gaussian blur, i.e. "heavy"
smoothing). Finally re-anchors the curve so it starts exactly at
(0, 0), since smoothing can nudge the very first point slightly, and
rotates it about that origin so the initial segment points straight
up (+y, zero lateral/x component) regardless of which direction the
stroke was actually drawn in.
"""
pts = resample_by_arclength(points, n_samples)
if len(pts) < 2:
return pts
window = max(int(n_samples * window_frac), 3)
if window % 2 == 0:
window += 1 # odd window keeps the filter centered
kernel = np.ones(window) / window
smoothed = pts.copy()
for _ in range(passes):
# Edge-pad so the curve doesn't shrink/shorten at the endpoints.
pad = window // 2
padded_x = np.pad(smoothed[:, 0], pad, mode="edge")
padded_y = np.pad(smoothed[:, 1], pad, mode="edge")
smoothed = np.stack(
[np.convolve(padded_x, kernel, mode="valid"),
np.convolve(padded_y, kernel, mode="valid")],
axis=1,
)
smoothed -= smoothed[0] # force the path to start exactly at (0, 0)
# Rotate about the origin so the initial segment points straight up (+y,
# zero lateral/x component) -- mirrors the heading-alignment rotation
# trajectory_to_model_frame() applies downstream, but at the canvas-space
# level so the on-screen overlay/plot already reflects it, not just the
# file eventually written for the model.
heading0 = np.arctan2(smoothed[1, 1] - smoothed[0, 1], smoothed[1, 0] - smoothed[0, 0])
delta = np.pi / 2 - heading0
cos_d, sin_d = np.cos(delta), np.sin(delta)
rotation = np.array([[cos_d, -sin_d], [sin_d, cos_d]])
smoothed = smoothed @ rotation.T
return smoothed
def parse_points_json(raw_json: str) -> np.ndarray:
"""Parse the JSON string written by JS (list of [x, y] in real meters,
x in roughly [TRAJ_X_MIN_M, TRAJ_X_MAX_M], y in roughly [TRAJ_Y_MIN_M,
TRAJ_Y_MAX_M]) into an (N, 2) numpy array. Returns an empty (0, 2) array
if the input is empty/invalid, so callers can handle the "nothing drawn
yet" state without try/except everywhere.
"""
try:
data = json.loads(raw_json) if raw_json else []
pts = np.array(data, dtype=float)
return pts.reshape(-1, 2) if pts.size else np.empty((0, 2))
except (json.JSONDecodeError, ValueError):
return np.empty((0, 2))
def points_to_trajectory(raw_json: str, n_samples: int = TRAJ_N_SAMPLES) -> np.ndarray:
"""End-to-end: raw freehand stroke JSON -> authoritative smoothed
(n_samples, 2) trajectory array. This is the function whose output
should actually be fed into the model.
"""
pts = parse_points_json(raw_json)
if len(pts) < 2:
return np.empty((0, 2))
return smooth_trajectory(pts, n_samples)
def trajectory_to_model_frame(trajectory: np.ndarray) -> np.ndarray:
"""Convert a canvas-drawn trajectory ([x=lateral(right+), y=forward], in canvas
meters) into the [forward, lateral] array rollout_demo_v2.py's --trajectory_file
expects, applying three corrections the canvas alone doesn't account for:
1. Column swap: canvas is [lateral, forward]; the model wants [forward, lateral].
2. Sign flip on lateral: the canvas's x-axis is positive-right (screen convention),
but rollout_demo_v2.py's own overlay rendering (_panel_coords_fit_trajectory:
`px = margin + (lateral_max - lateral) * scale`) shows increasing lateral moving
LEFT on screen, i.e. this codebase's convention is positive-lateral = left.
Without negating, a rightward-drawn stroke would condition/render as a left turn.
3. Heading alignment: trajectory_to_speed_yawrate/its reconstruction always treats
the trajectory's own first segment as pointing exactly along local +forward. A
freehand stroke's first smoothed segment is essentially never *exactly* aligned
with forward (mouse imprecision, smoothing edge effects) -- even a couple of
degrees of initial misalignment, left uncorrected, gets amplified into meters of
apparent lateral drift by the far end of a long trajectory once scaled up by
TRAJ_REAL_WORLD_SCALE. Rotating so segment 0 already points along +forward before
conversion matches what the reconstruction assumes, eliminating that drift.
Verified against a real canvas-drawn+smoothed trajectory: this reduces the
round-trip (speed/yaw_rate -> reconstructed position) error from several meters
down to ~0.16m over a ~97m trajectory.
"""
traj = np.asarray(trajectory)[:, [1, 0]] * TRAJ_REAL_WORLD_SCALE # [forward, lateral(right+)]
traj = traj.copy()
traj[:, 1] *= -1 # canvas right+ -> model left+
traj -= traj[0]
heading0 = np.arctan2(traj[1, 1] - traj[0, 1], traj[1, 0] - traj[0, 0])
cos_h, sin_h = np.cos(-heading0), np.sin(-heading0)
rotation = np.array([[cos_h, -sin_h], [sin_h, cos_h]])
return traj @ rotation.T
def render_trajectory_plot(raw_json: str):
"""Return the dense trajectory array into a gr.State for downstream use,
plus that same trajectory as a JSON string -- the JSON string is consumed
client-side to draw the smoothed curve back onto the canvas (see the
`smoothed_points_json` wiring in build_demo()).
"""
trajectory = points_to_trajectory(raw_json)
smoothed_json = json.dumps(trajectory.tolist()) if len(trajectory) else "[]"
return trajectory, smoothed_json
# --------------------------------------------------------------------------
# Canvas markup only -- NO <script> here. gr.HTML renders its content via
# innerHTML, and browsers do not execute <script> tags inserted that way.
# The actual JS lives in TRAJ_HEAD_SCRIPT below instead, which runs normally
# because it's inserted as real page <head> content.
# --------------------------------------------------------------------------
TRAJ_CANVAS_HTML_JS = f"""
<div style="position:relative; width:{TRAJ_CANVAS_SIZE_PX}px; margin:0 auto;">
<canvas id="traj-canvas" width="{TRAJ_CANVAS_SIZE_PX}" height="{TRAJ_CANVAS_SIZE_PX}"
style="display:block; border:1px solid #999; border-radius:8px; cursor:crosshair; background:#fafafa;">
</canvas>
<div style="position:absolute; top:8px; left:8px; display:flex; align-items:center; gap:8px;
background:rgba(255,255,255,0.85); border-radius:6px; padding:6px 8px;
font-size:11px; color:#555; line-height:1.4; pointer-events:none;">
<button id="traj-clear-btn"
style="font-size:13px; font-weight:600; padding:5px 14px; cursor:pointer;
background:#dc2626; color:#fff; border:none; border-radius:6px;
box-shadow:0 1px 3px rgba(0,0,0,0.35); pointer-events:auto; white-space:nowrap;">
Clear
</button>
<div style="text-align:left;">
<div><span style="color:#dc2626;">Red</span> = what you drew</div>
<div><span style="color:#4f46e5;">Blue</span> = model input (smoothed and aligned)</div>
</div>
</div>
</div>
"""
# --------------------------------------------------------------------------
# JS canvas logic: freehand paint a stroke, translate it to start at the
# fixed origin, and send the raw points to Python once the stroke finishes.
# All smoothing happens server-side (see smooth_trajectory above) -- this
# script only ever renders the raw, unsmoothed stroke for feedback.
#
# Placed in gr.Blocks(head=...) rather than inside the gr.HTML block above
# so the browser actually executes it (see comment above). Because gr.HTML
# content can render slightly after the head script runs,
# initTrajectoryCanvas() is invoked via a short poll that waits for the
# canvas element to exist in the DOM, then attaches listeners exactly once.
# --------------------------------------------------------------------------
TRAJ_HEAD_SCRIPT = f"""
<script>
function initTrajectoryCanvas() {{
const canvas = document.getElementById("traj-canvas");
const ctx = canvas.getContext("2d");
const clearBtn = document.getElementById("traj-clear-btn");
let isDrawing = false;
let rawStroke = []; // pixel-space points recorded during the current/last stroke
let smoothedOverlay = []; // pixel-space smoothed curve, received back from Python
const MIN_POINT_DIST = 3; // px; throttles how densely we sample while dragging
// Coordinate helpers: canvas pixel space has y=0 at the TOP, but
// trajectories are usually thought of with y=0 at the BOTTOM ("up" is
// positive) -- these flip that axis. They also map to/from real-world
// meters using the physical extent [X_MIN_M, X_MAX_M] x [Y_MIN_M, Y_MAX_M].
const X_MIN_M = {TRAJ_X_MIN_M}, X_MAX_M = {TRAJ_X_MAX_M};
const Y_MIN_M = {TRAJ_Y_MIN_M}, Y_MAX_M = {TRAJ_Y_MAX_M};
const SAMPLES_PER_SECOND = {TRAJ_SAMPLES_PER_SECOND}; // mark one overlay point per second, not all of them
function metersToPixel(mx, my) {{
const px = (mx - X_MIN_M) / (X_MAX_M - X_MIN_M) * canvas.width;
const py = canvas.height - (my - Y_MIN_M) / (Y_MAX_M - Y_MIN_M) * canvas.height;
return {{ x: px, y: py }};
}}
function pixelToMeters(px, py) {{
const mx = X_MIN_M + (px / canvas.width) * (X_MAX_M - X_MIN_M);
const my = Y_MIN_M + ((canvas.height - py) / canvas.height) * (Y_MAX_M - Y_MIN_M);
return {{ x: mx, y: my }};
}}
const ORIGIN_PX = metersToPixel(0, 0); // fixed start point, in pixel space
// Called from the `smoothed_points_json.change(..., js=...)` listener
// below whenever Python sends back a newly-smoothed curve. Converts it
// to pixel space and stores it for draw() to render as an overlay.
window.trajUpdateOverlay = function (jsonStr) {{
try {{
const data = JSON.parse(jsonStr || "[]");
smoothedOverlay = data.map(([mx, my]) => metersToPixel(mx, my));
}} catch (e) {{
smoothedOverlay = [];
}}
draw();
}};
// Redraw the meter-scale grid (with axis labels), reference axis,
// the raw stroke, the smoothed overlay, and the origin marker.
function draw() {{
ctx.clearRect(0, 0, canvas.width, canvas.height);
// Grid lines every 1m in x, every 2m in y, with tick labels along
// the bottom/left edges so the scale is readable at a glance.
ctx.strokeStyle = "#eee";
ctx.fillStyle = "#888";
ctx.font = "10px sans-serif";
for (let mx = X_MIN_M; mx <= X_MAX_M; mx += 1) {{
const px = metersToPixel(mx, 0).x;
ctx.beginPath(); ctx.moveTo(px, 0); ctx.lineTo(px, canvas.height); ctx.stroke();
ctx.fillText(mx.toFixed(0), px + 2, canvas.height - 4);
}}
for (let my = Y_MIN_M; my <= Y_MAX_M; my += 2) {{
const py = metersToPixel(0, my).y;
ctx.beginPath(); ctx.moveTo(0, py); ctx.lineTo(canvas.width, py); ctx.stroke();
ctx.fillText(my.toFixed(0), 2, py - 2);
}}
ctx.fillText("x (m)", canvas.width - 32, canvas.height - 4);
ctx.fillText("y (m)", 2, 10);
// Highlight x=0: the reference axis for left vs. right turns, won't
// generally line up with the meter grid's tick spacing above.
const xZeroPx = metersToPixel(0, 0).x;
ctx.strokeStyle = "#94a3b8";
ctx.lineWidth = 1.5;
ctx.beginPath();
ctx.moveTo(xZeroPx, 0);
ctx.lineTo(xZeroPx, canvas.height);
ctx.stroke();
ctx.lineWidth = 1;
if (rawStroke.length >= 2) {{
ctx.strokeStyle = "#dc2626";
ctx.lineWidth = 2;
ctx.beginPath();
ctx.moveTo(rawStroke[0].x, rawStroke[0].y);
for (const pt of rawStroke.slice(1)) ctx.lineTo(pt.x, pt.y);
ctx.stroke();
}}
if (smoothedOverlay.length >= 2) {{
ctx.strokeStyle = "#4f46e5";
ctx.lineWidth = 2.5;
ctx.setLineDash([7, 5]);
ctx.beginPath();
ctx.moveTo(smoothedOverlay[0].x, smoothedOverlay[0].y);
for (const pt of smoothedOverlay.slice(1)) ctx.lineTo(pt.x, pt.y);
ctx.stroke();
ctx.setLineDash([]); // reset so it doesn't leak into other strokes
// Mark one point per second of implied travel time (not every one of the
// TRAJ_N_SAMPLES resampled points -- that's too dense to read), so the
// pacing along the curve is visible.
ctx.fillStyle = "#4f46e5";
for (let i = 0; i < smoothedOverlay.length; i += SAMPLES_PER_SECOND) {{
const pt = smoothedOverlay[i];
ctx.beginPath();
ctx.arc(pt.x, pt.y, 3, 0, 2 * Math.PI);
ctx.fill();
}}
}}
// Fixed origin marker, drawn last so it's always on top of everything.
ctx.fillStyle = "#16a34a";
ctx.fillRect(ORIGIN_PX.x - 6, ORIGIN_PX.y - 6, 12, 12);
}}
// Send the current raw stroke (in meters) to Python and trigger its
// .change() listener, which resamples + heavily smooths it there.
function syncToGradio() {{
const inMeters = rawStroke.map(p => {{
const m = pixelToMeters(p.x, p.y);
return [m.x, m.y];
}});
const hidden = document.querySelector("#raw_points_json textarea");
if (hidden) {{
hidden.value = JSON.stringify(inMeters);
hidden.dispatchEvent(new Event("input", {{ bubbles: true }}));
}}
}}
function getMousePos(evt) {{
const rect = canvas.getBoundingClientRect();
return {{ x: evt.clientX - rect.left, y: evt.clientY - rect.top }};
}}
function distance(a, b) {{
return Math.hypot(a.x - b.x, a.y - b.y);
}}
// Shift every point in the stroke so it starts exactly at the fixed
// origin, preserving the drawn shape relative to that anchor.
function anchorStrokeToOrigin() {{
if (rawStroke.length === 0) return;
const dx = ORIGIN_PX.x - rawStroke[0].x;
const dy = ORIGIN_PX.y - rawStroke[0].y;
rawStroke = rawStroke.map(p => ({{ x: p.x + dx, y: p.y + dy }}));
}}
canvas.addEventListener("mousedown", (evt) => {{
isDrawing = true;
rawStroke = [getMousePos(evt)]; // start a fresh stroke, discard the old one
draw();
}});
canvas.addEventListener("mousemove", (evt) => {{
if (!isDrawing) return;
const pos = getMousePos(evt);
const last = rawStroke[rawStroke.length - 1];
if (!last || distance(last, pos) >= MIN_POINT_DIST) {{
rawStroke.push(pos);
draw(); // live feedback while painting, no server round-trip
}}
}});
function finishStroke() {{
if (!isDrawing) return;
isDrawing = false;
if (rawStroke.length >= 2) {{
anchorStrokeToOrigin();
draw();
syncToGradio(); // only sync once the stroke is complete
}}
}}
canvas.addEventListener("mouseup", finishStroke);
canvas.addEventListener("mouseleave", finishStroke); // in case the drag exits the canvas
clearBtn.addEventListener("click", () => {{
rawStroke = [];
draw();
syncToGradio();
}});
draw();
// Deliberately not calling syncToGradio() here (unlike the standalone
// app_traj.py demo this was copied from): "Generate rollouts" must stay
// unconditional until the user actually draws or clears the canvas.
canvas.dataset.trajInitialized = "true"; // guard against double-init
}}
// Poll for the canvas element since gr.HTML content can render slightly
// after this head script runs. Stops as soon as it's found and initialized.
const trajPollInterval = setInterval(() => {{
const canvas = document.getElementById("traj-canvas");
if (canvas && !canvas.dataset.trajInitialized) {{
clearInterval(trajPollInterval);
initTrajectoryCanvas();
}}
}}, 200);
</script>
"""
# ----------------------------------------------------------------------------
# Input-video cursor tracking. Rather than physically cropping the upload,
# "Generate rollouts" uses wherever the playhead is left as the end of the
# context window (see run_rollout / RolloutEngine.roll_out's context_end_frame).
#
# Gradio tears down and recreates the <video> element inside #input_video on
# every upload, so listeners are attached once at the document level in the
# capture phase (video events like `seeked`/`pause` don't bubble) instead of
# polling for the element like the trajectory canvas above does -- capture-phase
# delegation keeps working across re-renders with no re-attachment needed.
# ----------------------------------------------------------------------------
CURSOR_HEAD_SCRIPT = f"""
<script>
const MIN_CONTEXT_TIME_S = {MIN_CONTEXT_TIME_S};
function isInputVideo(target) {{
return target instanceof HTMLVideoElement
&& target.closest("#input_video") !== null;
}}
function writeCursorTime(seconds) {{
const hidden = document.querySelector("#cursor_time_s textarea");
if (hidden) {{
hidden.value = String(seconds);
hidden.dispatchEvent(new Event("input", {{ bubbles: true }}));
}}
}}
document.addEventListener("loadedmetadata", (evt) => {{
if (!isInputVideo(evt.target)) return;
// Auto-seek away from time 0 (never a valid context end) to the earliest
// frame with enough lookback. If this seek has no effect for some reason
// (e.g. a non-seekable stream), no "seeked" event follows and the hidden
// field is simply never written -- run_rollout then falls back to using
// the tail of the whole video, same as before this feature existed.
evt.target.currentTime = Math.min(MIN_CONTEXT_TIME_S, evt.target.duration || MIN_CONTEXT_TIME_S);
}}, true);
document.addEventListener("seeked", (evt) => {{
if (!isInputVideo(evt.target)) return;
writeCursorTime(evt.target.currentTime);
}}, true);
document.addEventListener("timeupdate", (evt) => {{
if (!isInputVideo(evt.target)) return;
writeCursorTime(evt.target.currentTime);
}}, true);
</script>
"""
# ----------------------------------------------------------------------------
# Video helpers (CPU)
# ----------------------------------------------------------------------------
def reencode_to_fps(src_path: str, out_path: Path, fps: int = CONTEXT_FPS):
"""Re-encode a clip to a constant `fps` H.264 mp4.
The rollout script requires the video's native frame rate to be an exact integer
multiple of the L1/L2 sampling rates. Forcing a constant fps here (via -vsync cfr)
makes an arbitrary upload satisfy that check.
"""
out_path.parent.mkdir(parents=True, exist_ok=True)
subprocess.run(
["ffmpeg", "-y", "-i", str(src_path),
"-vsync", "cfr", "-r", str(fps),
"-an", "-c:v", "libx264", "-pix_fmt", "yuv420p",
"-movflags", "+faststart", str(out_path)],
check=True, capture_output=True,
)
EXAMPLES_DIR = Path(__file__).parent / "examples"
EXAMPLE_THUMB_DIR = EXAMPLES_DIR / ".thumbs"
def example_thumbnail(video_path: Path) -> str:
"""First-frame jpg thumbnail for an example clip, generated once and cached
alongside it (regenerated if the source clip is newer than the cached thumb)."""
thumb_path = EXAMPLE_THUMB_DIR / f"{video_path.stem}.jpg"
if not thumb_path.exists() or thumb_path.stat().st_mtime < video_path.stat().st_mtime:
thumb_path.parent.mkdir(parents=True, exist_ok=True)
cap = cv2.VideoCapture(str(video_path))
ok, frame = cap.read()
cap.release()
if not ok:
raise RuntimeError(f"Could not read a frame from example clip {video_path}")
cv2.imwrite(str(thumb_path), frame)
return str(thumb_path)
def discover_examples() -> list[Path]:
"""Example clips to show as clickable thumbnails next to the upload box.
Returns whatever's in examples/ (possibly none) rather than a fixed list, so
clips can be dropped in/removed without a code change."""
return sorted(EXAMPLES_DIR.glob("*.mp4"))
def frames_to_mp4(frame_paths: list[Path], out_path: Path, fps: int = OUTPUT_FPS):
first = cv2.imread(str(frame_paths[0]))
h, w = first.shape[:2]
# mp4v then re-encode with ffmpeg for browser-compatible H.264
tmp = out_path.with_suffix(".raw.mp4")
writer = cv2.VideoWriter(str(tmp), cv2.VideoWriter_fourcc(*"mp4v"), fps, (w, h))
for p in frame_paths:
writer.write(cv2.imread(str(p)))
writer.release()
subprocess.run(
["ffmpeg", "-y", "-i", str(tmp), "-c:v", "libx264",
"-pix_fmt", "yuv420p", "-movflags", "+faststart", str(out_path)],
check=True, capture_output=True,
)
tmp.unlink(missing_ok=True)
def _parse_cursor_time_s(raw) -> float | None:
"""Parse the hidden cursor_time_s textbox. Returns None for the "-1" sentinel
(cursor never moved, or the browser couldn't seek/report it) and for anything
unparsable, so callers can fall back to the tail-of-video default."""
try:
t = float(raw)
except (TypeError, ValueError):
return None
return t if t >= 0 else None
# ----------------------------------------------------------------------------
# GPU inference
# ----------------------------------------------------------------------------
@spaces.GPU(duration=60)
def run_rollout(video_path: str, gen_seconds: float, l1_steps: int, l2_steps: int,
seed: int, trajectory, cursor_time_s: str, progress=gr.Progress()):
if video_path is None:
raise gr.Error("Please upload a short mp4 clip first.")
if UI_ONLY:
print(f"[ui-only] cursor_time_s = {cursor_time_s!r}")
raise gr.Error(
"Running with --ui-only: no model is loaded, so rollouts can't be "
"generated. Restart without --ui-only (on a GPU machine) to run inference."
)
def say(frac, msg):
# progress=None hides the bar (and, it turns out, the desc text along with it --
# Gradio falls back to its generic loading spinner instead), so always pass a
# real fraction: 0.0 through "Loading", the real per-step fraction through
# "Rolling out", 1.0 through "Saving"/"Done".
progress(frac, desc=msg)
try:
# seed < 0 means "surprise me" — draw a fresh positive seed. The rollout script
# only seeds when seed > 0, hence the clamp. Within a batch each sample still
# draws its own noise, so one seed yields NUM_VIDEOS distinct rollouts; setting
# it explicitly lets a run be reproduced exactly.
seed = random.randrange(1, 2**31 - 1) if int(seed) < 0 else max(1, int(seed))
# eta is not exposed: 0.0 keeps sampling on the deterministic ODE.
eta = DEFAULT_ETA
# The model itself only understands whole rollout steps -- convert the
# seconds the user picked back into a step count (see SECONDS_PER_GEN_STEP).
num_gen_steps = max(1, round(gen_seconds / SECONDS_PER_GEN_STEP))
job = Path(tempfile.gettempdir()) / f"orbis2_{uuid.uuid4().hex[:8]}"
say(0.0, "⏳ Loading…")
print(f"[load] re-encoding uploaded clip to {CONTEXT_FPS} fps…")
ctx_video = job / "context.mp4"
reencode_to_fps(video_path, ctx_video, CONTEXT_FPS)
print(f"[load] context video ready at {ctx_video}")
# The playhead position (wherever the user left it) marks where the context
# window ends -- convert it from seconds to a frame index in the re-encoded
# (CONTEXT_FPS) video. The "-1" sentinel (cursor never moved / not reported)
# falls back to None, which makes RolloutEngine.roll_out use the tail of the
# whole video, same as before this feature existed.
print("[load] resolving context window from cursor position…")
parsed_cursor_s = _parse_cursor_time_s(cursor_time_s)
context_end_frame = None
if parsed_cursor_s is not None:
cap = cv2.VideoCapture(str(ctx_video))
ctx_frame_count = int(cap.get(cv2.CAP_PROP_FRAME_COUNT))
cap.release()
context_end_frame = max(0, min(round(parsed_cursor_s * CONTEXT_FPS), ctx_frame_count - 1))
print(f"context @ {CONTEXT_FPS} fps | cursor {parsed_cursor_s} s -> "
f"end frame {context_end_frame} | seed {seed} | "
f"{int(num_gen_steps)} rollout steps | L1 {int(l1_steps)} / L2 {int(l2_steps)} "
f"sampler steps | eta {eta}")
output_dir = job / "rollout"
traj_path = None
if trajectory is not None and len(trajectory) >= 2:
print("[load] preparing trajectory file…")
traj_path = job / "trajectory.npy"
np.save(traj_path, trajectory_to_model_frame(trajectory))
# First request in this worker moves the model to CUDA and compiles it, loading
# this environment's cached artifacts if they're already there from a previous
# run, or compiling from scratch and saving them if not; every later request
# just reuses that already-ready model (see orbis2_app_engine.py).
cache_path = compile_cache_path()
print(f"[load] ensure_ready: moving model to GPU / compiling if needed… (compile cache: {cache_path})")
ENGINE.ensure_ready(device="cuda", compile=True,
compile_artifacts=str(cache_path))
# Deliberately no say(..., "Rolling out…") here: between this point and the
# first real tick below, roll_out is still doing L2's upfront forecast (and,
# on a fresh environment, a possibly multi-minute from-scratch torch.compile)
# with no progress signal of its own -- showing "Rolling out" for that silent
# stretch would be misleading, so the bar just stays on "Loading…" until
# progress_cb's first call proves a rollout step has actually completed.
#
# ENGINE.roll_out reports one tick per autoregressive rollout step (out of
# num_gen_steps total) via progress_cb -- the bar climbs 0 -> 1 across those
# ticks, so it's already sitting at 100% by the time "Saving" replaces the text
# (decode + trajectory overlay + frame writing, which happen after the last
# tick with no progress signal of their own, don't need their own fraction).
def _rollout_progress(step, total):
frac = step / total if total else 1.0
say(frac, f"🚗 Rolling out… ({step}/{total})")
ENGINE.roll_out(
video_path=str(ctx_video),
output_dir=str(output_dir),
num_gen_frames=int(num_gen_steps),
# L1 and L2 sample at different step counts; num_steps is L1's NFE,
# l2_nfe overrides the config's l2_pred_NFE for the L2 predictor.
num_steps=int(l1_steps),
l2_nfe=int(l2_steps),
l1_frame_rate=L1_FRAME_RATE,
# Given explicitly since roll_out() no longer resolves size from a
# training config -- it always uses the height/width passed in here.
height=FRAME_H,
width=FRAME_W,
eta=float(eta),
seed=seed,
num_videos=NUM_VIDEOS,
trajectory_file=str(traj_path) if traj_path is not None else None,
vis_mode="trajectory_ego",
device="cuda",
context_end_frame=context_end_frame,
progress_cb=_rollout_progress,
)
say(1.0, "💾 Saving…")
out_mp4s = []
for i in range(NUM_VIDEOS):
seq_dir = output_dir / "fake_images" / f"sequence_{i:04d}"
gen_frames = sorted(seq_dir.glob("*.jpg"))
if not gen_frames:
raise gr.Error(f"No generated frames found in {seq_dir} — "
"check rollout output path.")
out_mp4 = job / f"rollout_{i}.mp4"
frames_to_mp4(gen_frames, out_mp4)
out_mp4s.append(str(out_mp4))
say(1.0, f"✅ Done — {NUM_VIDEOS} rollouts of {int(num_gen_steps)} steps.")
return tuple(out_mp4s)
except gr.Error:
raise
except Exception as e:
tb = traceback.format_exc()
print(tb)
raise gr.Error(f"Rollout failed: {e}\n\n{tb[-1500:]}")
# ----------------------------------------------------------------------------
# UI
# ----------------------------------------------------------------------------
def build_demo():
with gr.Blocks(title="Orbis 2: A Hierarchical World Model for Driving",
head=TRAJ_HEAD_SCRIPT + CURSOR_HEAD_SCRIPT,
# The UI is designed light-only (forced light .gradio-container
# background, hardcoded dark hint text, light canvas). A visitor
# whose browser is in dark mode would otherwise get Gradio's light
# text on our forced-light background -> invisible text. Force the
# page to always load in light theme via ?__theme=light.
js="""
() => {
const url = new URL(window.location.href);
if (url.searchParams.get('__theme') !== 'light') {
url.searchParams.set('__theme', 'light');
window.location.replace(url.href);
}
}
""",
css=f"""
/* Fixes the upload box's height so it doesn't jump between the empty
dropzone and a loaded video of some other aspect ratio; object-fit
letterboxes videos that don't match instead of growing the box. */
#input_video {{ height: 320px !important; }}
#input_video video {{ height: 100% !important; object-fit: contain !important; }}
/* Vertical rail of example-clip thumbnails to the left of the input
video; height matches #input_video above so the two line up (the
"Examples" label below it, like the hint text below the video,
sits outside this fixed height). */
#example_rail {{ display: flex !important; flex-direction: column !important;
flex-wrap: nowrap !important;
justify-content: space-between; height: 320px; gap: 6px; }}
#example_rail > div {{ flex: 1 1 0 !important; min-height: 0; }}
#example_label {{ text-align: center; }}
.example-thumb {{ cursor: pointer; border-radius: 6px; overflow: hidden;
height: 100% !important; }}
.example-thumb img {{ height: 100% !important; width: 100% !important;
object-fit: cover !important; }}
.example-thumb:hover {{ outline: 2px solid {THEME_ACCENT}; }}
/* Blue theme: very light cornflower page background, one darker
cornflower blue reused for interactive widgets (native range/number/
checkbox accent color covers sliders without fragile DOM targeting)
and the primary action button.
We force a LIGHT page background but Gradio still picks its light/dark
*theme* from the visitor's browser. In dark mode it paints panels
(Groups, Video, Accordions, inputs) with dark fills + light text while
our page stays light -> mismatched panels and invisible text. Rather
than depend on the ?__theme=light redirect firing, pin Gradio's own
theme CSS variables to light values for BOTH the light and .dark
container scopes, so every panel is light with dark text regardless of
the browser's mode. Because these are the same variables Gradio's
theme sets, everything (blocks, labels, inputs, borders) stays
consistent instead of being patched element by element. */
.gradio-container, .gradio-container.dark {{
--body-background-fill: {THEME_BG};
--background-fill-primary: #ffffff;
--background-fill-secondary: {THEME_BG};
--block-background-fill: #ffffff;
--panel-background-fill: #ffffff;
--input-background-fill: #ffffff;
--block-label-background-fill: #ffffff;
--block-title-background-fill: #ffffff;
--body-text-color: #111827;
--body-text-color-subdued: #4b5563;
--block-label-text-color: #111827;
--block-title-text-color: #111827;
--block-info-text-color: #4b5563;
--code-background-fill: #f3f4f6;
--border-color-primary: #d0d7e6;
--link-text-color: #2563eb;
--link-text-color-hover: #1d4ed8;
--link-text-color-active: #1d4ed8;
--link-text-color-visited: #2563eb;
background: {THEME_BG} !important;
color: #111827 !important;
}}
/* Belt-and-suspenders for text nodes that read a hardcoded color rather
than the variables above; links stay blue (rule listed last, same
specificity, so it wins). */
.gradio-container p, .gradio-container li, .gradio-container span,
.gradio-container .prose, .gradio-container .prose p,
.gradio-container .prose li, .gradio-container .prose span {{
color: #111827 !important;
}}
.gradio-container a, .gradio-container .prose a {{
color: #2563eb !important;
}}
/* Code / bibtex blocks: Gradio's dark-theme code fill leaks through in
dark mode. Pin a light fill + dark text directly on pre/code too, in
case the block reads a hardcoded color instead of --code-background-fill. */
.gradio-container pre, .gradio-container code,
.gradio-container .prose pre, .gradio-container .prose code {{
background: #f3f4f6 !important;
color: #111827 !important;
}}
/* Video player controls sit on a dark bar; our forced-dark body text
above turns the timer text and play/fullscreen icons black-on-black.
Force just the timer text + icon SVGs white. Deliberately scoped to
.time and svg only (color/fill, never background) so the seek/progress
bar -- which is already white -- is left untouched. */
.gradio-container .controls .time {{ color: #fff !important; }}
.gradio-container .controls svg {{
fill: #fff !important; color: #fff !important;
}}
input[type="range"], input[type="number"], input[type="checkbox"] {{
accent-color: {THEME_ACCENT};
}}
#generate_btn {{ background: {THEME_ACCENT} !important; border-color: {THEME_ACCENT} !important; }}
"""
) as demo:
gr.HTML(TITLE_HTML)
gr.HTML(AUTHOR_HTML)
gr.Markdown(DESCRIPTION)
with gr.Row():
with gr.Column(scale=2):
with gr.Group():
n_seconds_gen = gr.Slider(GEN_SECONDS_MIN, GEN_SECONDS_MAX, value=DEFAULT_GEN_SECONDS,
step=GEN_SECONDS_STEP,
label="Generated video length (s)",
info=f"Generated in chunks of {SECONDS_PER_GEN_STEP:.2g}s each")
l1_steps = gr.Slider(5, 20, value=DEFAULT_L1_STEPS, step=1,
label="L1 sampler steps",
info="Detail predictor — more steps, sharper frames")
with gr.Row():
l2_steps = gr.Slider(4, 8, value=DEFAULT_L2_STEPS, step=1, scale=2,
label="L2 sampler steps",
info="Abstract predictor — distilled, needs very few")
seed_in = gr.Number(value=-1, precision=0, label="Seed", scale=1,
info="-1 draws a fresh seed each run")
with gr.Column(scale=4):
with gr.Row():
example_paths = discover_examples()
example_thumbs = []
if example_paths:
with gr.Column(scale=1, min_width=90):
with gr.Column(elem_id="example_rail"):
for p in example_paths:
# Each thumbnail gets its own Row so Gradio can't auto-group
# these bare sibling Images into a 2-column form grid (the
# cause of the "two side by side, one below" layout bug).
with gr.Row():
example_thumbs.append((
gr.Image(value=example_thumbnail(p), interactive=False,
show_label=False,
elem_classes=["example-thumb"]),
str(p),
))
gr.Markdown("**Examples**", elem_id="example_label")
with gr.Column(scale=3):
inp = gr.Video(label="Upload a short mp4 clip", sources=["upload"],
elem_id="input_video")
cursor_time_s = gr.Textbox(elem_id="cursor_time_s", value="-1", visible=False)
gr.Markdown(
"Use the **video cursor** to pick where the model generation should "
"start (last context frame)."
)
btn = gr.Button("Generate rollouts", variant="primary", elem_id="generate_btn")
# Wired up here, after `inp` exists, so clicking a thumbnail loads that
# clip into the upload box exactly like gr.Examples would.
for thumb, path in example_thumbs:
thumb.select(lambda path=path: path, outputs=inp)
with gr.Column(scale=3):
gr.HTML(TRAJ_CANVAS_HTML_JS)
gr.HTML(
'<p style="margin:0 0 4px 0;"><strong>Draw a steering trajectory '
"(optional).</strong> Freehand-draw a path on the canvas (starts at "
"the green square, forward is up). Leave it untouched, or hit Clear, "
"to run the rollout unconditionally.</p>"
'<p style="font-size:14px; color:#555; margin:0;">'
"Trajectories are smoothed, centered, and heading-aligned. "
"Hand-drawn trajectories are out-of-distribution for the model, "
"so expect erratic outputs. Have fun crashing."
"</p>"
)
raw_points_json = gr.Textbox(elem_id="raw_points_json", visible=False)
smoothed_points_json = gr.Textbox(elem_id="smoothed_points_json", visible=False)
trajectory_state = gr.State()
raw_points_json.change(
fn=render_trajectory_plot,
inputs=raw_points_json,
outputs=[trajectory_state, smoothed_points_json],
)
smoothed_points_json.change(
fn=None,
inputs=[smoothed_points_json],
outputs=[],
js="(v) => { if (window.trajUpdateOverlay) { window.trajUpdateOverlay(v); } }",
)
with gr.Row():
outs = [gr.Video(label=f"Rollout {i + 1}", autoplay=True, loop=True)
for i in range(NUM_VIDEOS)]
with gr.Accordion("About this demo", open=False):
gr.Markdown(DEMO_MD)
with gr.Accordion("About the paper", open=False):
gr.Markdown(PAPER_MD)
btn.click(
run_rollout,
inputs=[inp, n_seconds_gen, l1_steps, l2_steps, seed_in, trajectory_state, cursor_time_s],
outputs=outs,
concurrency_limit=1,
)
return demo
if __name__ == "__main__":
build_demo().queue(max_size=8).launch()
|