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()