File size: 20,655 Bytes
3bd48a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c9ff8ab
3bd48a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b9a3a43
3bd48a2
 
 
 
 
 
 
c9ff8ab
3bd48a2
 
 
 
 
99f1eda
 
 
 
 
 
 
 
 
 
5d82a35
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
c9ff8ab
 
37bc768
 
 
 
 
 
c9ff8ab
 
 
 
 
 
 
 
 
 
 
 
 
 
 
37bc768
c9ff8ab
 
 
 
 
 
 
 
 
 
 
 
3bd48a2
 
 
 
 
 
 
 
 
 
 
 
71ca7ab
 
3bd48a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
b9a3a43
3bd48a2
 
 
 
 
 
 
 
 
 
86d8d8b
 
 
 
 
 
 
 
3bd48a2
 
71ca7ab
3bd48a2
71ca7ab
 
3bd48a2
71ca7ab
3bd48a2
 
 
 
 
 
 
 
71ca7ab
3bd48a2
 
 
99f1eda
3bd48a2
 
 
 
 
 
 
 
71ca7ab
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3bd48a2
 
 
 
71ca7ab
3bd48a2
 
 
 
 
b9a3a43
3bd48a2
c9ff8ab
3bd48a2
 
 
 
 
 
 
 
 
 
 
 
 
b9a3a43
3bd48a2
99f1eda
 
 
 
 
 
b9a3a43
3bd48a2
b9a3a43
3bd48a2
 
 
 
 
 
3bfcce9
 
3bd48a2
 
 
 
 
 
 
 
 
 
 
 
3bfcce9
 
3bd48a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3bfcce9
 
3bd48a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
3bfcce9
 
 
 
 
 
 
3bd48a2
c9ff8ab
5d82a35
 
3bd48a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
71ca7ab
 
 
 
 
 
 
 
3bd48a2
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
"""
Persistent, in-process rollout engine for app2.py.

app2.py used to shell out to `python evaluate/rollout_demo_v2.py` on every click of
"Generate rollouts", which re-read the config, re-instantiated the model, re-loaded
the ~20GB checkpoint from disk, and re-established torch.compile from scratch on
every single request. This module keeps one loaded (and, after the first request,
compiled) model alive for the life of the process instead: `load()` does the
one-time CPU setup, `ensure_ready()` moves to GPU and compiles exactly once, and
`roll_out()` is the cheap per-request call.

This relies on the Hugging Face ZeroGPU worker that runs `@spaces.GPU`-decorated
calls being reused across requests rather than restarted per call -- if the worker
is ever recycled (idle timeout, scale-to-zero), the next request after that pays
the load/compile cost again, same as today's cold start, just no longer on *every*
request.

`orbis2/evaluate/rollout_demo_v2.py` is vendored from another repo and can be
overwritten by a future sync, so nothing app-critical lives there: this module only
imports its small, generic, model-agnostic helpers (context sampling, trajectory
math, config resolution) and owns model construction / the rollout body itself.
"""

import contextlib
import logging
import os
import sys
import threading
from pathlib import Path

import torch
from omegaconf import OmegaConf
from pytorch_lightning import seed_everything
from torchvision.utils import save_image

_REPO_DIR = Path(__file__).parent / "orbis2"
if str(_REPO_DIR) not in sys.path:
    sys.path.insert(0, str(_REPO_DIR))

from evaluate.rollout_demo_v2 import (  # noqa: E402
    STEERING_FORMAT,
    _L1L2FrameIndexer,
    get_rollout_future_frame_count,
    load_l1_l2_context,
    load_trajectory_points,
    make_unconditional_steering,
    maybe_apply_condition_preprocessor_scales,
    maybe_apply_l2_nfe,
    overlay_trajectory_on_images,
    require_l1l2_model,
    resample_trajectory_by_arclength,
    resolve_l2_frame_rate,
    resolve_video_backend,
    trajectory_to_speed_yawrate,
)
from modules import fm_samplers  # noqa: E402
from util import instantiate_from_config  # noqa: E402

logger = logging.getLogger(__name__)


def _describe_sampler(sampler, nfe, eta):
    """Human-readable solver name + init params (e.g. FlowMatchingSamplerHeunPlusEuler
    with its timescale/integration_t_eps/step_schedule/...), plus this call's NFE/eta --
    the latter two are sample()-time args, not stored on the sampler itself."""
    if sampler is None:
        return "none"
    params = {k: v for k, v in vars(sampler).items() if k != "module"}
    return f"{type(sampler).__name__}(NFE={nfe}, eta={eta}, {params})"


@contextlib.contextmanager
def _l1_nfe_override(model, num_steps):
    """Let the "L1 sampler steps" UI control actually change L1's step count.

    FlowMatchingSamplerHeunPlusEuler._resolve_schedule ignores NFE whenever the
    sampler has an explicit `step_schedule` configured (see fm_samplers.py) --
    step_schedule always wins, so a fixed-schedule L1 config (like this app's
    config_distill.yaml) would otherwise make num_steps/l1_steps a no-op. Since
    that's vendored code we don't touch, this instead mutates the already-
    constructed sampler's `step_schedule` attribute directly -- same pattern as
    maybe_apply_l2_nfe's post-init override of l2_pred_NFE, which relies on the
    same fact: these are read at sample time, not at construction.

    Uses the config's tuned step_schedule as-is whenever the requested
    num_steps matches its own step count -- that's the "default config steps"
    case. Only when a genuinely different l1_nfe is requested does this clear
    step_schedule for the call, so the sampler falls back to its own NFE-driven
    schedule (NFE-1 uniform Heun steps + 1 uniform Euler step). Restores the
    original schedule afterward either way.
    """
    sampler = getattr(model, "sampler", None)
    is_fixed_schedule = (
        isinstance(sampler, fm_samplers.FlowMatchingSamplerHeunPlusEuler)
        and sampler.step_schedule is not None
        and len(sampler.step_schedule) != int(num_steps)
    )
    if not is_fixed_schedule:
        yield
        return

    original_schedule = sampler.step_schedule
    print(f"[solver] L1 step_schedule has {len(original_schedule)} steps; requested "
          f"l1_steps={num_steps} differs -- using {int(num_steps) - 1} Heun step(s) + "
          "1 Euler step for this request instead.")
    sampler.step_schedule = None
    try:
        yield
    finally:
        sampler.step_schedule = original_schedule


@contextlib.contextmanager
def _rollout_progress_tqdm(progress_cb):
    """Temporarily wrap fm_samplers.tqdm -- used only by FlowMatchingSampler.roll_out's
    outer per-rollout-step loop (one iteration per autoregressive chunk, i.e. per
    num_gen_frames) -- so it still prints its normal console progress bar while also
    reporting each completed step to progress_cb(step, total). Real per-step progress
    without editing the vendored fm_samplers.py; restores the original tqdm on exit
    either way."""
    if progress_cb is None:
        yield
        return

    real_tqdm = fm_samplers.tqdm

    def _tqdm_with_progress(iterable, *args, **kwargs):
        total = kwargs.get("total")
        if total is None:
            try:
                total = len(iterable)
            except TypeError:
                total = None

        def _gen():
            for i, item in enumerate(real_tqdm(iterable, *args, **kwargs)):
                yield item
                progress_cb(i + 1, total)

        return _gen()

    fm_samplers.tqdm = _tqdm_with_progress
    try:
        yield
    finally:
        fm_samplers.tqdm = real_tqdm


class RolloutEngine:
    """Owns one loaded L1-L2 hierarchical model and serves repeated rollout requests."""

    def __init__(self, exp_dir, config_name, ckpt_name, evaluate_ema=True):
        self.exp_dir = str(exp_dir)
        self.config_path = os.path.join(self.exp_dir, config_name)
        self.ckpt_path = os.path.join(self.exp_dir, ckpt_name)
        self.evaluate_ema = evaluate_ema

        self.model = None
        self._ready = False
        self._lock = threading.Lock()
        self._compile_artifacts = None
        self._artifacts_saved = False

    def load(self):
        """CPU-only one-time setup: read the config, build the model, load the
        checkpoint. Call before any GPU is attached (i.e. outside a @spaces.GPU
        function) -- mirrors the model-construction prefix of
        rollout_demo_v2.generate_images()."""
        config = OmegaConf.load(self.config_path)
        model = instantiate_from_config(config.model)

        ckpt_result = model.load_state_dict(
            torch.load(self.ckpt_path, map_location="cpu")["state_dict"], strict=False
        )
        exempt_prefixes = tuple(getattr(model, "checkpoint_exempt_key_prefixes", ()))
        unexpected_missing = [k for k in ckpt_result.missing_keys if not k.startswith(exempt_prefixes)]
        assert unexpected_missing == [], unexpected_missing
        model.eval()

        require_l1l2_model(model)

        self.model = model
        logger.info(f"Loaded Orbis 2 model from {self.ckpt_path} (CPU, not yet compiled)")
        return self

    def min_context_end_frame(self, l1_frame_rate):
        """Earliest frame index (0-based, in a video encoded at exactly `l1_frame_rate`
        fps) that has enough L1+L2 lookback to be a valid end of the context window.
        CPU-only (only reads model metadata) -- callable right after load(), before
        ensure_ready() puts anything on the GPU.

        `frame_interval=1` below assumes the video's native fps equals l1_frame_rate
        exactly, which is app2.py's re-encode invariant (CONTEXT_FPS == L1_FRAME_RATE);
        under that assumption l1_span == l1_context_frames, matching load_l1_l2_context's
        own math for the same case.
        """
        model = self.model
        l1_context_frames = int(model.vit.num_context_frames)
        l2_context_frames = int(model.condition_preprocessor.num_context_frames)
        l2_frame_rate = resolve_l2_frame_rate(model)

        indexer = _L1L2FrameIndexer(
            frame_interval=1,
            stored_data_frame_rate=l1_frame_rate,
            num_l2_context=l2_context_frames,
            l2_frame_rate=l2_frame_rate,
            l1_context_frames=l1_context_frames,
        )
        return indexer.get_required_l1_start_offset() + l1_context_frames - 1

    @property
    def frames_per_rollout_step(self):
        """Number of output frames each autoregressive rollout step yields
        (model.num_pred_frames). CPU-only (model metadata), callable right after
        load(). Lets callers convert the "rollout steps" the model actually takes
        into seconds of generated video (frames_per_rollout_step / l1_frame_rate)."""
        return int(self.model.num_pred_frames)

    def ensure_ready(self, device="cuda", compile=True, compile_mode="reduce-overhead",
                      compile_artifacts=None, speed_scale=1.0, yaw_rate_scale=1.0):
        """Move the model to `device` and wrap it with torch.compile, exactly once.
        Safe to call on every request: after the first call this is a no-op, because
        the model lives in this persistent process and keeps its GPU/compiled state
        between calls."""
        if self._ready:
            print("[compile] model already on-device and compiled from a prior request -- reusing it as-is.")
            return self.model

        with self._lock:
            if self._ready:
                return self.model

            model = self.model.to(device)

            if compile:
                def _maybe_compile(module, attr):
                    net = getattr(module, attr, None)
                    if net is not None:
                        print(f"[compile] wrapping {type(module).__name__}.{attr} in torch.compile(mode={compile_mode!r})")
                        setattr(module, attr, torch.compile(net, mode=compile_mode))
                        logger.info(f"Compiled {type(module).__name__}.{attr} with mode={compile_mode!r}")

                _maybe_compile(model, "ema_vit" if self.evaluate_ema else "vit")
                l2_predictor = getattr(getattr(model, "condition_preprocessor", None), "l2_predictor", None)
                if l2_predictor is not None:
                    _maybe_compile(l2_predictor, "ema_vit")

                if compile_artifacts:
                    if os.path.exists(compile_artifacts):
                        with open(compile_artifacts, "rb") as f:
                            torch.compiler.load_cache_artifacts(f.read())
                        print(f"[compile] loaded cached compile artifacts from {compile_artifacts!r} "
                              "-- first call reuses these instead of tracing from scratch.")
                        logger.info(f"Loaded compile artifacts from {compile_artifacts!r}")
                    else:
                        print(f"[compile] no compile artifacts found at {compile_artifacts!r} "
                              "-- compiling from scratch; will save artifacts after the first rollout.")
                        logger.info(
                            f"Compile artifacts not found at {compile_artifacts!r}; "
                            "will save after the first rollout."
                        )
                else:
                    print("[compile] no compile_artifacts path given -- compiling from scratch, nothing cached to disk.")
                self._compile_artifacts = compile_artifacts

            maybe_apply_condition_preprocessor_scales(model, speed_scale, yaw_rate_scale)

            self.model = model
            self._ready = True
            logger.info(f"Model moved to {device} and ready -- subsequent requests reuse this instance.")

        return self.model

    @torch.no_grad()
    def roll_out(self, video_path, output_dir, num_gen_frames, num_steps, l1_frame_rate,
                 height, width, seed, num_videos=1, l2_nfe=None, eta=0.0,
                 trajectory_file=None, vis_mode="none", decode_device="cpu", device="cuda",
                 context_end_frame=None, progress_cb=None):
        """One rollout request against the already-loaded, already-compiled model.
        Mirrors rollout_demo_v2.generate_images() from context loading onward -- the
        model-construction prefix runs once in load()/ensure_ready(), not here."""
        if not self._ready:
            raise RuntimeError("RolloutEngine.ensure_ready() must be called before roll_out().")

        model = self.model

        if int(seed) > 0:
            torch.backends.cudnn.enable = False
            torch.backends.cudnn.deterministic = True
            seed_everything(int(seed))

        maybe_apply_l2_nfe(model, l2_nfe)

        print(f"[solver] L1: {_describe_sampler(getattr(model, 'sampler', None), num_steps, eta)}")
        l2_predictor = getattr(model.condition_preprocessor, "l2_predictor", None)
        if l2_predictor is not None:
            l2_nfe = getattr(model.condition_preprocessor, "l2_pred_NFE", None)
            print(f"[solver] L2: {_describe_sampler(getattr(l2_predictor, 'sampler', None), l2_nfe, eta)}")

        height, width = int(height), int(width)
        backend = resolve_video_backend()
        l2_frame_rate = resolve_l2_frame_rate(model)

        l1_context_frames = int(model.vit.num_context_frames)
        l2_context_frames = int(model.condition_preprocessor.num_context_frames)

        start_frame = None if context_end_frame is None else int(context_end_frame) - l1_context_frames + 1

        print(f"[load] reading context from video: L1 {l1_context_frames} frames @ "
              f"{l1_frame_rate}fps, L2 {l2_context_frames} frames @ {l2_frame_rate}fps…")
        l1_tensor, l2_tensor = load_l1_l2_context(
            video_path=video_path,
            start_frame=start_frame,
            l1_frame_rate=l1_frame_rate,
            l2_frame_rate=l2_frame_rate,
            l1_context_frames=l1_context_frames,
            l2_context_frames=l2_context_frames,
            height=height,
            width=width,
            backend=backend,
            device=device,
        )
        print(f"[load] context ready: L1 tensor {tuple(l1_tensor.shape)}, "
              f"L2 tensor {tuple(l2_tensor.shape)}")

        num_future_frames = get_rollout_future_frame_count(model, num_gen_frames)
        frame_rate = torch.tensor(float(l1_frame_rate), device=device)
        data_batch = {"images": l1_tensor, "l2_context": l2_tensor, "frame_rate": frame_rate}

        get_required_steps = getattr(model.condition_preprocessor, "get_required_rollout_odometry_steps", None)
        min_odo_steps = None
        if callable(get_required_steps):
            min_odo_steps = get_required_steps(
                validation_params=None,
                num_condition_frames=l1_context_frames,
                num_gen_frames=num_future_frames,
                rollout_steps=num_gen_frames,
            )

        print(f"[load] preparing {'steering trajectory' if trajectory_file is not None else 'unconditional'} "
              "odometry conditioning…")
        if trajectory_file is not None:
            if min_odo_steps is None:
                raise ValueError(
                    "trajectory_file requires the model's condition_preprocessor to report a "
                    "required odometry length (get_required_rollout_odometry_steps returned None)."
                )
            odometry_steps_per_image_frame = getattr(model.condition_preprocessor, "odometry_steps_per_image_frame", 1)
            dt = 1.0 / (l1_frame_rate * odometry_steps_per_image_frame)
            traj_xy = load_trajectory_points(trajectory_file)
            traj_xy = resample_trajectory_by_arclength(traj_xy, min_odo_steps + 1)
            speed_yawrate = trajectory_to_speed_yawrate(traj_xy, dt)
            data_batch["steering"] = torch.as_tensor(
                speed_yawrate, dtype=l1_tensor.dtype, device=device
            ).unsqueeze(0)
        else:
            data_batch["steering"] = make_unconditional_steering(min_odo_steps, dtype=l1_tensor.dtype, device=device)
        data_batch["steering_format"] = STEERING_FORMAT

        num_videos = max(1, int(num_videos))
        if num_videos > 1:
            def _tile_batch(t):
                return t.repeat(num_videos, *([1] * (t.dim() - 1)))

            l1_tensor = _tile_batch(l1_tensor)
            l2_tensor = _tile_batch(l2_tensor)
            data_batch["images"] = l1_tensor
            data_batch["l2_context"] = l2_tensor
            data_batch["steering"] = _tile_batch(data_batch["steering"])

        condition_kwargs = model.condition_preprocessor.get_condition_kwargs_from_batch(data_batch, split="rollout")

        # The next call encodes the L1 context into latents and runs L2's own upfront
        # forecast (L1's autoregressive loop conditions on it, so it has to happen
        # before L1 can take a single step) -- neither has a progress hook, and on a
        # fresh torch/CUDA/GPU combo this is also where a from-scratch torch.compile
        # can silently eat a minute or more. No further output until this returns.
        print("[load] encoding context + L2 forecast (first call on this GPU/torch/CUDA "
              "combo may compile from scratch here -- can take a while, silently)…")
        autocast_enabled = device.startswith("cuda")
        with torch.autocast(dtype=torch.float16, device_type="cuda", enabled=autocast_enabled), \
             _rollout_progress_tqdm(progress_cb), \
             _l1_nfe_override(model, num_steps):
            _latents, gen_frames = model.roll_out(
                x_0={"images": l1_tensor},
                num_gen_frames=num_gen_frames,
                latent_input=False,
                NFE=num_steps,
                eta=eta,
                sample_with_ema=self.evaluate_ema,
                num_samples=l1_tensor.size(0),
                frame_rate=frame_rate.reshape(1).repeat(l1_tensor.size(0)),
                condition_kwargs=condition_kwargs,
                decode_device=decode_device,
                num_condition_frames=l1_tensor.size(1),
            )

        if vis_mode in {"trajectory", "trajectory_ego"}:
            overlay_trajectory = model.condition_preprocessor.get_rollout_visualization_trajectory(
                condition_kwargs=model.condition_preprocessor.get_condition_kwargs_from_batch(data_batch, split="rollout"),
                num_condition_frames=l1_context_frames,
                num_gen_steps=num_gen_frames,
                num_pred_frames=model.num_pred_frames,
            )
            if overlay_trajectory is not None:
                gen_frames = overlay_trajectory_on_images(gen_frames, overlay_trajectory, mode=vis_mode)

        del _latents, condition_kwargs, l1_tensor, l2_tensor

        os.makedirs(output_dir, exist_ok=True)
        num_out = gen_frames.shape[0]
        num_frames = gen_frames.shape[1]
        for b in range(num_out):
            seq_dir = os.path.join(output_dir, "fake_images", f"sequence_{b:04d}")
            os.makedirs(seq_dir, exist_ok=True)
            for f in range(num_frames):
                save_image(
                    (gen_frames[b, f] + 1.0) / 2.0,
                    os.path.join(seq_dir, f"frame_{f:04d}.jpg"),
                )

        if device.startswith("cuda"):
            logger.info(f"Max memory: {torch.cuda.max_memory_allocated() / 1024**3:.02f} GB")

        if self._compile_artifacts and not self._artifacts_saved and not os.path.exists(self._compile_artifacts):
            artifacts = torch.compiler.save_cache_artifacts()
            if artifacts is not None:
                with open(self._compile_artifacts, "wb") as f:
                    f.write(artifacts[0])
                logger.info(f"Saved compile artifacts to {self._compile_artifacts!r}")
            self._artifacts_saved = True


_engine = None
_engine_lock = threading.Lock()


def get_engine(exp_dir, config_name, ckpt_name, evaluate_ema=True):
    """Return the process-wide RolloutEngine, constructing and CPU-loading it on
    first call. Subsequent calls (with the same or different args -- args are only
    used for the first, initializing call) just return the cached instance."""
    global _engine
    if _engine is None:
        with _engine_lock:
            if _engine is None:
                _engine = RolloutEngine(exp_dir, config_name, ckpt_name, evaluate_ema=evaluate_ema).load()
    return _engine