Spaces:
Sleeping
Sleeping
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
|