From 59e9a8575cd0a73e174a217efc626c9dd9b63484 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 00:28:20 +0000 Subject: [PATCH 01/51] Add fake demo runtime vertical slice --- .../flashdreams/runtime/demo/__init__.py | 63 ++ .../flashdreams/runtime/demo/drivers.py | 222 ++++++ flashdreams/flashdreams/runtime/demo/host.py | 66 ++ .../flashdreams/runtime/demo/outputs.py | 116 ++- .../flashdreams/runtime/demo/pipeline.py | 75 ++ .../flashdreams/runtime/demo/run_modes.py | 413 +++++++++++ .../runtime/demo/session_inputs.py | 150 ++++ .../tests/test_demo_runtime_vertical_slice.py | 682 ++++++++++++++++++ 8 files changed, 1785 insertions(+), 2 deletions(-) create mode 100644 flashdreams/flashdreams/runtime/demo/drivers.py create mode 100644 flashdreams/flashdreams/runtime/demo/host.py create mode 100644 flashdreams/flashdreams/runtime/demo/pipeline.py create mode 100644 flashdreams/flashdreams/runtime/demo/run_modes.py create mode 100644 flashdreams/flashdreams/runtime/demo/session_inputs.py create mode 100644 flashdreams/tests/test_demo_runtime_vertical_slice.py diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index 7a0535556..04883d8bf 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -3,8 +3,43 @@ """Experimental shared demo API above the inference runtime API.""" +from flashdreams.runtime.demo.drivers import ( + BatchSessionDriver, + DriverInvariantError, + run_demo_session, +) +from flashdreams.runtime.demo.host import RuntimeHost from flashdreams.runtime.demo.outputs import build_output_target +from flashdreams.runtime.demo.outputs import ( + NullOutputSink, + OutputDecision, + OutputSink, + SessionInfo, +) +from flashdreams.runtime.demo.pipeline import StepOutcome, StepPipeline from flashdreams.runtime.demo.replay import run_replay_demo +from flashdreams.runtime.demo.run_modes import ( + DefaultErrorPolicy, + ErrorAction, + InMemorySessionMetricsRecorder, + MetricsSnapshot, + NoopTransportService, + RunContext, + RunMode, + RunResult, + RunSummary, + SessionEdges, + SingleSessionAdmissionPolicy, +) +from flashdreams.runtime.demo.session_inputs import ( + BatchInputSource, + ControlDecision, + InputSource, + ModelInputProvider, + PreparedStep, + RealtimeInputSource, + UserInputWindow, +) from flashdreams.runtime.demo.spec import ( DemoAdapter, DemoSpec, @@ -17,14 +52,42 @@ ) __all__ = [ + "BatchInputSource", + "BatchSessionDriver", + "ControlDecision", + "DefaultErrorPolicy", "DemoAdapter", "DemoSpec", + "DriverInvariantError", + "ErrorAction", + "InMemorySessionMetricsRecorder", + "InputSource", + "MetricsSnapshot", + "ModelInputProvider", "Mp4OutputSpec", + "NoopTransportService", "NullOutputSpec", + "NullOutputSink", + "OutputDecision", "OutputSpec", + "OutputSink", "PreparedScenario", + "PreparedStep", + "RealtimeInputSource", + "RunContext", + "RunMode", + "RunResult", + "RunSummary", + "RuntimeHost", + "SessionEdges", + "SessionInfo", + "SingleSessionAdmissionPolicy", + "StepOutcome", + "StepPipeline", + "UserInputWindow", "WebRTCAppResources", "WebRTCOutputSpec", "build_output_target", + "run_demo_session", "run_replay_demo", ] diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py new file mode 100644 index 000000000..e6050f6ec --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -0,0 +1,222 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Session drivers and helpers for demo runtime vertical slices.""" + +from __future__ import annotations + +from typing import Any + +from flashdreams.runtime.interfaces import InferenceSession + +from .host import RuntimeHost +from .outputs import SessionInfo +from .pipeline import StepPipeline +from .run_modes import ( + DriverStatus, + RunContext, + RunMode, + RunResult, + SessionEdges, + SessionReservation, +) +from .session_inputs import ModelInputProvider +from .spec import DemoAdapter, DemoSpec, PreparedScenario + + +class DriverInvariantError(RuntimeError): + """A driver invariant was violated; this is a driver bug, not a run result.""" + + +class BatchSessionDriver: + """Minimal finite-session driver for Phase 2 fake-model coverage.""" + + def run_one_session( + self, + *, + host: RuntimeHost, + provider: ModelInputProvider, + session_edges: SessionEdges, + pipeline: StepPipeline, + ) -> RunResult: + session: InferenceSession | None = None + final_status: DriverStatus = "completed" + final_reason: str | None = None + final_error: Exception | None = None + setup_ok = False + try: + try: + initial_input = host.call(provider.prepare_initial_input) + session = host.call(host.start_session, initial_input) + session_info = host.call(_session_info, session) + session_edges.output_sink.open(session_info) + setup_ok = True + except Exception as exc: + action = session_edges.error_policy.handle_setup_error(exc) + if action.drop_chunk or action.result_status == "completed": + raise DriverInvariantError( + "Setup failures must resolve to failed or skipped." + ) from exc + session_edges.metrics.record_error(exc, action) + final_status = action.result_status + final_reason = str(exc) + final_error = exc if action.result_status == "failed" else None + + while setup_ok: + if session is None: + raise DriverInvariantError("setup_ok was set without a session.") + try: + if session_edges.input_source.is_finished(): + break + request = host.call(session.next_step_request) + if request is None: + break + user_window = session_edges.input_source.next_window(request) + outcome = host.call( + pipeline.execute_step, + request=request, + user_window=user_window, + provider=provider, + session=session, + output=session_edges.output_sink, + metrics=session_edges.metrics, + ) + if outcome.control.reset: + host.call(session.reset, outcome.control.reset_input) + if not outcome.control.provider_already_reset: + host.call(provider.reset, outcome.control.reset_input) + continue + if outcome.control.close_session: + break + if outcome.output.should_stop: + break + except DriverInvariantError: + raise + except Exception as exc: + action = session_edges.error_policy.handle(exc) + session_edges.metrics.record_error(exc, action) + if action.drop_chunk: + continue + final_status = action.result_status + final_reason = str(exc) + final_error = exc if action.result_status == "failed" else None + break + except DriverInvariantError: + raise + except Exception as exc: + final_status = "failed" + final_reason = str(exc) + final_error = exc + finally: + if session is not None: + host.call(_close_safely, session.close, session_edges) + host.call(_close_safely, provider.close, session_edges) + + return session_edges.close_result( + status=final_status, + reason=final_reason, + error=final_error, + ) + + +def run_demo_session( + *, + context: RunContext, + spec: DemoSpec, + scenario: PreparedScenario, + adapter: DemoAdapter, + run_mode: RunMode, + pipeline: StepPipeline, + reservation: SessionReservation | None = None, +) -> RunResult: + """Run one prepared demo session through a selected run mode.""" + reservation = reservation or context.admission.try_reserve() + if reservation is None: + result = RunResult.rejected(reason="busy") + context.run_metrics.record_session(result) + return result + + provider: Any | None = None + session_edges: SessionEdges | None = None + driver_started = False + try: + create_provider = getattr(adapter, "create_model_input_provider") + provider = context.host.call(create_provider, spec, scenario) + run_mode.validate_session( + spec=spec, + scenario=scenario, + adapter=adapter, + provider=provider, + ) + session_edges = run_mode.create_session_edges( + context=context, + spec=spec, + scenario=scenario, + provider=provider, + adapter=adapter, + ) + driver = run_mode.select_driver() + if not isinstance(driver, BatchSessionDriver): + raise TypeError( + "Phase 2 run_demo_session supports BatchSessionDriver only, " + f"got {type(driver).__name__}." + ) + driver_started = True + result = driver.run_one_session( + host=context.host, + provider=provider, + session_edges=session_edges, + pipeline=pipeline, + ) + context.run_metrics.record_session(result) + return result + except DriverInvariantError: + raise + except Exception as exc: + if provider is not None and not driver_started: + try: + context.host.call(provider.close) + except Exception as close_exc: + if session_edges is not None: + session_edges.metrics.record_cleanup_error(close_exc) + else: + context.run_metrics.record_cleanup_error(close_exc) + if session_edges is not None: + result = session_edges.close_result( + status="failed", + reason=str(exc), + error=exc, + ) + else: + result = RunResult(status="failed", reason=str(exc), error=exc) + context.run_metrics.record_session(result) + return result + finally: + reservation.release() + + +def _session_info(session: InferenceSession) -> SessionInfo: + session_info = getattr(session, "session_info", None) + if not callable(session_info): + return SessionInfo() + value = session_info() + if not isinstance(value, SessionInfo): + raise TypeError( + "session.session_info() must return SessionInfo, " + f"got {type(value).__name__}." + ) + return value + + +def _close_safely(close: Any, session_edges: SessionEdges) -> None: + try: + close() + except Exception as exc: + session_edges.metrics.record_cleanup_error(exc) + + +__all__ = [ + "BatchSessionDriver", + "DriverInvariantError", + "run_demo_session", +] diff --git a/flashdreams/flashdreams/runtime/demo/host.py b/flashdreams/flashdreams/runtime/demo/host.py new file mode 100644 index 000000000..ee870ebc5 --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/host.py @@ -0,0 +1,66 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Minimal runtime host for the Phase 2 demo-session vertical slice.""" + +from __future__ import annotations + +from collections.abc import Callable +from typing import TypeVar + +from flashdreams.runtime.inputs import InferenceInput +from flashdreams.runtime.interfaces import InferenceRuntime, InferenceSession + +_T = TypeVar("_T") + + +class RuntimeHost: + """Thin synchronous host around an :class:`InferenceRuntime`. + + Phase 3 moves the thread-affine worker boundary here. Phase 2 keeps the + dispatch direct so fake-model CPU tests can prove the session-driver shape + without introducing worker behavior early. + """ + + def __init__(self, runtime: InferenceRuntime) -> None: + self._runtime = runtime + self._healthy = True + + @property + def runtime(self) -> InferenceRuntime: + """Return the hosted runtime.""" + return self._runtime + + @property + def is_healthy(self) -> bool: + """Return whether admission should continue accepting sessions.""" + return self._healthy + + def mark_unhealthy(self) -> None: + """Latch the host as unhealthy.""" + self._healthy = False + + def call(self, func: Callable[..., _T], /, *args: object, **kwargs: object) -> _T: + """Run one model-affine callable synchronously.""" + return func(*args, **kwargs) + + async def call_async( + self, + func: Callable[..., _T], + /, + *args: object, + **kwargs: object, + ) -> _T: + """Async-compatible direct dispatch placeholder for Phase 3.""" + return self.call(func, *args, **kwargs) + + def start_session(self, inputs: InferenceInput) -> InferenceSession: + """Start one inference session through the hosted runtime.""" + return self._runtime.start_session(inputs) + + def close(self) -> None: + """Close the hosted runtime.""" + self._runtime.close() + + +__all__ = ["RuntimeHost"] diff --git a/flashdreams/flashdreams/runtime/demo/outputs.py b/flashdreams/flashdreams/runtime/demo/outputs.py index 421ec3bb4..be4fd2bda 100644 --- a/flashdreams/flashdreams/runtime/demo/outputs.py +++ b/flashdreams/flashdreams/runtime/demo/outputs.py @@ -1,18 +1,124 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Shared demo output-target construction.""" +"""Shared demo output contracts and output-target construction.""" from __future__ import annotations +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field from pathlib import Path +from typing import Literal, Protocol, runtime_checkable +from flashdreams.runtime._utils import freeze_mapping +from flashdreams.runtime.output import OutputArtifact from flashdreams.runtime.output import NullOutputTarget, OutputTarget +from flashdreams.runtime.types import StepResult from flashdreams.runtime.video_output import Mp4VideoOutputTarget, VideoWriter from .spec import Mp4OutputSpec, NullOutputSpec, OutputSpec, WebRTCOutputSpec +@dataclass(frozen=True, kw_only=True, slots=True) +class SessionInfo: + """Output-facing metadata known after session setup.""" + + output_layout: str | None = None + steady_output_frame_count: int | None = None + metadata: Mapping[str, object] = field(default_factory=dict) + + def __post_init__(self) -> None: + if self.output_layout is not None and not self.output_layout.strip(): + raise ValueError("SessionInfo.output_layout must be non-empty when set.") + if ( + self.steady_output_frame_count is not None + and self.steady_output_frame_count < 0 + ): + raise ValueError( + "SessionInfo.steady_output_frame_count must be >= 0 when set." + ) + object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class OutputDecision: + """Flow-control decision returned by an output sink after one step.""" + + should_stop: bool = False + dropped: bool = False + drop_policy: Literal["none", "drop_newest", "drop_oldest"] = "none" + backpressure_s: float = 0.0 + metadata: Mapping[str, object] = field(default_factory=dict) + + def __post_init__(self) -> None: + if self.drop_policy not in {"none", "drop_newest", "drop_oldest"}: + raise ValueError(f"Unsupported drop_policy={self.drop_policy!r}.") + if self.backpressure_s < 0: + raise ValueError("OutputDecision.backpressure_s must be >= 0.") + object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) + + +@runtime_checkable +class OutputSink(Protocol): + """Consumes generated session outputs for a demo run mode.""" + + produces_artifacts: bool + + def open(self, session_info: SessionInfo) -> None: + """Prepare output resources for a session.""" + ... + + def begin_generation(self, generation: int) -> None: + """Start an output generation, discarding stale live output if needed.""" + ... + + def write(self, result: StepResult) -> OutputDecision: + """Consume one generated result and return output flow-control state.""" + ... + + def close(self) -> Sequence[OutputArtifact]: + """Finalize output resources and return produced artifacts.""" + ... + + +@dataclass(slots=True) +class NullOutputSink: + """Output sink for headless runs and fake-model vertical-slice tests.""" + + store_results: bool = False + produces_artifacts: bool = False + output_count: int = field(default=0, init=False) + results: list[StepResult] = field(default_factory=list, init=False) + opened: bool = field(default=False, init=False) + closed: bool = field(default=False, init=False) + session_info: SessionInfo | None = field(default=None, init=False) + generation: int | None = field(default=None, init=False) + + def open(self, session_info: SessionInfo) -> None: + self.session_info = session_info + self.output_count = 0 + self.results.clear() + self.opened = True + self.closed = False + + def begin_generation(self, generation: int) -> None: + if generation < 0: + raise ValueError("generation must be >= 0.") + self.generation = generation + + def write(self, result: StepResult) -> OutputDecision: + if not self.opened or self.closed: + raise RuntimeError("Cannot write to a closed output sink.") + self.output_count += 1 + if self.store_results: + self.results.append(result) + return OutputDecision() + + def close(self) -> Sequence[OutputArtifact]: + self.closed = True + return () + + def build_output_target( output: OutputSpec, *, @@ -42,4 +148,10 @@ def build_output_target( raise TypeError(f"Unsupported demo output spec: {type(output).__name__}.") -__all__ = ["build_output_target"] +__all__ = [ + "NullOutputSink", + "OutputDecision", + "OutputSink", + "SessionInfo", + "build_output_target", +] diff --git a/flashdreams/flashdreams/runtime/demo/pipeline.py b/flashdreams/flashdreams/runtime/demo/pipeline.py new file mode 100644 index 000000000..bca09ba18 --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/pipeline.py @@ -0,0 +1,75 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Shared per-step pipeline for demo session drivers.""" + +from __future__ import annotations + +from dataclasses import dataclass, field + +from flashdreams.runtime.interfaces import InferenceSession +from flashdreams.runtime.types import StepRequest, StepResult + +from .outputs import OutputDecision, OutputSink +from .run_modes import SessionMetricsRecorder +from .session_inputs import ControlDecision, ModelInputProvider, UserInputWindow + + +@dataclass(frozen=True, kw_only=True, slots=True) +class StepOutcome: + """Combined output and control result from one shared model step.""" + + output: OutputDecision = field(default_factory=OutputDecision) + control: ControlDecision = field(default_factory=ControlDecision) + + +class StepPipeline: + """Shared invariant for provider conversion, model step, output, and metrics.""" + + def execute_step( + self, + *, + request: StepRequest, + user_window: UserInputWindow, + provider: ModelInputProvider, + session: InferenceSession, + output: OutputSink, + metrics: SessionMetricsRecorder, + ) -> StepOutcome: + prepared = provider.prepare_step( + request=request, + user_window=user_window, + ) + if prepared.control.reset or prepared.control.close_session: + metrics.record_control( + request=request, + user_window=user_window, + control=prepared.control, + ) + return StepOutcome(control=prepared.control) + if prepared.inference_input is None: + raise RuntimeError("ModelInputProvider returned no inference input.") + + result = session.step(prepared.inference_input) + if not isinstance(result, StepResult): + raise TypeError( + "InferenceSession.step must return StepResult, " + f"got {type(result).__name__}." + ) + decision = output.write(result) + if not isinstance(decision, OutputDecision): + raise TypeError( + "OutputSink.write must return OutputDecision, " + f"got {type(decision).__name__}." + ) + metrics.record_step( + request=request, + user_window=user_window, + inference_input=prepared.inference_input, + result=result, + decision=decision, + ) + return StepOutcome(output=decision) + + +__all__ = ["StepOutcome", "StepPipeline"] diff --git a/flashdreams/flashdreams/runtime/demo/run_modes.py b/flashdreams/flashdreams/runtime/demo/run_modes.py new file mode 100644 index 000000000..335a7fb96 --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/run_modes.py @@ -0,0 +1,413 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Run/session result and policy helpers for demo session drivers.""" + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from threading import Lock +from typing import Any, Literal, Protocol, TYPE_CHECKING, runtime_checkable + +from flashdreams.runtime._utils import freeze_mapping +from flashdreams.runtime.output import OutputArtifact + +from .outputs import OutputDecision, OutputSink + +if TYPE_CHECKING: + from .host import RuntimeHost + from .session_inputs import BatchInputSource, ModelInputProvider + from .spec import DemoAdapter, DemoSpec, PreparedScenario + +SessionStatus = Literal[ + "completed", + "failed", + "skipped", + "cancelled", + "rejected", + "not_activated", +] + +DriverStatus = Literal[ + "completed", + "failed", + "skipped", + "cancelled", + "not_activated", +] + + +@dataclass(frozen=True, kw_only=True, slots=True) +class MetricsSnapshot: + """Closed session or run metrics summary.""" + + counters: Mapping[str, int | float] = field(default_factory=dict) + timings: Mapping[str, Sequence[float]] = field(default_factory=dict) + session_statuses: Sequence[str] = () + errors: Sequence[str] = () + + def __post_init__(self) -> None: + object.__setattr__(self, "counters", freeze_mapping(self.counters)) + object.__setattr__( + self, + "timings", + freeze_mapping({key: tuple(values) for key, values in self.timings.items()}), + ) + object.__setattr__(self, "session_statuses", tuple(self.session_statuses)) + object.__setattr__(self, "errors", tuple(self.errors)) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class RunResult: + """Outcome of one demo session.""" + + __hash__ = None + + status: SessionStatus + artifacts: Sequence[OutputArtifact] = () + metrics: MetricsSnapshot | None = None + reason: str | None = None + error: Exception | None = None + + @classmethod + def rejected(cls, reason: str) -> "RunResult": + """Admission refused the session. The only no-session result helper.""" + return cls(status="rejected", reason=reason) + + def __post_init__(self) -> None: + object.__setattr__(self, "artifacts", tuple(self.artifacts)) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class RunSummary: + """Summary for a run context after one or more sessions.""" + + metrics: MetricsSnapshot + sessions: Sequence[RunResult] = () + + def __post_init__(self) -> None: + object.__setattr__(self, "sessions", tuple(self.sessions)) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class ErrorAction: + """Driver policy decision for an operational error.""" + + close_session: bool = True + drop_chunk: bool = False + continue_next_scenario: bool = False + result_status: Literal["completed", "failed", "skipped"] = "failed" + + +class DefaultErrorPolicy: + """Default policy: operational errors fail the current session.""" + + def handle_setup_error(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status="failed") + + def handle(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status="failed") + + +@runtime_checkable +class ErrorPolicy(Protocol): + """Maps driver-observed exceptions to session outcomes.""" + + def handle_setup_error(self, exc: Exception) -> ErrorAction: ... + + def handle(self, exc: Exception) -> ErrorAction: ... + + +@runtime_checkable +class SessionMetricsRecorder(Protocol): + """Metrics callbacks consumed by the Phase 2 drivers and pipeline.""" + + def record_step( + self, + *, + request: object, + user_window: object, + inference_input: object, + result: object, + decision: OutputDecision, + ) -> None: ... + + def record_control( + self, + *, + request: object, + user_window: object, + control: object, + ) -> None: ... + + def record_error(self, exc: Exception, action: ErrorAction) -> None: ... + + def record_cleanup_error(self, exc: Exception) -> None: ... + + def record_session(self, result: RunResult) -> None: ... + + def close(self) -> MetricsSnapshot: ... + + +@dataclass(slots=True) +class InMemorySessionMetricsRecorder: + """Small non-raising metrics recorder for driver tests and fake demos.""" + + step_count: int = 0 + control_count: int = 0 + errors: list[str] = field(default_factory=list) + cleanup_errors: list[str] = field(default_factory=list) + sessions: list[RunResult] = field(default_factory=list) + closed: bool = False + + def record_step( + self, + *, + request: object, + user_window: object, + inference_input: object, + result: object, + decision: OutputDecision, + ) -> None: + del request, user_window, inference_input, result, decision + if not self.closed: + self.step_count += 1 + + def record_control( + self, + *, + request: object, + user_window: object, + control: object, + ) -> None: + del request, user_window, control + if not self.closed: + self.control_count += 1 + + def record_error(self, exc: Exception, action: ErrorAction) -> None: + del action + if not self.closed: + self.errors.append(str(exc)) + + def record_cleanup_error(self, exc: Exception) -> None: + if not self.closed: + self.cleanup_errors.append(str(exc)) + + def record_session(self, result: RunResult) -> None: + if not self.closed: + self.sessions.append(result) + + def close(self) -> MetricsSnapshot: + self.closed = True + return MetricsSnapshot( + counters={ + "steps": self.step_count, + "controls": self.control_count, + "sessions": len(self.sessions), + "cleanup_errors": len(self.cleanup_errors), + }, + session_statuses=tuple(result.status for result in self.sessions), + errors=tuple((*self.errors, *self.cleanup_errors)), + ) + + +class NoopTransportService: + """Idempotent placeholder transport for batch sessions.""" + + def __init__(self) -> None: + self.closed = False + + def is_active(self) -> bool: + return not self.closed + + def close(self) -> None: + self.closed = True + + +@runtime_checkable +class TransportService(Protocol): + """Per-session transport lifecycle hook.""" + + def is_active(self) -> bool: ... + + def close(self) -> None: ... + + +@runtime_checkable +class SessionReservation(Protocol): + """Admission reservation for one session.""" + + def release(self) -> None: ... + + +class SingleSessionAdmissionPolicy: + """Atomic single-session admission policy.""" + + def __init__(self, *, health_check: Any | None = None) -> None: + self._lock = Lock() + self._reserved = False + self._health_check = health_check + + def try_reserve(self) -> SessionReservation | None: + with self._lock: + if self._reserved or not self._is_healthy(): + return None + self._reserved = True + return _SingleSessionReservation(self) + + def _release(self) -> None: + with self._lock: + self._reserved = False + + def _is_healthy(self) -> bool: + if self._health_check is None: + return True + return bool(self._health_check()) + + +class _SingleSessionReservation: + def __init__(self, policy: SingleSessionAdmissionPolicy) -> None: + self._policy = policy + self._released = False + self.release_count = 0 + + def release(self) -> None: + if self._released: + return + self._released = True + self.release_count += 1 + self._policy._release() + + +@runtime_checkable +class AdmissionPolicy(Protocol): + """Atomically reserves session capacity or rejects.""" + + def try_reserve(self) -> SessionReservation | None: ... + + +@dataclass(slots=True) +class RunContext: + """Run-scoped services shared by one or more demo sessions.""" + + host: "RuntimeHost" + run_metrics: SessionMetricsRecorder + admission: AdmissionPolicy + services: Mapping[str, object] = field(default_factory=dict) + + def __post_init__(self) -> None: + self.services = freeze_mapping(self.services) + + def close(self) -> RunSummary: + for service in self.services.values(): + close = getattr(service, "close", None) + if callable(close): + try: + close() + except Exception as exc: + self.run_metrics.record_cleanup_error(exc) + return RunSummary( + metrics=self.run_metrics.close(), + sessions=tuple(getattr(self.run_metrics, "sessions", ())), + ) + + +@dataclass(slots=True) +class SessionEdges: + """Per-session input/output/policy bundle consumed by drivers.""" + + input_source: "BatchInputSource" + output_sink: OutputSink + metrics: SessionMetricsRecorder = field( + default_factory=InMemorySessionMetricsRecorder + ) + error_policy: ErrorPolicy = field(default_factory=DefaultErrorPolicy) + transport: TransportService = field(default_factory=NoopTransportService) + clock: object | None = None + activation: object | None = None + cleanup_tasks: set[object] = field(default_factory=set) + _closed_result: RunResult | None = field(default=None, init=False, repr=False) + + def close_result( + self, + *, + status: DriverStatus = "completed", + reason: str | None = None, + error: Exception | None = None, + ) -> RunResult: + """Idempotently close output, transport, and metrics once.""" + if self._closed_result is not None: + return self._closed_result + + artifacts: Sequence[OutputArtifact] = () + try: + artifacts = tuple(self.output_sink.close()) + except Exception as exc: + self.metrics.record_cleanup_error(exc) + try: + self.transport.close() + except Exception as exc: + self.metrics.record_cleanup_error(exc) + try: + metrics = self.metrics.close() + except Exception as exc: + metrics = MetricsSnapshot(errors=(f"metrics.close failed: {exc}",)) + self._closed_result = RunResult( + status=status, + artifacts=artifacts, + metrics=metrics, + reason=reason, + error=error, + ) + return self._closed_result + + +@runtime_checkable +class RunMode(Protocol): + """Minimal Phase 2 run-mode surface used by ``run_demo_session``.""" + + def validate_session( + self, + *, + spec: "DemoSpec", + scenario: "PreparedScenario", + adapter: "DemoAdapter", + provider: "ModelInputProvider", + ) -> None: ... + + def create_session_edges( + self, + *, + context: RunContext, + spec: "DemoSpec", + scenario: "PreparedScenario", + provider: "ModelInputProvider", + adapter: "DemoAdapter", + ) -> SessionEdges: ... + + def select_driver(self) -> object: ... + + +__all__ = [ + "AdmissionPolicy", + "DefaultErrorPolicy", + "DriverStatus", + "ErrorAction", + "ErrorPolicy", + "InMemorySessionMetricsRecorder", + "MetricsSnapshot", + "NoopTransportService", + "RunContext", + "RunMode", + "RunResult", + "RunSummary", + "SessionEdges", + "SessionMetricsRecorder", + "SessionReservation", + "SessionStatus", + "SingleSessionAdmissionPolicy", + "TransportService", +] diff --git a/flashdreams/flashdreams/runtime/demo/session_inputs.py b/flashdreams/flashdreams/runtime/demo/session_inputs.py new file mode 100644 index 000000000..cf6838961 --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/session_inputs.py @@ -0,0 +1,150 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Input-source and model-input-provider contracts for demo sessions.""" + +from __future__ import annotations + +import math +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from typing import Protocol, runtime_checkable + +from flashdreams.runtime._utils import freeze_mapping +from flashdreams.runtime.inputs import InferenceInput, UserInputs +from flashdreams.runtime.types import StepRequest + + +@dataclass(frozen=True, kw_only=True, slots=True) +class ControlDecision: + """Provider-authored control request for the current session.""" + + reset: bool = False + close_session: bool = False + reset_input: InferenceInput | None = None + provider_already_reset: bool = False + reason: str | None = None + + def __post_init__(self) -> None: + if self.reason is not None and not self.reason.strip(): + raise ValueError("ControlDecision.reason must be non-empty when set.") + + +@dataclass(frozen=True, kw_only=True, slots=True) +class UserInputWindow: + """User/app inputs selected by a driver for one model step.""" + + __hash__ = None + + start_s: float + end_s: float + frame_times: Sequence[float] = () + inputs: UserInputs = field(default_factory=UserInputs) + control: ControlDecision | None = None + metadata: Mapping[str, object] = field(default_factory=dict) + + def __post_init__(self) -> None: + if not math.isfinite(self.start_s) or self.start_s < 0: + raise ValueError("UserInputWindow.start_s must be finite and >= 0.") + if not math.isfinite(self.end_s) or self.end_s < self.start_s: + raise ValueError( + "UserInputWindow.end_s must be finite and >= start_s." + ) + previous = -math.inf + for frame_time in self.frame_times: + if not math.isfinite(float(frame_time)): + raise ValueError("UserInputWindow.frame_times must be finite.") + if float(frame_time) < previous: + raise ValueError( + "UserInputWindow.frame_times must be sorted in ascending order." + ) + previous = float(frame_time) + object.__setattr__(self, "frame_times", tuple(float(t) for t in self.frame_times)) + object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class PreparedStep: + """Model-facing input plus optional provider-authored control decision.""" + + __hash__ = None + + inference_input: InferenceInput | None = None + control: ControlDecision = field(default_factory=ControlDecision) + + +@runtime_checkable +class InputSource(Protocol): + """Facts common to every demo session input source.""" + + is_finite: bool + is_deterministic: bool + + def is_finished(self) -> bool: + """Return whether the driver should stop requesting windows.""" + ... + + +@runtime_checkable +class BatchInputSource(InputSource, Protocol): + """Finite input source consumed by the batch driver.""" + + def next_window(self, request: StepRequest) -> UserInputWindow: + """Return the next batch input window for ``request``.""" + ... + + +@runtime_checkable +class RealtimeInputSource(InputSource, Protocol): + """Realtime input source consumed by a future realtime driver.""" + + async def next_realtime_window( + self, + *, + request: StepRequest, + clock: object, + ) -> object: + """Return the next realtime window result. + + The concrete realtime result shape lands with the realtime clock phase. + Keeping this protocol separate now prevents batch sources from stubbing + async behavior they never serve. + """ + ... + + +@runtime_checkable +class ModelInputProvider(Protocol): + """Model-owned conversion from user windows into model-facing inputs.""" + + def prepare_initial_input(self) -> InferenceInput: + """Prepare session-global model inputs.""" + ... + + def prepare_step( + self, + *, + request: StepRequest, + user_window: UserInputWindow, + ) -> PreparedStep: + """Prepare one model step from a driver-owned user input window.""" + ... + + def reset(self, inputs: InferenceInput | None = None) -> None: + """Reset provider-owned session state.""" + ... + + def close(self) -> None: + """Release provider-owned resources.""" + ... + + +__all__ = [ + "BatchInputSource", + "ControlDecision", + "InputSource", + "ModelInputProvider", + "PreparedStep", + "RealtimeInputSource", + "UserInputWindow", +] diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py new file mode 100644 index 000000000..e7a3dc303 --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -0,0 +1,682 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from collections.abc import Callable, Sequence +from typing import Any, Literal + +import pytest + +from flashdreams.runtime import ( + CanonicalInputSchema, + IdentityInputMapping, + InferenceConfig, + InferenceInput, + InferenceInputSchema, + InferenceRuntime, + InferenceSession, + InputMapping, + OutputArtifact, + StepRequest, + StepResult, + UserInputs, +) +from flashdreams.runtime.demo import ( + BatchSessionDriver, + ControlDecision, + DemoSpec, + DriverInvariantError, + ErrorAction, + InMemorySessionMetricsRecorder, + NullOutputSpec, + OutputDecision, + PreparedScenario, + PreparedStep, + RunContext, + RunResult, + RuntimeHost, + SessionEdges, + SessionInfo, + SingleSessionAdmissionPolicy, + StepPipeline, + UserInputWindow, + run_demo_session, +) + +pytestmark = pytest.mark.ci_cpu + + +def test_step_pipeline_passes_provider_input_to_session_and_sink() -> None: + provider = _FakeVideoModelInputProvider() + session = _FakeVideoSession(num_steps=1) + output = _RecordingOutputSink() + output.open(SessionInfo(output_layout="fake-video", steady_output_frame_count=1)) + metrics = InMemorySessionMetricsRecorder() + request = StepRequest(step_index=0) + user_window = _window(0) + + outcome = StepPipeline().execute_step( + request=request, + user_window=user_window, + provider=provider, + session=session, + output=output, + metrics=metrics, + ) + + assert outcome == _empty_step_outcome() + assert session.step_inputs == provider.prepared_step_inputs + assert [result.output for result in output.results] == ["frame-0"] + assert metrics.step_count == 1 + + +def test_batch_driver_runs_fake_video_demo_through_runtime_host() -> None: + session = _FakeVideoSession(num_steps=2) + runtime = _FakeVideoRuntime(session=session) + host = _RecordingRuntimeHost(runtime) + provider = _FakeVideoModelInputProvider() + output = _RecordingOutputSink() + metrics = InMemorySessionMetricsRecorder() + edges = SessionEdges( + input_source=_FakeBatchInputSource(num_windows=2), + output_sink=output, + metrics=metrics, + ) + + result = BatchSessionDriver().run_one_session( + host=host, + provider=provider, + session_edges=edges, + pipeline=StepPipeline(), + ) + + assert result.status == "completed" + assert result.metrics is not None + assert result.metrics.counters["steps"] == 2 + assert runtime.start_session_inputs == [provider.initial_input] + assert [dict(inputs.step) for inputs in session.step_inputs] == [ + {"request_step": 0, "window": (0.0, 1.0)}, + {"request_step": 1, "window": (1.0, 2.0)}, + ] + assert [result.output for result in output.results] == ["frame-0", "frame-1"] + assert output.opened_with == SessionInfo( + output_layout="fake-video", + steady_output_frame_count=1, + ) + assert session.close_count == 1 + assert provider.close_count == 1 + assert host.calls.count("execute_step") == 2 + assert "prepare_initial_input" in host.calls + assert "start_session" in host.calls + + +def test_run_demo_session_builds_edges_and_records_session_once() -> None: + session = _FakeVideoSession(num_steps=1) + runtime = _FakeVideoRuntime(session=session) + run_metrics = InMemorySessionMetricsRecorder() + context = _run_context(runtime, run_metrics=run_metrics) + provider = _FakeVideoModelInputProvider() + adapter = _FakeDemoAdapter(provider=provider) + output = _RecordingOutputSink() + factory_calls: list[tuple[DemoSpec, PreparedScenario]] = [] + run_mode = _FakeRunMode( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink_factory=lambda spec, scenario: _record_output_factory_call( + factory_calls, + spec, + scenario, + output, + ), + ) + spec = _spec() + scenario = _scenario() + + result = run_demo_session( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + run_mode=run_mode, + pipeline=StepPipeline(), + ) + + assert result.status == "completed" + assert adapter.provider_calls == [(spec, scenario)] + assert factory_calls == [(spec, scenario)] + assert run_metrics.sessions == [result] + assert len(run_metrics.sessions) == 1 + new_reservation = context.admission.try_reserve() + assert new_reservation is not None + new_reservation.release() + + +def test_busy_admission_returns_rejected_and_records_once() -> None: + runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) + admission = SingleSessionAdmissionPolicy() + held = admission.try_reserve() + assert held is not None + run_metrics = InMemorySessionMetricsRecorder() + context = _run_context(runtime, admission=admission, run_metrics=run_metrics) + adapter = _FakeDemoAdapter(provider=_FakeVideoModelInputProvider()) + + result = run_demo_session( + context=context, + spec=_spec(), + scenario=_scenario(), + adapter=adapter, + run_mode=_FakeRunMode(input_source=_FakeBatchInputSource(num_windows=1)), + pipeline=StepPipeline(), + ) + + held.release() + assert result == RunResult.rejected(reason="busy") + assert run_metrics.sessions == [result] + assert adapter.provider_calls == [] + assert runtime.start_session_inputs == [] + + +def test_setup_failure_returns_failed_before_runtime_session_creation() -> None: + runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) + provider = _FakeVideoModelInputProvider( + fail_initial=ValueError("invalid provider compatibility") + ) + metrics = InMemorySessionMetricsRecorder() + + result = BatchSessionDriver().run_one_session( + host=RuntimeHost(runtime), + provider=provider, + session_edges=SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=_RecordingOutputSink(), + metrics=metrics, + ), + pipeline=StepPipeline(), + ) + + assert result.status == "failed" + assert isinstance(result.error, ValueError) + assert result.reason == "invalid provider compatibility" + assert runtime.start_session_inputs == [] + assert provider.close_count == 1 + assert metrics.errors == ["invalid provider compatibility"] + + +def test_run_demo_session_closes_provider_when_validation_fails() -> None: + runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) + run_metrics = InMemorySessionMetricsRecorder() + context = _run_context(runtime, run_metrics=run_metrics) + provider = _FakeVideoModelInputProvider() + + result = run_demo_session( + context=context, + spec=_spec(), + scenario=_scenario(), + adapter=_FakeDemoAdapter(provider=provider), + run_mode=_FakeRunMode( + input_source=_FakeBatchInputSource(num_windows=1), + validate_error=ValueError("provider incompatible"), + ), + pipeline=StepPipeline(), + ) + + assert result.status == "failed" + assert result.reason == "provider incompatible" + assert provider.close_count == 1 + assert runtime.start_session_inputs == [] + assert run_metrics.sessions == [result] + + +def test_setup_failure_can_return_skipped_but_not_completed() -> None: + skipped = BatchSessionDriver().run_one_session( + host=RuntimeHost(_FakeVideoRuntime(session=_FakeVideoSession(num_steps=1))), + provider=_FakeVideoModelInputProvider(fail_initial=RuntimeError("skip me")), + session_edges=SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=_RecordingOutputSink(), + error_policy=_SetupPolicy(result_status="skipped"), + ), + pipeline=StepPipeline(), + ) + assert skipped.status == "skipped" + assert skipped.error is None + + with pytest.raises(DriverInvariantError, match="Setup failures"): + BatchSessionDriver().run_one_session( + host=RuntimeHost(_FakeVideoRuntime(session=_FakeVideoSession(num_steps=1))), + provider=_FakeVideoModelInputProvider( + fail_initial=RuntimeError("bad policy") + ), + session_edges=SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=_RecordingOutputSink(), + error_policy=_SetupPolicy(result_status="completed"), + ), + pipeline=StepPipeline(), + ) + + +def test_input_source_finished_error_returns_failed_not_completed() -> None: + metrics = InMemorySessionMetricsRecorder() + + result = BatchSessionDriver().run_one_session( + host=RuntimeHost(_FakeVideoRuntime(session=_FakeVideoSession(num_steps=1))), + provider=_FakeVideoModelInputProvider(), + session_edges=SessionEdges( + input_source=_FakeBatchInputSource( + num_windows=1, + fail_is_finished=RuntimeError("input source failed"), + ), + output_sink=_RecordingOutputSink(), + metrics=metrics, + ), + pipeline=StepPipeline(), + ) + + assert result.status == "failed" + assert result.reason == "input source failed" + assert metrics.errors == ["input source failed"] + + +def test_step_failure_returns_failed_from_driver() -> None: + session = _FakeVideoSession(num_steps=1, fail_step=0) + output = _RecordingOutputSink() + + result = BatchSessionDriver().run_one_session( + host=RuntimeHost(_FakeVideoRuntime(session=session)), + provider=_FakeVideoModelInputProvider(), + session_edges=SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=output, + ), + pipeline=StepPipeline(), + ) + + assert result.status == "failed" + assert isinstance(result.error, RuntimeError) + assert result.reason == "step failed" + assert output.results == [] + assert session.close_count == 1 + + +def test_session_edges_close_result_is_idempotent_and_first_result_wins() -> None: + output = _RecordingOutputSink( + artifacts=(OutputArtifact(kind="test/artifact", uri="memory://artifact"),) + ) + transport = _RecordingTransport() + metrics = InMemorySessionMetricsRecorder() + edges = SessionEdges( + input_source=_FakeBatchInputSource(num_windows=0), + output_sink=output, + metrics=metrics, + transport=transport, + ) + first_error = RuntimeError("first") + + first = edges.close_result( + status="failed", + reason="first", + error=first_error, + ) + second = edges.close_result(status="completed") + + assert second is first + assert first.status == "failed" + assert first.reason == "first" + assert first.error is first_error + assert tuple(first.artifacts) == ( + OutputArtifact(kind="test/artifact", uri="memory://artifact"), + ) + assert output.close_count == 1 + assert transport.close_count == 1 + assert metrics.closed + + +def test_run_result_rejected_is_the_only_convenience_constructor() -> None: + constructors = { + name + for name, value in RunResult.__dict__.items() + if isinstance(value, classmethod) + } + + assert constructors == {"rejected"} + assert RunResult.rejected(reason="busy").status == "rejected" + + +def _empty_step_outcome() -> Any: + from flashdreams.runtime.demo import StepOutcome + + return StepOutcome(output=OutputDecision(), control=ControlDecision()) + + +def _window(index: int) -> UserInputWindow: + start_s = float(index) + return UserInputWindow( + start_s=start_s, + end_s=start_s + 1.0, + frame_times=(start_s + 1.0,), + inputs=UserInputs(), + ) + + +def _spec() -> DemoSpec: + return DemoSpec( + model_id="fake-video-demo", + input_mode="replay", + output=NullOutputSpec(), + config=InferenceConfig(model_id="fake-video-demo"), + ) + + +def _scenario() -> PreparedScenario: + return PreparedScenario(initial_inputs=InferenceInput()) + + +def _run_context( + runtime: _FakeVideoRuntime, + *, + admission: SingleSessionAdmissionPolicy | None = None, + run_metrics: InMemorySessionMetricsRecorder | None = None, +) -> RunContext: + host = RuntimeHost(runtime) + return RunContext( + host=host, + run_metrics=run_metrics or InMemorySessionMetricsRecorder(), + admission=admission or SingleSessionAdmissionPolicy( + health_check=lambda: host.is_healthy + ), + ) + + +def _record_output_factory_call( + calls: list[tuple[DemoSpec, PreparedScenario]], + spec: DemoSpec, + scenario: PreparedScenario, + output: "_RecordingOutputSink", +) -> "_RecordingOutputSink": + calls.append((spec, scenario)) + return output + + +class _FakeVideoModelInputProvider: + def __init__(self, *, fail_initial: Exception | None = None) -> None: + self.fail_initial = fail_initial + self.initial_input = InferenceInput( + global_conditioning={"prompt": "fake video prompt"} + ) + self.prepared_step_inputs: list[InferenceInput] = [] + self.reset_inputs: list[InferenceInput | None] = [] + self.close_count = 0 + + def prepare_initial_input(self) -> InferenceInput: + if self.fail_initial is not None: + raise self.fail_initial + return self.initial_input + + def prepare_step( + self, + *, + request: StepRequest, + user_window: UserInputWindow, + ) -> PreparedStep: + inference_input = InferenceInput( + step={ + "request_step": request.step_index, + "window": (user_window.start_s, user_window.end_s), + } + ) + self.prepared_step_inputs.append(inference_input) + return PreparedStep(inference_input=inference_input) + + def reset(self, inputs: InferenceInput | None = None) -> None: + self.reset_inputs.append(inputs) + + def close(self) -> None: + self.close_count += 1 + + +class _FakeBatchInputSource: + is_finite = True + is_deterministic = True + + def __init__( + self, + *, + num_windows: int, + fail_is_finished: Exception | None = None, + ) -> None: + self.windows = [_window(index) for index in range(num_windows)] + self.fail_is_finished = fail_is_finished + self.next_window_requests: list[StepRequest] = [] + self.index = 0 + + def is_finished(self) -> bool: + if self.fail_is_finished is not None: + raise self.fail_is_finished + return self.index >= len(self.windows) + + def next_window(self, request: StepRequest) -> UserInputWindow: + self.next_window_requests.append(request) + window = self.windows[self.index] + self.index += 1 + return window + + +class _FakeVideoRuntime: + def __init__(self, *, session: "_FakeVideoSession") -> None: + self.session = session + self.start_session_inputs: list[InferenceInput] = [] + self.close_count = 0 + + def start_session(self, inputs: InferenceInput) -> InferenceSession: + self.start_session_inputs.append(inputs) + return self.session + + def close(self) -> None: + self.close_count += 1 + + +class _FakeVideoSession: + def __init__( + self, + *, + num_steps: int, + fail_step: int | None = None, + ) -> None: + self.num_steps = num_steps + self.fail_step = fail_step + self.next_request_index = 0 + self.step_inputs: list[InferenceInput] = [] + self.close_count = 0 + + def session_info(self) -> SessionInfo: + return SessionInfo(output_layout="fake-video", steady_output_frame_count=1) + + def next_step_request(self) -> StepRequest | None: + if self.next_request_index >= self.num_steps: + return None + request = StepRequest(step_index=self.next_request_index) + self.next_request_index += 1 + return request + + def step(self, inputs: InferenceInput) -> StepResult: + step_index = len(self.step_inputs) + if self.fail_step == step_index: + raise RuntimeError("step failed") + self.step_inputs.append(inputs) + return StepResult( + step_index=step_index, + output=f"frame-{step_index}", + frame_count=1, + metrics={"model_step_s": 0.01}, + ) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self.next_request_index = 0 + self.step_inputs.clear() + + def close(self) -> None: + self.close_count += 1 + + +class _RecordingRuntimeHost(RuntimeHost): + def __init__(self, runtime: _FakeVideoRuntime) -> None: + super().__init__(runtime) + self.calls: list[str] = [] + + def call(self, func: Callable[..., Any], /, *args: object, **kwargs: object) -> Any: + self.calls.append(getattr(func, "__name__", type(func).__name__)) + return super().call(func, *args, **kwargs) + + +class _RecordingOutputSink: + produces_artifacts = True + + def __init__( + self, + *, + artifacts: Sequence[OutputArtifact] = (), + decision: OutputDecision | None = None, + ) -> None: + self.artifacts = tuple(artifacts) + self.decision = decision or OutputDecision() + self.opened_with: SessionInfo | None = None + self.results: list[StepResult] = [] + self.close_count = 0 + + def open(self, session_info: SessionInfo) -> None: + self.opened_with = session_info + + def begin_generation(self, generation: int) -> None: + del generation + + def write(self, result: StepResult) -> OutputDecision: + self.results.append(result) + return self.decision + + def close(self) -> Sequence[OutputArtifact]: + self.close_count += 1 + return self.artifacts + + +class _RecordingTransport: + def __init__(self) -> None: + self.close_count = 0 + + def is_active(self) -> bool: + return self.close_count == 0 + + def close(self) -> None: + self.close_count += 1 + + +class _SetupPolicy: + def __init__( + self, + *, + result_status: Literal["completed", "failed", "skipped"], + ) -> None: + self.result_status = result_status + + def handle_setup_error(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status=self.result_status) + + def handle(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status="failed") + + +class _FakeDemoAdapter: + model_id = "fake-video-demo" + inference_input_schema = InferenceInputSchema() + canonical_input_schema = CanonicalInputSchema() + + def __init__(self, *, provider: _FakeVideoModelInputProvider) -> None: + self.provider = provider + self.provider_calls: list[tuple[DemoSpec, PreparedScenario]] = [] + + def supported_input_modes(self) -> tuple[str, ...]: + return ("replay",) + + def supported_output_modes(self) -> tuple[str, ...]: + return ("null",) + + def default_input_mapping(self) -> InputMapping: + return IdentityInputMapping() + + def validate_config(self, config: InferenceConfig) -> None: + if config.model_id != self.model_id: + raise ValueError(f"Unsupported model_id={config.model_id!r}.") + + def create_runtime(self, config: InferenceConfig) -> InferenceRuntime: + del config + raise NotImplementedError("FakeVideoDemo uses an explicit RuntimeHost.") + + def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: + del spec + return _scenario() + + def create_model_input_provider( + self, + spec: DemoSpec, + scenario: PreparedScenario, + ) -> _FakeVideoModelInputProvider: + self.provider_calls.append((spec, scenario)) + return self.provider + + +class _FakeRunMode: + def __init__( + self, + *, + input_source: _FakeBatchInputSource, + output_sink: _RecordingOutputSink | None = None, + output_sink_factory: ( + Callable[[DemoSpec, PreparedScenario], _RecordingOutputSink] | None + ) = None, + metrics: InMemorySessionMetricsRecorder | None = None, + validate_error: Exception | None = None, + ) -> None: + self.input_source = input_source + self.output_sink = output_sink or _RecordingOutputSink() + self.output_sink_factory = output_sink_factory + self.metrics = metrics or InMemorySessionMetricsRecorder() + self.validate_error = validate_error + + def validate_session( + self, + *, + spec: DemoSpec, + scenario: PreparedScenario, + adapter: Any, + provider: Any, + ) -> None: + del spec, scenario, adapter, provider + if self.validate_error is not None: + raise self.validate_error + + def create_session_edges( + self, + *, + context: RunContext, + spec: DemoSpec, + scenario: PreparedScenario, + provider: Any, + adapter: Any, + ) -> SessionEdges: + del context, provider, adapter + output_sink = ( + self.output_sink_factory(spec, scenario) + if self.output_sink_factory is not None + else self.output_sink + ) + return SessionEdges( + input_source=self.input_source, + output_sink=output_sink, + metrics=self.metrics, + ) + + def select_driver(self) -> BatchSessionDriver: + return BatchSessionDriver() From 42c7859e3f5912946c241cdf510c28cdae3e0579 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 00:40:34 +0000 Subject: [PATCH 02/51] Fix demo session invariant cleanup Close session edges and record a failed run result when DriverInvariantError escapes after edge creation, and apply the ruff formatting fixes required by CPU CI. --- .../flashdreams/runtime/demo/__init__.py | 2 +- .../flashdreams/runtime/demo/drivers.py | 17 ++++++- .../flashdreams/runtime/demo/outputs.py | 3 +- .../flashdreams/runtime/demo/run_modes.py | 6 ++- .../runtime/demo/session_inputs.py | 8 ++-- .../tests/test_demo_runtime_vertical_slice.py | 47 +++++++++++++++++-- 6 files changed, 70 insertions(+), 13 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index 04883d8bf..ce933720c 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -9,12 +9,12 @@ run_demo_session, ) from flashdreams.runtime.demo.host import RuntimeHost -from flashdreams.runtime.demo.outputs import build_output_target from flashdreams.runtime.demo.outputs import ( NullOutputSink, OutputDecision, OutputSink, SessionInfo, + build_output_target, ) from flashdreams.runtime.demo.pipeline import StepOutcome, StepPipeline from flashdreams.runtime.demo.replay import run_replay_demo diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index e6050f6ec..61b55220e 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -170,7 +170,22 @@ def run_demo_session( ) context.run_metrics.record_session(result) return result - except DriverInvariantError: + except DriverInvariantError as exc: + if provider is not None and not driver_started: + try: + context.host.call(provider.close) + except Exception as close_exc: + if session_edges is not None: + session_edges.metrics.record_cleanup_error(close_exc) + else: + context.run_metrics.record_cleanup_error(close_exc) + if session_edges is not None: + result = session_edges.close_result( + status="failed", + reason=str(exc), + error=exc, + ) + context.run_metrics.record_session(result) raise except Exception as exc: if provider is not None and not driver_started: diff --git a/flashdreams/flashdreams/runtime/demo/outputs.py b/flashdreams/flashdreams/runtime/demo/outputs.py index be4fd2bda..d24f99da7 100644 --- a/flashdreams/flashdreams/runtime/demo/outputs.py +++ b/flashdreams/flashdreams/runtime/demo/outputs.py @@ -11,8 +11,7 @@ from typing import Literal, Protocol, runtime_checkable from flashdreams.runtime._utils import freeze_mapping -from flashdreams.runtime.output import OutputArtifact -from flashdreams.runtime.output import NullOutputTarget, OutputTarget +from flashdreams.runtime.output import NullOutputTarget, OutputArtifact, OutputTarget from flashdreams.runtime.types import StepResult from flashdreams.runtime.video_output import Mp4VideoOutputTarget, VideoWriter diff --git a/flashdreams/flashdreams/runtime/demo/run_modes.py b/flashdreams/flashdreams/runtime/demo/run_modes.py index 335a7fb96..973749d9d 100644 --- a/flashdreams/flashdreams/runtime/demo/run_modes.py +++ b/flashdreams/flashdreams/runtime/demo/run_modes.py @@ -8,7 +8,7 @@ from collections.abc import Mapping, Sequence from dataclasses import dataclass, field from threading import Lock -from typing import Any, Literal, Protocol, TYPE_CHECKING, runtime_checkable +from typing import TYPE_CHECKING, Any, Literal, Protocol, runtime_checkable from flashdreams.runtime._utils import freeze_mapping from flashdreams.runtime.output import OutputArtifact @@ -52,7 +52,9 @@ def __post_init__(self) -> None: object.__setattr__( self, "timings", - freeze_mapping({key: tuple(values) for key, values in self.timings.items()}), + freeze_mapping( + {key: tuple(values) for key, values in self.timings.items()} + ), ) object.__setattr__(self, "session_statuses", tuple(self.session_statuses)) object.__setattr__(self, "errors", tuple(self.errors)) diff --git a/flashdreams/flashdreams/runtime/demo/session_inputs.py b/flashdreams/flashdreams/runtime/demo/session_inputs.py index cf6838961..47bbbad6e 100644 --- a/flashdreams/flashdreams/runtime/demo/session_inputs.py +++ b/flashdreams/flashdreams/runtime/demo/session_inputs.py @@ -47,9 +47,7 @@ def __post_init__(self) -> None: if not math.isfinite(self.start_s) or self.start_s < 0: raise ValueError("UserInputWindow.start_s must be finite and >= 0.") if not math.isfinite(self.end_s) or self.end_s < self.start_s: - raise ValueError( - "UserInputWindow.end_s must be finite and >= start_s." - ) + raise ValueError("UserInputWindow.end_s must be finite and >= start_s.") previous = -math.inf for frame_time in self.frame_times: if not math.isfinite(float(frame_time)): @@ -59,7 +57,9 @@ def __post_init__(self) -> None: "UserInputWindow.frame_times must be sorted in ascending order." ) previous = float(frame_time) - object.__setattr__(self, "frame_times", tuple(float(t) for t in self.frame_times)) + object.__setattr__( + self, "frame_times", tuple(float(t) for t in self.frame_times) + ) object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index e7a3dc303..88a646078 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -256,6 +256,42 @@ def test_setup_failure_can_return_skipped_but_not_completed() -> None: ) +def test_run_demo_session_closes_edges_when_driver_invariant_escapes() -> None: + runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) + run_metrics = InMemorySessionMetricsRecorder() + context = _run_context(runtime, run_metrics=run_metrics) + provider = _FakeVideoModelInputProvider( + fail_initial=RuntimeError("bad setup policy") + ) + output = _RecordingOutputSink() + transport = _RecordingTransport() + session_metrics = InMemorySessionMetricsRecorder() + + with pytest.raises(DriverInvariantError, match="Setup failures"): + run_demo_session( + context=context, + spec=_spec(), + scenario=_scenario(), + adapter=_FakeDemoAdapter(provider=provider), + run_mode=_FakeRunMode( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=output, + metrics=session_metrics, + transport=transport, + error_policy=_SetupPolicy(result_status="completed"), + ), + pipeline=StepPipeline(), + ) + + assert output.close_count == 1 + assert transport.close_count == 1 + assert session_metrics.closed + assert provider.close_count == 1 + assert len(run_metrics.sessions) == 1 + assert run_metrics.sessions[0].status == "failed" + assert isinstance(run_metrics.sessions[0].error, DriverInvariantError) + + def test_input_source_finished_error_returns_failed_not_completed() -> None: metrics = InMemorySessionMetricsRecorder() @@ -382,9 +418,8 @@ def _run_context( return RunContext( host=host, run_metrics=run_metrics or InMemorySessionMetricsRecorder(), - admission=admission or SingleSessionAdmissionPolicy( - health_check=lambda: host.is_healthy - ), + admission=admission + or SingleSessionAdmissionPolicy(health_check=lambda: host.is_healthy), ) @@ -637,12 +672,16 @@ def __init__( Callable[[DemoSpec, PreparedScenario], _RecordingOutputSink] | None ) = None, metrics: InMemorySessionMetricsRecorder | None = None, + transport: _RecordingTransport | None = None, + error_policy: _SetupPolicy | None = None, validate_error: Exception | None = None, ) -> None: self.input_source = input_source self.output_sink = output_sink or _RecordingOutputSink() self.output_sink_factory = output_sink_factory self.metrics = metrics or InMemorySessionMetricsRecorder() + self.transport = transport + self.error_policy = error_policy self.validate_error = validate_error def validate_session( @@ -676,6 +715,8 @@ def create_session_edges( input_source=self.input_source, output_sink=output_sink, metrics=self.metrics, + error_policy=self.error_policy or _SetupPolicy(result_status="failed"), + transport=self.transport or _RecordingTransport(), ) def select_driver(self) -> BatchSessionDriver: From 539911e59398f3f51361f67cb91f4b3f9f296195 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 00:58:25 +0000 Subject: [PATCH 03/51] Finalize direct demo driver invariant failures --- .../flashdreams/runtime/demo/drivers.py | 19 ++++++++++--- .../tests/test_demo_runtime_vertical_slice.py | 27 +++++++++++++------ 2 files changed, 34 insertions(+), 12 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 61b55220e..7bee95a96 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -43,6 +43,7 @@ def run_one_session( final_status: DriverStatus = "completed" final_reason: str | None = None final_error: Exception | None = None + invariant_closed = False setup_ok = False try: try: @@ -101,16 +102,26 @@ def run_one_session( final_reason = str(exc) final_error = exc if action.result_status == "failed" else None break - except DriverInvariantError: + except DriverInvariantError as exc: + if session is not None: + host.call(_close_safely, session.close, session_edges) + host.call(_close_safely, provider.close, session_edges) + session_edges.close_result( + status="failed", + reason=str(exc), + error=exc, + ) + invariant_closed = True raise except Exception as exc: final_status = "failed" final_reason = str(exc) final_error = exc finally: - if session is not None: - host.call(_close_safely, session.close, session_edges) - host.call(_close_safely, provider.close, session_edges) + if not invariant_closed: + if session is not None: + host.call(_close_safely, session.close, session_edges) + host.call(_close_safely, provider.close, session_edges) return session_edges.close_result( status=final_status, diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index 88a646078..8e19a489b 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -241,20 +241,31 @@ def test_setup_failure_can_return_skipped_but_not_completed() -> None: assert skipped.status == "skipped" assert skipped.error is None + provider = _FakeVideoModelInputProvider(fail_initial=RuntimeError("bad policy")) + output = _RecordingOutputSink() + transport = _RecordingTransport() + metrics = InMemorySessionMetricsRecorder() + edges = SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=output, + metrics=metrics, + error_policy=_SetupPolicy(result_status="completed"), + transport=transport, + ) + with pytest.raises(DriverInvariantError, match="Setup failures"): BatchSessionDriver().run_one_session( host=RuntimeHost(_FakeVideoRuntime(session=_FakeVideoSession(num_steps=1))), - provider=_FakeVideoModelInputProvider( - fail_initial=RuntimeError("bad policy") - ), - session_edges=SessionEdges( - input_source=_FakeBatchInputSource(num_windows=1), - output_sink=_RecordingOutputSink(), - error_policy=_SetupPolicy(result_status="completed"), - ), + provider=provider, + session_edges=edges, pipeline=StepPipeline(), ) + assert output.close_count == 1 + assert transport.close_count == 1 + assert metrics.closed + assert provider.close_count == 1 + def test_run_demo_session_closes_edges_when_driver_invariant_escapes() -> None: runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) From d06df9f82fd4714701791b60f24a69393f5a6e3b Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 01:25:20 +0000 Subject: [PATCH 04/51] Add RuntimeHost worker boundary --- flashdreams/flashdreams/runtime/__init__.py | 3 +- .../flashdreams/runtime/demo/__init__.py | 8 +- flashdreams/flashdreams/runtime/demo/host.py | 150 ++++++++++-- flashdreams/flashdreams/runtime/worker.py | 56 ++++- flashdreams/tests/test_demo_runtime_host.py | 213 ++++++++++++++++++ .../tests/test_demo_runtime_vertical_slice.py | 2 + flashdreams/tests/test_runtime_worker.py | 35 ++- 7 files changed, 432 insertions(+), 35 deletions(-) create mode 100644 flashdreams/tests/test_demo_runtime_host.py diff --git a/flashdreams/flashdreams/runtime/__init__.py b/flashdreams/flashdreams/runtime/__init__.py index fb6eb4b05..e3e574594 100644 --- a/flashdreams/flashdreams/runtime/__init__.py +++ b/flashdreams/flashdreams/runtime/__init__.py @@ -59,7 +59,7 @@ from flashdreams.runtime.runner import run_inference_session from flashdreams.runtime.types import StepRequest, StepResult from flashdreams.runtime.video_output import Mp4VideoOutputTarget -from flashdreams.runtime.worker import ThreadAffineRuntimeWorker +from flashdreams.runtime.worker import ModelExecutionWorker, ThreadAffineRuntimeWorker __all__ = [ "CanonicalInputs", @@ -91,6 +91,7 @@ "MappingCompatibility", "MetricsRecorder", "ModelAdapter", + "ModelExecutionWorker", "Mp4VideoOutputTarget", "NullMetricsRecorder", "NullOutputTarget", diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index ce933720c..c324bc6b2 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -8,7 +8,11 @@ DriverInvariantError, run_demo_session, ) -from flashdreams.runtime.demo.host import RuntimeHost +from flashdreams.runtime.demo.host import ( + ModelWarmupPlan, + RuntimeHost, + WarmupSessionInputs, +) from flashdreams.runtime.demo.outputs import ( NullOutputSink, OutputDecision, @@ -63,6 +67,7 @@ "InMemorySessionMetricsRecorder", "InputSource", "MetricsSnapshot", + "ModelWarmupPlan", "ModelInputProvider", "Mp4OutputSpec", "NoopTransportService", @@ -85,6 +90,7 @@ "StepOutcome", "StepPipeline", "UserInputWindow", + "WarmupSessionInputs", "WebRTCAppResources", "WebRTCOutputSpec", "build_output_target", diff --git a/flashdreams/flashdreams/runtime/demo/host.py b/flashdreams/flashdreams/runtime/demo/host.py index ee870ebc5..6a6e93ce5 100644 --- a/flashdreams/flashdreams/runtime/demo/host.py +++ b/flashdreams/flashdreams/runtime/demo/host.py @@ -1,48 +1,128 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Minimal runtime host for the Phase 2 demo-session vertical slice.""" +"""Runtime host and model-execution boundary for shared demos.""" from __future__ import annotations -from collections.abc import Callable +from collections.abc import Callable, Mapping, Sequence +from dataclasses import dataclass, field from typing import TypeVar +from flashdreams.runtime._utils import freeze_mapping from flashdreams.runtime.inputs import InferenceInput from flashdreams.runtime.interfaces import InferenceRuntime, InferenceSession +from flashdreams.runtime.worker import ModelExecutionWorker _T = TypeVar("_T") -class RuntimeHost: - """Thin synchronous host around an :class:`InferenceRuntime`. +@dataclass(frozen=True, kw_only=True, slots=True) +class WarmupSessionInputs: + """Inputs used to warm one temporary runtime session.""" + + initial_input: InferenceInput + step_inputs: Sequence[InferenceInput] = () + + def __post_init__(self) -> None: + object.__setattr__(self, "step_inputs", tuple(self.step_inputs)) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class ModelWarmupPlan: + """Host-owned model warmup plan built by a demo adapter or run mode.""" + + sessions: Sequence[WarmupSessionInputs] = () + measured: bool = False + metadata: Mapping[str, object] = field(default_factory=dict) - Phase 3 moves the thread-affine worker boundary here. Phase 2 keeps the - dispatch direct so fake-model CPU tests can prove the session-driver shape - without introducing worker behavior early. - """ + def __post_init__(self) -> None: + object.__setattr__(self, "sessions", tuple(self.sessions)) + object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) - def __init__(self, runtime: InferenceRuntime) -> None: + +class RuntimeHost: + """Own one runtime and the worker used for model-affine calls.""" + + def __init__( + self, + runtime: InferenceRuntime, + *, + worker: ModelExecutionWorker | None = None, + is_control_rank: bool = True, + worker_loop: Callable[[], None] | None = None, + ) -> None: self._runtime = runtime + self._worker = worker or ModelExecutionWorker() + self._is_control_rank = is_control_rank + self._worker_loop = worker_loop self._healthy = True + self._closed = False + self._unhealthy_reason: str | None = None + self._unhealthy_error: Exception | None = None @property def runtime(self) -> InferenceRuntime: """Return the hosted runtime.""" return self._runtime + @property + def worker(self) -> ModelExecutionWorker: + """Return the host's model-execution worker.""" + return self._worker + + @property + def is_control_rank(self) -> bool: + """Whether this process owns run modes, providers, sinks, and metrics.""" + return self._is_control_rank + @property def is_healthy(self) -> bool: """Return whether admission should continue accepting sessions.""" - return self._healthy + return self._healthy and not self._closed - def mark_unhealthy(self) -> None: - """Latch the host as unhealthy.""" + @property + def unhealthy_reason(self) -> str | None: + """Return the first latched unhealthy reason, if any.""" + return self._unhealthy_reason + + @property + def unhealthy_error(self) -> Exception | None: + """Return the first latched unhealthy error, if any.""" + return self._unhealthy_error + + def mark_unhealthy( + self, + reason: str = "marked unhealthy", + error: Exception | None = None, + ) -> None: + """Latch the host as unhealthy without overwriting the first reason.""" + if not self._healthy: + return self._healthy = False + self._unhealthy_reason = reason + self._unhealthy_error = error + + def preload(self) -> None: + """Initialize optional distributed state and preload runtime resources.""" + self._call_optional_runtime_hook("initialize_distributed") + self._call_optional_runtime_hook("preload") + + def warmup(self, plan: ModelWarmupPlan | None = None) -> None: + """Run warmup sessions through the same worker boundary as real sessions.""" + plan = plan or ModelWarmupPlan() + for warmup_session in plan.sessions: + session = self.call(self.start_session, warmup_session.initial_input) + try: + for step_input in warmup_session.step_inputs: + self.call(session.step, step_input) + finally: + self.call(session.close) def call(self, func: Callable[..., _T], /, *args: object, **kwargs: object) -> _T: - """Run one model-affine callable synchronously.""" - return func(*args, **kwargs) + """Run one model-affine callable synchronously on the worker.""" + self._require_open() + return self._worker.call_blocking(func, *args, **kwargs) async def call_async( self, @@ -51,16 +131,44 @@ async def call_async( *args: object, **kwargs: object, ) -> _T: - """Async-compatible direct dispatch placeholder for Phase 3.""" - return self.call(func, *args, **kwargs) + """Run model-affine work without blocking realtime event loops.""" + self._require_open() + return await self._worker.call(func, *args, **kwargs) def start_session(self, inputs: InferenceInput) -> InferenceSession: """Start one inference session through the hosted runtime.""" + self._require_open() return self._runtime.start_session(inputs) - def close(self) -> None: - """Close the hosted runtime.""" - self._runtime.close() + def run_worker_loop(self) -> None: + """Serve control-rank work on non-control ranks until runtime shutdown.""" + worker_loop = self._worker_loop + if worker_loop is None: + worker_loop = getattr(self._runtime, "run_worker_loop", None) + if worker_loop is None: + worker_loop = getattr(self._runtime, "wait_for_termination", None) + if callable(worker_loop): + worker_loop() - -__all__ = ["RuntimeHost"] + def close(self) -> None: + """Close runtime-owned state and stop the model-execution worker.""" + if self._closed: + return + try: + self._worker.call_blocking(self._runtime.close) + self._call_optional_runtime_hook("close_distributed") + finally: + self._closed = True + self._worker.close_blocking() + + def _call_optional_runtime_hook(self, name: str) -> None: + hook = getattr(self._runtime, name, None) + if callable(hook): + self.call(hook) + + def _require_open(self) -> None: + if self._closed: + raise RuntimeError("runtime host is closed") + + +__all__ = ["ModelWarmupPlan", "RuntimeHost", "WarmupSessionInputs"] diff --git a/flashdreams/flashdreams/runtime/worker.py b/flashdreams/flashdreams/runtime/worker.py index af4576c09..a545054b0 100644 --- a/flashdreams/flashdreams/runtime/worker.py +++ b/flashdreams/flashdreams/runtime/worker.py @@ -6,6 +6,7 @@ from __future__ import annotations import asyncio +import threading from concurrent.futures import ThreadPoolExecutor from typing import Any, Callable, TypeVar, cast @@ -14,7 +15,7 @@ _T = TypeVar("_T") -class ThreadAffineRuntimeWorker: +class ModelExecutionWorker: """Run ordered runtime lifecycle calls on one owned OS thread. CUDA graphs, Triton launchers, and some backend contexts are thread-local. @@ -38,14 +39,23 @@ def __init__( thread_name_prefix=thread_name, initializer=self._initialize_thread, ) + self._state_lock = threading.Lock() self._accepting = True self._closed = False - self._close_lock = asyncio.Lock() + self._thread_id: int | None = None @property def closed(self) -> bool: return self._closed + @property + def worker_thread_id(self) -> int | None: + return self._thread_id + + @property + def is_worker_thread(self) -> bool: + return self._thread_id == threading.get_ident() + async def call( self, func: Callable[..., _T], @@ -54,8 +64,8 @@ async def call( **kwargs: Any, ) -> _T: """Run one callable after all previously submitted worker calls.""" - if not self._accepting: - raise RuntimeError("runtime worker is closed") + self._require_not_worker_thread() + self._require_accepting() future = self._submit(func, args, kwargs) try: return await asyncio.shield(future) @@ -71,21 +81,30 @@ def call_blocking( **kwargs: Any, ) -> _T: """Run one callable from synchronous code on the owned worker thread.""" - if not self._accepting: - raise RuntimeError("runtime worker is closed") + self._require_not_worker_thread() + self._require_accepting() future = self._executor.submit(_invoke, func, args, kwargs) return cast(_T, future.result()) async def close(self) -> None: """Drain submitted work and stop accepting lifecycle calls.""" - async with self._close_lock: + self._require_not_worker_thread() + await asyncio.to_thread(self.close_blocking) + + def close_blocking(self) -> None: + """Synchronous close for non-async setup and teardown paths.""" + self._require_not_worker_thread() + with self._state_lock: if self._closed: return self._accepting = False - barrier = self._submit(_noop, (), {}) - await asyncio.shield(barrier) + try: + barrier = self._executor.submit(_noop) + barrier.result() + finally: self._executor.shutdown(wait=True, cancel_futures=False) - self._closed = True + with self._state_lock: + self._closed = True def _submit( self, @@ -97,9 +116,24 @@ def _submit( return loop.run_in_executor(self._executor, _invoke, func, args, kwargs) def _initialize_thread(self) -> None: + self._thread_id = threading.get_ident() if self._device is not None and self._device.type == "cuda": torch.cuda.set_device(self._device) + def _require_accepting(self) -> None: + if not self._accepting: + raise RuntimeError("runtime worker is closed") + + def _require_not_worker_thread(self) -> None: + if self.is_worker_thread: + raise RuntimeError( + "Cannot dispatch to the model execution worker from its own " + "thread; call the function directly." + ) + + +ThreadAffineRuntimeWorker = ModelExecutionWorker + def _invoke( func: Callable[..., _T], @@ -118,4 +152,4 @@ def _consume_exception(future: asyncio.Future[Any]) -> None: future.exception() -__all__ = ["ThreadAffineRuntimeWorker"] +__all__ = ["ModelExecutionWorker", "ThreadAffineRuntimeWorker"] diff --git a/flashdreams/tests/test_demo_runtime_host.py b/flashdreams/tests/test_demo_runtime_host.py new file mode 100644 index 000000000..eed7635b0 --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_host.py @@ -0,0 +1,213 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +import threading + +import pytest + +from flashdreams.runtime import InferenceInput, InferenceSession, StepRequest, StepResult +from flashdreams.runtime.demo import ModelWarmupPlan, RuntimeHost, WarmupSessionInputs + +pytestmark = pytest.mark.ci_cpu + + +def test_runtime_host_latches_health_and_runs_lifecycle_in_order() -> None: + runtime = _LifecycleRuntime() + host = RuntimeHost(runtime) + error = RuntimeError("runtime wedged") + + host.mark_unhealthy("first failure", error) + host.mark_unhealthy("second failure") + + assert not host.is_healthy + assert host.unhealthy_reason == "first failure" + assert host.unhealthy_error is error + + initial_input = InferenceInput(global_conditioning={"session": "warmup"}) + step_inputs = ( + InferenceInput(step={"step": 0}), + InferenceInput(step={"step": 1}), + ) + host.preload() + host.warmup( + ModelWarmupPlan( + sessions=( + WarmupSessionInputs( + initial_input=initial_input, + step_inputs=step_inputs, + ), + ), + ) + ) + host.close() + host.close() + + assert runtime.events == [ + "initialize_distributed", + "preload", + ("start_session", initial_input), + ("step", step_inputs[0]), + ("step", step_inputs[1]), + "session.close", + "runtime.close", + "close_distributed", + ] + assert not host.is_healthy + with pytest.raises(RuntimeError, match="closed"): + host.call(lambda: None) + + +@pytest.mark.asyncio +async def test_runtime_host_call_async_does_not_block_event_loop() -> None: + runtime = _LifecycleRuntime() + host = RuntimeHost(runtime) + loop_thread_id = threading.get_ident() + started = threading.Event() + release = threading.Event() + heartbeat_ticks = 0 + + def _slow_model_call() -> int: + started.set() + assert release.wait(timeout=2.0) + return threading.get_ident() + + async def _heartbeat_until_done(task: asyncio.Task[int]) -> None: + nonlocal heartbeat_ticks + while not task.done(): + heartbeat_ticks += 1 + await asyncio.sleep(0) + + try: + model_task = asyncio.create_task(host.call_async(_slow_model_call)) + assert await asyncio.to_thread(started.wait, 1.0) + heartbeat_task = asyncio.create_task(_heartbeat_until_done(model_task)) + for _ in range(5): + await asyncio.sleep(0) + assert heartbeat_ticks > 0 + + release.set() + worker_thread_id = await model_task + await heartbeat_task + finally: + host.close() + + assert worker_thread_id != loop_thread_id + + +def test_runtime_host_reentrant_sync_and_async_dispatch_raise() -> None: + host = RuntimeHost(_LifecycleRuntime()) + + def _nested_sync_dispatch() -> None: + host.call(lambda: None) + + def _nested_async_dispatch() -> None: + async def _dispatch() -> None: + await host.call_async(lambda: None) + + asyncio.run(_dispatch()) + + try: + with pytest.raises(RuntimeError, match="own thread"): + host.call(_nested_sync_dispatch) + with pytest.raises(RuntimeError, match="own thread"): + host.call(_nested_async_dispatch) + finally: + host.close() + + +def test_non_control_rank_setup_returns_after_worker_loop_without_demo_edges() -> None: + runtime = _LifecycleRuntime() + host = RuntimeHost( + runtime, + is_control_rank=False, + worker_loop=runtime.run_worker_loop, + ) + constructed: list[str] = [] + + result = _fake_run_setup(host, constructed) + + assert result == "worker-rank" + assert constructed == [] + assert runtime.events == [ + "initialize_distributed", + "preload", + "run_worker_loop", + "runtime.close", + "close_distributed", + ] + + +def test_control_rank_setup_reaches_demo_assembly() -> None: + runtime = _LifecycleRuntime() + host = RuntimeHost(runtime) + constructed: list[str] = [] + + try: + result = _fake_run_setup(host, constructed) + finally: + host.close() + + assert result == "control-rank" + assert constructed == ["run_mode", "provider", "input_source", "output_sink"] + assert runtime.events[:2] == ["initialize_distributed", "preload"] + assert "run_worker_loop" not in runtime.events + + +def _fake_run_setup(host: RuntimeHost, constructed: list[str]) -> str: + host.preload() + if not host.is_control_rank: + host.run_worker_loop() + host.close() + return "worker-rank" + + constructed.extend(["run_mode", "provider", "input_source", "output_sink"]) + return "control-rank" + + +class _LifecycleRuntime: + def __init__(self) -> None: + self.events: list[object] = [] + + def initialize_distributed(self) -> None: + self.events.append("initialize_distributed") + + def preload(self) -> None: + self.events.append("preload") + + def start_session(self, inputs: InferenceInput) -> InferenceSession: + self.events.append(("start_session", inputs)) + return _LifecycleSession(self.events) + + def run_worker_loop(self) -> None: + self.events.append("run_worker_loop") + + def close(self) -> None: + self.events.append("runtime.close") + + def close_distributed(self) -> None: + self.events.append("close_distributed") + + +class _LifecycleSession: + def __init__(self, events: list[object]) -> None: + self._events = events + self._next_step = 0 + + def next_step_request(self) -> StepRequest | None: + request = StepRequest(step_index=self._next_step) + self._next_step += 1 + return request + + def step(self, inputs: InferenceInput) -> StepResult: + self._events.append(("step", inputs)) + return StepResult(step_index=self._next_step, output=None) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self._next_step = 0 + + def close(self) -> None: + self._events.append("session.close") diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index 8e19a489b..e550e1d26 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -109,6 +109,8 @@ def test_batch_driver_runs_fake_video_demo_through_runtime_host() -> None: assert host.calls.count("execute_step") == 2 assert "prepare_initial_input" in host.calls assert "start_session" in host.calls + assert "prepare_step" not in host.calls + assert "step" not in host.calls def test_run_demo_session_builds_edges_and_records_session_once() -> None: diff --git a/flashdreams/tests/test_runtime_worker.py b/flashdreams/tests/test_runtime_worker.py index f6fbf84aa..d76558005 100644 --- a/flashdreams/tests/test_runtime_worker.py +++ b/flashdreams/tests/test_runtime_worker.py @@ -8,11 +8,15 @@ import pytest -from flashdreams.runtime import ThreadAffineRuntimeWorker +from flashdreams.runtime import ModelExecutionWorker, ThreadAffineRuntimeWorker pytestmark = pytest.mark.ci_cpu +def test_model_execution_worker_keeps_legacy_worker_alias() -> None: + assert ThreadAffineRuntimeWorker is ModelExecutionWorker + + @pytest.mark.asyncio async def test_worker_preserves_order_and_thread_affinity() -> None: worker = ThreadAffineRuntimeWorker(thread_name="test-runtime") @@ -95,3 +99,32 @@ async def test_worker_sets_cuda_device_when_thread_starts( await worker.close() assert [str(device) for device in seen] == ["cuda:3"] + + +def test_blocking_worker_call_is_not_reentrant() -> None: + worker = ModelExecutionWorker() + + def _nested_dispatch() -> None: + worker.call_blocking(lambda: None) + + try: + with pytest.raises(RuntimeError, match="own thread"): + worker.call_blocking(_nested_dispatch) + finally: + worker.close_blocking() + + +def test_async_worker_call_is_not_reentrant_from_worker_thread() -> None: + worker = ModelExecutionWorker() + + def _nested_async_dispatch() -> None: + async def _dispatch() -> None: + await worker.call(lambda: None) + + asyncio.run(_dispatch()) + + try: + with pytest.raises(RuntimeError, match="own thread"): + worker.call_blocking(_nested_async_dispatch) + finally: + worker.close_blocking() From d661edb35db6bd3d7d2bbfa3fa457d20f9d9bc0b Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 01:42:45 +0000 Subject: [PATCH 05/51] Apply Phase 3 test import formatting --- flashdreams/tests/test_demo_runtime_host.py | 7 ++++++- 1 file changed, 6 insertions(+), 1 deletion(-) diff --git a/flashdreams/tests/test_demo_runtime_host.py b/flashdreams/tests/test_demo_runtime_host.py index eed7635b0..b14dafcb6 100644 --- a/flashdreams/tests/test_demo_runtime_host.py +++ b/flashdreams/tests/test_demo_runtime_host.py @@ -8,7 +8,12 @@ import pytest -from flashdreams.runtime import InferenceInput, InferenceSession, StepRequest, StepResult +from flashdreams.runtime import ( + InferenceInput, + InferenceSession, + StepRequest, + StepResult, +) from flashdreams.runtime.demo import ModelWarmupPlan, RuntimeHost, WarmupSessionInputs pytestmark = pytest.mark.ci_cpu From dbf33a9a9d68dd6d2ffdad8af8626d9baf2facbc Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 02:39:20 +0000 Subject: [PATCH 06/51] runtime: extract demo run mode session helpers --- .../flashdreams/runtime/demo/__init__.py | 8 + .../flashdreams/runtime/demo/drivers.py | 246 ++++++- .../flashdreams/runtime/demo/run_modes.py | 94 ++- .../tests/test_demo_runtime_run_modes.py | 626 ++++++++++++++++++ .../tests/test_demo_runtime_vertical_slice.py | 39 +- 5 files changed, 995 insertions(+), 18 deletions(-) create mode 100644 flashdreams/tests/test_demo_runtime_run_modes.py diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index c324bc6b2..64621ac56 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -7,6 +7,7 @@ BatchSessionDriver, DriverInvariantError, run_demo_session, + run_demo_session_async, ) from flashdreams.runtime.demo.host import ( ModelWarmupPlan, @@ -23,6 +24,7 @@ from flashdreams.runtime.demo.pipeline import StepOutcome, StepPipeline from flashdreams.runtime.demo.replay import run_replay_demo from flashdreams.runtime.demo.run_modes import ( + AsyncSessionDriver, DefaultErrorPolicy, ErrorAction, InMemorySessionMetricsRecorder, @@ -30,8 +32,10 @@ NoopTransportService, RunContext, RunMode, + RunModeWarmup, RunResult, RunSummary, + SessionDriver, SessionEdges, SingleSessionAdmissionPolicy, ) @@ -64,6 +68,7 @@ "DemoSpec", "DriverInvariantError", "ErrorAction", + "AsyncSessionDriver", "InMemorySessionMetricsRecorder", "InputSource", "MetricsSnapshot", @@ -81,10 +86,12 @@ "RealtimeInputSource", "RunContext", "RunMode", + "RunModeWarmup", "RunResult", "RunSummary", "RuntimeHost", "SessionEdges", + "SessionDriver", "SessionInfo", "SingleSessionAdmissionPolicy", "StepOutcome", @@ -95,5 +102,6 @@ "WebRTCOutputSpec", "build_output_target", "run_demo_session", + "run_demo_session_async", "run_replay_demo", ] diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 7bee95a96..e4ecc0750 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -5,7 +5,9 @@ from __future__ import annotations -from typing import Any +import asyncio +import inspect +from typing import Any, cast from flashdreams.runtime.interfaces import InferenceSession @@ -20,7 +22,7 @@ SessionEdges, SessionReservation, ) -from .session_inputs import ModelInputProvider +from .session_inputs import BatchInputSource, ModelInputProvider from .spec import DemoAdapter, DemoSpec, PreparedScenario @@ -63,6 +65,7 @@ def run_one_session( final_reason = str(exc) final_error = exc if action.result_status == "failed" else None + input_source = cast(BatchInputSource, session_edges.input_source) while setup_ok: if session is None: raise DriverInvariantError("setup_ok was set without a session.") @@ -72,7 +75,7 @@ def run_one_session( request = host.call(session.next_step_request) if request is None: break - user_window = session_edges.input_source.next_window(request) + user_window = input_source.next_window(request) outcome = host.call( pipeline.execute_step, request=request, @@ -141,7 +144,8 @@ def run_demo_session( reservation: SessionReservation | None = None, ) -> RunResult: """Run one prepared demo session through a selected run mode.""" - reservation = reservation or context.admission.try_reserve() + if reservation is None: + reservation = context.admission.try_reserve() if reservation is None: result = RunResult.rejected(reason="busy") context.run_metrics.record_session(result) @@ -166,14 +170,15 @@ def run_demo_session( provider=provider, adapter=adapter, ) - driver = run_mode.select_driver() - if not isinstance(driver, BatchSessionDriver): - raise TypeError( - "Phase 2 run_demo_session supports BatchSessionDriver only, " - f"got {type(driver).__name__}." + if session_edges.is_closed: + raise DriverInvariantError( + "RunMode returned already closed SessionEdges; session edges " + "must not be reused." ) + driver = run_mode.select_driver() driver_started = True - result = driver.run_one_session( + result = _run_sync_driver( + driver=driver, host=context.host, provider=provider, session_edges=session_edges, @@ -190,7 +195,9 @@ def run_demo_session( session_edges.metrics.record_cleanup_error(close_exc) else: context.run_metrics.record_cleanup_error(close_exc) - if session_edges is not None: + if session_edges is not None and ( + driver_started or not session_edges.is_closed + ): result = session_edges.close_result( status="failed", reason=str(exc), @@ -207,7 +214,9 @@ def run_demo_session( session_edges.metrics.record_cleanup_error(close_exc) else: context.run_metrics.record_cleanup_error(close_exc) - if session_edges is not None: + if session_edges is not None and ( + driver_started or not session_edges.is_closed + ): result = session_edges.close_result( status="failed", reason=str(exc), @@ -221,6 +230,106 @@ def run_demo_session( reservation.release() +async def run_demo_session_async( + *, + context: RunContext, + spec: DemoSpec, + scenario: PreparedScenario, + adapter: DemoAdapter, + run_mode: RunMode, + pipeline: StepPipeline, + reservation: SessionReservation | None = None, +) -> RunResult: + """Run one prepared async/realtime demo session through a selected run mode.""" + if reservation is None: + reservation = context.admission.try_reserve() + if reservation is None: + result = RunResult.rejected(reason="busy") + context.run_metrics.record_session(result) + return result + + provider: Any | None = None + session_edges: SessionEdges | None = None + driver_started = False + try: + try: + create_provider = getattr(adapter, "create_model_input_provider") + provider = await context.host.call_async(create_provider, spec, scenario) + run_mode.validate_session( + spec=spec, + scenario=scenario, + adapter=adapter, + provider=provider, + ) + session_edges = run_mode.create_session_edges( + context=context, + spec=spec, + scenario=scenario, + provider=provider, + adapter=adapter, + ) + if session_edges.is_closed: + raise DriverInvariantError( + "RunMode returned already closed SessionEdges; session edges " + "must not be reused." + ) + driver = run_mode.select_driver() + driver_started = True + result = await _run_async_driver( + driver=driver, + host=context.host, + provider=provider, + session_edges=session_edges, + pipeline=pipeline, + ) + context.run_metrics.record_session(result) + return result + except asyncio.CancelledError: + _uncancel_current_task() + result = await _close_partial_session_async( + context=context, + provider=provider, + session_edges=session_edges, + status="cancelled", + reason="cancelled during session assembly", + error=None, + close_provider=not driver_started, + ) + context.run_metrics.record_session(result) + return result + except DriverInvariantError as exc: + if provider is not None and not driver_started: + await _close_provider_async( + context=context, + provider=provider, + session_edges=session_edges, + ) + if session_edges is not None and ( + driver_started or not session_edges.is_closed + ): + result = session_edges.close_result( + status="failed", + reason=str(exc), + error=exc, + ) + context.run_metrics.record_session(result) + raise + except Exception as exc: + result = await _close_partial_session_async( + context=context, + provider=provider, + session_edges=session_edges, + status="failed", + reason=str(exc), + error=exc, + close_provider=not driver_started, + ) + context.run_metrics.record_session(result) + return result + finally: + reservation.release() + + def _session_info(session: InferenceSession) -> SessionInfo: session_info = getattr(session, "session_info", None) if not callable(session_info): @@ -241,8 +350,121 @@ def _close_safely(close: Any, session_edges: SessionEdges) -> None: session_edges.metrics.record_cleanup_error(exc) +def _run_sync_driver( + *, + driver: object, + host: RuntimeHost, + provider: ModelInputProvider, + session_edges: SessionEdges, + pipeline: StepPipeline, +) -> RunResult: + run_one_session = getattr(driver, "run_one_session", None) + if not callable(run_one_session): + raise TypeError( + "RunMode.select_driver() must return an object with run_one_session(...)." + ) + result = run_one_session( + host=host, + provider=provider, + session_edges=session_edges, + pipeline=pipeline, + ) + if inspect.isawaitable(result): + raise TypeError( + "run_demo_session(...) requires a synchronous session driver; " + "use run_demo_session_async(...) for async drivers." + ) + if not isinstance(result, RunResult): + raise TypeError( + "Session driver run_one_session(...) must return RunResult, " + f"got {type(result).__name__}." + ) + return result + + +async def _run_async_driver( + *, + driver: object, + host: RuntimeHost, + provider: ModelInputProvider, + session_edges: SessionEdges, + pipeline: StepPipeline, +) -> RunResult: + run_one_session = getattr(driver, "run_one_session", None) + if not callable(run_one_session): + raise TypeError( + "RunMode.select_driver() must return an object with run_one_session(...)." + ) + result = run_one_session( + host=host, + provider=provider, + session_edges=session_edges, + pipeline=pipeline, + ) + if not inspect.isawaitable(result): + raise TypeError("run_demo_session_async(...) requires an async session driver.") + resolved = await result + if not isinstance(resolved, RunResult): + raise TypeError( + "Async session driver run_one_session(...) must return RunResult, " + f"got {type(resolved).__name__}." + ) + return resolved + + +async def _close_partial_session_async( + *, + context: RunContext, + provider: Any | None, + session_edges: SessionEdges | None, + status: DriverStatus, + reason: str | None, + error: Exception | None, + close_provider: bool, +) -> RunResult: + if provider is not None and close_provider: + await _close_provider_async( + context=context, + provider=provider, + session_edges=session_edges, + ) + if session_edges is not None: + return session_edges.close_result(status=status, reason=reason, error=error) + return RunResult(status=status, reason=reason, error=error) + + +async def _close_provider_async( + *, + context: RunContext, + provider: Any, + session_edges: SessionEdges | None, +) -> None: + try: + await context.host.call_async(provider.close) + except Exception as close_exc: + if session_edges is not None: + session_edges.metrics.record_cleanup_error(close_exc) + else: + context.run_metrics.record_cleanup_error(close_exc) + + +def _uncancel_current_task() -> None: + task = asyncio.current_task() + if task is None: + return + uncancel = getattr(task, "uncancel", None) + if not callable(uncancel): + return + cancelling = getattr(task, "cancelling", None) + if not callable(cancelling): + return + while cancelling(): + uncancel() + + __all__ = [ "BatchSessionDriver", "DriverInvariantError", "run_demo_session", + "run_demo_session_async", ] diff --git a/flashdreams/flashdreams/runtime/demo/run_modes.py b/flashdreams/flashdreams/runtime/demo/run_modes.py index 973749d9d..bba8ed650 100644 --- a/flashdreams/flashdreams/runtime/demo/run_modes.py +++ b/flashdreams/flashdreams/runtime/demo/run_modes.py @@ -5,6 +5,7 @@ from __future__ import annotations +import asyncio from collections.abc import Mapping, Sequence from dataclasses import dataclass, field from threading import Lock @@ -13,11 +14,13 @@ from flashdreams.runtime._utils import freeze_mapping from flashdreams.runtime.output import OutputArtifact +from .host import ModelWarmupPlan from .outputs import OutputDecision, OutputSink if TYPE_CHECKING: from .host import RuntimeHost - from .session_inputs import BatchInputSource, ModelInputProvider + from .pipeline import StepPipeline + from .session_inputs import InputSource, ModelInputProvider from .spec import DemoAdapter, DemoSpec, PreparedScenario SessionStatus = Literal[ @@ -291,6 +294,34 @@ class AdmissionPolicy(Protocol): def try_reserve(self) -> SessionReservation | None: ... +@runtime_checkable +class SessionDriver(Protocol): + """Synchronous one-session driver selected by a run mode.""" + + def run_one_session( + self, + *, + host: "RuntimeHost", + provider: "ModelInputProvider", + session_edges: "SessionEdges", + pipeline: "StepPipeline", + ) -> RunResult: ... + + +@runtime_checkable +class AsyncSessionDriver(Protocol): + """Async one-session driver selected by realtime run modes.""" + + async def run_one_session( + self, + *, + host: "RuntimeHost", + provider: "ModelInputProvider", + session_edges: "SessionEdges", + pipeline: "StepPipeline", + ) -> RunResult: ... + + @dataclass(slots=True) class RunContext: """Run-scoped services shared by one or more demo sessions.""" @@ -298,12 +329,18 @@ class RunContext: host: "RuntimeHost" run_metrics: SessionMetricsRecorder admission: AdmissionPolicy + model_warmup_plan: ModelWarmupPlan = field(default_factory=ModelWarmupPlan) services: Mapping[str, object] = field(default_factory=dict) + cleanup_tasks: set[asyncio.Task[RunResult]] = field(default_factory=set) def __post_init__(self) -> None: self.services = freeze_mapping(self.services) def close(self) -> RunSummary: + if self.cleanup_tasks: + raise RuntimeError( + "Pending session cleanup tasks; async runs must await close_async()." + ) for service in self.services.values(): close = getattr(service, "close", None) if callable(close): @@ -316,13 +353,21 @@ def close(self) -> RunSummary: sessions=tuple(getattr(self.run_metrics, "sessions", ())), ) + async def close_async(self) -> RunSummary: + while self.cleanup_tasks: + pending = tuple(self.cleanup_tasks) + await asyncio.gather(*pending, return_exceptions=True) + self.cleanup_tasks.difference_update(pending) + return self.close() + @dataclass(slots=True) class SessionEdges: """Per-session input/output/policy bundle consumed by drivers.""" - input_source: "BatchInputSource" + input_source: "InputSource" output_sink: OutputSink + cleanup_tasks: set[asyncio.Task[RunResult]] metrics: SessionMetricsRecorder = field( default_factory=InMemorySessionMetricsRecorder ) @@ -330,9 +375,13 @@ class SessionEdges: transport: TransportService = field(default_factory=NoopTransportService) clock: object | None = None activation: object | None = None - cleanup_tasks: set[object] = field(default_factory=set) _closed_result: RunResult | None = field(default=None, init=False, repr=False) + @property + def is_closed(self) -> bool: + """Return whether ``close_result(...)`` has already finalized this session.""" + return self._closed_result is not None + def close_result( self, *, @@ -369,7 +418,16 @@ def close_result( @runtime_checkable class RunMode(Protocol): - """Minimal Phase 2 run-mode surface used by ``run_demo_session``.""" + """Run/session construction strategy consumed by shared helpers.""" + + name: str + + def validate_run( + self, + *, + spec: "DemoSpec", + adapter: "DemoAdapter", + ) -> None: ... def validate_session( self, @@ -380,6 +438,15 @@ def validate_session( provider: "ModelInputProvider", ) -> None: ... + def create_run_context( + self, + *, + spec: "DemoSpec", + adapter: "DemoAdapter", + host: "RuntimeHost", + model_warmup_plan: ModelWarmupPlan, + ) -> RunContext: ... + def create_session_edges( self, *, @@ -390,11 +457,26 @@ def create_session_edges( adapter: "DemoAdapter", ) -> SessionEdges: ... - def select_driver(self) -> object: ... + def select_driver(self) -> SessionDriver | AsyncSessionDriver: ... + + +@runtime_checkable +class RunModeWarmup(Protocol): + """Optional run-mode warmup for output or transport services.""" + + def warmup_context( + self, + *, + context: RunContext, + spec: "DemoSpec", + scenario: "PreparedScenario", + adapter: "DemoAdapter", + ) -> None: ... __all__ = [ "AdmissionPolicy", + "AsyncSessionDriver", "DefaultErrorPolicy", "DriverStatus", "ErrorAction", @@ -404,9 +486,11 @@ def select_driver(self) -> object: ... "NoopTransportService", "RunContext", "RunMode", + "RunModeWarmup", "RunResult", "RunSummary", "SessionEdges", + "SessionDriver", "SessionMetricsRecorder", "SessionReservation", "SessionStatus", diff --git a/flashdreams/tests/test_demo_runtime_run_modes.py b/flashdreams/tests/test_demo_runtime_run_modes.py new file mode 100644 index 000000000..d5a6e4f04 --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_run_modes.py @@ -0,0 +1,626 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +import contextlib +from collections.abc import Callable, Coroutine, Mapping, Sequence +from pathlib import Path +from typing import Any, cast + +import pytest + +from flashdreams.runtime import ( + CanonicalInputSchema, + IdentityInputMapping, + InferenceConfig, + InferenceInput, + InferenceInputSchema, + InferenceRuntime, + InferenceSession, + InputMapping, +) +from flashdreams.runtime.demo import ( + AsyncSessionDriver, + DemoSpec, + DriverInvariantError, + InMemorySessionMetricsRecorder, + ModelWarmupPlan, + Mp4OutputSpec, + NullOutputSpec, + OutputDecision, + PreparedScenario, + RunContext, + RunResult, + RuntimeHost, + SessionDriver, + SessionEdges, + SessionInfo, + StepPipeline, + WebRTCOutputSpec, + run_demo_session, + run_demo_session_async, +) + +pytestmark = pytest.mark.ci_cpu + + +def test_fake_mp4_run_mode_calls_session_helper_once(tmp_path: Path) -> None: + spec = DemoSpec( + model_id="fake-demo", + input_mode="replay", + output=Mp4OutputSpec(path=tmp_path / "fake.mp4", fps=12), + ) + adapter = _FakeAdapter() + mode = _FakeRunMode(name="mp4", driver=_ClosingSyncDriver()) + runtime = _UnusedRuntime() + helper_calls: list[DemoSpec] = [] + + result = _run_fake_single_session_mode( + spec=spec, + adapter=adapter, + mode=mode, + runtime=runtime, + helper=lambda **kwargs: _record_sync_helper(helper_calls, **kwargs), + ) + + assert result.status == "completed" + assert helper_calls == [spec] + assert len(mode.created_edges) == 1 + assert mode.created_edges[0].is_closed + assert mode.created_edges[0].cleanup_tasks is mode.require_context().cleanup_tasks + assert adapter.providers[0].close_count == 1 + assert mode.admission.reservations[0].release_count == 1 + + +def test_fake_benchmark_run_mode_calls_helper_once_per_scenario() -> None: + specs = [ + DemoSpec( + model_id="fake-demo", + input_mode="replay", + output=NullOutputSpec(), + scenario=f"scenario-{index}", + ) + for index in range(2) + ] + adapter = _FakeAdapter() + mode = _FakeRunMode(name="benchmark", driver=_ClosingSyncDriver()) + helper_calls: list[DemoSpec] = [] + + results = _run_fake_benchmark_mode( + specs=specs, + adapter=adapter, + mode=mode, + runtime=_UnusedRuntime(), + helper=lambda **kwargs: _record_sync_helper(helper_calls, **kwargs), + ) + + assert [result.status for result in results] == ["completed", "completed"] + assert helper_calls == specs + assert adapter.prepare_scenario_calls == specs + assert len(adapter.providers) == 2 + assert len({id(provider) for provider in adapter.providers}) == 2 + assert all(provider.close_count == 1 for provider in adapter.providers) + assert len(mode.created_edges) == 2 + assert len({id(edges) for edges in mode.created_edges}) == 2 + assert all(edges.is_closed for edges in mode.created_edges) + assert all( + reservation.release_count == 1 for reservation in mode.admission.reservations + ) + + +@pytest.mark.asyncio +async def test_fake_webrtc_offer_reserves_before_prepare_or_negotiation() -> None: + spec = DemoSpec( + model_id="fake-demo", + input_mode="keyboard-driving", + output=WebRTCOutputSpec(port=8081), + ) + events: list[str] = [] + blocking_io = _BlockingIOService(events) + webrtc = _FakeWebRTCService(events) + adapter = _FakeAdapter(events=events) + mode = _FakeRunMode( + name="webrtc", + driver=_ClosingAsyncDriver(), + admission=_RecordingAdmission(events=events), + services={"blocking_io": blocking_io, "webrtc": webrtc}, + ) + context = mode.create_run_context( + spec=spec, + adapter=adapter, + host=RuntimeHost(_UnusedRuntime()), + model_warmup_plan=ModelWarmupPlan(), + ) + helper_calls: list[DemoSpec] = [] + + answer = await _handle_fake_webrtc_offer( + context=context, + spec=spec, + adapter=adapter, + mode=mode, + helper=lambda **kwargs: _record_async_helper(helper_calls, **kwargs), + events=events, + ) + + assert answer == "answer" + assert events.index("admission.reserve") < events.index("blocking_io.run") + assert events.index("blocking_io.run") < events.index("webrtc.answer") + assert blocking_io.run_count == 1 + assert helper_calls == [spec] + assert mode.admission.reservations[0].release_count == 1 + assert adapter.providers[0].close_count == 1 + assert mode.created_edges[0].is_closed + + +def test_run_demo_session_rejects_reused_closed_session_edges() -> None: + spec = DemoSpec( + model_id="fake-demo", + input_mode="replay", + output=NullOutputSpec(), + ) + adapter = _FakeAdapter() + mode = _ReusingRunMode(name="mp4", driver=_ClosingSyncDriver()) + context = mode.create_run_context( + spec=spec, + adapter=adapter, + host=RuntimeHost(_UnusedRuntime()), + model_warmup_plan=ModelWarmupPlan(), + ) + scenario = adapter.prepare_scenario(spec) + + first = run_demo_session( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + run_mode=mode, + pipeline=StepPipeline(), + ) + with pytest.raises(DriverInvariantError, match="must not be reused"): + run_demo_session( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + run_mode=mode, + pipeline=StepPipeline(), + ) + + assert first.status == "completed" + run_metrics = cast(InMemorySessionMetricsRecorder, context.run_metrics) + assert len(run_metrics.sessions) == 1 + assert run_metrics.sessions[0] is first + assert adapter.providers[1].close_count == 1 + + +@pytest.mark.asyncio +async def test_run_context_close_async_drains_cleanup_tasks() -> None: + metrics = InMemorySessionMetricsRecorder() + context = RunContext( + host=RuntimeHost(_UnusedRuntime()), + run_metrics=metrics, + admission=_RecordingAdmission(events=[]), + ) + task = asyncio.create_task(_finished_cleanup_result()) + context.cleanup_tasks.add(task) + + with pytest.raises(RuntimeError, match="Pending session cleanup tasks"): + context.close() + + summary = await context.close_async() + + assert not context.cleanup_tasks + assert task.done() + assert summary.metrics.counters["sessions"] == 0 + assert metrics.closed + + +def _run_fake_single_session_mode( + *, + spec: DemoSpec, + adapter: "_FakeAdapter", + mode: "_FakeRunMode", + runtime: "_UnusedRuntime", + helper: Callable[..., RunResult], +) -> RunResult: + mode.validate_run(spec=spec, adapter=adapter) + scenario = adapter.prepare_scenario(spec) + context = mode.create_run_context( + spec=spec, + adapter=adapter, + host=RuntimeHost(runtime), + model_warmup_plan=ModelWarmupPlan(), + ) + mode.warmup_context( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + ) + return helper( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + run_mode=mode, + pipeline=StepPipeline(), + ) + + +def _run_fake_benchmark_mode( + *, + specs: Sequence[DemoSpec], + adapter: "_FakeAdapter", + mode: "_FakeRunMode", + runtime: "_UnusedRuntime", + helper: Callable[..., RunResult], +) -> list[RunResult]: + mode.validate_run(spec=specs[0], adapter=adapter) + context = mode.create_run_context( + spec=specs[0], + adapter=adapter, + host=RuntimeHost(runtime), + model_warmup_plan=ModelWarmupPlan(), + ) + results: list[RunResult] = [] + for spec in specs: + scenario = adapter.prepare_scenario(spec) + results.append( + helper( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + run_mode=mode, + pipeline=StepPipeline(), + ) + ) + return results + + +async def _handle_fake_webrtc_offer( + *, + context: RunContext, + spec: DemoSpec, + adapter: "_FakeAdapter", + mode: "_FakeRunMode", + helper: Callable[..., Coroutine[Any, Any, RunResult]], + events: list[str], +) -> str: + events.append("handler.start") + reservation = context.admission.try_reserve() + if reservation is None: + return "busy" + + task: asyncio.Task[RunResult] | None = None + try: + blocking_io = cast(_BlockingIOService, context.services["blocking_io"]) + scenario = await blocking_io.run(adapter.prepare_scenario, spec) + task = asyncio.create_task( + helper( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + run_mode=mode, + pipeline=StepPipeline(), + reservation=reservation, + ) + ) + webrtc = cast(_FakeWebRTCService, context.services["webrtc"]) + return await webrtc.answer(task) + except Exception: + if task is not None: + task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await task + reservation.release() + raise + + +def _record_sync_helper(calls: list[DemoSpec], **kwargs: Any) -> RunResult: + calls.append(kwargs["spec"]) + return run_demo_session(**kwargs) + + +async def _record_async_helper(calls: list[DemoSpec], **kwargs: Any) -> RunResult: + calls.append(kwargs["spec"]) + return await run_demo_session_async(**kwargs) + + +async def _finished_cleanup_result() -> RunResult: + await asyncio.sleep(0) + return RunResult.rejected(reason="test cleanup") + + +class _FakeAdapter: + model_id = "fake-demo" + inference_input_schema = InferenceInputSchema() + canonical_input_schema = CanonicalInputSchema() + + def __init__(self, *, events: list[str] | None = None) -> None: + self.events = events + self.prepare_scenario_calls: list[DemoSpec] = [] + self.providers: list[_FakeProvider] = [] + + def supported_input_modes(self) -> tuple[str, ...]: + return ("replay", "keyboard-driving") + + def supported_output_modes(self) -> tuple[str, ...]: + return ("null", "mp4", "webrtc") + + def default_input_mapping(self) -> InputMapping: + return IdentityInputMapping() + + def validate_config(self, config: InferenceConfig) -> None: + if config.model_id != self.model_id: + raise ValueError(f"Unsupported model_id={config.model_id!r}.") + + def create_runtime(self, config: InferenceConfig) -> InferenceRuntime: + self.validate_config(config) + return _UnusedRuntime() + + def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: + if self.events is not None: + self.events.append("adapter.prepare_scenario") + self.prepare_scenario_calls.append(spec) + return PreparedScenario(initial_inputs=InferenceInput()) + + def create_model_input_provider( + self, + spec: DemoSpec, + scenario: PreparedScenario, + ) -> "_FakeProvider": + del spec, scenario + provider = _FakeProvider() + self.providers.append(provider) + return provider + + +class _FakeProvider: + def __init__(self) -> None: + self.close_count = 0 + + def close(self) -> None: + self.close_count += 1 + + +class _FakeRunMode: + def __init__( + self, + *, + name: str, + driver: SessionDriver | AsyncSessionDriver, + admission: "_RecordingAdmission | None" = None, + services: Mapping[str, object] | None = None, + ) -> None: + self.name = name + self.driver = driver + self.created_edges: list[SessionEdges] = [] + self.validate_run_count = 0 + self.warmup_count = 0 + self.admission = admission or _RecordingAdmission(events=[]) + self.services = services or {} + self.context: RunContext | None = None + + def require_context(self) -> RunContext: + if self.context is None: + raise AssertionError("Run context was not created.") + return self.context + + def validate_run(self, *, spec: DemoSpec, adapter: Any) -> None: + del spec, adapter + self.validate_run_count += 1 + + def validate_session( + self, + *, + spec: DemoSpec, + scenario: Any, + adapter: Any, + provider: Any, + ) -> None: + del spec, scenario, adapter, provider + + def create_run_context( + self, + *, + spec: DemoSpec, + adapter: Any, + host: RuntimeHost, + model_warmup_plan: ModelWarmupPlan, + ) -> RunContext: + del spec, adapter + self.context = RunContext( + host=host, + run_metrics=InMemorySessionMetricsRecorder(), + admission=self.admission, + model_warmup_plan=model_warmup_plan, + services=self.services, + ) + return self.context + + def create_session_edges( + self, + *, + context: RunContext, + spec: DemoSpec, + scenario: Any, + provider: Any, + adapter: Any, + ) -> SessionEdges: + del spec, scenario, provider, adapter + edges = SessionEdges( + input_source=_FinishedInputSource(), + output_sink=_RecordingOutputSink(), + cleanup_tasks=context.cleanup_tasks, + metrics=InMemorySessionMetricsRecorder(), + transport=_RecordingTransport(), + ) + self.created_edges.append(edges) + return edges + + def select_driver(self) -> SessionDriver | AsyncSessionDriver: + return self.driver + + def warmup_context( + self, + *, + context: RunContext, + spec: DemoSpec, + scenario: Any, + adapter: Any, + ) -> None: + del context, spec, scenario, adapter + self.warmup_count += 1 + + +class _ReusingRunMode(_FakeRunMode): + def __init__( + self, + *, + name: str, + driver: SessionDriver | AsyncSessionDriver, + ) -> None: + super().__init__(name=name, driver=driver) + self._edges: SessionEdges | None = None + + def create_session_edges( + self, + *, + context: RunContext, + spec: DemoSpec, + scenario: Any, + provider: Any, + adapter: Any, + ) -> SessionEdges: + if self._edges is None: + self._edges = super().create_session_edges( + context=context, + spec=spec, + scenario=scenario, + provider=provider, + adapter=adapter, + ) + return self._edges + + +class _ClosingSyncDriver: + def run_one_session( + self, + *, + host: RuntimeHost, + provider: Any, + session_edges: SessionEdges, + pipeline: StepPipeline, + ) -> RunResult: + del host, pipeline + provider.close() + return session_edges.close_result(status="completed") + + +class _ClosingAsyncDriver: + async def run_one_session( + self, + *, + host: RuntimeHost, + provider: Any, + session_edges: SessionEdges, + pipeline: StepPipeline, + ) -> RunResult: + del pipeline + await host.call_async(provider.close) + return session_edges.close_result(status="completed") + + +class _RecordingAdmission: + def __init__(self, *, events: list[str]) -> None: + self.events = events + self.reservations: list[_RecordingReservation] = [] + + def try_reserve(self) -> "_RecordingReservation": + self.events.append("admission.reserve") + reservation = _RecordingReservation() + self.reservations.append(reservation) + return reservation + + +class _RecordingReservation: + def __init__(self) -> None: + self.release_count = 0 + + def release(self) -> None: + if self.release_count: + return + self.release_count += 1 + + +class _FinishedInputSource: + is_finite = True + is_deterministic = True + + def is_finished(self) -> bool: + return True + + +class _RecordingOutputSink: + produces_artifacts = False + + def __init__(self) -> None: + self.close_count = 0 + + def open(self, session_info: SessionInfo) -> None: + del session_info + + def begin_generation(self, generation: int) -> None: + del generation + + def write(self, result: Any) -> OutputDecision: + del result + return OutputDecision() + + def close(self) -> Sequence[Any]: + self.close_count += 1 + return () + + +class _RecordingTransport: + def close(self) -> None: + return + + def is_active(self) -> bool: + return True + + +class _BlockingIOService: + def __init__(self, events: list[str]) -> None: + self.events = events + self.run_count = 0 + + async def run(self, func: Callable[..., Any], *args: object) -> Any: + self.run_count += 1 + self.events.append("blocking_io.run") + await asyncio.sleep(0) + return func(*args) + + +class _FakeWebRTCService: + def __init__(self, events: list[str]) -> None: + self.events = events + + async def answer(self, task: asyncio.Task[RunResult]) -> str: + self.events.append("webrtc.answer") + result = await task + assert result.status == "completed" + return "answer" + + +class _UnusedRuntime: + def start_session(self, inputs: InferenceInput) -> InferenceSession: + del inputs + raise AssertionError("The fake Phase 4 drivers do not start sessions.") + + def close(self) -> None: + return diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index e550e1d26..084176773 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -29,6 +29,7 @@ DriverInvariantError, ErrorAction, InMemorySessionMetricsRecorder, + ModelWarmupPlan, NullOutputSpec, OutputDecision, PreparedScenario, @@ -81,6 +82,7 @@ def test_batch_driver_runs_fake_video_demo_through_runtime_host() -> None: edges = SessionEdges( input_source=_FakeBatchInputSource(num_windows=2), output_sink=output, + cleanup_tasks=set(), metrics=metrics, ) @@ -191,6 +193,7 @@ def test_setup_failure_returns_failed_before_runtime_session_creation() -> None: session_edges=SessionEdges( input_source=_FakeBatchInputSource(num_windows=1), output_sink=_RecordingOutputSink(), + cleanup_tasks=set(), metrics=metrics, ), pipeline=StepPipeline(), @@ -236,6 +239,7 @@ def test_setup_failure_can_return_skipped_but_not_completed() -> None: session_edges=SessionEdges( input_source=_FakeBatchInputSource(num_windows=1), output_sink=_RecordingOutputSink(), + cleanup_tasks=set(), error_policy=_SetupPolicy(result_status="skipped"), ), pipeline=StepPipeline(), @@ -250,6 +254,7 @@ def test_setup_failure_can_return_skipped_but_not_completed() -> None: edges = SessionEdges( input_source=_FakeBatchInputSource(num_windows=1), output_sink=output, + cleanup_tasks=set(), metrics=metrics, error_policy=_SetupPolicy(result_status="completed"), transport=transport, @@ -317,6 +322,7 @@ def test_input_source_finished_error_returns_failed_not_completed() -> None: fail_is_finished=RuntimeError("input source failed"), ), output_sink=_RecordingOutputSink(), + cleanup_tasks=set(), metrics=metrics, ), pipeline=StepPipeline(), @@ -337,6 +343,7 @@ def test_step_failure_returns_failed_from_driver() -> None: session_edges=SessionEdges( input_source=_FakeBatchInputSource(num_windows=1), output_sink=output, + cleanup_tasks=set(), ), pipeline=StepPipeline(), ) @@ -357,6 +364,7 @@ def test_session_edges_close_result_is_idempotent_and_first_result_wins() -> Non edges = SessionEdges( input_source=_FakeBatchInputSource(num_windows=0), output_sink=output, + cleanup_tasks=set(), metrics=metrics, transport=transport, ) @@ -676,6 +684,8 @@ def create_model_input_provider( class _FakeRunMode: + name = "fake" + def __init__( self, *, @@ -697,6 +707,32 @@ def __init__( self.error_policy = error_policy self.validate_error = validate_error + def validate_run( + self, + *, + spec: DemoSpec, + adapter: Any, + ) -> None: + del spec, adapter + + def create_run_context( + self, + *, + spec: DemoSpec, + adapter: Any, + host: RuntimeHost, + model_warmup_plan: ModelWarmupPlan, + ) -> RunContext: + del spec, adapter + return RunContext( + host=host, + run_metrics=InMemorySessionMetricsRecorder(), + admission=SingleSessionAdmissionPolicy( + health_check=lambda: host.is_healthy + ), + model_warmup_plan=model_warmup_plan, + ) + def validate_session( self, *, @@ -718,7 +754,7 @@ def create_session_edges( provider: Any, adapter: Any, ) -> SessionEdges: - del context, provider, adapter + del provider, adapter output_sink = ( self.output_sink_factory(spec, scenario) if self.output_sink_factory is not None @@ -727,6 +763,7 @@ def create_session_edges( return SessionEdges( input_source=self.input_source, output_sink=output_sink, + cleanup_tasks=context.cleanup_tasks, metrics=self.metrics, error_policy=self.error_policy or _SetupPolicy(result_status="failed"), transport=self.transport or _RecordingTransport(), From 6f83fdf1153b99c370ddf534337c8730872b85a9 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 02:49:45 +0000 Subject: [PATCH 07/51] runtime: finalize session edges when invariant cleanup dispatch fails --- .../flashdreams/runtime/demo/drivers.py | 24 ++++++++++++-- .../tests/test_demo_runtime_vertical_slice.py | 33 +++++++++++++++++++ 2 files changed, 55 insertions(+), 2 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index e4ecc0750..7b34d792b 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -107,8 +107,16 @@ def run_one_session( break except DriverInvariantError as exc: if session is not None: - host.call(_close_safely, session.close, session_edges) - host.call(_close_safely, provider.close, session_edges) + _close_on_host_best_effort( + host=host, + close=session.close, + session_edges=session_edges, + ) + _close_on_host_best_effort( + host=host, + close=provider.close, + session_edges=session_edges, + ) session_edges.close_result( status="failed", reason=str(exc), @@ -350,6 +358,18 @@ def _close_safely(close: Any, session_edges: SessionEdges) -> None: session_edges.metrics.record_cleanup_error(exc) +def _close_on_host_best_effort( + *, + host: RuntimeHost, + close: Any, + session_edges: SessionEdges, +) -> None: + try: + host.call(_close_safely, close, session_edges) + except Exception as exc: + session_edges.metrics.record_cleanup_error(exc) + + def _run_sync_driver( *, driver: object, diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index 084176773..7578a5063 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -274,6 +274,39 @@ def test_setup_failure_can_return_skipped_but_not_completed() -> None: assert provider.close_count == 1 +def test_batch_driver_invariant_finalizes_edges_when_host_closed() -> None: + runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) + host = RuntimeHost(runtime) + host.close() + provider = _FakeVideoModelInputProvider() + output = _RecordingOutputSink() + transport = _RecordingTransport() + metrics = InMemorySessionMetricsRecorder() + edges = SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=output, + cleanup_tasks=set(), + metrics=metrics, + error_policy=_SetupPolicy(result_status="completed"), + transport=transport, + ) + + with pytest.raises(DriverInvariantError, match="Setup failures"): + BatchSessionDriver().run_one_session( + host=host, + provider=provider, + session_edges=edges, + pipeline=StepPipeline(), + ) + + assert edges.is_closed + assert output.close_count == 1 + assert transport.close_count == 1 + assert metrics.closed + assert metrics.cleanup_errors == ["runtime host is closed"] + assert provider.close_count == 0 + + def test_run_demo_session_closes_edges_when_driver_invariant_escapes() -> None: runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) run_metrics = InMemorySessionMetricsRecorder() From 631963bd4a5da713b7e06389836f130a8b42fd1c Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 02:58:48 +0000 Subject: [PATCH 08/51] runtime: guard cleanup error recording during teardown --- .../flashdreams/runtime/demo/drivers.py | 10 ++-- .../flashdreams/runtime/demo/run_modes.py | 11 ++++- .../tests/test_demo_runtime_vertical_slice.py | 48 +++++++++++++++++++ 3 files changed, 62 insertions(+), 7 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 7b34d792b..326ba9692 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -200,7 +200,7 @@ def run_demo_session( context.host.call(provider.close) except Exception as close_exc: if session_edges is not None: - session_edges.metrics.record_cleanup_error(close_exc) + session_edges.record_cleanup_error(close_exc) else: context.run_metrics.record_cleanup_error(close_exc) if session_edges is not None and ( @@ -219,7 +219,7 @@ def run_demo_session( context.host.call(provider.close) except Exception as close_exc: if session_edges is not None: - session_edges.metrics.record_cleanup_error(close_exc) + session_edges.record_cleanup_error(close_exc) else: context.run_metrics.record_cleanup_error(close_exc) if session_edges is not None and ( @@ -355,7 +355,7 @@ def _close_safely(close: Any, session_edges: SessionEdges) -> None: try: close() except Exception as exc: - session_edges.metrics.record_cleanup_error(exc) + session_edges.record_cleanup_error(exc) def _close_on_host_best_effort( @@ -367,7 +367,7 @@ def _close_on_host_best_effort( try: host.call(_close_safely, close, session_edges) except Exception as exc: - session_edges.metrics.record_cleanup_error(exc) + session_edges.record_cleanup_error(exc) def _run_sync_driver( @@ -463,7 +463,7 @@ async def _close_provider_async( await context.host.call_async(provider.close) except Exception as close_exc: if session_edges is not None: - session_edges.metrics.record_cleanup_error(close_exc) + session_edges.record_cleanup_error(close_exc) else: context.run_metrics.record_cleanup_error(close_exc) diff --git a/flashdreams/flashdreams/runtime/demo/run_modes.py b/flashdreams/flashdreams/runtime/demo/run_modes.py index bba8ed650..55ad870d8 100644 --- a/flashdreams/flashdreams/runtime/demo/run_modes.py +++ b/flashdreams/flashdreams/runtime/demo/run_modes.py @@ -382,6 +382,13 @@ def is_closed(self) -> bool: """Return whether ``close_result(...)`` has already finalized this session.""" return self._closed_result is not None + def record_cleanup_error(self, exc: Exception) -> None: + """Record a cleanup error without letting metrics failures block teardown.""" + try: + self.metrics.record_cleanup_error(exc) + except Exception: + return + def close_result( self, *, @@ -397,11 +404,11 @@ def close_result( try: artifacts = tuple(self.output_sink.close()) except Exception as exc: - self.metrics.record_cleanup_error(exc) + self.record_cleanup_error(exc) try: self.transport.close() except Exception as exc: - self.metrics.record_cleanup_error(exc) + self.record_cleanup_error(exc) try: metrics = self.metrics.close() except Exception as exc: diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index 7578a5063..3a5e81062 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -307,6 +307,41 @@ def test_batch_driver_invariant_finalizes_edges_when_host_closed() -> None: assert provider.close_count == 0 +def test_batch_driver_invariant_finalizes_edges_when_cleanup_metrics_fail() -> None: + runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) + host = RuntimeHost(runtime) + host.close() + provider = _FakeVideoModelInputProvider() + output = _RecordingOutputSink() + transport = _RecordingTransport() + metrics = _FailingCleanupMetrics() + edges = SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=output, + cleanup_tasks=set(), + metrics=metrics, + error_policy=_SetupPolicy(result_status="completed"), + transport=transport, + ) + + with pytest.raises(DriverInvariantError, match="Setup failures") as raised: + BatchSessionDriver().run_one_session( + host=host, + provider=provider, + session_edges=edges, + pipeline=StepPipeline(), + ) + + result = edges.close_result() + assert result.status == "failed" + assert result.error is raised.value + assert output.close_count == 1 + assert transport.close_count == 1 + assert metrics.closed + assert metrics.cleanup_error_attempts == 1 + assert provider.close_count == 0 + + def test_run_demo_session_closes_edges_when_driver_invariant_escapes() -> None: runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) run_metrics = InMemorySessionMetricsRecorder() @@ -660,6 +695,19 @@ def close(self) -> None: self.close_count += 1 +class _FailingCleanupMetrics(InMemorySessionMetricsRecorder): + cleanup_error_attempts: int + + def __init__(self) -> None: + super().__init__() + self.cleanup_error_attempts = 0 + + def record_cleanup_error(self, exc: Exception) -> None: + del exc + self.cleanup_error_attempts += 1 + raise RuntimeError("cleanup metrics failed") + + class _SetupPolicy: def __init__( self, From dbfd430bd224b55dbaabff185802529afb6f5ec8 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 03:38:22 +0000 Subject: [PATCH 09/51] runtime: add realtime timing contracts --- .../flashdreams/runtime/demo/__init__.py | 30 ++ .../flashdreams/runtime/demo/run_modes.py | 5 +- .../runtime/demo/session_inputs.py | 9 +- .../flashdreams/runtime/demo/timing.py | 371 ++++++++++++++++++ flashdreams/tests/test_demo_runtime_timing.py | 236 +++++++++++ 5 files changed, 646 insertions(+), 5 deletions(-) create mode 100644 flashdreams/flashdreams/runtime/demo/timing.py create mode 100644 flashdreams/tests/test_demo_runtime_timing.py diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index 64621ac56..1637abfd3 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -58,6 +58,22 @@ WebRTCAppResources, WebRTCOutputSpec, ) +from flashdreams.runtime.demo.timing import ( + SPARSE_KEY_SEGMENTS_METADATA_KEY, + ActivationPolicy, + ActivationResult, + ActivationSignal, + AlwaysActiveActivationPolicy, + CatchUpDecision, + CatchUpPolicy, + DeterministicClock, + KeyboardRealtimeInputSource, + RealtimeClock, + RealtimeWindowResult, + ResamplerRealtimeClock, + SignalActivationPolicy, + input_frame_count_from_request, +) __all__ = [ "BatchInputSource", @@ -69,11 +85,19 @@ "DriverInvariantError", "ErrorAction", "AsyncSessionDriver", + "ActivationPolicy", + "ActivationResult", + "ActivationSignal", + "AlwaysActiveActivationPolicy", "InMemorySessionMetricsRecorder", "InputSource", + "CatchUpDecision", + "CatchUpPolicy", + "DeterministicClock", "MetricsSnapshot", "ModelWarmupPlan", "ModelInputProvider", + "KeyboardRealtimeInputSource", "Mp4OutputSpec", "NoopTransportService", "NullOutputSpec", @@ -84,6 +108,9 @@ "PreparedScenario", "PreparedStep", "RealtimeInputSource", + "RealtimeClock", + "RealtimeWindowResult", + "ResamplerRealtimeClock", "RunContext", "RunMode", "RunModeWarmup", @@ -93,7 +120,9 @@ "SessionEdges", "SessionDriver", "SessionInfo", + "SignalActivationPolicy", "SingleSessionAdmissionPolicy", + "SPARSE_KEY_SEGMENTS_METADATA_KEY", "StepOutcome", "StepPipeline", "UserInputWindow", @@ -101,6 +130,7 @@ "WebRTCAppResources", "WebRTCOutputSpec", "build_output_target", + "input_frame_count_from_request", "run_demo_session", "run_demo_session_async", "run_replay_demo", diff --git a/flashdreams/flashdreams/runtime/demo/run_modes.py b/flashdreams/flashdreams/runtime/demo/run_modes.py index 55ad870d8..c815e6170 100644 --- a/flashdreams/flashdreams/runtime/demo/run_modes.py +++ b/flashdreams/flashdreams/runtime/demo/run_modes.py @@ -22,6 +22,7 @@ from .pipeline import StepPipeline from .session_inputs import InputSource, ModelInputProvider from .spec import DemoAdapter, DemoSpec, PreparedScenario + from .timing import ActivationPolicy, DeterministicClock, RealtimeClock SessionStatus = Literal[ "completed", @@ -373,8 +374,8 @@ class SessionEdges: ) error_policy: ErrorPolicy = field(default_factory=DefaultErrorPolicy) transport: TransportService = field(default_factory=NoopTransportService) - clock: object | None = None - activation: object | None = None + clock: "RealtimeClock | DeterministicClock | None" = None + activation: "ActivationPolicy | None" = None _closed_result: RunResult | None = field(default=None, init=False, repr=False) @property diff --git a/flashdreams/flashdreams/runtime/demo/session_inputs.py b/flashdreams/flashdreams/runtime/demo/session_inputs.py index 47bbbad6e..b724c0d03 100644 --- a/flashdreams/flashdreams/runtime/demo/session_inputs.py +++ b/flashdreams/flashdreams/runtime/demo/session_inputs.py @@ -8,12 +8,15 @@ import math from collections.abc import Mapping, Sequence from dataclasses import dataclass, field -from typing import Protocol, runtime_checkable +from typing import TYPE_CHECKING, Protocol, runtime_checkable from flashdreams.runtime._utils import freeze_mapping from flashdreams.runtime.inputs import InferenceInput, UserInputs from flashdreams.runtime.types import StepRequest +if TYPE_CHECKING: + from .timing import RealtimeClock, RealtimeWindowResult + @dataclass(frozen=True, kw_only=True, slots=True) class ControlDecision: @@ -102,8 +105,8 @@ async def next_realtime_window( self, *, request: StepRequest, - clock: object, - ) -> object: + clock: "RealtimeClock", + ) -> "RealtimeWindowResult": """Return the next realtime window result. The concrete realtime result shape lands with the realtime clock phase. diff --git a/flashdreams/flashdreams/runtime/demo/timing.py b/flashdreams/flashdreams/runtime/demo/timing.py new file mode 100644 index 000000000..0261e60fb --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/timing.py @@ -0,0 +1,371 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Realtime activation, clock, and input-window primitives for demo run modes.""" + +from __future__ import annotations + +import asyncio +import math +import time +from collections.abc import Awaitable, Callable, Sequence +from dataclasses import dataclass, field +from typing import Literal, Protocol, runtime_checkable + +from flashdreams.runtime.inputs import UserInputs +from flashdreams.runtime.types import StepRequest +from flashdreams.serving.realtime.input import KeyboardResampler + +from .session_inputs import UserInputWindow + +CatchUpPolicy = Literal["drop", "fold", "compress"] + +SPARSE_KEY_SEGMENTS_METADATA_KEY = "sparse_key_segments" + + +@dataclass(frozen=True, kw_only=True, slots=True) +class CatchUpDecision: + """How a realtime clock bounded stale virtual input time.""" + + skipped_s: float = 0.0 + skipped_windows: int = 0 + input_policy: CatchUpPolicy | None = None + reason: str | None = None + + def __post_init__(self) -> None: + if not math.isfinite(self.skipped_s) or self.skipped_s < 0.0: + raise ValueError("CatchUpDecision.skipped_s must be finite and >= 0.") + if self.skipped_windows < 0: + raise ValueError("CatchUpDecision.skipped_windows must be >= 0.") + if self.input_policy not in {None, "drop", "fold", "compress"}: + raise ValueError( + f"Unsupported catch-up input_policy={self.input_policy!r}." + ) + if self.reason is not None and not self.reason.strip(): + raise ValueError("CatchUpDecision.reason must be non-empty when set.") + + +@dataclass(frozen=True, kw_only=True, slots=True) +class RealtimeWindowResult: + """Realtime input window plus any catch-up decision that preceded it.""" + + window: UserInputWindow + catch_up: CatchUpDecision = field(default_factory=CatchUpDecision) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class ActivationResult: + """Result of waiting for a realtime activation gate.""" + + activated: bool + reason: str | None = None + + def __post_init__(self) -> None: + if self.reason is not None and not self.reason.strip(): + raise ValueError("ActivationResult.reason must be non-empty when set.") + + +@runtime_checkable +class DeterministicClock(Protocol): + """Clock facts for finite deterministic run modes.""" + + is_realtime: bool + is_deterministic: bool + + +@runtime_checkable +class RealtimeClock(Protocol): + """Realtime virtual clock used by realtime drivers and input sources.""" + + is_realtime: bool + is_deterministic: bool + + def now(self) -> float: ... + + def anchor(self, wall_time_s: float) -> None: ... + + async def wait_until_window_end(self, end_s: float) -> None: ... + + async def apply_backpressure(self, requested_s: float) -> None: ... + + def catch_up( + self, + *, + request: StepRequest, + max_lag_s: float, + policy: CatchUpPolicy, + ) -> CatchUpDecision: ... + + +@runtime_checkable +class ActivationPolicy(Protocol): + """Wait until a realtime session should start generating.""" + + timeout_s: float | None + + async def wait_until_active( + self, + clock: RealtimeClock | DeterministicClock, + ) -> ActivationResult: ... + + +@runtime_checkable +class ActivationSignal(Protocol): + """Event-like object accepted by ``SignalActivationPolicy``.""" + + def is_set(self) -> bool: ... + + async def wait(self) -> object: ... + + +@dataclass(slots=True) +class AlwaysActiveActivationPolicy: + """Activation policy for batch/null modes or already-ready realtime modes.""" + + timeout_s: float | None = None + anchor_clock: bool = False + + async def wait_until_active( + self, + clock: RealtimeClock | DeterministicClock, + ) -> ActivationResult: + _anchor_if_realtime(clock, anchor=self.anchor_clock) + return ActivationResult(activated=True) + + +@dataclass(slots=True) +class SignalActivationPolicy: + """Activate when any supplied signal fires, with optional timeout.""" + + signals: Sequence[ActivationSignal] + timeout_s: float | None = None + timeout_reason: str = "activation timed out" + anchor_clock: bool = True + + def __post_init__(self) -> None: + if not self.signals: + raise ValueError("SignalActivationPolicy.signals must be non-empty.") + self.signals = tuple(self.signals) + if self.timeout_s is not None and self.timeout_s <= 0.0: + raise ValueError("SignalActivationPolicy.timeout_s must be > 0 when set.") + if not self.timeout_reason.strip(): + raise ValueError("SignalActivationPolicy.timeout_reason must be non-empty.") + + async def wait_until_active( + self, + clock: RealtimeClock | DeterministicClock, + ) -> ActivationResult: + if any(signal.is_set() for signal in self.signals): + _anchor_if_realtime(clock, anchor=self.anchor_clock) + return ActivationResult(activated=True) + + tasks = [asyncio.create_task(signal.wait()) for signal in self.signals] + try: + done, pending = await asyncio.wait( + tasks, + timeout=self.timeout_s, + return_when=asyncio.FIRST_COMPLETED, + ) + if not done: + return ActivationResult( + activated=False, + reason=self.timeout_reason, + ) + for task in done: + task.result() + _anchor_if_realtime(clock, anchor=self.anchor_clock) + return ActivationResult(activated=True) + finally: + for task in tasks: + if not task.done(): + task.cancel() + await asyncio.gather(*tasks, return_exceptions=True) + + +class _RealtimeTimeline(Protocol): + dt: float + next_chunk_start_v: float + + +@dataclass(slots=True) +class ResamplerRealtimeClock: + """Realtime clock that reuses ``KeyboardResampler``'s virtual timeline.""" + + resampler: _RealtimeTimeline + now_fn: Callable[[], float] = time.monotonic + sleep_fn: Callable[[float], Awaitable[None]] = asyncio.sleep + is_realtime: bool = True + is_deterministic: bool = False + _pending_backpressure_s: float = field(default=0.0, init=False, repr=False) + + @property + def pending_backpressure_s(self) -> float: + return self._pending_backpressure_s + + def now(self) -> float: + return float(self.now_fn()) + + def anchor(self, wall_time_s: float) -> None: + if not math.isfinite(wall_time_s): + raise ValueError("wall_time_s must be finite.") + self.resampler.next_chunk_start_v = float(wall_time_s) + self._pending_backpressure_s = 0.0 + + async def wait_until_window_end(self, end_s: float) -> None: + if not math.isfinite(end_s): + raise ValueError("end_s must be finite.") + delay_s = float(end_s) - self.now() + if delay_s > 0.0: + await self.sleep_fn(delay_s) + + async def apply_backpressure(self, requested_s: float) -> None: + if not math.isfinite(requested_s) or requested_s < 0.0: + raise ValueError("requested_s must be finite and >= 0.") + self._pending_backpressure_s += float(requested_s) + + def catch_up( + self, + *, + request: StepRequest, + max_lag_s: float, + policy: CatchUpPolicy, + ) -> CatchUpDecision: + if policy != "fold": + raise NotImplementedError( + f"Catch-up policy {policy!r} has no existing resampler analog yet." + ) + if not math.isfinite(max_lag_s) or max_lag_s < 0.0: + raise ValueError("max_lag_s must be finite and >= 0.") + + input_frame_count = input_frame_count_from_request(request) + chunk_duration_s = input_frame_count * float(self.resampler.dt) + if chunk_duration_s <= 0.0: + raise ValueError("Realtime resampler dt must produce a positive window.") + + effective_now_s = self.now() + self._pending_backpressure_s + self._pending_backpressure_s = 0.0 + current_start_s = float(self.resampler.next_chunk_start_v) + lag_s = effective_now_s - (current_start_s + chunk_duration_s) + if lag_s <= max_lag_s: + return CatchUpDecision() + + latest_start_s = effective_now_s - chunk_duration_s + if latest_start_s <= current_start_s: + return CatchUpDecision() + + skipped_s = latest_start_s - current_start_s + skipped_windows = max(1, math.ceil(skipped_s / chunk_duration_s)) + self.resampler.next_chunk_start_v = latest_start_s + return CatchUpDecision( + skipped_s=skipped_s, + skipped_windows=skipped_windows, + input_policy=policy, + reason="lag exceeded max_lag_s", + ) + + +@dataclass(slots=True) +class KeyboardRealtimeInputSource: + """Realtime input source backed by the existing keyboard resampler.""" + + resampler: KeyboardResampler + max_lag_s: float | None = None + catch_up_policy: CatchUpPolicy = "fold" + is_finite: bool = False + is_deterministic: bool = False + + def __post_init__(self) -> None: + if self.max_lag_s is not None and ( + not math.isfinite(self.max_lag_s) or self.max_lag_s < 0.0 + ): + raise ValueError( + "KeyboardRealtimeInputSource.max_lag_s must be finite and >= 0." + ) + if self.catch_up_policy != "fold": + raise NotImplementedError( + f"Catch-up policy {self.catch_up_policy!r} has no existing " + "KeyboardResampler analog yet." + ) + + def is_finished(self) -> bool: + return False + + def on_edge(self, *, arrival_t: float, event: str, key: str) -> None: + self.resampler.on_edge(arrival_t=arrival_t, event=event, key=key) + + def reset(self, *, start_v: float) -> None: + self.resampler.reset(start_v=start_v) + + async def next_realtime_window( + self, + *, + request: StepRequest, + clock: RealtimeClock, + ) -> RealtimeWindowResult: + input_frame_count = input_frame_count_from_request(request) + chunk_duration_s = input_frame_count * self.resampler.dt + window_end_s = self.resampler.next_chunk_start_v + chunk_duration_s + await clock.wait_until_window_end(window_end_s) + catch_up = clock.catch_up( + request=request, + max_lag_s=self.max_lag_s + if self.max_lag_s is not None + else chunk_duration_s, + policy=self.catch_up_policy, + ) + start_s = self.resampler.next_chunk_start_v + segments, frame_times = self.resampler.sample_chunk(input_frame_count) + end_s = self.resampler.next_chunk_start_v + window = UserInputWindow( + start_s=start_s, + end_s=end_s, + frame_times=tuple(frame_times), + inputs=UserInputs(), + metadata={SPARSE_KEY_SEGMENTS_METADATA_KEY: tuple(segments)}, + ) + return RealtimeWindowResult(window=window, catch_up=catch_up) + + +def input_frame_count_from_request(request: StepRequest) -> int: + """Return the positive realtime input frame count declared on a request.""" + + value = request.metadata.get("input_frame_count") + if isinstance(value, bool) or not isinstance(value, int): + raise ValueError( + "StepRequest.metadata['input_frame_count'] must be an integer." + ) + parsed = value + if parsed <= 0: + raise ValueError("StepRequest.metadata['input_frame_count'] must be > 0.") + return parsed + + +def _anchor_if_realtime( + clock: RealtimeClock | DeterministicClock, + *, + anchor: bool, +) -> None: + if not anchor or not getattr(clock, "is_realtime", False): + return + now = getattr(clock, "now", None) + clock_anchor = getattr(clock, "anchor", None) + if callable(now) and callable(clock_anchor): + clock_anchor(float(now())) + + +__all__ = [ + "ActivationPolicy", + "ActivationResult", + "ActivationSignal", + "AlwaysActiveActivationPolicy", + "CatchUpDecision", + "CatchUpPolicy", + "DeterministicClock", + "KeyboardRealtimeInputSource", + "RealtimeClock", + "RealtimeWindowResult", + "ResamplerRealtimeClock", + "SPARSE_KEY_SEGMENTS_METADATA_KEY", + "SignalActivationPolicy", + "input_frame_count_from_request", +] diff --git a/flashdreams/tests/test_demo_runtime_timing.py b/flashdreams/tests/test_demo_runtime_timing.py new file mode 100644 index 000000000..02e151680 --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_timing.py @@ -0,0 +1,236 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +from typing import Any, cast + +import pytest + +from flashdreams.runtime import StepRequest +from flashdreams.runtime.demo import NullOutputSink, RunResult, SessionEdges +from flashdreams.runtime.demo.timing import ( + SPARSE_KEY_SEGMENTS_METADATA_KEY, + CatchUpDecision, + CatchUpPolicy, + KeyboardRealtimeInputSource, + ResamplerRealtimeClock, + SignalActivationPolicy, +) +from flashdreams.serving.realtime.input import KeyboardResampler + +pytestmark = pytest.mark.ci_cpu + + +@pytest.mark.asyncio +async def test_signal_activation_waits_for_first_input_and_anchors_clock() -> None: + event = asyncio.Event() + resampler = KeyboardResampler(fps=30.0, start_v=0.0) + clock = ResamplerRealtimeClock(resampler=resampler, now_fn=lambda: 12.0) + policy = SignalActivationPolicy(signals=(event,), timeout_s=1.0) + + wait_task = asyncio.create_task(policy.wait_until_active(clock)) + await asyncio.sleep(0) + + assert not wait_task.done() + + event.set() + result = await wait_task + + assert result.activated + assert result.reason is None + assert resampler.next_chunk_start_v == pytest.approx(12.0) + + +@pytest.mark.asyncio +async def test_activation_timeout_can_close_edges_as_not_activated() -> None: + event = asyncio.Event() + resampler = KeyboardResampler(fps=30.0, start_v=0.0) + clock = ResamplerRealtimeClock(resampler=resampler, now_fn=lambda: 12.0) + policy = SignalActivationPolicy( + signals=(event,), + timeout_s=0.001, + timeout_reason="no first input", + ) + cleanup_tasks: set[asyncio.Task[RunResult]] = set() + edges = SessionEdges( + input_source=_OpenRealtimeInputSource(), + output_sink=NullOutputSink(), + cleanup_tasks=cleanup_tasks, + activation=policy, + clock=clock, + ) + + activation = await policy.wait_until_active(clock) + result = edges.close_result( + status="not_activated", + reason=activation.reason, + ) + + assert not activation.activated + assert activation.reason == "no first input" + assert result.status == "not_activated" + assert result.reason == "no first input" + assert edges.is_closed + assert resampler.next_chunk_start_v == pytest.approx(0.0) + + +def test_resampler_clock_catch_up_bounds_latency() -> None: + resampler = KeyboardResampler(fps=1.0, start_v=0.0) + clock = ResamplerRealtimeClock(resampler=resampler, now_fn=lambda: 5.0) + + decision = clock.catch_up( + request=_request(input_frame_count=1), + max_lag_s=1.0, + policy="fold", + ) + + assert decision == CatchUpDecision( + skipped_s=4.0, + skipped_windows=4, + input_policy="fold", + reason="lag exceeded max_lag_s", + ) + assert resampler.next_chunk_start_v == pytest.approx(4.0) + + +@pytest.mark.asyncio +async def test_realtime_input_source_matches_resampler_for_recorded_trace() -> None: + expected_resampler = _resampler_with_recorded_trace() + expected_resampler.next_chunk_start_v = 2.0 + expected_segments, expected_frame_times = expected_resampler.sample_chunk(2) + source_resampler = _resampler_with_recorded_trace() + sleep = _RecordingSleep() + clock = ResamplerRealtimeClock( + resampler=source_resampler, + now_fn=lambda: 3.0, + sleep_fn=sleep, + ) + source = KeyboardRealtimeInputSource(resampler=source_resampler) + + result = await source.next_realtime_window( + request=_request(input_frame_count=2), + clock=clock, + ) + + assert sleep.delays == [] + assert result.catch_up == CatchUpDecision( + skipped_s=2.0, + skipped_windows=2, + input_policy="fold", + reason="lag exceeded max_lag_s", + ) + assert result.window.start_s == pytest.approx(2.0) + assert result.window.end_s == pytest.approx(3.0) + assert result.window.frame_times == tuple(expected_frame_times) + assert result.window.metadata[SPARSE_KEY_SEGMENTS_METADATA_KEY] == tuple( + expected_segments + ) + + +@pytest.mark.asyncio +async def test_backpressure_is_clock_adjustment_not_blocking_sleep() -> None: + resampler = KeyboardResampler(fps=1.0, start_v=0.0) + sleep = _RecordingSleep() + clock = ResamplerRealtimeClock( + resampler=resampler, + now_fn=lambda: 2.2, + sleep_fn=sleep, + ) + + await clock.apply_backpressure(0.3) + decision = clock.catch_up( + request=_request(input_frame_count=1), + max_lag_s=1.0, + policy="fold", + ) + + assert sleep.delays == [] + assert clock.pending_backpressure_s == pytest.approx(0.0) + assert decision.skipped_s == pytest.approx(1.5) + assert decision.skipped_windows == 2 + assert decision.input_policy == "fold" + assert resampler.next_chunk_start_v == pytest.approx(1.5) + + +@pytest.mark.asyncio +async def test_window_floor_sleeps_only_when_virtual_time_is_ahead() -> None: + resampler = KeyboardResampler(fps=1.0, start_v=0.0) + sleep = _RecordingSleep() + clock = ResamplerRealtimeClock( + resampler=resampler, + now_fn=lambda: 1.0, + sleep_fn=sleep, + ) + + await clock.wait_until_window_end(1.25) + await clock.wait_until_window_end(0.75) + + assert sleep.delays == [0.25] + + +@pytest.mark.parametrize("policy", ["drop", "compress"]) +def test_keyboard_resampler_defers_unsupported_catch_up_policies( + policy: str, +) -> None: + resampler = KeyboardResampler(fps=1.0, start_v=0.0) + clock = ResamplerRealtimeClock(resampler=resampler, now_fn=lambda: 5.0) + unsupported_policy = cast(CatchUpPolicy, policy) + + with pytest.raises(NotImplementedError, match="no existing resampler analog"): + clock.catch_up( + request=_request(input_frame_count=1), + max_lag_s=1.0, + policy=unsupported_policy, + ) + + with pytest.raises(NotImplementedError, match="KeyboardResampler analog"): + KeyboardRealtimeInputSource( + resampler=resampler, + catch_up_policy=unsupported_policy, + ) + + +def _request(*, input_frame_count: int) -> StepRequest: + return StepRequest( + step_index=0, + metadata={"input_frame_count": input_frame_count}, + ) + + +def _resampler_with_recorded_trace() -> KeyboardResampler: + resampler = KeyboardResampler(fps=2.0, start_v=0.0) + for arrival_t, event, key in ( + (0.25, "keydown", "w"), + (1.25, "keydown", "a"), + (2.25, "keyup", "w"), + (2.75, "keydown", "d"), + ): + resampler.on_edge(arrival_t=arrival_t, event=event, key=key) + return resampler + + +class _RecordingSleep: + def __init__(self) -> None: + self.delays: list[float] = [] + + async def __call__(self, delay_s: float) -> None: + self.delays.append(delay_s) + + +class _OpenRealtimeInputSource: + is_finite = False + is_deterministic = False + + def is_finished(self) -> bool: + return False + + async def next_realtime_window( + self, + *, + request: StepRequest, + clock: Any, + ) -> object: + del request, clock + raise AssertionError("Activation timeout must close before requesting input.") From d1cb9146531e03285b41282ce95e43536fb98547 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 04:06:49 +0000 Subject: [PATCH 10/51] runtime: add realtime session driver --- .../flashdreams/runtime/demo/__init__.py | 8 + .../flashdreams/runtime/demo/drivers.py | 307 ++++++- .../test_demo_runtime_realtime_driver.py | 790 ++++++++++++++++++ 3 files changed, 1103 insertions(+), 2 deletions(-) create mode 100644 flashdreams/tests/test_demo_runtime_realtime_driver.py diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index 1637abfd3..1b9a01492 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -5,9 +5,13 @@ from flashdreams.runtime.demo.drivers import ( BatchSessionDriver, + CLEANUP_TIMEOUT_S, DriverInvariantError, + RealtimeSessionDriver, run_demo_session, run_demo_session_async, + shielded_session_cleanup, + uncancel_current_task, ) from flashdreams.runtime.demo.host import ( ModelWarmupPlan, @@ -78,6 +82,7 @@ __all__ = [ "BatchInputSource", "BatchSessionDriver", + "CLEANUP_TIMEOUT_S", "ControlDecision", "DefaultErrorPolicy", "DemoAdapter", @@ -109,6 +114,7 @@ "PreparedStep", "RealtimeInputSource", "RealtimeClock", + "RealtimeSessionDriver", "RealtimeWindowResult", "ResamplerRealtimeClock", "RunContext", @@ -134,4 +140,6 @@ "run_demo_session", "run_demo_session_async", "run_replay_demo", + "shielded_session_cleanup", + "uncancel_current_task", ] diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 326ba9692..f98f9b1a0 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -24,6 +24,9 @@ ) from .session_inputs import BatchInputSource, ModelInputProvider from .spec import DemoAdapter, DemoSpec, PreparedScenario +from .timing import ActivationPolicy, RealtimeClock + +CLEANUP_TIMEOUT_S = 30.0 class DriverInvariantError(RuntimeError): @@ -141,6 +144,161 @@ def run_one_session( ) +class RealtimeSessionDriver: + """Async realtime session driver built on shared Phase 5 primitives.""" + + cleanup_timeout_s: float + + def __init__(self, *, cleanup_timeout_s: float = CLEANUP_TIMEOUT_S) -> None: + if cleanup_timeout_s <= 0: + raise ValueError("cleanup_timeout_s must be > 0.") + self.cleanup_timeout_s = float(cleanup_timeout_s) + + async def run_one_session( + self, + *, + host: RuntimeHost, + provider: ModelInputProvider, + session_edges: SessionEdges, + pipeline: StepPipeline, + ) -> RunResult: + session: InferenceSession | None = None + final_status: DriverStatus = "completed" + final_reason: str | None = None + final_error: Exception | None = None + setup_ok = False + generation = 0 + first_step_started = False + invariant_error: DriverInvariantError | None = None + try: + activation, clock = _realtime_activation_and_clock(session_edges) + input_source = _realtime_input_source(session_edges) + activation_result = await activation.wait_until_active(clock) + if not activation_result.activated: + final_status = "not_activated" + final_reason = activation_result.reason + elif not session_edges.transport.is_active(): + final_status = "not_activated" + final_reason = "transport closed before first step" + else: + try: + initial_input = await host.call_async( + provider.prepare_initial_input + ) + session = await host.call_async(host.start_session, initial_input) + session_info = await host.call_async(_session_info, session) + session_edges.output_sink.open(session_info) + session_edges.output_sink.begin_generation(generation) + setup_ok = True + except Exception as exc: + action = session_edges.error_policy.handle_setup_error(exc) + if action.drop_chunk or action.result_status == "completed": + raise DriverInvariantError( + "Setup failures must resolve to failed or skipped." + ) from exc + session_edges.metrics.record_error(exc, action) + final_status = action.result_status + final_reason = str(exc) + final_error = exc if action.result_status == "failed" else None + + while setup_ok: + if session is None: + raise DriverInvariantError("setup_ok was set without a session.") + if not session_edges.transport.is_active(): + if not first_step_started: + final_status = "not_activated" + final_reason = "transport closed before first step" + break + try: + request = await host.call_async(session.next_step_request) + if request is None: + break + window_result = await input_source.next_realtime_window( + request=request, + clock=clock, + ) + if ( + not session_edges.transport.is_active() + and not first_step_started + ): + final_status = "not_activated" + final_reason = "transport closed before first step" + break + outcome = await host.call_async( + pipeline.execute_step, + request=request, + user_window=window_result.window, + provider=provider, + session=session, + output=session_edges.output_sink, + metrics=session_edges.metrics, + ) + first_step_started = True + if outcome.control.reset: + await host.call_async( + session.reset, + outcome.control.reset_input, + ) + if not outcome.control.provider_already_reset: + await host.call_async( + provider.reset, + outcome.control.reset_input, + ) + generation += 1 + session_edges.output_sink.begin_generation(generation) + continue + if outcome.control.close_session: + break + if outcome.output.should_stop: + break + if outcome.output.backpressure_s > 0: + await clock.apply_backpressure(outcome.output.backpressure_s) + except Exception as exc: + action = session_edges.error_policy.handle(exc) + session_edges.metrics.record_error(exc, action) + if action.close_session: + final_status = action.result_status + final_reason = str(exc) + final_error = exc if action.result_status == "failed" else None + break + if action.drop_chunk: + continue + final_status = "failed" + final_reason = str(exc) + final_error = exc + break + except asyncio.CancelledError: + uncancel_current_task() + final_status = "cancelled" + final_reason = ( + "cancelled before first step" if session is None else "cancelled" + ) + final_error = None + except DriverInvariantError as exc: + invariant_error = exc + final_status = "failed" + final_reason = str(exc) + final_error = exc + except Exception as exc: + final_status = "failed" + final_reason = str(exc) + final_error = exc + + result = await shielded_session_cleanup( + host=host, + session=session, + provider=provider, + session_edges=session_edges, + status=final_status, + reason=final_reason, + error=final_error, + timeout_s=self.cleanup_timeout_s, + ) + if invariant_error is not None: + raise invariant_error + return result + + def run_demo_session( *, context: RunContext, @@ -293,7 +451,7 @@ async def run_demo_session_async( context.run_metrics.record_session(result) return result except asyncio.CancelledError: - _uncancel_current_task() + uncancel_current_task() result = await _close_partial_session_async( context=context, provider=provider, @@ -370,6 +528,73 @@ def _close_on_host_best_effort( session_edges.record_cleanup_error(exc) +async def shielded_session_cleanup( + *, + host: RuntimeHost, + session: InferenceSession | None, + provider: ModelInputProvider, + session_edges: SessionEdges, + status: DriverStatus, + reason: str | None, + error: Exception | None, + timeout_s: float = CLEANUP_TIMEOUT_S, +) -> RunResult: + """Close realtime session resources exactly once without leaking cancellation.""" + + if timeout_s <= 0: + session_edges.record_cleanup_error(ValueError("timeout_s must be > 0.")) + return session_edges.close_result(status=status, reason=reason, error=error) + + async def cleanup() -> RunResult: + if session is not None: + session_closed = await _close_model_resource_async( + host=host, + close=session.close, + session_edges=session_edges, + timeout_s=timeout_s, + ) + if not session_closed: + host.mark_unhealthy("model-affine cleanup timed out") + return session_edges.close_result( + status=status, + reason=reason, + error=error, + ) + provider_closed = await _close_model_resource_async( + host=host, + close=provider.close, + session_edges=session_edges, + timeout_s=timeout_s, + ) + if not provider_closed: + host.mark_unhealthy("model-affine cleanup timed out") + return session_edges.close_result( + status=status, + reason=reason, + error=error, + ) + + try: + cleanup_task = asyncio.create_task(cleanup()) + except RuntimeError as exc: + session_edges.record_cleanup_error(exc) + return session_edges.close_result(status=status, reason=reason, error=error) + + session_edges.cleanup_tasks.add(cleanup_task) + try: + while not cleanup_task.done(): + try: + await asyncio.shield(cleanup_task) + except asyncio.CancelledError: + uncancel_current_task() + continue + except Exception: + break + return _cleanup_result(cleanup_task, session_edges, status, reason, error) + finally: + session_edges.cleanup_tasks.discard(cleanup_task) + + def _run_sync_driver( *, driver: object, @@ -442,6 +667,16 @@ async def _close_partial_session_async( error: Exception | None, close_provider: bool, ) -> RunResult: + if provider is not None and close_provider and session_edges is not None: + return await shielded_session_cleanup( + host=context.host, + session=None, + provider=provider, + session_edges=session_edges, + status=status, + reason=reason, + error=error, + ) if provider is not None and close_provider: await _close_provider_async( context=context, @@ -468,7 +703,71 @@ async def _close_provider_async( context.run_metrics.record_cleanup_error(close_exc) -def _uncancel_current_task() -> None: +async def _close_model_resource_async( + *, + host: RuntimeHost, + close: Any, + session_edges: SessionEdges, + timeout_s: float, +) -> bool: + try: + await asyncio.wait_for( + host.call_async(_close_safely, close, session_edges), + timeout=timeout_s, + ) + except asyncio.TimeoutError as exc: + session_edges.record_cleanup_error(exc) + return False + except Exception as exc: + session_edges.record_cleanup_error(exc) + return True + + +def _cleanup_result( + cleanup_task: asyncio.Task[RunResult], + session_edges: SessionEdges, + status: DriverStatus, + reason: str | None, + error: Exception | None, +) -> RunResult: + if cleanup_task.done() and not cleanup_task.cancelled(): + exc = cleanup_task.exception() + if exc is None: + return cleanup_task.result() + if isinstance(exc, Exception): + session_edges.record_cleanup_error(exc) + else: + session_edges.record_cleanup_error( + RuntimeError(f"cleanup failed with {type(exc).__name__}") + ) + return session_edges.close_result(status=status, reason=reason, error=error) + + +def _realtime_activation_and_clock( + session_edges: SessionEdges, +) -> tuple[ActivationPolicy, RealtimeClock]: + activation = session_edges.activation + if activation is None: + raise DriverInvariantError( + "RealtimeSessionDriver requires SessionEdges.activation." + ) + clock = session_edges.clock + if not isinstance(clock, RealtimeClock): + raise DriverInvariantError("RealtimeSessionDriver requires a RealtimeClock.") + return activation, clock + + +def _realtime_input_source(session_edges: SessionEdges) -> Any: + input_source = session_edges.input_source + next_realtime_window = getattr(input_source, "next_realtime_window", None) + if not callable(next_realtime_window): + raise DriverInvariantError( + "RealtimeSessionDriver requires a RealtimeInputSource." + ) + return input_source + + +def uncancel_current_task() -> None: task = asyncio.current_task() if task is None: return @@ -484,7 +783,11 @@ def _uncancel_current_task() -> None: __all__ = [ "BatchSessionDriver", + "CLEANUP_TIMEOUT_S", "DriverInvariantError", + "RealtimeSessionDriver", "run_demo_session", "run_demo_session_async", + "shielded_session_cleanup", + "uncancel_current_task", ] diff --git a/flashdreams/tests/test_demo_runtime_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py new file mode 100644 index 000000000..93bff4f9d --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -0,0 +1,790 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +import time +from collections.abc import Callable, Sequence +from typing import Any, Literal, cast + +import pytest + +from flashdreams.runtime import ( + InferenceInput, + InferenceRuntime, + StepRequest, + StepResult, +) +from flashdreams.runtime.demo import ( + ActivationResult, + ErrorAction, + InMemorySessionMetricsRecorder, + OutputDecision, + PreparedStep, + RealtimeSessionDriver, + RealtimeWindowResult, + RunContext, + RunResult, + RuntimeHost, + SessionEdges, + SessionInfo, + SingleSessionAdmissionPolicy, + StepPipeline, + UserInputWindow, + DriverInvariantError, + shielded_session_cleanup, +) + +pytestmark = pytest.mark.ci_cpu + + +@pytest.mark.asyncio +async def test_realtime_driver_non_activation_returns_not_activated() -> None: + runtime = _FakeRealtimeRuntime(session=_FakeRealtimeSession(num_steps=1)) + host = RuntimeHost(runtime) + provider = _FakeRealtimeProvider() + output = _RecordingOutputSink() + transport = _RecordingTransport() + metrics = InMemorySessionMetricsRecorder() + edges = _edges( + input_source=_RealtimeInputSource(), + output=output, + transport=transport, + metrics=metrics, + activation=_ActivationPolicy(ActivationResult(activated=False, reason="idle")), + ) + + try: + result = await RealtimeSessionDriver().run_one_session( + host=host, + provider=provider, + session_edges=edges, + pipeline=StepPipeline(), + ) + finally: + host.close() + + assert result.status == "not_activated" + assert result.reason == "idle" + assert runtime.start_session_inputs == [] + assert provider.prepare_initial_count == 0 + assert provider.close_count == 1 + assert output.close_count == 1 + assert transport.close_count == 1 + assert metrics.closed + + +@pytest.mark.asyncio +async def test_realtime_driver_transport_close_before_first_step_is_not_activated() -> ( + None +): + session = _FakeRealtimeSession(num_steps=1) + runtime = _FakeRealtimeRuntime(session=session) + host = RuntimeHost(runtime) + transport = _RecordingTransport() + edges = _edges( + input_source=_RealtimeInputSource(transport_to_close=transport), + transport=transport, + ) + + try: + result = await RealtimeSessionDriver().run_one_session( + host=host, + provider=_FakeRealtimeProvider(), + session_edges=edges, + pipeline=StepPipeline(), + ) + finally: + host.close() + + assert result.status == "not_activated" + assert result.reason == "transport closed before first step" + assert session.step_inputs == [] + + +@pytest.mark.asyncio +async def test_repeated_cancellation_during_cleanup_still_closes_edges() -> None: + entered_window = asyncio.Event() + session = _FakeRealtimeSession(num_steps=1) + runtime = _FakeRealtimeRuntime(session=session) + host = RuntimeHost(runtime) + provider = _FakeRealtimeProvider(close_delay_s=0.05) + output = _RecordingOutputSink() + transport = _RecordingTransport() + metrics = InMemorySessionMetricsRecorder() + edges = _edges( + input_source=_RealtimeInputSource( + entered=entered_window, + wait_forever=True, + ), + output=output, + transport=transport, + metrics=metrics, + ) + task = asyncio.create_task( + RealtimeSessionDriver().run_one_session( + host=host, + provider=provider, + session_edges=edges, + pipeline=StepPipeline(), + ) + ) + await entered_window.wait() + + task.cancel() + await asyncio.sleep(0) + task.cancel() + result = await task + host.close() + + assert result.status == "cancelled" + assert result.reason == "cancelled" + assert session.close_count == 1 + assert provider.close_count == 1 + assert output.close_count == 1 + assert transport.close_count == 1 + assert metrics.closed + assert not edges.cleanup_tasks + + +@pytest.mark.asyncio +async def test_cancelled_realtime_driver_inside_timeout_returns_result() -> None: + timeout_context = getattr(asyncio, "timeout", None) + if timeout_context is None: + pytest.skip("asyncio.timeout is unavailable on this Python version.") + entered_window = asyncio.Event() + runtime = _FakeRealtimeRuntime(session=_FakeRealtimeSession(num_steps=1)) + host = RuntimeHost(runtime) + edges = _edges( + input_source=_RealtimeInputSource( + entered=entered_window, + wait_forever=True, + ) + ) + + try: + async with timeout_context(0.01): + result = await RealtimeSessionDriver().run_one_session( + host=host, + provider=_FakeRealtimeProvider(), + session_edges=edges, + pipeline=StepPipeline(), + ) + finally: + host.close() + + assert entered_window.is_set() + assert result.status == "cancelled" + + +@pytest.mark.asyncio +async def test_realtime_driver_invariant_finalizes_edges_before_reraising() -> None: + runtime = _FakeRealtimeRuntime(session=_FakeRealtimeSession(num_steps=1)) + host = RuntimeHost(runtime) + provider = _FakeRealtimeProvider(fail_initial=RuntimeError("bad setup policy")) + output = _RecordingOutputSink() + transport = _RecordingTransport() + metrics = InMemorySessionMetricsRecorder() + edges = _edges( + output=output, + transport=transport, + metrics=metrics, + error_policy=_SetupPolicy(result_status="completed"), + ) + + try: + with pytest.raises(DriverInvariantError, match="Setup failures") as raised: + await RealtimeSessionDriver().run_one_session( + host=host, + provider=provider, + session_edges=edges, + pipeline=StepPipeline(), + ) + finally: + host.close() + + result = edges.close_result() + assert result.status == "failed" + assert result.error is raised.value + assert provider.close_count == 1 + assert output.close_count == 1 + assert transport.close_count == 1 + assert metrics.closed + + +@pytest.mark.asyncio +async def test_run_context_close_async_drains_registered_cleanup_task() -> None: + runtime = _FakeRealtimeRuntime(session=_FakeRealtimeSession(num_steps=1)) + host = RuntimeHost(runtime) + context = RunContext( + host=host, + run_metrics=InMemorySessionMetricsRecorder(), + admission=SingleSessionAdmissionPolicy(), + ) + edges = _edges(cleanup_tasks=context.cleanup_tasks) + cleanup_task = asyncio.create_task( + shielded_session_cleanup( + host=host, + session=runtime.session, + provider=_FakeRealtimeProvider(close_delay_s=0.02), + session_edges=edges, + status="cancelled", + reason="test", + error=None, + ) + ) + await asyncio.sleep(0) + + summary = await context.close_async() + cleanup_result = await cleanup_task + + assert cleanup_task.done() + assert cleanup_result.status == "cancelled" + assert not context.cleanup_tasks + assert edges.is_closed + assert summary.metrics.counters["sessions"] == 0 + + +@pytest.mark.asyncio +async def test_shielded_cleanup_never_raises_and_returns_result_on_close_errors() -> ( + None +): + runtime = _FakeRealtimeRuntime( + session=_FakeRealtimeSession(num_steps=1, fail_close=RuntimeError("session")) + ) + host = RuntimeHost(runtime) + provider = _FakeRealtimeProvider(fail_close=RuntimeError("provider")) + metrics = InMemorySessionMetricsRecorder() + edges = _edges(metrics=metrics) + + try: + result = await shielded_session_cleanup( + host=host, + session=runtime.session, + provider=provider, + session_edges=edges, + status="failed", + reason="test failure", + error=RuntimeError("original"), + ) + finally: + host.close() + + assert result.status == "failed" + assert result.reason == "test failure" + assert metrics.closed + assert metrics.cleanup_errors == ["session", "provider"] + + +@pytest.mark.asyncio +async def test_shielded_cleanup_timeout_bounds_shutdown() -> None: + host = _NeverReturningHost() + metrics = InMemorySessionMetricsRecorder() + edges = _edges(metrics=metrics) + + result = await shielded_session_cleanup( + host=cast(RuntimeHost, host), + session=None, + provider=_FakeRealtimeProvider(), + session_edges=edges, + status="cancelled", + reason="timeout test", + error=None, + timeout_s=0.001, + ) + + assert result.status == "cancelled" + assert host.unhealthy_reason == "model-affine cleanup timed out" + assert len(metrics.cleanup_errors) == 1 + assert metrics.closed + + +@pytest.mark.asyncio +async def test_realtime_driver_applies_backpressure_through_clock() -> None: + session = _FakeRealtimeSession(num_steps=2) + runtime = _FakeRealtimeRuntime(session=session) + host = RuntimeHost(runtime) + clock = _RecordingRealtimeClock() + output = _RecordingOutputSink( + decisions=( + OutputDecision(backpressure_s=0.25), + OutputDecision(should_stop=True), + ) + ) + edges = _edges(clock=clock, output=output) + + try: + result = await RealtimeSessionDriver().run_one_session( + host=host, + provider=_FakeRealtimeProvider(), + session_edges=edges, + pipeline=StepPipeline(), + ) + finally: + host.close() + + assert result.status == "completed" + assert clock.backpressure == [0.25] + assert len(output.results) == 2 + + +@pytest.mark.asyncio +async def test_realtime_driver_calls_step_pipeline_on_runtime_host() -> None: + session = _FakeRealtimeSession(num_steps=1) + runtime = _FakeRealtimeRuntime(session=session) + host = _RecordingRuntimeHost(runtime) + edges = _edges( + output=_RecordingOutputSink(decisions=(OutputDecision(should_stop=True),)) + ) + + try: + result = await RealtimeSessionDriver().run_one_session( + host=host, + provider=_FakeRealtimeProvider(), + session_edges=edges, + pipeline=StepPipeline(), + ) + finally: + host.close() + + assert result.status == "completed" + assert "execute_step" in host.async_calls + assert "prepare_step" not in host.async_calls + assert "step" not in host.async_calls + + +@pytest.mark.asyncio +async def test_slow_fake_model_step_does_not_block_event_loop() -> None: + session = _FakeRealtimeSession(num_steps=1, step_delay_s=0.05) + runtime = _FakeRealtimeRuntime(session=session) + host = RuntimeHost(runtime) + edges = _edges( + output=_RecordingOutputSink(decisions=(OutputDecision(should_stop=True),)) + ) + ticks = 0 + finished = False + + async def heartbeat() -> None: + nonlocal ticks + while not finished: + ticks += 1 + await asyncio.sleep(0.005) + + heartbeat_task = asyncio.create_task(heartbeat()) + try: + result = await RealtimeSessionDriver().run_one_session( + host=host, + provider=_FakeRealtimeProvider(), + session_edges=edges, + pipeline=StepPipeline(), + ) + finally: + finished = True + await heartbeat_task + host.close() + + assert result.status == "completed" + assert ticks >= 2 + + +@pytest.mark.asyncio +async def test_realtime_driver_fatal_model_error_returns_failed() -> None: + session = _FakeRealtimeSession(num_steps=1, fail_step=0) + runtime = _FakeRealtimeRuntime(session=session) + host = RuntimeHost(runtime) + metrics = InMemorySessionMetricsRecorder() + edges = _edges(metrics=metrics) + + try: + result = await RealtimeSessionDriver().run_one_session( + host=host, + provider=_FakeRealtimeProvider(), + session_edges=edges, + pipeline=StepPipeline(), + ) + finally: + host.close() + + assert result.status == "failed" + assert result.reason == "step failed" + assert isinstance(result.error, RuntimeError) + assert session.close_count == 1 + assert metrics.errors == ["step failed"] + + +@pytest.mark.asyncio +async def test_realtime_driver_can_drop_recoverable_output_error() -> None: + session = _FakeRealtimeSession(num_steps=2) + runtime = _FakeRealtimeRuntime(session=session) + host = RuntimeHost(runtime) + output = _RecordingOutputSink( + fail_first_write=RuntimeError("output queue full"), + decisions=(OutputDecision(should_stop=True),), + ) + metrics = InMemorySessionMetricsRecorder() + edges = _edges( + output=output, + metrics=metrics, + error_policy=_DropOutputErrorPolicy(), + ) + + try: + result = await RealtimeSessionDriver().run_one_session( + host=host, + provider=_FakeRealtimeProvider(), + session_edges=edges, + pipeline=StepPipeline(), + ) + finally: + host.close() + + assert result.status == "completed" + assert metrics.errors == ["output queue full"] + assert [step.step_index for step in output.results] == [1] + assert len(session.step_inputs) == 2 + + +def _edges( + *, + input_source: "_RealtimeInputSource | None" = None, + output: "_RecordingOutputSink | None" = None, + transport: "_RecordingTransport | None" = None, + metrics: InMemorySessionMetricsRecorder | None = None, + activation: "_ActivationPolicy | None" = None, + clock: "_RecordingRealtimeClock | None" = None, + cleanup_tasks: set[asyncio.Task[RunResult]] | None = None, + error_policy: Any | None = None, +) -> SessionEdges: + return SessionEdges( + input_source=input_source or _RealtimeInputSource(), + output_sink=output + or _RecordingOutputSink(decisions=(OutputDecision(should_stop=True),)), + cleanup_tasks=cleanup_tasks or set(), + metrics=metrics or InMemorySessionMetricsRecorder(), + error_policy=error_policy or _DefaultTestErrorPolicy(), + transport=transport or _RecordingTransport(), + clock=clock or _RecordingRealtimeClock(), + activation=activation or _ActivationPolicy(ActivationResult(activated=True)), + ) + + +def _window(index: int) -> UserInputWindow: + start_s = float(index) + return UserInputWindow( + start_s=start_s, + end_s=start_s + 1.0, + frame_times=(start_s + 1.0,), + ) + + +class _ActivationPolicy: + timeout_s: float | None = None + + def __init__(self, result: ActivationResult) -> None: + self.result = result + self.calls = 0 + + async def wait_until_active(self, clock: Any) -> ActivationResult: + del clock + self.calls += 1 + await asyncio.sleep(0) + return self.result + + +class _RecordingRealtimeClock: + is_realtime = True + is_deterministic = False + + def __init__(self) -> None: + self.backpressure: list[float] = [] + self.anchors: list[float] = [] + + def now(self) -> float: + return 0.0 + + def anchor(self, wall_time_s: float) -> None: + self.anchors.append(wall_time_s) + + async def wait_until_window_end(self, end_s: float) -> None: + del end_s + + async def apply_backpressure(self, requested_s: float) -> None: + self.backpressure.append(requested_s) + await asyncio.sleep(0) + + def catch_up(self, **kwargs: Any) -> object: + del kwargs + return object() + + +class _RealtimeInputSource: + is_finite = False + is_deterministic = False + + def __init__( + self, + *, + entered: asyncio.Event | None = None, + wait_forever: bool = False, + transport_to_close: "_RecordingTransport | None" = None, + ) -> None: + self.entered = entered + self.wait_forever = wait_forever + self.transport_to_close = transport_to_close + self.requests: list[StepRequest] = [] + + def is_finished(self) -> bool: + return False + + async def next_realtime_window( + self, + *, + request: StepRequest, + clock: Any, + ) -> RealtimeWindowResult: + del clock + self.requests.append(request) + if self.entered is not None: + self.entered.set() + if self.wait_forever: + await asyncio.Event().wait() + if self.transport_to_close is not None: + self.transport_to_close.close() + return RealtimeWindowResult(window=_window(request.step_index)) + + +class _FakeRealtimeProvider: + def __init__( + self, + *, + fail_initial: Exception | None = None, + fail_close: Exception | None = None, + close_delay_s: float = 0.0, + ) -> None: + self.fail_initial = fail_initial + self.fail_close = fail_close + self.close_delay_s = close_delay_s + self.prepare_initial_count = 0 + self.close_count = 0 + self.reset_inputs: list[InferenceInput | None] = [] + + def prepare_initial_input(self) -> InferenceInput: + if self.fail_initial is not None: + raise self.fail_initial + self.prepare_initial_count += 1 + return InferenceInput(global_conditioning={"prompt": "realtime"}) + + def prepare_step( + self, + *, + request: StepRequest, + user_window: UserInputWindow, + ) -> PreparedStep: + return PreparedStep( + inference_input=InferenceInput( + step={ + "request_step": request.step_index, + "window": (user_window.start_s, user_window.end_s), + } + ) + ) + + def reset(self, inputs: InferenceInput | None = None) -> None: + self.reset_inputs.append(inputs) + + def close(self) -> None: + if self.close_delay_s: + time.sleep(self.close_delay_s) + self.close_count += 1 + if self.fail_close is not None: + raise self.fail_close + + +class _FakeRealtimeRuntime: + def __init__(self, *, session: "_FakeRealtimeSession") -> None: + self.session = session + self.start_session_inputs: list[InferenceInput] = [] + self.close_count = 0 + + def start_session(self, inputs: InferenceInput) -> "_FakeRealtimeSession": + self.start_session_inputs.append(inputs) + return self.session + + def close(self) -> None: + self.close_count += 1 + + +class _FakeRealtimeSession: + def __init__( + self, + *, + num_steps: int, + fail_step: int | None = None, + fail_close: Exception | None = None, + step_delay_s: float = 0.0, + ) -> None: + self.num_steps = num_steps + self.fail_step = fail_step + self.fail_close = fail_close + self.step_delay_s = step_delay_s + self.next_request_index = 0 + self.step_inputs: list[InferenceInput] = [] + self.close_count = 0 + + def session_info(self) -> SessionInfo: + return SessionInfo(output_layout="fake-realtime", steady_output_frame_count=1) + + def next_step_request(self) -> StepRequest | None: + if self.next_request_index >= self.num_steps: + return None + request = StepRequest(step_index=self.next_request_index) + self.next_request_index += 1 + return request + + def step(self, inputs: InferenceInput) -> StepResult: + step_index = len(self.step_inputs) + if self.step_delay_s: + time.sleep(self.step_delay_s) + if self.fail_step == step_index: + raise RuntimeError("step failed") + self.step_inputs.append(inputs) + return StepResult( + step_index=step_index, + output=f"frame-{step_index}", + frame_count=1, + ) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self.next_request_index = 0 + self.step_inputs.clear() + + def close(self) -> None: + self.close_count += 1 + if self.fail_close is not None: + raise self.fail_close + + +class _RecordingRuntimeHost(RuntimeHost): + def __init__(self, runtime: InferenceRuntime) -> None: + super().__init__(runtime) + self.async_calls: list[str] = [] + + async def call_async( + self, + func: Callable[..., Any], + /, + *args: object, + **kwargs: object, + ) -> Any: + self.async_calls.append(getattr(func, "__name__", type(func).__name__)) + return await super().call_async(func, *args, **kwargs) + + +class _RecordingOutputSink: + produces_artifacts = False + + def __init__( + self, + *, + decisions: Sequence[OutputDecision] = (), + fail_first_write: Exception | None = None, + ) -> None: + self.decisions = list(decisions) + self.fail_first_write = fail_first_write + self.opened_with: SessionInfo | None = None + self.generations: list[int] = [] + self.results: list[StepResult] = [] + self.close_count = 0 + self.write_attempts = 0 + + def open(self, session_info: SessionInfo) -> None: + self.opened_with = session_info + + def begin_generation(self, generation: int) -> None: + self.generations.append(generation) + + def write(self, result: StepResult) -> OutputDecision: + self.write_attempts += 1 + if self.fail_first_write is not None: + exc = self.fail_first_write + self.fail_first_write = None + raise exc + self.results.append(result) + if self.decisions: + return self.decisions.pop(0) + return OutputDecision() + + def close(self) -> Sequence[Any]: + self.close_count += 1 + return () + + +class _RecordingTransport: + def __init__(self) -> None: + self.active = True + self.close_count = 0 + + def is_active(self) -> bool: + return self.active + + def close(self) -> None: + self.active = False + self.close_count += 1 + + +class _DefaultTestErrorPolicy: + def handle_setup_error(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status="failed") + + def handle(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status="failed") + + +class _SetupPolicy(_DefaultTestErrorPolicy): + def __init__( + self, + *, + result_status: Literal["completed", "failed", "skipped"], + ) -> None: + self.result_status = result_status + + def handle_setup_error(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status=self.result_status) + + +class _DropOutputErrorPolicy(_DefaultTestErrorPolicy): + def handle(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction( + close_session=False, + drop_chunk=True, + result_status="failed", + ) + + +class _NeverReturningHost: + def __init__(self) -> None: + self.unhealthy_reason: str | None = None + + async def call_async( + self, + func: Callable[..., Any], + /, + *args: object, + **kwargs: object, + ) -> Any: + del func, args, kwargs + await asyncio.Event().wait() + + def mark_unhealthy( + self, + reason: str = "marked unhealthy", + error: Exception | None = None, + ) -> None: + del error + self.unhealthy_reason = reason From 3b08bd3f5700dbd12639dfe5be861916e87b455f Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 04:17:48 +0000 Subject: [PATCH 11/51] runtime: clarify unavailable host cleanup --- flashdreams/flashdreams/runtime/demo/__init__.py | 2 +- flashdreams/flashdreams/runtime/demo/drivers.py | 3 +++ flashdreams/tests/test_demo_runtime_realtime_driver.py | 2 +- 3 files changed, 5 insertions(+), 2 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index 1b9a01492..1a3f76892 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -4,8 +4,8 @@ """Experimental shared demo API above the inference runtime API.""" from flashdreams.runtime.demo.drivers import ( - BatchSessionDriver, CLEANUP_TIMEOUT_S, + BatchSessionDriver, DriverInvariantError, RealtimeSessionDriver, run_demo_session, diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index f98f9b1a0..729e610ca 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -525,6 +525,9 @@ def _close_on_host_best_effort( try: host.call(_close_safely, close, session_edges) except Exception as exc: + # If the host/worker is already unavailable, do not fall back to calling + # model-affine cleanup directly on the caller thread. Record the loss and + # let close_result finalize output, transport, and metrics. session_edges.record_cleanup_error(exc) diff --git a/flashdreams/tests/test_demo_runtime_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py index 93bff4f9d..7cb813bda 100644 --- a/flashdreams/tests/test_demo_runtime_realtime_driver.py +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -18,6 +18,7 @@ ) from flashdreams.runtime.demo import ( ActivationResult, + DriverInvariantError, ErrorAction, InMemorySessionMetricsRecorder, OutputDecision, @@ -32,7 +33,6 @@ SingleSessionAdmissionPolicy, StepPipeline, UserInputWindow, - DriverInvariantError, shielded_session_cleanup, ) From b856f48191439f7be91f0b21d861eb5bb9d2d96f Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 04:27:00 +0000 Subject: [PATCH 12/51] runtime: harden batch cleanup finalization --- .../flashdreams/runtime/demo/drivers.py | 12 ++++- .../tests/test_demo_runtime_vertical_slice.py | 45 +++++++++++++++++++ 2 files changed, 55 insertions(+), 2 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 729e610ca..5728f46dc 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -134,8 +134,16 @@ def run_one_session( finally: if not invariant_closed: if session is not None: - host.call(_close_safely, session.close, session_edges) - host.call(_close_safely, provider.close, session_edges) + _close_on_host_best_effort( + host=host, + close=session.close, + session_edges=session_edges, + ) + _close_on_host_best_effort( + host=host, + close=provider.close, + session_edges=session_edges, + ) return session_edges.close_result( status=final_status, diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index 3a5e81062..b8a5231f7 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -307,6 +307,43 @@ def test_batch_driver_invariant_finalizes_edges_when_host_closed() -> None: assert provider.close_count == 0 +def test_batch_driver_ordinary_cleanup_finalizes_edges_when_host_closed() -> None: + session = _FakeVideoSession(num_steps=1) + runtime = _FakeVideoRuntime(session=session) + host = _ClosingAfterStepRuntimeHost(runtime) + provider = _FakeVideoModelInputProvider() + output = _RecordingOutputSink() + transport = _RecordingTransport() + metrics = InMemorySessionMetricsRecorder() + + result = BatchSessionDriver().run_one_session( + host=host, + provider=provider, + session_edges=SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=output, + cleanup_tasks=set(), + metrics=metrics, + transport=transport, + ), + pipeline=StepPipeline(), + ) + + assert result.status == "completed" + assert result.metrics is not None + assert result.metrics.counters["steps"] == 1 + assert result.metrics.counters["cleanup_errors"] == 2 + assert output.close_count == 1 + assert transport.close_count == 1 + assert metrics.closed + assert metrics.cleanup_errors == [ + "runtime host is closed", + "runtime host is closed", + ] + assert session.close_count == 0 + assert provider.close_count == 0 + + def test_batch_driver_invariant_finalizes_edges_when_cleanup_metrics_fail() -> None: runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) host = RuntimeHost(runtime) @@ -654,6 +691,14 @@ def call(self, func: Callable[..., Any], /, *args: object, **kwargs: object) -> return super().call(func, *args, **kwargs) +class _ClosingAfterStepRuntimeHost(_RecordingRuntimeHost): + def call(self, func: Callable[..., Any], /, *args: object, **kwargs: object) -> Any: + result = super().call(func, *args, **kwargs) + if getattr(func, "__name__", type(func).__name__) == "execute_step": + self.close() + return result + + class _RecordingOutputSink: produces_artifacts = True From 4a3b80787333d92dff508490cc62d3d963fab55f Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 04:56:45 +0000 Subject: [PATCH 13/51] runtime: split demo step requirements from input windows Add StepRequirements for model-authored per-step requirements and adapt legacy StepRequest values at the shared demo driver boundary. Move demo input, pipeline, and realtime timing contracts to driver-owned input windows, with CPU coverage for validation, deterministic slicing, realtime slicing, and legacy compatibility. --- flashdreams/flashdreams/runtime/__init__.py | 9 +- .../flashdreams/runtime/demo/drivers.py | 61 ++++++++++- .../flashdreams/runtime/demo/pipeline.py | 4 +- .../runtime/demo/session_inputs.py | 8 +- .../flashdreams/runtime/demo/timing.py | 20 ++-- flashdreams/flashdreams/runtime/types.py | 89 ++++++++++++++- .../test_demo_runtime_realtime_driver.py | 23 +++- flashdreams/tests/test_demo_runtime_timing.py | 10 +- .../tests/test_demo_runtime_vertical_slice.py | 103 +++++++++++++++++- .../tests/test_inference_runtime_api.py | 48 ++++++++ 10 files changed, 337 insertions(+), 38 deletions(-) diff --git a/flashdreams/flashdreams/runtime/__init__.py b/flashdreams/flashdreams/runtime/__init__.py index e3e574594..1613863dc 100644 --- a/flashdreams/flashdreams/runtime/__init__.py +++ b/flashdreams/flashdreams/runtime/__init__.py @@ -57,7 +57,12 @@ ) from flashdreams.runtime.output import NullOutputTarget, OutputArtifact, OutputTarget from flashdreams.runtime.runner import run_inference_session -from flashdreams.runtime.types import StepRequest, StepResult +from flashdreams.runtime.types import ( + StepRequest, + StepRequirements, + StepResult, + step_requirements_from_request, +) from flashdreams.runtime.video_output import Mp4VideoOutputTarget from flashdreams.runtime.worker import ModelExecutionWorker, ThreadAffineRuntimeWorker @@ -101,10 +106,12 @@ "RuntimeMetricSample", "ScriptedModality", "StepRequest", + "StepRequirements", "StepResult", "TimeWindow", "ThreadAffineRuntimeWorker", "run_inference_session", + "step_requirements_from_request", "undeclared_inference_inputs", "UserInputCapability", "UserInputEvent", diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 5728f46dc..abb06776c 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -10,6 +10,11 @@ from typing import Any, cast from flashdreams.runtime.interfaces import InferenceSession +from flashdreams.runtime.types import ( + StepRequest, + StepRequirements, + step_requirements_from_request, +) from .host import RuntimeHost from .outputs import SessionInfo @@ -75,7 +80,7 @@ def run_one_session( try: if session_edges.input_source.is_finished(): break - request = host.call(session.next_step_request) + request = _next_step_requirements(host=host, session=session) if request is None: break user_window = input_source.next_window(request) @@ -218,7 +223,10 @@ async def run_one_session( final_reason = "transport closed before first step" break try: - request = await host.call_async(session.next_step_request) + request = await _next_step_requirements_async( + host=host, + session=session, + ) if request is None: break window_result = await input_source.next_realtime_window( @@ -517,6 +525,55 @@ def _session_info(session: InferenceSession) -> SessionInfo: return value +def _next_step_requirements( + *, + host: RuntimeHost, + session: InferenceSession, +) -> StepRequirements | None: + next_requirements = getattr(session, "next_step_requirements", None) + if callable(next_requirements): + return _coerce_step_requirements(host.call(next_requirements)) + + next_request = getattr(session, "next_step_request", None) + if not callable(next_request): + raise TypeError( + "InferenceSession must provide next_step_requirements() or " + "legacy next_step_request()." + ) + return _coerce_step_requirements(host.call(next_request)) + + +async def _next_step_requirements_async( + *, + host: RuntimeHost, + session: InferenceSession, +) -> StepRequirements | None: + next_requirements = getattr(session, "next_step_requirements", None) + if callable(next_requirements): + return _coerce_step_requirements(await host.call_async(next_requirements)) + + next_request = getattr(session, "next_step_request", None) + if not callable(next_request): + raise TypeError( + "InferenceSession must provide next_step_requirements() or " + "legacy next_step_request()." + ) + return _coerce_step_requirements(await host.call_async(next_request)) + + +def _coerce_step_requirements(value: object) -> StepRequirements | None: + if value is None: + return None + if isinstance(value, StepRequirements): + return value + if isinstance(value, StepRequest): + return step_requirements_from_request(value) + raise TypeError( + "Session next-step method must return StepRequirements, legacy " + f"StepRequest, or None; got {type(value).__name__}." + ) + + def _close_safely(close: Any, session_edges: SessionEdges) -> None: try: close() diff --git a/flashdreams/flashdreams/runtime/demo/pipeline.py b/flashdreams/flashdreams/runtime/demo/pipeline.py index bca09ba18..5a9da08b0 100644 --- a/flashdreams/flashdreams/runtime/demo/pipeline.py +++ b/flashdreams/flashdreams/runtime/demo/pipeline.py @@ -8,7 +8,7 @@ from dataclasses import dataclass, field from flashdreams.runtime.interfaces import InferenceSession -from flashdreams.runtime.types import StepRequest, StepResult +from flashdreams.runtime.types import StepRequirements, StepResult from .outputs import OutputDecision, OutputSink from .run_modes import SessionMetricsRecorder @@ -29,7 +29,7 @@ class StepPipeline: def execute_step( self, *, - request: StepRequest, + request: StepRequirements, user_window: UserInputWindow, provider: ModelInputProvider, session: InferenceSession, diff --git a/flashdreams/flashdreams/runtime/demo/session_inputs.py b/flashdreams/flashdreams/runtime/demo/session_inputs.py index b724c0d03..ae976dd24 100644 --- a/flashdreams/flashdreams/runtime/demo/session_inputs.py +++ b/flashdreams/flashdreams/runtime/demo/session_inputs.py @@ -12,7 +12,7 @@ from flashdreams.runtime._utils import freeze_mapping from flashdreams.runtime.inputs import InferenceInput, UserInputs -from flashdreams.runtime.types import StepRequest +from flashdreams.runtime.types import StepRequirements if TYPE_CHECKING: from .timing import RealtimeClock, RealtimeWindowResult @@ -92,7 +92,7 @@ def is_finished(self) -> bool: class BatchInputSource(InputSource, Protocol): """Finite input source consumed by the batch driver.""" - def next_window(self, request: StepRequest) -> UserInputWindow: + def next_window(self, request: StepRequirements) -> UserInputWindow: """Return the next batch input window for ``request``.""" ... @@ -104,7 +104,7 @@ class RealtimeInputSource(InputSource, Protocol): async def next_realtime_window( self, *, - request: StepRequest, + request: StepRequirements, clock: "RealtimeClock", ) -> "RealtimeWindowResult": """Return the next realtime window result. @@ -127,7 +127,7 @@ def prepare_initial_input(self) -> InferenceInput: def prepare_step( self, *, - request: StepRequest, + request: StepRequirements, user_window: UserInputWindow, ) -> PreparedStep: """Prepare one model step from a driver-owned user input window.""" diff --git a/flashdreams/flashdreams/runtime/demo/timing.py b/flashdreams/flashdreams/runtime/demo/timing.py index 0261e60fb..a07af6d50 100644 --- a/flashdreams/flashdreams/runtime/demo/timing.py +++ b/flashdreams/flashdreams/runtime/demo/timing.py @@ -13,7 +13,7 @@ from typing import Literal, Protocol, runtime_checkable from flashdreams.runtime.inputs import UserInputs -from flashdreams.runtime.types import StepRequest +from flashdreams.runtime.types import StepRequirements from flashdreams.serving.realtime.input import KeyboardResampler from .session_inputs import UserInputWindow @@ -91,7 +91,7 @@ async def apply_backpressure(self, requested_s: float) -> None: ... def catch_up( self, *, - request: StepRequest, + request: StepRequirements, max_lag_s: float, policy: CatchUpPolicy, ) -> CatchUpDecision: ... @@ -226,7 +226,7 @@ async def apply_backpressure(self, requested_s: float) -> None: def catch_up( self, *, - request: StepRequest, + request: StepRequirements, max_lag_s: float, policy: CatchUpPolicy, ) -> CatchUpDecision: @@ -299,7 +299,7 @@ def reset(self, *, start_v: float) -> None: async def next_realtime_window( self, *, - request: StepRequest, + request: StepRequirements, clock: RealtimeClock, ) -> RealtimeWindowResult: input_frame_count = input_frame_count_from_request(request) @@ -326,17 +326,15 @@ async def next_realtime_window( return RealtimeWindowResult(window=window, catch_up=catch_up) -def input_frame_count_from_request(request: StepRequest) -> int: - """Return the positive realtime input frame count declared on a request.""" +def input_frame_count_from_request(request: StepRequirements) -> int: + """Return the positive input frame count declared by a step requirement.""" - value = request.metadata.get("input_frame_count") + value = request.input_frame_count if isinstance(value, bool) or not isinstance(value, int): - raise ValueError( - "StepRequest.metadata['input_frame_count'] must be an integer." - ) + raise ValueError("StepRequirements.input_frame_count must be an integer.") parsed = value if parsed <= 0: - raise ValueError("StepRequest.metadata['input_frame_count'] must be > 0.") + raise ValueError("StepRequirements.input_frame_count must be > 0.") return parsed diff --git a/flashdreams/flashdreams/runtime/types.py b/flashdreams/flashdreams/runtime/types.py index 4130a925f..1773f0cae 100644 --- a/flashdreams/flashdreams/runtime/types.py +++ b/flashdreams/flashdreams/runtime/types.py @@ -13,6 +13,65 @@ from flashdreams.runtime._utils import freeze_mapping from flashdreams.runtime.inputs import InferenceInputSchema, TimeWindow +_STEP_REQUIREMENTS_INPUT_COUNT_METADATA_KEY = "input_frame_count" +_STEP_REQUIREMENTS_STEADY_OUTPUT_COUNT_METADATA_KEY = "steady_output_frame_count" +_STEP_REQUIREMENTS_USER_INPUT_METADATA_KEYS = frozenset( + { + "input_window", + "user_input", + "user_input_window", + "user_inputs", + } +) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class StepRequirements: + """Model-authored per-step requirements consumed by shared demo drivers.""" + + __hash__ = None + + step_index: int + input_frame_count: int = 1 + steady_output_frame_count: int | None = None + inference_input_schema: InferenceInputSchema | None = None + metadata: Mapping[str, Any] = field(default_factory=dict) + + def __post_init__(self) -> None: + if isinstance(self.step_index, bool) or not isinstance(self.step_index, int): + raise TypeError("StepRequirements.step_index must be an integer.") + if self.step_index < 0: + raise ValueError("StepRequirements.step_index must be >= 0.") + if isinstance(self.input_frame_count, bool) or not isinstance( + self.input_frame_count, int + ): + raise TypeError("StepRequirements.input_frame_count must be an integer.") + if self.input_frame_count <= 0: + raise ValueError("StepRequirements.input_frame_count must be > 0.") + if self.steady_output_frame_count is not None: + if isinstance(self.steady_output_frame_count, bool) or not isinstance( + self.steady_output_frame_count, int + ): + raise TypeError( + "StepRequirements.steady_output_frame_count must be an integer." + ) + if self.steady_output_frame_count < 0: + raise ValueError( + "StepRequirements.steady_output_frame_count must be >= 0." + ) + user_input_keys = sorted( + key + for key in self.metadata + if key in _STEP_REQUIREMENTS_USER_INPUT_METADATA_KEYS + ) + if user_input_keys: + joined = ", ".join(user_input_keys) + raise ValueError( + "StepRequirements.metadata must not include driver-owned user " + f"input keys: {joined}." + ) + object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) + @dataclass(frozen=True, kw_only=True, slots=True) class StepRequest: @@ -36,4 +95,32 @@ def __post_init__(self) -> None: object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) -__all__ = ["StepRequest", "StepResult"] +def step_requirements_from_request(request: StepRequest) -> StepRequirements: + """Adapt a legacy ``StepRequest`` that did not carry driver-owned inputs.""" + + if request.user_input_window is not None: + raise ValueError( + "StepRequest.user_input_window cannot be adapted to StepRequirements; " + "user input windows are driver-owned." + ) + metadata = dict(request.metadata) + input_frame_count = metadata.pop(_STEP_REQUIREMENTS_INPUT_COUNT_METADATA_KEY, 1) + steady_output_frame_count = metadata.pop( + _STEP_REQUIREMENTS_STEADY_OUTPUT_COUNT_METADATA_KEY, + None, + ) + return StepRequirements( + step_index=request.step_index, + input_frame_count=input_frame_count, + steady_output_frame_count=steady_output_frame_count, + inference_input_schema=request.inference_input_schema, + metadata=metadata, + ) + + +__all__ = [ + "StepRequest", + "StepRequirements", + "StepResult", + "step_requirements_from_request", +] diff --git a/flashdreams/tests/test_demo_runtime_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py index 7cb813bda..57d2a2550 100644 --- a/flashdreams/tests/test_demo_runtime_realtime_driver.py +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -14,6 +14,7 @@ InferenceInput, InferenceRuntime, StepRequest, + StepRequirements, StepResult, ) from flashdreams.runtime.demo import ( @@ -163,14 +164,21 @@ async def test_cancelled_realtime_driver_inside_timeout_returns_result() -> None ) ) + async def timeout_after_window_entry(timeout: Any) -> None: + await entered_window.wait() + timeout.reschedule(asyncio.get_running_loop().time()) + try: - async with timeout_context(0.01): + async with timeout_context(None) as timeout: + timeout_task = asyncio.create_task(timeout_after_window_entry(timeout)) result = await RealtimeSessionDriver().run_one_session( host=host, provider=_FakeRealtimeProvider(), session_edges=edges, pipeline=StepPipeline(), ) + timeout_task.cancel() + await asyncio.gather(timeout_task, return_exceptions=True) finally: host.close() @@ -532,7 +540,7 @@ def __init__( self.entered = entered self.wait_forever = wait_forever self.transport_to_close = transport_to_close - self.requests: list[StepRequest] = [] + self.requests: list[StepRequirements] = [] def is_finished(self) -> bool: return False @@ -540,7 +548,7 @@ def is_finished(self) -> bool: async def next_realtime_window( self, *, - request: StepRequest, + request: StepRequirements, clock: Any, ) -> RealtimeWindowResult: del clock @@ -578,7 +586,7 @@ def prepare_initial_input(self) -> InferenceInput: def prepare_step( self, *, - request: StepRequest, + request: StepRequirements, user_window: UserInputWindow, ) -> PreparedStep: return PreparedStep( @@ -635,13 +643,16 @@ def __init__( def session_info(self) -> SessionInfo: return SessionInfo(output_layout="fake-realtime", steady_output_frame_count=1) - def next_step_request(self) -> StepRequest | None: + def next_step_requirements(self) -> StepRequirements | None: if self.next_request_index >= self.num_steps: return None - request = StepRequest(step_index=self.next_request_index) + request = StepRequirements(step_index=self.next_request_index) self.next_request_index += 1 return request + def next_step_request(self) -> StepRequest | None: + raise AssertionError("demo driver should request StepRequirements") + def step(self, inputs: InferenceInput) -> StepResult: step_index = len(self.step_inputs) if self.step_delay_s: diff --git a/flashdreams/tests/test_demo_runtime_timing.py b/flashdreams/tests/test_demo_runtime_timing.py index 02e151680..a30c2752f 100644 --- a/flashdreams/tests/test_demo_runtime_timing.py +++ b/flashdreams/tests/test_demo_runtime_timing.py @@ -8,7 +8,7 @@ import pytest -from flashdreams.runtime import StepRequest +from flashdreams.runtime import StepRequirements from flashdreams.runtime.demo import NullOutputSink, RunResult, SessionEdges from flashdreams.runtime.demo.timing import ( SPARSE_KEY_SEGMENTS_METADATA_KEY, @@ -192,10 +192,10 @@ def test_keyboard_resampler_defers_unsupported_catch_up_policies( ) -def _request(*, input_frame_count: int) -> StepRequest: - return StepRequest( +def _request(*, input_frame_count: int) -> StepRequirements: + return StepRequirements( step_index=0, - metadata={"input_frame_count": input_frame_count}, + input_frame_count=input_frame_count, ) @@ -229,7 +229,7 @@ def is_finished(self) -> bool: async def next_realtime_window( self, *, - request: StepRequest, + request: StepRequirements, clock: Any, ) -> object: del request, clock diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index b8a5231f7..866998523 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -19,6 +19,7 @@ InputMapping, OutputArtifact, StepRequest, + StepRequirements, StepResult, UserInputs, ) @@ -54,7 +55,7 @@ def test_step_pipeline_passes_provider_input_to_session_and_sink() -> None: output = _RecordingOutputSink() output.open(SessionInfo(output_layout="fake-video", steady_output_frame_count=1)) metrics = InMemorySessionMetricsRecorder() - request = StepRequest(step_index=0) + request = StepRequirements(step_index=0) user_window = _window(0) outcome = StepPipeline().execute_step( @@ -115,6 +116,43 @@ def test_batch_driver_runs_fake_video_demo_through_runtime_host() -> None: assert "step" not in host.calls +def test_batch_driver_slices_windows_from_step_requirements() -> None: + session = _FakeVideoSession(num_steps=2, input_frame_counts=(3, 2)) + runtime = _FakeVideoRuntime(session=session) + host = _RecordingRuntimeHost(runtime) + provider = _FakeVideoModelInputProvider() + input_source = _SlicingBatchInputSource(fps=2.0, num_windows=2) + + result = BatchSessionDriver().run_one_session( + host=host, + provider=provider, + session_edges=SessionEdges( + input_source=input_source, + output_sink=_RecordingOutputSink(), + cleanup_tasks=set(), + metrics=InMemorySessionMetricsRecorder(), + ), + pipeline=StepPipeline(), + ) + + assert result.status == "completed" + assert [request.step_index for request in input_source.next_window_requests] == [ + 0, + 1, + ] + assert [ + request.input_frame_count for request in input_source.next_window_requests + ] == [3, 2] + assert input_source.windows == [ + _window_with_frame_times(start_s=0.0, frame_times=(0.0, 0.5, 1.0)), + _window_with_frame_times(start_s=1.5, frame_times=(1.5, 2.0)), + ] + assert [dict(inputs.step) for inputs in session.step_inputs] == [ + {"request_step": 0, "window": (0.0, 1.5)}, + {"request_step": 1, "window": (1.5, 2.5)}, + ] + + def test_run_demo_session_builds_edges_and_records_session_once() -> None: session = _FakeVideoSession(num_steps=1) runtime = _FakeVideoRuntime(session=session) @@ -521,6 +559,19 @@ def _window(index: int) -> UserInputWindow: ) +def _window_with_frame_times( + *, + start_s: float, + frame_times: Sequence[float], +) -> UserInputWindow: + return UserInputWindow( + start_s=start_s, + end_s=start_s + len(frame_times) * 0.5, + frame_times=frame_times, + inputs=UserInputs(), + ) + + def _spec() -> DemoSpec: return DemoSpec( model_id="fake-video-demo", @@ -577,7 +628,7 @@ def prepare_initial_input(self) -> InferenceInput: def prepare_step( self, *, - request: StepRequest, + request: StepRequirements, user_window: UserInputWindow, ) -> PreparedStep: inference_input = InferenceInput( @@ -608,7 +659,7 @@ def __init__( ) -> None: self.windows = [_window(index) for index in range(num_windows)] self.fail_is_finished = fail_is_finished - self.next_window_requests: list[StepRequest] = [] + self.next_window_requests: list[StepRequirements] = [] self.index = 0 def is_finished(self) -> bool: @@ -616,13 +667,45 @@ def is_finished(self) -> bool: raise self.fail_is_finished return self.index >= len(self.windows) - def next_window(self, request: StepRequest) -> UserInputWindow: + def next_window(self, request: StepRequirements) -> UserInputWindow: self.next_window_requests.append(request) window = self.windows[self.index] self.index += 1 return window +class _SlicingBatchInputSource: + is_finite = True + is_deterministic = True + + def __init__(self, *, fps: float, num_windows: int) -> None: + self.fps = fps + self.num_windows = num_windows + self.next_window_requests: list[StepRequirements] = [] + self.windows: list[UserInputWindow] = [] + self.window_index = 0 + self.next_frame_index = 0 + + def is_finished(self) -> bool: + return self.window_index >= self.num_windows + + def next_window(self, request: StepRequirements) -> UserInputWindow: + self.next_window_requests.append(request) + start_frame = self.next_frame_index + self.next_frame_index += request.input_frame_count + self.window_index += 1 + frame_times = tuple( + frame_index / self.fps + for frame_index in range(start_frame, self.next_frame_index) + ) + window = _window_with_frame_times( + start_s=start_frame / self.fps, + frame_times=frame_times, + ) + self.windows.append(window) + return window + + class _FakeVideoRuntime: def __init__(self, *, session: "_FakeVideoSession") -> None: self.session = session @@ -642,9 +725,11 @@ def __init__( self, *, num_steps: int, + input_frame_counts: Sequence[int] | None = None, fail_step: int | None = None, ) -> None: self.num_steps = num_steps + self.input_frame_counts = tuple(input_frame_counts or (1,) * num_steps) self.fail_step = fail_step self.next_request_index = 0 self.step_inputs: list[InferenceInput] = [] @@ -653,13 +738,19 @@ def __init__( def session_info(self) -> SessionInfo: return SessionInfo(output_layout="fake-video", steady_output_frame_count=1) - def next_step_request(self) -> StepRequest | None: + def next_step_requirements(self) -> StepRequirements | None: if self.next_request_index >= self.num_steps: return None - request = StepRequest(step_index=self.next_request_index) + request = StepRequirements( + step_index=self.next_request_index, + input_frame_count=self.input_frame_counts[self.next_request_index], + ) self.next_request_index += 1 return request + def next_step_request(self) -> StepRequest | None: + raise AssertionError("demo driver should request StepRequirements") + def step(self, inputs: InferenceInput) -> StepResult: step_index = len(self.step_inputs) if self.fail_step == step_index: diff --git a/flashdreams/tests/test_inference_runtime_api.py b/flashdreams/tests/test_inference_runtime_api.py index 42f75d688..da00789c3 100644 --- a/flashdreams/tests/test_inference_runtime_api.py +++ b/flashdreams/tests/test_inference_runtime_api.py @@ -20,11 +20,13 @@ OutputArtifact, RuntimeMetricSample, StepRequest, + StepRequirements, StepResult, TimeWindow, UserInputEvent, UserInputs, UserInputSchema, + step_requirements_from_request, ) pytestmark = pytest.mark.ci_cpu @@ -67,6 +69,12 @@ def test_inference_config_rejects_empty_model_id() -> None: ), (lambda: UserInputEvent(timestamp_s=0.0, event_type=" "), "event_type"), (lambda: StepRequest(step_index=-1), "step_index"), + (lambda: StepRequirements(step_index=-1), "step_index"), + (lambda: StepRequirements(step_index=0, input_frame_count=0), "input_frame"), + ( + lambda: StepRequirements(step_index=0, steady_output_frame_count=-1), + "steady_output", + ), (lambda: StepResult(step_index=-1), "step_index"), (lambda: StepResult(step_index=0, frame_count=-1), "frame_count"), (lambda: RuntimeMetricSample(name=" ", value=1.0), "name"), @@ -185,6 +193,46 @@ def test_identity_input_mapping_leaves_inference_input_unchanged() -> None: ) +def test_step_requirements_adapt_legacy_request_metadata() -> None: + schema = InferenceInputSchema(step_fields=(InputField(name="camera_poses"),)) + request = StepRequest( + step_index=3, + inference_input_schema=schema, + metadata={ + "input_frame_count": 4, + "steady_output_frame_count": 2, + "model": "fake-video-demo", + }, + ) + + requirements = step_requirements_from_request(request) + + assert requirements == StepRequirements( + step_index=3, + input_frame_count=4, + steady_output_frame_count=2, + inference_input_schema=schema, + metadata={"model": "fake-video-demo"}, + ) + with pytest.raises(TypeError): + cast(Any, requirements.metadata)["model"] = "changed" + + +def test_step_requirements_keep_user_inputs_driver_owned() -> None: + requirements = StepRequirements(step_index=0, metadata={"model": "fake"}) + + assert not hasattr(requirements, "user_input_window") + with pytest.raises(ValueError, match="driver-owned user input"): + StepRequirements(step_index=0, metadata={"user_inputs": UserInputs()}) + with pytest.raises(ValueError, match="driver-owned"): + step_requirements_from_request( + StepRequest( + step_index=0, + user_input_window=TimeWindow(start_s=0.0, end_s=1.0), + ) + ) + + def test_null_output_target_counts_and_optionally_stores_results() -> None: target = NullOutputTarget(store_results=True) result = StepResult(step_index=0, output=b"frame") From 19a44496eb0c43520aee7c3443427e0bbe814b87 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 05:26:19 +0000 Subject: [PATCH 14/51] runtime: add MP4 and null demo output sinks Add Mp4OutputSink and build_output_sink for the shared demo runtime while preserving the legacy OutputTarget builder for replay. Make null sink step recording lightweight, keep sink close idempotent, and cover MP4 artifacts, sink construction, payload ownership, setup-open failures, and cleanup-close failures with CPU tests. --- .../flashdreams/runtime/demo/__init__.py | 4 + .../flashdreams/runtime/demo/outputs.py | 155 +++++++++++- .../tests/test_demo_runtime_output_sinks.py | 221 ++++++++++++++++++ .../tests/test_demo_runtime_vertical_slice.py | 61 +++++ 4 files changed, 438 insertions(+), 3 deletions(-) create mode 100644 flashdreams/tests/test_demo_runtime_output_sinks.py diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index 1a3f76892..680cfbb30 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -19,10 +19,12 @@ WarmupSessionInputs, ) from flashdreams.runtime.demo.outputs import ( + Mp4OutputSink, NullOutputSink, OutputDecision, OutputSink, SessionInfo, + build_output_sink, build_output_target, ) from flashdreams.runtime.demo.pipeline import StepOutcome, StepPipeline @@ -103,6 +105,7 @@ "ModelWarmupPlan", "ModelInputProvider", "KeyboardRealtimeInputSource", + "Mp4OutputSink", "Mp4OutputSpec", "NoopTransportService", "NullOutputSpec", @@ -135,6 +138,7 @@ "WarmupSessionInputs", "WebRTCAppResources", "WebRTCOutputSpec", + "build_output_sink", "build_output_target", "input_frame_count_from_request", "run_demo_session", diff --git a/flashdreams/flashdreams/runtime/demo/outputs.py b/flashdreams/flashdreams/runtime/demo/outputs.py index d24f99da7..866acbb00 100644 --- a/flashdreams/flashdreams/runtime/demo/outputs.py +++ b/flashdreams/flashdreams/runtime/demo/outputs.py @@ -1,7 +1,7 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""Shared demo output contracts and output-target construction.""" +"""Shared demo output contracts and output construction.""" from __future__ import annotations @@ -10,6 +10,12 @@ from pathlib import Path from typing import Literal, Protocol, runtime_checkable +from flashdreams.infra.postprocess import VideoTensorLayout +from flashdreams.infra.runner_io import ( + DEFAULT_RUNNER_INSTALL_HINT, + write_video_tensor, +) +from flashdreams.infra.video_output import VideoResultCollector, prepare_video_for_mp4 from flashdreams.runtime._utils import freeze_mapping from flashdreams.runtime.output import NullOutputTarget, OutputArtifact, OutputTarget from flashdreams.runtime.types import StepResult @@ -87,7 +93,7 @@ class NullOutputSink: store_results: bool = False produces_artifacts: bool = False output_count: int = field(default=0, init=False) - results: list[StepResult] = field(default_factory=list, init=False) + results: list[Mapping[str, object]] = field(default_factory=list, init=False) opened: bool = field(default=False, init=False) closed: bool = field(default=False, init=False) session_info: SessionInfo | None = field(default=None, init=False) @@ -110,7 +116,7 @@ def write(self, result: StepResult) -> OutputDecision: raise RuntimeError("Cannot write to a closed output sink.") self.output_count += 1 if self.store_results: - self.results.append(result) + self.results.append(_result_record(result)) return OutputDecision() def close(self) -> Sequence[OutputArtifact]: @@ -118,6 +124,147 @@ def close(self) -> Sequence[OutputArtifact]: return () +@dataclass(slots=True) +class Mp4OutputSink: + """MP4 artifact sink for shared demo drivers.""" + + output_path: Path + fps: int | float + output_layout: VideoTensorLayout = "bvtchw" + writer: VideoWriter = field(default=write_video_tensor, repr=False) + install_hint: str = DEFAULT_RUNNER_INSTALL_HINT + move_to_cpu: bool = True + enabled: bool = True + produces_artifacts: bool = True + _opened: bool = field(default=False, init=False, repr=False) + _closed: bool = field(default=True, init=False, repr=False) + _collector: VideoResultCollector | None = field( + default=None, + init=False, + repr=False, + ) + _artifacts: tuple[OutputArtifact, ...] | None = field( + default=None, + init=False, + repr=False, + ) + session_info: SessionInfo | None = field(default=None, init=False) + + def __post_init__(self) -> None: + if float(self.fps) <= 0: + raise ValueError("Mp4OutputSink.fps must be > 0.") + self.output_path = Path(self.output_path) + + def open(self, session_info: SessionInfo) -> None: + self.session_info = session_info + self._collector = VideoResultCollector( + output_layout=self.output_layout, + enabled=self.enabled, + move_to_cpu=self.move_to_cpu, + ) + self._artifacts = None + self._opened = True + self._closed = False + + def begin_generation(self, generation: int) -> None: + if generation < 0: + raise ValueError("generation must be >= 0.") + + def write(self, result: StepResult) -> OutputDecision: + if not self._opened or self._closed or self._collector is None: + raise RuntimeError("Cannot write to a closed output sink.") + if result.layout is None: + raise TypeError("Mp4OutputSink requires a video StepResult with layout.") + if result.layout != self.output_layout: + raise ValueError( + "Mp4OutputSink received layout " + f"{result.layout!r}; expected {self.output_layout!r}." + ) + self._collector.add(result) + return OutputDecision() + + def close(self) -> Sequence[OutputArtifact]: + if self._artifacts is not None: + return self._artifacts + if self._collector is None: + self._opened = False + self._closed = True + self._artifacts = () + return self._artifacts + + collector = self._collector + self._collector = None + self._opened = False + self._closed = True + video = collector.finish() + if video is None: + self._artifacts = () + return self._artifacts + writable_video, writable_layout = prepare_video_for_mp4( + video, + layout=self.output_layout, + ) + path = self.writer( + writable_video, + self.output_path, + fps=self.fps, + layout=writable_layout, + install_hint=self.install_hint, + ) + self._artifacts = ( + OutputArtifact( + kind="video/mp4", + uri=str(path), + metadata={ + "fps": self.fps, + "source_layout": self.output_layout, + "shape": tuple(int(dim) for dim in video.shape), + "stats_history": tuple(collector.stats_history), + }, + ), + ) + return self._artifacts + + +def build_output_sink( + output: OutputSpec, + *, + mp4_writer: VideoWriter | None = None, +) -> OutputSink: + """Build a shared demo output sink from a demo output spec.""" + if isinstance(output, NullOutputSpec): + return NullOutputSink(store_results=output.store_results) + if isinstance(output, Mp4OutputSpec): + writer = mp4_writer or write_video_tensor + return Mp4OutputSink( + output_path=Path(output.path), + fps=output.fps, + output_layout=output.output_layout, + writer=writer, + move_to_cpu=output.move_to_cpu, + ) + if isinstance(output, WebRTCOutputSpec): + raise ValueError("WebRTC output requires a realtime transport sink.") + raise TypeError(f"Unsupported demo output spec: {type(output).__name__}.") + + +def _result_record(result: StepResult) -> Mapping[str, object]: + record: dict[str, object] = { + "step_index": result.step_index, + "frame_count": result.frame_count, + "metrics": dict(result.metrics), + "metadata": dict(result.metadata), + } + if result.layout is not None: + record["layout"] = result.layout + if result.output_window is not None: + record["output_window"] = ( + result.output_window.start_s, + result.output_window.end_s, + ) + return freeze_mapping(record) + + def build_output_target( output: OutputSpec, *, @@ -148,9 +295,11 @@ def build_output_target( __all__ = [ + "Mp4OutputSink", "NullOutputSink", "OutputDecision", "OutputSink", "SessionInfo", + "build_output_sink", "build_output_target", ] diff --git a/flashdreams/tests/test_demo_runtime_output_sinks.py b/flashdreams/tests/test_demo_runtime_output_sinks.py new file mode 100644 index 000000000..147e86407 --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_output_sinks.py @@ -0,0 +1,221 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from collections.abc import Mapping +from dataclasses import fields, is_dataclass +from pathlib import Path +from typing import Any + +import pytest +import torch + +from flashdreams.runtime import OutputArtifact, StepResult, TimeWindow +from flashdreams.runtime.demo import ( + DemoSpec, + Mp4OutputSink, + Mp4OutputSpec, + NullOutputSink, + NullOutputSpec, + OutputDecision, + SessionInfo, + WebRTCOutputSpec, + build_output_sink, +) + +pytestmark = pytest.mark.ci_cpu + + +def test_mp4_output_sink_writes_artifact_and_close_is_idempotent( + tmp_path: Path, +) -> None: + writer_calls: list[dict[str, Any]] = [] + + def fake_writer( + video: torch.Tensor, + path: Path, + *, + fps: int | float, + layout: str, + install_hint: str, + ) -> Path: + del install_hint + writer_calls.append( + { + "shape": tuple(video.shape), + "path": path, + "fps": fps, + "layout": layout, + } + ) + return path + + sink = Mp4OutputSink( + output_path=tmp_path / "out.mp4", + fps=24, + writer=fake_writer, + move_to_cpu=False, + ) + sink.open(SessionInfo(output_layout="bvtchw", steady_output_frame_count=1)) + sink.begin_generation(0) + + decision = sink.write( + StepResult.from_video_chunk( + step_index=2, + video_chunk=torch.zeros((1, 2, 3, 3, 4, 5)), + layout="bvtchw", + metrics={"model_step_s": 0.25}, + output_window=TimeWindow(start_s=1.0, end_s=2.0), + ) + ) + artifacts = tuple(sink.close()) + second_close = tuple(sink.close()) + + assert decision == OutputDecision() + assert artifacts == second_close + assert artifacts == ( + OutputArtifact( + kind="video/mp4", + uri=str(tmp_path / "out.mp4"), + metadata={ + "fps": 24, + "source_layout": "bvtchw", + "shape": (1, 2, 3, 3, 4, 5), + "stats_history": ( + { + "step_index": 2, + "frames": 3, + "model_step_s": 0.25, + "output_start_s": 1.0, + "output_end_s": 2.0, + }, + ), + }, + ), + ) + assert writer_calls == [ + { + "shape": (3, 4, 10, 3), + "path": tmp_path / "out.mp4", + "fps": 24, + "layout": "thwc", + } + ] + + +def test_output_sink_is_built_from_demo_spec(tmp_path: Path) -> None: + def fake_writer(*args: Any, **kwargs: Any) -> Path: + del args, kwargs + return tmp_path / "demo.mp4" + + spec = DemoSpec( + model_id="fake-demo", + input_mode="replay", + output=Mp4OutputSpec(path=tmp_path / "demo.mp4", fps=12), + ) + + mp4_sink = build_output_sink(spec.output, mp4_writer=fake_writer) + null_sink = build_output_sink(NullOutputSpec(store_results=True)) + + assert isinstance(mp4_sink, Mp4OutputSink) + assert mp4_sink.output_path == tmp_path / "demo.mp4" + assert mp4_sink.fps == 12 + assert mp4_sink.writer is fake_writer + assert isinstance(null_sink, NullOutputSink) + assert null_sink.store_results + with pytest.raises(ValueError, match="realtime transport sink"): + build_output_sink(WebRTCOutputSpec()) + + +def test_sinks_do_not_retain_step_result_references(tmp_path: Path) -> None: + mp4_sink = Mp4OutputSink( + output_path=tmp_path / "out.mp4", + fps=24, + writer=lambda *args: tmp_path / "out.mp4", + move_to_cpu=False, + ) + mp4_sink.open(SessionInfo(output_layout="bvtchw", steady_output_frame_count=1)) + mp4_result = StepResult.from_video_chunk( + step_index=0, + video_chunk=torch.zeros((1, 1, 1, 3, 2, 2)), + layout="bvtchw", + ) + + mp4_sink.write(mp4_result) + + null_sink = NullOutputSink(store_results=True) + null_sink.open(SessionInfo()) + null_result = StepResult( + step_index=1, + output=object(), + frame_count=2, + metrics={"model_step_s": 0.1}, + metadata={"source": "fake"}, + ) + + null_sink.write(null_result) + + assert not _object_graph_contains(mp4_sink, mp4_result) + assert not _object_graph_contains(null_sink, null_result) + assert null_sink.results == [ + { + "step_index": 1, + "frame_count": 2, + "metrics": {"model_step_s": 0.1}, + "metadata": {"source": "fake"}, + } + ] + + +def test_null_output_sink_records_steps_without_artifacts() -> None: + sink = NullOutputSink(store_results=True) + sink.open(SessionInfo(output_layout="fake-video", steady_output_frame_count=1)) + + sink.write(StepResult(step_index=0, output="first", frame_count=1)) + sink.write(StepResult(step_index=1, output="second", frame_count=2)) + artifacts = tuple(sink.close()) + + assert artifacts == () + assert tuple(sink.close()) == () + assert sink.output_count == 2 + assert sink.results == [ + {"step_index": 0, "frame_count": 1, "metrics": {}, "metadata": {}}, + {"step_index": 1, "frame_count": 2, "metrics": {}, "metadata": {}}, + ] + + +def _object_graph_contains( + root: object, + needle: object, + *, + seen: set[int] | None = None, +) -> bool: + if root is needle: + return True + if seen is None: + seen = set() + root_id = id(root) + if root_id in seen: + return False + seen.add(root_id) + if root is None or isinstance(root, str | bytes | int | float | bool | Path): + return False + if isinstance(root, torch.Tensor): + return False + if callable(root): + return False + if isinstance(root, Mapping): + return any( + _object_graph_contains(key, needle, seen=seen) + or _object_graph_contains(value, needle, seen=seen) + for key, value in root.items() + ) + if isinstance(root, list | tuple | set | frozenset): + return any(_object_graph_contains(value, needle, seen=seen) for value in root) + if is_dataclass(root): + return any( + _object_graph_contains(getattr(root, field.name), needle, seen=seen) + for field in fields(root) + ) + return False diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index 866998523..d7f4bb492 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -245,6 +245,31 @@ def test_setup_failure_returns_failed_before_runtime_session_creation() -> None: assert metrics.errors == ["invalid provider compatibility"] +def test_output_sink_open_failure_returns_failed_before_step_loop() -> None: + session = _FakeVideoSession(num_steps=1) + metrics = InMemorySessionMetricsRecorder() + output = _RecordingOutputSink(fail_open=RuntimeError("open failed")) + + result = BatchSessionDriver().run_one_session( + host=RuntimeHost(_FakeVideoRuntime(session=session)), + provider=_FakeVideoModelInputProvider(), + session_edges=SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=output, + cleanup_tasks=set(), + metrics=metrics, + ), + pipeline=StepPipeline(), + ) + + assert result.status == "failed" + assert result.reason == "open failed" + assert metrics.errors == ["open failed"] + assert session.step_inputs == [] + assert output.results == [] + assert output.close_count == 1 + + def test_run_demo_session_closes_provider_when_validation_fails() -> None: runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) run_metrics = InMemorySessionMetricsRecorder() @@ -532,6 +557,34 @@ def test_session_edges_close_result_is_idempotent_and_first_result_wins() -> Non assert metrics.closed +def test_output_sink_close_failure_records_cleanup_error_without_losing_result() -> ( + None +): + session = _FakeVideoSession(num_steps=1) + metrics = InMemorySessionMetricsRecorder() + output = _RecordingOutputSink(fail_close=RuntimeError("close failed")) + + result = BatchSessionDriver().run_one_session( + host=RuntimeHost(_FakeVideoRuntime(session=session)), + provider=_FakeVideoModelInputProvider(), + session_edges=SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=output, + cleanup_tasks=set(), + metrics=metrics, + ), + pipeline=StepPipeline(), + ) + + assert result.status == "completed" + assert result.reason is None + assert result.metrics is not None + assert result.metrics.counters["steps"] == 1 + assert result.metrics.counters["cleanup_errors"] == 1 + assert result.metrics.errors == ("close failed",) + assert output.close_count == 1 + + def test_run_result_rejected_is_the_only_convenience_constructor() -> None: constructors = { name @@ -798,14 +851,20 @@ def __init__( *, artifacts: Sequence[OutputArtifact] = (), decision: OutputDecision | None = None, + fail_open: Exception | None = None, + fail_close: Exception | None = None, ) -> None: self.artifacts = tuple(artifacts) self.decision = decision or OutputDecision() + self.fail_open = fail_open + self.fail_close = fail_close self.opened_with: SessionInfo | None = None self.results: list[StepResult] = [] self.close_count = 0 def open(self, session_info: SessionInfo) -> None: + if self.fail_open is not None: + raise self.fail_open self.opened_with = session_info def begin_generation(self, generation: int) -> None: @@ -817,6 +876,8 @@ def write(self, result: StepResult) -> OutputDecision: def close(self) -> Sequence[OutputArtifact]: self.close_count += 1 + if self.fail_close is not None: + raise self.fail_close return self.artifacts From 9d3afd921e1224781068a0ba7ea348e0dc72d24b Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 06:04:13 +0000 Subject: [PATCH 15/51] runtime: add demo capability validation Add provider and run-mode capability contracts, resolved run validation, and first-class InferenceConfig.seed for deterministic demo runs. Wire validation into shared demo session helpers and cover raw input schemas, mapping-backed providers, MP4/WebRTC compatibility checks, realtime/batch source validation, determinism resolution, and reset coordination with CPU fakes. --- flashdreams/flashdreams/runtime/config.py | 8 + .../flashdreams/runtime/demo/__init__.py | 12 + .../flashdreams/runtime/demo/drivers.py | 27 + .../flashdreams/runtime/demo/run_modes.py | 13 + .../runtime/demo/session_inputs.py | 37 +- .../flashdreams/runtime/demo/timing.py | 3 +- .../flashdreams/runtime/demo/validation.py | 201 ++++++ .../test_demo_runtime_realtime_driver.py | 9 + .../tests/test_demo_runtime_run_modes.py | 17 + flashdreams/tests/test_demo_runtime_timing.py | 3 +- .../tests/test_demo_runtime_validation.py | 675 ++++++++++++++++++ .../tests/test_demo_runtime_vertical_slice.py | 15 + .../tests/test_inference_runtime_api.py | 9 + 13 files changed, 1024 insertions(+), 5 deletions(-) create mode 100644 flashdreams/flashdreams/runtime/demo/validation.py create mode 100644 flashdreams/tests/test_demo_runtime_validation.py diff --git a/flashdreams/flashdreams/runtime/config.py b/flashdreams/flashdreams/runtime/config.py index 4b8752f13..f1c0c2a0c 100644 --- a/flashdreams/flashdreams/runtime/config.py +++ b/flashdreams/flashdreams/runtime/config.py @@ -61,6 +61,9 @@ class InferenceConfig: cache_policy: str | None = None """Optional cache policy selector; ``None`` leaves the choice to the adapter.""" + seed: int | None = None + """Optional seed used when resolving deterministic demo/runtime behavior.""" + runtime_options: Mapping[str, Any] = field(default_factory=dict) """Adapter/backend-specific runtime options.""" @@ -70,6 +73,11 @@ class InferenceConfig: def __post_init__(self) -> None: if not self.model_id.strip(): raise ValueError("InferenceConfig.model_id must be non-empty.") + if self.seed is not None: + if isinstance(self.seed, bool) or not isinstance(self.seed, int): + raise TypeError("InferenceConfig.seed must be an integer.") + if self.seed < 0: + raise ValueError("InferenceConfig.seed must be >= 0.") object.__setattr__( self, "runtime_options", freeze_mapping(self.runtime_options) ) diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index 680cfbb30..cefcba4a2 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -38,6 +38,7 @@ NoopTransportService, RunContext, RunMode, + RunModeCapabilities, RunModeWarmup, RunResult, RunSummary, @@ -51,6 +52,7 @@ InputSource, ModelInputProvider, PreparedStep, + ProviderCapabilities, RealtimeInputSource, UserInputWindow, ) @@ -80,6 +82,11 @@ SignalActivationPolicy, input_frame_count_from_request, ) +from flashdreams.runtime.demo.validation import ( + ResolvedRunCapabilities, + resolve_run_capabilities, + validate_resolved_run, +) __all__ = [ "BatchInputSource", @@ -115,13 +122,16 @@ "OutputSink", "PreparedScenario", "PreparedStep", + "ProviderCapabilities", "RealtimeInputSource", "RealtimeClock", "RealtimeSessionDriver", "RealtimeWindowResult", + "ResolvedRunCapabilities", "ResamplerRealtimeClock", "RunContext", "RunMode", + "RunModeCapabilities", "RunModeWarmup", "RunResult", "RunSummary", @@ -141,9 +151,11 @@ "build_output_sink", "build_output_target", "input_frame_count_from_request", + "resolve_run_capabilities", "run_demo_session", "run_demo_session_async", "run_replay_demo", "shielded_session_cleanup", "uncancel_current_task", + "validate_resolved_run", ] diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index abb06776c..eabb5307c 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -30,6 +30,7 @@ from .session_inputs import BatchInputSource, ModelInputProvider from .spec import DemoAdapter, DemoSpec, PreparedScenario from .timing import ActivationPolicy, RealtimeClock +from .validation import resolve_run_capabilities, validate_resolved_run CLEANUP_TIMEOUT_S = 30.0 @@ -352,6 +353,19 @@ def run_demo_session( provider=provider, adapter=adapter, ) + resolved_capabilities = resolve_run_capabilities( + spec=spec, + provider=provider, + session_edges=session_edges, + ) + validate_resolved_run( + spec=spec, + adapter=adapter, + provider=provider, + run_mode=run_mode, + session_edges=session_edges, + resolved=resolved_capabilities, + ) if session_edges.is_closed: raise DriverInvariantError( "RunMode returned already closed SessionEdges; session edges " @@ -450,6 +464,19 @@ async def run_demo_session_async( provider=provider, adapter=adapter, ) + resolved_capabilities = resolve_run_capabilities( + spec=spec, + provider=provider, + session_edges=session_edges, + ) + validate_resolved_run( + spec=spec, + adapter=adapter, + provider=provider, + run_mode=run_mode, + session_edges=session_edges, + resolved=resolved_capabilities, + ) if session_edges.is_closed: raise DriverInvariantError( "RunMode returned already closed SessionEdges; session edges " diff --git a/flashdreams/flashdreams/runtime/demo/run_modes.py b/flashdreams/flashdreams/runtime/demo/run_modes.py index c815e6170..78c7b837b 100644 --- a/flashdreams/flashdreams/runtime/demo/run_modes.py +++ b/flashdreams/flashdreams/runtime/demo/run_modes.py @@ -127,6 +127,17 @@ def handle_setup_error(self, exc: Exception) -> ErrorAction: ... def handle(self, exc: Exception) -> ErrorAction: ... +@dataclass(frozen=True, kw_only=True, slots=True) +class RunModeCapabilities: + """Run-mode requirements and output/transport capabilities.""" + + realtime: bool = False + requires_finite_input: bool = False + supports_backpressure: bool = False + supports_interactive_events: bool = False + supports_artifacts: bool = False + + @runtime_checkable class SessionMetricsRecorder(Protocol): """Metrics callbacks consumed by the Phase 2 drivers and pipeline.""" @@ -429,6 +440,7 @@ class RunMode(Protocol): """Run/session construction strategy consumed by shared helpers.""" name: str + capabilities: RunModeCapabilities def validate_run( self, @@ -494,6 +506,7 @@ def warmup_context( "NoopTransportService", "RunContext", "RunMode", + "RunModeCapabilities", "RunModeWarmup", "RunResult", "RunSummary", diff --git a/flashdreams/flashdreams/runtime/demo/session_inputs.py b/flashdreams/flashdreams/runtime/demo/session_inputs.py index ae976dd24..dda42995c 100644 --- a/flashdreams/flashdreams/runtime/demo/session_inputs.py +++ b/flashdreams/flashdreams/runtime/demo/session_inputs.py @@ -11,13 +11,32 @@ from typing import TYPE_CHECKING, Protocol, runtime_checkable from flashdreams.runtime._utils import freeze_mapping -from flashdreams.runtime.inputs import InferenceInput, UserInputs +from flashdreams.runtime.inputs import ( + InferenceInput, + InferenceInputSchema, + UserInputs, + UserInputSchema, +) from flashdreams.runtime.types import StepRequirements if TYPE_CHECKING: from .timing import RealtimeClock, RealtimeWindowResult +@dataclass(frozen=True, kw_only=True, slots=True) +class ProviderCapabilities: + """Model-provider capabilities used to validate run-mode compatibility.""" + + supports_realtime_clock: bool = False + supports_recorded_input: bool = False + supports_reset: bool = False + deterministic_given_inputs: bool = False + user_input_schema: UserInputSchema = field(default_factory=UserInputSchema) + inference_input_schema: InferenceInputSchema = field( + default_factory=InferenceInputSchema + ) + + @dataclass(frozen=True, kw_only=True, slots=True) class ControlDecision: """Provider-authored control request for the current session.""" @@ -82,6 +101,7 @@ class InputSource(Protocol): is_finite: bool is_deterministic: bool + user_input_schema: UserInputSchema def is_finished(self) -> bool: """Return whether the driver should stop requesting windows.""" @@ -120,6 +140,8 @@ async def next_realtime_window( class ModelInputProvider(Protocol): """Model-owned conversion from user windows into model-facing inputs.""" + capabilities: ProviderCapabilities + def prepare_initial_input(self) -> InferenceInput: """Prepare session-global model inputs.""" ... @@ -134,11 +156,19 @@ def prepare_step( ... def reset(self, inputs: InferenceInput | None = None) -> None: - """Reset provider-owned session state.""" + """Reset provider-owned session state. + + Implementations must be idempotent so driver cleanup and reset control + paths can safely converge after failures. + """ ... def close(self) -> None: - """Release provider-owned resources.""" + """Release provider-owned resources. + + Implementations must be idempotent and tolerate cleanup after partial + setup or earlier reset failures. + """ ... @@ -148,6 +178,7 @@ def close(self) -> None: "InputSource", "ModelInputProvider", "PreparedStep", + "ProviderCapabilities", "RealtimeInputSource", "UserInputWindow", ] diff --git a/flashdreams/flashdreams/runtime/demo/timing.py b/flashdreams/flashdreams/runtime/demo/timing.py index a07af6d50..9c57237ea 100644 --- a/flashdreams/flashdreams/runtime/demo/timing.py +++ b/flashdreams/flashdreams/runtime/demo/timing.py @@ -12,7 +12,7 @@ from dataclasses import dataclass, field from typing import Literal, Protocol, runtime_checkable -from flashdreams.runtime.inputs import UserInputs +from flashdreams.runtime.inputs import UserInputs, UserInputSchema from flashdreams.runtime.types import StepRequirements from flashdreams.serving.realtime.input import KeyboardResampler @@ -273,6 +273,7 @@ class KeyboardRealtimeInputSource: catch_up_policy: CatchUpPolicy = "fold" is_finite: bool = False is_deterministic: bool = False + user_input_schema: UserInputSchema = field(default_factory=UserInputSchema) def __post_init__(self) -> None: if self.max_lag_s is not None and ( diff --git a/flashdreams/flashdreams/runtime/demo/validation.py b/flashdreams/flashdreams/runtime/demo/validation.py new file mode 100644 index 000000000..27b0ae757 --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/validation.py @@ -0,0 +1,201 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Capability resolution and validation for shared demo runs.""" + +from __future__ import annotations + +from collections.abc import Sequence +from dataclasses import dataclass + +from flashdreams.runtime.inputs import UserInputCapability, UserInputSchema + +from .run_modes import RunMode, RunModeCapabilities, SessionEdges +from .session_inputs import ( + BatchInputSource, + ModelInputProvider, + ProviderCapabilities, + RealtimeInputSource, +) +from .spec import DemoAdapter, DemoSpec + + +@dataclass(frozen=True, kw_only=True, slots=True) +class ResolvedRunCapabilities: + """Capabilities of one concrete provider/run-mode/session-edges pairing.""" + + finite: bool + deterministic: bool + realtime: bool + resettable: bool + produces_artifacts: bool + + +def resolve_run_capabilities( + *, + spec: DemoSpec, + provider: ModelInputProvider, + session_edges: SessionEdges, +) -> ResolvedRunCapabilities: + """Resolve concrete run capabilities from provider, edges, and config.""" + + provider_capabilities = _provider_capabilities(provider) + clock = session_edges.clock + realtime = bool(getattr(clock, "is_realtime", False)) + deterministic_clock = ( + bool(getattr(clock, "is_deterministic", False)) if clock is not None else True + ) + config = spec.config + seeded = config is not None and config.seed is not None + return ResolvedRunCapabilities( + finite=bool(session_edges.input_source.is_finite), + deterministic=( + provider_capabilities.deterministic_given_inputs + and bool(session_edges.input_source.is_deterministic) + and deterministic_clock + and seeded + ), + realtime=realtime, + resettable=provider_capabilities.supports_reset, + produces_artifacts=bool(session_edges.output_sink.produces_artifacts), + ) + + +def validate_resolved_run( + *, + spec: DemoSpec, + adapter: DemoAdapter, + provider: ModelInputProvider, + run_mode: RunMode, + session_edges: SessionEdges, + resolved: ResolvedRunCapabilities, +) -> None: + """Reject structurally incompatible provider/input/run-mode combinations.""" + + del spec, adapter + provider_capabilities = _provider_capabilities(provider) + run_mode_capabilities = _run_mode_capabilities(run_mode) + _validate_input_source_shape( + run_mode_capabilities=run_mode_capabilities, + session_edges=session_edges, + resolved=resolved, + ) + _validate_provider_modes( + provider_capabilities=provider_capabilities, + run_mode_capabilities=run_mode_capabilities, + resolved=resolved, + ) + _validate_user_input_schema( + provider_schema=provider_capabilities.user_input_schema, + source_schema=_input_source_user_input_schema(session_edges.input_source), + ) + + +def _validate_input_source_shape( + *, + run_mode_capabilities: RunModeCapabilities, + session_edges: SessionEdges, + resolved: ResolvedRunCapabilities, +) -> None: + input_source = session_edges.input_source + if run_mode_capabilities.realtime: + if not isinstance(input_source, RealtimeInputSource): + raise ValueError("Realtime run modes require a RealtimeInputSource.") + if session_edges.clock is None or not resolved.realtime: + raise ValueError("Realtime run modes require a realtime clock.") + return + if not isinstance(input_source, BatchInputSource): + raise ValueError("Batch run modes require a BatchInputSource.") + if resolved.realtime: + raise ValueError("Batch run modes cannot use a realtime clock.") + + +def _validate_provider_modes( + *, + provider_capabilities: ProviderCapabilities, + run_mode_capabilities: RunModeCapabilities, + resolved: ResolvedRunCapabilities, +) -> None: + if run_mode_capabilities.realtime and not ( + provider_capabilities.supports_realtime_clock + ): + raise ValueError("Provider does not support realtime input.") + if run_mode_capabilities.requires_finite_input: + if not resolved.finite: + raise ValueError("Run mode requires finite input.") + if not provider_capabilities.supports_recorded_input: + raise ValueError("Provider does not support recorded input.") + if resolved.produces_artifacts and not run_mode_capabilities.supports_artifacts: + raise ValueError("Run mode does not support artifact output.") + if ( + run_mode_capabilities.supports_interactive_events + and not provider_capabilities.supports_realtime_clock + ): + raise ValueError("Interactive run mode requires realtime provider support.") + + +def _validate_user_input_schema( + *, + provider_schema: UserInputSchema, + source_schema: UserInputSchema, +) -> None: + missing = _missing_capabilities( + required=provider_schema.declared_capabilities(), + provided=source_schema, + ) + if missing: + names = ", ".join( + f"{capability.event_type}[{','.join(sorted(capability.payload_fields))}]" + for capability in missing + ) + raise ValueError( + "Input source does not satisfy provider raw user input schema: " + f"{names}." + ) + + +def _missing_capabilities( + *, + required: Sequence[UserInputCapability], + provided: UserInputSchema, +) -> tuple[UserInputCapability, ...]: + return tuple( + capability for capability in required if not provided.supports(capability) + ) + + +def _provider_capabilities(provider: ModelInputProvider) -> ProviderCapabilities: + capabilities = getattr(provider, "capabilities", None) + if not isinstance(capabilities, ProviderCapabilities): + raise TypeError( + "ModelInputProvider.capabilities must be a ProviderCapabilities " + f"instance, got {type(capabilities).__name__}." + ) + return capabilities + + +def _run_mode_capabilities(run_mode: RunMode) -> RunModeCapabilities: + capabilities = getattr(run_mode, "capabilities", None) + if not isinstance(capabilities, RunModeCapabilities): + raise TypeError( + "RunMode.capabilities must be a RunModeCapabilities instance, " + f"got {type(capabilities).__name__}." + ) + return capabilities + + +def _input_source_user_input_schema(input_source: object) -> UserInputSchema: + schema = getattr(input_source, "user_input_schema", None) + if not isinstance(schema, UserInputSchema): + raise TypeError( + "InputSource.user_input_schema must be a UserInputSchema instance, " + f"got {type(schema).__name__}." + ) + return schema + + +__all__ = [ + "ResolvedRunCapabilities", + "resolve_run_capabilities", + "validate_resolved_run", +] diff --git a/flashdreams/tests/test_demo_runtime_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py index 57d2a2550..26974ed98 100644 --- a/flashdreams/tests/test_demo_runtime_realtime_driver.py +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -16,6 +16,7 @@ StepRequest, StepRequirements, StepResult, + UserInputSchema, ) from flashdreams.runtime.demo import ( ActivationResult, @@ -24,6 +25,7 @@ InMemorySessionMetricsRecorder, OutputDecision, PreparedStep, + ProviderCapabilities, RealtimeSessionDriver, RealtimeWindowResult, RunContext, @@ -529,6 +531,7 @@ def catch_up(self, **kwargs: Any) -> object: class _RealtimeInputSource: is_finite = False is_deterministic = False + user_input_schema = UserInputSchema() def __init__( self, @@ -563,6 +566,12 @@ async def next_realtime_window( class _FakeRealtimeProvider: + capabilities = ProviderCapabilities( + supports_realtime_clock=True, + supports_reset=True, + deterministic_given_inputs=False, + ) + def __init__( self, *, diff --git a/flashdreams/tests/test_demo_runtime_run_modes.py b/flashdreams/tests/test_demo_runtime_run_modes.py index d5a6e4f04..8dd6e26af 100644 --- a/flashdreams/tests/test_demo_runtime_run_modes.py +++ b/flashdreams/tests/test_demo_runtime_run_modes.py @@ -20,6 +20,8 @@ InferenceRuntime, InferenceSession, InputMapping, + StepRequirements, + UserInputSchema, ) from flashdreams.runtime.demo import ( AsyncSessionDriver, @@ -31,13 +33,16 @@ NullOutputSpec, OutputDecision, PreparedScenario, + ProviderCapabilities, RunContext, + RunModeCapabilities, RunResult, RuntimeHost, SessionDriver, SessionEdges, SessionInfo, StepPipeline, + UserInputWindow, WebRTCOutputSpec, run_demo_session, run_demo_session_async, @@ -380,6 +385,12 @@ def create_model_input_provider( class _FakeProvider: + capabilities = ProviderCapabilities( + supports_recorded_input=True, + supports_reset=True, + deterministic_given_inputs=True, + ) + def __init__(self) -> None: self.close_count = 0 @@ -401,6 +412,7 @@ def __init__( self.created_edges: list[SessionEdges] = [] self.validate_run_count = 0 self.warmup_count = 0 + self.capabilities = RunModeCapabilities(requires_finite_input=True) self.admission = admission or _RecordingAdmission(events=[]) self.services = services or {} self.context: RunContext | None = None @@ -560,10 +572,15 @@ def release(self) -> None: class _FinishedInputSource: is_finite = True is_deterministic = True + user_input_schema = UserInputSchema() def is_finished(self) -> bool: return True + def next_window(self, request: StepRequirements) -> UserInputWindow: + del request + return UserInputWindow(start_s=0.0, end_s=0.0) + class _RecordingOutputSink: produces_artifacts = False diff --git a/flashdreams/tests/test_demo_runtime_timing.py b/flashdreams/tests/test_demo_runtime_timing.py index a30c2752f..98e3cf56d 100644 --- a/flashdreams/tests/test_demo_runtime_timing.py +++ b/flashdreams/tests/test_demo_runtime_timing.py @@ -8,7 +8,7 @@ import pytest -from flashdreams.runtime import StepRequirements +from flashdreams.runtime import StepRequirements, UserInputSchema from flashdreams.runtime.demo import NullOutputSink, RunResult, SessionEdges from flashdreams.runtime.demo.timing import ( SPARSE_KEY_SEGMENTS_METADATA_KEY, @@ -222,6 +222,7 @@ async def __call__(self, delay_s: float) -> None: class _OpenRealtimeInputSource: is_finite = False is_deterministic = False + user_input_schema = UserInputSchema() def is_finished(self) -> bool: return False diff --git a/flashdreams/tests/test_demo_runtime_validation.py b/flashdreams/tests/test_demo_runtime_validation.py new file mode 100644 index 000000000..0dbe7505a --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_validation.py @@ -0,0 +1,675 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from collections.abc import Mapping, Sequence +from typing import Any + +import pytest + +import flashdreams.runtime.demo as demo_api +from flashdreams.runtime import ( + DRIVER_COMMAND, + CanonicalInputSchema, + CanonicalInputs, + IdentityInputMapping, + InferenceConfig, + InferenceInput, + InferenceInputSchema, + InferenceRuntime, + InputCanonicalizer, + InputField, + InputMapping, + InputMappingSchema, + KeyboardToDriverCommand, + OutputArtifact, + StepRequest, + StepRequirements, + StepResult, + TimeWindow, + UserInputCapability, + UserInputEvent, + UserInputSchema, + UserInputs, +) +from flashdreams.runtime.demo import ( + DemoSpec, + Mp4OutputSpec, + NullOutputSpec, + OutputDecision, + PreparedScenario, + PreparedStep, + ProviderCapabilities, + ResolvedRunCapabilities, + RunContext, + RunModeCapabilities, + RunResult, + RuntimeHost, + SessionEdges, + SessionInfo, + UserInputWindow, + resolve_run_capabilities, + validate_resolved_run, +) +from flashdreams.runtime.demo.timing import RealtimeWindowResult + +pytestmark = pytest.mark.ci_cpu + +KEY_SCHEMA = UserInputSchema( + capabilities=( + UserInputCapability( + event_type="key_down", + payload_fields=frozenset({"key"}), + ), + UserInputCapability( + event_type="key_up", + payload_fields=frozenset({"key"}), + ), + ) +) + + +def test_provider_capabilities_declare_raw_and_inference_schemas() -> None: + capabilities = ProviderCapabilities( + supports_recorded_input=True, + deterministic_given_inputs=True, + user_input_schema=KEY_SCHEMA, + inference_input_schema=InferenceInputSchema( + step_fields=(InputField(name="driver_command"),), + ), + ) + + assert capabilities.user_input_schema.supports( + UserInputCapability(event_type="key_down", payload_fields=frozenset({"key"})) + ) + assert capabilities.inference_input_schema.missing_step(InferenceInput()) == ( + "driver_command", + ) + + +def test_provider_that_wraps_mapping_canonicalizes_raw_inputs_first() -> None: + mapping = _DriverCommandMapping() + provider = _MappingBackedProvider(mapping=mapping, source_schema=KEY_SCHEMA) + source = _BatchInputSource(user_input_schema=KEY_SCHEMA) + edges = _edges(input_source=source) + run_mode = _RunMode( + RunModeCapabilities(requires_finite_input=True, supports_artifacts=True) + ) + resolved = resolve_run_capabilities( + spec=_spec(seed=7, output=Mp4OutputSpec(path="out.mp4", fps=12)), + provider=provider, + session_edges=edges, + ) + + validate_resolved_run( + spec=_spec(seed=7), + adapter=_Adapter(), + provider=provider, + run_mode=run_mode, + session_edges=edges, + resolved=resolved, + ) + prepared = provider.prepare_step( + request=StepRequirements(step_index=4), + user_window=UserInputWindow( + start_s=0.0, + end_s=1.0, + inputs=UserInputs( + events=( + UserInputEvent( + timestamp_s=0.1, + event_type="key_down", + payload={"key": "w"}, + ), + ) + ), + ), + ) + + assert mapping.validated_with is not None + assert mapping.validated_with[0] == CanonicalInputSchema( + modalities=(DRIVER_COMMAND,), + description=KEY_SCHEMA.description, + ) + assert prepared.inference_input is not None + assert prepared.inference_input.step["request_step"] == 4 + assert prepared.inference_input.step["driver_command"]["throttle"] == 1.0 + + +def test_raw_user_schema_validation_is_not_canonical_schema_validation() -> None: + provider = _Provider( + ProviderCapabilities( + supports_recorded_input=True, + user_input_schema=KEY_SCHEMA, + inference_input_schema=InferenceInputSchema( + step_fields=(InputField(name="driver_command"),), + ), + ) + ) + source = _BatchInputSource( + user_input_schema=UserInputSchema(event_types=frozenset({"key_down"})) + ) + run_mode = _RunMode( + RunModeCapabilities(requires_finite_input=True, supports_artifacts=True) + ) + edges = _edges(input_source=source) + resolved = resolve_run_capabilities( + spec=_spec(seed=1), + provider=provider, + session_edges=edges, + ) + + with pytest.raises(ValueError, match="raw user input schema"): + validate_resolved_run( + spec=_spec(seed=1), + adapter=_Adapter(), + provider=provider, + run_mode=run_mode, + session_edges=edges, + resolved=resolved, + ) + + +def test_mp4_rejects_provider_without_recorded_input_support() -> None: + provider = _Provider( + ProviderCapabilities( + supports_recorded_input=False, + supports_realtime_clock=True, + ) + ) + run_mode = _RunMode( + RunModeCapabilities(requires_finite_input=True, supports_artifacts=True) + ) + edges = _edges(output_sink=_OutputSink(produces_artifacts=True)) + resolved = resolve_run_capabilities( + spec=_spec(seed=1, output=Mp4OutputSpec(path="out.mp4", fps=12)), + provider=provider, + session_edges=edges, + ) + + with pytest.raises(ValueError, match="recorded input"): + validate_resolved_run( + spec=_spec(seed=1), + adapter=_Adapter(), + provider=provider, + run_mode=run_mode, + session_edges=edges, + resolved=resolved, + ) + + +def test_webrtc_rejects_provider_without_realtime_input_support() -> None: + provider = _Provider(ProviderCapabilities(supports_recorded_input=True)) + run_mode = _RunMode( + RunModeCapabilities( + realtime=True, + supports_backpressure=True, + supports_interactive_events=True, + ) + ) + edges = _edges(input_source=_RealtimeInputSource(), clock=_Clock(realtime=True)) + resolved = resolve_run_capabilities( + spec=_spec(seed=1), + provider=provider, + session_edges=edges, + ) + + with pytest.raises(ValueError, match="realtime input"): + validate_resolved_run( + spec=_spec(seed=1), + adapter=_Adapter(), + provider=provider, + run_mode=run_mode, + session_edges=edges, + resolved=resolved, + ) + + +def test_realtime_run_mode_rejects_batch_input_source() -> None: + provider = _Provider(ProviderCapabilities(supports_realtime_clock=True)) + run_mode = _RunMode(RunModeCapabilities(realtime=True)) + edges = _edges(input_source=_BatchInputSource(), clock=_Clock(realtime=True)) + resolved = resolve_run_capabilities( + spec=_spec(seed=1), + provider=provider, + session_edges=edges, + ) + + with pytest.raises(ValueError, match="RealtimeInputSource"): + validate_resolved_run( + spec=_spec(seed=1), + adapter=_Adapter(), + provider=provider, + run_mode=run_mode, + session_edges=edges, + resolved=resolved, + ) + + +def test_determinism_resolves_from_provider_source_clock_and_seed() -> None: + provider = _Provider( + ProviderCapabilities( + supports_recorded_input=True, + supports_reset=True, + deterministic_given_inputs=True, + ) + ) + + deterministic = resolve_run_capabilities( + spec=_spec(seed=123), + provider=provider, + session_edges=_edges(clock=_Clock(deterministic=True)), + ) + unseeded = resolve_run_capabilities( + spec=_spec(seed=None), + provider=provider, + session_edges=_edges(clock=_Clock(deterministic=True)), + ) + nondeterministic_source = resolve_run_capabilities( + spec=_spec(seed=123), + provider=provider, + session_edges=_edges( + input_source=_BatchInputSource(deterministic=False), + clock=_Clock(deterministic=True), + ), + ) + + assert deterministic == ResolvedRunCapabilities( + finite=True, + deterministic=True, + realtime=False, + resettable=True, + produces_artifacts=True, + ) + assert not unseeded.deterministic + assert not nondeterministic_source.deterministic + + +def test_no_general_purpose_input_mapping_provider_is_exported() -> None: + assert not hasattr(demo_api, "InputMappingProvider") + + +def test_reset_control_updates_provider_and_session_together() -> None: + reset_input = InferenceInput(global_conditioning={"prompt": "reset"}) + provider = _ResettingProvider(reset_input=reset_input) + session = _ResettableSession(num_steps=2) + edges = _edges(input_source=_BatchInputSource(num_windows=2)) + + result = demo_api.BatchSessionDriver().run_one_session( + host=RuntimeHost(_Runtime(session=session)), + provider=provider, + session_edges=edges, + pipeline=demo_api.StepPipeline(), + ) + + assert result.status == "completed" + assert session.reset_inputs == [reset_input] + assert provider.reset_inputs == [reset_input] + assert len(session.step_inputs) == 1 + + +def _spec( + *, + seed: int | None, + output: Any | None = None, +) -> DemoSpec: + return DemoSpec( + model_id="fake-demo", + input_mode="replay", + output=output or NullOutputSpec(), + config=InferenceConfig(model_id="fake-demo", seed=seed), + ) + + +def _edges( + *, + input_source: Any | None = None, + output_sink: Any | None = None, + clock: Any | None = None, +) -> SessionEdges: + return SessionEdges( + input_source=input_source or _BatchInputSource(), + output_sink=output_sink or _OutputSink(produces_artifacts=True), + cleanup_tasks=set(), + clock=clock, + ) + + +class _Adapter: + model_id = "fake-demo" + inference_input_schema = InferenceInputSchema() + canonical_input_schema = None + + def supported_input_modes(self) -> tuple[str, ...]: + return ("replay",) + + def supported_output_modes(self) -> tuple[str, ...]: + return ("null", "mp4") + + def default_input_mapping(self) -> InputMapping: + return IdentityInputMapping() + + def validate_config(self, config: InferenceConfig) -> None: + del config + + def create_runtime(self, config: InferenceConfig) -> InferenceRuntime: + del config + raise NotImplementedError + + def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: + del spec + return PreparedScenario(initial_inputs=InferenceInput()) + + +class _Provider: + def __init__(self, capabilities: ProviderCapabilities) -> None: + self.capabilities = capabilities + + def prepare_initial_input(self) -> InferenceInput: + return InferenceInput() + + def prepare_step( + self, + *, + request: StepRequirements, + user_window: UserInputWindow, + ) -> PreparedStep: + del request, user_window + return PreparedStep(inference_input=InferenceInput()) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + + def close(self) -> None: + return + + +class _MappingBackedProvider(_Provider): + def __init__(self, *, mapping: "_DriverCommandMapping", source_schema: UserInputSchema) -> None: + self.mapping = mapping + self.canonicalizer = InputCanonicalizer((KeyboardToDriverCommand(),)) + self.source_schema = source_schema + capabilities = ProviderCapabilities( + supports_recorded_input=True, + supports_reset=True, + deterministic_given_inputs=True, + user_input_schema=source_schema, + inference_input_schema=InferenceInputSchema( + step_fields=(InputField(name="driver_command"),), + ), + ) + super().__init__(capabilities) + self.mapping.validate( + canonical_schema=self.canonicalizer.canonical_schema(source_schema), + inference_input_schema=capabilities.inference_input_schema, + ) + + def prepare_step( + self, + *, + request: StepRequirements, + user_window: UserInputWindow, + ) -> PreparedStep: + canonical_inputs = self.canonicalizer.canonicalize( + user_window.inputs, + window=TimeWindow(start_s=user_window.start_s, end_s=user_window.end_s), + source_schema=self.source_schema, + ) + inference_input = self.mapping.map_step_inputs( + canonical_inputs=canonical_inputs, + inference_input=InferenceInput(), + request=StepRequest(step_index=request.step_index), + ) + return PreparedStep(inference_input=inference_input) + + +class _DriverCommandMapping: + mapping_schema = InputMappingSchema( + name="driver-command", + consumes=(DRIVER_COMMAND,), + produces_step=(InputField(name="driver_command"),), + ) + + def __init__(self) -> None: + self.validated_with: tuple[ + CanonicalInputSchema | None, + InferenceInputSchema | None, + ] | None = None + + def validate( + self, + *, + canonical_schema: CanonicalInputSchema | None = None, + inference_input_schema: InferenceInputSchema | None = None, + ) -> None: + if canonical_schema is not None and not canonical_schema.supports( + DRIVER_COMMAND + ): + raise ValueError("mapping cannot be fed") + if inference_input_schema is not None: + inference_input_schema.require_step( + InferenceInput(step={"driver_command": object()}) + ) + self.validated_with = (canonical_schema, inference_input_schema) + + def map_global_conditioning_inputs( + self, + *, + canonical_inputs: CanonicalInputs, + inference_input: InferenceInput, + ) -> InferenceInput: + del canonical_inputs + return inference_input + + def map_step_inputs( + self, + *, + canonical_inputs: CanonicalInputs, + inference_input: InferenceInput, + request: StepRequest, + ) -> InferenceInput: + del inference_input + return InferenceInput( + step={ + "driver_command": canonical_inputs.values["driver_command"], + "request_step": request.step_index, + } + ) + + +class _RunMode: + name = "fake" + + def __init__(self, capabilities: RunModeCapabilities) -> None: + self.capabilities = capabilities + + def validate_run(self, *, spec: DemoSpec, adapter: Any) -> None: + del spec, adapter + + def validate_session( + self, + *, + spec: DemoSpec, + scenario: PreparedScenario, + adapter: Any, + provider: Any, + ) -> None: + del spec, scenario, adapter, provider + + def create_run_context( + self, + *, + spec: DemoSpec, + adapter: Any, + host: RuntimeHost, + model_warmup_plan: Any, + ) -> RunContext: + del spec, adapter, model_warmup_plan + return RunContext( + host=host, + run_metrics=demo_api.InMemorySessionMetricsRecorder(), + admission=demo_api.SingleSessionAdmissionPolicy(), + ) + + def create_session_edges( + self, + *, + context: RunContext, + spec: DemoSpec, + scenario: PreparedScenario, + provider: Any, + adapter: Any, + ) -> SessionEdges: + del context, spec, scenario, provider, adapter + return _edges() + + def select_driver(self) -> Any: + raise NotImplementedError + + +class _BatchInputSource: + is_finite = True + + def __init__( + self, + *, + user_input_schema: UserInputSchema | None = None, + deterministic: bool = True, + num_windows: int = 1, + ) -> None: + self.user_input_schema = user_input_schema or UserInputSchema() + self.is_deterministic = deterministic + self.num_windows = num_windows + self.index = 0 + + def is_finished(self) -> bool: + return self.index >= self.num_windows + + def next_window(self, request: StepRequirements) -> UserInputWindow: + del request + self.index += 1 + return UserInputWindow(start_s=0.0, end_s=1.0) + + +class _RealtimeInputSource: + is_finite = False + is_deterministic = False + user_input_schema = UserInputSchema() + + def is_finished(self) -> bool: + return False + + async def next_realtime_window( + self, + *, + request: StepRequirements, + clock: Any, + ) -> RealtimeWindowResult: + del request, clock + return RealtimeWindowResult(window=UserInputWindow(start_s=0.0, end_s=1.0)) + + +class _Clock: + def __init__(self, *, realtime: bool = False, deterministic: bool = True) -> None: + self.is_realtime = realtime + self.is_deterministic = deterministic + + +class _OutputSink: + def __init__(self, *, produces_artifacts: bool) -> None: + self.produces_artifacts = produces_artifacts + + def open(self, session_info: SessionInfo) -> None: + del session_info + + def begin_generation(self, generation: int) -> None: + del generation + + def write(self, result: StepResult) -> OutputDecision: + del result + return OutputDecision() + + def close(self) -> Sequence[OutputArtifact]: + return () + + +class _ResettingProvider(_Provider): + def __init__(self, *, reset_input: InferenceInput) -> None: + self.reset_input = reset_input + self.prepare_count = 0 + self.reset_inputs: list[InferenceInput | None] = [] + super().__init__( + ProviderCapabilities( + supports_recorded_input=True, + supports_reset=True, + deterministic_given_inputs=True, + ) + ) + + def prepare_step( + self, + *, + request: StepRequirements, + user_window: UserInputWindow, + ) -> PreparedStep: + del request, user_window + self.prepare_count += 1 + if self.prepare_count == 1: + return PreparedStep( + control=demo_api.ControlDecision( + reset=True, + reset_input=self.reset_input, + ) + ) + return PreparedStep(inference_input=InferenceInput(step={"after_reset": True})) + + def reset(self, inputs: InferenceInput | None = None) -> None: + self.reset_inputs.append(inputs) + + +class _Runtime: + def __init__(self, *, session: "_ResettableSession") -> None: + self.session = session + + def start_session(self, inputs: InferenceInput) -> "_ResettableSession": + del inputs + return self.session + + def close(self) -> None: + return + + +class _ResettableSession: + def __init__(self, *, num_steps: int) -> None: + self.num_steps = num_steps + self.next_request_index = 0 + self.reset_inputs: list[InferenceInput | None] = [] + self.step_inputs: list[InferenceInput] = [] + + def session_info(self) -> SessionInfo: + return SessionInfo() + + def next_step_requirements(self) -> StepRequirements | None: + if self.next_request_index >= self.num_steps: + return None + request = StepRequirements(step_index=self.next_request_index) + self.next_request_index += 1 + return request + + def next_step_request(self) -> StepRequest | None: + requirements = self.next_step_requirements() + if requirements is None: + return None + return StepRequest(step_index=requirements.step_index) + + def step(self, inputs: InferenceInput) -> StepResult: + self.step_inputs.append(inputs) + return StepResult(step_index=len(self.step_inputs) - 1, output=None) + + def reset(self, inputs: InferenceInput | None = None) -> None: + self.reset_inputs.append(inputs) + self.next_request_index = 0 + + def close(self) -> None: + return diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index d7f4bb492..f48966c16 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -22,6 +22,7 @@ StepRequirements, StepResult, UserInputs, + UserInputSchema, ) from flashdreams.runtime.demo import ( BatchSessionDriver, @@ -35,7 +36,9 @@ OutputDecision, PreparedScenario, PreparedStep, + ProviderCapabilities, RunContext, + RunModeCapabilities, RunResult, RuntimeHost, SessionEdges, @@ -664,6 +667,12 @@ def _record_output_factory_call( class _FakeVideoModelInputProvider: + capabilities = ProviderCapabilities( + supports_recorded_input=True, + supports_reset=True, + deterministic_given_inputs=True, + ) + def __init__(self, *, fail_initial: Exception | None = None) -> None: self.fail_initial = fail_initial self.initial_input = InferenceInput( @@ -703,6 +712,7 @@ def close(self) -> None: class _FakeBatchInputSource: is_finite = True is_deterministic = True + user_input_schema = UserInputSchema() def __init__( self, @@ -730,6 +740,7 @@ def next_window(self, request: StepRequirements) -> UserInputWindow: class _SlicingBatchInputSource: is_finite = True is_deterministic = True + user_input_schema = UserInputSchema() def __init__(self, *, fps: float, num_windows: int) -> None: self.fps = fps @@ -984,6 +995,10 @@ def __init__( self.transport = transport self.error_policy = error_policy self.validate_error = validate_error + self.capabilities = RunModeCapabilities( + requires_finite_input=True, + supports_artifacts=True, + ) def validate_run( self, diff --git a/flashdreams/tests/test_inference_runtime_api.py b/flashdreams/tests/test_inference_runtime_api.py index da00789c3..d1235c905 100644 --- a/flashdreams/tests/test_inference_runtime_api.py +++ b/flashdreams/tests/test_inference_runtime_api.py @@ -40,11 +40,13 @@ def test_inference_config_keeps_runtime_settings_separate() -> None: backend="local", precision="bf16", compile=False, + seed=123, runtime_options={"chunk_size": 3}, ) assert config.model_id == "lingbot-world" assert config.preset_id == "fast-taehv" + assert config.seed == 123 assert config.runtime_options["chunk_size"] == 3 assert denied_app_fields.isdisjoint(field.name for field in fields(InferenceConfig)) with pytest.raises(TypeError): @@ -56,6 +58,13 @@ def test_inference_config_rejects_empty_model_id() -> None: InferenceConfig(model_id=" ") +def test_inference_config_rejects_invalid_seed() -> None: + with pytest.raises(TypeError, match="seed"): + InferenceConfig(model_id="fake", seed=True) + with pytest.raises(ValueError, match="seed"): + InferenceConfig(model_id="fake", seed=-1) + + @pytest.mark.parametrize( ("factory", "match"), [ From b16920ee519653d371fcb309b9eceefd6f451585 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 06:16:49 +0000 Subject: [PATCH 16/51] runtime: avoid serving dependency in demo timing Replace the demo timing module's concrete KeyboardResampler import with a structural resampler protocol, keeping runtime/demo independent of new serving imports. Also include Ruff formatting fixes needed by the PR lint check. --- .../flashdreams/runtime/demo/timing.py | 17 +++++++++++++++-- .../flashdreams/runtime/demo/validation.py | 3 +-- .../tests/test_demo_runtime_validation.py | 19 ++++++++++++------- 3 files changed, 28 insertions(+), 11 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/timing.py b/flashdreams/flashdreams/runtime/demo/timing.py index 9c57237ea..d48e1e213 100644 --- a/flashdreams/flashdreams/runtime/demo/timing.py +++ b/flashdreams/flashdreams/runtime/demo/timing.py @@ -14,7 +14,6 @@ from flashdreams.runtime.inputs import UserInputs, UserInputSchema from flashdreams.runtime.types import StepRequirements -from flashdreams.serving.realtime.input import KeyboardResampler from .session_inputs import UserInputWindow @@ -182,11 +181,25 @@ async def wait_until_active( await asyncio.gather(*tasks, return_exceptions=True) +SparseKeySegment = tuple[float, float, frozenset[str]] + + class _RealtimeTimeline(Protocol): dt: float next_chunk_start_v: float +class _SparseInputResampler(_RealtimeTimeline, Protocol): + def on_edge(self, *, arrival_t: float, event: str, key: str) -> None: ... + + def reset(self, *, start_v: float) -> None: ... + + def sample_chunk( + self, + num_frames: int, + ) -> tuple[Sequence[SparseKeySegment], Sequence[float]]: ... + + @dataclass(slots=True) class ResamplerRealtimeClock: """Realtime clock that reuses ``KeyboardResampler``'s virtual timeline.""" @@ -268,7 +281,7 @@ def catch_up( class KeyboardRealtimeInputSource: """Realtime input source backed by the existing keyboard resampler.""" - resampler: KeyboardResampler + resampler: _SparseInputResampler max_lag_s: float | None = None catch_up_policy: CatchUpPolicy = "fold" is_finite: bool = False diff --git a/flashdreams/flashdreams/runtime/demo/validation.py b/flashdreams/flashdreams/runtime/demo/validation.py index 27b0ae757..c5058b148 100644 --- a/flashdreams/flashdreams/runtime/demo/validation.py +++ b/flashdreams/flashdreams/runtime/demo/validation.py @@ -149,8 +149,7 @@ def _validate_user_input_schema( for capability in missing ) raise ValueError( - "Input source does not satisfy provider raw user input schema: " - f"{names}." + f"Input source does not satisfy provider raw user input schema: {names}." ) diff --git a/flashdreams/tests/test_demo_runtime_validation.py b/flashdreams/tests/test_demo_runtime_validation.py index 0dbe7505a..c5780bd72 100644 --- a/flashdreams/tests/test_demo_runtime_validation.py +++ b/flashdreams/tests/test_demo_runtime_validation.py @@ -11,8 +11,8 @@ import flashdreams.runtime.demo as demo_api from flashdreams.runtime import ( DRIVER_COMMAND, - CanonicalInputSchema, CanonicalInputs, + CanonicalInputSchema, IdentityInputMapping, InferenceConfig, InferenceInput, @@ -30,8 +30,8 @@ TimeWindow, UserInputCapability, UserInputEvent, - UserInputSchema, UserInputs, + UserInputSchema, ) from flashdreams.runtime.demo import ( DemoSpec, @@ -386,7 +386,9 @@ def close(self) -> None: class _MappingBackedProvider(_Provider): - def __init__(self, *, mapping: "_DriverCommandMapping", source_schema: UserInputSchema) -> None: + def __init__( + self, *, mapping: "_DriverCommandMapping", source_schema: UserInputSchema + ) -> None: self.mapping = mapping self.canonicalizer = InputCanonicalizer((KeyboardToDriverCommand(),)) self.source_schema = source_schema @@ -432,10 +434,13 @@ class _DriverCommandMapping: ) def __init__(self) -> None: - self.validated_with: tuple[ - CanonicalInputSchema | None, - InferenceInputSchema | None, - ] | None = None + self.validated_with: ( + tuple[ + CanonicalInputSchema | None, + InferenceInputSchema | None, + ] + | None + ) = None def validate( self, From a9d38604cb5ec011949801ac8c3d21cc2f66f42f Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 06:41:11 +0000 Subject: [PATCH 17/51] runtime: split demo warmup setup Add an optional adapter-owned model warmup hook, worker-dispatched warmup plan construction, and a shared run-context warmup helper that separates model runtime warmup from run-mode output/transport warmup. Cover temporary providers, worker-thread input construction, runtime/session warmup calls, transport-only warmup, and metrics exclusion with fake CPU tests. --- .../flashdreams/runtime/demo/__init__.py | 6 + .../flashdreams/runtime/demo/run_modes.py | 65 ++- flashdreams/flashdreams/runtime/demo/spec.py | 17 +- flashdreams/tests/test_demo_runtime_warmup.py | 416 ++++++++++++++++++ 4 files changed, 502 insertions(+), 2 deletions(-) create mode 100644 flashdreams/tests/test_demo_runtime_warmup.py diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index cefcba4a2..eedeeeb84 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -45,6 +45,8 @@ SessionDriver, SessionEdges, SingleSessionAdmissionPolicy, + build_model_warmup_plan, + warmup_run_context, ) from flashdreams.runtime.demo.session_inputs import ( BatchInputSource, @@ -59,6 +61,7 @@ from flashdreams.runtime.demo.spec import ( DemoAdapter, DemoSpec, + ModelWarmupAdapter, Mp4OutputSpec, NullOutputSpec, OutputSpec, @@ -109,6 +112,7 @@ "CatchUpPolicy", "DeterministicClock", "MetricsSnapshot", + "ModelWarmupAdapter", "ModelWarmupPlan", "ModelInputProvider", "KeyboardRealtimeInputSource", @@ -150,6 +154,7 @@ "WebRTCOutputSpec", "build_output_sink", "build_output_target", + "build_model_warmup_plan", "input_frame_count_from_request", "resolve_run_capabilities", "run_demo_session", @@ -158,4 +163,5 @@ "shielded_session_cleanup", "uncancel_current_task", "validate_resolved_run", + "warmup_run_context", ] diff --git a/flashdreams/flashdreams/runtime/demo/run_modes.py b/flashdreams/flashdreams/runtime/demo/run_modes.py index 78c7b837b..aa2331605 100644 --- a/flashdreams/flashdreams/runtime/demo/run_modes.py +++ b/flashdreams/flashdreams/runtime/demo/run_modes.py @@ -14,7 +14,7 @@ from flashdreams.runtime._utils import freeze_mapping from flashdreams.runtime.output import OutputArtifact -from .host import ModelWarmupPlan +from .host import ModelWarmupPlan, WarmupSessionInputs from .outputs import OutputDecision, OutputSink if TYPE_CHECKING: @@ -494,6 +494,67 @@ def warmup_context( ) -> None: ... +def build_model_warmup_plan( + *, + host: "RuntimeHost", + adapter: "DemoAdapter", + spec: "DemoSpec", + scenario: "PreparedScenario", +) -> ModelWarmupPlan: + """Build a host-owned warmup plan through the model-affine worker.""" + + create_sessions = getattr(adapter, "create_model_warmup_sessions", None) + if create_sessions is None: + return ModelWarmupPlan() + if not callable(create_sessions): + raise TypeError( + "Demo adapter create_model_warmup_sessions attribute must be callable." + ) + sessions = host.call(create_sessions, spec, scenario) + return ModelWarmupPlan(sessions=_coerce_warmup_sessions(sessions)) + + +def warmup_run_context( + *, + context: RunContext, + spec: "DemoSpec", + scenario: "PreparedScenario", + adapter: "DemoAdapter", + run_mode: object, +) -> None: + """Run model warmup, then optional output/transport warmup for a context.""" + + context.host.warmup(context.model_warmup_plan) + warmup_context = getattr(run_mode, "warmup_context", None) + if warmup_context is None: + return + if not callable(warmup_context): + raise TypeError("RunMode.warmup_context attribute must be callable.") + warmup_context( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + ) + + +def _coerce_warmup_sessions(value: object) -> tuple[WarmupSessionInputs, ...]: + if not isinstance(value, Sequence): + raise TypeError( + "Demo adapter create_model_warmup_sessions(...) must return a sequence " + f"of WarmupSessionInputs, got {type(value).__name__}." + ) + sessions: list[WarmupSessionInputs] = [] + for session in value: + if not isinstance(session, WarmupSessionInputs): + raise TypeError( + "Demo adapter create_model_warmup_sessions(...) must return only " + f"WarmupSessionInputs, got {type(session).__name__}." + ) + sessions.append(session) + return tuple(sessions) + + __all__ = [ "AdmissionPolicy", "AsyncSessionDriver", @@ -517,4 +578,6 @@ def warmup_context( "SessionStatus", "SingleSessionAdmissionPolicy", "TransportService", + "build_model_warmup_plan", + "warmup_run_context", ] diff --git a/flashdreams/flashdreams/runtime/demo/spec.py b/flashdreams/flashdreams/runtime/demo/spec.py index 6ba652f38..d7a6a309b 100644 --- a/flashdreams/flashdreams/runtime/demo/spec.py +++ b/flashdreams/flashdreams/runtime/demo/spec.py @@ -5,7 +5,7 @@ from __future__ import annotations -from collections.abc import Callable, Mapping +from collections.abc import Callable, Mapping, Sequence from dataclasses import dataclass, field, replace from pathlib import Path from typing import Any, Literal, Protocol, TypeAlias @@ -18,6 +18,8 @@ from flashdreams.runtime.interfaces import ModelAdapter from flashdreams.runtime.mapping import InputMapping +from .host import WarmupSessionInputs + @dataclass(frozen=True, kw_only=True, slots=True) class NullOutputSpec: @@ -170,9 +172,22 @@ def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: ... +class ModelWarmupAdapter(Protocol): + """Optional adapter hook for model-affine runtime warmup inputs.""" + + def create_model_warmup_sessions( + self, + spec: DemoSpec, + scenario: PreparedScenario, + ) -> Sequence[WarmupSessionInputs]: + """Return temporary synthetic or loopback sessions for model warmup.""" + ... + + __all__ = [ "DemoAdapter", "DemoSpec", + "ModelWarmupAdapter", "Mp4OutputSpec", "NullOutputSpec", "OutputSpec", diff --git a/flashdreams/tests/test_demo_runtime_warmup.py b/flashdreams/tests/test_demo_runtime_warmup.py new file mode 100644 index 000000000..4c3951ed0 --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_warmup.py @@ -0,0 +1,416 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import threading +from typing import Any + +import pytest + +from flashdreams.runtime import ( + CanonicalInputSchema, + IdentityInputMapping, + InferenceConfig, + InferenceInput, + InferenceInputSchema, + InferenceRuntime, + InferenceSession, + InputMapping, + StepRequest, + StepRequirements, + StepResult, +) +from flashdreams.runtime.demo import ( + DemoSpec, + InMemorySessionMetricsRecorder, + ModelWarmupPlan, + NullOutputSpec, + PreparedScenario, + PreparedStep, + ProviderCapabilities, + RunContext, + RuntimeHost, + UserInputWindow, + WarmupSessionInputs, + build_model_warmup_plan, + warmup_run_context, +) + +pytestmark = pytest.mark.ci_cpu + + +def test_model_warmup_plan_uses_temporary_provider_on_worker_thread() -> None: + setup_thread_id = threading.get_ident() + runtime = _WarmupRuntime() + host = RuntimeHost(runtime) + adapter = _WarmupAdapter(warmup_steps=2) + spec = _spec() + scenario = adapter.prepare_scenario(spec) + + try: + plan = build_model_warmup_plan( + host=host, + adapter=adapter, + spec=spec, + scenario=scenario, + ) + real_provider = host.call(adapter.create_model_input_provider, spec, scenario) + finally: + host.close() + + warmup_provider = adapter.warmup_providers[0] + assert adapter.warmup_thread_id == host.worker.worker_thread_id + assert adapter.warmup_thread_id != setup_thread_id + assert warmup_provider is not real_provider + assert warmup_provider.close_count == 1 + assert real_provider.close_count == 0 + assert plan == ModelWarmupPlan( + sessions=( + WarmupSessionInputs( + initial_input=InferenceInput( + global_conditioning={"provider": "warmup"} + ), + step_inputs=( + InferenceInput(step={"provider": "warmup", "step": 0}), + InferenceInput(step={"provider": "warmup", "step": 1}), + ), + ), + ), + ) + + +def test_runtime_host_warmup_uses_runtime_session_api() -> None: + runtime = _WarmupRuntime() + host = RuntimeHost(runtime) + adapter = _WarmupAdapter(warmup_steps=2) + spec = _spec() + scenario = adapter.prepare_scenario(spec) + + try: + plan = build_model_warmup_plan( + host=host, + adapter=adapter, + spec=spec, + scenario=scenario, + ) + host.warmup(plan) + finally: + host.close() + + assert runtime.events[:4] == [ + ( + "start_session", + InferenceInput(global_conditioning={"provider": "warmup"}), + ), + ("step", InferenceInput(step={"provider": "warmup", "step": 0})), + ("step", InferenceInput(step={"provider": "warmup", "step": 1})), + "session.close", + ] + + +def test_run_mode_warmup_context_warms_transport_without_model_session() -> None: + runtime = _WarmupRuntime() + host = RuntimeHost(runtime) + adapter = _WarmupAdapter(warmup_steps=0) + spec = _spec() + scenario = adapter.prepare_scenario(spec) + transport = _TransportWarmupService() + mode = _TransportWarmupRunMode() + context = RunContext( + host=host, + run_metrics=InMemorySessionMetricsRecorder(), + admission=_Admission(), + model_warmup_plan=ModelWarmupPlan(), + services={"transport": transport}, + ) + + try: + warmup_run_context( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + run_mode=mode, + ) + + assert runtime.events == [] + assert mode.warmup_calls == 1 + assert transport.warmup_calls == 1 + finally: + host.close() + + +def test_model_warmup_is_excluded_from_run_metrics() -> None: + runtime = _WarmupRuntime() + host = RuntimeHost(runtime) + adapter = _WarmupAdapter(warmup_steps=1) + spec = _spec() + scenario = adapter.prepare_scenario(spec) + metrics = InMemorySessionMetricsRecorder() + + try: + plan = build_model_warmup_plan( + host=host, + adapter=adapter, + spec=spec, + scenario=scenario, + ) + context = RunContext( + host=host, + run_metrics=metrics, + admission=_Admission(), + model_warmup_plan=plan, + ) + + warmup_run_context( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + run_mode=object(), + ) + + assert runtime.events[:3] == [ + ( + "start_session", + InferenceInput(global_conditioning={"provider": "warmup"}), + ), + ("step", InferenceInput(step={"provider": "warmup", "step": 0})), + "session.close", + ] + assert metrics.sessions == [] + assert metrics.step_count == 0 + assert metrics.control_count == 0 + finally: + host.close() + + +def test_adapter_without_model_warmup_hook_gets_empty_plan() -> None: + host = RuntimeHost(_WarmupRuntime()) + adapter = _NoWarmupAdapter() + spec = _spec() + scenario = adapter.prepare_scenario(spec) + + try: + plan = build_model_warmup_plan( + host=host, + adapter=adapter, + spec=spec, + scenario=scenario, + ) + finally: + host.close() + + assert plan == ModelWarmupPlan() + + +def _spec() -> DemoSpec: + return DemoSpec( + model_id="fake-demo", + input_mode="replay", + output=NullOutputSpec(), + ) + + +def _scenario() -> PreparedScenario: + return PreparedScenario(initial_inputs=InferenceInput()) + + +class _WarmupAdapter: + model_id = "fake-demo" + inference_input_schema = InferenceInputSchema() + canonical_input_schema = CanonicalInputSchema() + + def __init__(self, *, warmup_steps: int) -> None: + self.warmup_steps = warmup_steps + self.warmup_thread_id: int | None = None + self.warmup_providers: list[_WarmupProvider] = [] + self.real_providers: list[_WarmupProvider] = [] + + def supported_input_modes(self) -> tuple[str, ...]: + return ("replay",) + + def supported_output_modes(self) -> tuple[str, ...]: + return ("null",) + + def default_input_mapping(self) -> InputMapping: + return IdentityInputMapping() + + def validate_config(self, config: InferenceConfig) -> None: + if config.model_id != self.model_id: + raise ValueError(f"Unsupported model_id={config.model_id!r}.") + + def create_runtime(self, config: InferenceConfig) -> InferenceRuntime: + self.validate_config(config) + return _WarmupRuntime() + + def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: + del spec + return _scenario() + + def create_model_warmup_sessions( + self, + spec: DemoSpec, + scenario: PreparedScenario, + ) -> tuple[WarmupSessionInputs, ...]: + del spec, scenario + self.warmup_thread_id = threading.get_ident() + provider = _WarmupProvider(name="warmup") + self.warmup_providers.append(provider) + try: + initial_input = provider.prepare_initial_input() + step_inputs = [] + for step_index in range(self.warmup_steps): + prepared = provider.prepare_step( + request=StepRequirements(step_index=step_index), + user_window=UserInputWindow( + start_s=float(step_index), + end_s=float(step_index + 1), + ), + ) + if prepared.inference_input is None: + raise RuntimeError("Warmup provider returned no step input.") + step_inputs.append(prepared.inference_input) + return ( + WarmupSessionInputs( + initial_input=initial_input, + step_inputs=tuple(step_inputs), + ), + ) + finally: + provider.close() + + def create_model_input_provider( + self, + spec: DemoSpec, + scenario: PreparedScenario, + ) -> "_WarmupProvider": + del spec, scenario + provider = _WarmupProvider(name="real") + self.real_providers.append(provider) + return provider + + +class _NoWarmupAdapter: + model_id = "fake-demo" + inference_input_schema = InferenceInputSchema() + canonical_input_schema = CanonicalInputSchema() + + def supported_input_modes(self) -> tuple[str, ...]: + return ("replay",) + + def supported_output_modes(self) -> tuple[str, ...]: + return ("null",) + + def default_input_mapping(self) -> InputMapping: + return IdentityInputMapping() + + def validate_config(self, config: InferenceConfig) -> None: + if config.model_id != self.model_id: + raise ValueError(f"Unsupported model_id={config.model_id!r}.") + + def create_runtime(self, config: InferenceConfig) -> InferenceRuntime: + self.validate_config(config) + return _WarmupRuntime() + + def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: + del spec + return _scenario() + + +class _WarmupProvider: + capabilities = ProviderCapabilities(supports_recorded_input=True) + + def __init__(self, *, name: str) -> None: + self.name = name + self.close_count = 0 + + def prepare_initial_input(self) -> InferenceInput: + return InferenceInput(global_conditioning={"provider": self.name}) + + def prepare_step( + self, + *, + request: StepRequirements, + user_window: UserInputWindow, + ) -> PreparedStep: + del user_window + return PreparedStep( + inference_input=InferenceInput( + step={"provider": self.name, "step": request.step_index} + ) + ) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + + def close(self) -> None: + self.close_count += 1 + + +class _WarmupRuntime: + def __init__(self) -> None: + self.events: list[object] = [] + + def start_session(self, inputs: InferenceInput) -> InferenceSession: + self.events.append(("start_session", inputs)) + return _WarmupSession(events=self.events) + + def close(self) -> None: + self.events.append("runtime.close") + + +class _WarmupSession: + def __init__(self, *, events: list[object]) -> None: + self.events = events + self.next_step = 0 + + def next_step_request(self) -> StepRequest | None: + request = StepRequest(step_index=self.next_step) + self.next_step += 1 + return request + + def step(self, inputs: InferenceInput) -> StepResult: + self.events.append(("step", inputs)) + return StepResult(step_index=self.next_step, output=None) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self.next_step = 0 + + def close(self) -> None: + self.events.append("session.close") + + +class _TransportWarmupService: + def __init__(self) -> None: + self.warmup_calls = 0 + + def warmup(self) -> None: + self.warmup_calls += 1 + + +class _TransportWarmupRunMode: + def __init__(self) -> None: + self.warmup_calls = 0 + + def warmup_context( + self, + *, + context: RunContext, + spec: DemoSpec, + scenario: PreparedScenario, + adapter: Any, + ) -> None: + del spec, scenario, adapter + transport = context.services["transport"] + if not isinstance(transport, _TransportWarmupService): + raise TypeError("Expected fake transport warmup service.") + transport.warmup() + self.warmup_calls += 1 + + +class _Admission: + def try_reserve(self) -> None: + return None From c8314f90af2843b4231fed7a2e0c1a8aad56da76 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 06:52:14 +0000 Subject: [PATCH 18/51] runtime: guard run cleanup telemetry failures Make run-level cleanup error recording best-effort so provider cleanup failures cannot replace the original session assembly outcome. Cover sync and async no-edge cleanup paths where provider close and cleanup telemetry both fail. --- .../flashdreams/runtime/demo/drivers.py | 13 +++- .../tests/test_demo_runtime_vertical_slice.py | 70 ++++++++++++++++++- 2 files changed, 79 insertions(+), 4 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index eabb5307c..7f3bcb064 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -390,7 +390,7 @@ def run_demo_session( if session_edges is not None: session_edges.record_cleanup_error(close_exc) else: - context.run_metrics.record_cleanup_error(close_exc) + _record_run_cleanup_error(context, close_exc) if session_edges is not None and ( driver_started or not session_edges.is_closed ): @@ -409,7 +409,7 @@ def run_demo_session( if session_edges is not None: session_edges.record_cleanup_error(close_exc) else: - context.run_metrics.record_cleanup_error(close_exc) + _record_run_cleanup_error(context, close_exc) if session_edges is not None and ( driver_started or not session_edges.is_closed ): @@ -795,7 +795,14 @@ async def _close_provider_async( if session_edges is not None: session_edges.record_cleanup_error(close_exc) else: - context.run_metrics.record_cleanup_error(close_exc) + _record_run_cleanup_error(context, close_exc) + + +def _record_run_cleanup_error(context: RunContext, exc: Exception) -> None: + try: + context.run_metrics.record_cleanup_error(exc) + except Exception: + return async def _close_model_resource_async( diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index f48966c16..fe3d17e6d 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -47,6 +47,7 @@ StepPipeline, UserInputWindow, run_demo_session, + run_demo_session_async, ) pytestmark = pytest.mark.ci_cpu @@ -298,6 +299,65 @@ def test_run_demo_session_closes_provider_when_validation_fails() -> None: assert run_metrics.sessions == [result] +def test_run_demo_session_keeps_failure_when_run_cleanup_metrics_fail() -> None: + runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) + run_metrics = _FailingCleanupMetrics() + context = _run_context(runtime, run_metrics=run_metrics) + provider = _FakeVideoModelInputProvider( + fail_close=RuntimeError("provider close failed") + ) + + result = run_demo_session( + context=context, + spec=_spec(), + scenario=_scenario(), + adapter=_FakeDemoAdapter(provider=provider), + run_mode=_FakeRunMode( + input_source=_FakeBatchInputSource(num_windows=1), + validate_error=ValueError("provider incompatible"), + ), + pipeline=StepPipeline(), + ) + + assert result.status == "failed" + assert result.reason == "provider incompatible" + assert provider.close_count == 1 + assert run_metrics.cleanup_error_attempts == 1 + assert run_metrics.sessions == [result] + assert runtime.start_session_inputs == [] + + +@pytest.mark.asyncio +async def test_run_demo_session_async_keeps_failure_when_run_cleanup_metrics_fail() -> ( + None +): + runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) + run_metrics = _FailingCleanupMetrics() + context = _run_context(runtime, run_metrics=run_metrics) + provider = _FakeVideoModelInputProvider( + fail_close=RuntimeError("provider close failed") + ) + + result = await run_demo_session_async( + context=context, + spec=_spec(), + scenario=_scenario(), + adapter=_FakeDemoAdapter(provider=provider), + run_mode=_FakeRunMode( + input_source=_FakeBatchInputSource(num_windows=1), + validate_error=ValueError("provider incompatible"), + ), + pipeline=StepPipeline(), + ) + + assert result.status == "failed" + assert result.reason == "provider incompatible" + assert provider.close_count == 1 + assert run_metrics.cleanup_error_attempts == 1 + assert run_metrics.sessions == [result] + assert runtime.start_session_inputs == [] + + def test_setup_failure_can_return_skipped_but_not_completed() -> None: skipped = BatchSessionDriver().run_one_session( host=RuntimeHost(_FakeVideoRuntime(session=_FakeVideoSession(num_steps=1))), @@ -673,8 +733,14 @@ class _FakeVideoModelInputProvider: deterministic_given_inputs=True, ) - def __init__(self, *, fail_initial: Exception | None = None) -> None: + def __init__( + self, + *, + fail_initial: Exception | None = None, + fail_close: Exception | None = None, + ) -> None: self.fail_initial = fail_initial + self.fail_close = fail_close self.initial_input = InferenceInput( global_conditioning={"prompt": "fake video prompt"} ) @@ -707,6 +773,8 @@ def reset(self, inputs: InferenceInput | None = None) -> None: def close(self) -> None: self.close_count += 1 + if self.fail_close is not None: + raise self.fail_close class _FakeBatchInputSource: From 7f36767c67253dd4e41a8c8d0c3c1c552290a7a3 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 07:33:29 +0000 Subject: [PATCH 19/51] Phase 10: unify demo metrics and error policies --- flashdreams/flashdreams/runtime/__init__.py | 2 + .../flashdreams/runtime/demo/__init__.py | 10 + .../flashdreams/runtime/demo/drivers.py | 12 + .../flashdreams/runtime/demo/run_modes.py | 163 +++++-------- flashdreams/flashdreams/runtime/metrics.py | 224 +++++++++++++++++- .../flashdreams/serving/realtime/timing.py | 50 ++++ .../test_demo_runtime_realtime_driver.py | 4 +- .../tests/test_demo_runtime_run_modes.py | 42 ++++ .../tests/test_demo_runtime_vertical_slice.py | 12 +- .../tests/test_inference_runtime_api.py | 77 ++++++ .../tests/test_realtime_timing_metrics.py | 71 ++++++ flashdreams/tests/test_runtime_runner.py | 42 +++- 12 files changed, 589 insertions(+), 120 deletions(-) create mode 100644 flashdreams/tests/test_realtime_timing_metrics.py diff --git a/flashdreams/flashdreams/runtime/__init__.py b/flashdreams/flashdreams/runtime/__init__.py index 1613863dc..6e205823c 100644 --- a/flashdreams/flashdreams/runtime/__init__.py +++ b/flashdreams/flashdreams/runtime/__init__.py @@ -52,6 +52,7 @@ from flashdreams.runtime.metrics import ( InMemoryMetricsRecorder, MetricsRecorder, + MetricsSnapshot, NullMetricsRecorder, RuntimeMetricSample, ) @@ -95,6 +96,7 @@ "KeyboardToDriverCommand", "MappingCompatibility", "MetricsRecorder", + "MetricsSnapshot", "ModelAdapter", "ModelExecutionWorker", "Mp4VideoOutputTarget", diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index eedeeeb84..4b3ad3cf9 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -31,11 +31,15 @@ from flashdreams.runtime.demo.replay import run_replay_demo from flashdreams.runtime.demo.run_modes import ( AsyncSessionDriver, + BenchmarkErrorPolicy, DefaultErrorPolicy, ErrorAction, InMemorySessionMetricsRecorder, MetricsSnapshot, + Mp4ErrorPolicy, + NativeWindowErrorPolicy, NoopTransportService, + NullErrorPolicy, RunContext, RunMode, RunModeCapabilities, @@ -45,6 +49,7 @@ SessionDriver, SessionEdges, SingleSessionAdmissionPolicy, + WebRTCErrorPolicy, build_model_warmup_plan, warmup_run_context, ) @@ -106,6 +111,7 @@ "ActivationResult", "ActivationSignal", "AlwaysActiveActivationPolicy", + "BenchmarkErrorPolicy", "InMemorySessionMetricsRecorder", "InputSource", "CatchUpDecision", @@ -116,11 +122,14 @@ "ModelWarmupPlan", "ModelInputProvider", "KeyboardRealtimeInputSource", + "Mp4ErrorPolicy", "Mp4OutputSink", "Mp4OutputSpec", + "NativeWindowErrorPolicy", "NoopTransportService", "NullOutputSpec", "NullOutputSink", + "NullErrorPolicy", "OutputDecision", "OutputSpec", "OutputSink", @@ -151,6 +160,7 @@ "UserInputWindow", "WarmupSessionInputs", "WebRTCAppResources", + "WebRTCErrorPolicy", "WebRTCOutputSpec", "build_output_sink", "build_output_target", diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 7f3bcb064..d87912feb 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -234,6 +234,7 @@ async def run_one_session( request=request, clock=clock, ) + session_edges.metrics.record_catch_up(window_result.catch_up) if ( not session_edges.transport.is_active() and not first_step_started @@ -383,6 +384,7 @@ def run_demo_session( context.run_metrics.record_session(result) return result except DriverInvariantError as exc: + _record_run_session_error(context, exc) if provider is not None and not driver_started: try: context.host.call(provider.close) @@ -402,6 +404,7 @@ def run_demo_session( context.run_metrics.record_session(result) raise except Exception as exc: + _record_run_session_error(context, exc) if provider is not None and not driver_started: try: context.host.call(provider.close) @@ -507,6 +510,7 @@ async def run_demo_session_async( context.run_metrics.record_session(result) return result except DriverInvariantError as exc: + _record_run_session_error(context, exc) if provider is not None and not driver_started: await _close_provider_async( context=context, @@ -524,6 +528,7 @@ async def run_demo_session_async( context.run_metrics.record_session(result) raise except Exception as exc: + _record_run_session_error(context, exc) result = await _close_partial_session_async( context=context, provider=provider, @@ -805,6 +810,13 @@ def _record_run_cleanup_error(context: RunContext, exc: Exception) -> None: return +def _record_run_session_error(context: RunContext, exc: Exception) -> None: + try: + context.run_metrics.record_session_error(exc) + except Exception: + return + + async def _close_model_resource_async( *, host: RuntimeHost, diff --git a/flashdreams/flashdreams/runtime/demo/run_modes.py b/flashdreams/flashdreams/runtime/demo/run_modes.py index aa2331605..ee71be193 100644 --- a/flashdreams/flashdreams/runtime/demo/run_modes.py +++ b/flashdreams/flashdreams/runtime/demo/run_modes.py @@ -12,10 +12,15 @@ from typing import TYPE_CHECKING, Any, Literal, Protocol, runtime_checkable from flashdreams.runtime._utils import freeze_mapping +from flashdreams.runtime.metrics import ( + InMemoryMetricsRecorder, + MetricsRecorder, + MetricsSnapshot, +) from flashdreams.runtime.output import OutputArtifact from .host import ModelWarmupPlan, WarmupSessionInputs -from .outputs import OutputDecision, OutputSink +from .outputs import OutputSink if TYPE_CHECKING: from .host import RuntimeHost @@ -42,28 +47,6 @@ ] -@dataclass(frozen=True, kw_only=True, slots=True) -class MetricsSnapshot: - """Closed session or run metrics summary.""" - - counters: Mapping[str, int | float] = field(default_factory=dict) - timings: Mapping[str, Sequence[float]] = field(default_factory=dict) - session_statuses: Sequence[str] = () - errors: Sequence[str] = () - - def __post_init__(self) -> None: - object.__setattr__(self, "counters", freeze_mapping(self.counters)) - object.__setattr__( - self, - "timings", - freeze_mapping( - {key: tuple(values) for key, values in self.timings.items()} - ), - ) - object.__setattr__(self, "session_statuses", tuple(self.session_statuses)) - object.__setattr__(self, "errors", tuple(self.errors)) - - @dataclass(frozen=True, kw_only=True, slots=True) class RunResult: """Outcome of one demo session.""" @@ -127,108 +110,65 @@ def handle_setup_error(self, exc: Exception) -> ErrorAction: ... def handle(self, exc: Exception) -> ErrorAction: ... -@dataclass(frozen=True, kw_only=True, slots=True) -class RunModeCapabilities: - """Run-mode requirements and output/transport capabilities.""" +class Mp4ErrorPolicy(DefaultErrorPolicy): + """Abort MP4 sessions on setup or step errors.""" - realtime: bool = False - requires_finite_input: bool = False - supports_backpressure: bool = False - supports_interactive_events: bool = False - supports_artifacts: bool = False +class NullErrorPolicy(DefaultErrorPolicy): + """Abort headless/null sessions on setup or step errors.""" -@runtime_checkable -class SessionMetricsRecorder(Protocol): - """Metrics callbacks consumed by the Phase 2 drivers and pipeline.""" - def record_step( - self, - *, - request: object, - user_window: object, - inference_input: object, - result: object, - decision: OutputDecision, - ) -> None: ... +class NativeWindowErrorPolicy(DefaultErrorPolicy): + """Abort native-window sessions unless a future UI policy overrides it.""" - def record_control( - self, - *, - request: object, - user_window: object, - control: object, - ) -> None: ... - def record_error(self, exc: Exception, action: ErrorAction) -> None: ... +class BenchmarkErrorPolicy(DefaultErrorPolicy): + """Close failed scenarios while letting benchmark loops continue.""" - def record_cleanup_error(self, exc: Exception) -> None: ... + def handle_setup_error(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status="failed", continue_next_scenario=True) - def record_session(self, result: RunResult) -> None: ... + def handle(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status="failed", continue_next_scenario=True) - def close(self) -> MetricsSnapshot: ... +@dataclass(frozen=True, kw_only=True, slots=True) +class WebRTCErrorPolicy: + """Drop configured recoverable realtime errors, otherwise close the session.""" -@dataclass(slots=True) -class InMemorySessionMetricsRecorder: - """Small non-raising metrics recorder for driver tests and fake demos.""" + recoverable_exception_types: tuple[type[Exception], ...] = () - step_count: int = 0 - control_count: int = 0 - errors: list[str] = field(default_factory=list) - cleanup_errors: list[str] = field(default_factory=list) - sessions: list[RunResult] = field(default_factory=list) - closed: bool = False + def handle_setup_error(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status="failed") - def record_step( - self, - *, - request: object, - user_window: object, - inference_input: object, - result: object, - decision: OutputDecision, - ) -> None: - del request, user_window, inference_input, result, decision - if not self.closed: - self.step_count += 1 - - def record_control( - self, - *, - request: object, - user_window: object, - control: object, - ) -> None: - del request, user_window, control - if not self.closed: - self.control_count += 1 - - def record_error(self, exc: Exception, action: ErrorAction) -> None: - del action - if not self.closed: - self.errors.append(str(exc)) + def handle(self, exc: Exception) -> ErrorAction: + if self.recoverable_exception_types and isinstance( + exc, self.recoverable_exception_types + ): + return ErrorAction( + close_session=False, + drop_chunk=True, + result_status="failed", + ) + return ErrorAction(result_status="failed") - def record_cleanup_error(self, exc: Exception) -> None: - if not self.closed: - self.cleanup_errors.append(str(exc)) - def record_session(self, result: RunResult) -> None: - if not self.closed: - self.sessions.append(result) +@dataclass(frozen=True, kw_only=True, slots=True) +class RunModeCapabilities: + """Run-mode requirements and output/transport capabilities.""" - def close(self) -> MetricsSnapshot: - self.closed = True - return MetricsSnapshot( - counters={ - "steps": self.step_count, - "controls": self.control_count, - "sessions": len(self.sessions), - "cleanup_errors": len(self.cleanup_errors), - }, - session_statuses=tuple(result.status for result in self.sessions), - errors=tuple((*self.errors, *self.cleanup_errors)), - ) + realtime: bool = False + requires_finite_input: bool = False + supports_backpressure: bool = False + supports_interactive_events: bool = False + supports_artifacts: bool = False + + +SessionMetricsRecorder = MetricsRecorder +InMemorySessionMetricsRecorder = InMemoryMetricsRecorder class NoopTransportService: @@ -558,13 +498,17 @@ def _coerce_warmup_sessions(value: object) -> tuple[WarmupSessionInputs, ...]: __all__ = [ "AdmissionPolicy", "AsyncSessionDriver", + "BenchmarkErrorPolicy", "DefaultErrorPolicy", "DriverStatus", "ErrorAction", "ErrorPolicy", "InMemorySessionMetricsRecorder", "MetricsSnapshot", + "Mp4ErrorPolicy", + "NativeWindowErrorPolicy", "NoopTransportService", + "NullErrorPolicy", "RunContext", "RunMode", "RunModeCapabilities", @@ -578,6 +522,7 @@ def _coerce_warmup_sessions(value: object) -> tuple[WarmupSessionInputs, ...]: "SessionStatus", "SingleSessionAdmissionPolicy", "TransportService", + "WebRTCErrorPolicy", "build_model_warmup_plan", "warmup_run_context", ] diff --git a/flashdreams/flashdreams/runtime/metrics.py b/flashdreams/flashdreams/runtime/metrics.py index 4286204f6..8a8fd6b1b 100644 --- a/flashdreams/flashdreams/runtime/metrics.py +++ b/flashdreams/flashdreams/runtime/metrics.py @@ -6,7 +6,8 @@ from __future__ import annotations import math -from collections.abc import Mapping +from collections import Counter, defaultdict +from collections.abc import Mapping, Sequence from dataclasses import dataclass, field from typing import Any, Protocol, runtime_checkable @@ -45,9 +46,31 @@ def __post_init__(self) -> None: object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) +@dataclass(frozen=True, kw_only=True, slots=True) +class MetricsSnapshot: + """Closed session or run metrics summary.""" + + counters: Mapping[str, int | float] = field(default_factory=dict) + timings: Mapping[str, Sequence[float]] = field(default_factory=dict) + session_statuses: Sequence[str] = () + errors: Sequence[str] = () + + def __post_init__(self) -> None: + object.__setattr__(self, "counters", freeze_mapping(self.counters)) + object.__setattr__( + self, + "timings", + freeze_mapping( + {key: tuple(values) for key, values in self.timings.items()} + ), + ) + object.__setattr__(self, "session_statuses", tuple(self.session_statuses)) + object.__setattr__(self, "errors", tuple(self.errors)) + + @runtime_checkable class MetricsRecorder(Protocol): - """Collector for runtime metrics.""" + """Collector for runtime, session, and run metrics.""" def record(self, sample: RuntimeMetricSample) -> None: """Record one metric sample.""" @@ -64,7 +87,53 @@ def record_timing( """Record one timing sample in seconds.""" ... - def close(self) -> None: + def record_step( + self, + *, + request: object, + user_window: object, + inference_input: object, + result: object, + decision: object, + ) -> None: + """Record one successful model step.""" + ... + + def record_control( + self, + *, + request: object, + user_window: object, + control: object, + ) -> None: + """Record one provider-authored control decision.""" + ... + + def record_error(self, exc: Exception, action: object) -> None: + """Record a driver-observed operational error.""" + ... + + def record_catch_up(self, decision: object) -> None: + """Record a realtime catch-up decision.""" + ... + + def record_cleanup_error(self, exc: Exception) -> None: + """Record a cleanup failure without interrupting teardown.""" + ... + + def record_orphaned_cleanup(self, exc: Exception) -> None: + """Record cleanup that timed out and is still queued.""" + ... + + def record_session(self, result: object) -> None: + """Record one closed session result.""" + ... + + def record_session_error(self, exc: Exception) -> None: + """Record diagnostic session assembly failure detail.""" + ... + + def close(self) -> MetricsSnapshot: """Finalize metric collection.""" ... @@ -74,6 +143,14 @@ class InMemoryMetricsRecorder: """Simple metrics recorder useful for tests, smoke runs, and adapters.""" samples: list[RuntimeMetricSample] = field(default_factory=list) + step_count: int = 0 + control_count: int = 0 + catch_up_count: int = 0 + errors: list[str] = field(default_factory=list) + cleanup_errors: list[str] = field(default_factory=list) + orphaned_cleanup_errors: list[str] = field(default_factory=list) + session_errors: list[str] = field(default_factory=list) + sessions: list[object] = field(default_factory=list) closed: bool = False def record(self, sample: RuntimeMetricSample) -> None: @@ -100,8 +177,96 @@ def record_timing( ) ) - def close(self) -> None: + def record_step( + self, + *, + request: object, + user_window: object, + inference_input: object, + result: object, + decision: object, + ) -> None: + del request, user_window, inference_input, result, decision + if not self.closed: + self.step_count += 1 + + def record_control( + self, + *, + request: object, + user_window: object, + control: object, + ) -> None: + del request, user_window, control + if not self.closed: + self.control_count += 1 + + def record_error(self, exc: Exception, action: object) -> None: + del action + if not self.closed: + self.errors.append(str(exc)) + + def record_catch_up(self, decision: object) -> None: + del decision + if not self.closed: + self.catch_up_count += 1 + + def record_cleanup_error(self, exc: Exception) -> None: + if not self.closed: + self.cleanup_errors.append(str(exc)) + + def record_orphaned_cleanup(self, exc: Exception) -> None: + if not self.closed: + self.orphaned_cleanup_errors.append(str(exc)) + + def record_session(self, result: object) -> None: + if not self.closed: + self.sessions.append(result) + + def record_session_error(self, exc: Exception) -> None: + if not self.closed: + self.session_errors.append(str(exc)) + + def close(self) -> MetricsSnapshot: self.closed = True + return self.snapshot() + + def snapshot(self) -> MetricsSnapshot: + timings: defaultdict[str, list[float]] = defaultdict(list) + for sample in self.samples: + if sample.category == "timing": + timings[sample.name].append(float(sample.value)) + session_statuses = tuple( + str(getattr(result, "status", "unknown")) for result in self.sessions + ) + session_status_counts = Counter(session_statuses) + return MetricsSnapshot( + counters={ + "samples": len(self.samples), + "steps": self.step_count, + "controls": self.control_count, + "catch_ups": self.catch_up_count, + "sessions": len(self.sessions), + "errors": len(self.errors), + "cleanup_errors": len(self.cleanup_errors), + "orphaned_cleanup_errors": len(self.orphaned_cleanup_errors), + "session_errors": len(self.session_errors), + **{ + f"sessions.{status}": count + for status, count in sorted(session_status_counts.items()) + }, + }, + timings=timings, + session_statuses=session_statuses, + errors=tuple( + ( + *self.errors, + *self.cleanup_errors, + *self.orphaned_cleanup_errors, + *self.session_errors, + ) + ), + ) class NullMetricsRecorder: @@ -120,5 +285,52 @@ def record_timing( ) -> None: del name, duration_s, step_index, metadata - def close(self) -> None: - return None + def record_step( + self, + *, + request: object, + user_window: object, + inference_input: object, + result: object, + decision: object, + ) -> None: + del request, user_window, inference_input, result, decision + + def record_control( + self, + *, + request: object, + user_window: object, + control: object, + ) -> None: + del request, user_window, control + + def record_error(self, exc: Exception, action: object) -> None: + del exc, action + + def record_catch_up(self, decision: object) -> None: + del decision + + def record_cleanup_error(self, exc: Exception) -> None: + del exc + + def record_orphaned_cleanup(self, exc: Exception) -> None: + del exc + + def record_session(self, result: object) -> None: + del result + + def record_session_error(self, exc: Exception) -> None: + del exc + + def close(self) -> MetricsSnapshot: + return MetricsSnapshot() + + +__all__ = [ + "InMemoryMetricsRecorder", + "MetricsRecorder", + "MetricsSnapshot", + "NullMetricsRecorder", + "RuntimeMetricSample", +] diff --git a/flashdreams/flashdreams/serving/realtime/timing.py b/flashdreams/flashdreams/serving/realtime/timing.py index 9891cfcdc..30147f97d 100644 --- a/flashdreams/flashdreams/serving/realtime/timing.py +++ b/flashdreams/flashdreams/serving/realtime/timing.py @@ -11,6 +11,8 @@ from threading import Lock from typing import Protocol +from flashdreams.runtime.metrics import MetricsRecorder + TraceComponentValue = str | int | float | bool | None @@ -402,6 +404,36 @@ def summarize_chunk_history(chunks: Iterable[ChunkTimes]) -> RecentTimingSummary ) +def record_chunk_timing_metrics( + metrics: MetricsRecorder, + chunk: ChunkTimes, +) -> None: + """Record available chunk timing durations to a session metrics recorder.""" + + _record_stage_timing_metrics( + metrics, + prefix="realtime.chunk", + durations_ms=chunk.stage_durations_ms(), + step_index=chunk.chunk_index, + ) + + +def record_video_model_timing_metrics( + metrics: MetricsRecorder, + timings: VideoModelTimings, + *, + chunk_index: int | None = None, +) -> None: + """Record backend-visible video model stage durations to session metrics.""" + + _record_stage_timing_metrics( + metrics, + prefix="realtime.model", + durations_ms=timings.stage_durations_ms(), + step_index=chunk_index, + ) + + class RollingChunkTimingSummary: def __init__(self, capacity: int) -> None: self._chunks: deque[ChunkTimes] = deque(maxlen=capacity) @@ -518,6 +550,24 @@ def _add_optional_trace_range( ) +def _record_stage_timing_metrics( + metrics: MetricsRecorder, + *, + prefix: str, + durations_ms: Mapping[str, float], + step_index: int | None, +) -> None: + for stage_name, duration_ms in durations_ms.items(): + try: + metrics.record_timing( + f"{prefix}.{stage_name}", + float(duration_ms) / 1000.0, + step_index=step_index, + ) + except Exception: + return + + def _summarize_values(values: list[float]) -> StageDurationSummary: ordered = sorted(values) count = len(ordered) diff --git a/flashdreams/tests/test_demo_runtime_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py index 26974ed98..f561c7646 100644 --- a/flashdreams/tests/test_demo_runtime_realtime_driver.py +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -316,13 +316,14 @@ async def test_realtime_driver_applies_backpressure_through_clock() -> None: runtime = _FakeRealtimeRuntime(session=session) host = RuntimeHost(runtime) clock = _RecordingRealtimeClock() + metrics = InMemorySessionMetricsRecorder() output = _RecordingOutputSink( decisions=( OutputDecision(backpressure_s=0.25), OutputDecision(should_stop=True), ) ) - edges = _edges(clock=clock, output=output) + edges = _edges(clock=clock, output=output, metrics=metrics) try: result = await RealtimeSessionDriver().run_one_session( @@ -336,6 +337,7 @@ async def test_realtime_driver_applies_backpressure_through_clock() -> None: assert result.status == "completed" assert clock.backpressure == [0.25] + assert metrics.catch_up_count == 2 assert len(output.results) == 2 diff --git a/flashdreams/tests/test_demo_runtime_run_modes.py b/flashdreams/tests/test_demo_runtime_run_modes.py index 8dd6e26af..de01b2771 100644 --- a/flashdreams/tests/test_demo_runtime_run_modes.py +++ b/flashdreams/tests/test_demo_runtime_run_modes.py @@ -25,11 +25,15 @@ ) from flashdreams.runtime.demo import ( AsyncSessionDriver, + BenchmarkErrorPolicy, DemoSpec, DriverInvariantError, InMemorySessionMetricsRecorder, ModelWarmupPlan, + Mp4ErrorPolicy, Mp4OutputSpec, + NativeWindowErrorPolicy, + NullErrorPolicy, NullOutputSpec, OutputDecision, PreparedScenario, @@ -43,6 +47,7 @@ SessionInfo, StepPipeline, UserInputWindow, + WebRTCErrorPolicy, WebRTCOutputSpec, run_demo_session, run_demo_session_async, @@ -222,6 +227,43 @@ async def test_run_context_close_async_drains_cleanup_tasks() -> None: assert metrics.closed +def test_error_policy_implementations_keep_setup_failures_terminal() -> None: + exc = RuntimeError("setup failed") + policies = ( + Mp4ErrorPolicy(), + BenchmarkErrorPolicy(), + WebRTCErrorPolicy(recoverable_exception_types=(RuntimeError,)), + NativeWindowErrorPolicy(), + NullErrorPolicy(), + ) + + for policy in policies: + action = policy.handle_setup_error(exc) + assert action.result_status == "failed" + assert action.close_session + assert not action.drop_chunk + + +def test_benchmark_error_policy_marks_failed_scenario_continuable() -> None: + action = BenchmarkErrorPolicy().handle(RuntimeError("scenario failed")) + + assert action.result_status == "failed" + assert action.close_session + assert action.continue_next_scenario + assert not action.drop_chunk + + +def test_webrtc_error_policy_can_drop_recoverable_step_errors() -> None: + action = WebRTCErrorPolicy( + recoverable_exception_types=(RuntimeError,), + ).handle(RuntimeError("output queue full")) + + assert action.result_status == "failed" + assert not action.close_session + assert action.drop_chunk + assert not action.continue_next_scenario + + def _run_fake_single_session_mode( *, spec: DemoSpec, diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index fe3d17e6d..bbe0a4abd 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -4,7 +4,7 @@ from __future__ import annotations from collections.abc import Callable, Sequence -from typing import Any, Literal +from typing import Any, Literal, cast import pytest @@ -297,6 +297,11 @@ def test_run_demo_session_closes_provider_when_validation_fails() -> None: assert provider.close_count == 1 assert runtime.start_session_inputs == [] assert run_metrics.sessions == [result] + assert run_metrics.session_errors == ["provider incompatible"] + snapshot = run_metrics.close() + assert snapshot.counters["sessions"] == 1 + assert snapshot.counters["session_errors"] == 1 + assert snapshot.session_statuses == ("failed",) def test_run_demo_session_keeps_failure_when_run_cleanup_metrics_fail() -> None: @@ -537,8 +542,9 @@ def test_run_demo_session_closes_edges_when_driver_invariant_escapes() -> None: assert session_metrics.closed assert provider.close_count == 1 assert len(run_metrics.sessions) == 1 - assert run_metrics.sessions[0].status == "failed" - assert isinstance(run_metrics.sessions[0].error, DriverInvariantError) + recorded = cast(RunResult, run_metrics.sessions[0]) + assert recorded.status == "failed" + assert isinstance(recorded.error, DriverInvariantError) def test_input_source_finished_error_returns_failed_not_completed() -> None: diff --git a/flashdreams/tests/test_inference_runtime_api.py b/flashdreams/tests/test_inference_runtime_api.py index d1235c905..6c939a7ed 100644 --- a/flashdreams/tests/test_inference_runtime_api.py +++ b/flashdreams/tests/test_inference_runtime_api.py @@ -4,6 +4,7 @@ from __future__ import annotations from dataclasses import fields +from types import SimpleNamespace from typing import Any, cast import pytest @@ -16,6 +17,8 @@ InferenceInputSchema, InMemoryMetricsRecorder, InputField, + MetricsSnapshot, + NullMetricsRecorder, NullOutputTarget, OutputArtifact, RuntimeMetricSample, @@ -290,6 +293,80 @@ def test_in_memory_metrics_recorder_uses_seconds_for_timing() -> None: assert sample.unit == "s" assert sample.category == "timing" assert sample.step_index == 2 + snapshot = recorder.close() + assert isinstance(snapshot, MetricsSnapshot) + assert recorder.closed + assert snapshot.counters["samples"] == 1 + assert snapshot.timings["model_step"] == (pytest.approx(0.125),) + + +def test_in_memory_metrics_recorder_rolls_up_sessions_and_diagnostics() -> None: + recorder = InMemoryMetricsRecorder() + + recorder.record_session(SimpleNamespace(status="completed")) + recorder.record_session_error(RuntimeError("assembly failed")) + recorder.record_error(RuntimeError("step failed"), object()) + recorder.record_catch_up(object()) + recorder.record_cleanup_error(RuntimeError("cleanup failed")) + recorder.record_orphaned_cleanup(RuntimeError("orphaned cleanup")) + snapshot = recorder.close() + + assert snapshot.counters["sessions"] == 1 + assert snapshot.counters["sessions.completed"] == 1 + assert snapshot.counters.get("sessions.failed", 0) == 0 + assert snapshot.counters["session_errors"] == 1 + assert snapshot.counters["catch_ups"] == 1 + assert snapshot.session_statuses == ("completed",) + assert snapshot.errors == ( + "step failed", + "cleanup failed", + "orphaned cleanup", + "assembly failed", + ) + + +def test_cancelled_session_rollup_does_not_count_as_failed() -> None: + recorder = InMemoryMetricsRecorder() + + recorder.record_session(SimpleNamespace(status="cancelled")) + snapshot = recorder.close() + + assert snapshot.counters["sessions"] == 1 + assert snapshot.counters["sessions.cancelled"] == 1 + assert snapshot.counters.get("sessions.failed", 0) == 0 + assert snapshot.session_statuses == ("cancelled",) + + +def test_null_metrics_recorder_keeps_old_and_new_calls_noop() -> None: + recorder = NullMetricsRecorder() + + recorder.record(RuntimeMetricSample(name="runtime", value=1.0)) + recorder.record_timing("model_step", 0.125, step_index=2) + recorder.record_step( + request=object(), + user_window=object(), + inference_input=object(), + result=object(), + decision=object(), + ) + recorder.record_control( + request=object(), + user_window=object(), + control=object(), + ) + recorder.record_error(RuntimeError("step failed"), object()) + recorder.record_catch_up(object()) + recorder.record_cleanup_error(RuntimeError("cleanup failed")) + recorder.record_orphaned_cleanup(RuntimeError("orphaned cleanup")) + recorder.record_session(SimpleNamespace(status="failed")) + recorder.record_session_error(RuntimeError("assembly failed")) + snapshot = recorder.close() + + assert isinstance(snapshot, MetricsSnapshot) + assert snapshot.counters == {} + assert snapshot.timings == {} + assert snapshot.session_statuses == () + assert snapshot.errors == () def test_timing_metric_samples_must_use_seconds() -> None: diff --git a/flashdreams/tests/test_realtime_timing_metrics.py b/flashdreams/tests/test_realtime_timing_metrics.py new file mode 100644 index 000000000..db157f195 --- /dev/null +++ b/flashdreams/tests/test_realtime_timing_metrics.py @@ -0,0 +1,71 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import pytest + +from flashdreams.runtime import InMemoryMetricsRecorder +from flashdreams.serving.realtime.timing import ( + ChunkTimes, + VideoModelTimings, + record_chunk_timing_metrics, + record_video_model_timing_metrics, +) + +pytestmark = pytest.mark.ci_cpu + + +def test_chunk_timing_records_feed_session_metrics() -> None: + chunk = ChunkTimes.create( + chunk_index=2, + input_sample_time=0.0, + request_time=0.010, + request_poses_ready_time=0.030, + intended_present_times=[0.100], + ) + chunk.chunk_render_start_time = 0.050 + chunk.chunk_ready_time = 0.090 + chunk.frames[0].image_ready_time = 0.110 + chunk.frames[0].present_time = 0.140 + metrics = InMemoryMetricsRecorder() + + record_chunk_timing_metrics(metrics, chunk) + snapshot = metrics.close() + + assert snapshot.timings["realtime.chunk.input_to_request"] == ( + pytest.approx(0.010), + ) + assert snapshot.timings["realtime.chunk.request_to_poses_ready"] == ( + pytest.approx(0.020), + ) + assert snapshot.timings["realtime.chunk.queue_wait"] == (pytest.approx(0.020),) + assert snapshot.timings["realtime.chunk.chunk_render"] == (pytest.approx(0.040),) + assert metrics.samples[0].step_index == 2 + + +def test_video_model_timing_records_feed_session_metrics() -> None: + timings = VideoModelTimings( + condition_start_time=1.0, + condition_ready_time=1.010, + model_start_time=1.020, + model_ready_time=1.070, + cache_update_start_time=1.075, + cache_update_ready_time=1.080, + decode_start_time=1.085, + decode_ready_time=1.095, + merge_start_time=1.100, + merge_ready_time=1.115, + ) + metrics = InMemoryMetricsRecorder() + + record_video_model_timing_metrics(metrics, timings, chunk_index=3) + snapshot = metrics.close() + + assert snapshot.timings["realtime.model.condition"] == (pytest.approx(0.010),) + assert snapshot.timings["realtime.model.model"] == (pytest.approx(0.050),) + assert snapshot.timings["realtime.model.cache_update"] == (pytest.approx(0.005),) + assert snapshot.timings["realtime.model.decode"] == (pytest.approx(0.010),) + assert snapshot.timings["realtime.model.merge"] == (pytest.approx(0.015),) + assert snapshot.timings["realtime.model.total"] == (pytest.approx(0.115),) + assert metrics.samples[0].step_index == 3 diff --git a/flashdreams/tests/test_runtime_runner.py b/flashdreams/tests/test_runtime_runner.py index b755ff48a..0750c8cb9 100644 --- a/flashdreams/tests/test_runtime_runner.py +++ b/flashdreams/tests/test_runtime_runner.py @@ -25,6 +25,7 @@ InputField, InputMapping, InputMappingSchema, + MetricsSnapshot, NullOutputTarget, OutputArtifact, RuntimeMetricSample, @@ -656,5 +657,44 @@ def record_timing( ) ) - def close(self) -> None: + def record_step( + self, + *, + request: object, + user_window: object, + inference_input: object, + result: object, + decision: object, + ) -> None: + del request, user_window, inference_input, result, decision + + def record_control( + self, + *, + request: object, + user_window: object, + control: object, + ) -> None: + del request, user_window, control + + def record_error(self, exc: Exception, action: object) -> None: + del exc, action + + def record_catch_up(self, decision: object) -> None: + del decision + + def record_cleanup_error(self, exc: Exception) -> None: + del exc + + def record_orphaned_cleanup(self, exc: Exception) -> None: + del exc + + def record_session(self, result: object) -> None: + del result + + def record_session_error(self, exc: Exception) -> None: + del exc + + def close(self) -> MetricsSnapshot: self._events.append("metrics.close") + return MetricsSnapshot() From 5abea73b995470df9a93ec2d7fe687f9f6d8d194 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 08:04:48 +0000 Subject: [PATCH 20/51] Phase 11: extract legacy runner onto shared batch path Refactor run_inference_session into a compatibility wrapper over the shared BatchSessionDriver and StepPipeline while preserving legacy mapping, output target, timing metrics, and cleanup behavior. Add an explicit legacy StepRequest adaptation opt-in and focused CPU coverage for wrapper delegation and driver-owned user-window handling. --- flashdreams/flashdreams/runtime/runner.py | 645 +++++++++++++++--- flashdreams/flashdreams/runtime/types.py | 8 +- .../tests/test_inference_runtime_api.py | 20 + flashdreams/tests/test_runtime_runner.py | 52 ++ 4 files changed, 634 insertions(+), 91 deletions(-) diff --git a/flashdreams/flashdreams/runtime/runner.py b/flashdreams/flashdreams/runtime/runner.py index c73a4f67c..f9a83411f 100644 --- a/flashdreams/flashdreams/runtime/runner.py +++ b/flashdreams/flashdreams/runtime/runner.py @@ -6,6 +6,8 @@ from __future__ import annotations import math +from collections.abc import Mapping, Sequence +from typing import TYPE_CHECKING, Any from flashdreams.runtime.canonical import InputCanonicalizer from flashdreams.runtime.config import InferenceConfig @@ -13,6 +15,7 @@ CanonicalInputs, CanonicalInputSchema, InferenceInput, + InferenceInputSchema, TimeWindow, UserInputs, UserInputSchema, @@ -29,7 +32,18 @@ ) from flashdreams.runtime.metrics import MetricsRecorder from flashdreams.runtime.output import OutputArtifact, OutputTarget -from flashdreams.runtime.types import StepResult +from flashdreams.runtime.types import ( + StepRequest, + StepRequirements, + StepResult, + step_requirements_from_request, +) + +if TYPE_CHECKING: + from flashdreams.runtime.demo.host import RuntimeHost + from flashdreams.runtime.demo.outputs import OutputDecision, SessionInfo + from flashdreams.runtime.demo.run_modes import RunResult + from flashdreams.runtime.demo.session_inputs import PreparedStep, UserInputWindow _DEFAULT_SESSION_HORIZON_S = 3600.0 @@ -46,73 +60,567 @@ def run_inference_session( output: OutputTarget, metrics: MetricsRecorder, ) -> tuple[OutputArtifact, ...]: - """Run one sequential inference session through the standard loop. + """Run one sequential inference session through the shared batch driver. - This v0 loop intentionally handles one adapter/runtime/session, one selected - input mapping, one replay/live input batch, one output target, and one - metrics recorder. It is synchronous and owns only orchestration. + The signature and failure semantics remain compatible with the original + replay runner while the implementation delegates the step loop to the shared + demo runtime pipeline. """ - runtime: InferenceRuntime | None = None - session: InferenceSession | None = None - output_opened = False - output_artifacts: tuple[OutputArtifact, ...] = () + return _run_inference_session_with_shared_batch( + adapter=adapter, + config=config, + mapping=mapping, + canonicalizer=canonicalizer, + source_schema=source_schema, + user_inputs=user_inputs, + initial_inputs=initial_inputs, + output=output, + metrics=metrics, + ) + + +def _run_inference_session_with_shared_batch( + *, + adapter: ModelAdapter, + config: InferenceConfig, + mapping: InputMapping, + canonicalizer: InputCanonicalizer, + source_schema: UserInputSchema, + user_inputs: UserInputs, + initial_inputs: InferenceInput, + output: OutputTarget, + metrics: MetricsRecorder, +) -> tuple[OutputArtifact, ...]: + lifecycle = _LegacyBatchLifecycle() + request_state = _LegacyStepRequestState() + runtime = _LegacyLazyRuntime( + adapter=adapter, + config=config, + lifecycle=lifecycle, + request_state=request_state, + ) + from flashdreams.runtime.demo.host import RuntimeHost + + host = RuntimeHost(runtime) + metrics_recorder = _LegacyRunnerMetricsRecorder(metrics) + output_sink = _LegacyOutputTargetSink( + output=output, + lifecycle=lifecycle, + host=host, + ) primary_error: BaseException | None = None try: - adapter.validate_config(config) - canonical_schema = canonicalizer.canonical_schema(source_schema) - _check_declared_mapping_compatibility( - mapping=mapping, - canonical_schema=canonical_schema, + _validate_legacy_runner_inputs( adapter=adapter, + config=config, + mapping=mapping, + canonicalizer=canonicalizer, + source_schema=source_schema, ) - mapping.validate( - canonical_schema=canonical_schema, + provider = _LegacyMappedModelInputProvider( + mapping=mapping, + canonicalizer=canonicalizer, + source_schema=source_schema, + user_inputs=user_inputs, + initial_inputs=initial_inputs, inference_input_schema=adapter.inference_input_schema, + request_state=request_state, ) - canonicalizer.reset() - mapped_initial_inputs = mapping.map_global_conditioning_inputs( - canonical_inputs=CanonicalInputs(), - inference_input=initial_inputs, + input_source = _LegacyBatchInputSource( + source_schema=source_schema, + user_inputs=user_inputs, + request_state=request_state, ) - runtime = adapter.create_runtime(config) - session = runtime.start_session(mapped_initial_inputs) - output.open() - output_opened = True - step_base_inputs = InferenceInput( - step=initial_inputs.step, - metadata=initial_inputs.metadata, + result = _run_shared_batch_session( + host=host, + provider=provider, + input_source=input_source, + output_sink=output_sink, + metrics=metrics_recorder, ) - - while (request := session.next_step_request()) is not None: - step_inputs = mapping.map_step_inputs( - canonical_inputs=canonicalizer.canonicalize( - user_inputs, - window=request.user_input_window - or _all_user_inputs_window(user_inputs), - source_schema=source_schema, - ), - inference_input=step_base_inputs, - request=request, - ) - result = session.step(step_inputs) - output.write(result) - _record_timing_metrics(metrics, result) + _raise_legacy_runner_error( + result=result, + output_sink=output_sink, + metrics=metrics_recorder, + ) + return tuple(result.artifacts) except BaseException as exc: primary_error = exc + if not metrics_recorder.closed: + _close_metrics_suppressing_secondary(metrics_recorder) raise finally: - cleanup_error, output_artifacts = _close_run_resources( - output=output if output_opened else None, - session=session, - runtime=runtime, + try: + host.close() + except BaseException: + if primary_error is None: + raise + + +def _validate_legacy_runner_inputs( + *, + adapter: ModelAdapter, + config: InferenceConfig, + mapping: InputMapping, + canonicalizer: InputCanonicalizer, + source_schema: UserInputSchema, +) -> CanonicalInputSchema: + adapter.validate_config(config) + canonical_schema = canonicalizer.canonical_schema(source_schema) + _check_declared_mapping_compatibility( + mapping=mapping, + canonical_schema=canonical_schema, + adapter=adapter, + ) + mapping.validate( + canonical_schema=canonical_schema, + inference_input_schema=adapter.inference_input_schema, + ) + return canonical_schema + + +def _run_shared_batch_session( + *, + host: RuntimeHost, + provider: "_LegacyMappedModelInputProvider", + input_source: "_LegacyBatchInputSource", + output_sink: "_LegacyOutputTargetSink", + metrics: "_LegacyRunnerMetricsRecorder", +) -> RunResult: + from flashdreams.runtime.demo.drivers import BatchSessionDriver + from flashdreams.runtime.demo.pipeline import StepPipeline + from flashdreams.runtime.demo.run_modes import SessionEdges + + return BatchSessionDriver().run_one_session( + host=host, + provider=provider, + session_edges=SessionEdges( + input_source=input_source, + output_sink=output_sink, + cleanup_tasks=set(), metrics=metrics, + ), + pipeline=StepPipeline(), + ) + + +class _LegacyLazyRuntime: + """Create the legacy runtime only after global inputs are mapped.""" + + def __init__( + self, + *, + adapter: ModelAdapter, + config: InferenceConfig, + lifecycle: "_LegacyBatchLifecycle", + request_state: "_LegacyStepRequestState", + ) -> None: + self._adapter = adapter + self._config = config + self._lifecycle = lifecycle + self._request_state = request_state + + def start_session(self, inputs: InferenceInput) -> InferenceSession: + runtime = self._lifecycle.runtime + if runtime is None: + runtime = self._adapter.create_runtime(self._config) + self._lifecycle.set_runtime(runtime) + session = runtime.start_session(inputs) + self._lifecycle.set_session(session) + return _LegacySessionAdapter( + session=session, + request_state=self._request_state, + ) + + def close(self) -> None: + self._lifecycle.close_runtime_direct() + + +class _LegacyBatchLifecycle: + """Own legacy model resources whose close order is caller-visible.""" + + def __init__(self) -> None: + self.runtime: InferenceRuntime | None = None + self.session: InferenceSession | None = None + self._session_close_attempted = False + self._runtime_close_attempted = False + + def set_runtime(self, runtime: InferenceRuntime) -> None: + self.runtime = runtime + + def set_session(self, session: InferenceSession) -> None: + self.session = session + self._session_close_attempted = False + + def close_session_via_host(self, host: RuntimeHost) -> None: + host.call(self.close_session_direct) + + def close_runtime_via_host(self, host: RuntimeHost) -> None: + host.call(self.close_runtime_direct) + + def close_session_direct(self) -> None: + if self.session is None or self._session_close_attempted: + return + self._session_close_attempted = True + self.session.close() + + def close_runtime_direct(self) -> None: + if self.runtime is None or self._runtime_close_attempted: + return + self._runtime_close_attempted = True + self.runtime.close() + + +class _LegacySessionAdapter: + """Expose old sessions through the new StepRequirements boundary.""" + + def __init__( + self, + *, + session: InferenceSession, + request_state: "_LegacyStepRequestState", + ) -> None: + self._session = session + self._request_state = request_state + + def next_step_requirements(self) -> StepRequirements | None: + request = self._session.next_step_request() + if request is None: + self._request_state.clear() + return None + self._request_state.store(request) + return step_requirements_from_request( + request, + allow_user_input_window=True, + ) + + def next_step_request(self) -> StepRequest | None: + return self._session.next_step_request() + + def session_info(self) -> SessionInfo: + from flashdreams.runtime.demo.outputs import SessionInfo + + session_info = getattr(self._session, "session_info", None) + if not callable(session_info): + return SessionInfo() + value = session_info() + if not isinstance(value, SessionInfo): + raise TypeError( + "session.session_info() must return SessionInfo, " + f"got {type(value).__name__}." + ) + return value + + def step(self, inputs: InferenceInput) -> StepResult: + return self._session.step(inputs) + + def reset(self, inputs: InferenceInput | None = None) -> None: + self._session.reset(inputs) + + def close(self) -> None: + return None + + +class _LegacyStepRequestState: + """Share the current legacy request between the session, source, and provider.""" + + def __init__(self) -> None: + self._request: StepRequest | None = None + + def store(self, request: StepRequest) -> None: + self._request = request + + def require_for_window(self, step_index: int) -> StepRequest: + request = self._request + if request is None: + raise RuntimeError("Legacy input source has no active step request.") + if request.step_index != step_index: + raise RuntimeError( + "Legacy input source request mismatch: " + f"expected step {request.step_index}, got {step_index}." + ) + return request + + def consume_for_step(self, step_index: int) -> StepRequest: + request = self.require_for_window(step_index) + self._request = None + return request + + def clear(self) -> None: + self._request = None + + +class _LegacyBatchInputSource: + is_finite = True + is_deterministic = True + + def __init__( + self, + *, + source_schema: UserInputSchema, + user_inputs: UserInputs, + request_state: _LegacyStepRequestState, + ) -> None: + self.user_input_schema = source_schema + self._user_inputs = user_inputs + self._request_state = request_state + + def is_finished(self) -> bool: + return False + + def next_window(self, request: StepRequirements) -> UserInputWindow: + from flashdreams.runtime.demo.session_inputs import UserInputWindow + + legacy_request = self._request_state.require_for_window(request.step_index) + window = legacy_request.user_input_window or _all_user_inputs_window( + self._user_inputs + ) + return UserInputWindow( + start_s=window.start_s, + end_s=window.end_s, + inputs=self._user_inputs, + ) + + +class _LegacyMappedModelInputProvider: + def __init__( + self, + *, + mapping: InputMapping, + canonicalizer: InputCanonicalizer, + source_schema: UserInputSchema, + user_inputs: UserInputs, + initial_inputs: InferenceInput, + inference_input_schema: InferenceInputSchema, + request_state: _LegacyStepRequestState, + ) -> None: + from flashdreams.runtime.demo.session_inputs import ProviderCapabilities + + self.capabilities = ProviderCapabilities( + supports_recorded_input=True, + deterministic_given_inputs=True, + user_input_schema=source_schema, + inference_input_schema=inference_input_schema, + ) + self._mapping = mapping + self._canonicalizer = canonicalizer + self._source_schema = source_schema + self._user_inputs = user_inputs + self._initial_inputs = initial_inputs + self._request_state = request_state + self._step_base_inputs = InferenceInput( + step=initial_inputs.step, + metadata=initial_inputs.metadata, + ) + + def prepare_initial_input(self) -> InferenceInput: + self._canonicalizer.reset() + return self._mapping.map_global_conditioning_inputs( + canonical_inputs=CanonicalInputs(), + inference_input=self._initial_inputs, ) - if cleanup_error is not None and primary_error is None: + + def prepare_step( + self, + *, + request: StepRequirements, + user_window: UserInputWindow, + ) -> PreparedStep: + from flashdreams.runtime.demo.session_inputs import PreparedStep + + legacy_request = self._request_state.consume_for_step(request.step_index) + canonical_inputs = self._canonicalizer.canonicalize( + self._user_inputs, + window=TimeWindow(start_s=user_window.start_s, end_s=user_window.end_s), + source_schema=self._source_schema, + ) + return PreparedStep( + inference_input=self._mapping.map_step_inputs( + canonical_inputs=canonical_inputs, + inference_input=self._step_base_inputs, + request=legacy_request, + ) + ) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self._canonicalizer.reset() + + def close(self) -> None: + return None + + +class _LegacyOutputTargetSink: + produces_artifacts = True + + def __init__( + self, + *, + output: OutputTarget, + lifecycle: _LegacyBatchLifecycle, + host: RuntimeHost, + ) -> None: + self._output = output + self._lifecycle = lifecycle + self._host = host + self._opened = False + self._closed = False + self._artifacts: tuple[OutputArtifact, ...] = () + self.cleanup_error: BaseException | None = None + + def open(self, session_info: SessionInfo) -> None: + del session_info + self._output.open() + self._opened = True + self._closed = False + self._artifacts = () + self.cleanup_error = None + + def begin_generation(self, generation: int) -> None: + del generation + + def write(self, result: StepResult) -> OutputDecision: + from flashdreams.runtime.demo.outputs import OutputDecision + + self._output.write(result) + return OutputDecision() + + def close(self) -> Sequence[OutputArtifact]: + if self._closed: + if self.cleanup_error is not None: + raise self.cleanup_error + return self._artifacts + + self._closed = True + cleanup_error: BaseException | None = None + artifacts: tuple[OutputArtifact, ...] = () + + def remember_error(exc: BaseException) -> None: + nonlocal cleanup_error + if cleanup_error is None: + cleanup_error = exc + + if self._opened: + try: + artifacts = tuple(self._output.close()) + except BaseException as exc: + remember_error(exc) + try: + self._lifecycle.close_session_via_host(self._host) + except BaseException as exc: + remember_error(exc) + try: + self._lifecycle.close_runtime_via_host(self._host) + except BaseException as exc: + remember_error(exc) + + self._artifacts = artifacts + self.cleanup_error = cleanup_error + if cleanup_error is not None: raise cleanup_error + return self._artifacts + + +class _LegacyRunnerMetricsRecorder: + def __init__(self, metrics: MetricsRecorder) -> None: + self._metrics = metrics + self.closed = False + self.close_error: BaseException | None = None - return output_artifacts + def record(self, sample: Any) -> None: + self._metrics.record(sample) + + def record_timing( + self, + name: str, + duration_s: float, + *, + step_index: int | None = None, + metadata: Mapping[str, Any] | None = None, + ) -> None: + self._metrics.record_timing( + name, + duration_s, + step_index=step_index, + metadata=metadata, + ) + + def record_step( + self, + *, + request: object, + user_window: object, + inference_input: object, + result: object, + decision: object, + ) -> None: + del request, user_window, inference_input, decision + if isinstance(result, StepResult): + _record_timing_metrics(self._metrics, result) + + def record_control( + self, + *, + request: object, + user_window: object, + control: object, + ) -> None: + del request, user_window, control + + def record_error(self, exc: Exception, action: object) -> None: + del exc, action + + def record_catch_up(self, decision: object) -> None: + del decision + + def record_cleanup_error(self, exc: Exception) -> None: + del exc + + def record_orphaned_cleanup(self, exc: Exception) -> None: + del exc + + def record_session(self, result: object) -> None: + del result + + def record_session_error(self, exc: Exception) -> None: + del exc + + def close(self) -> Any: + if self.closed: + if self.close_error is not None: + raise self.close_error + return None + self.closed = True + try: + return self._metrics.close() + except BaseException as exc: + self.close_error = exc + raise + + +def _raise_legacy_runner_error( + *, + result: RunResult, + output_sink: _LegacyOutputTargetSink, + metrics: _LegacyRunnerMetricsRecorder, +) -> None: + if result.error is not None: + raise result.error + if output_sink.cleanup_error is not None: + raise output_sink.cleanup_error + if metrics.close_error is not None: + raise metrics.close_error + + +def _close_metrics_suppressing_secondary( + metrics: _LegacyRunnerMetricsRecorder, +) -> None: + try: + metrics.close() + except BaseException: + return def _check_declared_mapping_compatibility( @@ -158,45 +666,4 @@ def _record_timing_metrics( ) -def _close_run_resources( - *, - output: OutputTarget | None, - session: InferenceSession | None, - runtime: InferenceRuntime | None, - metrics: MetricsRecorder, -) -> tuple[BaseException | None, tuple[OutputArtifact, ...]]: - cleanup_error: BaseException | None = None - artifacts: tuple[OutputArtifact, ...] = () - - def remember_error(exc: BaseException) -> None: - nonlocal cleanup_error - if cleanup_error is None: - cleanup_error = exc - - if output is not None: - try: - artifacts = tuple(output.close()) - except BaseException as exc: - remember_error(exc) - - if session is not None: - try: - session.close() - except BaseException as exc: - remember_error(exc) - - if runtime is not None: - try: - runtime.close() - except BaseException as exc: - remember_error(exc) - - try: - metrics.close() - except BaseException as exc: - remember_error(exc) - - return cleanup_error, artifacts - - __all__ = ["run_inference_session"] diff --git a/flashdreams/flashdreams/runtime/types.py b/flashdreams/flashdreams/runtime/types.py index 1773f0cae..49ee958a0 100644 --- a/flashdreams/flashdreams/runtime/types.py +++ b/flashdreams/flashdreams/runtime/types.py @@ -95,10 +95,14 @@ def __post_init__(self) -> None: object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) -def step_requirements_from_request(request: StepRequest) -> StepRequirements: +def step_requirements_from_request( + request: StepRequest, + *, + allow_user_input_window: bool = False, +) -> StepRequirements: """Adapt a legacy ``StepRequest`` that did not carry driver-owned inputs.""" - if request.user_input_window is not None: + if request.user_input_window is not None and not allow_user_input_window: raise ValueError( "StepRequest.user_input_window cannot be adapted to StepRequirements; " "user input windows are driver-owned." diff --git a/flashdreams/tests/test_inference_runtime_api.py b/flashdreams/tests/test_inference_runtime_api.py index 6c939a7ed..531f2f218 100644 --- a/flashdreams/tests/test_inference_runtime_api.py +++ b/flashdreams/tests/test_inference_runtime_api.py @@ -245,6 +245,26 @@ def test_step_requirements_keep_user_inputs_driver_owned() -> None: ) +def test_step_requirements_can_drop_legacy_user_window_when_source_owns_it() -> None: + request = StepRequest( + step_index=2, + user_input_window=TimeWindow(start_s=1.0, end_s=2.0), + metadata={"input_frame_count": 3, "model": "fake"}, + ) + + requirements = step_requirements_from_request( + request, + allow_user_input_window=True, + ) + + assert requirements == StepRequirements( + step_index=2, + input_frame_count=3, + metadata={"model": "fake"}, + ) + assert not hasattr(requirements, "user_input_window") + + def test_null_output_target_counts_and_optionally_stores_results() -> None: target = NullOutputTarget(store_results=True) result = StepResult(step_index=0, output=b"frame") diff --git a/flashdreams/tests/test_runtime_runner.py b/flashdreams/tests/test_runtime_runner.py index 0750c8cb9..7ba92e836 100644 --- a/flashdreams/tests/test_runtime_runner.py +++ b/flashdreams/tests/test_runtime_runner.py @@ -8,6 +8,7 @@ import pytest +import flashdreams.runtime.runner as runner_module from flashdreams.runtime import ( DRIVER_COMMAND, CanonicalInputs, @@ -74,6 +75,57 @@ def test_run_inference_session_completes_two_step_run() -> None: assert metrics.closed +def test_run_inference_session_delegates_to_shared_batch_helper( + monkeypatch: pytest.MonkeyPatch, +) -> None: + calls: list[Mapping[str, object]] = [] + artifact = OutputArtifact(kind="test/artifact", uri="memory://artifact") + + def _fake_helper(**kwargs: object) -> tuple[OutputArtifact, ...]: + calls.append(kwargs) + return (artifact,) + + monkeypatch.setattr( + runner_module, + "_run_inference_session_with_shared_batch", + _fake_helper, + ) + adapter = _FakeAdapter() + config = InferenceConfig(model_id="fake-model") + mapping = _ChunkIndexMapping() + canonicalizer = InputCanonicalizer() + source_schema = UserInputSchema() + user_inputs = UserInputs() + initial_inputs = InferenceInput(global_conditioning={"prompt": "drive forward"}) + output = NullOutputTarget() + metrics = InMemoryMetricsRecorder() + + artifacts = runner_module.run_inference_session( + adapter=adapter, + config=config, + mapping=mapping, + canonicalizer=canonicalizer, + source_schema=source_schema, + user_inputs=user_inputs, + initial_inputs=initial_inputs, + output=output, + metrics=metrics, + ) + + assert artifacts == (artifact,) + assert len(calls) == 1 + call = calls[0] + assert call["adapter"] is adapter + assert call["config"] is config + assert call["mapping"] is mapping + assert call["canonicalizer"] is canonicalizer + assert call["source_schema"] is source_schema + assert call["user_inputs"] is user_inputs + assert call["initial_inputs"] is initial_inputs + assert call["output"] is output + assert call["metrics"] is metrics + + def test_runner_preserves_initial_step_inputs_for_identity_mapping() -> None: adapter = _FakeAdapter() From 386e9b0d6d05f0f69c8192d881364cf7f44ec438 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 08:37:38 +0000 Subject: [PATCH 21/51] Phase 11.5: adopt shared replay output sinks Move the default replay demo path onto a private replay RunMode backed by RunContext, SessionEdges, BatchSessionDriver, StepPipeline, and the shared MP4/null OutputSink implementations. Preserve old runner/output-target injection as a compatibility path, return RunResult from replay, and make DemoApplication map failed replay results to a non-zero exit. Add CPU coverage for sink adoption, legacy payload parity, runner compatibility, and replay failure exits. --- .../flashdreams/runtime/demo/__init__.py | 3 +- flashdreams/flashdreams/runtime/demo/app.py | 11 +- .../flashdreams/runtime/demo/replay.py | 508 +++++++++++++++++- flashdreams/tests/test_runtime_demo_api.py | 184 ++++++- integrations/lingbot/tests/test_demo_api.py | 7 +- .../omnidreams/tests/test_demo_api.py | 7 +- 6 files changed, 694 insertions(+), 26 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/__init__.py b/flashdreams/flashdreams/runtime/demo/__init__.py index 4b3ad3cf9..60187237c 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -28,7 +28,7 @@ build_output_target, ) from flashdreams.runtime.demo.pipeline import StepOutcome, StepPipeline -from flashdreams.runtime.demo.replay import run_replay_demo +from flashdreams.runtime.demo.replay import OutputSinkFactory, run_replay_demo from flashdreams.runtime.demo.run_modes import ( AsyncSessionDriver, BenchmarkErrorPolicy, @@ -131,6 +131,7 @@ "NullOutputSink", "NullErrorPolicy", "OutputDecision", + "OutputSinkFactory", "OutputSpec", "OutputSink", "PreparedScenario", diff --git a/flashdreams/flashdreams/runtime/demo/app.py b/flashdreams/flashdreams/runtime/demo/app.py index b84e59163..d44dd81bb 100644 --- a/flashdreams/flashdreams/runtime/demo/app.py +++ b/flashdreams/flashdreams/runtime/demo/app.py @@ -6,6 +6,7 @@ from __future__ import annotations import argparse +import sys from abc import ABC, abstractmethod from typing import Any @@ -29,10 +30,18 @@ def main(self, argv: list[str] | None = None) -> None: configure_logging() args = self.parse_args(argv) if args.command == "replay": - run_replay_demo( + result = run_replay_demo( spec=self.replay_spec(args), adapter=self.replay_adapter(), ) + if result.status != "completed": + reason = result.reason or ( + str(result.error) if result.error is not None else None + ) + if reason is None: + reason = f"Replay demo ended with status {result.status!r}." + print(reason, file=sys.stderr) + raise SystemExit(1) return if args.command == "webrtc": context = initialize_cuda_distributed( diff --git a/flashdreams/flashdreams/runtime/demo/replay.py b/flashdreams/flashdreams/runtime/demo/replay.py index 18b873254..d33d6478a 100644 --- a/flashdreams/flashdreams/runtime/demo/replay.py +++ b/flashdreams/flashdreams/runtime/demo/replay.py @@ -5,28 +5,72 @@ from __future__ import annotations +import math from collections.abc import Callable, Sequence +from flashdreams.runtime.canonical import InputCanonicalizer +from flashdreams.runtime.config import InferenceConfig +from flashdreams.runtime.inputs import ( + CanonicalInputs, + CanonicalInputSchema, + InferenceInput, + InferenceInputSchema, + TimeWindow, + UserInputs, + UserInputSchema, +) +from flashdreams.runtime.interfaces import InferenceRuntime, InferenceSession +from flashdreams.runtime.mapping import ( + DeclaresMappingSchema, + InputMapping, + check_mapping_compatibility, +) from flashdreams.runtime.metrics import MetricsRecorder, NullMetricsRecorder from flashdreams.runtime.output import OutputArtifact, OutputTarget -from flashdreams.runtime.runner import run_inference_session +from flashdreams.runtime.types import ( + StepRequest, + StepRequirements, + StepResult, + step_requirements_from_request, +) -from .outputs import build_output_target -from .spec import DemoAdapter, DemoSpec, OutputSpec, WebRTCOutputSpec +from .drivers import BatchSessionDriver, run_demo_session +from .host import ModelWarmupPlan, RuntimeHost +from .outputs import OutputSink, build_output_sink, build_output_target +from .pipeline import StepPipeline +from .run_modes import ( + Mp4ErrorPolicy, + NullErrorPolicy, + RunContext, + RunModeCapabilities, + RunResult, + SessionEdges, + SingleSessionAdmissionPolicy, +) +from .session_inputs import PreparedStep, ProviderCapabilities, UserInputWindow +from .spec import ( + DemoAdapter, + DemoSpec, + OutputSpec, + PreparedScenario, + WebRTCOutputSpec, +) OutputTargetFactory = Callable[[OutputSpec], OutputTarget] InferenceSessionRunner = Callable[..., Sequence[OutputArtifact]] +OutputSinkFactory = Callable[[OutputSpec], OutputSink] def run_replay_demo( *, spec: DemoSpec, adapter: DemoAdapter, - output_target_factory: OutputTargetFactory = build_output_target, + output_target_factory: OutputTargetFactory | None = None, + output_sink_factory: OutputSinkFactory = build_output_sink, metrics: MetricsRecorder | None = None, - runner: InferenceSessionRunner = run_inference_session, -) -> tuple[OutputArtifact, ...]: - """Run one prepared demo scenario through the shared runtime runner.""" + runner: InferenceSessionRunner | None = None, +) -> RunResult: + """Run one prepared replay scenario through the shared batch demo path.""" _require_supported_mode( mode=spec.input_mode, supported=adapter.supported_input_modes(), @@ -55,12 +99,43 @@ def run_replay_demo( if spec.config is None: raise RuntimeError("DemoSpec.config was not initialized.") + if runner is not None or output_target_factory is not None: + return _run_replay_demo_with_compat_runner( + spec=spec, + adapter=adapter, + prepared=prepared, + mapping=mapping, + output_target_factory=output_target_factory or build_output_target, + metrics=metrics, + runner=runner or _default_inference_session_runner(), + ) + + return _run_replay_demo_with_run_mode( + spec=spec, + adapter=adapter, + prepared=prepared, + mapping=mapping, + output_sink_factory=output_sink_factory, + metrics=metrics, + ) + + +def _run_replay_demo_with_compat_runner( + *, + spec: DemoSpec, + adapter: DemoAdapter, + prepared: "PreparedScenario", + mapping: InputMapping, + output_target_factory: OutputTargetFactory, + metrics: MetricsRecorder | None, + runner: InferenceSessionRunner, +) -> RunResult: output = output_target_factory(spec.output) metrics_recorder = metrics or NullMetricsRecorder() - return tuple( + artifacts = tuple( runner( adapter=adapter, - config=spec.config, + config=_require_config(spec), mapping=mapping, canonicalizer=prepared.canonicalizer, source_schema=prepared.source_schema, @@ -70,6 +145,420 @@ def run_replay_demo( metrics=metrics_recorder, ) ) + return RunResult(status="completed", artifacts=artifacts) + + +def _run_replay_demo_with_run_mode( + *, + spec: DemoSpec, + adapter: DemoAdapter, + prepared: "PreparedScenario", + mapping: InputMapping, + output_sink_factory: OutputSinkFactory, + metrics: MetricsRecorder | None, +) -> RunResult: + config = _require_config(spec) + _validate_replay_mapping( + adapter=adapter, + config=config, + mapping=mapping, + source_schema=prepared.source_schema, + canonicalizer=prepared.canonicalizer, + ) + request_state = _ReplayStepRequestState() + runtime = _ReplayRuntimeAdapter( + runtime=adapter.create_runtime(config), + request_state=request_state, + ) + host = RuntimeHost(runtime) + mode = _ReplayRunMode( + request_state=request_state, + output_sink_factory=output_sink_factory, + run_metrics=metrics or NullMetricsRecorder(), + ) + replay_adapter = _ReplayProviderAdapter( + adapter=adapter, + mapping=mapping, + request_state=request_state, + ) + context = mode.create_run_context( + spec=spec, + adapter=replay_adapter, + host=host, + model_warmup_plan=ModelWarmupPlan(), + ) + try: + return run_demo_session( + context=context, + spec=spec, + scenario=prepared, + adapter=replay_adapter, + run_mode=mode, + pipeline=StepPipeline(), + ) + finally: + context.close() + host.close() + + +class _ReplayRunMode: + name = "replay" + capabilities = RunModeCapabilities( + requires_finite_input=True, + supports_artifacts=True, + ) + + def __init__( + self, + *, + request_state: "_ReplayStepRequestState", + output_sink_factory: OutputSinkFactory, + run_metrics: MetricsRecorder, + ) -> None: + self._request_state = request_state + self._output_sink_factory = output_sink_factory + self._run_metrics = run_metrics + + def validate_run(self, *, spec: DemoSpec, adapter: DemoAdapter) -> None: + del spec, adapter + + def validate_session( + self, + *, + spec: DemoSpec, + scenario: "PreparedScenario", + adapter: DemoAdapter, + provider: object, + ) -> None: + del spec, scenario, adapter, provider + + def create_run_context( + self, + *, + spec: DemoSpec, + adapter: DemoAdapter, + host: RuntimeHost, + model_warmup_plan: ModelWarmupPlan, + ) -> RunContext: + del spec, adapter + return RunContext( + host=host, + run_metrics=self._run_metrics, + admission=SingleSessionAdmissionPolicy( + health_check=lambda: host.is_healthy + ), + model_warmup_plan=model_warmup_plan, + ) + + def create_session_edges( + self, + *, + context: RunContext, + spec: DemoSpec, + scenario: "PreparedScenario", + provider: object, + adapter: DemoAdapter, + ) -> SessionEdges: + del provider, adapter + return SessionEdges( + input_source=_ReplayBatchInputSource( + scenario=scenario, + request_state=self._request_state, + ), + output_sink=self._output_sink_factory(spec.output), + cleanup_tasks=context.cleanup_tasks, + error_policy=( + Mp4ErrorPolicy() if spec.output.mode == "mp4" else NullErrorPolicy() + ), + ) + + def select_driver(self) -> BatchSessionDriver: + return BatchSessionDriver() + + +class _ReplayProviderAdapter: + def __init__( + self, + *, + adapter: DemoAdapter, + mapping: InputMapping, + request_state: "_ReplayStepRequestState", + ) -> None: + self._adapter = adapter + self._mapping = mapping + self._request_state = request_state + + @property + def model_id(self) -> str: + return self._adapter.model_id + + @property + def inference_input_schema(self) -> InferenceInputSchema: + return self._adapter.inference_input_schema + + @property + def canonical_input_schema(self) -> CanonicalInputSchema | None: + return self._adapter.canonical_input_schema + + def default_input_mapping(self) -> InputMapping | None: + return self._adapter.default_input_mapping() + + def supported_input_modes(self) -> tuple[str, ...]: + return self._adapter.supported_input_modes() + + def supported_output_modes(self) -> tuple[str, ...]: + return self._adapter.supported_output_modes() + + def validate_config(self, config: InferenceConfig) -> None: + self._adapter.validate_config(config) + + def create_runtime(self, config: InferenceConfig) -> InferenceRuntime: + return self._adapter.create_runtime(config) + + def prepare_scenario(self, spec: DemoSpec) -> "PreparedScenario": + return self._adapter.prepare_scenario(spec) + + def create_model_input_provider( + self, + spec: DemoSpec, + scenario: "PreparedScenario", + ) -> "_ReplayMappingModelInputProvider": + del spec + return _ReplayMappingModelInputProvider( + adapter=self._adapter, + scenario=scenario, + mapping=self._mapping, + request_state=self._request_state, + ) + + +class _ReplayRuntimeAdapter: + def __init__( + self, + *, + runtime: InferenceRuntime, + request_state: "_ReplayStepRequestState", + ) -> None: + self._runtime = runtime + self._request_state = request_state + + def start_session(self, inputs: InferenceInput) -> InferenceSession: + return _ReplaySessionAdapter( + session=self._runtime.start_session(inputs), + request_state=self._request_state, + ) + + def close(self) -> None: + self._runtime.close() + + +class _ReplaySessionAdapter: + def __init__( + self, + *, + session: InferenceSession, + request_state: "_ReplayStepRequestState", + ) -> None: + self._session = session + self._request_state = request_state + + def next_step_requirements(self) -> StepRequirements | None: + next_requirements = getattr(self._session, "next_step_requirements", None) + if callable(next_requirements): + value = next_requirements() + self._request_state.clear() + return value + + request = self._session.next_step_request() + if request is None: + self._request_state.clear() + return None + self._request_state.store(request) + return step_requirements_from_request( + request, + allow_user_input_window=True, + ) + + def next_step_request(self) -> StepRequest | None: + return self._session.next_step_request() + + def step(self, inputs: InferenceInput) -> StepResult: + return self._session.step(inputs) + + def reset(self, inputs: InferenceInput | None = None) -> None: + self._session.reset(inputs) + + def close(self) -> None: + self._session.close() + + +class _ReplayStepRequestState: + def __init__(self) -> None: + self._request: StepRequest | None = None + + def store(self, request: StepRequest) -> None: + self._request = request + + def request_for_window(self, step_index: int) -> StepRequest | None: + request = self._request + if request is None: + return None + if request.step_index != step_index: + raise RuntimeError( + "Replay input source request mismatch: " + f"expected step {request.step_index}, got {step_index}." + ) + return request + + def consume_for_step(self, request: StepRequirements) -> StepRequest: + legacy_request = self.request_for_window(request.step_index) + if legacy_request is not None: + self._request = None + return legacy_request + return StepRequest( + step_index=request.step_index, + inference_input_schema=request.inference_input_schema, + metadata=request.metadata, + ) + + def clear(self) -> None: + self._request = None + + +class _ReplayBatchInputSource: + is_finite = True + is_deterministic = True + + def __init__( + self, + *, + scenario: "PreparedScenario", + request_state: _ReplayStepRequestState, + ) -> None: + self.user_input_schema = scenario.source_schema + self._user_inputs = scenario.user_inputs + self._request_state = request_state + + def is_finished(self) -> bool: + return False + + def next_window(self, request: StepRequirements) -> UserInputWindow: + legacy_request = self._request_state.request_for_window(request.step_index) + window = ( + legacy_request.user_input_window if legacy_request is not None else None + ) + if window is None: + window = _all_user_inputs_window(self._user_inputs) + return UserInputWindow( + start_s=window.start_s, + end_s=window.end_s, + inputs=self._user_inputs, + ) + + +class _ReplayMappingModelInputProvider: + def __init__( + self, + *, + adapter: DemoAdapter, + scenario: "PreparedScenario", + mapping: InputMapping, + request_state: _ReplayStepRequestState, + ) -> None: + self.capabilities = ProviderCapabilities( + supports_recorded_input=True, + deterministic_given_inputs=True, + user_input_schema=scenario.source_schema, + inference_input_schema=adapter.inference_input_schema, + ) + self._scenario = scenario + self._mapping = mapping + self._request_state = request_state + self._step_base_inputs = InferenceInput( + step=scenario.initial_inputs.step, + metadata=scenario.initial_inputs.metadata, + ) + + def prepare_initial_input(self) -> InferenceInput: + self._scenario.canonicalizer.reset() + return self._mapping.map_global_conditioning_inputs( + canonical_inputs=CanonicalInputs(), + inference_input=self._scenario.initial_inputs, + ) + + def prepare_step( + self, + *, + request: StepRequirements, + user_window: UserInputWindow, + ) -> PreparedStep: + legacy_request = self._request_state.consume_for_step(request) + canonical_inputs = self._scenario.canonicalizer.canonicalize( + self._scenario.user_inputs, + window=TimeWindow(start_s=user_window.start_s, end_s=user_window.end_s), + source_schema=self._scenario.source_schema, + ) + return PreparedStep( + inference_input=self._mapping.map_step_inputs( + canonical_inputs=canonical_inputs, + inference_input=self._step_base_inputs, + request=legacy_request, + ) + ) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self._scenario.canonicalizer.reset() + + def close(self) -> None: + return None + + +def _validate_replay_mapping( + *, + adapter: DemoAdapter, + config: InferenceConfig, + mapping: InputMapping, + source_schema: UserInputSchema, + canonicalizer: InputCanonicalizer, +) -> None: + adapter.validate_config(config) + canonical_schema = canonicalizer.canonical_schema(source_schema) + if isinstance(mapping, DeclaresMappingSchema): + compatibility = check_mapping_compatibility( + canonical_schema=canonical_schema, + inference_input_schema=adapter.inference_input_schema, + mapping_schema=mapping.mapping_schema, + ) + compatibility.raise_if_incompatible() + mapping.validate( + canonical_schema=canonical_schema, + inference_input_schema=adapter.inference_input_schema, + ) + + +def _require_config(spec: DemoSpec) -> InferenceConfig: + if spec.config is None: + raise RuntimeError("DemoSpec.config was not initialized.") + return spec.config + + +def _all_user_inputs_window(user_inputs: UserInputs) -> TimeWindow: + if not user_inputs.events: + return TimeWindow(start_s=0.0, end_s=3600.0) + return TimeWindow( + start_s=0.0, + end_s=max( + 3600.0, + math.nextafter(user_inputs.events[-1].timestamp_s, math.inf), + ), + ) + + +def _default_inference_session_runner() -> InferenceSessionRunner: + from flashdreams.runtime.runner import run_inference_session + + return run_inference_session def _require_supported_mode( @@ -88,6 +577,7 @@ def _require_supported_mode( __all__ = [ "InferenceSessionRunner", + "OutputSinkFactory", "OutputTargetFactory", "run_replay_demo", ] diff --git a/flashdreams/tests/test_runtime_demo_api.py b/flashdreams/tests/test_runtime_demo_api.py index 6cdf772e5..e40b62b51 100644 --- a/flashdreams/tests/test_runtime_demo_api.py +++ b/flashdreams/tests/test_runtime_demo_api.py @@ -3,6 +3,7 @@ from __future__ import annotations +import argparse from collections.abc import Sequence from pathlib import Path from types import SimpleNamespace @@ -36,14 +37,21 @@ ) from flashdreams.runtime.demo import ( DemoSpec, + Mp4OutputSink, Mp4OutputSpec, + NullOutputSink, NullOutputSpec, + OutputSink, + OutputSpec, PreparedScenario, + RunResult, WebRTCAppResources, WebRTCOutputSpec, + build_output_sink, build_output_target, run_replay_demo, ) +from flashdreams.runtime.demo.app import DemoApplication from flashdreams.runtime.demo.webrtc import ( serve_webrtc_demo, ) @@ -52,7 +60,42 @@ pytestmark = pytest.mark.ci_cpu -def test_replay_demo_uses_shared_runner() -> None: +def test_replay_demo_uses_shared_batch_path_by_default() -> None: + adapter = _FakeDemoAdapter() + sinks: list[OutputSink] = [] + + def output_sink_factory(output_spec: OutputSpec) -> OutputSink: + sink = build_output_sink(output_spec) + sinks.append(sink) + return sink + + spec = DemoSpec( + model_id="fake-demo", + scenario="valid-scenario", + input_mode="replay", + output=NullOutputSpec(), + ) + + result = run_replay_demo( + spec=spec, + adapter=adapter, + metrics=NullMetricsRecorder(), + output_sink_factory=output_sink_factory, + ) + + assert result.status == "completed" + assert result.artifacts == () + assert len(sinks) == 1 + assert isinstance(sinks[0], NullOutputSink) + assert adapter.create_runtime_called + assert adapter.runtime is not None + assert adapter.runtime.closed + assert adapter.runtime.session is not None + assert adapter.runtime.session.closed + assert adapter.prepare_scenario_calls == [spec] + + +def test_replay_demo_keeps_compat_runner_injection() -> None: adapter = _FakeDemoAdapter() output = _RecordingOutputTarget() calls: list[dict[str, Any]] = [] @@ -68,7 +111,7 @@ def fake_runner(**kwargs: Any) -> Sequence[OutputArtifact]: output=NullOutputSpec(), ) - artifacts = run_replay_demo( + result = run_replay_demo( spec=spec, adapter=adapter, output_target_factory=lambda output_spec: output, @@ -76,7 +119,10 @@ def fake_runner(**kwargs: Any) -> Sequence[OutputArtifact]: runner=fake_runner, ) - assert artifacts == (OutputArtifact(kind="test/artifact", uri="memory://artifact"),) + assert result == RunResult( + status="completed", + artifacts=(OutputArtifact(kind="test/artifact", uri="memory://artifact"),), + ) assert len(calls) == 1 assert calls[0]["adapter"] is adapter assert calls[0]["config"] == spec.config @@ -90,8 +136,9 @@ def fake_runner(**kwargs: Any) -> Sequence[OutputArtifact]: assert not adapter.create_runtime_called -def test_replay_demo_builds_output_target_from_spec(tmp_path: Path) -> None: +def test_replay_demo_builds_output_sink_from_spec(tmp_path: Path) -> None: writer_calls: list[dict[str, Any]] = [] + sinks: list[OutputSink] = [] def fake_writer( video: torch.Tensor, @@ -119,18 +166,24 @@ def fake_writer( output=Mp4OutputSpec(path=tmp_path / "demo.mp4", fps=12), ) - artifacts = run_replay_demo( + result = run_replay_demo( spec=spec, adapter=_FakeDemoAdapter(video_output=True), - output_target_factory=lambda output_spec: build_output_target( - output_spec, - mp4_writer=fake_writer, + output_sink_factory=lambda output_spec: _record_output_sink( + sinks, + build_output_sink( + output_spec, + mp4_writer=fake_writer, + ), ), ) - assert len(artifacts) == 1 - assert artifacts[0].kind == "video/mp4" - assert artifacts[0].uri == str(tmp_path / "demo.mp4") + assert result.status == "completed" + assert len(sinks) == 1 + assert isinstance(sinks[0], Mp4OutputSink) + assert len(result.artifacts) == 1 + assert result.artifacts[0].kind == "video/mp4" + assert result.artifacts[0].uri == str(tmp_path / "demo.mp4") assert writer_calls == [ { "shape": (2, 2, 2, 3), @@ -141,6 +194,63 @@ def fake_writer( ] +def test_replay_demo_mp4_sink_matches_legacy_output_target_payload( + tmp_path: Path, +) -> None: + def writer(records: list[dict[str, Any]]): + def fake_writer( + video: torch.Tensor, + path: Path, + *, + fps: int | float, + layout: str, + install_hint: str, + ) -> Path: + del install_hint + records.append( + { + "bytes": video.detach().cpu().numpy().tobytes(), + "shape": tuple(video.shape), + "path": path, + "fps": fps, + "layout": layout, + } + ) + return path + + return fake_writer + + spec = DemoSpec( + model_id="fake-demo", + scenario="valid-scenario", + input_mode="replay", + output=Mp4OutputSpec(path=tmp_path / "demo.mp4", fps=12), + ) + sink_records: list[dict[str, Any]] = [] + target_records: list[dict[str, Any]] = [] + + sink_result = run_replay_demo( + spec=spec, + adapter=_FakeDemoAdapter(video_output=True), + output_sink_factory=lambda output_spec: build_output_sink( + output_spec, + mp4_writer=writer(sink_records), + ), + ) + target_result = run_replay_demo( + spec=spec, + adapter=_FakeDemoAdapter(video_output=True), + output_target_factory=lambda output_spec: build_output_target( + output_spec, + mp4_writer=writer(target_records), + ), + ) + + assert sink_result.status == "completed" + assert target_result.status == "completed" + assert sink_records == target_records + + def test_replay_demo_fails_before_runtime_creation_when_scenario_invalid() -> None: adapter = _FakeDemoAdapter(scenario_valid=False) output_factory_calls = 0 @@ -170,6 +280,18 @@ def output_factory(output_spec: object) -> OutputTarget: assert output_factory_calls == 0 +def test_replay_demo_step_failure_exits_nonzero_and_prints_reason( + capsys: pytest.CaptureFixture[str], +) -> None: + app = _ReplayOnlyDemoApplication(adapter=_FakeDemoAdapter(fail_step=0)) + + with pytest.raises(SystemExit) as raised: + app.main(["replay"]) + + assert raised.value.code == 1 + assert "step failed" in capsys.readouterr().err + + def test_demo_adapter_declares_supported_modes() -> None: adapter = _FakeDemoAdapter( input_modes=("replay",), @@ -291,6 +413,36 @@ def map_step_inputs( ) +def _record_output_sink(sinks: list[OutputSink], sink: OutputSink) -> OutputSink: + sinks.append(sink) + return sink + + +class _ReplayOnlyDemoApplication(DemoApplication): + def __init__(self, *, adapter: "_FakeDemoAdapter") -> None: + self._adapter = adapter + + def parse_args(self, argv: list[str] | None = None) -> argparse.Namespace: + del argv + return argparse.Namespace(command="replay") + + def replay_spec(self, args: argparse.Namespace) -> DemoSpec: + del args + return DemoSpec( + model_id="fake-demo", + scenario="valid-scenario", + input_mode="replay", + output=NullOutputSpec(), + ) + + def replay_adapter(self) -> "_FakeDemoAdapter": + return self._adapter + + def serve_webrtc(self, args: argparse.Namespace, *, context: Any) -> None: + del args, context + raise AssertionError("webrtc should not run") + + class _FakeDemoAdapter: model_id = "fake-demo" inference_input_schema = InferenceInputSchema( @@ -304,11 +456,13 @@ def __init__( *, scenario_valid: bool = True, video_output: bool = False, + fail_step: int | None = None, input_modes: tuple[str, ...] = ("replay",), output_modes: tuple[str, ...] = ("null", "mp4"), ) -> None: self._scenario_valid = scenario_valid self._video_output = video_output + self._fail_step = fail_step self._input_modes = input_modes self._output_modes = output_modes self.mapping = _ChunkIndexMapping() @@ -344,6 +498,7 @@ def create_runtime(self, config: InferenceConfig) -> InferenceRuntime: self.runtime = _FakeRuntime( inference_input_schema=self.inference_input_schema, video_output=self._video_output, + fail_step=self._fail_step, ) return self.runtime @@ -360,9 +515,11 @@ def __init__( *, inference_input_schema: InferenceInputSchema, video_output: bool, + fail_step: int | None, ) -> None: self._inference_input_schema = inference_input_schema self._video_output = video_output + self._fail_step = fail_step self.session: _FakeSession | None = None self.closed = False @@ -371,6 +528,7 @@ def start_session(self, inputs: InferenceInput) -> InferenceSession: self.session = _FakeSession( inference_input_schema=self._inference_input_schema, video_output=self._video_output, + fail_step=self._fail_step, ) return self.session @@ -384,9 +542,11 @@ def __init__( *, inference_input_schema: InferenceInputSchema, video_output: bool, + fail_step: int | None, ) -> None: self._inference_input_schema = inference_input_schema self._video_output = video_output + self._fail_step = fail_step self.step_index = 0 self.closed = False @@ -403,6 +563,8 @@ def next_step_request(self) -> StepRequest | None: def step(self, inputs: InferenceInput) -> StepResult: self._inference_input_schema.require_step(inputs) + if self._fail_step == self.step_index: + raise RuntimeError("step failed") if self._video_output: result = StepResult.from_video_chunk( step_index=self.step_index, diff --git a/integrations/lingbot/tests/test_demo_api.py b/integrations/lingbot/tests/test_demo_api.py index 2cd5edfac..6682dfa5e 100644 --- a/integrations/lingbot/tests/test_demo_api.py +++ b/integrations/lingbot/tests/test_demo_api.py @@ -135,14 +135,17 @@ def fake_runner(**kwargs: Any) -> Sequence[OutputArtifact]: ), ) - artifacts = run_replay_demo( + result = run_replay_demo( spec=spec, adapter=adapter, output_target_factory=lambda output_spec: output, runner=fake_runner, ) - assert artifacts == (OutputArtifact(kind="video/mp4", uri="memory://lingbot"),) + assert result.status == "completed" + assert result.artifacts == ( + OutputArtifact(kind="video/mp4", uri="memory://lingbot"), + ) assert len(calls) == 1 assert calls[0]["adapter"] is adapter assert calls[0]["config"] == spec.config diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index 1d2f81ef3..cb8d710d0 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -107,14 +107,17 @@ def fake_runner(**kwargs: Any) -> Sequence[OutputArtifact]: ), ) - artifacts = run_replay_demo( + result = run_replay_demo( spec=spec, adapter=adapter, output_target_factory=lambda output_spec: output, runner=fake_runner, ) - assert artifacts == (OutputArtifact(kind="video/mp4", uri="memory://omnidreams"),) + assert result.status == "completed" + assert result.artifacts == ( + OutputArtifact(kind="video/mp4", uri="memory://omnidreams"), + ) assert len(calls) == 1 assert calls[0]["adapter"] is adapter assert calls[0]["config"] == spec.config From 42eeaaf67083228bed5c35fd3200f15fcd820c6b Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 09:43:36 +0000 Subject: [PATCH 22/51] Phase 12: decompose shared WebRTC session edges --- .../flashdreams/serving/webrtc/services.py | 1017 +++++++++++++++++ flashdreams/tests/test_webrtc_services.py | 625 ++++++++++ 2 files changed, 1642 insertions(+) create mode 100644 flashdreams/flashdreams/serving/webrtc/services.py create mode 100644 flashdreams/tests/test_webrtc_services.py diff --git a/flashdreams/flashdreams/serving/webrtc/services.py b/flashdreams/flashdreams/serving/webrtc/services.py new file mode 100644 index 000000000..4602aca58 --- /dev/null +++ b/flashdreams/flashdreams/serving/webrtc/services.py @@ -0,0 +1,1017 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Session-edge services for the shared WebRTC demo run mode. + +These classes are the Phase 12 decomposition layer: they translate WebRTC +transport facts into the shared demo runtime contracts without owning model +execution. The production manager still uses its legacy execution hook until +the realtime-driver adoption phase. +""" + +from __future__ import annotations + +import asyncio +import contextlib +import inspect +import json +import math +import threading +from collections import deque +from collections.abc import Callable, Coroutine, Mapping, MutableSet, Sequence +from concurrent.futures import Future +from dataclasses import dataclass, field +from typing import Any, Literal, Protocol, runtime_checkable + +from flashdreams.runtime._utils import freeze_mapping +from flashdreams.runtime import ( + StepRequirements, + StepResult, + UserInputCapability, + UserInputEvent, + UserInputs, + UserInputSchema, +) +from flashdreams.runtime.demo import ( + AsyncSessionDriver, + DemoAdapter, + DemoSpec, + InMemorySessionMetricsRecorder, + ModelInputProvider, + ModelWarmupPlan, + OutputDecision, + PreparedScenario, + RealtimeSessionDriver, + RealtimeWindowResult, + RunContext, + RunModeCapabilities, + RunResult, + RuntimeHost, + SessionEdges, + SessionInfo, + SingleSessionAdmissionPolicy, + StepPipeline, + UserInputWindow, + WebRTCErrorPolicy, + input_frame_count_from_request, + run_demo_session_async, +) +from flashdreams.runtime.demo.timing import ( + SPARSE_KEY_SEGMENTS_METADATA_KEY, + ActivationResult, + CatchUpPolicy, + DeterministicClock, + RealtimeClock, +) + +from .messages import ( + MESSAGE_TYPE_ACTION, + MESSAGE_TYPE_DISCONNECT, + MESSAGE_TYPE_EVENT, + MESSAGE_TYPE_HEARTBEAT, +) +from .server import SessionBusyError + +WebRTCMessageKind = Literal[ + "action", + "disconnect", + "event", + "heartbeat", + "error", +] + +WebRTCDropPolicy = Literal["none", "drop_newest", "drop_oldest"] + +_CLEAR_EVENT_STATES = frozenset({"clear", "release", "off", "none"}) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class WebRTCMessageResult: + """Result of translating one browser data-channel message.""" + + kind: WebRTCMessageKind + activated: bool = False + error: str | None = None + + +@dataclass(frozen=True, kw_only=True, slots=True) +class WebRTCOfferRequest: + """Browser SDP offer passed to the shared offer/session handler.""" + + sdp: str + type: str + + def __post_init__(self) -> None: + if not self.sdp.strip(): + raise ValueError("WebRTCOfferRequest.sdp must be non-empty.") + if not self.type.strip(): + raise ValueError("WebRTCOfferRequest.type must be non-empty.") + + +@dataclass(frozen=True, kw_only=True, slots=True) +class WebRTCOutputBridgeDecision: + """Immediate delivery decision from a nonblocking WebRTC output bridge.""" + + accepted: bool = True + should_stop: bool = False + dropped: bool = False + drop_policy: WebRTCDropPolicy = "none" + backpressure_s: float = 0.0 + metadata: Mapping[str, object] = field(default_factory=dict) + + def __post_init__(self) -> None: + if self.drop_policy not in {"none", "drop_newest", "drop_oldest"}: + raise ValueError(f"Unsupported drop_policy={self.drop_policy!r}.") + if not math.isfinite(self.backpressure_s) or self.backpressure_s < 0.0: + raise ValueError( + "WebRTCOutputBridgeDecision.backpressure_s must be finite and >= 0." + ) + object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) + + +@runtime_checkable +class WebRTCOfferAnswerer(Protocol): + """Creates an SDP answer after the shared session task has been scheduled.""" + + async def create_answer( + self, + *, + offer: WebRTCOfferRequest, + session_task: asyncio.Task[RunResult], + ) -> Mapping[str, str]: ... + + +@runtime_checkable +class BlockingPreparationService(Protocol): + """Runs blocking scenario preparation outside the aiohttp event loop.""" + + async def run( + self, + func: Callable[..., object], + *args: object, + **kwargs: object, + ) -> object: ... + + +@runtime_checkable +class WebRTCOutputBridge(Protocol): + """Thread-safe bridge from model-worker output writes to WebRTC delivery.""" + + def submit_chunk( + self, + result: StepResult, + *, + generation: int, + force_keyframe: bool = False, + ) -> WebRTCOutputBridgeDecision: ... + + def close(self) -> None: ... + + +@runtime_checkable +class WebRTCSessionEdgeFactory(Protocol): + """Builds per-peer shared session edges on the WebRTC control rank.""" + + def create_session_edges( + self, + *, + context: RunContext, + spec: DemoSpec, + scenario: PreparedScenario, + provider: ModelInputProvider, + adapter: DemoAdapter, + ) -> SessionEdges: ... + + +class AsyncioBlockingPreparationService: + """Default blocking-prep service backed by ``asyncio.to_thread``.""" + + async def run( + self, + func: Callable[..., object], + *args: object, + **kwargs: object, + ) -> object: + return await asyncio.to_thread(func, *args, **kwargs) + + +class WebRTCTransportService: + """Idempotent per-peer transport lifecycle for realtime session edges.""" + + def __init__( + self, + *, + loop: asyncio.AbstractEventLoop | None = None, + on_close: Callable[[str | None], None] | None = None, + ) -> None: + self._loop = loop + self._on_close = on_close + self._closed_signal = _ThreadSafeActivationSignal(loop=loop) + self._lock = threading.Lock() + self._closed = False + self._close_reason: str | None = None + self._close_count = 0 + self.last_client_message_at: float | None = None + + @property + def close_count(self) -> int: + """Number of effective transport closes, after idempotency.""" + + return self._close_count + + @property + def close_reason(self) -> str | None: + return self._close_reason + + @property + def closed_signal(self) -> "_ThreadSafeActivationSignal": + return self._closed_signal + + def mark_client_message(self, timestamp_s: float) -> None: + if not math.isfinite(timestamp_s) or timestamp_s < 0.0: + raise ValueError("timestamp_s must be finite and >= 0.") + self.last_client_message_at = float(timestamp_s) + + def is_active(self) -> bool: + return not self._closed + + def disconnect(self, reason: str = "client disconnected") -> None: + self.close(reason=reason) + + def close(self, reason: str | None = None) -> None: + callback: Callable[[str | None], None] | None = None + with self._lock: + if self._closed: + return + self._closed = True + self._close_reason = reason + self._close_count += 1 + callback = self._on_close + self._closed_signal.set() + if callback is not None: + callback(reason) + + +@dataclass(slots=True) +class WebRTCActivationPolicy: + """Activate on the first browser action/event or stop on disconnect.""" + + input_source: "WebRTCInputSource" + transport: WebRTCTransportService + timeout_s: float | None = None + timeout_reason: str = "activation timed out" + anchor_clock: bool = True + + def __post_init__(self) -> None: + if self.timeout_s is not None and self.timeout_s <= 0.0: + raise ValueError("timeout_s must be > 0 when set.") + if not self.timeout_reason.strip(): + raise ValueError("timeout_reason must be non-empty.") + + async def wait_until_active( + self, + clock: RealtimeClock | DeterministicClock, + ) -> ActivationResult: + if self.input_source.activation_signal.is_set(): + self._anchor(clock) + return ActivationResult(activated=True) + if not self.transport.is_active(): + return ActivationResult( + activated=False, + reason=self.transport.close_reason or "transport closed", + ) + + activation_task = asyncio.create_task( + self.input_source.activation_signal.wait() + ) + closed_task = asyncio.create_task(self.transport.closed_signal.wait()) + try: + done, pending = await asyncio.wait( + {activation_task, closed_task}, + timeout=self.timeout_s, + return_when=asyncio.FIRST_COMPLETED, + ) + if not done: + return ActivationResult( + activated=False, + reason=self.timeout_reason, + ) + for task in done: + task.result() + if not self.transport.is_active(): + return ActivationResult( + activated=False, + reason=self.transport.close_reason or "transport closed", + ) + self._anchor(clock) + return ActivationResult(activated=True) + finally: + for task in (activation_task, closed_task): + if not task.done(): + task.cancel() + await asyncio.gather( + activation_task, + closed_task, + return_exceptions=True, + ) + + def _anchor(self, clock: RealtimeClock | DeterministicClock) -> None: + if not self.anchor_clock or not clock.is_realtime: + return + now = getattr(clock, "now", None) + anchor = getattr(clock, "anchor", None) + if callable(now) and callable(anchor): + anchor(now()) + + +@dataclass(slots=True) +class WebRTCInputSource: + """Realtime source fed by browser data-channel events.""" + + resampler: Any + max_lag_s: float | None = None + catch_up_policy: CatchUpPolicy = "fold" + user_input_schema: UserInputSchema = field( + default_factory=lambda: WEBRTC_USER_INPUT_SCHEMA + ) + is_finite: bool = False + is_deterministic: bool = False + _activation_signal: "_ThreadSafeActivationSignal" = field( + default_factory=lambda: _ThreadSafeActivationSignal(), + init=False, + repr=False, + ) + _events: deque[UserInputEvent] = field( + default_factory=deque, + init=False, + repr=False, + ) + + def __post_init__(self) -> None: + if self.max_lag_s is not None and ( + not math.isfinite(self.max_lag_s) or self.max_lag_s < 0.0 + ): + raise ValueError("max_lag_s must be finite and >= 0.") + if self.catch_up_policy != "fold": + raise NotImplementedError( + f"Catch-up policy {self.catch_up_policy!r} has no WebRTC analog yet." + ) + + @property + def activation_signal(self) -> "_ThreadSafeActivationSignal": + return self._activation_signal + + def is_finished(self) -> bool: + return False + + def reset(self, *, start_v: float) -> None: + self.resampler.reset(start_v=start_v) + self._events.clear() + self._activation_signal.clear() + + def handle_browser_message( + self, + raw_message: object, + *, + timestamp_s: float, + ) -> WebRTCMessageResult: + """Translate one browser data-channel message into typed user inputs.""" + + if not isinstance(raw_message, str): + return WebRTCMessageResult(kind="error", error="Expected text payload.") + try: + payload = json.loads(raw_message) + except json.JSONDecodeError: + return WebRTCMessageResult(kind="error", error="Invalid JSON payload.") + if not isinstance(payload, dict): + return WebRTCMessageResult( + kind="error", + error="Payload must be a JSON object.", + ) + return self.handle_browser_payload(payload, timestamp_s=timestamp_s) + + def handle_browser_payload( + self, + payload: Mapping[str, object], + *, + timestamp_s: float, + ) -> WebRTCMessageResult: + message_type = str(payload.get("type", "")).strip().lower() + if message_type == MESSAGE_TYPE_HEARTBEAT: + return WebRTCMessageResult(kind="heartbeat") + if message_type == MESSAGE_TYPE_DISCONNECT: + return WebRTCMessageResult(kind="disconnect") + if message_type == MESSAGE_TYPE_EVENT: + return self._record_text_event(payload, timestamp_s=timestamp_s) + if message_type == MESSAGE_TYPE_ACTION: + action_payload = payload.get("action", payload) + if not isinstance(action_payload, Mapping): + return WebRTCMessageResult( + kind="error", + error="'action' must be an object.", + ) + return self._record_action( + {str(key): value for key, value in action_payload.items()}, + timestamp_s=timestamp_s, + ) + return WebRTCMessageResult( + kind="error", + error=( + "Unsupported message type, expected " + "'action', 'event', 'heartbeat', or 'disconnect'." + ), + ) + + def record_user_event( + self, + *, + timestamp_s: float, + event_type: str, + payload: Mapping[str, object], + source_event_id: str | None = None, + activate: bool = True, + ) -> None: + event = UserInputEvent( + timestamp_s=timestamp_s, + event_type=event_type, + payload=dict(payload), + source="webrtc", + source_event_id=source_event_id, + ) + self.user_input_schema.validate_event(event) + self._events.append(event) + if activate: + self._activation_signal.set() + + async def next_realtime_window( + self, + *, + request: StepRequirements, + clock: RealtimeClock, + ) -> RealtimeWindowResult: + input_frame_count = input_frame_count_from_request(request) + chunk_duration_s = input_frame_count * float(self.resampler.dt) + if chunk_duration_s <= 0.0: + raise ValueError("Realtime resampler dt must produce a positive window.") + + window_end_s = float(self.resampler.next_chunk_start_v) + chunk_duration_s + await clock.wait_until_window_end(window_end_s) + catch_up = clock.catch_up( + request=request, + max_lag_s=self.max_lag_s + if self.max_lag_s is not None + else chunk_duration_s, + policy=self.catch_up_policy, + ) + start_s = float(self.resampler.next_chunk_start_v) + segments, frame_times = self.resampler.sample_chunk(input_frame_count) + end_s = float(self.resampler.next_chunk_start_v) + window = RealtimeWindowResult( + window=_user_input_window( + start_s=start_s, + end_s=end_s, + frame_times=tuple(frame_times), + inputs=UserInputs(events=self._events_for_window(start_s, end_s)), + metadata={SPARSE_KEY_SEGMENTS_METADATA_KEY: tuple(segments)}, + ), + catch_up=catch_up, + ) + self._prune_events(before_s=start_s) + return window + + def _record_action( + self, + payload: Mapping[str, object], + *, + timestamp_s: float, + ) -> WebRTCMessageResult: + event = str(payload.get("event", "")).strip().lower() + if event == "step": + self._activation_signal.set() + return WebRTCMessageResult(kind="action", activated=True) + if event not in {"keydown", "keyup"}: + return WebRTCMessageResult( + kind="error", + error=f"Unsupported event={event!r}; expected 'keydown' or 'keyup'.", + ) + key = str(payload.get("key", "")).strip() + if not key: + return WebRTCMessageResult( + kind="error", + error="Action payload must include non-empty 'key'.", + ) + self.resampler.on_edge(arrival_t=timestamp_s, event=event, key=key) + self.record_user_event( + timestamp_s=timestamp_s, + event_type="key_down" if event == "keydown" else "key_up", + payload={"key": key}, + ) + return WebRTCMessageResult(kind="action", activated=True) + + def _record_text_event( + self, + payload: Mapping[str, object], + *, + timestamp_s: float, + ) -> WebRTCMessageResult: + state = str(payload.get("state", "trigger")).strip().lower() or "trigger" + event_id = str(payload.get("event_id", payload.get("id", ""))).strip() + clears = state in _CLEAR_EVENT_STATES + if not event_id and not clears: + return WebRTCMessageResult( + kind="error", + error=( + "Event payload must include non-empty 'event_id' unless state " + "clears the active event." + ), + ) + active_event_id = None if clears else event_id + self.record_user_event( + timestamp_s=timestamp_s, + event_type="text_event", + payload={"event_id": active_event_id, "state": state}, + source_event_id=active_event_id, + ) + return WebRTCMessageResult(kind="event", activated=True) + + def _events_for_window( + self, + start_s: float, + end_s: float, + ) -> tuple[UserInputEvent, ...]: + return tuple( + sorted( + ( + event + for event in self._events + if start_s <= event.timestamp_s < end_s + ), + key=lambda event: event.timestamp_s, + ) + ) + + def _prune_events(self, *, before_s: float) -> None: + self._events = deque( + event for event in self._events if event.timestamp_s >= before_s + ) + + +class WebRTCOutputSink: + """Output sink that schedules WebRTC media delivery without blocking.""" + + produces_artifacts = False + + def __init__(self, *, bridge: WebRTCOutputBridge) -> None: + self._bridge = bridge + self._opened = False + self._closed = True + self._bridge_closed = False + self._generation = 0 + self._force_keyframe = False + self.session_info: SessionInfo | None = None + + def open(self, session_info: SessionInfo) -> None: + self.session_info = session_info + self._opened = True + self._closed = False + self._generation = 0 + self._force_keyframe = True + + def begin_generation(self, generation: int) -> None: + if generation < 0: + raise ValueError("generation must be >= 0.") + self._generation = generation + self._force_keyframe = True + + def write(self, result: StepResult) -> OutputDecision: + if not self._opened or self._closed: + raise RuntimeError("Cannot write to a closed output sink.") + decision = self._bridge.submit_chunk( + result, + generation=self._generation, + force_keyframe=self._force_keyframe, + ) + self._force_keyframe = False + return OutputDecision( + should_stop=decision.should_stop, + dropped=decision.dropped, + drop_policy=decision.drop_policy, + backpressure_s=decision.backpressure_s, + metadata=decision.metadata, + ) + + def close(self) -> Sequence[object]: + if self._bridge_closed: + return () + self._closed = True + self._opened = False + self._bridge.close() + self._bridge_closed = True + return () + + +class ThreadSafeWebRTCOutputBridge: + """Schedule async encoder delivery from any thread without blocking writes.""" + + def __init__( + self, + *, + loop: asyncio.AbstractEventLoop, + video_encoder: Any, + video_track: Any, + max_pending_chunks: int = 2, + close_track: bool = True, + on_delivery: Callable[[object], None] | None = None, + on_error: Callable[[BaseException], None] | None = None, + ) -> None: + if max_pending_chunks <= 0: + raise ValueError("max_pending_chunks must be > 0.") + self._loop = loop + self._video_encoder = video_encoder + self._video_track = video_track + self._max_pending_chunks = max_pending_chunks + self._close_track = close_track + self._on_delivery = on_delivery + self._on_error = on_error + self._pending: set[Future[object]] = set() + self._lock = threading.Lock() + self._closed = False + + @property + def pending_count(self) -> int: + with self._lock: + return len(self._pending) + + def submit_chunk( + self, + result: StepResult, + *, + generation: int, + force_keyframe: bool = False, + ) -> WebRTCOutputBridgeDecision: + del generation + with self._lock: + if self._closed: + return WebRTCOutputBridgeDecision( + accepted=False, + should_stop=True, + dropped=True, + drop_policy="drop_newest", + metadata={"reason": "closed"}, + ) + if len(self._pending) >= self._max_pending_chunks: + return WebRTCOutputBridgeDecision( + accepted=False, + dropped=True, + drop_policy="drop_newest", + metadata={"reason": "pending queue full"}, + ) + future = asyncio.run_coroutine_threadsafe( + self._deliver(result, force_keyframe=force_keyframe), + self._loop, + ) + self._pending.add(future) + future.add_done_callback(self._on_done) + + return WebRTCOutputBridgeDecision( + accepted=True, + backpressure_s=self._track_backpressure_s(), + ) + + def close(self) -> None: + with self._lock: + if self._closed: + return + self._closed = True + pending = tuple(self._pending) + for future in pending: + future.cancel() + if self._close_track: + self._schedule_track_close() + + async def _deliver( + self, + result: StepResult, + *, + force_keyframe: bool, + ) -> object: + return await self._video_encoder.deliver_chunk( + result, + self._video_track, + force_keyframe=force_keyframe, + ) + + def _on_done(self, future: Future[object]) -> None: + with self._lock: + self._pending.discard(future) + if future.cancelled(): + return + try: + result = future.result() + except BaseException as exc: + if self._on_error is not None: + self._on_error(exc) + return + if self._on_delivery is not None: + self._on_delivery(result) + + def _track_backpressure_s(self) -> float: + qsize = getattr(self._video_track, "qsize", None) + fps = getattr(self._video_track, "fps", None) or getattr( + self._video_encoder, + "fps", + None, + ) + if not callable(qsize) or fps is None: + return 0.0 + try: + queue_depth = int(qsize()) + frames_per_second = float(fps) + except (TypeError, ValueError): + return 0.0 + if frames_per_second <= 0.0: + return 0.0 + return max(0.0, queue_depth / frames_per_second) + + def _schedule_track_close(self) -> None: + close = getattr(self._video_track, "close", None) + if not callable(close): + return + try: + result = close() + if inspect.isawaitable(result): + asyncio.run_coroutine_threadsafe(result, self._loop) + except BaseException as exc: + if self._on_error is not None: + self._on_error(exc) + + +class WebRTCRunMode: + """Shared realtime run mode that delegates peer-specific edges to WebRTC.""" + + name = "webrtc" + capabilities = RunModeCapabilities( + realtime=True, + supports_backpressure=True, + supports_interactive_events=True, + ) + + def __init__( + self, + *, + edge_factory: WebRTCSessionEdgeFactory, + blocking_preparation: BlockingPreparationService | None = None, + driver: AsyncSessionDriver | None = None, + error_policy: WebRTCErrorPolicy | None = None, + ) -> None: + self._edge_factory = edge_factory + self._blocking_preparation = ( + blocking_preparation or AsyncioBlockingPreparationService() + ) + self._driver = driver or RealtimeSessionDriver() + self._error_policy = error_policy or WebRTCErrorPolicy() + + @property + def blocking_preparation(self) -> BlockingPreparationService: + return self._blocking_preparation + + @property + def error_policy(self) -> WebRTCErrorPolicy: + return self._error_policy + + def validate_run(self, *, spec: DemoSpec, adapter: DemoAdapter) -> None: + del adapter + if spec.output.mode != "webrtc": + raise ValueError("WebRTCRunMode requires WebRTC output.") + + def validate_session( + self, + *, + spec: DemoSpec, + scenario: PreparedScenario, + adapter: DemoAdapter, + provider: ModelInputProvider, + ) -> None: + del spec, scenario, adapter + if not provider.capabilities.supports_realtime_clock: + raise ValueError("WebRTC providers must support realtime clocks.") + + def create_run_context( + self, + *, + spec: DemoSpec, + adapter: DemoAdapter, + host: RuntimeHost, + model_warmup_plan: ModelWarmupPlan, + ) -> RunContext: + del spec, adapter + services: dict[str, object] = {} + if host.is_control_rank: + services["blocking_preparation"] = self._blocking_preparation + return RunContext( + host=host, + run_metrics=InMemorySessionMetricsRecorder(), + admission=SingleSessionAdmissionPolicy( + health_check=lambda: host.is_control_rank and host.is_healthy + ), + model_warmup_plan=model_warmup_plan, + services=services, + ) + + def create_session_edges( + self, + *, + context: RunContext, + spec: DemoSpec, + scenario: PreparedScenario, + provider: ModelInputProvider, + adapter: DemoAdapter, + ) -> SessionEdges: + if not context.host.is_control_rank: + raise RuntimeError("WebRTC session edges are control-rank only.") + edges = self._edge_factory.create_session_edges( + context=context, + spec=spec, + scenario=scenario, + provider=provider, + adapter=adapter, + ) + if not isinstance(edges, SessionEdges): + raise TypeError( + "WebRTC edge factory must return SessionEdges, " + f"got {type(edges).__name__}." + ) + return edges + + def select_driver(self) -> AsyncSessionDriver: + return self._driver + + +class WebRTCSessionOfferHandler: + """Reserve, prepare, and launch one WebRTC session before SDP negotiation.""" + + def __init__( + self, + *, + context: RunContext, + spec: DemoSpec, + adapter: DemoAdapter, + run_mode: WebRTCRunMode, + answerer: WebRTCOfferAnswerer, + pipeline: StepPipeline | None = None, + session_helper: Callable[..., Coroutine[Any, Any, RunResult]] | None = None, + busy_message: str = "Another WebRTC session is already active.", + session_tasks: MutableSet[asyncio.Task[RunResult]] | None = None, + ) -> None: + self._context = context + self._spec = spec + self._adapter = adapter + self._run_mode = run_mode + self._answerer = answerer + self._pipeline = pipeline or StepPipeline() + self._session_helper = session_helper or run_demo_session_async + self._busy_message = busy_message + self._session_tasks = session_tasks if session_tasks is not None else set() + + async def handle_offer( + self, + *, + offer_sdp: str, + offer_type: str, + ) -> Mapping[str, str]: + if not self._context.host.is_control_rank: + raise RuntimeError("WebRTC offers are handled only on the control rank.") + reservation = self._context.admission.try_reserve() + if reservation is None: + raise SessionBusyError(self._busy_message) + + task: asyncio.Task[RunResult] | None = None + try: + scenario = await self._run_blocking_prepare(self._spec) + task = asyncio.create_task( + self._session_helper( + context=self._context, + spec=self._spec, + scenario=scenario, + adapter=self._adapter, + run_mode=self._run_mode, + pipeline=self._pipeline, + reservation=reservation, + ) + ) + self._track_task(task) + answer = await self._answerer.create_answer( + offer=WebRTCOfferRequest(sdp=offer_sdp, type=offer_type), + session_task=task, + ) + return dict(answer) + except Exception: + if task is not None: + task.cancel() + with contextlib.suppress(asyncio.CancelledError): + await task + reservation.release() + raise + + async def _run_blocking_prepare(self, spec: DemoSpec) -> PreparedScenario: + service = self._run_mode.blocking_preparation + result = await service.run(self._adapter.prepare_scenario, spec) + if not isinstance(result, PreparedScenario): + raise TypeError( + "DemoAdapter.prepare_scenario must return PreparedScenario, " + f"got {type(result).__name__}." + ) + return result + + def _track_task(self, task: asyncio.Task[RunResult]) -> None: + self._session_tasks.add(task) + task.add_done_callback(self._discard_task) + + def _discard_task(self, task: asyncio.Task[RunResult]) -> None: + self._session_tasks.discard(task) + if task.cancelled(): + return + with contextlib.suppress(Exception): + task.exception() + + +WEBRTC_USER_INPUT_SCHEMA = UserInputSchema( + capabilities=( + UserInputCapability( + event_type="key_down", + input_modality="keyboard", + payload_fields=frozenset({"key"}), + ), + UserInputCapability( + event_type="key_up", + input_modality="keyboard", + payload_fields=frozenset({"key"}), + ), + UserInputCapability( + event_type="text_event", + input_modality="text", + payload_fields=frozenset({"event_id", "state"}), + ), + ), + description="browser WebRTC data-channel events", +) + + +def _user_input_window( + *, + start_s: float, + end_s: float, + frame_times: Sequence[float], + inputs: UserInputs, + metadata: Mapping[str, object], +) -> UserInputWindow: + return UserInputWindow( + start_s=start_s, + end_s=end_s, + frame_times=frame_times, + inputs=inputs, + metadata=metadata, + ) + + +class _ThreadSafeActivationSignal: + def __init__(self, *, loop: asyncio.AbstractEventLoop | None = None) -> None: + self._loop = loop + self._event = asyncio.Event() + + def is_set(self) -> bool: + return self._event.is_set() + + async def wait(self) -> object: + return await self._event.wait() + + def set(self) -> None: + if self._event.is_set(): + return + if self._loop is not None and self._loop.is_running(): + self._loop.call_soon_threadsafe(self._event.set) + return + self._event.set() + + def clear(self) -> None: + self._event.clear() + + +__all__ = [ + "AsyncioBlockingPreparationService", + "BlockingPreparationService", + "ThreadSafeWebRTCOutputBridge", + "WEBRTC_USER_INPUT_SCHEMA", + "WebRTCActivationPolicy", + "WebRTCInputSource", + "WebRTCMessageResult", + "WebRTCOfferAnswerer", + "WebRTCOfferRequest", + "WebRTCOutputBridge", + "WebRTCOutputBridgeDecision", + "WebRTCOutputSink", + "WebRTCRunMode", + "WebRTCSessionEdgeFactory", + "WebRTCSessionOfferHandler", + "WebRTCTransportService", +] diff --git a/flashdreams/tests/test_webrtc_services.py b/flashdreams/tests/test_webrtc_services.py new file mode 100644 index 000000000..2755045f6 --- /dev/null +++ b/flashdreams/tests/test_webrtc_services.py @@ -0,0 +1,625 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +import asyncio +import json +import threading +from collections.abc import Mapping, Sequence +from typing import Any + +import pytest + +from flashdreams.runtime import ( + CanonicalInputSchema, + IdentityInputMapping, + InferenceConfig, + InferenceInput, + InferenceInputSchema, + InferenceRuntime, + InferenceSession, + InputMapping, + StepRequirements, + StepResult, +) +from flashdreams.runtime.demo import ( + DemoAdapter, + DemoSpec, + InMemorySessionMetricsRecorder, + ModelInputProvider, + ModelWarmupPlan, + OutputDecision, + PreparedScenario, + ProviderCapabilities, + RealtimeSessionDriver, + RunContext, + RunResult, + RuntimeHost, + SessionEdges, + SessionInfo, + StepPipeline, + UserInputWindow, + WebRTCOutputSpec, + run_demo_session_async, +) +from flashdreams.runtime.demo.timing import ResamplerRealtimeClock +from flashdreams.serving.webrtc.server import SessionBusyError +from flashdreams.serving.webrtc.services import ( + AsyncioBlockingPreparationService, + ThreadSafeWebRTCOutputBridge, + WebRTCActivationPolicy, + WebRTCInputSource, + WebRTCOfferRequest, + WebRTCOutputSink, + WebRTCRunMode, + WebRTCSessionOfferHandler, + WebRTCTransportService, +) + +pytestmark = pytest.mark.ci_cpu + + +@pytest.mark.asyncio +async def test_webrtc_offer_handler_calls_shared_session_helper() -> None: + spec = _webrtc_spec() + adapter = _FakeAdapter() + mode = WebRTCRunMode(edge_factory=_FinishedEdgeFactory()) + context = mode.create_run_context( + spec=spec, + adapter=adapter, + host=RuntimeHost(_UnusedRuntime()), + model_warmup_plan=ModelWarmupPlan(), + ) + answerer = _RecordingAnswerer() + helper_calls: list[DemoSpec] = [] + handler = WebRTCSessionOfferHandler( + context=context, + spec=spec, + adapter=adapter, + run_mode=mode, + answerer=answerer, + session_helper=lambda **kwargs: _record_completed_helper( + helper_calls, + **kwargs, + ), + ) + + answer = await handler.handle_offer(offer_sdp="v=0\r\n", offer_type="offer") + + assert answer == {"sdp": "answer-sdp", "type": "answer"} + assert helper_calls == [spec] + assert answerer.offers == [WebRTCOfferRequest(sdp="v=0\r\n", type="offer")] + + +@pytest.mark.asyncio +async def test_webrtc_busy_rejects_before_prepare_provider_or_answer() -> None: + spec = _webrtc_spec() + adapter = _FakeAdapter() + mode = WebRTCRunMode(edge_factory=_FinishedEdgeFactory()) + context = RunContext( + host=RuntimeHost(_UnusedRuntime()), + run_metrics=InMemorySessionMetricsRecorder(), + admission=_BusyAdmission(), + ) + answerer = _RecordingAnswerer() + handler = WebRTCSessionOfferHandler( + context=context, + spec=spec, + adapter=adapter, + run_mode=mode, + answerer=answerer, + ) + + with pytest.raises(SessionBusyError): + await handler.handle_offer(offer_sdp="v=0\r\n", offer_type="offer") + + assert adapter.prepare_thread_id is None + assert adapter.providers == [] + assert answerer.offers == [] + + +@pytest.mark.asyncio +async def test_webrtc_scenario_prepare_runs_off_event_loop_thread() -> None: + loop_thread_id = threading.get_ident() + spec = _webrtc_spec() + adapter = _FakeAdapter() + + result = await AsyncioBlockingPreparationService().run( + adapter.prepare_scenario, + spec, + ) + + assert isinstance(result, PreparedScenario) + assert adapter.prepare_thread_id is not None + assert adapter.prepare_thread_id != loop_thread_id + + +@pytest.mark.asyncio +async def test_webrtc_input_source_emits_typed_user_inputs() -> None: + resampler = _FakeResampler(dt=0.1, start_v=0.0) + source = WebRTCInputSource(resampler=resampler) + source.handle_browser_message( + json.dumps({"type": "action", "action": {"event": "keydown", "key": "w"}}), + timestamp_s=0.05, + ) + source.handle_browser_message( + json.dumps({"type": "event", "event_id": "prompt-1", "state": "trigger"}), + timestamp_s=0.06, + ) + clock = ResamplerRealtimeClock( + resampler=resampler, + now_fn=lambda: 0.2, + sleep_fn=_record_sleep, + ) + + result = await source.next_realtime_window( + request=StepRequirements(step_index=0, input_frame_count=2), + clock=clock, + ) + + assert source.activation_signal.is_set() + assert resampler.edges == [(0.05, "keydown", "w")] + assert result.window.start_s == pytest.approx(0.0) + assert result.window.end_s == pytest.approx(0.2) + assert [event.event_type for event in result.window.inputs.events] == [ + "key_down", + "text_event", + ] + assert result.window.inputs.events[0].payload == {"key": "w"} + assert result.window.inputs.events[1].payload == { + "event_id": "prompt-1", + "state": "trigger", + } + + +@pytest.mark.asyncio +async def test_webrtc_output_sink_uses_nonblocking_threadsafe_bridge() -> None: + loop = asyncio.get_running_loop() + encoder = _BlockingEncoder() + track = _FakeVideoTrack() + deliveries: list[object] = [] + bridge = ThreadSafeWebRTCOutputBridge( + loop=loop, + video_encoder=encoder, + video_track=track, + on_delivery=deliveries.append, + ) + sink = WebRTCOutputSink(bridge=bridge) + sink.open(SessionInfo()) + + decision = sink.write(StepResult(step_index=0, frame_count=1)) + + assert isinstance(decision, OutputDecision) + assert not decision.dropped + await asyncio.wait_for(encoder.started.wait(), timeout=1.0) + assert not encoder.release.is_set() + assert bridge.pending_count == 1 + + encoder.release.set() + await asyncio.wait_for(encoder.done.wait(), timeout=1.0) + await asyncio.sleep(0) + + assert deliveries == ["delivered"] + assert bridge.pending_count == 0 + sink.close() + + +@pytest.mark.asyncio +async def test_disconnect_closes_transport_and_releases_reservation_once() -> None: + spec = _webrtc_spec() + adapter = _FakeAdapter() + transport_closed: list[str | None] = [] + transport = WebRTCTransportService(on_close=transport_closed.append) + mode = WebRTCRunMode( + edge_factory=_DisconnectedEdgeFactory(transport=transport), + driver=RealtimeSessionDriver(cleanup_timeout_s=1.0), + ) + context = mode.create_run_context( + spec=spec, + adapter=adapter, + host=RuntimeHost(_SessionRuntime()), + model_warmup_plan=ModelWarmupPlan(), + ) + reservation = context.admission.try_reserve() + assert reservation is not None + transport.disconnect("browser disconnect") + + result = await _record_async_helper( + [], + context=context, + spec=spec, + scenario=adapter.prepare_scenario(spec), + adapter=adapter, + run_mode=mode, + pipeline=StepPipeline(), + reservation=reservation, + ) + transport.close("cleanup close") + + assert result.status == "not_activated" + assert result.reason == "browser disconnect" + assert transport.close_count == 1 + assert transport_closed == ["browser disconnect"] + assert reservation.release_count == 1 # ty:ignore[unresolved-attribute] + assert adapter.providers[0].close_count == 1 + + +def test_webrtc_run_mode_objects_are_control_rank_only() -> None: + spec = _webrtc_spec() + adapter = _FakeAdapter() + edge_factory = _FinishedEdgeFactory() + mode = WebRTCRunMode(edge_factory=edge_factory) + worker_context = mode.create_run_context( + spec=spec, + adapter=adapter, + host=RuntimeHost(_UnusedRuntime(), is_control_rank=False), + model_warmup_plan=ModelWarmupPlan(), + ) + + assert worker_context.services == {} + assert worker_context.admission.try_reserve() is None + with pytest.raises(RuntimeError, match="control-rank only"): + mode.create_session_edges( + context=worker_context, + spec=spec, + scenario=adapter.prepare_scenario(spec), + provider=_FakeProvider(), + adapter=adapter, + ) + + control_context = mode.create_run_context( + spec=spec, + adapter=adapter, + host=RuntimeHost(_UnusedRuntime(), is_control_rank=True), + model_warmup_plan=ModelWarmupPlan(), + ) + assert set(control_context.services) == {"blocking_preparation"} + assert control_context.admission.try_reserve() is not None + + +def _webrtc_spec() -> DemoSpec: + return DemoSpec( + model_id="fake-demo", + input_mode="keyboard-driving", + output=WebRTCOutputSpec(port=8081), + ) + + +async def _record_sleep(delay_s: float) -> None: + del delay_s + + +async def _record_async_helper( + calls: list[DemoSpec], + **kwargs: Any, +) -> RunResult: + calls.append(kwargs["spec"]) + return await run_demo_session_async(**kwargs) + + +async def _record_completed_helper( + calls: list[DemoSpec], + **kwargs: Any, +) -> RunResult: + calls.append(kwargs["spec"]) + return RunResult(status="completed") + + +class _FakeAdapter: + model_id = "fake-demo" + inference_input_schema = InferenceInputSchema() + canonical_input_schema = CanonicalInputSchema() + + def __init__(self) -> None: + self.prepare_thread_id: int | None = None + self.providers: list[_FakeProvider] = [] + + def supported_input_modes(self) -> tuple[str, ...]: + return ("keyboard-driving",) + + def supported_output_modes(self) -> tuple[str, ...]: + return ("webrtc",) + + def default_input_mapping(self) -> InputMapping: + return IdentityInputMapping() + + def validate_config(self, config: InferenceConfig) -> None: + if config.model_id != self.model_id: + raise ValueError(f"Unsupported model_id={config.model_id!r}.") + + def create_runtime(self, config: InferenceConfig) -> InferenceRuntime: + self.validate_config(config) + return _SessionRuntime() + + def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: + self.prepare_thread_id = threading.get_ident() + assert spec.model_id == self.model_id + return PreparedScenario(initial_inputs=InferenceInput()) + + def create_model_input_provider( + self, + spec: DemoSpec, + scenario: PreparedScenario, + ) -> "_FakeProvider": + del spec, scenario + provider = _FakeProvider() + self.providers.append(provider) + return provider + + +class _FakeProvider: + capabilities = ProviderCapabilities( + supports_realtime_clock=True, + supports_reset=True, + deterministic_given_inputs=False, + ) + + def __init__(self) -> None: + self.close_count = 0 + + def prepare_initial_input(self) -> InferenceInput: + return InferenceInput() + + def prepare_step( + self, + *, + request: StepRequirements, + user_window: UserInputWindow, + ) -> Any: + del request, user_window + raise AssertionError("disconnected tests must stop before step prep") + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + + def close(self) -> None: + self.close_count += 1 + + +class _SessionRuntime: + def start_session(self, inputs: InferenceInput) -> InferenceSession: + del inputs + return _NeverSteppedSession() + + def close(self) -> None: + return + + +class _UnusedRuntime: + def start_session(self, inputs: InferenceInput) -> InferenceSession: + del inputs + raise AssertionError("runtime should not be used") + + def close(self) -> None: + return + + +class _NeverSteppedSession: + def next_step_requirements(self) -> StepRequirements | None: + return StepRequirements(step_index=0, input_frame_count=1) + + def next_step_request(self) -> None: + return None + + def step(self, inputs: InferenceInput) -> StepResult: + del inputs + raise AssertionError("disconnected tests must stop before stepping") + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + + def close(self) -> None: + return + + def session_info(self) -> SessionInfo: + return SessionInfo() + + +class _FinishedEdgeFactory: + def __init__(self) -> None: + self.edges: list[SessionEdges] = [] + + def create_session_edges( + self, + *, + context: RunContext, + spec: DemoSpec, + scenario: PreparedScenario, + provider: ModelInputProvider, + adapter: DemoAdapter, + ) -> SessionEdges: + del spec, scenario, provider, adapter + edges = SessionEdges( + input_source=_FinishedRealtimeInputSource(), + output_sink=_RecordingOutputSink(), + cleanup_tasks=context.cleanup_tasks, + metrics=InMemorySessionMetricsRecorder(), + transport=WebRTCTransportService(), + clock=_InstantClock(), + activation=_AlreadyActive(), + ) + self.edges.append(edges) + return edges + + +class _DisconnectedEdgeFactory: + def __init__(self, *, transport: WebRTCTransportService) -> None: + self.transport = transport + + def create_session_edges( + self, + *, + context: RunContext, + spec: DemoSpec, + scenario: PreparedScenario, + provider: ModelInputProvider, + adapter: DemoAdapter, + ) -> SessionEdges: + del spec, scenario, provider, adapter + resampler = _FakeResampler(dt=0.1, start_v=0.0) + source = WebRTCInputSource(resampler=resampler) + return SessionEdges( + input_source=source, + output_sink=_RecordingOutputSink(), + cleanup_tasks=context.cleanup_tasks, + metrics=InMemorySessionMetricsRecorder(), + transport=self.transport, + clock=ResamplerRealtimeClock( + resampler=resampler, + now_fn=lambda: 0.0, + sleep_fn=_record_sleep, + ), + activation=WebRTCActivationPolicy( + input_source=source, + transport=self.transport, + ), + ) + + +class _FinishedRealtimeInputSource: + is_finite = False + is_deterministic = False + user_input_schema = _FakeProvider.capabilities.user_input_schema + + def is_finished(self) -> bool: + return True + + async def next_realtime_window( + self, + *, + request: StepRequirements, + clock: Any, + ) -> Any: + del request, clock + raise AssertionError("finished async driver should not request windows") + + +class _RecordingOutputSink: + produces_artifacts = False + + def __init__(self) -> None: + self.close_count = 0 + + def open(self, session_info: SessionInfo) -> None: + del session_info + + def begin_generation(self, generation: int) -> None: + del generation + + def write(self, result: StepResult) -> OutputDecision: + del result + return OutputDecision() + + def close(self) -> Sequence[Any]: + self.close_count += 1 + return () + + +class _AlreadyActive: + timeout_s = None + + async def wait_until_active(self, clock: Any) -> Any: + del clock + return type("Activation", (), {"activated": True, "reason": None})() + + +class _InstantClock: + is_realtime = True + is_deterministic = False + + def now(self) -> float: + return 0.0 + + def anchor(self, wall_time_s: float) -> None: + del wall_time_s + + async def wait_until_window_end(self, end_s: float) -> None: + del end_s + + async def apply_backpressure(self, requested_s: float) -> None: + del requested_s + + def catch_up( + self, + *, + request: StepRequirements, + max_lag_s: float, + policy: str, + ) -> Any: + del request, max_lag_s, policy + return type("CatchUp", (), {"skipped_s": 0.0})() + + +class _BusyAdmission: + def try_reserve(self) -> None: + return None + + +class _RecordingAnswerer: + def __init__(self) -> None: + self.offers: list[WebRTCOfferRequest] = [] + + async def create_answer( + self, + *, + offer: WebRTCOfferRequest, + session_task: asyncio.Task[RunResult], + ) -> Mapping[str, str]: + self.offers.append(offer) + result = await session_task + assert result.status == "completed" + return {"sdp": "answer-sdp", "type": "answer"} + + +class _FakeResampler: + def __init__(self, *, dt: float, start_v: float) -> None: + self.dt = dt + self.next_chunk_start_v = start_v + self.edges: list[tuple[float, str, str]] = [] + + def reset(self, *, start_v: float) -> None: + self.next_chunk_start_v = start_v + self.edges.clear() + + def on_edge(self, *, arrival_t: float, event: str, key: str) -> None: + self.edges.append((arrival_t, event, key)) + + def sample_chunk( + self, + num_frames: int, + ) -> tuple[tuple[tuple[float, float, frozenset[str]], ...], tuple[float, ...]]: + start = self.next_chunk_start_v + frame_times = tuple(start + index * self.dt for index in range(num_frames)) + end = start + num_frames * self.dt + self.next_chunk_start_v = end + return (((start, end, frozenset({"w"})),), frame_times) + + +class _BlockingEncoder: + fps = 30 + + def __init__(self) -> None: + self.started = asyncio.Event() + self.release = asyncio.Event() + self.done = asyncio.Event() + + async def deliver_chunk( + self, + result: StepResult, + track: Any, + *, + force_keyframe: bool = False, + ) -> str: + del result, track, force_keyframe + self.started.set() + await self.release.wait() + self.done.set() + return "delivered" + + +class _FakeVideoTrack: + fps = 30 + + def qsize(self) -> int: + return 0 From 75a2a7d55e773b4348d0507f71489551b300ea47 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 10:01:09 +0000 Subject: [PATCH 23/51] runtime: fix WebRTC service import ordering --- flashdreams/flashdreams/serving/webrtc/services.py | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/flashdreams/flashdreams/serving/webrtc/services.py b/flashdreams/flashdreams/serving/webrtc/services.py index 4602aca58..caedf71c1 100644 --- a/flashdreams/flashdreams/serving/webrtc/services.py +++ b/flashdreams/flashdreams/serving/webrtc/services.py @@ -23,7 +23,6 @@ from dataclasses import dataclass, field from typing import Any, Literal, Protocol, runtime_checkable -from flashdreams.runtime._utils import freeze_mapping from flashdreams.runtime import ( StepRequirements, StepResult, @@ -32,6 +31,7 @@ UserInputs, UserInputSchema, ) +from flashdreams.runtime._utils import freeze_mapping from flashdreams.runtime.demo import ( AsyncSessionDriver, DemoAdapter, From 60bef842ee60a4549f6fe92fe2d74293d202d30a Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 10:45:11 +0000 Subject: [PATCH 24/51] runtime: harden realtime WebRTC service contracts --- .../flashdreams/runtime/demo/drivers.py | 4 +- .../flashdreams/runtime/demo/run_modes.py | 7 ++ .../flashdreams/serving/webrtc/encoders.py | 49 +++++++- .../flashdreams/serving/webrtc/media.py | 14 ++- .../flashdreams/serving/webrtc/nvenc.py | 109 ++++++++++++++++++ .../flashdreams/serving/webrtc/services.py | 39 ++++++- .../test_demo_runtime_realtime_driver.py | 40 ++++++- flashdreams/tests/test_encoders.py | 21 +++- flashdreams/tests/test_webrtc_services.py | 87 +++++++++++++- 9 files changed, 350 insertions(+), 20 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index d87912feb..7ec4e2e27 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -271,6 +271,8 @@ async def run_one_session( break if outcome.output.backpressure_s > 0: await clock.apply_backpressure(outcome.output.backpressure_s) + except DriverInvariantError: + raise except Exception as exc: action = session_edges.error_policy.handle(exc) session_edges.metrics.record_error(exc, action) @@ -830,7 +832,7 @@ async def _close_model_resource_async( timeout=timeout_s, ) except asyncio.TimeoutError as exc: - session_edges.record_cleanup_error(exc) + session_edges.record_orphaned_cleanup(exc) return False except Exception as exc: session_edges.record_cleanup_error(exc) diff --git a/flashdreams/flashdreams/runtime/demo/run_modes.py b/flashdreams/flashdreams/runtime/demo/run_modes.py index ee71be193..cb93e4e7f 100644 --- a/flashdreams/flashdreams/runtime/demo/run_modes.py +++ b/flashdreams/flashdreams/runtime/demo/run_modes.py @@ -341,6 +341,13 @@ def record_cleanup_error(self, exc: Exception) -> None: except Exception: return + def record_orphaned_cleanup(self, exc: Exception) -> None: + """Record timed-out worker cleanup without blocking teardown.""" + try: + self.metrics.record_orphaned_cleanup(exc) + except Exception: + return + def close_result( self, *, diff --git a/flashdreams/flashdreams/serving/webrtc/encoders.py b/flashdreams/flashdreams/serving/webrtc/encoders.py index 5eea69991..b07d9324c 100644 --- a/flashdreams/flashdreams/serving/webrtc/encoders.py +++ b/flashdreams/flashdreams/serving/webrtc/encoders.py @@ -20,7 +20,7 @@ import importlib.util from dataclasses import dataclass -from typing import TYPE_CHECKING, Literal, Protocol, runtime_checkable +from typing import TYPE_CHECKING, Any, Literal, Protocol, cast, runtime_checkable import torch from aiortc import MediaStreamTrack @@ -69,6 +69,20 @@ class VideoEncoder(Protocol): def create_track(self, *, maxsize: int) -> BufferedVideoTrack | NVENCVideoTrack: ... + def prepare_chunk_payload( + self, + result: StepResult, + track: MediaStreamTrack, + ) -> object: ... + + async def deliver_prepared_chunk( + self, + payload: object, + track: MediaStreamTrack, + *, + force_keyframe: bool = False, + ) -> ChunkDeliveryResult: ... + async def deliver_chunk( self, result: StepResult, @@ -110,10 +124,24 @@ def create_track(self, *, maxsize: int) -> BufferedVideoTrack: return BufferedVideoTrack(fps=self.fps, maxsize=maxsize) - async def deliver_chunk( + def prepare_chunk_payload( self, result: StepResult, track: MediaStreamTrack, + ) -> tuple[object, ...]: + from flashdreams.serving.webrtc.media import BufferedVideoTrack + + if not isinstance(track, BufferedVideoTrack): + raise TypeError( + "DefaultRTCEncoder requires a BufferedVideoTrack; got " + f"{type(track).__name__}. Create it via encoder.create_track()." + ) + return track.prepare_result_frames(result) + + async def deliver_prepared_chunk( + self, + payload: object, + track: MediaStreamTrack, *, force_keyframe: bool = False, ) -> ChunkDeliveryResult: @@ -128,7 +156,9 @@ async def deliver_chunk( "DefaultRTCEncoder requires a BufferedVideoTrack; got " f"{type(track).__name__}. Create it via encoder.create_track()." ) - enqueued = await track.enqueue_result(result) + if not isinstance(payload, tuple): + raise TypeError("DefaultRTCEncoder payload must be a tuple of RGB frames.") + enqueued = await track.enqueue_frames(cast(Any, payload)) return ChunkDeliveryResult( backend=self.backend, num_frames=enqueued, @@ -136,6 +166,19 @@ async def deliver_chunk( encode_ms=0.0, ) + async def deliver_chunk( + self, + result: StepResult, + track: MediaStreamTrack, + *, + force_keyframe: bool = False, + ) -> ChunkDeliveryResult: + return await self.deliver_prepared_chunk( + self.prepare_chunk_payload(result, track), + track, + force_keyframe=force_keyframe, + ) + def close(self) -> None: return diff --git a/flashdreams/flashdreams/serving/webrtc/media.py b/flashdreams/flashdreams/serving/webrtc/media.py index 470569132..7e358966f 100644 --- a/flashdreams/flashdreams/serving/webrtc/media.py +++ b/flashdreams/flashdreams/serving/webrtc/media.py @@ -81,16 +81,26 @@ def maxsize(self) -> int: def qsize(self) -> int: return self._frames.qsize() - async def enqueue_result(self, result: StepResult) -> int: + def prepare_result_frames(self, result: StepResult) -> tuple[np.ndarray, ...]: + if self._closed: + return () + return tuple(self._frame_converter(result)) + + async def enqueue_frames(self, frames: Sequence[np.ndarray]) -> int: if self._closed: return 0 - frames = await asyncio.to_thread(self._frame_converter, result) for i, frame in enumerate(frames): if self._closed: return i await self._frames.put(frame) return len(frames) + async def enqueue_result(self, result: StepResult) -> int: + if self._closed: + return 0 + frames = await asyncio.to_thread(self.prepare_result_frames, result) + return await self.enqueue_frames(frames) + async def recv(self) -> VideoFrame: if self._closed: raise MediaStreamError diff --git a/flashdreams/flashdreams/serving/webrtc/nvenc.py b/flashdreams/flashdreams/serving/webrtc/nvenc.py index 3a4cc93dc..94a52200f 100644 --- a/flashdreams/flashdreams/serving/webrtc/nvenc.py +++ b/flashdreams/flashdreams/serving/webrtc/nvenc.py @@ -23,6 +23,7 @@ import contextlib import time from collections.abc import Callable +from dataclasses import dataclass from fractions import Fraction from typing import TYPE_CHECKING, Any @@ -30,6 +31,7 @@ from aiortc import MediaStreamTrack from av.packet import Packet from loguru import logger +from torch import Tensor from flashdreams.runtime import StepResult from flashdreams.serving.webrtc.encoders import ChunkDeliveryResult @@ -57,6 +59,13 @@ _RTP_VIDEO_CLOCK = 90_000 +@dataclass(frozen=True, slots=True) +class NVENCChunkPayload: + """Encoder-owned CUDA frames prepared before async delivery is scheduled.""" + + frames: Tensor + + def _payload_contains_nal_type(payload: bytes, nal_type: int) -> bool: """Scan an Annex-B H.264 payload for the presence of a specific NAL type.""" i = 0 @@ -240,6 +249,35 @@ def create_track(self, *, maxsize: int) -> NVENCVideoTrack: return NVENCVideoTrack(fps=self.fps, maxsize=maxsize) + def prepare_chunk_payload( + self, + result: StepResult, + track: MediaStreamTrack, + ) -> NVENCChunkPayload: + from flashdreams.serving.webrtc.media import NVENCVideoTrack + + if not isinstance(track, NVENCVideoTrack): + raise TypeError( + "PyNvHardwareEncoder requires an NVENCVideoTrack; got " + f"{type(track).__name__}. Create it via encoder.create_track()." + ) + return NVENCChunkPayload(frames=_result_to_abgr_frames(result)) + + async def deliver_prepared_chunk( + self, + payload: object, + track: MediaStreamTrack, + *, + force_keyframe: bool = False, + ) -> ChunkDeliveryResult: + if not isinstance(payload, NVENCChunkPayload): + raise TypeError("PyNvHardwareEncoder payload must be an NVENCChunkPayload.") + return await self._deliver_prepared_frames( + payload.frames, + track, + force_keyframe=force_keyframe, + ) + async def deliver_chunk( self, result: StepResult, @@ -297,6 +335,63 @@ def _stream(packet: Packet) -> None: encode_ms=encode_ms, ) + async def _deliver_prepared_frames( + self, + frames: Tensor, + track: MediaStreamTrack, + *, + force_keyframe: bool = False, + ) -> ChunkDeliveryResult: + from flashdreams.serving.webrtc.media import NVENCVideoTrack + + if not isinstance(track, NVENCVideoTrack): + raise TypeError( + "PyNvHardwareEncoder requires an NVENCVideoTrack; got " + f"{type(track).__name__}. Create it via encoder.create_track()." + ) + loop = asyncio.get_running_loop() + emitted = 0 + enqueued = 0 + + def _stream(packet: Packet) -> None: + nonlocal emitted, enqueued + emitted += 1 + enqueue = track.enqueue_encoded_packet(packet) + try: + future = asyncio.run_coroutine_threadsafe( + enqueue, + loop, + ) + except RuntimeError: + enqueue.close() + return + try: + accepted = future.result() + except Exception: + return + if accepted: + enqueued += 1 + + _num_frames, num_keyframes, encode_ms = await asyncio.to_thread( + self.encode_frames_sync, + frames, + force_keyframe=force_keyframe, + on_packet=_stream, + ) + if enqueued < emitted: + logger.debug( + "NVENC track closed while enqueueing encoded chunk; " + "enqueued {} of {} packet(s).", + enqueued, + emitted, + ) + return ChunkDeliveryResult( + backend=self.backend, + num_frames=enqueued, + num_keyframes=num_keyframes, + encode_ms=encode_ms, + ) + def encode_chunk_sync( self, result: StepResult, @@ -311,6 +406,20 @@ def encode_chunk_sync( just to get access to the emitted packets. """ frames = _result_to_abgr_frames(result) + return self.encode_frames_sync( + frames, + force_keyframe=force_keyframe, + on_packet=on_packet, + ) + + def encode_frames_sync( + self, + frames: Tensor, + *, + force_keyframe: bool = False, + on_packet: Callable[[Packet], None] | None = None, + ) -> tuple[int, int, float]: + """Encode preconverted ``ABGR`` frames for prepared async delivery.""" if not frames.is_cuda: raise ValueError("expected CUDA tensor for hardware encode path") num_frames = frames.shape[0] diff --git a/flashdreams/flashdreams/serving/webrtc/services.py b/flashdreams/flashdreams/serving/webrtc/services.py index caedf71c1..3acc863a6 100644 --- a/flashdreams/flashdreams/serving/webrtc/services.py +++ b/flashdreams/flashdreams/serving/webrtc/services.py @@ -649,7 +649,30 @@ def submit_chunk( generation: int, force_keyframe: bool = False, ) -> WebRTCOutputBridgeDecision: - del generation + prepare = getattr(self._video_encoder, "prepare_chunk_payload", None) + deliver = getattr(self._video_encoder, "deliver_prepared_chunk", None) + if not callable(prepare) or not callable(deliver): + raise TypeError( + "ThreadSafeWebRTCOutputBridge requires a video encoder with " + "prepare_chunk_payload(...) and deliver_prepared_chunk(...)." + ) + with self._lock: + if self._closed: + return WebRTCOutputBridgeDecision( + accepted=False, + should_stop=True, + dropped=True, + drop_policy="drop_newest", + metadata={"reason": "closed"}, + ) + if len(self._pending) >= self._max_pending_chunks: + return WebRTCOutputBridgeDecision( + accepted=False, + dropped=True, + drop_policy="drop_newest", + metadata={"reason": "pending queue full"}, + ) + payload = prepare(result, self._video_track) with self._lock: if self._closed: return WebRTCOutputBridgeDecision( @@ -667,7 +690,11 @@ def submit_chunk( metadata={"reason": "pending queue full"}, ) future = asyncio.run_coroutine_threadsafe( - self._deliver(result, force_keyframe=force_keyframe), + self._deliver( + payload, + generation=generation, + force_keyframe=force_keyframe, + ), self._loop, ) self._pending.add(future) @@ -691,12 +718,14 @@ def close(self) -> None: async def _deliver( self, - result: StepResult, + payload: object, *, + generation: int, force_keyframe: bool, ) -> object: - return await self._video_encoder.deliver_chunk( - result, + del generation + return await self._video_encoder.deliver_prepared_chunk( + payload, self._video_track, force_keyframe=force_keyframe, ) diff --git a/flashdreams/tests/test_demo_runtime_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py index f561c7646..960b2baa2 100644 --- a/flashdreams/tests/test_demo_runtime_realtime_driver.py +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -223,6 +223,37 @@ async def test_realtime_driver_invariant_finalizes_edges_before_reraising() -> N assert metrics.closed +@pytest.mark.asyncio +async def test_realtime_step_invariant_reraises_without_error_policy() -> None: + runtime = _FakeRealtimeRuntime(session=_FakeRealtimeSession(num_steps=1)) + host = RuntimeHost(runtime) + output = _RecordingOutputSink() + transport = _RecordingTransport() + metrics = InMemorySessionMetricsRecorder() + edges = _edges( + output=output, + transport=transport, + metrics=metrics, + error_policy=_DropOutputErrorPolicy(), + ) + + try: + with pytest.raises(DriverInvariantError, match="step invariant"): + await RealtimeSessionDriver().run_one_session( + host=host, + provider=_FakeRealtimeProvider(), + session_edges=edges, + pipeline=_InvariantPipeline(), + ) + finally: + host.close() + + assert metrics.errors == [] + assert output.close_count == 1 + assert transport.close_count == 1 + assert metrics.closed + + @pytest.mark.asyncio async def test_run_context_close_async_drains_registered_cleanup_task() -> None: runtime = _FakeRealtimeRuntime(session=_FakeRealtimeSession(num_steps=1)) @@ -306,7 +337,8 @@ async def test_shielded_cleanup_timeout_bounds_shutdown() -> None: assert result.status == "cancelled" assert host.unhealthy_reason == "model-affine cleanup timed out" - assert len(metrics.cleanup_errors) == 1 + assert metrics.cleanup_errors == [] + assert len(metrics.orphaned_cleanup_errors) == 1 assert metrics.closed @@ -704,6 +736,12 @@ async def call_async( return await super().call_async(func, *args, **kwargs) +class _InvariantPipeline(StepPipeline): + def execute_step(self, **kwargs: object) -> Any: + del kwargs + raise DriverInvariantError("step invariant") + + class _RecordingOutputSink: produces_artifacts = False diff --git a/flashdreams/tests/test_encoders.py b/flashdreams/tests/test_encoders.py index 610217807..06fa38e85 100644 --- a/flashdreams/tests/test_encoders.py +++ b/flashdreams/tests/test_encoders.py @@ -30,7 +30,7 @@ import asyncio import sys import threading -from collections.abc import Callable +from collections.abc import Callable, Sequence from fractions import Fraction from types import ModuleType, SimpleNamespace from unittest.mock import MagicMock, patch @@ -371,10 +371,18 @@ class _FakeBufferedVideoTrack: def __init__(self) -> None: self.enqueued_results: list[StepResult] = [] + self.enqueued_frames: list[object] = [] - async def enqueue_result(self, result: StepResult) -> int: + def prepare_result_frames(self, result: StepResult) -> tuple[object, ...]: self.enqueued_results.append(result) - return result.frame_count + return tuple(object() for _ in range(result.frame_count)) + + async def enqueue_frames(self, frames: Sequence[object]) -> int: + self.enqueued_frames.extend(frames) + return len(frames) + + async def enqueue_result(self, result: StepResult) -> int: + return await self.enqueue_frames(self.prepare_result_frames(result)) class TestDefaultRTCEncoderDeliver: @@ -405,6 +413,7 @@ async def test_deliver_chunk_returns_frames_from_track( assert result.num_frames == 4 assert result.num_keyframes == 0 assert fake_track.enqueued_results == [step_result] + assert len(fake_track.enqueued_frames) == 4 @pytest.mark.parametrize( ("layout", "shape"), @@ -429,7 +438,7 @@ async def test_software_conversion_uses_declared_layout( await track.close() @pytest.mark.asyncio - async def test_software_path_defers_host_conversion_to_track(self) -> None: + async def test_software_path_prepares_host_frames_with_track(self) -> None: from flashdreams.serving.webrtc.media import BufferedVideoTrack source = torch.zeros((2, 3, 2, 2), dtype=torch.uint8) @@ -447,7 +456,9 @@ def _converter(delivered: StepResult) -> list[np.ndarray]: return [np.zeros((2, 2, 3), dtype=np.uint8) for _ in range(2)] track = BufferedVideoTrack(fps=30, maxsize=2, frame_converter=_converter) - delivery = await DefaultRTCEncoder(fps=30).deliver_chunk(step_result, track) + encoder = DefaultRTCEncoder(fps=30) + payload = encoder.prepare_chunk_payload(step_result, track) + delivery = await encoder.deliver_prepared_chunk(payload, track) assert delivery.num_frames == 2 assert seen == [step_result] diff --git a/flashdreams/tests/test_webrtc_services.py b/flashdreams/tests/test_webrtc_services.py index 2755045f6..25a1a1ba2 100644 --- a/flashdreams/tests/test_webrtc_services.py +++ b/flashdreams/tests/test_webrtc_services.py @@ -187,11 +187,13 @@ async def test_webrtc_output_sink_uses_nonblocking_threadsafe_bridge() -> None: ) sink = WebRTCOutputSink(bridge=bridge) sink.open(SessionInfo()) + step_result = StepResult(step_index=0, frame_count=1) - decision = sink.write(StepResult(step_index=0, frame_count=1)) + decision = sink.write(step_result) assert isinstance(decision, OutputDecision) assert not decision.dropped + assert encoder.prepared_payloads == [step_result.step_index] await asyncio.wait_for(encoder.started.wait(), timeout=1.0) assert not encoder.release.is_set() assert bridge.pending_count == 1 @@ -201,10 +203,64 @@ async def test_webrtc_output_sink_uses_nonblocking_threadsafe_bridge() -> None: await asyncio.sleep(0) assert deliveries == ["delivered"] + assert encoder.delivered_payloads == [{"step_index": step_result.step_index}] assert bridge.pending_count == 0 sink.close() +@pytest.mark.asyncio +async def test_webrtc_output_bridge_prepares_payload_before_async_delivery() -> None: + loop = asyncio.get_running_loop() + encoder = _BlockingEncoder() + track = _FakeVideoTrack() + bridge = ThreadSafeWebRTCOutputBridge( + loop=loop, + video_encoder=encoder, + video_track=track, + ) + sink = WebRTCOutputSink(bridge=bridge) + sink.open(SessionInfo()) + + decision = sink.write(StepResult(step_index=7, frame_count=1)) + + assert not decision.dropped + assert encoder.prepared_payloads == [7] + assert encoder.delivered_payloads == [] + + encoder.release.set() + await asyncio.wait_for(encoder.done.wait(), timeout=1.0) + await asyncio.sleep(0) + + assert encoder.delivered_payloads == [{"step_index": 7}] + sink.close() + + +@pytest.mark.asyncio +async def test_webrtc_output_bridge_drops_full_queue_before_payload_prepare() -> None: + loop = asyncio.get_running_loop() + encoder = _BlockingEncoder() + bridge = ThreadSafeWebRTCOutputBridge( + loop=loop, + video_encoder=encoder, + video_track=_FakeVideoTrack(), + max_pending_chunks=1, + ) + sink = WebRTCOutputSink(bridge=bridge) + sink.open(SessionInfo()) + + first = sink.write(StepResult(step_index=0, frame_count=1)) + second = sink.write(StepResult(step_index=1, frame_count=1)) + + assert not first.dropped + assert second.dropped + assert second.drop_policy == "drop_newest" + assert encoder.prepared_payloads == [0] + + encoder.release.set() + await asyncio.wait_for(encoder.done.wait(), timeout=1.0) + sink.close() + + @pytest.mark.asyncio async def test_disconnect_closes_transport_and_releases_reservation_once() -> None: spec = _webrtc_spec() @@ -603,20 +659,45 @@ def __init__(self) -> None: self.started = asyncio.Event() self.release = asyncio.Event() self.done = asyncio.Event() + self.prepared_payloads: list[int] = [] + self.delivered_payloads: list[object] = [] - async def deliver_chunk( + def prepare_chunk_payload( self, result: StepResult, track: Any, + ) -> object: + del track + self.prepared_payloads.append(result.step_index) + return {"step_index": result.step_index} + + async def deliver_prepared_chunk( + self, + payload: object, + track: Any, *, force_keyframe: bool = False, ) -> str: - del result, track, force_keyframe + del track, force_keyframe + self.delivered_payloads.append(payload) self.started.set() await self.release.wait() self.done.set() return "delivered" + async def deliver_chunk( + self, + result: StepResult, + track: Any, + *, + force_keyframe: bool = False, + ) -> str: + return await self.deliver_prepared_chunk( + self.prepare_chunk_payload(result, track), + track, + force_keyframe=force_keyframe, + ) + class _FakeVideoTrack: fps = 30 From 6ca870c170571f5d882d8d8754ce54064dc94f0b Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 20:05:11 +0000 Subject: [PATCH 25/51] runtime: route WebRTC sessions through realtime driver --- .../flashdreams/serving/webrtc/manager.py | 790 +++++++++++++++++- .../flashdreams/serving/webrtc/media.py | 22 + .../flashdreams/serving/webrtc/services.py | 135 ++- flashdreams/tests/test_webrtc_manager.py | 238 +++++- flashdreams/tests/test_webrtc_services.py | 45 + 5 files changed, 1175 insertions(+), 55 deletions(-) diff --git a/flashdreams/flashdreams/serving/webrtc/manager.py b/flashdreams/flashdreams/serving/webrtc/manager.py index a7c6c87a4..4e20959cc 100644 --- a/flashdreams/flashdreams/serving/webrtc/manager.py +++ b/flashdreams/flashdreams/serving/webrtc/manager.py @@ -23,12 +23,37 @@ ) from loguru import logger +from flashdreams.runtime.demo import ( + SPARSE_KEY_SEGMENTS_METADATA_KEY, + DemoSpec, + InMemorySessionMetricsRecorder, + ModelInputProvider, + PreparedScenario, + PreparedStep, + ProviderCapabilities, + ResamplerRealtimeClock, + RunContext, + RunResult, + RuntimeHost, + SessionEdges, + SessionInfo, + SingleSessionAdmissionPolicy, + StepPipeline, + UserInputWindow, + WebRTCErrorPolicy, + WebRTCOutputSpec, + run_demo_session_async, +) from flashdreams.runtime.inputs import ( + CanonicalInputSchema, InferenceInput, + InferenceInputSchema, TimeWindow, UserInputEvent, UserInputs, + UserInputSchema, ) +from flashdreams.runtime.mapping import InputMapping from flashdreams.runtime.types import StepRequest, StepResult from flashdreams.serving.realtime.input import ( DEFAULT_SUPPORTED_KEYS, @@ -55,6 +80,18 @@ WebRTCSessionRuntime, ) from flashdreams.serving.webrtc.server import SessionBusyError +from flashdreams.serving.webrtc.services import ( + WEBRTC_SKIPPED_INPUTS_METADATA_KEY, + WEBRTC_SKIPPED_WINDOW_METADATA_KEY, + WEBRTC_USER_INPUT_SCHEMA, + ThreadSafeWebRTCOutputBridge, + WebRTCActivationPolicy, + WebRTCChunkDelivery, + WebRTCInputSource, + WebRTCOutputSink, + WebRTCRunMode, + WebRTCTransportService, +) from flashdreams.serving.webrtc.warmup import ( run_loopback_warmup_session, wait_for_ice_gathering_complete, @@ -78,6 +115,10 @@ """Maximum unconsumed raw events kept for an ``InferenceSession`` step.""" _RELEASE_USER_EVENT_TYPES = frozenset({"key_up"}) _KEY_USER_EVENT_TYPES = frozenset({"key_down", "key_up"}) +_SESSION_INPUT_KEY = "webrtc_session_input" +_STEP_REQUEST_KEY = "webrtc_step_request" +_SEGMENTS_KEY = "webrtc_segments" +_FRAME_TIMES_KEY = "webrtc_frame_times" _RuntimeT = TypeVar("_RuntimeT", bound=WebRTCSessionRuntime) _RuntimeConfigT = TypeVar("_RuntimeConfigT", bound=WebRTCRuntimeConfig) @@ -140,6 +181,392 @@ def _stat_int(stats: Mapping[str, float | int], name: str) -> int: return int(round(_stat_float(stats, name))) +def _runtime_drives_inference_session(runtime: Any) -> bool: + return callable(getattr(runtime, "start_inference_session", None)) + + +def _run_on_event_loop(loop: asyncio.AbstractEventLoop, awaitable: Any) -> Any: + """Run one legacy async WebRTC runtime call from a RuntimeHost worker.""" + return asyncio.run_coroutine_threadsafe(awaitable, loop).result() + + +def _step_request_from_requirements( + request: Any, + *, + window: TimeWindow, +) -> StepRequest: + metadata = dict(getattr(request, "metadata", {})) + metadata["input_frame_count"] = request.input_frame_count + steady_output_frame_count = getattr(request, "steady_output_frame_count", None) + if steady_output_frame_count is not None: + metadata["steady_output_frame_count"] = steady_output_frame_count + return StepRequest( + step_index=request.step_index, + inference_input_schema=getattr(request, "inference_input_schema", None), + user_input_window=window, + metadata=metadata, + ) + + +class _LegacyWebRTCRuntimeAdapter: + """Shared compatibility adapter from old async WebRTC runtimes to RuntimeHost.""" + + def __init__(self, *, runtime: Any, loop: asyncio.AbstractEventLoop) -> None: + self._runtime = runtime + self._loop = loop + + def reset_for_new_session(self, session_input: Any = None) -> None: + _run_on_event_loop( + self._loop, + self._runtime.reset_for_new_session(session_input=session_input), + ) + + def start_session(self, inputs: InferenceInput) -> "_LegacyWebRTCSessionAdapter": + del inputs + inference_session = None + if _runtime_drives_inference_session(self._runtime): + inference_session = _run_on_event_loop( + self._loop, + self._runtime.start_inference_session(), + ) + return _LegacyWebRTCSessionAdapter( + runtime=self._runtime, + inference_session=inference_session, + loop=self._loop, + ) + + def close(self) -> None: + # The underlying async runtime is owned by BaseWebRTCSessionManager and + # closed from shutdown(); RuntimeHost only owns this adapter's worker. + return + + +class _LegacyWebRTCSessionAdapter: + """RuntimeHost-facing session view over a legacy WebRTC runtime/session.""" + + def __init__( + self, + *, + runtime: Any, + inference_session: Any | None, + loop: asyncio.AbstractEventLoop, + ) -> None: + self._runtime = runtime + self._inference_session = inference_session + self._loop = loop + + def session_info(self) -> SessionInfo: + steady_frames: int | None = None + try: + steady_frames = int(self._runtime.peek_steady_output_num_frames()) + except Exception: + steady_frames = None + return SessionInfo(steady_output_frame_count=steady_frames) + + def next_step_request(self) -> StepRequest | None: + if self._inference_session is not None: + return self._inference_session.next_step_request() + return self._runtime.next_step_request() + + def step(self, inputs: InferenceInput) -> StepResult: + if self._inference_session is not None: + result = self._inference_session.step(inputs) + else: + result = _run_on_event_loop( + self._loop, + self._runtime.step( + request=inputs.step[_STEP_REQUEST_KEY], + segments=list(inputs.step[_SEGMENTS_KEY]), + frame_times=list(inputs.step[_FRAME_TIMES_KEY]), + ), + ) + request = inputs.step[_STEP_REQUEST_KEY] + if result.step_index != request.step_index: + raise RuntimeError( + "Runtime result step does not match its request: " + f"requested {request.step_index}, got {result.step_index}." + ) + if not isinstance(result, StepResult): + raise TypeError( + "WebRTC session steps must produce StepResult, got " + f"{type(result).__name__}." + ) + return result + + def reset(self, inputs: InferenceInput | None = None) -> None: + session_input = None + if inputs is not None: + session_input = inputs.global_conditioning.get(_SESSION_INPUT_KEY) + _run_on_event_loop( + self._loop, + self._runtime.reset_for_new_session(session_input=session_input), + ) + if self._inference_session is not None: + self._inference_session = _run_on_event_loop( + self._loop, + self._runtime.start_inference_session(), + ) + + def close(self) -> None: + close = getattr(self._inference_session, "close", None) + if callable(close): + close() + + +class _LegacyWebRTCModelInputProvider: + """Shared provider used until model-specific WebRTC providers land.""" + + def __init__(self, *, runtime: Any, session_input: Any = None) -> None: + self._runtime = runtime + self._session_input = session_input + self._uses_inference_session = _runtime_drives_inference_session(runtime) + self._session_input_state_advanced = False + self.capabilities = ProviderCapabilities( + supports_realtime_clock=True, + supports_reset=True, + deterministic_given_inputs=False, + user_input_schema=self._user_input_schema(), + ) + + def prepare_initial_input(self) -> InferenceInput: + if self._session_input is None: + return InferenceInput() + return InferenceInput( + global_conditioning={_SESSION_INPUT_KEY: self._session_input} + ) + + def prepare_step( + self, + *, + request: Any, + user_window: UserInputWindow, + ) -> PreparedStep: + if self._uses_inference_session: + return PreparedStep( + inference_input=self._prepare_inference_session_step( + request=request, + user_window=user_window, + ) + ) + return PreparedStep( + inference_input=self._prepare_segment_step( + request=request, + user_window=user_window, + ) + ) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self._session_input_state_advanced = False + + def close(self) -> None: + return + + def _user_input_schema(self) -> UserInputSchema: + schema = getattr(self._runtime, "input_source_schema", None) + if isinstance(schema, UserInputSchema): + return schema + return WEBRTC_USER_INPUT_SCHEMA + + def _prepare_inference_session_step( + self, + *, + request: Any, + user_window: UserInputWindow, + ) -> InferenceInput: + self._advance_skipped_input_state(user_window) + window_start = user_window.start_s + if request.step_index == 0 and not self._session_input_state_advanced: + window_start = 0.0 + window = TimeWindow(start_s=window_start, end_s=user_window.end_s) + canonical_inputs = self._runtime.input_canonicalizer.canonicalize( + user_window.inputs, + window=window, + source_schema=self._runtime.input_source_schema, + ) + mapping = self._runtime.input_mapping + return mapping.map_step_inputs( + canonical_inputs=canonical_inputs, + inference_input=InferenceInput(), + request=_step_request_from_requirements(request, window=window), + ) + + def _advance_skipped_input_state(self, user_window: UserInputWindow) -> None: + skipped_inputs = user_window.metadata.get(WEBRTC_SKIPPED_INPUTS_METADATA_KEY) + skipped_window = user_window.metadata.get(WEBRTC_SKIPPED_WINDOW_METADATA_KEY) + if not isinstance(skipped_inputs, UserInputs): + return + if not isinstance(skipped_window, tuple) or len(skipped_window) != 2: + return + start_value, end_value = skipped_window + if not isinstance(start_value, int | float) or not isinstance( + end_value, + int | float, + ): + return + start_s = float(start_value) + end_s = float(end_value) + if end_s <= start_s: + return + self._runtime.input_canonicalizer.canonicalize( + skipped_inputs, + window=TimeWindow(start_s=start_s, end_s=end_s), + source_schema=self._runtime.input_source_schema, + ) + self._session_input_state_advanced = True + + @staticmethod + def _prepare_segment_step( + *, + request: Any, + user_window: UserInputWindow, + ) -> InferenceInput: + segments = user_window.metadata.get(SPARSE_KEY_SEGMENTS_METADATA_KEY) + if not isinstance(segments, tuple): + raise RuntimeError("WebRTC user window is missing resampled key segments.") + window = TimeWindow(start_s=user_window.start_s, end_s=user_window.end_s) + return InferenceInput( + step={ + _STEP_REQUEST_KEY: _step_request_from_requirements( + request, + window=window, + ), + _SEGMENTS_KEY: tuple(segments), + _FRAME_TIMES_KEY: tuple(user_window.frame_times), + } + ) + + +class _LegacyWebRTCDemoAdapter: + """Minimal adapter for the shared helper while WebRTC providers migrate.""" + + model_id: str + inference_input_schema = InferenceInputSchema() + canonical_input_schema = CanonicalInputSchema() + + def __init__( + self, + *, + runtime: Any, + identity: str, + session_input: Any = None, + ) -> None: + self._runtime = runtime + self.model_id = identity + self._session_input = session_input + + def supported_input_modes(self) -> tuple[str, ...]: + return ("webrtc",) + + def supported_output_modes(self) -> tuple[str, ...]: + return ("webrtc",) + + def default_input_mapping(self) -> InputMapping | None: + return None + + def validate_config(self, config: Any) -> None: + if config.model_id != self.model_id: + raise ValueError( + f"Expected WebRTC model_id={self.model_id!r}, got {config.model_id!r}." + ) + + def create_runtime(self, config: Any) -> Any: + self.validate_config(config) + return self._runtime + + def prepare_scenario(self, spec: Any) -> PreparedScenario: + del spec + return PreparedScenario(initial_inputs=self._initial_inputs()) + + def create_model_input_provider( + self, + spec: Any, + scenario: PreparedScenario, + ) -> _LegacyWebRTCModelInputProvider: + del spec, scenario + return _LegacyWebRTCModelInputProvider( + runtime=self._runtime, + session_input=self._session_input, + ) + + def _initial_inputs(self) -> InferenceInput: + if self._session_input is None: + return InferenceInput() + return InferenceInput( + global_conditioning={_SESSION_INPUT_KEY: self._session_input} + ) + + +class _ManagedWebRTCSessionEdgeFactory: + """Build shared realtime edges for one negotiated peer connection.""" + + def __init__( + self, + *, + manager: "BaseWebRTCSessionManager[Any, Any]", + managed_session: "ManagedWebRTCSession", + loop: asyncio.AbstractEventLoop, + ) -> None: + self._manager = manager + self._managed_session = managed_session + self._loop = loop + + def create_session_edges( + self, + *, + context: RunContext, + spec: Any, + scenario: PreparedScenario, + provider: ModelInputProvider, + adapter: Any, + ) -> SessionEdges: + del spec, scenario, provider, adapter + input_source = self._managed_session.input_source + transport = self._managed_session.transport + if input_source is None or transport is None: + raise RuntimeError("Managed WebRTC session is missing shared edges.") + bridge = ThreadSafeWebRTCOutputBridge( + loop=self._loop, + video_encoder=self._managed_session.video_encoder, + video_track=self._managed_session.video_track, + on_chunk_delivery=self._on_chunk_delivery, + on_error=self._on_delivery_error, + ) + return SessionEdges( + input_source=input_source, + output_sink=WebRTCOutputSink(bridge=bridge), + cleanup_tasks=context.cleanup_tasks, + metrics=InMemorySessionMetricsRecorder(), + error_policy=WebRTCErrorPolicy(), + transport=transport, + clock=ResamplerRealtimeClock( + resampler=self._managed_session.resampler, + now_fn=self._loop.time, + sleep_fn=asyncio.sleep, + ), + activation=WebRTCActivationPolicy( + input_source=input_source, + transport=transport, + ), + ) + + def _on_chunk_delivery(self, chunk: WebRTCChunkDelivery) -> None: + self._manager._handle_shared_chunk_delivery( + managed_session=self._managed_session, + chunk=chunk, + ) + + def _on_delivery_error(self, exc: BaseException) -> None: + self._manager._handle_shared_delivery_error( + managed_session=self._managed_session, + exc=exc, + ) + if self._manager.fatal_generation_errors: + self._loop.call_soon_threadsafe( + lambda: asyncio.create_task(self._manager.close_active_session()) + ) + + @dataclass(slots=True) class ManagedWebRTCSession: """Per-session state for the single active WebRTC peer connection.""" @@ -152,6 +579,9 @@ class ManagedWebRTCSession: control_channel: Any | None = None generation_task: asyncio.Task[Any] | None = None first_action_received: asyncio.Event = field(default_factory=asyncio.Event) + input_source: WebRTCInputSource | None = None + transport: WebRTCTransportService | None = None + reservation: Any | None = None pending_action_arrivals: deque[float] = field(default_factory=deque) inference_session: Any | None = None """Active ``InferenceSession``; ``None`` means call ``runtime.generate_chunk``.""" @@ -189,6 +619,11 @@ async def close(self) -> None: self.generation_task.cancel() with contextlib.suppress(asyncio.CancelledError): await self.generation_task + if self.generation_task is None or self.generation_task.done(): + reservation = self.reservation + self.reservation = None + if reservation is not None: + reservation.release() self.generation_task = None await self.video_track.close() @@ -234,6 +669,9 @@ def __init__( self._preload_lock = asyncio.Lock() self._session_lock = asyncio.Lock() self._pending_session_input: Any = None + self._shared_runtime_adapter: _LegacyWebRTCRuntimeAdapter | None = None + self._shared_host: RuntimeHost | None = None + self._shared_context: RunContext | None = None @property def pending_session_input(self) -> Any: @@ -334,6 +772,50 @@ def _resolve_video_encoder(self) -> VideoEncoder: encoder = DefaultRTCEncoder(fps=self.fps) return encoder + def _shared_run_context(self, loop: asyncio.AbstractEventLoop) -> RunContext: + if self._shared_context is not None: + return self._shared_context + runtime_adapter = _LegacyWebRTCRuntimeAdapter( + runtime=self._runtime, + loop=loop, + ) + host = RuntimeHost(runtime_adapter) + self._shared_runtime_adapter = runtime_adapter + self._shared_host = host + self._shared_context = RunContext( + host=host, + run_metrics=InMemorySessionMetricsRecorder(), + admission=SingleSessionAdmissionPolicy( + health_check=lambda: host.is_healthy + ), + ) + return self._shared_context + + def _shared_demo_spec(self) -> DemoSpec: + return DemoSpec( + model_id=self.identity, + input_mode="webrtc", + output=WebRTCOutputSpec( + fps=self.fps, + video_width=self.runtime_config.video_width, + video_height=self.runtime_config.video_height, + warmup_chunks=self.runtime_config.warmup_chunks, + warmup_timeout_s=self.runtime_config.warmup_timeout_s, + client_liveness_timeout_s=self.client_liveness_timeout_s, + ), + ) + + async def _reset_runtime_for_session( + self, + *, + context: RunContext, + session_input: Any, + ) -> None: + reset = getattr(context.host.runtime, "reset_for_new_session", None) + if not callable(reset): + raise RuntimeError("WebRTC runtime adapter cannot reset sessions.") + await context.host.call_async(reset, session_input) + def _prefer_h264_video_codec(self, *, transceiver: Any) -> None: """Constrain the transceiver's codec preferences to H.264 variants. @@ -496,7 +978,7 @@ async def _handle_event_message( @staticmethod def _drives_inference_session(runtime: Any) -> bool: """Return whether ``runtime`` should be driven through ``InferenceSession``.""" - return callable(getattr(runtime, "start_inference_session", None)) + return _runtime_drives_inference_session(runtime) def _record_user_event( self, @@ -808,46 +1290,61 @@ async def _create_answer_with_runtime_ready_locked( if not self._runtime_ready: raise RuntimeError("Runtime is not initialized.") - await self._runtime.reset_for_new_session(session_input=session_input) - - peer_connection = RTCPeerConnection(rtc_configuration) - # Bounded queue sized to one *steady-state* chunk so the producer - # is throttled to the consumer's drain rate. AR step 0 emits fewer - # frames than steady state; sizing to it would force a per-chunk - # stall, so we size to the steady-state count. - num_frames = self._runtime_steady_output_num_frames(self._runtime) - video_encoder = self._resolve_video_encoder() - video_track = video_encoder.create_track(maxsize=num_frames) - # Use ``addTransceiver`` (not ``addTrack``) so we can constrain the - # SDP m-line's codec list via ``setCodecPreferences`` when the - # encoder emits pre-encoded H.264 packets. - video_transceiver = peer_connection.addTransceiver( - video_track, - direction="sendonly", - ) - if video_encoder.prefers_codec == "h264": - self._prefer_h264_video_codec(transceiver=video_transceiver) - # Start the resampler's virtual clock at 0; the real anchor is set - # in the ``on_datachannel`` handler so chunk 0's window starts when - # input can actually arrive. - resampler = self._make_resampler_at_fps( - start_v=0.0, - fps=self._runtime_input_fps(self._runtime), - ) loop = asyncio.get_running_loop() - managed_session = ManagedWebRTCSession( - runtime=self._runtime, - video_track=video_track, - video_encoder=video_encoder, - peer_connection=peer_connection, - resampler=resampler, - last_client_message_at=loop.time(), - ) - session_runtime: Any = self._runtime - if self._drives_inference_session(session_runtime): - managed_session.inference_session = ( - await session_runtime.start_inference_session() + context = self._shared_run_context(loop) + reservation = context.admission.try_reserve() + if reservation is None: + raise SessionBusyError(self.busy_message) + try: + await self._reset_runtime_for_session( + context=context, + session_input=session_input, ) + except Exception: + reservation.release() + raise + + try: + peer_connection = RTCPeerConnection(rtc_configuration) + # Bounded queue sized to one *steady-state* chunk so the producer + # is throttled to the consumer's drain rate. AR step 0 emits fewer + # frames than steady state; sizing to it would force a per-chunk + # stall, so we size to the steady-state count. + num_frames = self._runtime_steady_output_num_frames(self._runtime) + video_encoder = self._resolve_video_encoder() + video_track = video_encoder.create_track(maxsize=num_frames) + # Use ``addTransceiver`` (not ``addTrack``) so we can constrain the + # SDP m-line's codec list via ``setCodecPreferences`` when the + # encoder emits pre-encoded H.264 packets. + video_transceiver = peer_connection.addTransceiver( + video_track, + direction="sendonly", + ) + if video_encoder.prefers_codec == "h264": + self._prefer_h264_video_codec(transceiver=video_transceiver) + # Start the resampler's virtual clock at 0; the real anchor is set + # in the ``on_datachannel`` handler so chunk 0's window starts when + # input can actually arrive. + resampler = self._make_resampler_at_fps( + start_v=0.0, + fps=self._runtime_input_fps(self._runtime), + ) + input_source = WebRTCInputSource(resampler=resampler) + transport = WebRTCTransportService(loop=loop) + managed_session = ManagedWebRTCSession( + runtime=self._runtime, + video_track=video_track, + video_encoder=video_encoder, + peer_connection=peer_connection, + resampler=resampler, + input_source=input_source, + transport=transport, + reservation=reservation, + last_client_message_at=loop.time(), + ) + except Exception: + reservation.release() + raise self._active_session = managed_session if enable_liveness_watchdog: managed_session.liveness_task = asyncio.create_task( @@ -858,10 +1355,13 @@ async def _create_answer_with_runtime_ready_locked( def on_datachannel(channel: Any) -> None: managed_session.control_channel = channel # Re-anchor the resampler at channel open. The real - # virtual-clock anchor happens in ``_generation_worker`` once - # the first keyboard event arrives. + # virtual-clock anchor happens in ``WebRTCActivationPolicy`` once + # the first browser event activates the shared realtime driver. channel_open_v = asyncio.get_running_loop().time() - managed_session.resampler.reset(start_v=channel_open_v) + if managed_session.input_source is not None: + managed_session.input_source.reset(start_v=channel_open_v) + else: + managed_session.resampler.reset(start_v=channel_open_v) @channel.on("message") def on_message(message: Any) -> None: @@ -872,15 +1372,21 @@ def on_message(message: Any) -> None: ) ) - # Spawn the generation worker once the channel is wired up so + # Spawn the shared realtime session once the channel is wired up so # ``chunk_done`` notifications have a channel to land on. managed_session.generation_task = asyncio.create_task( - self._generation_worker(managed_session=managed_session) + self._run_realtime_driver_session( + managed_session=managed_session, + context=context, + session_input=session_input, + ) ) @channel.on("close") def on_close() -> None: logger.info("Control data channel closed; closing active session.") + if managed_session.transport is not None: + managed_session.transport.disconnect("data channel closed") asyncio.create_task(self.close_active_session()) @peer_connection.on("connectionstatechange") @@ -993,6 +1499,13 @@ async def _client_liveness_watchdog( async def shutdown(self) -> None: await self.close_active_session() + if self._shared_context is not None: + await self._shared_context.close_async() + if self._shared_host is not None: + self._shared_host.close() + self._shared_context = None + self._shared_host = None + self._shared_runtime_adapter = None await self._runtime.close() self._runtime_ready = False self._warmup_complete = False @@ -1013,6 +1526,10 @@ async def _handle_datachannel_message( if channel is None or managed_session.closed: return managed_session.last_client_message_at = asyncio.get_running_loop().time() + if managed_session.transport is not None: + managed_session.transport.mark_client_message( + managed_session.last_client_message_at + ) if not isinstance(raw_message, str): self._send_json(channel, make_error_payload("Expected text payload.")) @@ -1029,6 +1546,12 @@ async def _handle_datachannel_message( channel, make_error_payload("Payload must be a JSON object.") ) return + if managed_session.input_source is not None: + await self._handle_shared_datachannel_payload( + managed_session=managed_session, + payload=payload, + ) + return message_type = str(payload.get("type", "")).strip().lower() if message_type == MESSAGE_TYPE_HEARTBEAT: return @@ -1106,6 +1629,185 @@ async def _handle_datachannel_message( # user actually interacts. Idempotent once already set. managed_session.first_action_received.set() + async def _handle_shared_datachannel_payload( + self, + *, + managed_session: ManagedWebRTCSession, + payload: dict[str, Any], + ) -> None: + channel = managed_session.control_channel + input_source = managed_session.input_source + if channel is None or input_source is None: + return + message_type = str(payload.get("type", "")).strip().lower() + if message_type == MESSAGE_TYPE_HEARTBEAT: + return + if message_type == MESSAGE_TYPE_DISCONNECT: + logger.info("Client requested disconnect; closing active session.") + if managed_session.transport is not None: + managed_session.transport.disconnect("client disconnected") + await self.close_active_session() + return + if message_type == MESSAGE_TYPE_EVENT: + handled = self._record_shared_event_payload( + managed_session=managed_session, + payload=payload, + ) + if handled: + managed_session.first_action_received.set() + return + result = input_source.handle_browser_payload( + payload, + timestamp_s=asyncio.get_running_loop().time(), + ) + if result.kind == "error": + self._send_json(channel, make_error_payload(result.error or "Bad input.")) + return + if result.activated: + managed_session.first_action_received.set() + + def _record_shared_event_payload( + self, + *, + managed_session: ManagedWebRTCSession, + payload: dict[str, Any], + ) -> bool: + channel = managed_session.control_channel + input_source = managed_session.input_source + if channel is None or input_source is None: + return False + event_id = str(payload.get("event_id", payload.get("id", ""))).strip() + state = str(payload.get("state", "trigger")).strip().lower() or "trigger" + clear_states = {"clear", "release", "off", "none"} + if not event_id and state not in clear_states: + self._send_json( + channel, + make_error_payload( + ( + "Event payload must include non-empty 'event_id' " + "unless state clears the active event." + ), + ), + ) + return False + clears = state in clear_states + try: + event_payload = self._validate_user_event_payload( + managed_session=managed_session, + event_type="text_event", + payload={ + "event_id": None if clears else event_id, + "state": state, + }, + ) + active_event_id = event_payload.get("event_id") + source_event_id = None if active_event_id is None else str(active_event_id) + input_source.record_user_event( + timestamp_s=asyncio.get_running_loop().time(), + event_type="text_event", + payload=event_payload, + source_event_id=source_event_id, + ) + except Exception as exc: + self._send_json(channel, make_error_payload(str(exc))) + return False + active_event_id = event_payload.get("event_id") + ack_event_id = None if active_event_id is None else str(active_event_id) + self._send_json( + channel, + make_event_ack_payload( + event_id=ack_event_id, + state=str(event_payload.get("state", state)), + result={"active_event_id": ack_event_id}, + ), + ) + return True + + async def _run_realtime_driver_session( + self, + *, + managed_session: ManagedWebRTCSession, + context: RunContext, + session_input: Any, + ) -> None: + adapter = _LegacyWebRTCDemoAdapter( + runtime=self._runtime, + identity=self.identity, + session_input=session_input, + ) + spec = self._shared_demo_spec() + scenario = adapter.prepare_scenario(spec) + run_mode = WebRTCRunMode( + edge_factory=_ManagedWebRTCSessionEdgeFactory( + manager=self, + managed_session=managed_session, + loop=asyncio.get_running_loop(), + ) + ) + try: + result = await run_demo_session_async( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + run_mode=run_mode, + pipeline=StepPipeline(), + reservation=managed_session.reservation, + ) + if result.status == "failed" and result.reason: + channel = managed_session.control_channel + if channel is not None: + self._send_json(channel, make_error_payload(result.reason)) + finally: + managed_session.reservation = None + if self._active_session is managed_session: + await self.close_active_session() + + def _handle_shared_chunk_delivery( + self, + *, + managed_session: ManagedWebRTCSession, + chunk: WebRTCChunkDelivery, + ) -> None: + channel = managed_session.control_channel + if channel is None or managed_session.closed: + return + delivery = chunk.delivery + enqueued_frames = int(getattr(delivery, "num_frames", chunk.frame_count)) + encode_ms = float(getattr(delivery, "encode_ms", 0.0)) + play_ms = chunk.frame_count * 1000.0 / managed_session.video_track.fps + queue_depth = managed_session.video_track.qsize() + self._send_json( + channel, + make_chunk_done_payload( + chunk_index=chunk.step_index, + num_frames=chunk.frame_count, + enqueued_frames=enqueued_frames, + fps=managed_session.video_track.fps, + width=self.runtime_config.video_width, + height=self.runtime_config.video_height, + model=self.identity, + gen_ms=_stat_ms(chunk.metrics, "model_step_s"), + enqueue_ms=encode_ms, + play_ms=play_ms, + queue_depth=queue_depth, + lag_ms=0.0, + control_latency_ms=None, + consumed_actions=0, + extra=chunk.metadata, + ), + ) + + def _handle_shared_delivery_error( + self, + *, + managed_session: ManagedWebRTCSession, + exc: BaseException, + ) -> None: + channel = managed_session.control_channel + if channel is not None: + self._send_json(channel, make_error_payload(str(exc))) + async def _generation_worker( self, *, managed_session: ManagedWebRTCSession ) -> None: diff --git a/flashdreams/flashdreams/serving/webrtc/media.py b/flashdreams/flashdreams/serving/webrtc/media.py index 7e358966f..25fa09437 100644 --- a/flashdreams/flashdreams/serving/webrtc/media.py +++ b/flashdreams/flashdreams/serving/webrtc/media.py @@ -101,6 +101,17 @@ async def enqueue_result(self, result: StepResult) -> int: frames = await asyncio.to_thread(self.prepare_result_frames, result) return await self.enqueue_frames(frames) + async def flush(self) -> None: + """Drop queued frames while keeping the RTP timestamp sequence alive.""" + if self._closed: + return + while True: + try: + self._frames.get_nowait() + except asyncio.QueueEmpty: + break + self._next_deadline_s = None + async def recv(self) -> VideoFrame: if self._closed: raise MediaStreamError @@ -244,6 +255,17 @@ def enqueue_encoded_packet_nowait(self, packet: Packet) -> bool: self._packets.put_nowait(packet) return True + async def flush(self) -> None: + """Drop queued encoded packets while preserving the open media track.""" + if self._closed: + return + while True: + try: + self._packets.get_nowait() + except asyncio.QueueEmpty: + break + self._next_deadline_s = None + async def recv(self) -> Packet: if self._closed: raise MediaStreamError diff --git a/flashdreams/flashdreams/serving/webrtc/services.py b/flashdreams/flashdreams/serving/webrtc/services.py index 3acc863a6..3ce38e4be 100644 --- a/flashdreams/flashdreams/serving/webrtc/services.py +++ b/flashdreams/flashdreams/serving/webrtc/services.py @@ -83,6 +83,8 @@ WebRTCDropPolicy = Literal["none", "drop_newest", "drop_oldest"] _CLEAR_EVENT_STATES = frozenset({"clear", "release", "off", "none"}) +WEBRTC_SKIPPED_INPUTS_METADATA_KEY = "webrtc_skipped_inputs" +WEBRTC_SKIPPED_WINDOW_METADATA_KEY = "webrtc_skipped_window" @dataclass(frozen=True, kw_only=True, slots=True) @@ -129,6 +131,23 @@ def __post_init__(self) -> None: object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) +@dataclass(frozen=True, kw_only=True, slots=True) +class WebRTCChunkDelivery: + """Completed WebRTC chunk delivery plus the model chunk summary.""" + + delivery: object + step_index: int + frame_count: int + generation: int + force_keyframe: bool + metadata: Mapping[str, object] = field(default_factory=dict) + metrics: Mapping[str, float | int] = field(default_factory=dict) + + def __post_init__(self) -> None: + object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) + object.__setattr__(self, "metrics", freeze_mapping(self.metrics)) + + @runtime_checkable class WebRTCOfferAnswerer(Protocol): """Creates an SDP answer after the shared session task has been scheduled.""" @@ -157,6 +176,8 @@ async def run( class WebRTCOutputBridge(Protocol): """Thread-safe bridge from model-worker output writes to WebRTC delivery.""" + def begin_generation(self, generation: int) -> None: ... + def submit_chunk( self, result: StepResult, @@ -456,6 +477,7 @@ async def next_realtime_window( window_end_s = float(self.resampler.next_chunk_start_v) + chunk_duration_s await clock.wait_until_window_end(window_end_s) + pre_catch_up_start_s = float(self.resampler.next_chunk_start_v) catch_up = clock.catch_up( request=request, max_lag_s=self.max_lag_s @@ -466,13 +488,24 @@ async def next_realtime_window( start_s = float(self.resampler.next_chunk_start_v) segments, frame_times = self.resampler.sample_chunk(input_frame_count) end_s = float(self.resampler.next_chunk_start_v) + metadata: dict[str, object] = { + SPARSE_KEY_SEGMENTS_METADATA_KEY: tuple(segments), + } + if start_s > pre_catch_up_start_s: + metadata[WEBRTC_SKIPPED_INPUTS_METADATA_KEY] = UserInputs( + events=self._events_for_window(pre_catch_up_start_s, start_s) + ) + metadata[WEBRTC_SKIPPED_WINDOW_METADATA_KEY] = ( + pre_catch_up_start_s, + start_s, + ) window = RealtimeWindowResult( window=_user_input_window( start_s=start_s, end_s=end_s, frame_times=tuple(frame_times), inputs=UserInputs(events=self._events_for_window(start_s, end_s)), - metadata={SPARSE_KEY_SEGMENTS_METADATA_KEY: tuple(segments)}, + metadata=metadata, ), catch_up=catch_up, ) @@ -576,12 +609,14 @@ def open(self, session_info: SessionInfo) -> None: self._closed = False self._generation = 0 self._force_keyframe = True + self._bridge.begin_generation(0) def begin_generation(self, generation: int) -> None: if generation < 0: raise ValueError("generation must be >= 0.") self._generation = generation self._force_keyframe = True + self._bridge.begin_generation(generation) def write(self, result: StepResult) -> OutputDecision: if not self._opened or self._closed: @@ -600,7 +635,7 @@ def write(self, result: StepResult) -> OutputDecision: metadata=decision.metadata, ) - def close(self) -> Sequence[object]: + def close(self) -> Sequence[Any]: if self._bridge_closed: return () self._closed = True @@ -622,6 +657,7 @@ def __init__( max_pending_chunks: int = 2, close_track: bool = True, on_delivery: Callable[[object], None] | None = None, + on_chunk_delivery: Callable[[WebRTCChunkDelivery], None] | None = None, on_error: Callable[[BaseException], None] | None = None, ) -> None: if max_pending_chunks <= 0: @@ -632,16 +668,34 @@ def __init__( self._max_pending_chunks = max_pending_chunks self._close_track = close_track self._on_delivery = on_delivery + self._on_chunk_delivery = on_chunk_delivery self._on_error = on_error - self._pending: set[Future[object]] = set() + self._pending: dict[Future[WebRTCChunkDelivery], int] = {} self._lock = threading.Lock() self._closed = False + self._generation = 0 @property def pending_count(self) -> int: with self._lock: return len(self._pending) + def begin_generation(self, generation: int) -> None: + if generation < 0: + raise ValueError("generation must be >= 0.") + with self._lock: + if self._closed or generation <= self._generation: + return + self._generation = generation + stale = tuple( + future + for future, future_generation in self._pending.items() + if future_generation < generation + ) + for future in stale: + future.cancel() + self._schedule_track_flush() + def submit_chunk( self, result: StepResult, @@ -665,6 +719,13 @@ def submit_chunk( drop_policy="drop_newest", metadata={"reason": "closed"}, ) + if generation < self._generation: + return WebRTCOutputBridgeDecision( + accepted=False, + dropped=True, + drop_policy="drop_newest", + metadata={"reason": "stale generation"}, + ) if len(self._pending) >= self._max_pending_chunks: return WebRTCOutputBridgeDecision( accepted=False, @@ -673,6 +734,15 @@ def submit_chunk( metadata={"reason": "pending queue full"}, ) payload = prepare(result, self._video_track) + chunk = WebRTCChunkDelivery( + delivery=None, + step_index=result.step_index, + frame_count=result.frame_count, + generation=generation, + force_keyframe=force_keyframe, + metadata=result.metadata, + metrics=result.metrics, + ) with self._lock: if self._closed: return WebRTCOutputBridgeDecision( @@ -682,6 +752,13 @@ def submit_chunk( drop_policy="drop_newest", metadata={"reason": "closed"}, ) + if generation < self._generation: + return WebRTCOutputBridgeDecision( + accepted=False, + dropped=True, + drop_policy="drop_newest", + metadata={"reason": "stale generation"}, + ) if len(self._pending) >= self._max_pending_chunks: return WebRTCOutputBridgeDecision( accepted=False, @@ -692,12 +769,13 @@ def submit_chunk( future = asyncio.run_coroutine_threadsafe( self._deliver( payload, + chunk=chunk, generation=generation, force_keyframe=force_keyframe, ), self._loop, ) - self._pending.add(future) + self._pending[future] = generation future.add_done_callback(self._on_done) return WebRTCOutputBridgeDecision( @@ -720,19 +798,39 @@ async def _deliver( self, payload: object, *, + chunk: WebRTCChunkDelivery, generation: int, force_keyframe: bool, - ) -> object: - del generation - return await self._video_encoder.deliver_prepared_chunk( + ) -> WebRTCChunkDelivery: + with self._lock: + if self._closed or generation < self._generation: + raise asyncio.CancelledError + delivery = await self._video_encoder.deliver_prepared_chunk( payload, self._video_track, force_keyframe=force_keyframe, ) + with self._lock: + if self._closed or generation < self._generation: + stale_after_delivery = True + else: + stale_after_delivery = False + if stale_after_delivery: + self._schedule_track_flush() + raise asyncio.CancelledError + return WebRTCChunkDelivery( + delivery=delivery, + step_index=chunk.step_index, + frame_count=chunk.frame_count, + generation=chunk.generation, + force_keyframe=chunk.force_keyframe, + metadata=chunk.metadata, + metrics=chunk.metrics, + ) - def _on_done(self, future: Future[object]) -> None: + def _on_done(self, future: Future[WebRTCChunkDelivery]) -> None: with self._lock: - self._pending.discard(future) + self._pending.pop(future, None) if future.cancelled(): return try: @@ -742,7 +840,9 @@ def _on_done(self, future: Future[object]) -> None: self._on_error(exc) return if self._on_delivery is not None: - self._on_delivery(result) + self._on_delivery(result.delivery) + if self._on_chunk_delivery is not None: + self._on_chunk_delivery(result) def _track_backpressure_s(self) -> float: qsize = getattr(self._video_track, "qsize", None) @@ -774,6 +874,18 @@ def _schedule_track_close(self) -> None: if self._on_error is not None: self._on_error(exc) + def _schedule_track_flush(self) -> None: + flush = getattr(self._video_track, "flush", None) + if not callable(flush): + return + try: + result = flush() + if inspect.isawaitable(result): + asyncio.run_coroutine_threadsafe(result, self._loop) + except BaseException as exc: + if self._on_error is not None: + self._on_error(exc) + class WebRTCRunMode: """Shared realtime run mode that delegates peer-specific edges to WebRTC.""" @@ -1031,6 +1143,8 @@ def clear(self) -> None: "BlockingPreparationService", "ThreadSafeWebRTCOutputBridge", "WEBRTC_USER_INPUT_SCHEMA", + "WEBRTC_SKIPPED_INPUTS_METADATA_KEY", + "WEBRTC_SKIPPED_WINDOW_METADATA_KEY", "WebRTCActivationPolicy", "WebRTCInputSource", "WebRTCMessageResult", @@ -1038,6 +1152,7 @@ def clear(self) -> None: "WebRTCOfferRequest", "WebRTCOutputBridge", "WebRTCOutputBridgeDecision", + "WebRTCChunkDelivery", "WebRTCOutputSink", "WebRTCRunMode", "WebRTCSessionEdgeFactory", diff --git a/flashdreams/tests/test_webrtc_manager.py b/flashdreams/tests/test_webrtc_manager.py index 8e4839e5e..f392a910a 100644 --- a/flashdreams/tests/test_webrtc_manager.py +++ b/flashdreams/tests/test_webrtc_manager.py @@ -11,7 +11,14 @@ import pytest import torch -from flashdreams.runtime import StepRequest, StepResult +from flashdreams.runtime import ( + InferenceInput, + StepRequest, + StepRequirements, + StepResult, + UserInputEvent, + UserInputs, +) from flashdreams.serving.webrtc import manager as manager_module from flashdreams.serving.webrtc.controls import WSAD_SUPPORTED_KEYS from flashdreams.serving.webrtc.encoders import ChunkDeliveryResult @@ -20,6 +27,12 @@ ManagedWebRTCSession, ) from flashdreams.serving.webrtc.server import SessionBusyError +from flashdreams.serving.webrtc.services import ( + WEBRTC_SKIPPED_INPUTS_METADATA_KEY, + WEBRTC_SKIPPED_WINDOW_METADATA_KEY, + WebRTCInputSource, + WebRTCTransportService, +) pytestmark = pytest.mark.ci_cpu @@ -83,6 +96,29 @@ async def deliver_chunk( encode_ms=0.1, ) + def prepare_chunk_payload( + self, + result: StepResult, + track: Any, + ) -> StepResult: + del track + return result + + async def deliver_prepared_chunk( + self, + payload: object, + track: Any, + *, + force_keyframe: bool = False, + ) -> ChunkDeliveryResult: + if not isinstance(payload, StepResult): + raise TypeError("fake payload must be StepResult") + return await self.deliver_chunk( + payload, + track, + force_keyframe=force_keyframe, + ) + def close(self) -> None: return @@ -122,6 +158,28 @@ def on_edge(self, *, arrival_t: float, event: str, key: str) -> None: self.edges.append((arrival_t, event, key)) +class _SharedResampler: + def __init__(self, *, start_v: float = 0.0, dt: float = 0.001) -> None: + self.next_chunk_start_v = start_v + self.dt = dt + self.edges: list[tuple[float, str, str]] = [] + + def reset(self, *, start_v: float) -> None: + self.next_chunk_start_v = start_v + self.edges.clear() + + def on_edge(self, *, arrival_t: float, event: str, key: str) -> None: + self.edges.append((arrival_t, event, key)) + + def sample_chunk( + self, num_frames: int + ) -> tuple[list[tuple[float, float, frozenset[str]]], list[float]]: + start = self.next_chunk_start_v + end = start + num_frames * self.dt + self.next_chunk_start_v = end + return [(start, end, frozenset({"w"}))], [end] + + class _CountingVideoTrack(_FakeVideoTrack): async def enqueue_result(self, result: StepResult) -> int: return result.frame_count @@ -423,6 +481,81 @@ def test_catch_up_input_clock_snaps_legacy_path_without_canonicalizer() -> None: assert managed.resampler.next_chunk_start_v == pytest.approx(2.0) +def test_legacy_provider_advances_skipped_webrtc_input_state() -> None: + class _RecordingCanonicalizer: + def __init__(self) -> None: + self.windows: list[tuple[float, float]] = [] + self.event_batches: list[list[str]] = [] + + def canonicalize( + self, + user_inputs: UserInputs, + *, + window: Any, + source_schema: Any, + ) -> object: + del source_schema + self.windows.append((window.start_s, window.end_s)) + self.event_batches.append( + [event.event_type for event in user_inputs.events] + ) + return object() + + class _RecordingMapping: + def map_step_inputs( + self, + *, + canonical_inputs: object, + inference_input: InferenceInput, + request: StepRequest, + ) -> InferenceInput: + del canonical_inputs, inference_input + return InferenceInput(step={"mapped_step": request.step_index}) + + runtime = SimpleNamespace( + start_inference_session=lambda: object(), + input_canonicalizer=_RecordingCanonicalizer(), + input_source_schema=object(), + input_mapping=_RecordingMapping(), + ) + provider = manager_module._LegacyWebRTCModelInputProvider(runtime=runtime) + skipped_inputs = UserInputs( + events=( + UserInputEvent( + timestamp_s=0.5, + event_type="key_down", + payload={"key": "w"}, + ), + ) + ) + current_inputs = UserInputs( + events=( + UserInputEvent( + timestamp_s=2.5, + event_type="key_up", + payload={"key": "w"}, + ), + ) + ) + + prepared = provider.prepare_step( + request=StepRequirements(step_index=0, input_frame_count=1), + user_window=manager_module.UserInputWindow( + start_s=2.0, + end_s=3.0, + inputs=current_inputs, + metadata={ + WEBRTC_SKIPPED_INPUTS_METADATA_KEY: skipped_inputs, + WEBRTC_SKIPPED_WINDOW_METADATA_KEY: (0.0, 2.0), + }, + ), + ) + + assert prepared.inference_input == InferenceInput(step={"mapped_step": 0}) + assert runtime.input_canonicalizer.windows == [(0.0, 2.0), (2.0, 3.0)] + assert runtime.input_canonicalizer.event_batches == [["key_down"], ["key_up"]] + + @pytest.mark.asyncio async def test_action_keydown_reports_error_when_user_event_queue_full( monkeypatch: pytest.MonkeyPatch, @@ -841,6 +974,109 @@ class _FrequentLogManager(_BaseTestManager): assert perf_logs[0][1][-2:] == (13, 512) +@pytest.mark.asyncio +async def test_realtime_driver_session_uses_shared_step_pipeline( + monkeypatch: pytest.MonkeyPatch, +) -> None: + pipeline_calls = 0 + original_pipeline = manager_module.StepPipeline + + class _RecordingPipeline(original_pipeline): + def execute_step( + self, + *, + request: StepRequirements, + user_window: Any, + provider: Any, + session: Any, + output: Any, + metrics: Any, + ) -> Any: + nonlocal pipeline_calls + pipeline_calls += 1 + return original_pipeline.execute_step( + self, + request=request, + user_window=user_window, + provider=provider, + session=session, + output=output, + metrics=metrics, + ) + + class _SharedRuntime: + def __init__(self) -> None: + self.step_requests = 0 + self.step_calls: list[tuple[int, list[Any], list[float]]] = [] + + async def reset_for_new_session(self, session_input: Any = None) -> None: + del session_input + + def next_step_request(self) -> StepRequest | None: + if self.step_requests > 0: + return None + self.step_requests += 1 + return _step_request(step_index=0, input_frame_count=1) + + async def step( + self, + *, + request: StepRequest, + segments: list[Any], + frame_times: list[float], + ) -> StepResult: + self.step_calls.append((request.step_index, segments, frame_times)) + return StepResult(step_index=request.step_index, output="ok", frame_count=1) + + def peek_input_fps(self) -> float: + return 30.0 + + def peek_steady_output_num_frames(self) -> int: + return 1 + + monkeypatch.setattr(manager_module, "StepPipeline", _RecordingPipeline) + runtime = _SharedRuntime() + manager = _make_manager(_BaseTestManager, runtime) + context = manager._shared_run_context(asyncio.get_running_loop()) + reservation = context.admission.try_reserve() + assert reservation is not None + resampler = _SharedResampler(start_v=asyncio.get_running_loop().time()) + input_source = WebRTCInputSource(resampler=resampler) + input_source.handle_browser_payload( + {"type": "action", "action": {"event": "step"}}, + timestamp_s=asyncio.get_running_loop().time(), + ) + managed, video_track, peer, channel = _managed_session(runtime) + managed.resampler = resampler # ty:ignore[invalid-assignment] + managed.input_source = input_source + managed.transport = WebRTCTransportService(loop=asyncio.get_running_loop()) + managed.reservation = reservation + manager._active_session = managed + + managed.generation_task = asyncio.create_task( + manager._run_realtime_driver_session( + managed_session=managed, + context=context, + session_input=None, + ) + ) + await asyncio.wait_for(managed.generation_task, timeout=5.0) + + assert pipeline_calls == 1 + assert runtime.step_calls + assert runtime.step_calls[0][0] == 0 + assert not manager.has_active_session() + assert video_track.closed + assert peer.closed + chunk_done = [ + json.loads(message) + for message in channel.messages + if json.loads(message).get("type") == "chunk_done" + ] + assert len(chunk_done) == 1 + assert chunk_done[0]["model"] == "fake-model" + + @pytest.mark.asyncio async def test_create_answer_raises_busy_with_subclass_message() -> None: manager = _make_manager( diff --git a/flashdreams/tests/test_webrtc_services.py b/flashdreams/tests/test_webrtc_services.py index 25a1a1ba2..bda6eee4c 100644 --- a/flashdreams/tests/test_webrtc_services.py +++ b/flashdreams/tests/test_webrtc_services.py @@ -261,6 +261,45 @@ async def test_webrtc_output_bridge_drops_full_queue_before_payload_prepare() -> sink.close() +@pytest.mark.asyncio +async def test_webrtc_output_bridge_generation_reset_cancels_stale_delivery() -> None: + loop = asyncio.get_running_loop() + encoder = _BlockingEncoder() + track = _FakeVideoTrack() + deliveries: list[object] = [] + chunk_deliveries: list[int] = [] + bridge = ThreadSafeWebRTCOutputBridge( + loop=loop, + video_encoder=encoder, + video_track=track, + on_delivery=deliveries.append, + on_chunk_delivery=lambda chunk: chunk_deliveries.append(chunk.step_index), + ) + sink = WebRTCOutputSink(bridge=bridge) + sink.open(SessionInfo()) + + first = sink.write(StepResult(step_index=0, frame_count=1)) + await asyncio.wait_for(encoder.started.wait(), timeout=1.0) + sink.begin_generation(1) + for _ in range(10): + if track.flush_count: + break + await asyncio.sleep(0) + second = sink.write(StepResult(step_index=1, frame_count=1)) + + assert not first.dropped + assert not second.dropped + assert track.flush_count == 1 + + encoder.release.set() + await asyncio.wait_for(encoder.done.wait(), timeout=1.0) + await asyncio.sleep(0) + + assert deliveries == ["delivered"] + assert chunk_deliveries == [1] + sink.close() + + @pytest.mark.asyncio async def test_disconnect_closes_transport_and_releases_reservation_once() -> None: spec = _webrtc_spec() @@ -702,5 +741,11 @@ async def deliver_chunk( class _FakeVideoTrack: fps = 30 + def __init__(self) -> None: + self.flush_count = 0 + def qsize(self) -> int: return 0 + + async def flush(self) -> None: + self.flush_count += 1 From b137f750a310a25e6ca718bbb53dda132f48981b Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 20:21:25 +0000 Subject: [PATCH 26/51] runtime: shield async invariant cleanup --- .../flashdreams/runtime/demo/drivers.py | 25 +++--- .../tests/test_demo_runtime_vertical_slice.py | 83 +++++++++++++++++++ 2 files changed, 95 insertions(+), 13 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 7ec4e2e27..97d3cb48b 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -513,20 +513,19 @@ async def run_demo_session_async( return result except DriverInvariantError as exc: _record_run_session_error(context, exc) - if provider is not None and not driver_started: - await _close_provider_async( - context=context, - provider=provider, - session_edges=session_edges, - ) - if session_edges is not None and ( + should_record_session = session_edges is not None and ( driver_started or not session_edges.is_closed - ): - result = session_edges.close_result( - status="failed", - reason=str(exc), - error=exc, - ) + ) + result = await _close_partial_session_async( + context=context, + provider=provider, + session_edges=session_edges, + status="failed", + reason=str(exc), + error=exc, + close_provider=not driver_started, + ) + if should_record_session: context.run_metrics.record_session(result) raise except Exception as exc: diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index bbe0a4abd..ff578e2e1 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -3,6 +3,8 @@ from __future__ import annotations +import asyncio +import threading from collections.abc import Callable, Sequence from typing import Any, Literal, cast @@ -363,6 +365,66 @@ async def test_run_demo_session_async_keeps_failure_when_run_cleanup_metrics_fai assert runtime.start_session_inputs == [] +@pytest.mark.asyncio +async def test_run_demo_session_async_invariant_cancellation_finalizes_edges() -> None: + close_entered = threading.Event() + release_close = threading.Event() + runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) + run_metrics = InMemorySessionMetricsRecorder() + context = _run_context(runtime, run_metrics=run_metrics) + provider = _BlockingCloseVideoModelInputProvider( + close_entered=close_entered, + release_close=release_close, + ) + output = _RecordingOutputSink() + transport = _RecordingTransport() + session_metrics = InMemorySessionMetricsRecorder() + select_error = DriverInvariantError("select driver invariant") + task = asyncio.create_task( + run_demo_session_async( + context=context, + spec=_spec(), + scenario=_scenario(), + adapter=_FakeDemoAdapter(provider=provider), + run_mode=_FakeRunMode( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=output, + metrics=session_metrics, + transport=transport, + select_error=select_error, + ), + pipeline=StepPipeline(), + ) + ) + + try: + assert await asyncio.to_thread(close_entered.wait, 2.0) + task.cancel() + await asyncio.sleep(0) + release_close.set() + with pytest.raises( + DriverInvariantError, match="select driver invariant" + ) as raised: + await task + finally: + release_close.set() + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + context.host.close() + + assert raised.value is select_error + assert provider.close_count == 1 + assert output.close_count == 1 + assert transport.close_count == 1 + assert session_metrics.closed + assert len(run_metrics.sessions) == 1 + recorded = cast(RunResult, run_metrics.sessions[0]) + assert recorded.status == "failed" + assert recorded.error is select_error + assert runtime.start_session_inputs == [] + + def test_setup_failure_can_return_skipped_but_not_completed() -> None: skipped = BatchSessionDriver().run_one_session( host=RuntimeHost(_FakeVideoRuntime(session=_FakeVideoSession(num_steps=1))), @@ -783,6 +845,23 @@ def close(self) -> None: raise self.fail_close +class _BlockingCloseVideoModelInputProvider(_FakeVideoModelInputProvider): + def __init__( + self, + *, + close_entered: threading.Event, + release_close: threading.Event, + ) -> None: + super().__init__() + self.close_entered = close_entered + self.release_close = release_close + + def close(self) -> None: + self.close_count += 1 + self.close_entered.set() + assert self.release_close.wait(timeout=2.0) + + class _FakeBatchInputSource: is_finite = True is_deterministic = True @@ -1061,6 +1140,7 @@ def __init__( transport: _RecordingTransport | None = None, error_policy: _SetupPolicy | None = None, validate_error: Exception | None = None, + select_error: Exception | None = None, ) -> None: self.input_source = input_source self.output_sink = output_sink or _RecordingOutputSink() @@ -1069,6 +1149,7 @@ def __init__( self.transport = transport self.error_policy = error_policy self.validate_error = validate_error + self.select_error = select_error self.capabilities = RunModeCapabilities( requires_finite_input=True, supports_artifacts=True, @@ -1137,4 +1218,6 @@ def create_session_edges( ) def select_driver(self) -> BatchSessionDriver: + if self.select_error is not None: + raise self.select_error return BatchSessionDriver() From f021b99395e6cd55c8faa973af144dc2e35a066a Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 20:40:56 +0000 Subject: [PATCH 27/51] runtime: attempt provider cleanup after session timeout --- flashdreams/flashdreams/runtime/demo/drivers.py | 5 ----- .../tests/test_demo_runtime_realtime_driver.py | 14 ++++++++++---- 2 files changed, 10 insertions(+), 9 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 97d3cb48b..48842d2c5 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -656,11 +656,6 @@ async def cleanup() -> RunResult: ) if not session_closed: host.mark_unhealthy("model-affine cleanup timed out") - return session_edges.close_result( - status=status, - reason=reason, - error=error, - ) provider_closed = await _close_model_resource_async( host=host, close=provider.close, diff --git a/flashdreams/tests/test_demo_runtime_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py index 960b2baa2..80af11980 100644 --- a/flashdreams/tests/test_demo_runtime_realtime_driver.py +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -321,13 +321,15 @@ async def test_shielded_cleanup_never_raises_and_returns_result_on_close_errors( @pytest.mark.asyncio async def test_shielded_cleanup_timeout_bounds_shutdown() -> None: host = _NeverReturningHost() + session = _FakeRealtimeSession(num_steps=1) + provider = _FakeRealtimeProvider() metrics = InMemorySessionMetricsRecorder() edges = _edges(metrics=metrics) result = await shielded_session_cleanup( host=cast(RuntimeHost, host), - session=None, - provider=_FakeRealtimeProvider(), + session=session, + provider=provider, session_edges=edges, status="cancelled", reason="timeout test", @@ -337,8 +339,9 @@ async def test_shielded_cleanup_timeout_bounds_shutdown() -> None: assert result.status == "cancelled" assert host.unhealthy_reason == "model-affine cleanup timed out" + assert host.close_targets == [session, provider] assert metrics.cleanup_errors == [] - assert len(metrics.orphaned_cleanup_errors) == 1 + assert len(metrics.orphaned_cleanup_errors) == 2 assert metrics.closed @@ -830,6 +833,7 @@ def handle(self, exc: Exception) -> ErrorAction: class _NeverReturningHost: def __init__(self) -> None: self.unhealthy_reason: str | None = None + self.close_targets: list[Any] = [] async def call_async( self, @@ -838,7 +842,9 @@ async def call_async( *args: object, **kwargs: object, ) -> Any: - del func, args, kwargs + del func, kwargs + close = cast(Callable[[], None], args[0]) + self.close_targets.append(getattr(close, "__self__", close)) await asyncio.Event().wait() def mark_unhealthy( From eb57155f88d0b6609a11b6543355500138460c19 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 21:56:25 +0000 Subject: [PATCH 28/51] Reshape OmniDreams replay runtime toward shared contract --- .../omnidreams/omnidreams/demo/adapter.py | 27 +++++-- .../omnidreams/omnidreams/demo/replay.py | 49 +++++++++--- .../omnidreams/omnidreams/demo/runtime.py | 18 +++++ .../omnidreams/tests/test_demo_api.py | 75 ++++++++++++++++++- 4 files changed, 148 insertions(+), 21 deletions(-) create mode 100644 integrations/omnidreams/omnidreams/demo/runtime.py diff --git a/integrations/omnidreams/omnidreams/demo/adapter.py b/integrations/omnidreams/omnidreams/demo/adapter.py index a1ca8e8f7..a6838ed4a 100644 --- a/integrations/omnidreams/omnidreams/demo/adapter.py +++ b/integrations/omnidreams/omnidreams/demo/adapter.py @@ -27,9 +27,9 @@ ) from flashdreams.runtime.interfaces import InferenceRuntime -from .replay import ( - OmnidreamsReplayRuntime, - OmnidreamsReplayRuntimeOptions, +from .runtime import ( + OmnidreamsRuntime, + OmnidreamsRuntimeOptions, PipelineFactory, ) from .spec import ( @@ -38,7 +38,8 @@ resolve_replay_scenario, ) -ReplayRuntimeFactory = Callable[..., InferenceRuntime] +RuntimeFactory = Callable[..., InferenceRuntime] +ReplayRuntimeFactory = RuntimeFactory class OmnidreamsDemoAdapter: @@ -47,10 +48,19 @@ class OmnidreamsDemoAdapter: def __init__( self, *, - replay_runtime_factory: ReplayRuntimeFactory = OmnidreamsReplayRuntime, + runtime_factory: RuntimeFactory | None = None, + replay_runtime_factory: ReplayRuntimeFactory | None = None, pipeline_factory: PipelineFactory | None = None, ) -> None: - self._replay_runtime_factory = replay_runtime_factory + if runtime_factory is not None and replay_runtime_factory is not None: + raise ValueError( + "Specify either runtime_factory or replay_runtime_factory, not both." + ) + self._runtime_factory = ( + runtime_factory + if runtime_factory is not None + else replay_runtime_factory or OmnidreamsRuntime + ) self._pipeline_factory = pipeline_factory self._mapping = IdentityInputMapping() @@ -119,9 +129,9 @@ def validate_config(self, config: InferenceConfig) -> None: def create_runtime(self, config: InferenceConfig) -> InferenceRuntime: self.validate_config(config) - return self._replay_runtime_factory( + return self._runtime_factory( config=config, - options=OmnidreamsReplayRuntimeOptions( + options=OmnidreamsRuntimeOptions( pipeline_config=self._pipeline_config(config), pipeline_factory=self._pipeline_factory, ), @@ -156,4 +166,5 @@ def _default_replay_prompt(self, config: InferenceConfig | None) -> str: __all__ = [ "OmnidreamsDemoAdapter", "ReplayRuntimeFactory", + "RuntimeFactory", ] diff --git a/integrations/omnidreams/omnidreams/demo/replay.py b/integrations/omnidreams/omnidreams/demo/replay.py index 8ccb58650..a299139a1 100644 --- a/integrations/omnidreams/omnidreams/demo/replay.py +++ b/integrations/omnidreams/omnidreams/demo/replay.py @@ -26,7 +26,7 @@ from flashdreams.runtime.config import InferenceConfig from flashdreams.runtime.inputs import InferenceInput from flashdreams.runtime.interfaces import InferenceSession -from flashdreams.runtime.types import StepRequest, StepResult +from flashdreams.runtime.types import StepRequest, StepRequirements, StepResult from .spec import OmnidreamsReplayScenario @@ -34,22 +34,22 @@ @dataclass(frozen=True, kw_only=True, slots=True) -class OmnidreamsReplayRuntimeOptions: - """Construction knobs for the replay runtime.""" +class OmnidreamsRuntimeOptions: + """Construction knobs for the OmniDreams runtime.""" pipeline_config: Any pipeline_factory: PipelineFactory | None = None output_layout: VideoTensorLayout = "bvtchw" -class OmnidreamsReplayRuntime: - """Heavyweight OmniDreams runtime consumed by ``run_inference_session``.""" +class OmnidreamsRuntime: + """Heavyweight OmniDreams runtime consumed by shared demo run modes.""" def __init__( self, *, config: InferenceConfig, - options: OmnidreamsReplayRuntimeOptions, + options: OmnidreamsRuntimeOptions, ) -> None: self.config = config self.options = options @@ -73,7 +73,7 @@ def __init__( def start_session(self, inputs: InferenceInput) -> InferenceSession: scenario = _scenario_from_inputs(inputs) - return OmnidreamsReplaySession( + return OmnidreamsSession( pipeline=self.pipeline, scenario=scenario, device=torch.device(f"cuda:{self.local_rank}") @@ -95,8 +95,8 @@ def close(self) -> None: torch.cuda.empty_cache() -class OmnidreamsReplaySession: - """One MP4 replay rollout over a prepared scenario.""" +class OmnidreamsSession: + """One OmniDreams rollout over a prepared scenario.""" def __init__( self, @@ -129,7 +129,7 @@ def __init__( if dist.is_initialized(): dist.barrier() - def next_step_request(self) -> StepRequest | None: + def next_step_requirements(self) -> StepRequirements | None: if self._closed: return None step_index = self._model_session.step_index @@ -138,7 +138,26 @@ def next_step_request(self) -> StepRequest | None: num_frames = self._model_session.next_num_frames() if self._frame_start + num_frames > self._hdmap_videos.shape[2]: return None - return StepRequest(step_index=step_index) + return StepRequirements( + step_index=step_index, + input_frame_count=num_frames, + ) + + def next_step_request(self) -> StepRequest | None: + requirements = self.next_step_requirements() + if requirements is None: + return None + metadata = dict(requirements.metadata) + metadata["input_frame_count"] = requirements.input_frame_count + if requirements.steady_output_frame_count is not None: + metadata["steady_output_frame_count"] = ( + requirements.steady_output_frame_count + ) + return StepRequest( + step_index=requirements.step_index, + inference_input_schema=requirements.inference_input_schema, + metadata=metadata, + ) def step(self, inputs: InferenceInput) -> StepResult: del inputs @@ -238,7 +257,15 @@ def _is_torchrun_env() -> bool: return "RANK" in os.environ and "WORLD_SIZE" in os.environ +OmnidreamsReplayRuntimeOptions = OmnidreamsRuntimeOptions +OmnidreamsReplayRuntime = OmnidreamsRuntime +OmnidreamsReplaySession = OmnidreamsSession + + __all__ = [ + "OmnidreamsRuntime", + "OmnidreamsRuntimeOptions", + "OmnidreamsSession", "OmnidreamsReplayRuntime", "OmnidreamsReplayRuntimeOptions", "OmnidreamsReplaySession", diff --git a/integrations/omnidreams/omnidreams/demo/runtime.py b/integrations/omnidreams/omnidreams/demo/runtime.py new file mode 100644 index 000000000..eae00ff93 --- /dev/null +++ b/integrations/omnidreams/omnidreams/demo/runtime.py @@ -0,0 +1,18 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""OmniDreams runtime/session contracts for shared demo run modes.""" + +from omnidreams.demo.replay import ( + OmnidreamsRuntime, + OmnidreamsRuntimeOptions, + OmnidreamsSession, + PipelineFactory, +) + +__all__ = [ + "OmnidreamsRuntime", + "OmnidreamsRuntimeOptions", + "OmnidreamsSession", + "PipelineFactory", +] diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index cb8d710d0..3057b3d39 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -25,6 +25,12 @@ from omnidreams.demo.replay import ( OmnidreamsReplayRuntime, OmnidreamsReplayRuntimeOptions, + OmnidreamsReplaySession, +) +from omnidreams.demo.runtime import ( + OmnidreamsRuntime, + OmnidreamsRuntimeOptions, + OmnidreamsSession, ) from omnidreams.demo.webrtc import ( OmnidreamsWebRTCModelRuntime, @@ -38,6 +44,7 @@ OutputArtifact, OutputTarget, StepRequest, + StepRequirements, StepResult, ) from flashdreams.runtime.demo import ( @@ -67,6 +74,55 @@ def test_omnidreams_demo_adapter_declares_replay_modes_only() -> None: assert adapter.supported_output_modes() == ("mp4",) +def test_omnidreams_runtime_keeps_replay_aliases() -> None: + assert OmnidreamsReplayRuntime is OmnidreamsRuntime + assert OmnidreamsReplayRuntimeOptions is OmnidreamsRuntimeOptions + assert OmnidreamsReplaySession is OmnidreamsSession + + +def test_omnidreams_demo_adapter_accepts_shared_runtime_factory() -> None: + runtime = _FactoryRuntime() + pipeline_config = object() + calls: list[dict[str, Any]] = [] + + def pipeline_factory(config_value: Any, device: str) -> Any: + del config_value, device + return object() + + def runtime_factory(**kwargs: Any) -> Any: + calls.append(kwargs) + return runtime + + adapter = OmnidreamsDemoAdapter( + runtime_factory=runtime_factory, + pipeline_factory=pipeline_factory, + ) + config = InferenceConfig( + model_id=OMNIDREAMS_MODEL_ID, + runtime_options={"pipeline_config": pipeline_config}, + ) + + assert adapter.create_runtime(config) is runtime + assert len(calls) == 1 + assert calls[0]["config"] == config + options = calls[0]["options"] + assert isinstance(options, OmnidreamsRuntimeOptions) + assert options.pipeline_config is pipeline_config + assert options.pipeline_factory is pipeline_factory + + +def test_omnidreams_demo_adapter_rejects_ambiguous_runtime_factories() -> None: + def runtime_factory(**kwargs: Any) -> _FactoryRuntime: + del kwargs + return _FactoryRuntime() + + with pytest.raises(ValueError, match="runtime_factory"): + OmnidreamsDemoAdapter( + runtime_factory=runtime_factory, + replay_runtime_factory=runtime_factory, + ) + + def test_omnidreams_demo_does_not_import_legacy_webrtc_package() -> None: demo_dir = Path(demo_package.__file__).parent @@ -237,9 +293,9 @@ def test_omnidreams_replay_runtime_generates_video_step_result( lambda *args, **kwargs: torch.zeros(2, 3, 2, 2), ) - runtime = OmnidreamsReplayRuntime( + runtime = OmnidreamsRuntime( config=InferenceConfig(model_id=OMNIDREAMS_MODEL_ID, device="cpu"), - options=OmnidreamsReplayRuntimeOptions( + options=OmnidreamsRuntimeOptions( pipeline_config=object(), pipeline_factory=lambda pipeline_config, device: pipeline, ), @@ -257,10 +313,16 @@ def test_omnidreams_replay_runtime_generates_video_step_result( session = runtime.start_session( InferenceInput(global_conditioning={"scenario": scenario}) ) + assert isinstance(session, OmnidreamsSession) + requirements = session.next_step_requirements() + assert isinstance(requirements, StepRequirements) + assert requirements.step_index == 0 + assert requirements.input_frame_count == 1 request = session.next_step_request() assert request is not None assert request.step_index == 0 + assert request.metadata["input_frame_count"] == 1 result = session.step(InferenceInput()) assert result.step_index == 0 @@ -561,6 +623,15 @@ def close(self) -> Sequence[OutputArtifact]: return () +class _FactoryRuntime: + def start_session(self, inputs: InferenceInput) -> Any: + del inputs + raise NotImplementedError + + def close(self) -> None: + return None + + class _FakeOmnidreamsPipeline: def __init__(self) -> None: self.initialize_cache_calls: list[dict[str, Any]] = [] From 03de4c5960ceb2d911228328a70519d492d62490 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 22:40:07 +0000 Subject: [PATCH 29/51] Wire OmniDreams precomputed HDMaps through provider --- .../flashdreams/runtime/demo/replay.py | 54 +++- .../omnidreams/omnidreams/demo/__init__.py | 2 + .../omnidreams/omnidreams/demo/adapter.py | 33 ++- .../omnidreams/omnidreams/demo/providers.py | 208 ++++++++++++++ .../omnidreams/omnidreams/demo/replay.py | 156 +++++++---- .../omnidreams/tests/test_demo_api.py | 253 +++++++++++++++++- 6 files changed, 630 insertions(+), 76 deletions(-) create mode 100644 integrations/omnidreams/omnidreams/demo/providers.py diff --git a/flashdreams/flashdreams/runtime/demo/replay.py b/flashdreams/flashdreams/runtime/demo/replay.py index d33d6478a..fa2f507a9 100644 --- a/flashdreams/flashdreams/runtime/demo/replay.py +++ b/flashdreams/flashdreams/runtime/demo/replay.py @@ -36,7 +36,7 @@ from .drivers import BatchSessionDriver, run_demo_session from .host import ModelWarmupPlan, RuntimeHost -from .outputs import OutputSink, build_output_sink, build_output_target +from .outputs import OutputDecision, OutputSink, build_output_sink, build_output_target from .pipeline import StepPipeline from .run_modes import ( Mp4ErrorPolicy, @@ -99,7 +99,7 @@ def run_replay_demo( if spec.config is None: raise RuntimeError("DemoSpec.config was not initialized.") - if runner is not None or output_target_factory is not None: + if runner is not None: return _run_replay_demo_with_compat_runner( spec=spec, adapter=adapter, @@ -110,6 +110,9 @@ def run_replay_demo( runner=runner or _default_inference_session_runner(), ) + if output_target_factory is not None: + output_sink_factory = _output_target_sink_factory(output_target_factory) + return _run_replay_demo_with_run_mode( spec=spec, adapter=adapter, @@ -120,6 +123,47 @@ def run_replay_demo( ) +def _output_target_sink_factory( + output_target_factory: OutputTargetFactory, +) -> OutputSinkFactory: + def create_output_sink(output_spec: OutputSpec) -> "_OutputTargetSink": + return _OutputTargetSink(output_target_factory(output_spec)) + + return create_output_sink + + +class _OutputTargetSink: + produces_artifacts = True + + def __init__(self, output: OutputTarget) -> None: + self._output = output + self._closed = True + self._artifacts: tuple[OutputArtifact, ...] | None = None + + def open(self, session_info: object) -> None: + del session_info + self._output.open() + self._closed = False + self._artifacts = None + + def begin_generation(self, generation: int) -> None: + del generation + + def write(self, result: StepResult) -> OutputDecision: + self._output.write(result) + return OutputDecision() + + def close(self) -> Sequence[OutputArtifact]: + if self._artifacts is not None: + return self._artifacts + if self._closed: + self._artifacts = () + return self._artifacts + self._closed = True + self._artifacts = tuple(self._output.close()) + return self._artifacts + + def _run_replay_demo_with_compat_runner( *, spec: DemoSpec, @@ -322,8 +366,10 @@ def create_model_input_provider( self, spec: DemoSpec, scenario: "PreparedScenario", - ) -> "_ReplayMappingModelInputProvider": - del spec + ) -> object: + create_provider = getattr(self._adapter, "create_model_input_provider", None) + if callable(create_provider): + return create_provider(spec, scenario) return _ReplayMappingModelInputProvider( adapter=self._adapter, scenario=scenario, diff --git a/integrations/omnidreams/omnidreams/demo/__init__.py b/integrations/omnidreams/omnidreams/demo/__init__.py index 6fa3a9b21..4365b1051 100644 --- a/integrations/omnidreams/omnidreams/demo/__init__.py +++ b/integrations/omnidreams/omnidreams/demo/__init__.py @@ -4,6 +4,7 @@ """Experimental OmniDreams demo adapter built on ``flashdreams.runtime.demo``.""" from omnidreams.demo.adapter import OmnidreamsDemoAdapter +from omnidreams.demo.providers import PrecomputedHDMapProvider from omnidreams.demo.spec import ( DEFAULT_OMNIDREAMS_PRESET, OMNIDREAMS_MODEL_ID, @@ -17,4 +18,5 @@ "OmnidreamsDemoAdapter", "OmnidreamsReplayScenario", "OmnidreamsWebRTCScenario", + "PrecomputedHDMapProvider", ] diff --git a/integrations/omnidreams/omnidreams/demo/adapter.py b/integrations/omnidreams/omnidreams/demo/adapter.py index a6838ed4a..201618716 100644 --- a/integrations/omnidreams/omnidreams/demo/adapter.py +++ b/integrations/omnidreams/omnidreams/demo/adapter.py @@ -17,7 +17,6 @@ InferenceInput, InferenceInputSchema, InputCanonicalizer, - InputField, UserInputSchema, ) from flashdreams.runtime.demo import ( @@ -25,8 +24,13 @@ Mp4OutputSpec, PreparedScenario, ) +from flashdreams.runtime.demo.session_inputs import ModelInputProvider from flashdreams.runtime.interfaces import InferenceRuntime +from .providers import ( + PrecomputedHDMapProvider, + precomputed_hdmap_inference_input_schema, +) from .runtime import ( OmnidreamsRuntime, OmnidreamsRuntimeOptions, @@ -70,15 +74,7 @@ def model_id(self) -> str: @property def inference_input_schema(self) -> InferenceInputSchema: - return InferenceInputSchema( - global_conditioning_fields=( - InputField( - name="scenario", - input_modality="omnidreams/replay-scenario", - description="Resolved OmniDreams replay scenario.", - ), - ) - ) + return precomputed_hdmap_inference_input_schema() @property def canonical_input_schema(self) -> CanonicalInputSchema | None: @@ -137,6 +133,23 @@ def create_runtime(self, config: InferenceConfig) -> InferenceRuntime: ), ) + def create_model_input_provider( + self, + spec: DemoSpec, + scenario: PreparedScenario, + ) -> ModelInputProvider: + if spec.input_mode != "replay": + raise ValueError( + "OmniDreams precomputed HDMap provider currently supports only " + f"input_mode='replay', got {spec.input_mode!r}." + ) + if spec.config is None: + raise RuntimeError("DemoSpec.config was not initialized.") + return PrecomputedHDMapProvider( + scenario=scenario, + config=spec.config, + ) + def _preset_id(self, config: InferenceConfig | None) -> str: return ( DEFAULT_OMNIDREAMS_PRESET diff --git a/integrations/omnidreams/omnidreams/demo/providers.py b/integrations/omnidreams/omnidreams/demo/providers.py new file mode 100644 index 000000000..831330a85 --- /dev/null +++ b/integrations/omnidreams/omnidreams/demo/providers.py @@ -0,0 +1,208 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""OmniDreams model-input providers for shared demo run modes.""" + +from __future__ import annotations + +import os + +import torch +import torch.distributed as dist +from loguru import logger +from omnidreams.runner import _load_video + +from flashdreams.infra.runner_io import ( + DEFAULT_RUNNER_INSTALL_HINT, + load_first_frame_tensor, +) +from flashdreams.runtime.config import InferenceConfig +from flashdreams.runtime.inputs import InferenceInput, InferenceInputSchema, InputField +from flashdreams.runtime.demo import ( + PreparedScenario, + PreparedStep, + ProviderCapabilities, + UserInputWindow, +) +from flashdreams.runtime.demo.session_inputs import ControlDecision +from flashdreams.runtime.types import StepRequirements + +from .spec import OmnidreamsReplayScenario + + +class PrecomputedHDMapProvider: + """Prepare fixed OmniDreams HDMap conditioning for replay-style runs.""" + + def __init__( + self, + *, + scenario: PreparedScenario, + config: InferenceConfig, + ) -> None: + self._scenario = _scenario_from_prepared(scenario) + self._device = _device_from_config(config) + self._dtype = torch.bfloat16 + self._frame_start = 0 + self._closed = False + self.capabilities = ProviderCapabilities( + supports_recorded_input=True, + supports_reset=True, + deterministic_given_inputs=True, + user_input_schema=scenario.source_schema, + inference_input_schema=precomputed_hdmap_inference_input_schema(), + ) + self._hdmap_videos: torch.Tensor | None = self._load_hdmaps() + + def prepare_initial_input(self) -> InferenceInput: + self._require_open() + scenario = self._scenario + first_frames = [ + load_first_frame_tensor( + path, + pixel_height=scenario.pixel_height, + pixel_width=scenario.pixel_width, + device=self._device, + dtype=self._dtype, + allow_video=True, + install_hint=DEFAULT_RUNNER_INSTALL_HINT, + ) + for path in scenario.first_frame_paths + ] + return InferenceInput( + global_conditioning={ + "scenario": scenario, + "prompt": [list(scenario.prompts)], + "first_frame": torch.stack(first_frames, dim=0).unsqueeze(0), + }, + metadata={"view_names": tuple(scenario.camera_names)}, + ) + + def prepare_step( + self, + *, + request: StepRequirements, + user_window: UserInputWindow, + ) -> PreparedStep: + del user_window + self._require_open() + hdmap_videos = self._require_hdmaps() + frame_end = self._frame_start + request.input_frame_count + if frame_end > hdmap_videos.shape[2]: + return PreparedStep( + control=ControlDecision( + close_session=True, + reason="OmniDreams precomputed HDMap input exhausted.", + ) + ) + + frame_start = self._frame_start + self._frame_start = frame_end + return PreparedStep( + inference_input=InferenceInput( + step={"hdmap": hdmap_videos[:, :, frame_start:frame_end]}, + metadata={ + "hdmap_frame_start": frame_start, + "hdmap_frame_end": frame_end, + }, + ) + ) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self._require_open() + self._frame_start = 0 + + def close(self) -> None: + if self._closed: + return + self._closed = True + self._hdmap_videos = None + + def _load_hdmaps(self) -> torch.Tensor: + scenario = self._scenario + videos = [ + _load_video( + path, + pixel_height=scenario.pixel_height, + pixel_width=scenario.pixel_width, + device=self._device, + dtype=self._dtype, + ) + for path in scenario.hdmap_video_paths + ] + hdmap_videos = torch.stack(videos, dim=0).unsqueeze(0) + if _is_rank_zero(): + logger.info( + "Loaded OmniDreams demo HDMaps shape={} views={}", + tuple(hdmap_videos.shape), + len(scenario.camera_names), + ) + return hdmap_videos + + def _require_hdmaps(self) -> torch.Tensor: + hdmap_videos = self._hdmap_videos + if hdmap_videos is None: + raise RuntimeError("OmniDreams precomputed HDMap provider is closed.") + return hdmap_videos + + def _require_open(self) -> None: + if self._closed: + raise RuntimeError("OmniDreams precomputed HDMap provider is closed.") + + +def precomputed_hdmap_inference_input_schema() -> InferenceInputSchema: + return InferenceInputSchema( + global_conditioning_fields=( + InputField( + name="prompt", + input_modality="omnidreams/prompt", + description="OmniDreams prompt batch.", + ), + InputField( + name="first_frame", + input_modality="video/frame", + description="Initial OmniDreams conditioning frame tensor.", + ), + InputField( + name="scenario", + required=False, + input_modality="omnidreams/replay-scenario", + description="Resolved OmniDreams replay scenario metadata.", + ), + ), + step_fields=( + InputField( + name="hdmap", + input_modality="omnidreams/hdmap-video", + frequency_consumed="per_step", + description="Per-step HDMap conditioning chunk.", + ), + ), + ) + + +def _scenario_from_prepared(scenario: PreparedScenario) -> OmnidreamsReplayScenario: + value = scenario.initial_inputs.global_conditioning.get("scenario") + if not isinstance(value, OmnidreamsReplayScenario): + raise TypeError( + "OmniDreams precomputed HDMap provider requires " + "initial_inputs.global_conditioning['scenario'] to be an " + "OmnidreamsReplayScenario." + ) + return value + + +def _device_from_config(config: InferenceConfig) -> torch.device: + if dist.is_initialized(): + return torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}") + return torch.device(config.device or "cuda") + + +def _is_rank_zero() -> bool: + return not dist.is_initialized() or dist.get_rank() == 0 + + +__all__ = [ + "PrecomputedHDMapProvider", + "precomputed_hdmap_inference_input_schema", +] diff --git a/integrations/omnidreams/omnidreams/demo/replay.py b/integrations/omnidreams/omnidreams/demo/replay.py index a299139a1..42871a5c6 100644 --- a/integrations/omnidreams/omnidreams/demo/replay.py +++ b/integrations/omnidreams/omnidreams/demo/replay.py @@ -6,7 +6,7 @@ from __future__ import annotations import os -from collections.abc import Callable +from collections.abc import Callable, Sequence from dataclasses import dataclass from typing import Any @@ -14,7 +14,6 @@ import torch.distributed as dist from loguru import logger from omnidreams.model_session import OmnidreamsModelSessionCore -from omnidreams.runner import _load_video from flashdreams.core.distributed import init as init_distributed from flashdreams.infra.postprocess import VideoTensorLayout @@ -76,6 +75,7 @@ def start_session(self, inputs: InferenceInput) -> InferenceSession: return OmnidreamsSession( pipeline=self.pipeline, scenario=scenario, + initial_inputs=inputs, device=torch.device(f"cuda:{self.local_rank}") if dist.is_initialized() else torch.device(self.config.device or "cuda"), @@ -103,18 +103,19 @@ def __init__( *, pipeline: Any, scenario: OmnidreamsReplayScenario, + initial_inputs: InferenceInput, device: torch.device, is_rank_zero: bool, output_layout: VideoTensorLayout, ) -> None: self.pipeline = pipeline self.scenario = scenario + self._initial_inputs = initial_inputs self.device = device self.is_rank_zero = is_rank_zero self.output_layout = output_layout self.dtype = torch.bfloat16 self._closed = False - self._frame_start = 0 self._model_session = OmnidreamsModelSessionCore( pipeline=pipeline, output_stream_factory=lambda: VideoOutputStream( @@ -123,7 +124,6 @@ def __init__( ), ) self._model_session.reset(self._initialize_cache) - self._hdmap_videos = self._load_hdmaps() if self.device.type == "cuda" and torch.cuda.is_available(): torch.cuda.synchronize(device=self.device) if dist.is_initialized(): @@ -136,8 +136,6 @@ def next_step_requirements(self) -> StepRequirements | None: if step_index >= self.scenario.total_blocks: return None num_frames = self._model_session.next_num_frames() - if self._frame_start + num_frames > self._hdmap_videos.shape[2]: - return None return StepRequirements( step_index=step_index, input_frame_count=num_frames, @@ -160,32 +158,31 @@ def next_step_request(self) -> StepRequest | None: ) def step(self, inputs: InferenceInput) -> StepResult: - del inputs if self._closed: raise RuntimeError("OmniDreams replay session is closed.") step_index = self._model_session.step_index num_frames = self._model_session.next_num_frames() - frame_end = self._frame_start + num_frames + hdmap = _hdmap_from_inputs(inputs) + if hdmap.shape[2] != num_frames: + raise ValueError( + "OmniDreams step HDMap frame count mismatch: " + f"expected {num_frames}, got {hdmap.shape[2]}." + ) logger.info( - "OmniDreams demo replay step {} frames=[{}, {})", + "OmniDreams demo replay step {} frames={}", step_index, - self._frame_start, - frame_end, - ) - result = self._model_session.step( - self._hdmap_videos[:, :, self._frame_start : frame_end] + num_frames, ) - self._frame_start = frame_end - return result + return self._model_session.step(hdmap) def reset(self, inputs: InferenceInput | None = None) -> None: if inputs is not None: scenario = _scenario_from_inputs(inputs) if scenario != self.scenario: raise ValueError("OmniDreams replay reset cannot swap scenarios.") + self._initial_inputs = inputs self._model_session.reset(self._initialize_cache) - self._frame_start = 0 def close(self) -> None: self._closed = True @@ -193,51 +190,21 @@ def close(self) -> None: def _initialize_cache(self) -> Any: scenario = self.scenario - first_frames = [ - load_first_frame_tensor( - path, - pixel_height=scenario.pixel_height, - pixel_width=scenario.pixel_width, + cache = self.pipeline.initialize_cache( + text=_prompt_from_inputs(self._initial_inputs, scenario), + image=_first_frame_from_inputs( + self._initial_inputs, + scenario=scenario, device=self.device, dtype=self.dtype, - allow_video=True, - install_hint=DEFAULT_RUNNER_INSTALL_HINT, - ) - for path in scenario.first_frame_paths - ] - first_frames_t = torch.stack(first_frames, dim=0).unsqueeze(0) - cache = self.pipeline.initialize_cache( - text=[list(scenario.prompts)], - image=first_frames_t, - view_names=list(scenario.camera_names), + ), + view_names=_view_names_from_inputs(self._initial_inputs, scenario), ) release = getattr(self.pipeline, "release_oneshot_encoders", None) if callable(release): release() return cache - def _load_hdmaps(self) -> torch.Tensor: - scenario = self.scenario - videos = [ - _load_video( - path, - pixel_height=scenario.pixel_height, - pixel_width=scenario.pixel_width, - device=self.device, - dtype=self.dtype, - ) - for path in scenario.hdmap_video_paths - ] - # [B=1, V, T, C, H, W] - hdmap_videos = torch.stack(videos, dim=0).unsqueeze(0) - if self.is_rank_zero: - logger.info( - "Loaded OmniDreams demo HDMaps shape={} views={}", - tuple(hdmap_videos.shape), - len(scenario.camera_names), - ) - return hdmap_videos - def _default_pipeline_factory(pipeline_config: Any, device: str) -> Any: return pipeline_config.setup().to(device=device).eval() @@ -253,6 +220,87 @@ def _scenario_from_inputs(inputs: InferenceInput) -> OmnidreamsReplayScenario: return scenario +def _prompt_from_inputs( + inputs: InferenceInput, + scenario: OmnidreamsReplayScenario, +) -> list[list[str]]: + prompt = inputs.global_conditioning.get("prompt") + if prompt is None: + return [list(scenario.prompts)] + if isinstance(prompt, str): + return [[prompt]] + if isinstance(prompt, Sequence): + values = list(prompt) + if all(isinstance(value, str) for value in values): + return [[str(value) for value in values]] + batches: list[list[str]] = [] + for batch in values: + if not isinstance(batch, Sequence) or isinstance(batch, str): + raise TypeError( + "OmniDreams initial prompt batches must be string sequences." + ) + batches.append([str(item) for item in batch]) + return batches + raise TypeError( + "OmniDreams initial prompt must be a string or sequence of strings." + ) + + +def _first_frame_from_inputs( + inputs: InferenceInput, + *, + scenario: OmnidreamsReplayScenario, + device: torch.device, + dtype: torch.dtype, +) -> torch.Tensor: + first_frame = inputs.global_conditioning.get("first_frame") + if isinstance(first_frame, torch.Tensor): + return first_frame + if first_frame is not None: + raise TypeError("OmniDreams initial first_frame must be a torch.Tensor.") + first_frames = [ + load_first_frame_tensor( + path, + pixel_height=scenario.pixel_height, + pixel_width=scenario.pixel_width, + device=device, + dtype=dtype, + allow_video=True, + install_hint=DEFAULT_RUNNER_INSTALL_HINT, + ) + for path in scenario.first_frame_paths + ] + return torch.stack(first_frames, dim=0).unsqueeze(0) + + +def _view_names_from_inputs( + inputs: InferenceInput, + scenario: OmnidreamsReplayScenario, +) -> list[str]: + value = inputs.metadata.get("view_names") or inputs.global_conditioning.get( + "view_names" + ) + if value is None: + return list(scenario.camera_names) + if isinstance(value, str): + return [value] + if isinstance(value, Sequence): + return [str(item) for item in value] + raise TypeError("OmniDreams view_names metadata must be a string sequence.") + + +def _hdmap_from_inputs(inputs: InferenceInput) -> torch.Tensor: + hdmap = inputs.step.get("hdmap") + if not isinstance(hdmap, torch.Tensor): + raise TypeError("OmniDreams session step requires step['hdmap'] tensor.") + if hdmap.ndim != 6: + raise ValueError( + "OmniDreams step['hdmap'] must have shape [B, V, T, C, H, W], " + f"got {tuple(hdmap.shape)}." + ) + return hdmap + + def _is_torchrun_env() -> bool: return "RANK" in os.environ and "WORLD_SIZE" in os.environ diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index 3057b3d39..34c2e81b3 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -20,6 +20,7 @@ OmnidreamsDemoAdapter, OmnidreamsReplayScenario, OmnidreamsWebRTCScenario, + PrecomputedHDMapProvider, ) from omnidreams.demo.app import _replay_spec, _webrtc_spec, parse_args from omnidreams.demo.replay import ( @@ -50,6 +51,9 @@ from flashdreams.runtime.demo import ( DemoSpec, Mp4OutputSpec, + OutputDecision, + SessionInfo, + UserInputWindow, WebRTCOutputSpec, ) from flashdreams.runtime.demo.replay import run_replay_demo @@ -72,6 +76,13 @@ def test_omnidreams_demo_adapter_declares_replay_modes_only() -> None: assert adapter.model_id == OMNIDREAMS_MODEL_ID assert adapter.supported_input_modes() == ("replay",) assert adapter.supported_output_modes() == ("mp4",) + assert [ + field.name + for field in adapter.inference_input_schema.global_conditioning_fields + ] == ["prompt", "first_frame", "scenario"] + assert [field.name for field in adapter.inference_input_schema.step_fields] == [ + "hdmap" + ] def test_omnidreams_runtime_keeps_replay_aliases() -> None: @@ -271,26 +282,193 @@ def test_omnidreams_replay_cli_can_disable_example_data(tmp_path: Path) -> None: OmnidreamsDemoAdapter().prepare_scenario(spec) -def test_omnidreams_replay_runtime_generates_video_step_result( +def test_omnidreams_precomputed_hdmap_provider_prepares_inputs( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - import omnidreams.demo.replay as replay_module + import omnidreams.demo.providers as providers_module + + hdmap = tmp_path / "hdmap.mp4" + first_frame = tmp_path / "first.png" + hdmap.write_bytes(b"fake") + first_frame.write_bytes(b"fake") + loaded_hdmap = torch.arange(3 * 3 * 2 * 2).reshape(3, 3, 2, 2) + monkeypatch.setattr( + providers_module, + "load_first_frame_tensor", + lambda *args, **kwargs: torch.ones(1, 3, 2, 2), + ) + monkeypatch.setattr( + providers_module, + "_load_video", + lambda *args, **kwargs: loaded_hdmap, + ) + adapter = OmnidreamsDemoAdapter() + spec = _replay_demo_spec( + tmp_path=tmp_path, + hdmap=hdmap, + first_frame=first_frame, + total_blocks=2, + ) + prepared = adapter.prepare_scenario(spec) + + provider = adapter.create_model_input_provider(spec, prepared) + + assert isinstance(provider, PrecomputedHDMapProvider) + initial = provider.prepare_initial_input() + scenario = initial.global_conditioning["scenario"] + assert isinstance(scenario, OmnidreamsReplayScenario) + assert initial.global_conditioning["prompt"] == [["drive"]] + assert initial.global_conditioning["first_frame"].shape == (1, 1, 1, 3, 2, 2) + assert initial.metadata["view_names"] == ("camera_front_wide_120fov",) + + step = provider.prepare_step( + request=StepRequirements(step_index=0, input_frame_count=2), + user_window=UserInputWindow(start_s=0.0, end_s=1.0), + ) + + assert step.inference_input is not None + hdmap_chunk = step.inference_input.step["hdmap"] + assert isinstance(hdmap_chunk, torch.Tensor) + assert hdmap_chunk.shape == (1, 1, 2, 3, 2, 2) + torch.testing.assert_close(hdmap_chunk[0, 0], loaded_hdmap[:2]) + + exhausted = provider.prepare_step( + request=StepRequirements(step_index=1, input_frame_count=2), + user_window=UserInputWindow(start_s=1.0, end_s=2.0), + ) + + assert exhausted.inference_input is None + assert exhausted.control.close_session is True + provider.reset() + reset_step = provider.prepare_step( + request=StepRequirements(step_index=0, input_frame_count=1), + user_window=UserInputWindow(start_s=0.0, end_s=1.0), + ) + assert reset_step.inference_input is not None + torch.testing.assert_close( + reset_step.inference_input.step["hdmap"][0, 0], + loaded_hdmap[:1], + ) + provider.close() + + +def test_omnidreams_replay_run_mode_uses_precomputed_provider( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + import omnidreams.demo.providers as providers_module hdmap = tmp_path / "hdmap.mp4" first_frame = tmp_path / "first.png" hdmap.write_bytes(b"fake") first_frame.write_bytes(b"fake") + loaded_hdmap = torch.arange(2 * 3 * 2 * 2).reshape(2, 3, 2, 2) pipeline = _FakeOmnidreamsPipeline() + sink = _RecordingOutputSink() monkeypatch.setattr( - replay_module, + providers_module, "load_first_frame_tensor", lambda *args, **kwargs: torch.zeros(1, 3, 2, 2), ) monkeypatch.setattr( - replay_module, + providers_module, + "_load_video", + lambda *args, **kwargs: loaded_hdmap, + ) + adapter = OmnidreamsDemoAdapter( + pipeline_factory=lambda pipeline_config, device: pipeline, + ) + spec = _replay_demo_spec( + tmp_path=tmp_path, + hdmap=hdmap, + first_frame=first_frame, + total_blocks=2, + ) + + result = run_replay_demo( + spec=spec, + adapter=adapter, + output_sink_factory=lambda output_spec: sink, + ) + + assert result.status == "completed" + assert result.artifacts == ( + OutputArtifact(kind="video/mp4", uri="memory://omnidreams"), + ) + assert [result.step_index for result in sink.results] == [0, 1] + assert pipeline.initialize_cache_calls == [ + { + "text": [["drive"]], + "image_shape": (1, 1, 1, 3, 2, 2), + "view_names": ["camera_front_wide_120fov"], + } + ] + assert len(pipeline.generated_hdmaps) == 2 + torch.testing.assert_close(pipeline.generated_hdmaps[0][0, 0], loaded_hdmap[:1]) + torch.testing.assert_close(pipeline.generated_hdmaps[1][0, 0], loaded_hdmap[1:2]) + + +def test_omnidreams_replay_output_target_path_uses_precomputed_provider( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + import omnidreams.demo.providers as providers_module + + hdmap = tmp_path / "hdmap.mp4" + first_frame = tmp_path / "first.png" + hdmap.write_bytes(b"fake") + first_frame.write_bytes(b"fake") + loaded_hdmap = torch.arange(1 * 3 * 2 * 2).reshape(1, 3, 2, 2) + pipeline = _FakeOmnidreamsPipeline() + output = _RecordingOutputTarget() + monkeypatch.setattr( + providers_module, + "load_first_frame_tensor", + lambda *args, **kwargs: torch.zeros(1, 3, 2, 2), + ) + monkeypatch.setattr( + providers_module, "_load_video", - lambda *args, **kwargs: torch.zeros(2, 3, 2, 2), + lambda *args, **kwargs: loaded_hdmap, + ) + adapter = OmnidreamsDemoAdapter( + pipeline_factory=lambda pipeline_config, device: pipeline, + ) + spec = _replay_demo_spec( + tmp_path=tmp_path, + hdmap=hdmap, + first_frame=first_frame, + total_blocks=1, + ) + + result = run_replay_demo( + spec=spec, + adapter=adapter, + output_target_factory=lambda output_spec: output, + ) + + assert result.status == "completed" + assert [result.step_index for result in output.results] == [0] + assert len(pipeline.generated_hdmaps) == 1 + torch.testing.assert_close(pipeline.generated_hdmaps[0][0, 0], loaded_hdmap) + + +def test_omnidreams_replay_runtime_generates_video_step_result( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + import omnidreams.demo.replay as replay_module + + hdmap = tmp_path / "hdmap.mp4" + first_frame = tmp_path / "first.png" + hdmap.write_bytes(b"fake") + first_frame.write_bytes(b"fake") + pipeline = _FakeOmnidreamsPipeline() + monkeypatch.setattr( + replay_module, + "load_first_frame_tensor", + lambda *args, **kwargs: torch.zeros(1, 3, 2, 2), ) runtime = OmnidreamsRuntime( @@ -323,7 +501,7 @@ def test_omnidreams_replay_runtime_generates_video_step_result( assert request is not None assert request.step_index == 0 assert request.metadata["input_frame_count"] == 1 - result = session.step(InferenceInput()) + result = session.step(InferenceInput(step={"hdmap": torch.zeros(1, 1, 1, 3, 2, 2)})) assert result.step_index == 0 assert result.frame_count == 1 @@ -613,16 +791,73 @@ async def test_omnidreams_demo_runtime_generates_directly_from_controls() -> Non class _RecordingOutputTarget: + def __init__(self) -> None: + self.results: list[StepResult] = [] + def open(self) -> None: return None def write(self, result: StepResult) -> None: - del result + self.results.append(result) def close(self) -> Sequence[OutputArtifact]: return () +def _replay_demo_spec( + *, + tmp_path: Path, + hdmap: Path, + first_frame: Path, + total_blocks: int, +) -> DemoSpec: + return DemoSpec( + model_id=OMNIDREAMS_MODEL_ID, + preset_id=DEFAULT_OMNIDREAMS_PRESET, + input_mode="replay", + scenario={ + "prompt": "drive", + "hdmap_video_paths": (hdmap,), + "first_frame_paths": (first_frame,), + "camera_names": ("camera_front_wide_120fov",), + "total_blocks": total_blocks, + "pixel_height": 2, + "pixel_width": 2, + "fps": 30, + }, + output=Mp4OutputSpec(path=tmp_path / "demo.mp4", fps=30), + config=InferenceConfig( + model_id=OMNIDREAMS_MODEL_ID, + preset_id=DEFAULT_OMNIDREAMS_PRESET, + device="cpu", + runtime_options={"pipeline_config": object()}, + ), + ) + + +class _RecordingOutputSink: + produces_artifacts = True + + def __init__(self) -> None: + self.session_info: SessionInfo | None = None + self.results: list[StepResult] = [] + self.closed = False + + def open(self, session_info: SessionInfo) -> None: + self.session_info = session_info + + def begin_generation(self, generation: int) -> None: + del generation + + def write(self, result: StepResult) -> OutputDecision: + self.results.append(result) + return OutputDecision() + + def close(self) -> Sequence[OutputArtifact]: + self.closed = True + return (OutputArtifact(kind="video/mp4", uri="memory://omnidreams"),) + + class _FactoryRuntime: def start_session(self, inputs: InferenceInput) -> Any: del inputs @@ -635,6 +870,7 @@ def close(self) -> None: class _FakeOmnidreamsPipeline: def __init__(self) -> None: self.initialize_cache_calls: list[dict[str, Any]] = [] + self.generated_hdmaps: list[torch.Tensor] = [] self.released_encoders = False def initialize_cache( @@ -667,7 +903,8 @@ def generate( cache: object, hdmap: torch.Tensor, ) -> torch.Tensor: - del cache, hdmap + del cache + self.generated_hdmaps.append(hdmap.detach().clone()) return torch.full((1, 1, 1, 3, 2, 2), float(autoregressive_index)) def finalize(self, *, autoregressive_index: int, cache: object) -> dict[str, float]: From 3affc4f93f0565411c55dbd4dbe9fa89eaada4d8 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 22:52:50 +0000 Subject: [PATCH 30/51] Fix provider cleanup after session close timeout --- .../flashdreams/runtime/demo/drivers.py | 41 ++++++++++++++++--- .../test_demo_runtime_realtime_driver.py | 5 ++- .../omnidreams/omnidreams/demo/providers.py | 2 +- 3 files changed, 39 insertions(+), 9 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 48842d2c5..f7f672282 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -647,6 +647,7 @@ async def shielded_session_cleanup( return session_edges.close_result(status=status, reason=reason, error=error) async def cleanup() -> RunResult: + session_closed = True if session is not None: session_closed = await _close_model_resource_async( host=host, @@ -656,12 +657,21 @@ async def cleanup() -> RunResult: ) if not session_closed: host.mark_unhealthy("model-affine cleanup timed out") - provider_closed = await _close_model_resource_async( - host=host, - close=provider.close, - session_edges=session_edges, - timeout_s=timeout_s, - ) + if session_closed: + provider_closed = await _close_model_resource_async( + host=host, + close=provider.close, + session_edges=session_edges, + timeout_s=timeout_s, + ) + else: + # The model worker may still be occupied by an orphaned session.close. + # Do not queue provider cleanup behind it after marking the host unhealthy. + provider_closed = await _close_model_resource_direct_async( + close=provider.close, + session_edges=session_edges, + timeout_s=timeout_s, + ) if not provider_closed: host.mark_unhealthy("model-affine cleanup timed out") return session_edges.close_result( @@ -833,6 +843,25 @@ async def _close_model_resource_async( return True +async def _close_model_resource_direct_async( + *, + close: Any, + session_edges: SessionEdges, + timeout_s: float, +) -> bool: + try: + await asyncio.wait_for( + asyncio.to_thread(_close_safely, close, session_edges), + timeout=timeout_s, + ) + except asyncio.TimeoutError as exc: + session_edges.record_orphaned_cleanup(exc) + return False + except Exception as exc: + session_edges.record_cleanup_error(exc) + return True + + def _cleanup_result( cleanup_task: asyncio.Task[RunResult], session_edges: SessionEdges, diff --git a/flashdreams/tests/test_demo_runtime_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py index 80af11980..576f6fdce 100644 --- a/flashdreams/tests/test_demo_runtime_realtime_driver.py +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -339,9 +339,10 @@ async def test_shielded_cleanup_timeout_bounds_shutdown() -> None: assert result.status == "cancelled" assert host.unhealthy_reason == "model-affine cleanup timed out" - assert host.close_targets == [session, provider] + assert host.close_targets == [session] + assert provider.close_count == 1 assert metrics.cleanup_errors == [] - assert len(metrics.orphaned_cleanup_errors) == 2 + assert len(metrics.orphaned_cleanup_errors) == 1 assert metrics.closed diff --git a/integrations/omnidreams/omnidreams/demo/providers.py b/integrations/omnidreams/omnidreams/demo/providers.py index 831330a85..92383fcf2 100644 --- a/integrations/omnidreams/omnidreams/demo/providers.py +++ b/integrations/omnidreams/omnidreams/demo/providers.py @@ -17,7 +17,6 @@ load_first_frame_tensor, ) from flashdreams.runtime.config import InferenceConfig -from flashdreams.runtime.inputs import InferenceInput, InferenceInputSchema, InputField from flashdreams.runtime.demo import ( PreparedScenario, PreparedStep, @@ -25,6 +24,7 @@ UserInputWindow, ) from flashdreams.runtime.demo.session_inputs import ControlDecision +from flashdreams.runtime.inputs import InferenceInput, InferenceInputSchema, InputField from flashdreams.runtime.types import StepRequirements from .spec import OmnidreamsReplayScenario From 19a5727d6baecb7526242bc009a96129698677e0 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 23:19:33 +0000 Subject: [PATCH 31/51] Add OmniDreams replay null output mode --- .../omnidreams/omnidreams/demo/adapter.py | 10 +-- .../omnidreams/omnidreams/demo/app.py | 28 ++++++-- .../omnidreams/tests/test_demo_api.py | 66 ++++++++++++++++++- 3 files changed, 94 insertions(+), 10 deletions(-) diff --git a/integrations/omnidreams/omnidreams/demo/adapter.py b/integrations/omnidreams/omnidreams/demo/adapter.py index 201618716..ec2e0d343 100644 --- a/integrations/omnidreams/omnidreams/demo/adapter.py +++ b/integrations/omnidreams/omnidreams/demo/adapter.py @@ -21,7 +21,6 @@ ) from flashdreams.runtime.demo import ( DemoSpec, - Mp4OutputSpec, PreparedScenario, ) from flashdreams.runtime.demo.session_inputs import ModelInputProvider @@ -87,7 +86,7 @@ def supported_input_modes(self) -> tuple[str, ...]: return ("replay",) def supported_output_modes(self) -> tuple[str, ...]: - return ("mp4",) + return ("mp4", "null") def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: if spec.input_mode != "replay": @@ -95,8 +94,11 @@ def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: "OmniDreams prepare_scenario currently supports only " f"input_mode='replay', got {spec.input_mode!r}." ) - if not isinstance(spec.output, Mp4OutputSpec): - raise ValueError("OmniDreams replay demo currently requires MP4 output.") + if spec.output.mode not in self.supported_output_modes(): + raise ValueError( + "OmniDreams replay demo supports output modes " + f"{self.supported_output_modes()}, got {spec.output.mode!r}." + ) scenario = resolve_replay_scenario( spec.scenario, default_prompt=self._default_replay_prompt(spec.config), diff --git a/integrations/omnidreams/omnidreams/demo/app.py b/integrations/omnidreams/omnidreams/demo/app.py index 1366643f9..08158e4c4 100644 --- a/integrations/omnidreams/omnidreams/demo/app.py +++ b/integrations/omnidreams/omnidreams/demo/app.py @@ -15,6 +15,7 @@ from flashdreams.runtime.demo import ( DemoSpec, Mp4OutputSpec, + NullOutputSpec, WebRTCOutputSpec, ) from flashdreams.runtime.demo.app import DemoApplication @@ -33,7 +34,7 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: ) subparsers = parser.add_subparsers(dest="command", required=True) - replay = subparsers.add_parser("replay", help="Run an MP4 replay demo.") + replay = subparsers.add_parser("replay", help="Run a finite replay demo.") replay.add_argument("--preset-id", default=DEFAULT_OMNIDREAMS_PRESET) replay.add_argument("--device", default="cuda") replay.add_argument("--prompt", default=None) @@ -54,7 +55,8 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: replay.add_argument("--pixel-height", type=int, default=704) replay.add_argument("--pixel-width", type=int, default=1280) replay.add_argument("--fps", type=int, default=30) - replay.add_argument("--output", type=Path, required=True) + replay.add_argument("--output-mode", choices=("mp4", "null"), default="mp4") + replay.add_argument("--output", type=Path, default=None) webrtc = subparsers.add_parser("webrtc", help="Serve a WebRTC driving demo.") webrtc.add_argument("--preset-id", default=DEFAULT_OMNIDREAMS_PRESET) @@ -74,7 +76,13 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: webrtc.add_argument("--client-liveness-timeout-s", type=float, default=10.0) webrtc.add_argument("--debug-serve-hdmaps", action="store_true") webrtc.add_argument("--prefer-sw-encoder", action="store_true") - return parser.parse_args(argv) + args = parser.parse_args(argv) + if args.command == "replay": + if args.output_mode == "mp4" and args.output is None: + parser.error("replay --output is required when --output-mode=mp4.") + if args.output_mode == "null" and args.output is not None: + parser.error("replay --output is only valid when --output-mode=mp4.") + return args class OmnidreamsDemoApplication(DemoApplication): @@ -129,7 +137,7 @@ def _replay_spec(args: argparse.Namespace) -> DemoSpec: preset_id=args.preset_id, input_mode="replay", scenario=scenario, - output=Mp4OutputSpec(path=args.output, fps=args.fps), + output=_replay_output_spec(args), config=InferenceConfig( model_id=OMNIDREAMS_MODEL_ID, preset_id=args.preset_id, @@ -138,6 +146,18 @@ def _replay_spec(args: argparse.Namespace) -> DemoSpec: ) +def _replay_output_spec(args: argparse.Namespace) -> Mp4OutputSpec | NullOutputSpec: + if args.output_mode == "mp4": + if args.output is None: + raise ValueError("OmniDreams MP4 replay requires --output.") + return Mp4OutputSpec(path=args.output, fps=args.fps) + if args.output_mode == "null": + return NullOutputSpec() + raise ValueError( + f"Unsupported OmniDreams replay output mode: {args.output_mode!r}." + ) + + def _webrtc_spec(args: argparse.Namespace, *, device: str) -> DemoSpec: return DemoSpec( model_id=OMNIDREAMS_MODEL_ID, diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index 34c2e81b3..21c6e70ee 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -51,6 +51,7 @@ from flashdreams.runtime.demo import ( DemoSpec, Mp4OutputSpec, + NullOutputSpec, OutputDecision, SessionInfo, UserInputWindow, @@ -70,12 +71,23 @@ def test_omnidreams_demo_defaults_to_stable_non_perf_preset() -> None: assert not args.preset_id.endswith("-perf") +def test_omnidreams_replay_cli_builds_null_output_spec() -> None: + args = parse_args(["replay", "--output-mode", "null"]) + + spec = _replay_spec(args) + + assert spec.input_mode == "replay" + assert isinstance(spec.output, NullOutputSpec) + assert spec.config is not None + assert spec.config.model_id == OMNIDREAMS_MODEL_ID + + def test_omnidreams_demo_adapter_declares_replay_modes_only() -> None: adapter = OmnidreamsDemoAdapter() assert adapter.model_id == OMNIDREAMS_MODEL_ID assert adapter.supported_input_modes() == ("replay",) - assert adapter.supported_output_modes() == ("mp4",) + assert adapter.supported_output_modes() == ("mp4", "null") assert [ field.name for field in adapter.inference_input_schema.global_conditioning_fields @@ -409,6 +421,55 @@ def test_omnidreams_replay_run_mode_uses_precomputed_provider( torch.testing.assert_close(pipeline.generated_hdmaps[1][0, 0], loaded_hdmap[1:2]) +def test_omnidreams_replay_null_output_uses_precomputed_provider( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + import omnidreams.demo.providers as providers_module + + hdmap = tmp_path / "hdmap.mp4" + first_frame = tmp_path / "first.png" + hdmap.write_bytes(b"fake") + first_frame.write_bytes(b"fake") + loaded_hdmap = torch.arange(2 * 3 * 2 * 2).reshape(2, 3, 2, 2) + pipeline = _FakeOmnidreamsPipeline() + monkeypatch.setattr( + providers_module, + "load_first_frame_tensor", + lambda *args, **kwargs: torch.zeros(1, 3, 2, 2), + ) + monkeypatch.setattr( + providers_module, + "_load_video", + lambda *args, **kwargs: loaded_hdmap, + ) + adapter = OmnidreamsDemoAdapter( + pipeline_factory=lambda pipeline_config, device: pipeline, + ) + spec = _replay_demo_spec( + tmp_path=tmp_path, + hdmap=hdmap, + first_frame=first_frame, + total_blocks=2, + output=NullOutputSpec(), + ) + + result = run_replay_demo(spec=spec, adapter=adapter) + + assert result.status == "completed" + assert result.artifacts == () + assert pipeline.initialize_cache_calls == [ + { + "text": [["drive"]], + "image_shape": (1, 1, 1, 3, 2, 2), + "view_names": ["camera_front_wide_120fov"], + } + ] + assert len(pipeline.generated_hdmaps) == 2 + torch.testing.assert_close(pipeline.generated_hdmaps[0][0, 0], loaded_hdmap[:1]) + torch.testing.assert_close(pipeline.generated_hdmaps[1][0, 0], loaded_hdmap[1:2]) + + def test_omnidreams_replay_output_target_path_uses_precomputed_provider( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -810,6 +871,7 @@ def _replay_demo_spec( hdmap: Path, first_frame: Path, total_blocks: int, + output: Mp4OutputSpec | NullOutputSpec | None = None, ) -> DemoSpec: return DemoSpec( model_id=OMNIDREAMS_MODEL_ID, @@ -825,7 +887,7 @@ def _replay_demo_spec( "pixel_width": 2, "fps": 30, }, - output=Mp4OutputSpec(path=tmp_path / "demo.mp4", fps=30), + output=output or Mp4OutputSpec(path=tmp_path / "demo.mp4", fps=30), config=InferenceConfig( model_id=OMNIDREAMS_MODEL_ID, preset_id=DEFAULT_OMNIDREAMS_PRESET, From 62c3b79cc19e187b55397cc4bc94230d9ddfc44d Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 23:32:46 +0000 Subject: [PATCH 32/51] Shield pre-edge provider cleanup from cancellation --- .../flashdreams/runtime/demo/drivers.py | 49 ++++++++-- .../tests/test_demo_runtime_run_modes.py | 89 +++++++++++++++++++ 2 files changed, 133 insertions(+), 5 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index f7f672282..5680a794e 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -801,12 +801,51 @@ async def _close_provider_async( session_edges: SessionEdges | None, ) -> None: try: - await context.host.call_async(provider.close) + close_task = asyncio.create_task(context.host.call_async(provider.close)) + except RuntimeError as close_exc: + _record_provider_cleanup_error( + context=context, + session_edges=session_edges, + exc=close_exc, + ) + return + + while not close_task.done(): + try: + await asyncio.shield(close_task) + except asyncio.CancelledError: + uncancel_current_task() + continue + except Exception: + break + + try: + await close_task + except asyncio.CancelledError: + uncancel_current_task() + _record_provider_cleanup_error( + context=context, + session_edges=session_edges, + exc=RuntimeError("provider cleanup was cancelled"), + ) except Exception as close_exc: - if session_edges is not None: - session_edges.record_cleanup_error(close_exc) - else: - _record_run_cleanup_error(context, close_exc) + _record_provider_cleanup_error( + context=context, + session_edges=session_edges, + exc=close_exc, + ) + + +def _record_provider_cleanup_error( + *, + context: RunContext, + session_edges: SessionEdges | None, + exc: Exception, +) -> None: + if session_edges is not None: + session_edges.record_cleanup_error(exc) + else: + _record_run_cleanup_error(context, exc) def _record_run_cleanup_error(context: RunContext, exc: Exception) -> None: diff --git a/flashdreams/tests/test_demo_runtime_run_modes.py b/flashdreams/tests/test_demo_runtime_run_modes.py index de01b2771..8742eb0e9 100644 --- a/flashdreams/tests/test_demo_runtime_run_modes.py +++ b/flashdreams/tests/test_demo_runtime_run_modes.py @@ -5,6 +5,7 @@ import asyncio import contextlib +import threading from collections.abc import Callable, Coroutine, Mapping, Sequence from pathlib import Path from typing import Any, cast @@ -164,6 +165,53 @@ async def test_fake_webrtc_offer_reserves_before_prepare_or_negotiation() -> Non assert mode.created_edges[0].is_closed +@pytest.mark.asyncio +async def test_async_session_cancellation_shields_pre_edge_provider_cleanup() -> None: + spec = DemoSpec( + model_id="fake-demo", + input_mode="keyboard-driving", + output=WebRTCOutputSpec(port=8081), + ) + provider = _BlockingCloseProvider() + adapter = _BlockingCloseAdapter(provider=provider) + mode = _CancelBeforeEdgesRunMode(name="webrtc", driver=_ClosingAsyncDriver()) + context = mode.create_run_context( + spec=spec, + adapter=adapter, + host=RuntimeHost(_UnusedRuntime()), + model_warmup_plan=ModelWarmupPlan(), + ) + scenario = adapter.prepare_scenario(spec) + task = asyncio.create_task( + run_demo_session_async( + context=context, + spec=spec, + scenario=scenario, + adapter=adapter, + run_mode=mode, + pipeline=StepPipeline(), + ) + ) + + close_started = await asyncio.to_thread(provider.close_started.wait, 1.0) + assert close_started + task.cancel() + provider.release_close.set() + try: + result = await asyncio.wait_for(task, timeout=1.0) + finally: + provider.release_close.set() + context.host.close() + + assert result.status == "cancelled" + assert result.reason == "cancelled during session assembly" + assert provider.close_count == 1 + assert mode.created_edges == [] + assert mode.admission.reservations[0].release_count == 1 + run_metrics = cast(InMemorySessionMetricsRecorder, context.run_metrics) + assert run_metrics.sessions == [result] + + def test_run_demo_session_rejects_reused_closed_session_edges() -> None: spec = DemoSpec( model_id="fake-demo", @@ -426,6 +474,21 @@ def create_model_input_provider( return provider +class _BlockingCloseAdapter(_FakeAdapter): + def __init__(self, *, provider: "_BlockingCloseProvider") -> None: + super().__init__() + self.provider = provider + + def create_model_input_provider( + self, + spec: DemoSpec, + scenario: PreparedScenario, + ) -> "_BlockingCloseProvider": + del spec, scenario + self.providers.append(self.provider) + return self.provider + + class _FakeProvider: capabilities = ProviderCapabilities( supports_recorded_input=True, @@ -440,6 +503,19 @@ def close(self) -> None: self.close_count += 1 +class _BlockingCloseProvider(_FakeProvider): + def __init__(self) -> None: + super().__init__() + self.close_started = threading.Event() + self.release_close = threading.Event() + + def close(self) -> None: + self.close_started.set() + if not self.release_close.wait(timeout=1.0): + raise RuntimeError("timed out waiting to release provider close") + super().close() + + class _FakeRunMode: def __init__( self, @@ -561,6 +637,19 @@ def create_session_edges( return self._edges +class _CancelBeforeEdgesRunMode(_FakeRunMode): + def validate_session( + self, + *, + spec: DemoSpec, + scenario: Any, + adapter: Any, + provider: Any, + ) -> None: + del spec, scenario, adapter, provider + raise asyncio.CancelledError + + class _ClosingSyncDriver: def run_one_session( self, From f988ce44670888bbadabbef9ba9d2d0b9535f269 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 23:39:19 +0000 Subject: [PATCH 33/51] Preserve worker affinity for cleanup timeouts --- .../flashdreams/runtime/demo/drivers.py | 69 +++++++------------ .../test_demo_runtime_realtime_driver.py | 10 +-- 2 files changed, 30 insertions(+), 49 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 5680a794e..cdad9a9f5 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -647,32 +647,14 @@ async def shielded_session_cleanup( return session_edges.close_result(status=status, reason=reason, error=error) async def cleanup() -> RunResult: - session_closed = True - if session is not None: - session_closed = await _close_model_resource_async( - host=host, - close=session.close, - session_edges=session_edges, - timeout_s=timeout_s, - ) - if not session_closed: - host.mark_unhealthy("model-affine cleanup timed out") - if session_closed: - provider_closed = await _close_model_resource_async( - host=host, - close=provider.close, - session_edges=session_edges, - timeout_s=timeout_s, - ) - else: - # The model worker may still be occupied by an orphaned session.close. - # Do not queue provider cleanup behind it after marking the host unhealthy. - provider_closed = await _close_model_resource_direct_async( - close=provider.close, - session_edges=session_edges, - timeout_s=timeout_s, - ) - if not provider_closed: + resources_closed = await _close_model_resources_async( + host=host, + session=session, + provider=provider, + session_edges=session_edges, + timeout_s=timeout_s, + ) + if not resources_closed: host.mark_unhealthy("model-affine cleanup timed out") return session_edges.close_result( status=status, @@ -862,16 +844,22 @@ def _record_run_session_error(context: RunContext, exc: Exception) -> None: return -async def _close_model_resource_async( +async def _close_model_resources_async( *, host: RuntimeHost, - close: Any, + session: InferenceSession | None, + provider: ModelInputProvider, session_edges: SessionEdges, timeout_s: float, ) -> bool: try: await asyncio.wait_for( - host.call_async(_close_safely, close, session_edges), + host.call_async( + _close_model_resources_safely, + session.close if session is not None else None, + provider.close, + session_edges, + ), timeout=timeout_s, ) except asyncio.TimeoutError as exc: @@ -882,23 +870,14 @@ async def _close_model_resource_async( return True -async def _close_model_resource_direct_async( - *, - close: Any, +def _close_model_resources_safely( + session_close: Any | None, + provider_close: Any, session_edges: SessionEdges, - timeout_s: float, -) -> bool: - try: - await asyncio.wait_for( - asyncio.to_thread(_close_safely, close, session_edges), - timeout=timeout_s, - ) - except asyncio.TimeoutError as exc: - session_edges.record_orphaned_cleanup(exc) - return False - except Exception as exc: - session_edges.record_cleanup_error(exc) - return True +) -> None: + if session_close is not None: + _close_safely(session_close, session_edges) + _close_safely(provider_close, session_edges) def _cleanup_result( diff --git a/flashdreams/tests/test_demo_runtime_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py index 576f6fdce..3ab873e37 100644 --- a/flashdreams/tests/test_demo_runtime_realtime_driver.py +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -339,8 +339,8 @@ async def test_shielded_cleanup_timeout_bounds_shutdown() -> None: assert result.status == "cancelled" assert host.unhealthy_reason == "model-affine cleanup timed out" - assert host.close_targets == [session] - assert provider.close_count == 1 + assert host.close_targets == [session, provider] + assert provider.close_count == 0 assert metrics.cleanup_errors == [] assert len(metrics.orphaned_cleanup_errors) == 1 assert metrics.closed @@ -844,8 +844,10 @@ async def call_async( **kwargs: object, ) -> Any: del func, kwargs - close = cast(Callable[[], None], args[0]) - self.close_targets.append(getattr(close, "__self__", close)) + for arg in args: + if callable(arg): + close = cast(Callable[[], None], arg) + self.close_targets.append(getattr(close, "__self__", close)) await asyncio.Event().wait() def mark_unhealthy( From c3609f720b732241fc738b84aa6a0829ab256ecc Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 23:46:21 +0000 Subject: [PATCH 34/51] Mark cleanup dispatch failures unhealthy --- .../flashdreams/runtime/demo/drivers.py | 1 + .../test_demo_runtime_realtime_driver.py | 55 +++++++++++++++++++ 2 files changed, 56 insertions(+) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index cdad9a9f5..2ebdb589e 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -867,6 +867,7 @@ async def _close_model_resources_async( return False except Exception as exc: session_edges.record_cleanup_error(exc) + return False return True diff --git a/flashdreams/tests/test_demo_runtime_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py index 3ab873e37..a9ccea100 100644 --- a/flashdreams/tests/test_demo_runtime_realtime_driver.py +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -346,6 +346,35 @@ async def test_shielded_cleanup_timeout_bounds_shutdown() -> None: assert metrics.closed +@pytest.mark.asyncio +async def test_shielded_cleanup_dispatch_failure_marks_host_unhealthy() -> None: + host = _RejectingHost(RuntimeError("worker rejected cleanup")) + session = _FakeRealtimeSession(num_steps=1) + provider = _FakeRealtimeProvider() + metrics = InMemorySessionMetricsRecorder() + edges = _edges(metrics=metrics) + + result = await shielded_session_cleanup( + host=cast(RuntimeHost, host), + session=session, + provider=provider, + session_edges=edges, + status="cancelled", + reason="dispatch failure", + error=None, + timeout_s=0.001, + ) + + assert result.status == "cancelled" + assert host.unhealthy_reason == "model-affine cleanup timed out" + assert host.cleanup_dispatch_count == 1 + assert session.close_count == 0 + assert provider.close_count == 0 + assert metrics.cleanup_errors == ["worker rejected cleanup"] + assert metrics.orphaned_cleanup_errors == [] + assert metrics.closed + + @pytest.mark.asyncio async def test_realtime_driver_applies_backpressure_through_clock() -> None: session = _FakeRealtimeSession(num_steps=2) @@ -857,3 +886,29 @@ def mark_unhealthy( ) -> None: del error self.unhealthy_reason = reason + + +class _RejectingHost: + def __init__(self, exc: Exception) -> None: + self.exc = exc + self.cleanup_dispatch_count = 0 + self.unhealthy_reason: str | None = None + + async def call_async( + self, + func: Callable[..., Any], + /, + *args: object, + **kwargs: object, + ) -> Any: + del func, args, kwargs + self.cleanup_dispatch_count += 1 + raise self.exc + + def mark_unhealthy( + self, + reason: str = "marked unhealthy", + error: Exception | None = None, + ) -> None: + del error + self.unhealthy_reason = reason From e976b4445d031c4e48c4e4dbbb5287d0c37dbf19 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Sun, 9 Aug 2026 23:54:38 +0000 Subject: [PATCH 35/51] Document worker-affine cleanup timeout tradeoff --- flashdreams/flashdreams/runtime/demo/drivers.py | 4 ++++ 1 file changed, 4 insertions(+) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 2ebdb589e..75f86aab0 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -863,6 +863,10 @@ async def _close_model_resources_async( timeout=timeout_s, ) except asyncio.TimeoutError as exc: + # Keep provider cleanup ordered behind session cleanup on the model worker. + # A timed-out session close may still hold CUDA/Triton state, so running + # provider cleanup on another thread or replacing the worker is unsafe. + # The caller marks the host unhealthy so future sessions reject instead. session_edges.record_orphaned_cleanup(exc) return False except Exception as exc: From d7b39fdc8218f0d47085b521fe394b2ed8f1d336 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 00:39:56 +0000 Subject: [PATCH 36/51] Phase 14: add OmniDreams Ludus replay provider --- .../omnidreams/omnidreams/demo/__init__.py | 16 +- .../omnidreams/omnidreams/demo/adapter.py | 48 +- .../omnidreams/omnidreams/demo/app.py | 62 ++- .../omnidreams/omnidreams/demo/providers.py | 492 +++++++++++++++++- .../omnidreams/omnidreams/demo/replay.py | 49 +- .../omnidreams/omnidreams/demo/spec.py | 282 +++++++++- .../omnidreams/tests/test_demo_api.py | 292 +++++++++++ 7 files changed, 1214 insertions(+), 27 deletions(-) diff --git a/integrations/omnidreams/omnidreams/demo/__init__.py b/integrations/omnidreams/omnidreams/demo/__init__.py index 4365b1051..a3d646989 100644 --- a/integrations/omnidreams/omnidreams/demo/__init__.py +++ b/integrations/omnidreams/omnidreams/demo/__init__.py @@ -4,18 +4,32 @@ """Experimental OmniDreams demo adapter built on ``flashdreams.runtime.demo``.""" from omnidreams.demo.adapter import OmnidreamsDemoAdapter -from omnidreams.demo.providers import PrecomputedHDMapProvider +from omnidreams.demo.providers import ( + LudusSceneConditioningProvider, + PrecomputedHDMapProvider, +) from omnidreams.demo.spec import ( DEFAULT_OMNIDREAMS_PRESET, + OMNIDREAMS_CONDITIONING_LUDUS, + OMNIDREAMS_CONDITIONING_MODES, + OMNIDREAMS_CONDITIONING_PRECOMPUTED, OMNIDREAMS_MODEL_ID, + OmnidreamsKeyboardTraceEvent, + OmnidreamsLudusReplayScenario, OmnidreamsReplayScenario, OmnidreamsWebRTCScenario, ) __all__ = [ "DEFAULT_OMNIDREAMS_PRESET", + "OMNIDREAMS_CONDITIONING_LUDUS", + "OMNIDREAMS_CONDITIONING_MODES", + "OMNIDREAMS_CONDITIONING_PRECOMPUTED", "OMNIDREAMS_MODEL_ID", + "LudusSceneConditioningProvider", "OmnidreamsDemoAdapter", + "OmnidreamsKeyboardTraceEvent", + "OmnidreamsLudusReplayScenario", "OmnidreamsReplayScenario", "OmnidreamsWebRTCScenario", "PrecomputedHDMapProvider", diff --git a/integrations/omnidreams/omnidreams/demo/adapter.py b/integrations/omnidreams/omnidreams/demo/adapter.py index ec2e0d343..3a7e58dc4 100644 --- a/integrations/omnidreams/omnidreams/demo/adapter.py +++ b/integrations/omnidreams/omnidreams/demo/adapter.py @@ -27,7 +27,9 @@ from flashdreams.runtime.interfaces import InferenceRuntime from .providers import ( + LudusSceneConditioningProvider, PrecomputedHDMapProvider, + keyboard_driving_user_input_schema, precomputed_hdmap_inference_input_schema, ) from .runtime import ( @@ -37,7 +39,12 @@ ) from .spec import ( DEFAULT_OMNIDREAMS_PRESET, + OMNIDREAMS_CONDITIONING_LUDUS, + OMNIDREAMS_CONDITIONING_MODES, + OMNIDREAMS_CONDITIONING_PRECOMPUTED, OMNIDREAMS_MODEL_ID, + conditioning_mode_from_scenario, + resolve_ludus_replay_scenario, resolve_replay_scenario, ) @@ -88,6 +95,9 @@ def supported_input_modes(self) -> tuple[str, ...]: def supported_output_modes(self) -> tuple[str, ...]: return ("mp4", "null") + def supported_conditioning_modes(self) -> tuple[str, ...]: + return OMNIDREAMS_CONDITIONING_MODES + def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: if spec.input_mode != "replay": raise ValueError( @@ -99,18 +109,29 @@ def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: "OmniDreams replay demo supports output modes " f"{self.supported_output_modes()}, got {spec.output.mode!r}." ) - scenario = resolve_replay_scenario( - spec.scenario, - default_prompt=self._default_replay_prompt(spec.config), - ) + conditioning_mode = conditioning_mode_from_scenario(spec.scenario) + if conditioning_mode == OMNIDREAMS_CONDITIONING_PRECOMPUTED: + scenario = resolve_replay_scenario( + spec.scenario, + default_prompt=self._default_replay_prompt(spec.config), + ) + source_schema = UserInputSchema(description="fixed OmniDreams replay input") + elif conditioning_mode == OMNIDREAMS_CONDITIONING_LUDUS: + scenario = resolve_ludus_replay_scenario(spec.scenario) + source_schema = keyboard_driving_user_input_schema() + else: + raise ValueError( + f"Unsupported OmniDreams conditioning mode: {conditioning_mode!r}." + ) return PreparedScenario( initial_inputs=InferenceInput( global_conditioning={"scenario": scenario}, ), - source_schema=UserInputSchema(description="fixed OmniDreams replay input"), + source_schema=source_schema, canonicalizer=InputCanonicalizer(), mapping=self._mapping, metadata={ + "conditioning_mode": conditioning_mode, "model_id": self.model_id, "preset_id": self._preset_id(spec.config), "num_views": len(scenario.camera_names), @@ -142,11 +163,26 @@ def create_model_input_provider( ) -> ModelInputProvider: if spec.input_mode != "replay": raise ValueError( - "OmniDreams precomputed HDMap provider currently supports only " + "OmniDreams replay providers currently support only " f"input_mode='replay', got {spec.input_mode!r}." ) if spec.config is None: raise RuntimeError("DemoSpec.config was not initialized.") + conditioning_mode = str( + scenario.metadata.get( + "conditioning_mode", + conditioning_mode_from_scenario(spec.scenario), + ) + ) + if conditioning_mode == OMNIDREAMS_CONDITIONING_LUDUS: + return LudusSceneConditioningProvider( + scenario=scenario, + config=spec.config, + ) + if conditioning_mode != OMNIDREAMS_CONDITIONING_PRECOMPUTED: + raise ValueError( + f"Unsupported OmniDreams conditioning mode: {conditioning_mode!r}." + ) return PrecomputedHDMapProvider( scenario=scenario, config=spec.config, diff --git a/integrations/omnidreams/omnidreams/demo/app.py b/integrations/omnidreams/omnidreams/demo/app.py index 08158e4c4..ac1641917 100644 --- a/integrations/omnidreams/omnidreams/demo/app.py +++ b/integrations/omnidreams/omnidreams/demo/app.py @@ -6,6 +6,7 @@ from __future__ import annotations import argparse +import math from pathlib import Path from typing import Any @@ -23,6 +24,10 @@ from .adapter import OmnidreamsDemoAdapter from .spec import ( DEFAULT_OMNIDREAMS_PRESET, + DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, + OMNIDREAMS_CONDITIONING_LUDUS, + OMNIDREAMS_CONDITIONING_MODES, + OMNIDREAMS_CONDITIONING_PRECOMPUTED, OMNIDREAMS_MODEL_ID, OmnidreamsWebRTCScenario, ) @@ -37,10 +42,29 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: replay = subparsers.add_parser("replay", help="Run a finite replay demo.") replay.add_argument("--preset-id", default=DEFAULT_OMNIDREAMS_PRESET) replay.add_argument("--device", default="cuda") + replay.add_argument("--seed", type=int, default=42) + replay.add_argument( + "--conditioning-mode", + choices=OMNIDREAMS_CONDITIONING_MODES, + default=OMNIDREAMS_CONDITIONING_PRECOMPUTED, + ) replay.add_argument("--prompt", default=None) replay.add_argument("--hdmap-video-paths", type=_split_paths, default=()) replay.add_argument("--first-frame-paths", type=_split_paths, default=()) replay.add_argument("--camera-names", type=_split_strings, default=()) + replay.add_argument("--keyboard-trace", type=Path, default=None) + replay.add_argument("--scene-path", type=Path, default=None) + replay.add_argument("--scene-dir", type=Path, default=None) + replay.add_argument("--scene-uuid", default=DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID) + replay.add_argument("--scene-variant", default="default") + replay.add_argument("--camera-name", default="camera_front_wide_120fov") + replay.add_argument("--move-speed-per-s", type=float, default=6.0) + replay.add_argument( + "--rotate-speed-rad-per-s", + type=float, + default=math.radians(35.0), + ) + replay.add_argument("--ludus-backend", choices=("cuda", "vulkan"), default="cuda") replay.add_argument( "--example-data", action=argparse.BooleanOptionalAction, @@ -82,6 +106,14 @@ def parse_args(argv: list[str] | None = None) -> argparse.Namespace: parser.error("replay --output is required when --output-mode=mp4.") if args.output_mode == "null" and args.output is not None: parser.error("replay --output is only valid when --output-mode=mp4.") + if ( + args.conditioning_mode == OMNIDREAMS_CONDITIONING_LUDUS + and args.keyboard_trace is None + ): + parser.error( + "replay --keyboard-trace is required when " + "--conditioning-mode=ludus-scene-driving." + ) return args @@ -116,6 +148,7 @@ def main(argv: list[str] | None = None) -> None: def _replay_spec(args: argparse.Namespace) -> DemoSpec: scenario: dict[str, object] = { + "conditioning_mode": args.conditioning_mode, "example_data": args.example_data, "example_data_uuid": args.example_data_uuid, "total_blocks": args.total_blocks, @@ -125,12 +158,27 @@ def _replay_spec(args: argparse.Namespace) -> DemoSpec: } if args.prompt: scenario["prompt"] = args.prompt - if args.hdmap_video_paths: - scenario["hdmap_video_paths"] = args.hdmap_video_paths - if args.first_frame_paths: - scenario["first_frame_paths"] = args.first_frame_paths - if args.camera_names: - scenario["camera_names"] = args.camera_names + if args.conditioning_mode == OMNIDREAMS_CONDITIONING_LUDUS: + scenario.update( + { + "keyboard_trace_path": args.keyboard_trace, + "scene_path": args.scene_path, + "scene_dir": args.scene_dir, + "scene_uuid": args.scene_uuid, + "scene_variant": args.scene_variant, + "camera_name": args.camera_name, + "move_speed_per_s": args.move_speed_per_s, + "rotate_speed_rad_per_s": args.rotate_speed_rad_per_s, + "ludus_backend": args.ludus_backend, + } + ) + else: + if args.hdmap_video_paths: + scenario["hdmap_video_paths"] = args.hdmap_video_paths + if args.first_frame_paths: + scenario["first_frame_paths"] = args.first_frame_paths + if args.camera_names: + scenario["camera_names"] = args.camera_names return DemoSpec( model_id=OMNIDREAMS_MODEL_ID, @@ -142,6 +190,8 @@ def _replay_spec(args: argparse.Namespace) -> DemoSpec: model_id=OMNIDREAMS_MODEL_ID, preset_id=args.preset_id, device=args.device, + seed=args.seed, + runtime_options={"seed": args.seed}, ), ) diff --git a/integrations/omnidreams/omnidreams/demo/providers.py b/integrations/omnidreams/omnidreams/demo/providers.py index 92383fcf2..72a174d3b 100644 --- a/integrations/omnidreams/omnidreams/demo/providers.py +++ b/integrations/omnidreams/omnidreams/demo/providers.py @@ -5,8 +5,12 @@ from __future__ import annotations +import contextlib import os +from pathlib import Path +from typing import Any +import numpy as np import torch import torch.distributed as dist from loguru import logger @@ -24,10 +28,27 @@ UserInputWindow, ) from flashdreams.runtime.demo.session_inputs import ControlDecision -from flashdreams.runtime.inputs import InferenceInput, InferenceInputSchema, InputField +from flashdreams.runtime.demo.timing import SPARSE_KEY_SEGMENTS_METADATA_KEY +from flashdreams.runtime.inputs import ( + InferenceInput, + InferenceInputSchema, + InputField, + UserInputCapability, + UserInputSchema, +) from flashdreams.runtime.types import StepRequirements +from flashdreams.serving.realtime.input import ( + WSAD_SUPPORTED_KEYS, + CameraPoseIntegrator, + KeyboardResampler, + PoseSegment, +) -from .spec import OmnidreamsReplayScenario +from .spec import ( + DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, + OmnidreamsLudusReplayScenario, + OmnidreamsReplayScenario, +) class PrecomputedHDMapProvider: @@ -39,7 +60,7 @@ def __init__( scenario: PreparedScenario, config: InferenceConfig, ) -> None: - self._scenario = _scenario_from_prepared(scenario) + self._scenario = _precomputed_scenario_from_prepared(scenario) self._device = _device_from_config(config) self._dtype = torch.bfloat16 self._frame_start = 0 @@ -150,6 +171,248 @@ def _require_open(self) -> None: raise RuntimeError("OmniDreams precomputed HDMap provider is closed.") +class LudusSceneConditioningProvider: + """Render finite Ludus keyboard-driving traces into OmniDreams HDMaps.""" + + def __init__( + self, + *, + scenario: PreparedScenario, + config: InferenceConfig, + ) -> None: + self._scenario = _ludus_scenario_from_prepared(scenario) + self._device = _device_from_config(config) + self._dtype = torch.bfloat16 + self._closed = False + self._scene: Any | None = None + self._rasterizer: Any | None = None + self._pose_integrator: CameraPoseIntegrator | None = None + self._keyboard_resampler: KeyboardResampler | None = None + self._next_timestamp_us = 0 + self._step_index = 0 + self.capabilities = ProviderCapabilities( + supports_realtime_clock=True, + supports_recorded_input=True, + supports_reset=True, + deterministic_given_inputs=True, + user_input_schema=keyboard_driving_user_input_schema(), + inference_input_schema=precomputed_hdmap_inference_input_schema(), + ) + + def prepare_initial_input(self) -> InferenceInput: + self._require_open() + scene = self._ensure_scene_loaded() + return InferenceInput( + global_conditioning={ + "scenario": self._scenario, + "prompt": [[str(scene.prompt)]], + "first_frame": _initial_rgb_tensor( + scene.initial_rgb, + device=self._device, + dtype=self._dtype, + ), + }, + metadata={ + "view_names": self._scenario.camera_names, + "scene_id": str(getattr(scene, "scene_id", "")), + }, + ) + + def prepare_step( + self, + *, + request: StepRequirements, + user_window: UserInputWindow, + ) -> PreparedStep: + self._require_open() + scenario = self._scenario + if request.step_index >= scenario.total_blocks: + return PreparedStep( + control=ControlDecision( + close_session=True, + reason="OmniDreams Ludus replay input exhausted.", + ) + ) + + self._ensure_scene_loaded() + pose_integrator = self._require_pose_integrator() + rasterizer = self._require_rasterizer() + segments, frame_times = self._sample_controls( + request=request, + user_window=user_window, + ) + rig_poses_world = pose_integrator.integrate_chunk( + segments=segments, + frame_times=frame_times, + ) + timestamps_us = self._consume_timestamps(request.input_frame_count) + raster_chunk = rasterizer.render_chunk( + rig_poses_world=rig_poses_world, + timestamps_us=timestamps_us, + ) + hdmap = _condition_frames_tensor( + raster_chunk.frames, + device=self._device, + dtype=self._dtype, + ) + self._step_index += 1 + return PreparedStep( + inference_input=InferenceInput( + step={"hdmap": hdmap}, + metadata={ + "frame_timestamps_us": tuple(int(t) for t in timestamps_us), + "keyboard_segments": _segments_metadata(segments), + "camera_name": scenario.camera_name, + "scene_uuid": scenario.scene_uuid, + }, + ) + ) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self._require_open() + if self._scene is not None: + self._reset_driving_state(self._scene) + else: + self._step_index = 0 + self._next_timestamp_us = 0 + + def close(self) -> None: + if self._closed: + return + self._closed = True + rasterizer = self._rasterizer + self._rasterizer = None + self._scene = None + self._pose_integrator = None + self._keyboard_resampler = None + _close_rasterizer(rasterizer) + + def _ensure_scene_loaded(self) -> Any: + if self._scene is not None: + return self._scene + scenario = self._scenario + scene_path = _resolve_ludus_scene_path(scenario) + scene = _load_ludus_scene_bundle(scenario, scene_path) + rasterizer = _new_ludus_rasterizer(scenario) + try: + rasterizer.load_scene(scene) + except Exception: + with contextlib.suppress(Exception): + _close_rasterizer(rasterizer) + raise + self._scene = scene + self._rasterizer = rasterizer + self._reset_driving_state(scene) + if _is_rank_zero(): + logger.info( + "Loaded OmniDreams Ludus replay scene={} camera={} trace_events={}", + scene_path, + scenario.camera_name, + len(scenario.keyboard_events), + ) + return scene + + def _reset_driving_state(self, scene: Any) -> None: + scenario = self._scenario + pose_integrator = CameraPoseIntegrator( + move_speed_per_s=scenario.move_speed_per_s, + rotate_speed_rad_per_s=scenario.rotate_speed_rad_per_s, + coordinate_system="FLU", + ) + pose_integrator.reset(np.asarray(scene.initial_rig_to_world, dtype=np.float32)) + keyboard_resampler = KeyboardResampler( + fps=float(scenario.fps), + supported_keys=WSAD_SUPPORTED_KEYS, + ) + for event in scenario.keyboard_events: + keyboard_resampler.on_edge( + arrival_t=event.timestamp_s, + event=event.event, + key=event.key, + ) + self._pose_integrator = pose_integrator + self._keyboard_resampler = keyboard_resampler + self._next_timestamp_us = int(scene.initial_timestamp_us) + self._step_index = 0 + + def _sample_controls( + self, + *, + request: StepRequirements, + user_window: UserInputWindow, + ) -> tuple[list[PoseSegment], list[float]]: + raw_segments = user_window.metadata.get(SPARSE_KEY_SEGMENTS_METADATA_KEY) + if isinstance(raw_segments, tuple): + frame_times = list(user_window.frame_times) + if len(frame_times) != request.input_frame_count: + raise RuntimeError( + "OmniDreams Ludus realtime window frame_times length does " + "not match the requested input frame count." + ) + return [_pose_segment(segment) for segment in raw_segments], frame_times + if raw_segments is not None: + raise RuntimeError( + "OmniDreams Ludus realtime key segments metadata must be a tuple." + ) + return self._require_keyboard_resampler().sample_chunk( + request.input_frame_count + ) + + def _consume_timestamps(self, num_frames: int) -> np.ndarray: + step_us = int(round(1_000_000 / float(self._scenario.fps))) + timestamps = np.array( + [ + self._next_timestamp_us + frame_index * step_us + for frame_index in range(num_frames) + ], + dtype=np.int64, + ) + self._next_timestamp_us += num_frames * step_us + return timestamps + + def _require_rasterizer(self) -> Any: + if self._rasterizer is None: + raise RuntimeError("OmniDreams Ludus rasterizer is not initialized.") + return self._rasterizer + + def _require_pose_integrator(self) -> CameraPoseIntegrator: + if self._pose_integrator is None: + raise RuntimeError("OmniDreams Ludus pose integrator is not initialized.") + return self._pose_integrator + + def _require_keyboard_resampler(self) -> KeyboardResampler: + if self._keyboard_resampler is None: + raise RuntimeError( + "OmniDreams Ludus keyboard resampler is not initialized." + ) + return self._keyboard_resampler + + def _require_open(self) -> None: + if self._closed: + raise RuntimeError("OmniDreams Ludus conditioning provider is closed.") + + +def keyboard_driving_user_input_schema() -> UserInputSchema: + return UserInputSchema( + capabilities=( + UserInputCapability( + event_type="keydown", + input_modality="keyboard", + payload_fields=frozenset({"key"}), + description="Keyboard key press edge.", + ), + UserInputCapability( + event_type="keyup", + input_modality="keyboard", + payload_fields=frozenset({"key"}), + description="Keyboard key release edge.", + ), + ), + description="Recorded or realtime WSAD keyboard driving controls.", + ) + + def precomputed_hdmap_inference_input_schema() -> InferenceInputSchema: return InferenceInputSchema( global_conditioning_fields=( @@ -181,7 +444,9 @@ def precomputed_hdmap_inference_input_schema() -> InferenceInputSchema: ) -def _scenario_from_prepared(scenario: PreparedScenario) -> OmnidreamsReplayScenario: +def _precomputed_scenario_from_prepared( + scenario: PreparedScenario, +) -> OmnidreamsReplayScenario: value = scenario.initial_inputs.global_conditioning.get("scenario") if not isinstance(value, OmnidreamsReplayScenario): raise TypeError( @@ -192,6 +457,223 @@ def _scenario_from_prepared(scenario: PreparedScenario) -> OmnidreamsReplayScena return value +def _ludus_scenario_from_prepared( + scenario: PreparedScenario, +) -> OmnidreamsLudusReplayScenario: + value = scenario.initial_inputs.global_conditioning.get("scenario") + if not isinstance(value, OmnidreamsLudusReplayScenario): + raise TypeError( + "OmniDreams Ludus conditioning provider requires " + "initial_inputs.global_conditioning['scenario'] to be an " + "OmnidreamsLudusReplayScenario." + ) + return value + + +def _resolve_ludus_scene_path(scenario: OmnidreamsLudusReplayScenario) -> Path: + if scenario.scene_path is not None: + if not scenario.scene_path.exists(): + raise FileNotFoundError( + f"OmniDreams Ludus scene_path missing: {scenario.scene_path}" + ) + return scenario.scene_path + if scenario.scene_dir is not None: + return _resolve_local_ludus_scene_path(scenario) + + from omnidreams.scenes import hf_hub_download_scene # noqa: PLC0415 + + return hf_hub_download_scene( + scenario.scene_uuid or DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, + scenario.scene_variant, + ) + + +def _resolve_local_ludus_scene_path(scenario: OmnidreamsLudusReplayScenario) -> Path: + scene_dir = scenario.scene_dir + if scene_dir is None: + raise RuntimeError("OmniDreams Ludus scene_dir is unexpectedly unset.") + if scene_dir.is_file(): + return scene_dir + if not scene_dir.is_dir(): + raise FileNotFoundError(f"OmniDreams Ludus scene_dir missing: {scene_dir}") + + candidates = _local_ludus_scene_candidates(scenario) + for candidate in candidates: + if candidate.is_file(): + return candidate + archives = sorted(scene_dir.glob("*.usdz")) + if scenario.scene_uuid is None and len(archives) == 1: + return archives[0] + expected = ", ".join(path.name for path in candidates) + raise FileNotFoundError( + f"No OmniDreams Ludus USDZ scene archive found in {scene_dir}. " + f"Expected one of: {expected}." + ) + + +def _local_ludus_scene_candidates( + scenario: OmnidreamsLudusReplayScenario, +) -> tuple[Path, ...]: + scene_dir = scenario.scene_dir + if scene_dir is None or scenario.scene_uuid is None: + return () + + from omnidreams.scenes import ( # noqa: PLC0415 + normalise_scene_uuid, + scene_variant_suffix, + ) + + bare_uuid = normalise_scene_uuid(scenario.scene_uuid) + suffix = scene_variant_suffix(scenario.scene_variant) + stems = [f"clipgt-{bare_uuid}{suffix}", f"{bare_uuid}{suffix}"] + if suffix: + stems.extend((f"clipgt-{bare_uuid}", bare_uuid)) + return tuple(scene_dir / f"{stem}.usdz" for stem in dict.fromkeys(stems)) + + +def _load_ludus_scene_bundle( + scenario: OmnidreamsLudusReplayScenario, + scene_path: Path, +) -> Any: + from omnidreams.interactive_drive.scene_loader import ( # noqa: PLC0415 + load_scene_bundle, + ) + + return load_scene_bundle( + scene_path=scene_path, + camera_name=scenario.camera_name, + variant=scenario.scene_variant, + prompt_override=scenario.prompt, + raster=_ludus_raster_config(scenario), + ) + + +def _new_ludus_rasterizer(scenario: OmnidreamsLudusReplayScenario) -> Any: + from omnidreams.interactive_drive.rasterizer import ( # noqa: PLC0415 + LudusConditionRasterizer, + ) + + return LudusConditionRasterizer(_ludus_raster_config(scenario), bev=None) + + +def _ludus_raster_config(scenario: OmnidreamsLudusReplayScenario) -> Any: + from omnidreams.interactive_drive.config import RasterConfig # noqa: PLC0415 + + return RasterConfig( + width=scenario.pixel_width, + height=scenario.pixel_height, + ludus_backend=scenario.ludus_backend, + ) + + +def _initial_rgb_tensor( + frame: object, + *, + device: torch.device, + dtype: torch.dtype, +) -> torch.Tensor: + tensor = torch.from_numpy(_rgb_hwc_uint8(frame)) + tensor = tensor.permute(2, 0, 1).unsqueeze(0).unsqueeze(0).unsqueeze(2) + return _to_model_range(tensor, device=device, dtype=dtype) + + +def _condition_frames_tensor( + frames: tuple[object, ...], + *, + device: torch.device, + dtype: torch.dtype, +) -> torch.Tensor: + cuda_video = _condition_cuda_video(frames) + if cuda_video is not None: + tensor = cuda_video.permute(0, 3, 1, 2).unsqueeze(0).unsqueeze(0) + return _to_model_range(tensor, device=device, dtype=dtype) + video = np.stack( + [_rgb_hwc_uint8(_frame_rgb(frame)) for frame in frames], + axis=0, + ) + tensor = torch.from_numpy(np.ascontiguousarray(video)) + tensor = tensor.permute(0, 3, 1, 2).unsqueeze(0).unsqueeze(0) + return _to_model_range(tensor, device=device, dtype=dtype) + + +def _condition_cuda_video(frames: tuple[object, ...]) -> torch.Tensor | None: + tensors: list[torch.Tensor] = [] + for frame in frames: + to_cuda_tensor = getattr(_frame_rgb(frame), "to_cuda_tensor", None) + if not callable(to_cuda_tensor): + return None + try: + tensor = to_cuda_tensor() + except RuntimeError: + return None + if ( + not torch.is_tensor(tensor) + or not tensor.is_cuda + or tensor.dtype != torch.uint8 + or tensor.ndim != 3 + or tensor.shape[-1] < 3 + ): + return None + tensors.append(tensor[..., :3]) + return torch.stack(tensors, dim=0) + + +def _frame_rgb(frame: object) -> object: + return getattr(frame, "rgb_host_uint8", frame) + + +def _rgb_hwc_uint8(frame: object) -> np.ndarray: + if torch.is_tensor(frame): + array = frame.detach().cpu().numpy() + else: + array = np.asarray(frame, dtype=np.uint8) + if array.ndim != 3 or array.shape[-1] < 3: + raise ValueError( + "OmniDreams Ludus rendered frames must be HWC RGB/RGBA uint8 arrays." + ) + return np.ascontiguousarray(np.array(array[..., :3], dtype=np.uint8, copy=True)) + + +def _to_model_range( + tensor: torch.Tensor, + *, + device: torch.device, + dtype: torch.dtype, +) -> torch.Tensor: + return tensor.to(device=device, dtype=dtype) / 127.5 - 1.0 + + +def _segments_metadata( + segments: list[PoseSegment], +) -> tuple[tuple[float, float, tuple[str, ...]], ...]: + return tuple( + (float(start), float(end), tuple(sorted(keys))) for start, end, keys in segments + ) + + +def _pose_segment(value: object) -> PoseSegment: + if not isinstance(value, tuple) or len(value) != 3: + raise RuntimeError("OmniDreams Ludus key segment must be a 3-tuple.") + start, end, keys = value + if not isinstance(start, int | float) or not isinstance(end, int | float): + raise RuntimeError("OmniDreams Ludus key segment bounds must be numeric.") + if not isinstance(keys, frozenset | set | tuple | list): + raise RuntimeError("OmniDreams Ludus key segment keys must be a sequence.") + return (float(start), float(end), frozenset(str(key) for key in keys)) + + +def _close_rasterizer(rasterizer: Any | None) -> None: + if rasterizer is None: + return + close = getattr(rasterizer, "cleanup", None) or getattr( + rasterizer, + "close", + None, + ) + if callable(close): + close() + + def _device_from_config(config: InferenceConfig) -> torch.device: if dist.is_initialized(): return torch.device(f"cuda:{int(os.environ.get('LOCAL_RANK', '0'))}") @@ -203,6 +685,8 @@ def _is_rank_zero() -> bool: __all__ = [ + "LudusSceneConditioningProvider", "PrecomputedHDMapProvider", + "keyboard_driving_user_input_schema", "precomputed_hdmap_inference_input_schema", ] diff --git a/integrations/omnidreams/omnidreams/demo/replay.py b/integrations/omnidreams/omnidreams/demo/replay.py index 42871a5c6..06c9cda0d 100644 --- a/integrations/omnidreams/omnidreams/demo/replay.py +++ b/integrations/omnidreams/omnidreams/demo/replay.py @@ -27,7 +27,9 @@ from flashdreams.runtime.interfaces import InferenceSession from flashdreams.runtime.types import StepRequest, StepRequirements, StepResult -from .spec import OmnidreamsReplayScenario +from .spec import OmnidreamsLudusReplayScenario, OmnidreamsReplayScenario + +OmnidreamsSessionScenario = OmnidreamsReplayScenario | OmnidreamsLudusReplayScenario PipelineFactory = Callable[[Any, str], Any] @@ -81,6 +83,7 @@ def start_session(self, inputs: InferenceInput) -> InferenceSession: else torch.device(self.config.device or "cuda"), is_rank_zero=self.is_rank_zero, output_layout=self.options.output_layout, + rollout_seed=self.config.seed, ) def close(self) -> None: @@ -102,11 +105,12 @@ def __init__( self, *, pipeline: Any, - scenario: OmnidreamsReplayScenario, + scenario: OmnidreamsSessionScenario, initial_inputs: InferenceInput, device: torch.device, is_rank_zero: bool, output_layout: VideoTensorLayout, + rollout_seed: int | None, ) -> None: self.pipeline = pipeline self.scenario = scenario @@ -114,6 +118,7 @@ def __init__( self.device = device self.is_rank_zero = is_rank_zero self.output_layout = output_layout + self.rollout_seed = rollout_seed self.dtype = torch.bfloat16 self._closed = False self._model_session = OmnidreamsModelSessionCore( @@ -190,6 +195,7 @@ def close(self) -> None: def _initialize_cache(self) -> Any: scenario = self.scenario + _seed_pipeline_for_rollout(self.pipeline, self.rollout_seed) cache = self.pipeline.initialize_cache( text=_prompt_from_inputs(self._initial_inputs, scenario), image=_first_frame_from_inputs( @@ -210,22 +216,30 @@ def _default_pipeline_factory(pipeline_config: Any, device: str) -> Any: return pipeline_config.setup().to(device=device).eval() -def _scenario_from_inputs(inputs: InferenceInput) -> OmnidreamsReplayScenario: +def _scenario_from_inputs(inputs: InferenceInput) -> OmnidreamsSessionScenario: scenario = inputs.global_conditioning.get("scenario") - if not isinstance(scenario, OmnidreamsReplayScenario): + if not isinstance( + scenario, + (OmnidreamsReplayScenario, OmnidreamsLudusReplayScenario), + ): raise TypeError( "OmniDreams replay runtime requires global_conditioning['scenario'] " - "to be an OmnidreamsReplayScenario." + "to be an OmnidreamsReplayScenario or OmnidreamsLudusReplayScenario." ) return scenario def _prompt_from_inputs( inputs: InferenceInput, - scenario: OmnidreamsReplayScenario, + scenario: OmnidreamsSessionScenario, ) -> list[list[str]]: prompt = inputs.global_conditioning.get("prompt") if prompt is None: + if not scenario.prompts: + raise ValueError( + "OmniDreams initial prompt is required when the scenario does " + "not carry fallback prompts." + ) return [list(scenario.prompts)] if isinstance(prompt, str): return [[prompt]] @@ -249,7 +263,7 @@ def _prompt_from_inputs( def _first_frame_from_inputs( inputs: InferenceInput, *, - scenario: OmnidreamsReplayScenario, + scenario: OmnidreamsSessionScenario, device: torch.device, dtype: torch.dtype, ) -> torch.Tensor: @@ -258,6 +272,12 @@ def _first_frame_from_inputs( return first_frame if first_frame is not None: raise TypeError("OmniDreams initial first_frame must be a torch.Tensor.") + first_frame_paths = getattr(scenario, "first_frame_paths", ()) + if not first_frame_paths: + raise ValueError( + "OmniDreams initial first_frame tensor is required when the " + "scenario does not carry fallback first_frame_paths." + ) first_frames = [ load_first_frame_tensor( path, @@ -268,14 +288,24 @@ def _first_frame_from_inputs( allow_video=True, install_hint=DEFAULT_RUNNER_INSTALL_HINT, ) - for path in scenario.first_frame_paths + for path in first_frame_paths ] return torch.stack(first_frames, dim=0).unsqueeze(0) +def _seed_pipeline_for_rollout(pipeline: Any, seed: int | None) -> None: + if seed is None: + return + diffusion_model = getattr(pipeline, "diffusion_model", None) + rng = getattr(diffusion_model, "rng", None) + if rng is None: + return + rng.manual_seed(int(seed)) + + def _view_names_from_inputs( inputs: InferenceInput, - scenario: OmnidreamsReplayScenario, + scenario: OmnidreamsSessionScenario, ) -> list[str]: value = inputs.metadata.get("view_names") or inputs.global_conditioning.get( "view_names" @@ -314,6 +344,7 @@ def _is_torchrun_env() -> bool: "OmnidreamsRuntime", "OmnidreamsRuntimeOptions", "OmnidreamsSession", + "OmnidreamsSessionScenario", "OmnidreamsReplayRuntime", "OmnidreamsReplayRuntimeOptions", "OmnidreamsReplaySession", diff --git a/integrations/omnidreams/omnidreams/demo/spec.py b/integrations/omnidreams/omnidreams/demo/spec.py index 6a4f147cf..8df653ecd 100644 --- a/integrations/omnidreams/omnidreams/demo/spec.py +++ b/integrations/omnidreams/omnidreams/demo/spec.py @@ -5,10 +5,12 @@ from __future__ import annotations +import json +import math from collections.abc import Mapping, Sequence from dataclasses import dataclass from pathlib import Path -from typing import Any +from typing import Any, Literal, TypeAlias, cast from omnidreams.runner import ( DEFAULT_EXAMPLE_DATA_UUID_1V, @@ -22,6 +24,43 @@ DEFAULT_OMNIDREAMS_PRESET = "omnidreams-sv-2steps-chunk2-loc6-lightvae-lighttae" OMNIDREAMS_MODEL_ID = "omnidreams" DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID = "0d404ff7-2b66-498c-b047-1ed8cded60d4" +OMNIDREAMS_CONDITIONING_PRECOMPUTED = "precomputed-hdmap" +OMNIDREAMS_CONDITIONING_LUDUS = "ludus-scene-driving" +OMNIDREAMS_CONDITIONING_MODES = ( + OMNIDREAMS_CONDITIONING_PRECOMPUTED, + OMNIDREAMS_CONDITIONING_LUDUS, +) +LudusBackendName: TypeAlias = Literal["cuda", "vulkan"] + +_KEY_EVENT_ALIASES = { + "down": "keydown", + "key_down": "keydown", + "keyup": "keyup", + "up": "keyup", + "key_up": "keyup", +} + + +@dataclass(frozen=True, kw_only=True, slots=True) +class OmnidreamsKeyboardTraceEvent: + """One recorded keyboard edge in a finite Ludus replay trace.""" + + timestamp_s: float + event: str + key: str + + def __post_init__(self) -> None: + timestamp_s = float(self.timestamp_s) + if not math.isfinite(timestamp_s) or timestamp_s < 0: + raise ValueError( + "OmnidreamsKeyboardTraceEvent.timestamp_s must be finite and >= 0." + ) + key = str(self.key).strip().lower() + if not key: + raise ValueError("OmnidreamsKeyboardTraceEvent.key must be non-empty.") + object.__setattr__(self, "timestamp_s", timestamp_s) + object.__setattr__(self, "event", _normalize_key_event_name(self.event)) + object.__setattr__(self, "key", key) @dataclass(frozen=True, kw_only=True, slots=True) @@ -69,6 +108,86 @@ def __post_init__(self) -> None: ) +@dataclass(frozen=True, kw_only=True, slots=True) +class OmnidreamsLudusReplayScenario: + """Resolved Ludus scene plus a finite recorded keyboard trace.""" + + keyboard_events: tuple[OmnidreamsKeyboardTraceEvent, ...] + scene_path: Path | None = None + scene_dir: Path | None = None + scene_uuid: str | None = DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID + scene_variant: str = SCENE_VARIANT_DEFAULT + camera_name: str = "camera_front_wide_120fov" + prompt: str | None = None + total_blocks: int = 60 + pixel_height: int = DEFAULT_VIDEO_HEIGHT + pixel_width: int = DEFAULT_VIDEO_WIDTH + fps: int = 30 + move_speed_per_s: float = 6.0 + rotate_speed_rad_per_s: float = math.radians(35.0) + ludus_backend: LudusBackendName = "cuda" + + @property + def camera_names(self) -> tuple[str, ...]: + return (self.camera_name,) + + @property + def prompts(self) -> tuple[str, ...]: + return () if self.prompt is None else (self.prompt,) + + def __post_init__(self) -> None: + if self.scene_path is not None: + object.__setattr__(self, "scene_path", Path(self.scene_path)) + if self.scene_dir is not None: + object.__setattr__(self, "scene_dir", Path(self.scene_dir)) + if not (self.scene_path or self.scene_dir or self.scene_uuid): + raise ValueError( + "OmnidreamsLudusReplayScenario requires scene_path, " + "scene_dir, or scene_uuid." + ) + if not self.scene_variant.strip(): + raise ValueError("OmnidreamsLudusReplayScenario.scene_variant is required.") + if not self.camera_name.strip(): + raise ValueError("OmnidreamsLudusReplayScenario.camera_name is required.") + if self.total_blocks <= 0: + raise ValueError("OmnidreamsLudusReplayScenario.total_blocks must be > 0.") + if self.pixel_height <= 0 or self.pixel_width <= 0: + raise ValueError( + "OmnidreamsLudusReplayScenario pixel dimensions must be > 0." + ) + if self.fps <= 0: + raise ValueError("OmnidreamsLudusReplayScenario.fps must be > 0.") + if self.move_speed_per_s <= 0: + raise ValueError( + "OmnidreamsLudusReplayScenario.move_speed_per_s must be > 0." + ) + if self.rotate_speed_rad_per_s <= 0: + raise ValueError( + "OmnidreamsLudusReplayScenario.rotate_speed_rad_per_s must be > 0." + ) + object.__setattr__( + self, + "ludus_backend", + _ludus_backend_name(self.ludus_backend), + ) + previous_timestamp_s = -math.inf + normalized_events: list[OmnidreamsKeyboardTraceEvent] = [] + for event in self.keyboard_events: + normalized = ( + event + if isinstance(event, OmnidreamsKeyboardTraceEvent) + else _keyboard_trace_event(event) + ) + if normalized.timestamp_s < previous_timestamp_s: + raise ValueError( + "OmnidreamsLudusReplayScenario.keyboard_events must be sorted " + "by non-decreasing timestamp_s." + ) + previous_timestamp_s = normalized.timestamp_s + normalized_events.append(normalized) + object.__setattr__(self, "keyboard_events", tuple(normalized_events)) + + @dataclass(frozen=True, kw_only=True, slots=True) class OmnidreamsWebRTCScenario: """Scene/options for the shared WebRTC demo path.""" @@ -89,6 +208,33 @@ def __post_init__(self) -> None: raise ValueError("OmnidreamsWebRTCScenario.camera_name is required.") +def conditioning_mode_from_scenario(value: Any) -> str: + """Return the resolved OmniDreams replay conditioning mode.""" + if isinstance(value, OmnidreamsLudusReplayScenario): + return OMNIDREAMS_CONDITIONING_LUDUS + if isinstance(value, OmnidreamsReplayScenario): + return OMNIDREAMS_CONDITIONING_PRECOMPUTED + if value is None or not isinstance(value, Mapping): + return OMNIDREAMS_CONDITIONING_PRECOMPUTED + + mode = ( + str(value.get("conditioning_mode", OMNIDREAMS_CONDITIONING_PRECOMPUTED)) + .strip() + .lower() + ) + if mode in {"precomputed", "hdmap", "precomputed-hdmaps"}: + mode = OMNIDREAMS_CONDITIONING_PRECOMPUTED + if mode in {"ludus", "keyboard-driving", "ludus-keyboard"}: + mode = OMNIDREAMS_CONDITIONING_LUDUS + if mode not in OMNIDREAMS_CONDITIONING_MODES: + supported = ", ".join(OMNIDREAMS_CONDITIONING_MODES) + raise ValueError( + f"Unsupported OmniDreams conditioning_mode={mode!r}. " + f"Supported modes: {supported}." + ) + return mode + + def resolve_replay_scenario( value: Any, *, @@ -151,6 +297,40 @@ def resolve_replay_scenario( ) +def resolve_ludus_replay_scenario(value: Any) -> OmnidreamsLudusReplayScenario: + """Normalize a user/demo scenario into a Ludus recorded-trace scenario.""" + if isinstance(value, OmnidreamsLudusReplayScenario): + _require_optional_existing_path(value.scene_path, label="scene_path") + return value + if value is None: + value = {} + if not isinstance(value, Mapping): + raise TypeError( + "OmniDreams Ludus replay scenario must be an " + "OmnidreamsLudusReplayScenario, a mapping, or None." + ) + return OmnidreamsLudusReplayScenario( + keyboard_events=_keyboard_trace_events(value), + scene_path=_optional_path(value.get("scene_path")), + scene_dir=_optional_path(value.get("scene_dir")), + scene_uuid=_optional_string( + value.get("scene_uuid", DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID) + ), + scene_variant=str(value.get("scene_variant", SCENE_VARIANT_DEFAULT)), + camera_name=str(value.get("camera_name", "camera_front_wide_120fov")), + prompt=_optional_string(value.get("prompt")), + total_blocks=int(value.get("total_blocks", 60)), + pixel_height=int(value.get("pixel_height", DEFAULT_VIDEO_HEIGHT)), + pixel_width=int(value.get("pixel_width", DEFAULT_VIDEO_WIDTH)), + fps=int(value.get("fps", 30)), + move_speed_per_s=float(value.get("move_speed_per_s", 6.0)), + rotate_speed_rad_per_s=float( + value.get("rotate_speed_rad_per_s", math.radians(35.0)) + ), + ludus_backend=_ludus_backend_name(value.get("ludus_backend", "cuda")), + ) + + def resolve_webrtc_scenario(value: Any) -> OmnidreamsWebRTCScenario: """Normalize a user/demo scenario into a WebRTC scenario.""" if value is None: @@ -205,6 +385,68 @@ def _resolve_example_data_default(value: Mapping[str, Any]) -> bool: ) +def _keyboard_trace_events( + value: Mapping[str, Any], +) -> tuple[OmnidreamsKeyboardTraceEvent, ...]: + events_value = value.get("keyboard_events") + if events_value is None: + trace_path = _optional_path(value.get("keyboard_trace_path")) + if trace_path is None: + return () + if not trace_path.exists(): + raise FileNotFoundError( + f"OmniDreams keyboard_trace_path missing: {trace_path}" + ) + loaded = json.loads(trace_path.read_text(encoding="utf-8")) + events_value = ( + loaded.get("events", ()) if isinstance(loaded, Mapping) else loaded + ) + if isinstance(events_value, (str, bytes)) or not isinstance(events_value, Sequence): + raise TypeError("OmniDreams keyboard trace must be a sequence of events.") + return tuple(_keyboard_trace_event(event) for event in events_value) + + +def _keyboard_trace_event(value: Any) -> OmnidreamsKeyboardTraceEvent: + if isinstance(value, OmnidreamsKeyboardTraceEvent): + return value + if not isinstance(value, Mapping): + raise TypeError( + "OmniDreams keyboard trace events must be mappings or " + "OmnidreamsKeyboardTraceEvent instances." + ) + timestamp = _first_present(value, ("timestamp_s", "time_s", "timestamp", "t")) + if timestamp is None: + raise ValueError("OmniDreams keyboard trace event missing timestamp_s.") + event = _first_present(value, ("event", "event_type", "type")) + if event is None: + raise ValueError("OmniDreams keyboard trace event missing event.") + key = value.get("key") + if key is None: + raise ValueError("OmniDreams keyboard trace event missing key.") + return OmnidreamsKeyboardTraceEvent( + timestamp_s=float(timestamp), + event=str(event), + key=str(key), + ) + + +def _first_present(value: Mapping[str, Any], keys: tuple[str, ...]) -> Any: + for key in keys: + if key in value: + return value[key] + return None + + +def _normalize_key_event_name(value: str) -> str: + event = str(value).strip().lower() + event = _KEY_EVENT_ALIASES.get(event, event) + if event not in {"keydown", "keyup"}: + raise ValueError( + "OmniDreams keyboard trace event must be 'keydown' or 'keyup'." + ) + return event + + def _bool_value(value: Any) -> bool: if isinstance(value, bool): return value @@ -248,6 +490,37 @@ def _string_tuple(value: Any) -> tuple[str, ...]: raise TypeError(f"Expected string or string sequence, got {type(value).__name__}.") +def _optional_path(value: Any) -> Path | None: + if value is None or value == "": + return None + return Path(value) + + +def _optional_string(value: Any) -> str | None: + if value is None: + return None + text = str(value).strip() + return text or None + + +def _ludus_backend_name(value: Any) -> LudusBackendName: + backend = str(value).strip().lower() + if backend not in {"cuda", "vulkan"}: + raise ValueError( + "OmnidreamsLudusReplayScenario.ludus_backend must be 'cuda' or 'vulkan'." + ) + return cast(LudusBackendName, backend) + + +def _require_optional_existing_path(path: Path | None, *, label: str) -> None: + if path is None: + return + if not path.exists(): + raise FileNotFoundError( + f"OmniDreams Ludus replay scenario missing {label}: {path}" + ) + + def _require_existing_paths(paths: tuple[Path, ...], *, label: str) -> None: if not paths: raise ValueError(f"OmniDreams replay scenario requires {label}.") @@ -262,9 +535,16 @@ def _require_existing_paths(paths: tuple[Path, ...], *, label: str) -> None: __all__ = [ "DEFAULT_OMNIDREAMS_PRESET", "DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID", + "OMNIDREAMS_CONDITIONING_LUDUS", + "OMNIDREAMS_CONDITIONING_MODES", + "OMNIDREAMS_CONDITIONING_PRECOMPUTED", "OMNIDREAMS_MODEL_ID", + "OmnidreamsKeyboardTraceEvent", + "OmnidreamsLudusReplayScenario", "OmnidreamsReplayScenario", "OmnidreamsWebRTCScenario", + "conditioning_mode_from_scenario", + "resolve_ludus_replay_scenario", "resolve_replay_scenario", "resolve_webrtc_scenario", ] diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index 21c6e70ee..9e7967926 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -3,11 +3,13 @@ from __future__ import annotations +import json from collections.abc import Sequence from pathlib import Path from types import SimpleNamespace from typing import Any +import numpy as np import omnidreams.demo as demo_package import omnidreams.demo.spec as spec_module import pytest @@ -16,8 +18,12 @@ from omnidreams.config import OMNIDREAMS_RUNNERS from omnidreams.demo import ( DEFAULT_OMNIDREAMS_PRESET, + OMNIDREAMS_CONDITIONING_LUDUS, + OMNIDREAMS_CONDITIONING_PRECOMPUTED, OMNIDREAMS_MODEL_ID, + LudusSceneConditioningProvider, OmnidreamsDemoAdapter, + OmnidreamsLudusReplayScenario, OmnidreamsReplayScenario, OmnidreamsWebRTCScenario, PrecomputedHDMapProvider, @@ -58,6 +64,7 @@ WebRTCOutputSpec, ) from flashdreams.runtime.demo.replay import run_replay_demo +from flashdreams.runtime.demo.timing import SPARSE_KEY_SEGMENTS_METADATA_KEY from flashdreams.serving.webrtc.manager import BaseWebRTCSessionManager from flashdreams.serving.webrtc.server import SESSION_MANAGER_KEY @@ -82,12 +89,72 @@ def test_omnidreams_replay_cli_builds_null_output_spec() -> None: assert spec.config.model_id == OMNIDREAMS_MODEL_ID +def test_omnidreams_replay_cli_builds_ludus_conditioning_spec( + tmp_path: Path, +) -> None: + trace_path = tmp_path / "trace.json" + trace_path.write_text( + json.dumps( + { + "events": [ + {"timestamp_s": 0.0, "event": "keydown", "key": "w"}, + {"timestamp_s": 0.5, "event": "keyup", "key": "w"}, + ] + } + ), + encoding="utf-8", + ) + output_path = tmp_path / "demo.mp4" + + args = parse_args( + [ + "replay", + "--conditioning-mode", + OMNIDREAMS_CONDITIONING_LUDUS, + "--keyboard-trace", + str(trace_path), + "--scene-uuid", + "scene-1", + "--scene-variant", + "rain", + "--camera-name", + "camera_front_wide_120fov", + "--seed", + "123", + "--total-blocks", + "3", + "--output", + str(output_path), + ] + ) + + spec = _replay_spec(args) + + assert spec.input_mode == "replay" + assert isinstance(spec.scenario, dict) + scenario = spec.scenario + assert scenario["conditioning_mode"] == OMNIDREAMS_CONDITIONING_LUDUS + assert scenario["keyboard_trace_path"] == trace_path + assert scenario["scene_uuid"] == "scene-1" + assert scenario["scene_variant"] == "rain" + assert scenario["total_blocks"] == 3 + assert isinstance(spec.output, Mp4OutputSpec) + assert spec.output.path == output_path + assert spec.config is not None + assert spec.config.seed == 123 + assert spec.config.runtime_options["seed"] == 123 + + def test_omnidreams_demo_adapter_declares_replay_modes_only() -> None: adapter = OmnidreamsDemoAdapter() assert adapter.model_id == OMNIDREAMS_MODEL_ID assert adapter.supported_input_modes() == ("replay",) assert adapter.supported_output_modes() == ("mp4", "null") + assert adapter.supported_conditioning_modes() == ( + OMNIDREAMS_CONDITIONING_PRECOMPUTED, + OMNIDREAMS_CONDITIONING_LUDUS, + ) assert [ field.name for field in adapter.inference_input_schema.global_conditioning_fields @@ -365,6 +432,81 @@ def test_omnidreams_precomputed_hdmap_provider_prepares_inputs( provider.close() +def test_omnidreams_ludus_provider_prepares_deterministic_hdmaps( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _scene, rasterizers = _install_fake_ludus_provider_dependencies(monkeypatch) + scene_path = tmp_path / "scene.usdz" + scene_path.write_bytes(b"fake") + adapter = OmnidreamsDemoAdapter() + spec = _ludus_replay_demo_spec( + tmp_path=tmp_path, + scene_path=scene_path, + total_blocks=2, + ) + prepared = adapter.prepare_scenario(spec) + + provider = adapter.create_model_input_provider(spec, prepared) + + assert isinstance(provider, LudusSceneConditioningProvider) + initial = provider.prepare_initial_input() + scenario = initial.global_conditioning["scenario"] + assert isinstance(scenario, OmnidreamsLudusReplayScenario) + assert scenario.camera_names == ("camera_front_wide_120fov",) + assert initial.global_conditioning["prompt"] == [["city scene"]] + assert initial.global_conditioning["first_frame"].shape == (1, 1, 1, 3, 2, 2) + assert initial.metadata["view_names"] == ("camera_front_wide_120fov",) + + first = provider.prepare_step( + request=StepRequirements(step_index=0, input_frame_count=2), + user_window=UserInputWindow(start_s=0.0, end_s=2 / 30), + ) + + assert first.inference_input is not None + first_hdmap = first.inference_input.step["hdmap"] + assert isinstance(first_hdmap, torch.Tensor) + assert first_hdmap.shape == (1, 1, 2, 3, 2, 2) + assert first.inference_input.metadata["frame_timestamps_us"] == (1_000, 34_333) + assert first.inference_input.metadata["keyboard_segments"] == ( + (0.0, 2 / 30, ("w",)), + ) + assert len(rasterizers) == 1 + assert rasterizers[0].calls[0]["timestamps_us"] == (1_000, 34_333) + assert rasterizers[0].calls[0]["rig_poses_world"].shape == (2, 4, 4) + assert rasterizers[0].calls[0]["rig_poses_world"][0, 0, 3] > 0 + + provider.reset() + reset_first = provider.prepare_step( + request=StepRequirements(step_index=0, input_frame_count=2), + user_window=UserInputWindow(start_s=0.0, end_s=2 / 30), + ) + + assert reset_first.inference_input is not None + torch.testing.assert_close(reset_first.inference_input.step["hdmap"], first_hdmap) + + provider.reset() + realtime_first = provider.prepare_step( + request=StepRequirements(step_index=0, input_frame_count=2), + user_window=UserInputWindow( + start_s=0.0, + end_s=2 / 30, + frame_times=(1 / 30, 2 / 30), + metadata={ + SPARSE_KEY_SEGMENTS_METADATA_KEY: ((0.0, 2 / 30, frozenset({"w"})),) + }, + ), + ) + + assert realtime_first.inference_input is not None + torch.testing.assert_close( + realtime_first.inference_input.step["hdmap"], + first_hdmap, + ) + provider.close() + assert rasterizers[0].closed is True + + def test_omnidreams_replay_run_mode_uses_precomputed_provider( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -421,6 +563,48 @@ def test_omnidreams_replay_run_mode_uses_precomputed_provider( torch.testing.assert_close(pipeline.generated_hdmaps[1][0, 0], loaded_hdmap[1:2]) +def test_omnidreams_replay_run_mode_uses_ludus_provider( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fake_ludus_provider_dependencies(monkeypatch) + scene_path = tmp_path / "scene.usdz" + scene_path.write_bytes(b"fake") + pipeline = _FakeOmnidreamsPipeline() + sink = _RecordingOutputSink() + adapter = OmnidreamsDemoAdapter( + pipeline_factory=lambda pipeline_config, device: pipeline, + ) + spec = _ludus_replay_demo_spec( + tmp_path=tmp_path, + scene_path=scene_path, + total_blocks=2, + ) + + result = run_replay_demo( + spec=spec, + adapter=adapter, + output_sink_factory=lambda output_spec: sink, + ) + + assert result.status == "completed" + assert result.artifacts == ( + OutputArtifact(kind="video/mp4", uri="memory://omnidreams"), + ) + assert [result.step_index for result in sink.results] == [0, 1] + assert pipeline.initialize_cache_calls == [ + { + "text": [["city scene"]], + "image_shape": (1, 1, 1, 3, 2, 2), + "view_names": ["camera_front_wide_120fov"], + } + ] + assert len(pipeline.generated_hdmaps) == 2 + assert pipeline.generated_hdmaps[0].shape == (1, 1, 1, 3, 2, 2) + assert pipeline.generated_hdmaps[1].shape == (1, 1, 1, 3, 2, 2) + assert not torch.equal(pipeline.generated_hdmaps[0], pipeline.generated_hdmaps[1]) + + def test_omnidreams_replay_null_output_uses_precomputed_provider( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, @@ -897,6 +1081,79 @@ def _replay_demo_spec( ) +def _ludus_replay_demo_spec( + *, + tmp_path: Path, + scene_path: Path, + total_blocks: int, + output: Mp4OutputSpec | NullOutputSpec | None = None, +) -> DemoSpec: + return DemoSpec( + model_id=OMNIDREAMS_MODEL_ID, + preset_id=DEFAULT_OMNIDREAMS_PRESET, + input_mode="replay", + scenario={ + "conditioning_mode": OMNIDREAMS_CONDITIONING_LUDUS, + "keyboard_events": ( + {"timestamp_s": 0.0, "event": "keydown", "key": "w"}, + {"timestamp_s": 0.5, "event": "keyup", "key": "w"}, + ), + "scene_path": scene_path, + "scene_variant": "default", + "camera_name": "camera_front_wide_120fov", + "total_blocks": total_blocks, + "pixel_height": 2, + "pixel_width": 2, + "fps": 30, + }, + output=output or Mp4OutputSpec(path=tmp_path / "demo.mp4", fps=30), + config=InferenceConfig( + model_id=OMNIDREAMS_MODEL_ID, + preset_id=DEFAULT_OMNIDREAMS_PRESET, + device="cpu", + seed=123, + runtime_options={"pipeline_config": object(), "seed": 123}, + ), + ) + + +def _install_fake_ludus_provider_dependencies( + monkeypatch: pytest.MonkeyPatch, +) -> tuple[SimpleNamespace, list["_FakeLudusRasterizer"]]: + import omnidreams.demo.providers as providers_module + + scene = SimpleNamespace( + scene_id="fake-scene", + prompt="city scene", + initial_rgb=np.zeros((2, 2, 3), dtype=np.uint8), + initial_rig_to_world=np.eye(4, dtype=np.float32), + initial_timestamp_us=1_000, + ) + rasterizers: list[_FakeLudusRasterizer] = [] + + def fake_load_scene_bundle(*args: Any, **kwargs: Any) -> SimpleNamespace: + del args, kwargs + return scene + + def fake_new_rasterizer(*args: Any, **kwargs: Any) -> "_FakeLudusRasterizer": + del args, kwargs + rasterizer = _FakeLudusRasterizer() + rasterizers.append(rasterizer) + return rasterizer + + monkeypatch.setattr( + providers_module, + "_load_ludus_scene_bundle", + fake_load_scene_bundle, + ) + monkeypatch.setattr( + providers_module, + "_new_ludus_rasterizer", + fake_new_rasterizer, + ) + return scene, rasterizers + + class _RecordingOutputSink: produces_artifacts = True @@ -982,6 +1239,41 @@ def cleanup(self) -> None: self.closed = True +class _FakeLudusRasterizer: + def __init__(self) -> None: + self.loaded_scene: object | None = None + self.calls: list[dict[str, Any]] = [] + self.closed = False + + def load_scene(self, scene: object) -> None: + self.loaded_scene = scene + + def render_chunk( + self, + *, + rig_poses_world: np.ndarray, + timestamps_us: np.ndarray, + ) -> SimpleNamespace: + self.calls.append( + { + "rig_poses_world": np.array(rig_poses_world, copy=True), + "timestamps_us": tuple(int(t) for t in timestamps_us), + } + ) + frames = [] + for timestamp_us in timestamps_us: + value = int(timestamp_us % 251) + frames.append( + SimpleNamespace( + rgb_host_uint8=np.full((2, 2, 3), value, dtype=np.uint8) + ) + ) + return SimpleNamespace(frames=tuple(frames)) + + def cleanup(self) -> None: + self.closed = True + + class _FakeConditioningWrapper: initial_frame_chunk_size = 2 frame_chunk_size = 3 From 26735b59042a688e1ab5083a018c290655bd6119 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 00:43:43 +0000 Subject: [PATCH 37/51] Add OmniDreams Ludus keyboard trace --- .../demo/traces/ludus_forward_sweep_60s.json | 86 +++++++++++++++++++ 1 file changed, 86 insertions(+) create mode 100644 integrations/omnidreams/omnidreams/demo/traces/ludus_forward_sweep_60s.json diff --git a/integrations/omnidreams/omnidreams/demo/traces/ludus_forward_sweep_60s.json b/integrations/omnidreams/omnidreams/demo/traces/ludus_forward_sweep_60s.json new file mode 100644 index 000000000..0f66222b6 --- /dev/null +++ b/integrations/omnidreams/omnidreams/demo/traces/ludus_forward_sweep_60s.json @@ -0,0 +1,86 @@ +{ + "name": "ludus_forward_sweep_60s", + "description": "Deterministic WSAD trace for OmniDreams Ludus replay MP4 validation.", + "events": [ + { + "timestamp_s": 0.0, + "event": "keydown", + "key": "w" + }, + { + "timestamp_s": 6.0, + "event": "keydown", + "key": "d" + }, + { + "timestamp_s": 10.0, + "event": "keyup", + "key": "d" + }, + { + "timestamp_s": 14.0, + "event": "keydown", + "key": "a" + }, + { + "timestamp_s": 18.0, + "event": "keyup", + "key": "a" + }, + { + "timestamp_s": 24.0, + "event": "keyup", + "key": "w" + }, + { + "timestamp_s": 25.0, + "event": "keydown", + "key": "w" + }, + { + "timestamp_s": 30.0, + "event": "keydown", + "key": "d" + }, + { + "timestamp_s": 34.0, + "event": "keyup", + "key": "d" + }, + { + "timestamp_s": 38.0, + "event": "keydown", + "key": "a" + }, + { + "timestamp_s": 42.0, + "event": "keyup", + "key": "a" + }, + { + "timestamp_s": 48.0, + "event": "keyup", + "key": "w" + }, + { + "timestamp_s": 50.0, + "event": "keydown", + "key": "w" + }, + { + "timestamp_s": 55.0, + "event": "keydown", + "key": "d" + }, + { + "timestamp_s": 58.0, + "event": "keyup", + "key": "d" + }, + { + "timestamp_s": 60.0, + "event": "keyup", + "key": "w" + } + ] +} From 758eabe51a3e0f49ba04b375a30f43ef84194efa Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 01:08:59 +0000 Subject: [PATCH 38/51] Mark runtime host unhealthy on cleanup close failures --- .../flashdreams/runtime/demo/drivers.py | 31 ++++++++++++------- .../test_demo_runtime_realtime_driver.py | 4 ++- 2 files changed, 22 insertions(+), 13 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 75f86aab0..b7dd0b325 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -33,6 +33,8 @@ from .validation import resolve_run_capabilities, validate_resolved_run CLEANUP_TIMEOUT_S = 30.0 +_MODEL_CLEANUP_FAILED_REASON = "model-affine cleanup failed" +_MODEL_CLEANUP_TIMED_OUT_REASON = "model-affine cleanup timed out" class DriverInvariantError(RuntimeError): @@ -607,11 +609,13 @@ def _coerce_step_requirements(value: object) -> StepRequirements | None: ) -def _close_safely(close: Any, session_edges: SessionEdges) -> None: +def _close_safely(close: Any, session_edges: SessionEdges) -> bool: try: close() except Exception as exc: session_edges.record_cleanup_error(exc) + return False + return True def _close_on_host_best_effort( @@ -647,15 +651,15 @@ async def shielded_session_cleanup( return session_edges.close_result(status=status, reason=reason, error=error) async def cleanup() -> RunResult: - resources_closed = await _close_model_resources_async( + unhealthy_reason = await _close_model_resources_async( host=host, session=session, provider=provider, session_edges=session_edges, timeout_s=timeout_s, ) - if not resources_closed: - host.mark_unhealthy("model-affine cleanup timed out") + if unhealthy_reason is not None: + host.mark_unhealthy(unhealthy_reason) return session_edges.close_result( status=status, reason=reason, @@ -851,9 +855,9 @@ async def _close_model_resources_async( provider: ModelInputProvider, session_edges: SessionEdges, timeout_s: float, -) -> bool: +) -> str | None: try: - await asyncio.wait_for( + resources_closed = await asyncio.wait_for( host.call_async( _close_model_resources_safely, session.close if session is not None else None, @@ -868,21 +872,24 @@ async def _close_model_resources_async( # provider cleanup on another thread or replacing the worker is unsafe. # The caller marks the host unhealthy so future sessions reject instead. session_edges.record_orphaned_cleanup(exc) - return False + return _MODEL_CLEANUP_TIMED_OUT_REASON except Exception as exc: session_edges.record_cleanup_error(exc) - return False - return True + return _MODEL_CLEANUP_FAILED_REASON + if not resources_closed: + return _MODEL_CLEANUP_FAILED_REASON + return None def _close_model_resources_safely( session_close: Any | None, provider_close: Any, session_edges: SessionEdges, -) -> None: +) -> bool: + resources_closed = True if session_close is not None: - _close_safely(session_close, session_edges) - _close_safely(provider_close, session_edges) + resources_closed = _close_safely(session_close, session_edges) + return _close_safely(provider_close, session_edges) and resources_closed def _cleanup_result( diff --git a/flashdreams/tests/test_demo_runtime_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py index a9ccea100..5c79b83f7 100644 --- a/flashdreams/tests/test_demo_runtime_realtime_driver.py +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -309,6 +309,8 @@ async def test_shielded_cleanup_never_raises_and_returns_result_on_close_errors( reason="test failure", error=RuntimeError("original"), ) + assert not host.is_healthy + assert host.unhealthy_reason == "model-affine cleanup failed" finally: host.close() @@ -366,7 +368,7 @@ async def test_shielded_cleanup_dispatch_failure_marks_host_unhealthy() -> None: ) assert result.status == "cancelled" - assert host.unhealthy_reason == "model-affine cleanup timed out" + assert host.unhealthy_reason == "model-affine cleanup failed" assert host.cleanup_dispatch_count == 1 assert session.close_count == 0 assert provider.close_count == 0 From 84b1a6cc21abc995c15f2e0550e93830aa724d1c Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 01:25:05 +0000 Subject: [PATCH 39/51] Quarantine runtime host on cleanup failures --- .../flashdreams/runtime/demo/drivers.py | 75 +++++++++++++++---- .../tests/test_demo_runtime_vertical_slice.py | 35 +++++++++ 2 files changed, 94 insertions(+), 16 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index b7dd0b325..1849cddef 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -321,6 +321,10 @@ async def run_one_session( return result +def _mark_host_cleanup_failed(host: RuntimeHost, exc: Exception | None = None) -> None: + host.mark_unhealthy(_MODEL_CLEANUP_FAILED_REASON, exc) + + def run_demo_session( *, context: RunContext, @@ -390,13 +394,11 @@ def run_demo_session( except DriverInvariantError as exc: _record_run_session_error(context, exc) if provider is not None and not driver_started: - try: - context.host.call(provider.close) - except Exception as close_exc: - if session_edges is not None: - session_edges.record_cleanup_error(close_exc) - else: - _record_run_cleanup_error(context, close_exc) + _close_partial_provider_sync( + context=context, + provider=provider, + session_edges=session_edges, + ) if session_edges is not None and ( driver_started or not session_edges.is_closed ): @@ -410,13 +412,11 @@ def run_demo_session( except Exception as exc: _record_run_session_error(context, exc) if provider is not None and not driver_started: - try: - context.host.call(provider.close) - except Exception as close_exc: - if session_edges is not None: - session_edges.record_cleanup_error(close_exc) - else: - _record_run_cleanup_error(context, close_exc) + _close_partial_provider_sync( + context=context, + provider=provider, + session_edges=session_edges, + ) if session_edges is not None and ( driver_started or not session_edges.is_closed ): @@ -623,14 +623,56 @@ def _close_on_host_best_effort( host: RuntimeHost, close: Any, session_edges: SessionEdges, -) -> None: +) -> bool: try: - host.call(_close_safely, close, session_edges) + cleanup_succeeded = host.call(_close_safely, close, session_edges) except Exception as exc: # If the host/worker is already unavailable, do not fall back to calling # model-affine cleanup directly on the caller thread. Record the loss and # let close_result finalize output, transport, and metrics. session_edges.record_cleanup_error(exc) + _mark_host_cleanup_failed(host, exc) + return False + if not cleanup_succeeded: + _mark_host_cleanup_failed(host) + return False + return True + + +def _close_partial_provider_sync( + *, + context: RunContext, + provider: Any, + session_edges: SessionEdges | None, +) -> None: + if session_edges is not None: + _close_on_host_best_effort( + host=context.host, + close=provider.close, + session_edges=session_edges, + ) + return + try: + cleanup_succeeded = context.host.call( + _close_run_provider_safely, + provider.close, + context, + ) + except Exception as exc: + _record_run_cleanup_error(context, exc) + _mark_host_cleanup_failed(context.host, exc) + return + if not cleanup_succeeded: + _mark_host_cleanup_failed(context.host) + + +def _close_run_provider_safely(close: Any, context: RunContext) -> bool: + try: + close() + except Exception as exc: + _record_run_cleanup_error(context, exc) + return False + return True async def shielded_session_cleanup( @@ -832,6 +874,7 @@ def _record_provider_cleanup_error( session_edges.record_cleanup_error(exc) else: _record_run_cleanup_error(context, exc) + _mark_host_cleanup_failed(context.host, exc) def _record_run_cleanup_error(context: RunContext, exc: Exception) -> None: diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index ff578e2e1..68f887d3a 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -122,6 +122,37 @@ def test_batch_driver_runs_fake_video_demo_through_runtime_host() -> None: assert "step" not in host.calls +def test_batch_driver_cleanup_failure_marks_host_unhealthy() -> None: + session = _FakeVideoSession(num_steps=1) + runtime = _FakeVideoRuntime(session=session) + host = RuntimeHost(runtime) + provider = _FakeVideoModelInputProvider( + fail_close=RuntimeError("provider close failed") + ) + metrics = InMemorySessionMetricsRecorder() + + try: + result = BatchSessionDriver().run_one_session( + host=host, + provider=provider, + session_edges=SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=_RecordingOutputSink(), + cleanup_tasks=set(), + metrics=metrics, + ), + pipeline=StepPipeline(), + ) + + assert result.status == "completed" + assert not host.is_healthy + assert host.unhealthy_reason == "model-affine cleanup failed" + assert provider.close_count == 1 + assert metrics.cleanup_errors == ["provider close failed"] + finally: + host.close() + + def test_batch_driver_slices_windows_from_step_requirements() -> None: session = _FakeVideoSession(num_steps=2, input_frame_counts=(3, 2)) runtime = _FakeVideoRuntime(session=session) @@ -329,6 +360,8 @@ def test_run_demo_session_keeps_failure_when_run_cleanup_metrics_fail() -> None: assert result.status == "failed" assert result.reason == "provider incompatible" assert provider.close_count == 1 + assert not context.host.is_healthy + assert context.host.unhealthy_reason == "model-affine cleanup failed" assert run_metrics.cleanup_error_attempts == 1 assert run_metrics.sessions == [result] assert runtime.start_session_inputs == [] @@ -360,6 +393,8 @@ async def test_run_demo_session_async_keeps_failure_when_run_cleanup_metrics_fai assert result.status == "failed" assert result.reason == "provider incompatible" assert provider.close_count == 1 + assert not context.host.is_healthy + assert context.host.unhealthy_reason == "model-affine cleanup failed" assert run_metrics.cleanup_error_attempts == 1 assert run_metrics.sessions == [result] assert runtime.start_session_inputs == [] From 26e4510cc11aa29da9a048b4a3f0cda0efd2bfd6 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 01:39:58 +0000 Subject: [PATCH 40/51] Document runtime cleanup quarantine rationale --- flashdreams/flashdreams/runtime/demo/drivers.py | 6 ++++++ 1 file changed, 6 insertions(+) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 1849cddef..087135c63 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -874,6 +874,9 @@ def _record_provider_cleanup_error( session_edges.record_cleanup_error(exc) else: _record_run_cleanup_error(context, exc) + # Partial async assembly may only have a provider to close. If that + # model-affine cleanup fails, quarantine the host instead of admitting a new + # session onto a worker that may still own model resources. _mark_host_cleanup_failed(context.host, exc) @@ -930,6 +933,9 @@ def _close_model_resources_safely( session_edges: SessionEdges, ) -> bool: resources_closed = True + # Session and provider close are intentionally ordered on the model worker. + # If session close hangs, timeout handling records orphaned cleanup and + # quarantines the host rather than moving provider close to another thread. if session_close is not None: resources_closed = _close_safely(session_close, session_edges) return _close_safely(provider_close, session_edges) and resources_closed From 1acd7d55853fb448b1f2474bdb913e56f2168060 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 02:12:49 +0000 Subject: [PATCH 41/51] Collapse OmniDreams WebRTC onto shared runtime session --- .../flashdreams/serving/webrtc/manager.py | 10 +- .../flashdreams/serving/webrtc/runtime.py | 1 + flashdreams/tests/test_webrtc_manager.py | 18 +- .../omnidreams/omnidreams/demo/webrtc.py | 851 ++++++++++++------ .../omnidreams/tests/test_demo_api.py | 166 ++-- 5 files changed, 682 insertions(+), 364 deletions(-) diff --git a/flashdreams/flashdreams/serving/webrtc/manager.py b/flashdreams/flashdreams/serving/webrtc/manager.py index 4e20959cc..6299cfb75 100644 --- a/flashdreams/flashdreams/serving/webrtc/manager.py +++ b/flashdreams/flashdreams/serving/webrtc/manager.py @@ -385,9 +385,17 @@ def _prepare_inference_session_step( source_schema=self._runtime.input_source_schema, ) mapping = self._runtime.input_mapping + inference_input = InferenceInput( + metadata={ + **dict(user_window.metadata), + "frame_times": tuple(user_window.frame_times), + "window_start_s": window.start_s, + "window_end_s": window.end_s, + } + ) return mapping.map_step_inputs( canonical_inputs=canonical_inputs, - inference_input=InferenceInput(), + inference_input=inference_input, request=_step_request_from_requirements(request, window=window), ) diff --git a/flashdreams/flashdreams/serving/webrtc/runtime.py b/flashdreams/flashdreams/serving/webrtc/runtime.py index 70c16fb38..32c7d4365 100644 --- a/flashdreams/flashdreams/serving/webrtc/runtime.py +++ b/flashdreams/flashdreams/serving/webrtc/runtime.py @@ -37,6 +37,7 @@ class WebRTCControlSignal(IntEnum): CLOSE = 3 EVENT = 4 SESSION_STEP = 5 + SESSION_CLOSE = 6 EXIT = 99 diff --git a/flashdreams/tests/test_webrtc_manager.py b/flashdreams/tests/test_webrtc_manager.py index f392a910a..5b60de566 100644 --- a/flashdreams/tests/test_webrtc_manager.py +++ b/flashdreams/tests/test_webrtc_manager.py @@ -19,6 +19,7 @@ UserInputEvent, UserInputs, ) +from flashdreams.runtime.demo.timing import SPARSE_KEY_SEGMENTS_METADATA_KEY from flashdreams.serving.webrtc import manager as manager_module from flashdreams.serving.webrtc.controls import WSAD_SUPPORTED_KEYS from flashdreams.serving.webrtc.encoders import ChunkDeliveryResult @@ -502,6 +503,9 @@ def canonicalize( return object() class _RecordingMapping: + def __init__(self) -> None: + self.inference_inputs: list[InferenceInput] = [] + def map_step_inputs( self, *, @@ -509,14 +513,16 @@ def map_step_inputs( inference_input: InferenceInput, request: StepRequest, ) -> InferenceInput: - del canonical_inputs, inference_input + del canonical_inputs + self.inference_inputs.append(inference_input) return InferenceInput(step={"mapped_step": request.step_index}) + mapping = _RecordingMapping() runtime = SimpleNamespace( start_inference_session=lambda: object(), input_canonicalizer=_RecordingCanonicalizer(), input_source_schema=object(), - input_mapping=_RecordingMapping(), + input_mapping=mapping, ) provider = manager_module._LegacyWebRTCModelInputProvider(runtime=runtime) skipped_inputs = UserInputs( @@ -543,8 +549,10 @@ def map_step_inputs( user_window=manager_module.UserInputWindow( start_s=2.0, end_s=3.0, + frame_times=(2.25, 2.75), inputs=current_inputs, metadata={ + SPARSE_KEY_SEGMENTS_METADATA_KEY: ((2.0, 3.0, frozenset({"w"})),), WEBRTC_SKIPPED_INPUTS_METADATA_KEY: skipped_inputs, WEBRTC_SKIPPED_WINDOW_METADATA_KEY: (0.0, 2.0), }, @@ -554,6 +562,12 @@ def map_step_inputs( assert prepared.inference_input == InferenceInput(step={"mapped_step": 0}) assert runtime.input_canonicalizer.windows == [(0.0, 2.0), (2.0, 3.0)] assert runtime.input_canonicalizer.event_batches == [["key_down"], ["key_up"]] + assert mapping.inference_inputs[0].metadata["frame_times"] == (2.25, 2.75) + assert mapping.inference_inputs[0].metadata["window_start_s"] == 2.0 + assert mapping.inference_inputs[0].metadata["window_end_s"] == 3.0 + assert mapping.inference_inputs[0].metadata[SPARSE_KEY_SEGMENTS_METADATA_KEY] == ( + (2.0, 3.0, frozenset({"w"})), + ) @pytest.mark.asyncio diff --git a/integrations/omnidreams/omnidreams/demo/webrtc.py b/integrations/omnidreams/omnidreams/demo/webrtc.py index 00857ea85..4468eda3d 100644 --- a/integrations/omnidreams/omnidreams/demo/webrtc.py +++ b/integrations/omnidreams/omnidreams/demo/webrtc.py @@ -17,64 +17,72 @@ from __future__ import annotations -import tempfile -import time -from collections.abc import Callable +import math +from collections.abc import Callable, Mapping from dataclasses import dataclass, replace from pathlib import Path from typing import Any -import cv2 -import numpy as np import torch from loguru import logger -from omnidreams.conditioning.conditioning_wrapper import ( - AV_POSITIVE_PROMPT, - OmnidreamsConditioningState, - OmnidreamsConditioningWrapper, - TextPrompt, -) -from omnidreams.conditioning.renderer import load_and_attach_ludus_scene -from omnidreams.conditioning.world_scenario.data_loaders import load_scene -from omnidreams.conditioning.world_scenario.settings import SETTINGS from omnidreams.config import OMNIDREAMS_CONFIGS -from omnidreams.scenes import ( - SCENE_CLIPGT_DIRNAME, - SCENE_PROMPT_FILENAME, - SCENE_VARIANT_DEFAULT, - ensure_hf_scene_synced, - extract_local_scene, - prepare_clipgt_dir, - resolve_scene_assets, -) +from omnidreams.scenes import SCENE_VARIANT_DEFAULT from omnidreams.transformer import CosmosTransformerConfig -from flashdreams.runtime import InferenceConfig, StepResult -from flashdreams.runtime.demo import DemoSpec, WebRTCAppResources, WebRTCOutputSpec +from flashdreams.core.distributed.rank_orchestration import distributed_op +from flashdreams.runtime import ( + CanonicalInputs, + CanonicalInputSchema, + InferenceConfig, + InferenceInput, + InferenceInputSchema, + InputCanonicalizer, + StepRequest, + StepRequirements, + StepResult, + TimeWindow, + step_requirements_from_request, +) +from flashdreams.runtime.demo import ( + DemoSpec, + PreparedScenario, + SessionInfo, + UserInputWindow, + WebRTCAppResources, + WebRTCOutputSpec, +) +from flashdreams.runtime.demo.timing import SPARSE_KEY_SEGMENTS_METADATA_KEY from flashdreams.runtime.demo.webrtc import ( CreateWebRTCApp, RunWebRTCServer, serve_webrtc_demo, ) from flashdreams.serving.webrtc.bootstrap import run_webrtc_server -from flashdreams.serving.webrtc.controls import ( - WSAD_SUPPORTED_KEYS, - CameraPoseIntegrator, - PoseSegment, -) +from flashdreams.serving.webrtc.controls import WSAD_SUPPORTED_KEYS, PoseSegment from flashdreams.serving.webrtc.encoders import EncoderBackend from flashdreams.serving.webrtc.manager import BaseWebRTCSessionManager -from flashdreams.serving.webrtc.runtime import ThreadAffineDistributedWebRTCRuntime +from flashdreams.serving.webrtc.runtime import ( + ThreadAffineDistributedWebRTCRuntime, + WebRTCControlSignal, +) from flashdreams.serving.webrtc.server import create_webrtc_app +from .providers import ( + LudusSceneConditioningProvider, + keyboard_driving_user_input_schema, +) +from .runtime import OmnidreamsRuntime, OmnidreamsRuntimeOptions, PipelineFactory from .spec import ( DEFAULT_OMNIDREAMS_PRESET, DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, OMNIDREAMS_MODEL_ID, + OmnidreamsLudusReplayScenario, resolve_webrtc_scenario, ) WebRTCRuntimeFactory = Callable[..., Any] +_WEBRTC_SESSION_TOTAL_BLOCKS = 2_147_483_647 +_WEBRTC_STEP_REQUEST_KEY = "omnidreams_webrtc_step_request" class OmnidreamsWebRTCModelRuntimeError(RuntimeError): @@ -121,7 +129,7 @@ class OmnidreamsWebRTCModelRuntimeConfig: move_speed_per_s: float = 6.0 """Forward and reverse translation speed in scene units per second.""" - rotate_speed_rad_per_s: float = float(np.deg2rad(35.0)) + rotate_speed_rad_per_s: float = math.radians(35.0) """Left and right rotation speed in radians per second.""" warmup_chunks: int = 10 @@ -142,6 +150,9 @@ class OmnidreamsWebRTCModelRuntimeConfig: encoder_gop: int = 30 """WebRTC video encoder group-of-pictures length.""" + pipeline_factory: PipelineFactory | None = None + """Optional test/runtime override for constructing the shared pipeline.""" + class OmnidreamsWebRTCModelRuntime( ThreadAffineDistributedWebRTCRuntime[ @@ -149,7 +160,7 @@ class OmnidreamsWebRTCModelRuntime( None, ] ): - """Run one single-view OmniDreams scene with browser camera controls.""" + """Compatibility WebRTC facade over the shared OmniDreams runtime/session.""" def __init__(self, *, config: OmnidreamsWebRTCModelRuntimeConfig) -> None: super().__init__( @@ -157,184 +168,96 @@ def __init__(self, *, config: OmnidreamsWebRTCModelRuntimeConfig) -> None: runtime_error_type=OmnidreamsWebRTCModelRuntimeError, thread_name="omnidreams-demo-runtime", ) - self.pose_integrator = self._new_pose_integrator() - self._wrapper: OmnidreamsConditioningWrapper | None = None - self._state: OmnidreamsConditioningState | None = None - self._renderer: Any | None = None - self._scene_data: Any | None = None - self._initial_rgb_frames: torch.Tensor | None = None - self._text_prompts: list[TextPrompt] | None = None - self._camera_to_rig: torch.Tensor | None = None - self._initial_ego_pose: np.ndarray | None = None - self._step_index = 0 - self._next_timestamp_us = 0 - self._clipgt_temp_dir: tempfile.TemporaryDirectory[str] | None = None - - def _new_pose_integrator(self) -> CameraPoseIntegrator: - return CameraPoseIntegrator( - move_speed_per_s=self.config.move_speed_per_s, - rotate_speed_rad_per_s=self.config.rotate_speed_rad_per_s, - coordinate_system="FLU", - ) + self.input_source_schema = keyboard_driving_user_input_schema() + self.input_canonicalizer = InputCanonicalizer() + self.input_mapping = _OmnidreamsWebRTCInputMapping() + self._runtime: OmnidreamsRuntime | None = None + self._active_provider: LudusSceneConditioningProvider | None = None + self._active_session: Any | None = None + self._debug_session: _OmnidreamsHDMapDebugSession | None = None + self._steady_output_frame_count_value = 1 def _is_runtime_initialized(self) -> bool: - return self._wrapper is not None and self._renderer is not None + return self._runtime is not None def _runtime_step_index(self) -> int: - return self._step_index + requirements = self._next_step_requirements_sync() + if requirements is None: + return 0 + return requirements.step_index def _next_input_frame_count(self) -> int: - wrapper = self._require_wrapper() - if self._state is None: - return int(wrapper.initial_frame_chunk_size) - return int(wrapper.frame_chunk_size) + requirements = self._next_step_requirements_sync() + if requirements is None: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC session is complete." + ) + return requirements.input_frame_count def _steady_output_frame_count(self) -> int: - return int(self._require_wrapper().frame_chunk_size) + return self._steady_output_frame_count_value def _initialize_sync(self) -> None: - if self._wrapper is not None: + if self._runtime is not None: return - - init_t0 = time.perf_counter() - cfg = self.config - transformer_cfg = cfg.pipeline_config.diffusion_model.transformer - if not isinstance(transformer_cfg, CosmosTransformerConfig): - raise TypeError( - "OmniDreams WebRTC requires a CosmosTransformerConfig pipeline." - ) - if transformer_cfg.num_views != 1: - raise ValueError( - "OmniDreams WebRTC supports only single-view configs; " - f"{cfg.pipeline_config_name!r} has num_views=" - f"{transformer_cfg.num_views}." - ) if self._device.type == "cuda" and not torch.cuda.is_available(): raise RuntimeError("CUDA is required for OmniDreams WebRTC inference.") - - scene_dir = self._prepare_scene() - clipgt_dir, first_frame_path, prompt_path = resolve_scene_assets( - scene_dir, - prompt_filename=SCENE_PROMPT_FILENAME, - clipgt_dirname=SCENE_CLIPGT_DIRNAME, - camera_name=cfg.camera_name, - variant=cfg.scene_variant, - ) - self._initial_rgb_frames = self._load_first_frame(first_frame_path) - prompt = prompt_path.read_text(encoding="utf-8").strip() or AV_POSITIVE_PROMPT - self._text_prompts = [TextPrompt(positive=prompt)] - - loadable_clipgt_dir, self._clipgt_temp_dir = prepare_clipgt_dir(clipgt_dir) - logger.info("Loading OmniDreams scene data from {}", loadable_clipgt_dir) - scene_data = load_scene( - loadable_clipgt_dir, - camera_names=[cfg.camera_name], - max_frames=-1, - input_pose_fps=SETTINGS["INPUT_POSE_FPS"], - resize_resolution_hw=(cfg.video_height, cfg.video_width), - ) - scene_data = load_and_attach_ludus_scene( - loadable_clipgt_dir, - scene_data, - device=self._device, - ) - self._validate_scene_data(scene_data, scene_dir=loadable_clipgt_dir) - + _validate_single_view_pipeline_config( + pipeline_config_name=self.config.pipeline_config_name, + pipeline_config=self.config.pipeline_config, + ) logger.info( - "Setting up OmniDreams pipeline {} on {}.", - cfg.pipeline_config_name, + "Setting up shared OmniDreams runtime {} on {} for WebRTC.", + self.config.pipeline_config_name, self._device, ) - wrapper = OmnidreamsConditioningWrapper( - pipeline_config_name=cfg.pipeline_config_name, - pipeline_config=cfg.pipeline_config, - resolution_wh=(cfg.video_width, cfg.video_height), - seed_for_every_rollout=cfg.seed, - device=self._device, - ) - renderer = wrapper.create_renderer(scene_data, [cfg.camera_name]) - - self._wrapper = wrapper - self._renderer = renderer - self._scene_data = scene_data - self._camera_to_rig = torch.as_tensor( - scene_data.camera_extrinsics[cfg.camera_name], - device=self._device, - dtype=torch.float32, - ) - self._initial_ego_pose = scene_data.ego_poses[0].transformation_matrix - self._next_timestamp_us = int(scene_data.ego_poses[0].timestamp) - self._reset_rollout_sync() - self._initialize_video_encoder_sync() - logger.info( - "OmniDreams runtime initialization complete in {:.1f}s.", - time.perf_counter() - init_t0, + self._runtime = OmnidreamsRuntime( + config=self._inference_config(), + options=OmnidreamsRuntimeOptions( + pipeline_config=self.config.pipeline_config, + pipeline_factory=self.config.pipeline_factory, + ), ) - - def _prepare_scene(self) -> Path: - cfg = self.config - if cfg.scene_dir is None: - return ensure_hf_scene_synced( - cfg.scene_uuid or DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, - variant=cfg.scene_variant, - clipgt_dirname=SCENE_CLIPGT_DIRNAME, - ) - return extract_local_scene( - cfg.scene_dir, - scene_uuid=cfg.scene_uuid, - variant=cfg.scene_variant, - clipgt_dirname=SCENE_CLIPGT_DIRNAME, - ) - - def _load_first_frame(self, path: Path) -> torch.Tensor: - logger.info("Loading OmniDreams first frame from {}", path) - image_bgr = cv2.imread(str(path), cv2.IMREAD_COLOR) - if image_bgr is None: - raise RuntimeError(f"Failed to read first frame from {path}") - image_rgb = cv2.cvtColor(image_bgr, cv2.COLOR_BGR2RGB) - image_rgb = cv2.resize( - image_rgb, - (self.config.video_width, self.config.video_height), - interpolation=cv2.INTER_CUBIC, - ) - return ( - torch.from_numpy(image_rgb) - .permute(2, 0, 1) - .contiguous() - .unsqueeze(0) - .unsqueeze(0) - .to(device=self._device, dtype=torch.uint8) - ) - - def _validate_scene_data(self, scene_data: Any, *, scene_dir: Path) -> None: - camera_name = self.config.camera_name - if not scene_data.ego_poses: - raise ValueError(f"Scene {scene_dir} has no ego poses.") - if camera_name not in scene_data.camera_models: - raise ValueError(f"Camera {camera_name!r} was not loaded from {scene_dir}.") - if camera_name not in scene_data.camera_extrinsics: - raise ValueError( - f"Camera {camera_name!r} has no extrinsics in {scene_dir}." - ) + self._initialize_video_encoder_sync() def _reset_rollout_sync(self, session_input: None = None) -> None: del session_input - wrapper = self._require_wrapper() - if self._renderer is None or self._scene_data is None: - raise OmnidreamsWebRTCModelRuntimeError("Scene state is not initialized.") - if self._initial_ego_pose is None: - raise OmnidreamsWebRTCModelRuntimeError( - "Initial camera pose is unavailable." + self._close_active_session_sync() + runtime = self._require_runtime() + scenario = self._session_scenario() + prepared = PreparedScenario( + initial_inputs=InferenceInput(global_conditioning={"scenario": scenario}), + source_schema=self.input_source_schema, + metadata={ + "conditioning_mode": "ludus-scene-driving", + "model_id": OMNIDREAMS_MODEL_ID, + "preset_id": self.config.pipeline_config_name, + }, + ) + provider = LudusSceneConditioningProvider( + scenario=prepared, + config=self._inference_config(), + ) + try: + initial_input = provider.prepare_initial_input() + session = runtime.start_session(initial_input) + except Exception: + provider.close() + raise + self._active_provider = provider + if self.config.debug_serve_hdmaps: + self._debug_session = _OmnidreamsHDMapDebugSession( + pipeline=runtime.pipeline, + scenario=scenario, ) - - if self._state is not None and self._state.pipeline_cache is not None: - del self._state.pipeline_cache - self._state = None - self._step_index = 0 - self.pose_integrator = self._new_pose_integrator() - self.pose_integrator.reset(self._initial_ego_pose) - self._next_timestamp_us = int(self._scene_data.ego_poses[0].timestamp) - wrapper.set_rollout_seed(self.config.seed) + self._active_session = self._debug_session + else: + self._debug_session = None + self._active_session = session + self._steady_output_frame_count_value = _steady_output_frame_count( + self._active_session, + fallback_pipeline=runtime.pipeline, + ) def _generate_one_chunk_sync( self, @@ -342,118 +265,478 @@ def _generate_one_chunk_sync( segments: list[PoseSegment], frame_times: list[float], ) -> StepResult: - wrapper = self._require_wrapper() - if ( - self._renderer is None - or self._initial_rgb_frames is None - or self._text_prompts is None - or self._camera_to_rig is None - ): - raise OmnidreamsWebRTCModelRuntimeError("Runtime is not initialized.") - if len(frame_times) != self._next_input_frame_count(): + request = self._next_step_request_sync() + if request is None: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC session is complete." + ) + inputs = self.input_mapping.map_step_inputs( + canonical_inputs=CanonicalInputs(), + inference_input=InferenceInput( + metadata={ + SPARSE_KEY_SEGMENTS_METADATA_KEY: tuple(segments), + "frame_times": tuple(frame_times), + "window_start_s": request.step_index / float(self.config.fps), + "window_end_s": (request.step_index + len(frame_times)) + / float(self.config.fps), + } + ), + request=request, + ) + return self._step_active_session_sync(inputs) + + def _close_sync(self) -> None: + self._close_active_session_sync() + runtime = self._runtime + self._runtime = None + if runtime is not None: + runtime.close() + if self._device.type == "cuda" and torch.cuda.is_available(): + torch.cuda.synchronize(device=self._device) + torch.cuda.empty_cache() + + async def start_inference_session(self) -> "_OmnidreamsWebRTCInferenceSession": + self._require_open_and_initialized() + if not await self._worker.call(self._has_active_session_sync): + await self.reset_for_new_session() + return _OmnidreamsWebRTCInferenceSession(self) + + def _next_step_request_sync(self) -> StepRequest | None: + requirements = self._next_step_requirements_sync() + if requirements is None: + return None + metadata = dict(requirements.metadata) + metadata["input_frame_count"] = requirements.input_frame_count + if requirements.steady_output_frame_count is not None: + metadata["steady_output_frame_count"] = ( + requirements.steady_output_frame_count + ) + return StepRequest( + step_index=requirements.step_index, + inference_input_schema=requirements.inference_input_schema, + metadata=metadata, + ) + + def _next_step_requirements_sync(self) -> StepRequirements | None: + session = self._require_active_session() + next_requirements = getattr(session, "next_step_requirements", None) + if callable(next_requirements): + result = next_requirements() + else: + next_request = session.next_step_request() + if next_request is None: + return None + result = step_requirements_from_request(next_request) + if result is None: + return None + if not isinstance(result, StepRequirements): + raise TypeError( + "OmniDreams WebRTC session requirements must be StepRequirements, " + f"got {type(result).__name__}." + ) + return result + + def _session_info_sync(self) -> SessionInfo: + return SessionInfo( + output_layout="bvtchw", + steady_output_frame_count=self._steady_output_frame_count(), + metadata={"model_id": OMNIDREAMS_MODEL_ID}, + ) + + def _step_active_session_sync(self, inputs: InferenceInput) -> StepResult: + provider = self._require_active_provider() + session = self._require_active_session() + request = _request_from_step_inputs(inputs) + requirements = step_requirements_from_request( + request, + allow_user_input_window=True, + ) + window = _user_window_from_step_inputs( + inputs, + request=request, + input_frame_count=requirements.input_frame_count, + ) + prepared = provider.prepare_step(request=requirements, user_window=window) + if prepared.control.close_session: raise OmnidreamsWebRTCModelRuntimeError( - f"Expected {self._next_input_frame_count()} frame times for " - f"step {self._step_index}, got {len(frame_times)}." + prepared.control.reason or "OmniDreams WebRTC input is exhausted." ) - if not segments: + if prepared.control.reset: + reset_input = prepared.control.reset_input + session.reset(reset_input) + if not prepared.control.provider_already_reset: + provider.reset(reset_input) raise OmnidreamsWebRTCModelRuntimeError( - f"Step {self._step_index} received no control segments." + prepared.control.reason or "OmniDreams WebRTC session reset requested." + ) + if prepared.inference_input is None: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC provider returned no inference input." + ) + result = session.step(prepared.inference_input) + if not isinstance(result, StepResult): + raise TypeError( + "OmniDreams WebRTC session steps must produce StepResult, got " + f"{type(result).__name__}." ) + return result + + @distributed_op(WebRTCControlSignal.SESSION_STEP) + def _step_active_session_sync_all_ranks( + self, + inputs: InferenceInput, + ) -> StepResult: + return self._step_active_session_sync(inputs) + + @distributed_op(WebRTCControlSignal.SESSION_CLOSE) + def _close_active_session_sync_all_ranks(self) -> None: + self._close_active_session_sync() + + def _close_active_session_sync(self) -> None: + session = self._active_session + provider = self._active_provider + self._active_session = None + self._debug_session = None + self._active_provider = None + first_error: Exception | None = None + close_session = getattr(session, "close", None) + if callable(close_session): + try: + close_session() + except Exception as exc: + first_error = exc + if provider is not None: + try: + provider.close() + except Exception as exc: + if first_error is None: + first_error = exc + if first_error is not None: + raise first_error + + def _has_active_session_sync(self) -> bool: + return self._active_session is not None and self._active_provider is not None + + def _require_runtime(self) -> OmnidreamsRuntime: + if self._runtime is None: + raise OmnidreamsWebRTCModelRuntimeError("Runtime is not initialized.") + return self._runtime - ego_poses = self.pose_integrator.integrate_chunk( - segments=segments, - frame_times=frame_times, - ) - ego_poses_t = torch.from_numpy(ego_poses).to( - device=self._device, - dtype=torch.float32, - ) - camera_poses = torch.einsum("nij,jk->nik", ego_poses_t, self._camera_to_rig) - frame_timestamps_us = self._consume_timestamps(len(frame_times)) - serve_hdmaps = self.config.debug_serve_hdmaps - - if self._state is None: - output = wrapper.start_generation( - text_prompts=self._text_prompts, - initial_rgb_frames=self._initial_rgb_frames, - renderer=self._renderer, - camera_names=[self.config.camera_name], - camera_poses_per_view={self.config.camera_name: camera_poses}, - frame_timestamps_us=frame_timestamps_us, - skip_video_generation=serve_hdmaps, + def _require_active_session(self) -> Any: + if self._active_session is None: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC session is not initialized." ) - else: - output = wrapper.continue_generation( - state=self._state, - camera_names=[self.config.camera_name], - camera_poses_per_view={self.config.camera_name: camera_poses}, - frame_timestamps_us=frame_timestamps_us, - skip_video_generation=serve_hdmaps, + return self._active_session + + def _require_active_provider(self) -> LudusSceneConditioningProvider: + if self._active_provider is None: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC provider is not initialized." ) - self._state = output.state - if self._state.pipeline_cache is not None: - wrapper.finalize_block_generation( - self._state.pipeline_cache, - output.finalization_state, + return self._active_provider + + def _inference_config(self) -> InferenceConfig: + return InferenceConfig( + model_id=OMNIDREAMS_MODEL_ID, + preset_id=self.config.pipeline_config_name, + device=str(self.config.device), + seed=self.config.seed, + runtime_options={"seed": self.config.seed}, + ) + + def _session_scenario(self) -> OmnidreamsLudusReplayScenario: + return OmnidreamsLudusReplayScenario( + keyboard_events=(), + scene_dir=self.config.scene_dir, + scene_uuid=self.config.scene_uuid or DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, + scene_variant=self.config.scene_variant, + camera_name=self.config.camera_name, + total_blocks=_WEBRTC_SESSION_TOTAL_BLOCKS, + pixel_height=self.config.video_height, + pixel_width=self.config.video_width, + fps=self.config.fps, + move_speed_per_s=self.config.move_speed_per_s, + rotate_speed_rad_per_s=self.config.rotate_speed_rad_per_s, + ) + + +class _OmnidreamsWebRTCInputMapping: + """Carry shared WebRTC window facts into the OmniDreams session facade.""" + + def validate( + self, + *, + canonical_schema: CanonicalInputSchema | None = None, + inference_input_schema: InferenceInputSchema | None = None, + ) -> None: + del canonical_schema, inference_input_schema + + def map_global_conditioning_inputs( + self, + *, + canonical_inputs: CanonicalInputs, + inference_input: InferenceInput, + ) -> InferenceInput: + del canonical_inputs + return inference_input + + def map_step_inputs( + self, + *, + canonical_inputs: CanonicalInputs, + inference_input: InferenceInput, + request: StepRequest, + ) -> InferenceInput: + del canonical_inputs + step = dict(inference_input.step) + step[_WEBRTC_STEP_REQUEST_KEY] = request + return InferenceInput( + global_conditioning=inference_input.global_conditioning, + step=step, + metadata=inference_input.metadata, + ) + + +class _OmnidreamsWebRTCInferenceSession: + """Synchronous session proxy consumed by the shared WebRTC compatibility path.""" + + def __init__(self, runtime: OmnidreamsWebRTCModelRuntime) -> None: + self._runtime = runtime + self._closed = False + + def session_info(self) -> SessionInfo: + self._require_open() + return self._runtime._worker.call_blocking(self._runtime._session_info_sync) + + def next_step_requirements(self) -> StepRequirements | None: + self._require_open() + return self._runtime._worker.call_blocking( + self._runtime._next_step_requirements_sync + ) + + def next_step_request(self) -> StepRequest | None: + self._require_open() + return self._runtime._worker.call_blocking( + self._runtime._next_step_request_sync + ) + + def step(self, inputs: InferenceInput) -> StepResult: + self._require_open() + return self._runtime._worker.call_blocking( + self._runtime._step_active_session_sync_all_ranks, + inputs, + ) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self._require_open() + self._runtime._worker.call_blocking(self._runtime._reset_rollout_sync_all_ranks) + + def close(self) -> None: + if self._closed: + return + self._closed = True + self._runtime._worker.call_blocking( + self._runtime._close_active_session_sync_all_ranks + ) + + def _require_open(self) -> None: + if self._closed: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC inference session is closed." ) - metadata = {"stream": "hdmap" if serve_hdmaps else "rgb"} - if serve_hdmaps: - video_chunk = output.condition_frames - else: - if output.rgb_frames is None: - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams generation produced no RGB frames." - ) - video_chunk = output.rgb_frames - result = StepResult.from_video_chunk( + +class _OmnidreamsHDMapDebugSession: + """Session-shaped debug path that streams rendered Ludus HDMaps.""" + + def __init__( + self, *, pipeline: Any, scenario: OmnidreamsLudusReplayScenario + ) -> None: + self._pipeline = pipeline + self._scenario = scenario + self._step_index = 0 + self._closed = False + + def session_info(self) -> SessionInfo: + return SessionInfo( + output_layout="bvtchw", + steady_output_frame_count=self._steady_output_frame_count(), + metadata={"stream": "hdmap"}, + ) + + def next_step_requirements(self) -> StepRequirements | None: + if self._closed or self._step_index >= self._scenario.total_blocks: + return None + return StepRequirements( step_index=self._step_index, - video_chunk=video_chunk.detach(), - layout="bvtchw", - metadata=metadata, + input_frame_count=self._num_frames(self._step_index), + steady_output_frame_count=self._steady_output_frame_count(), + ) + + def next_step_request(self) -> StepRequest | None: + requirements = self.next_step_requirements() + if requirements is None: + return None + return StepRequest( + step_index=requirements.step_index, + metadata={ + "input_frame_count": requirements.input_frame_count, + "steady_output_frame_count": requirements.steady_output_frame_count, + }, ) - expected_frames = len(frame_times) - if result.frame_count != expected_frames: + + def step(self, inputs: InferenceInput) -> StepResult: + requirements = self.next_step_requirements() + if requirements is None: raise OmnidreamsWebRTCModelRuntimeError( - f"Expected generated chunk to contain {expected_frames} frames, " - f"got {result.frame_count}." + "OmniDreams WebRTC debug session is complete." ) + hdmap = inputs.step.get("hdmap") + if not isinstance(hdmap, torch.Tensor): + raise TypeError("OmniDreams WebRTC debug session requires step['hdmap'].") + result = StepResult.from_video_chunk( + step_index=requirements.step_index, + video_chunk=hdmap.detach(), + layout="bvtchw", + metadata={"stream": "hdmap"}, + ) self._step_index += 1 return result - def _consume_timestamps(self, num_frames: int) -> list[int]: - step_us = int(round(1_000_000 / self.config.fps)) - timestamps = [ - self._next_timestamp_us + frame_index * step_us - for frame_index in range(num_frames) - ] - self._next_timestamp_us += num_frames * step_us - return timestamps + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self._step_index = 0 + self._closed = False - def _close_sync(self) -> None: - if self._wrapper is not None and self._state is not None: - self._wrapper.cleanup(self._state) - elif self._renderer is not None: - self._renderer.cleanup() - self._state = None - self._wrapper = None - self._renderer = None - self._scene_data = None - self._initial_rgb_frames = None - self._text_prompts = None - self._camera_to_rig = None - self._initial_ego_pose = None - if self._clipgt_temp_dir is not None: - self._clipgt_temp_dir.cleanup() - self._clipgt_temp_dir = None - if self._device.type == "cuda": - torch.cuda.synchronize(device=self._device) - torch.cuda.empty_cache() + def close(self) -> None: + self._closed = True - def _require_wrapper(self) -> OmnidreamsConditioningWrapper: - if self._wrapper is None: - raise OmnidreamsWebRTCModelRuntimeError("Runtime is not initialized.") - return self._wrapper + def _steady_output_frame_count(self) -> int: + return self._num_frames(1) + + def _num_frames(self, step_index: int) -> int: + get_num_frames = getattr(self._pipeline, "get_num_frames", None) + if not callable(get_num_frames): + return 1 + return int(get_num_frames(step_index)) + + +def _request_from_step_inputs(inputs: InferenceInput) -> StepRequest: + request = inputs.step.get(_WEBRTC_STEP_REQUEST_KEY) + if not isinstance(request, StepRequest): + raise TypeError( + "OmniDreams WebRTC step input is missing the shared StepRequest." + ) + return request + + +def _user_window_from_step_inputs( + inputs: InferenceInput, + *, + request: StepRequest, + input_frame_count: int, +) -> UserInputWindow: + frame_times = _frame_times_from_metadata(inputs.metadata, input_frame_count) + segments = _segments_from_metadata(inputs.metadata) + window = request.user_input_window or TimeWindow( + start_s=float(inputs.metadata.get("window_start_s", 0.0)), + end_s=float(inputs.metadata.get("window_end_s", frame_times[-1])), + ) + return UserInputWindow( + start_s=window.start_s, + end_s=window.end_s, + frame_times=frame_times, + metadata={SPARSE_KEY_SEGMENTS_METADATA_KEY: segments}, + ) + + +def _frame_times_from_metadata( + metadata: Mapping[str, object], + input_frame_count: int, +) -> tuple[float, ...]: + value = metadata.get("frame_times") + if not isinstance(value, tuple): + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC step input is missing frame_times metadata." + ) + frame_times = tuple( + _float_metadata_value(frame_time, label="frame_times") for frame_time in value + ) + if len(frame_times) != input_frame_count: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC frame_times length does not match " + f"input_frame_count={input_frame_count}." + ) + return frame_times + + +def _segments_from_metadata(metadata: Mapping[str, object]) -> tuple[PoseSegment, ...]: + value = metadata.get(SPARSE_KEY_SEGMENTS_METADATA_KEY) + if not isinstance(value, tuple): + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC step input is missing resampled key segments." + ) + segments: list[PoseSegment] = [] + for segment in value: + if not isinstance(segment, tuple) or len(segment) != 3: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC key segments must be 3-tuples." + ) + start_s, end_s, keys = segment + if not isinstance(keys, frozenset | set | tuple | list): + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC key segment keys must be a sequence." + ) + segments.append( + ( + _float_metadata_value(start_s, label="segment start"), + _float_metadata_value(end_s, label="segment end"), + frozenset(str(key) for key in keys), + ) + ) + return tuple(segments) + + +def _float_metadata_value(value: object, *, label: str) -> float: + if isinstance(value, bool) or not isinstance(value, int | float): + raise OmnidreamsWebRTCModelRuntimeError( + f"OmniDreams WebRTC {label} metadata must be numeric." + ) + return float(value) + + +def _steady_output_frame_count(session: Any, *, fallback_pipeline: Any) -> int: + session_info = getattr(session, "session_info", None) + if callable(session_info): + value = session_info() + if isinstance(value, SessionInfo) and value.steady_output_frame_count: + return int(value.steady_output_frame_count) + get_num_frames = getattr(fallback_pipeline, "get_num_frames", None) + if callable(get_num_frames): + return int(get_num_frames(1)) + return 1 + + +def _validate_single_view_pipeline_config( + *, + pipeline_config_name: str, + pipeline_config: Any, +) -> None: + diffusion_model = getattr(pipeline_config, "diffusion_model", None) + transformer_cfg = getattr(diffusion_model, "transformer", None) + if transformer_cfg is None: + return + if not isinstance(transformer_cfg, CosmosTransformerConfig): + raise TypeError( + "OmniDreams WebRTC requires a CosmosTransformerConfig pipeline." + ) + if transformer_cfg.num_views != 1: + raise ValueError( + "OmniDreams WebRTC supports only single-view configs; " + f"{pipeline_config_name!r} has num_views={transformer_cfg.num_views}." + ) def serve_omnidreams_webrtc_demo( diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index 9e7967926..e0796b740 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -46,6 +46,7 @@ ) from flashdreams.runtime import ( + CanonicalInputs, InferenceConfig, InferenceInput, OutputArtifact, @@ -996,43 +997,88 @@ def fake_server_runner(**kwargs: Any) -> None: @pytest.mark.asyncio -async def test_omnidreams_demo_runtime_generates_directly_from_controls() -> None: +async def test_omnidreams_webrtc_runtime_uses_shared_session( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _scene, rasterizers = _install_fake_ludus_provider_dependencies(monkeypatch) + scene_path = tmp_path / "scene.usdz" + scene_path.write_bytes(b"fake") + pipeline = _VariableFrameOmnidreamsPipeline((2, 3)) config = OmnidreamsWebRTCModelRuntimeConfig( pipeline_config_name="fake", pipeline_config=object(), + pipeline_factory=lambda pipeline_config, device: pipeline, + scene_dir=scene_path, device="cpu", fps=30, + video_height=2, + video_width=2, warmup_chunks=0, ) runtime = OmnidreamsWebRTCModelRuntime(config=config) - wrapper = _FakeConditioningWrapper() - runtime._wrapper = wrapper # ty:ignore[invalid-assignment] - runtime._renderer = _FakeRenderer() - runtime._scene_data = SimpleNamespace(ego_poses=[SimpleNamespace(timestamp=1_000)]) - runtime._initial_rgb_frames = torch.zeros((1, 1, 3, 4, 5), dtype=torch.uint8) - runtime._text_prompts = [] - runtime._camera_to_rig = torch.eye(4) - runtime._initial_ego_pose = torch.eye(4).numpy() - runtime.pose_integrator.reset() - runtime._next_timestamp_us = 1_000 - - first = runtime._generate_one_chunk_sync( - segments=[(0.0, 2 / 30, frozenset({"w"}))], - frame_times=[1 / 30, 2 / 30], - ) - second = runtime._generate_one_chunk_sync( - segments=[(2 / 30, 5 / 30, frozenset({"d"}))], - frame_times=[3 / 30, 4 / 30, 5 / 30], + await runtime.initialize() + await runtime.reset_for_new_session() + session = await runtime.start_inference_session() + + first_request = session.next_step_request() + assert first_request is not None + assert first_request.metadata["input_frame_count"] == 2 + first = session.step( + runtime.input_mapping.map_step_inputs( + canonical_inputs=CanonicalInputs(), + inference_input=InferenceInput( + metadata={ + SPARSE_KEY_SEGMENTS_METADATA_KEY: ( + (0.0, 2 / 30, frozenset({"w"})), + ), + "frame_times": (1 / 30, 2 / 30), + "window_start_s": 0.0, + "window_end_s": 2 / 30, + } + ), + request=first_request, + ) + ) + second_request = session.next_step_request() + assert second_request is not None + assert second_request.metadata["input_frame_count"] == 3 + second = session.step( + runtime.input_mapping.map_step_inputs( + canonical_inputs=CanonicalInputs(), + inference_input=InferenceInput( + metadata={ + SPARSE_KEY_SEGMENTS_METADATA_KEY: ( + (2 / 30, 5 / 30, frozenset({"d"})), + ), + "frame_times": (3 / 30, 4 / 30, 5 / 30), + "window_start_s": 2 / 30, + "window_end_s": 5 / 30, + } + ), + request=second_request, + ) ) assert (first.step_index, first.frame_count) == (0, 2) assert (second.step_index, second.frame_count) == (1, 3) - assert wrapper.calls == [ - ("start", (2, 4, 4), [1_000, 34_333]), - ("continue", (3, 4, 4), [67_666, 100_999, 134_332]), + assert isinstance(session, OmnidreamsSession) is False + assert pipeline.initialize_cache_calls == [ + { + "text": [["city scene"]], + "image_shape": (1, 1, 1, 3, 2, 2), + "view_names": ["camera_front_wide_120fov"], + } + ] + assert [tuple(hdmap.shape) for hdmap in pipeline.generated_hdmaps] == [ + (1, 1, 2, 3, 2, 2), + (1, 1, 3, 3, 2, 2), ] - assert wrapper.finalized == [0, 1] + assert rasterizers[0].calls[0]["timestamps_us"] == (1_000, 34_333) + assert rasterizers[0].calls[1]["timestamps_us"] == (67_666, 100_999, 134_332) + session.close() await runtime.close() + assert rasterizers[0].closed is True class _RecordingOutputTarget: @@ -1231,12 +1277,27 @@ def finalize(self, *, autoregressive_index: int, cache: object) -> dict[str, flo return {"denoise_s": 0.25} -class _FakeRenderer: - def __init__(self) -> None: - self.closed = False +class _VariableFrameOmnidreamsPipeline(_FakeOmnidreamsPipeline): + def __init__(self, frame_counts: tuple[int, ...]) -> None: + super().__init__() + self._frame_counts = frame_counts - def cleanup(self) -> None: - self.closed = True + def get_num_frames(self, autoregressive_index: int) -> int: + return self._frame_counts[ + min(autoregressive_index, len(self._frame_counts) - 1) + ] + + def generate( + self, + *, + autoregressive_index: int, + cache: object, + hdmap: torch.Tensor, + ) -> torch.Tensor: + del cache + self.generated_hdmaps.append(hdmap.detach().clone()) + frame_count = self.get_num_frames(autoregressive_index) + return torch.full((1, 1, frame_count, 3, 2, 2), float(autoregressive_index)) class _FakeLudusRasterizer: @@ -1274,55 +1335,6 @@ def cleanup(self) -> None: self.closed = True -class _FakeConditioningWrapper: - initial_frame_chunk_size = 2 - frame_chunk_size = 3 - - def __init__(self) -> None: - self.calls: list[tuple[str, tuple[int, ...], list[int]]] = [] - self.finalized: list[int] = [] - self.cleaned = False - - def start_generation(self, **kwargs: Any) -> SimpleNamespace: - return self._output("start", kwargs=kwargs, frame_count=2, step_index=0) - - def continue_generation(self, **kwargs: Any) -> SimpleNamespace: - return self._output("continue", kwargs=kwargs, frame_count=3, step_index=1) - - def _output( - self, - operation: str, - *, - kwargs: dict[str, Any], - frame_count: int, - step_index: int, - ) -> SimpleNamespace: - poses = kwargs["camera_poses_per_view"]["camera_front_wide_120fov"] - timestamps = kwargs["frame_timestamps_us"] - self.calls.append((operation, tuple(poses.shape), timestamps)) - state = kwargs.get("state") or SimpleNamespace(pipeline_cache=object()) - return SimpleNamespace( - state=state, - condition_frames=torch.zeros( - (1, 1, frame_count, 3, 4, 5), dtype=torch.uint8 - ), - rgb_frames=torch.zeros((1, 1, frame_count, 3, 4, 5), dtype=torch.uint8), - finalization_state={"autoregressive_index": step_index}, - ) - - def finalize_block_generation( - self, - pipeline_cache: object, - finalization_state: dict[str, int], - ) -> None: - del pipeline_cache - self.finalized.append(finalization_state["autoregressive_index"]) - - def cleanup(self, state: object) -> None: - del state - self.cleaned = True - - class _FakeWebRTCRuntime: def __init__(self, config: Any) -> None: self.config = config From 53c1b9d283f79d1b6580078307689f7186156a90 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 02:29:57 +0000 Subject: [PATCH 42/51] Fix OmniDreams WebRTC shared runtime warmup --- .../flashdreams/serving/webrtc/manager.py | 10 +- .../flashdreams/serving/webrtc/warmup.py | 29 ++- flashdreams/tests/test_webrtc_manager.py | 39 ++++ flashdreams/tests/test_webrtc_warmup.py | 154 +++++++++++++++ .../omnidreams/omnidreams/demo/webrtc.py | 12 +- .../omnidreams/tests/test_demo_api.py | 181 +++++++++++++++++- 6 files changed, 417 insertions(+), 8 deletions(-) create mode 100644 flashdreams/tests/test_webrtc_warmup.py diff --git a/flashdreams/flashdreams/serving/webrtc/manager.py b/flashdreams/flashdreams/serving/webrtc/manager.py index 6299cfb75..7e065a72c 100644 --- a/flashdreams/flashdreams/serving/webrtc/manager.py +++ b/flashdreams/flashdreams/serving/webrtc/manager.py @@ -1762,7 +1762,15 @@ async def _run_realtime_driver_session( pipeline=StepPipeline(), reservation=managed_session.reservation, ) - if result.status == "failed" and result.reason: + if result.status == "completed": + logger.info("Shared WebRTC session completed.") + else: + logger.warning( + "Shared WebRTC session ended with status={} reason={}", + result.status, + result.reason, + ) + if result.status != "completed" and result.reason: channel = managed_session.control_channel if channel is not None: self._send_json(channel, make_error_payload(result.reason)) diff --git a/flashdreams/flashdreams/serving/webrtc/warmup.py b/flashdreams/flashdreams/serving/webrtc/warmup.py index 05f3d1c10..5f59c4a86 100644 --- a/flashdreams/flashdreams/serving/webrtc/warmup.py +++ b/flashdreams/flashdreams/serving/webrtc/warmup.py @@ -13,6 +13,8 @@ from aiortc.mediastreams import MediaStreamError from loguru import logger as loguru_logger +from .messages import MESSAGE_TYPE_CHUNK_DONE, MESSAGE_TYPE_ERROR + class CreateAnswerCallback(Protocol): async def __call__(self, *, offer_sdp: str, offer_type: str) -> dict[str, str]: ... @@ -47,14 +49,30 @@ async def run_loopback_warmup_session( client_peer.addTransceiver("video", direction="recvonly") channel_open = asyncio.Event() warmup_done = asyncio.Event() + warmup_failure: str | None = None received_chunks = 0 drain_tasks: set[asyncio.Task[Any]] = set() heartbeat_task: asyncio.Task[Any] | None = None + def fail_warmup(reason: str) -> None: + nonlocal warmup_failure + if warmup_done.is_set(): + return + warmup_failure = reason + warmup_done.set() + @control_channel.on("open") def on_open() -> None: channel_open.set() + @control_channel.on("close") + def on_close() -> None: + if received_chunks < num_chunks: + fail_warmup( + f"{label} loopback warmup data channel closed before warmup " + f"completed ({received_chunks}/{num_chunks} chunk(s))." + ) + @control_channel.on("message") def on_message(message: Any) -> None: nonlocal received_chunks @@ -64,7 +82,14 @@ def on_message(message: Any) -> None: payload = json.loads(message) except json.JSONDecodeError: return - if not isinstance(payload, dict) or payload.get("type") != "chunk_done": + if not isinstance(payload, dict): + return + message_type = payload.get("type") + if message_type == MESSAGE_TYPE_ERROR: + message_text = str(payload.get("message", "unknown error")) + fail_warmup(f"{label} loopback warmup failed: {message_text}") + return + if message_type != MESSAGE_TYPE_CHUNK_DONE: return received_chunks += 1 logger.info( @@ -109,6 +134,8 @@ def on_track(track: Any) -> None: for action_payload in action_payloads: control_channel.send(json.dumps(action_payload)) await asyncio.wait_for(warmup_done.wait(), timeout=warmup_timeout_s) + if warmup_failure is not None: + raise RuntimeError(warmup_failure) finally: if heartbeat_task is not None: heartbeat_task.cancel() diff --git a/flashdreams/tests/test_webrtc_manager.py b/flashdreams/tests/test_webrtc_manager.py index 5b60de566..787a8ee7b 100644 --- a/flashdreams/tests/test_webrtc_manager.py +++ b/flashdreams/tests/test_webrtc_manager.py @@ -19,6 +19,7 @@ UserInputEvent, UserInputs, ) +from flashdreams.runtime.demo import RunResult from flashdreams.runtime.demo.timing import SPARSE_KEY_SEGMENTS_METADATA_KEY from flashdreams.serving.webrtc import manager as manager_module from flashdreams.serving.webrtc.controls import WSAD_SUPPORTED_KEYS @@ -1091,6 +1092,44 @@ def peek_steady_output_num_frames(self) -> int: assert chunk_done[0]["model"] == "fake-model" +@pytest.mark.asyncio +async def test_realtime_driver_session_reports_non_completed_result( + monkeypatch: pytest.MonkeyPatch, +) -> None: + async def fake_run_demo_session_async(**kwargs: Any) -> RunResult: + del kwargs + return RunResult( + status="not_activated", + reason="transport closed before first step", + ) + + monkeypatch.setattr( + manager_module, + "run_demo_session_async", + fake_run_demo_session_async, + ) + runtime = SimpleNamespace() + manager = _make_manager(_BaseTestManager, runtime) + context = manager._shared_run_context(asyncio.get_running_loop()) + reservation = context.admission.try_reserve() + assert reservation is not None + managed, _video_track, _peer, channel = _managed_session(runtime) + managed.reservation = reservation + manager._active_session = managed + + await manager._run_realtime_driver_session( + managed_session=managed, + context=context, + session_input=None, + ) + + assert json.loads(channel.messages[0]) == { + "type": "error", + "message": "transport closed before first step", + } + assert not manager.has_active_session() + + @pytest.mark.asyncio async def test_create_answer_raises_busy_with_subclass_message() -> None: manager = _make_manager( diff --git a/flashdreams/tests/test_webrtc_warmup.py b/flashdreams/tests/test_webrtc_warmup.py new file mode 100644 index 000000000..53ed09f1b --- /dev/null +++ b/flashdreams/tests/test_webrtc_warmup.py @@ -0,0 +1,154 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +from __future__ import annotations + +from types import SimpleNamespace +from typing import Any + +import pytest + +from flashdreams.serving.webrtc import warmup as warmup_module +from flashdreams.serving.webrtc.messages import make_error_payload +from flashdreams.serving.webrtc.warmup import run_loopback_warmup_session + +pytestmark = pytest.mark.ci_cpu + + +@pytest.mark.asyncio +async def test_loopback_warmup_fails_on_server_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + channel = _FakeLoopbackChannel( + incoming_on_send=(make_error_payload("shared driver failed"),) + ) + _install_fake_loopback_peer(monkeypatch, channel=channel) + + with pytest.raises(RuntimeError, match="shared driver failed"): + await run_loopback_warmup_session( + num_chunks=1, + warmup_timeout_s=1.0, + create_answer=_fake_create_answer, + action_payloads=(_step_action(),), + ) + + +@pytest.mark.asyncio +async def test_loopback_warmup_fails_on_early_channel_close( + monkeypatch: pytest.MonkeyPatch, +) -> None: + channel = _FakeLoopbackChannel(close_on_send=True) + _install_fake_loopback_peer(monkeypatch, channel=channel) + + with pytest.raises(RuntimeError, match=r"0/1 chunk"): + await run_loopback_warmup_session( + num_chunks=1, + warmup_timeout_s=1.0, + create_answer=_fake_create_answer, + action_payloads=(_step_action(),), + ) + + +async def _fake_create_answer(*, offer_sdp: str, offer_type: str) -> dict[str, str]: + del offer_sdp, offer_type + return {"sdp": "answer-sdp", "type": "answer"} + + +def _step_action() -> dict[str, object]: + return {"type": "action", "action": {"event": "step"}} + + +def _install_fake_loopback_peer( + monkeypatch: pytest.MonkeyPatch, + *, + channel: "_FakeLoopbackChannel", +) -> None: + monkeypatch.setattr( + warmup_module, + "RTCPeerConnection", + lambda _configuration: _FakeLoopbackPeer(channel=channel), + ) + + async def wait_for_ice_gathering_complete(*args: Any, **kwargs: Any) -> None: + del args, kwargs + + monkeypatch.setattr( + warmup_module, + "wait_for_ice_gathering_complete", + wait_for_ice_gathering_complete, + ) + + +class _FakeLoopbackPeer: + iceGatheringState = "complete" + + def __init__(self, *, channel: "_FakeLoopbackChannel") -> None: + self.localDescription: Any | None = None + self._channel = channel + + def createDataChannel(self, *args: Any, **kwargs: Any) -> "_FakeLoopbackChannel": + del args, kwargs + return self._channel + + def addTransceiver(self, *args: Any, **kwargs: Any) -> None: + del args, kwargs + + def on(self, event_name: str) -> Any: + del event_name + + def decorator(callback: Any) -> Any: + return callback + + return decorator + + async def createOffer(self) -> Any: + return SimpleNamespace(sdp="offer-sdp", type="offer") + + async def setLocalDescription(self, description: Any) -> None: + self.localDescription = description + + async def setRemoteDescription(self, description: Any) -> None: + del description + self._channel.open() + + async def close(self) -> None: + self._channel.close() + + +class _FakeLoopbackChannel: + readyState = "open" + + def __init__( + self, + *, + incoming_on_send: tuple[dict[str, str], ...] = (), + close_on_send: bool = False, + ) -> None: + self._incoming_on_send = list(incoming_on_send) + self._close_on_send = close_on_send + self._handlers: dict[str, Any] = {} + + def on(self, event_name: str) -> Any: + def decorator(callback: Any) -> Any: + self._handlers[event_name] = callback + return callback + + return decorator + + def send(self, message: str) -> None: + del message + if self._incoming_on_send: + self._handlers["message"]( + warmup_module.json.dumps(self._incoming_on_send.pop(0)) + ) + if self._close_on_send: + self.close() + + def open(self) -> None: + self._handlers["open"]() + + def close(self) -> None: + self.readyState = "closed" + close_handler = self._handlers.get("close") + if close_handler is not None: + close_handler() diff --git a/integrations/omnidreams/omnidreams/demo/webrtc.py b/integrations/omnidreams/omnidreams/demo/webrtc.py index 4468eda3d..5fa00869e 100644 --- a/integrations/omnidreams/omnidreams/demo/webrtc.py +++ b/integrations/omnidreams/omnidreams/demo/webrtc.py @@ -66,11 +66,9 @@ WebRTCControlSignal, ) from flashdreams.serving.webrtc.server import create_webrtc_app +from flashdreams.serving.webrtc.services import WEBRTC_USER_INPUT_SCHEMA -from .providers import ( - LudusSceneConditioningProvider, - keyboard_driving_user_input_schema, -) +from .providers import LudusSceneConditioningProvider from .runtime import OmnidreamsRuntime, OmnidreamsRuntimeOptions, PipelineFactory from .spec import ( DEFAULT_OMNIDREAMS_PRESET, @@ -168,7 +166,11 @@ def __init__(self, *, config: OmnidreamsWebRTCModelRuntimeConfig) -> None: runtime_error_type=OmnidreamsWebRTCModelRuntimeError, thread_name="omnidreams-demo-runtime", ) - self.input_source_schema = keyboard_driving_user_input_schema() + # The shared WebRTC input source emits normalized runtime events + # (``key_down``/``key_up``). The Ludus provider consumes sparse + # resampler metadata on this transitional path, so keep validation + # aligned with the WebRTC source rather than the replay trace schema. + self.input_source_schema = WEBRTC_USER_INPUT_SCHEMA self.input_canonicalizer = InputCanonicalizer() self.input_mapping = _OmnidreamsWebRTCInputMapping() self._runtime: OmnidreamsRuntime | None = None diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index e0796b740..5e9c9fbd2 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -3,6 +3,7 @@ from __future__ import annotations +import asyncio import json from collections.abc import Sequence from pathlib import Path @@ -66,8 +67,15 @@ ) from flashdreams.runtime.demo.replay import run_replay_demo from flashdreams.runtime.demo.timing import SPARSE_KEY_SEGMENTS_METADATA_KEY -from flashdreams.serving.webrtc.manager import BaseWebRTCSessionManager +from flashdreams.serving.webrtc.manager import ( + BaseWebRTCSessionManager, + ManagedWebRTCSession, +) from flashdreams.serving.webrtc.server import SESSION_MANAGER_KEY +from flashdreams.serving.webrtc.services import ( + WebRTCInputSource, + WebRTCTransportService, +) pytestmark = pytest.mark.ci_cpu @@ -1081,6 +1089,85 @@ async def test_omnidreams_webrtc_runtime_uses_shared_session( assert rasterizers[0].closed is True +@pytest.mark.asyncio +async def test_omnidreams_webrtc_manager_drives_shared_session( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _scene, rasterizers = _install_fake_ludus_provider_dependencies(monkeypatch) + scene_path = tmp_path / "scene.usdz" + scene_path.write_bytes(b"fake") + pipeline = _VariableFrameOmnidreamsPipeline((2, 3)) + config = OmnidreamsWebRTCModelRuntimeConfig( + pipeline_config_name="fake", + pipeline_config=object(), + pipeline_factory=lambda pipeline_config, device: pipeline, + scene_dir=scene_path, + device="cpu", + fps=30, + video_height=2, + video_width=2, + warmup_chunks=0, + ) + runtime = OmnidreamsWebRTCModelRuntime(config=config) + manager = BaseWebRTCSessionManager( + runtime=runtime, + runtime_config=config, + fps=config.fps, + identity=config.pipeline_config_name, + supported_control_keys=frozenset({"w", "a", "s", "d"}), + ) + await runtime.initialize() + manager._runtime_ready = True + await runtime.reset_for_new_session() + loop = asyncio.get_running_loop() + context = manager._shared_run_context(loop) + reservation = context.admission.try_reserve() + assert reservation is not None + resampler = _FakeWebRTCResampler(start_v=loop.time(), fps=config.fps) + input_source = WebRTCInputSource(resampler=resampler) + input_source.handle_browser_payload( + {"type": "action", "action": {"event": "step"}}, + timestamp_s=loop.time(), + ) + transport = WebRTCTransportService(loop=loop) + channel = _FakeWebRTCChannel() + managed_session = ManagedWebRTCSession( + runtime=runtime, + video_track=_FakeWebRTCVideoTrack(fps=config.fps), # ty:ignore[invalid-argument-type] + video_encoder=_FakeWebRTCVideoEncoder(), # ty:ignore[invalid-argument-type] + peer_connection=_FakeWebRTCPeerConnection(), + resampler=resampler, # ty:ignore[invalid-argument-type] + control_channel=channel, + input_source=input_source, + transport=transport, + reservation=reservation, + last_client_message_at=loop.time(), + ) + manager._active_session = managed_session + try: + managed_session.generation_task = asyncio.create_task( + manager._run_realtime_driver_session( + managed_session=managed_session, + context=context, + session_input=None, + ) + ) + chunk = await _wait_for_chunk_done(channel) + + assert chunk["type"] == "chunk_done" + assert chunk["model"] == "fake" + assert chunk["num_frames"] == 2 + assert [tuple(hdmap.shape) for hdmap in pipeline.generated_hdmaps][:1] == [ + (1, 1, 2, 3, 2, 2) + ] + assert rasterizers[0].calls[0]["timestamps_us"] == (1_000, 34_333) + finally: + transport.close("test complete") + await manager.close_active_session() + await runtime.close() + + class _RecordingOutputTarget: def __init__(self) -> None: self.results: list[StepResult] = [] @@ -1335,6 +1422,98 @@ def cleanup(self) -> None: self.closed = True +class _FakeWebRTCResampler: + def __init__(self, *, start_v: float, fps: int) -> None: + self.dt = 1.0 / fps + self.next_chunk_start_v = start_v + + def reset(self, *, start_v: float) -> None: + self.next_chunk_start_v = start_v + + def on_edge(self, *, arrival_t: float, event: str, key: str) -> None: + del arrival_t, event, key + + def sample_chunk( + self, + num_frames: int, + ) -> tuple[list[tuple[float, float, frozenset[str]]], list[float]]: + start = self.next_chunk_start_v + frame_times = [start + (index + 1) * self.dt for index in range(num_frames)] + end = frame_times[-1] + self.next_chunk_start_v = end + return [(start, end, frozenset({"w"}))], frame_times + + +class _FakeWebRTCVideoTrack: + def __init__(self, *, fps: int) -> None: + self.fps = fps + self.closed = False + self.enqueued: list[StepResult] = [] + + async def enqueue_result(self, result: StepResult) -> int: + self.enqueued.append(result) + return result.frame_count + + def qsize(self) -> int: + return 0 + + async def close(self) -> None: + self.closed = True + + +class _FakeWebRTCVideoEncoder: + backend = "fake" + prefers_codec: str | None = None + + def prepare_chunk_payload(self, result: StepResult, track: Any) -> StepResult: + del track + return result + + async def deliver_prepared_chunk( + self, + payload: object, + track: Any, + *, + force_keyframe: bool = False, + ) -> SimpleNamespace: + del force_keyframe + if not isinstance(payload, StepResult): + raise TypeError("Fake WebRTC encoder expected a StepResult payload.") + return SimpleNamespace( + num_frames=await track.enqueue_result(payload), + encode_ms=0.0, + ) + + +class _FakeWebRTCPeerConnection: + def __init__(self) -> None: + self.closed = False + + async def close(self) -> None: + self.closed = True + + +class _FakeWebRTCChannel: + def __init__(self) -> None: + self.messages: list[str] = [] + + def send(self, message: str) -> None: + self.messages.append(message) + + +async def _wait_for_chunk_done(channel: _FakeWebRTCChannel) -> dict[str, Any]: + for _ in range(100): + chunk_done = [ + json.loads(message) + for message in channel.messages + if json.loads(message).get("type") == "chunk_done" + ] + if chunk_done: + return chunk_done[0] + await asyncio.sleep(0.01) + pytest.fail("Timed out waiting for WebRTC chunk_done.") + + class _FakeWebRTCRuntime: def __init__(self, config: Any) -> None: self.config = config From 8c68f207c42304722fafb40e471da447069120d4 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 02:37:05 +0000 Subject: [PATCH 43/51] Keep OmniDreams WebRTC encoders across warmup --- .../omnidreams/omnidreams/demo/replay.py | 15 ++++-- .../omnidreams/omnidreams/demo/webrtc.py | 5 ++ .../omnidreams/tests/test_demo_api.py | 53 +++++++++++++++++++ 3 files changed, 70 insertions(+), 3 deletions(-) diff --git a/integrations/omnidreams/omnidreams/demo/replay.py b/integrations/omnidreams/omnidreams/demo/replay.py index 06c9cda0d..f7eb60cdd 100644 --- a/integrations/omnidreams/omnidreams/demo/replay.py +++ b/integrations/omnidreams/omnidreams/demo/replay.py @@ -41,6 +41,7 @@ class OmnidreamsRuntimeOptions: pipeline_config: Any pipeline_factory: PipelineFactory | None = None output_layout: VideoTensorLayout = "bvtchw" + release_oneshot_encoders_after_cache_init: bool = True class OmnidreamsRuntime: @@ -84,6 +85,9 @@ def start_session(self, inputs: InferenceInput) -> InferenceSession: is_rank_zero=self.is_rank_zero, output_layout=self.options.output_layout, rollout_seed=self.config.seed, + release_oneshot_encoders_after_cache_init=( + self.options.release_oneshot_encoders_after_cache_init + ), ) def close(self) -> None: @@ -111,6 +115,7 @@ def __init__( is_rank_zero: bool, output_layout: VideoTensorLayout, rollout_seed: int | None, + release_oneshot_encoders_after_cache_init: bool, ) -> None: self.pipeline = pipeline self.scenario = scenario @@ -119,6 +124,9 @@ def __init__( self.is_rank_zero = is_rank_zero self.output_layout = output_layout self.rollout_seed = rollout_seed + self.release_oneshot_encoders_after_cache_init = ( + release_oneshot_encoders_after_cache_init + ) self.dtype = torch.bfloat16 self._closed = False self._model_session = OmnidreamsModelSessionCore( @@ -206,9 +214,10 @@ def _initialize_cache(self) -> Any: ), view_names=_view_names_from_inputs(self._initial_inputs, scenario), ) - release = getattr(self.pipeline, "release_oneshot_encoders", None) - if callable(release): - release() + if self.release_oneshot_encoders_after_cache_init: + release = getattr(self.pipeline, "release_oneshot_encoders", None) + if callable(release): + release() return cache diff --git a/integrations/omnidreams/omnidreams/demo/webrtc.py b/integrations/omnidreams/omnidreams/demo/webrtc.py index 5fa00869e..afcb84f82 100644 --- a/integrations/omnidreams/omnidreams/demo/webrtc.py +++ b/integrations/omnidreams/omnidreams/demo/webrtc.py @@ -218,6 +218,11 @@ def _initialize_sync(self) -> None: options=OmnidreamsRuntimeOptions( pipeline_config=self.config.pipeline_config, pipeline_factory=self.config.pipeline_factory, + # WebRTC warms the same long-lived runtime before real browser + # sessions. Keep prompt/image encoders available for later + # peer connections until Phase 14 replaces loopback warmup with + # first-class model/runtime warmup. + release_oneshot_encoders_after_cache_init=False, ), ) self._initialize_video_encoder_sync() diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index 5e9c9fbd2..52572e40b 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -764,6 +764,7 @@ def test_omnidreams_replay_runtime_generates_video_step_result( assert result.video_chunk.shape == (1, 1, 1, 3, 2, 2) assert result.metrics["denoise_s"] == 0.25 assert session.next_step_request() is None + assert pipeline.released_encoders is True assert pipeline.initialize_cache_calls == [ { "text": [["drive"]], @@ -1089,6 +1090,41 @@ async def test_omnidreams_webrtc_runtime_uses_shared_session( assert rasterizers[0].closed is True +@pytest.mark.asyncio +async def test_omnidreams_webrtc_runtime_keeps_encoders_after_warmup_session( + tmp_path: Path, + monkeypatch: pytest.MonkeyPatch, +) -> None: + _install_fake_ludus_provider_dependencies(monkeypatch) + scene_path = tmp_path / "scene.usdz" + scene_path.write_bytes(b"fake") + pipeline = _FailsIfEncodersReleasedOmnidreamsPipeline() + config = OmnidreamsWebRTCModelRuntimeConfig( + pipeline_config_name="fake", + pipeline_config=object(), + pipeline_factory=lambda pipeline_config, device: pipeline, + scene_dir=scene_path, + device="cpu", + fps=30, + video_height=2, + video_width=2, + warmup_chunks=0, + ) + runtime = OmnidreamsWebRTCModelRuntime(config=config) + await runtime.initialize() + + await runtime.reset_for_new_session() + warmup_session = await runtime.start_inference_session() + warmup_session.close() + await runtime.reset_for_new_session() + browser_session = await runtime.start_inference_session() + + assert pipeline.released_encoders is False + assert len(pipeline.initialize_cache_calls) == 2 + browser_session.close() + await runtime.close() + + @pytest.mark.asyncio async def test_omnidreams_webrtc_manager_drives_shared_session( tmp_path: Path, @@ -1387,6 +1423,23 @@ def generate( return torch.full((1, 1, frame_count, 3, 2, 2), float(autoregressive_index)) +class _FailsIfEncodersReleasedOmnidreamsPipeline(_FakeOmnidreamsPipeline): + def initialize_cache( + self, + *, + text: list[list[str]], + image: torch.Tensor, + view_names: list[str], + ) -> object: + if self.released_encoders: + raise AssertionError("encoders were released before the next session") + return super().initialize_cache( + text=text, + image=image, + view_names=view_names, + ) + + class _FakeLudusRasterizer: def __init__(self) -> None: self.loaded_scene: object | None = None From a2ae69bc5793b74312c864d04309d1dbf322afcd Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 02:54:32 +0000 Subject: [PATCH 44/51] Cap OmniDreams WebRTC video display size --- .../omnidreams/demo/web/adapter.css | 15 ++++++++++++++ .../omnidreams/omnidreams/demo/web/adapter.js | 1 + integrations/omnidreams/pyproject.toml | 2 +- .../omnidreams/tests/test_demo_api.py | 20 +++++++++++++++++++ 4 files changed, 37 insertions(+), 1 deletion(-) create mode 100644 integrations/omnidreams/omnidreams/demo/web/adapter.css diff --git a/integrations/omnidreams/omnidreams/demo/web/adapter.css b/integrations/omnidreams/omnidreams/demo/web/adapter.css new file mode 100644 index 000000000..6fb6d6132 --- /dev/null +++ b/integrations/omnidreams/omnidreams/demo/web/adapter.css @@ -0,0 +1,15 @@ +/* +SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +SPDX-License-Identifier: Apache-2.0 +*/ + +/* Keep the browser from enlarging OmniDreams' native 1280x704 stream. */ +.stageVideo { + inset: 50% auto auto 50%; + width: min(100vw, 1280px, calc(100vh * 1280 / 704)); + height: min(100vh, 704px, calc(100vw * 704 / 1280)); + aspect-ratio: 1280 / 704; + transform: translate(-50%, -50%); + object-fit: contain; + object-position: center; +} diff --git a/integrations/omnidreams/omnidreams/demo/web/adapter.js b/integrations/omnidreams/omnidreams/demo/web/adapter.js index d07fb8cc1..a34ed03f7 100644 --- a/integrations/omnidreams/omnidreams/demo/web/adapter.js +++ b/integrations/omnidreams/omnidreams/demo/web/adapter.js @@ -3,6 +3,7 @@ export default { modelName: "OmniDreams", + stylesheet: "/model-static/adapter.css?v=omnidreams-ui-v1", controls: [ { label: "Drive / Turn", diff --git a/integrations/omnidreams/pyproject.toml b/integrations/omnidreams/pyproject.toml index 221cb042f..fe48be188 100644 --- a/integrations/omnidreams/pyproject.toml +++ b/integrations/omnidreams/pyproject.toml @@ -138,7 +138,7 @@ exclude = ["tests"] # workspace editable. Editable installs pick these up from the source # tree automatically. [tool.setuptools.package-data] -"omnidreams.demo" = ["web/adapter.js"] +"omnidreams.demo" = ["web/adapter.js", "web/adapter.css"] "omnidreams.interactive_drive" = [ "configs/*.yaml", "configs/wheels/*.yaml", diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index 52572e40b..acdb65eed 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -14,6 +14,7 @@ import omnidreams.demo as demo_package import omnidreams.demo.spec as spec_module import pytest +import tomli as tomllib import torch from aiohttp import web from omnidreams.config import OMNIDREAMS_RUNNERS @@ -951,6 +952,25 @@ def fake_create_packaged_webrtc_app(**kwargs: Any) -> web.Application: assert app_calls[0]["configure_app"] is None +def test_omnidreams_webrtc_adapter_caps_video_display_size() -> None: + web_dir = Path(demo_package.__file__).resolve().parent / "web" + adapter_js = (web_dir / "adapter.js").read_text(encoding="utf-8") + adapter_css = (web_dir / "adapter.css").read_text(encoding="utf-8") + + assert 'stylesheet: "/model-static/adapter.css?v=omnidreams-ui-v1"' in adapter_js + assert ".stageVideo" in adapter_css + assert "1280px" in adapter_css + assert "704px" in adapter_css + assert "object-fit: contain" in adapter_css + + pyproject = Path(__file__).resolve().parents[1] / "pyproject.toml" + with pyproject.open("rb") as fh: + meta = tomllib.load(fh) + package_data = meta["tool"]["setuptools"]["package-data"]["omnidreams.demo"] + assert "web/adapter.js" in package_data + assert "web/adapter.css" in package_data + + def test_omnidreams_webrtc_demo_serves_through_shared_runner( monkeypatch: pytest.MonkeyPatch, ) -> None: From 16cd15d82608078bc4669a8278aaa7b510f57f38 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 03:01:02 +0000 Subject: [PATCH 45/51] Load OmniDreams WebRTC video sizing stylesheet reliably --- .../flashdreams/serving/webrtc/server.py | 8 ++-- .../serving/webrtc/web/mock_ui_server.py | 16 ++++---- .../serving/webrtc/web/request_session.html | 2 +- .../serving/webrtc/web/request_session.js | 9 ++++- flashdreams/tests/test_webrtc_serving.py | 39 ++++++++++++++++++- .../omnidreams/demo/web/adapter.css | 5 ++- .../omnidreams/omnidreams/demo/web/adapter.js | 2 +- .../omnidreams/tests/test_demo_api.py | 3 +- 8 files changed, 66 insertions(+), 18 deletions(-) diff --git a/flashdreams/flashdreams/serving/webrtc/server.py b/flashdreams/flashdreams/serving/webrtc/server.py index bf257701d..7a2c6694d 100644 --- a/flashdreams/flashdreams/serving/webrtc/server.py +++ b/flashdreams/flashdreams/serving/webrtc/server.py @@ -91,10 +91,12 @@ async def healthz(request: web.Request) -> web.StreamResponse: ) async def ui_config(_: web.Request) -> web.StreamResponse: - adapter_module = None + payload: dict[str, str | None] = {"adapter_module": None} if model_web_dir is not None and (model_web_dir / "adapter.js").is_file(): - adapter_module = "/model-static/adapter.js?v=model-ui-v1" - return web.json_response({"adapter_module": adapter_module}) + payload["adapter_module"] = "/model-static/adapter.js?v=model-ui-v2" + if model_web_dir is not None and (model_web_dir / "adapter.css").is_file(): + payload["model_stylesheet"] = "/model-static/adapter.css?v=model-ui-v2" + return web.json_response(payload) async def on_startup(app: web.Application) -> None: manager = app[SESSION_MANAGER_KEY] diff --git a/flashdreams/flashdreams/serving/webrtc/web/mock_ui_server.py b/flashdreams/flashdreams/serving/webrtc/web/mock_ui_server.py index 735179a4a..0c62855b5 100644 --- a/flashdreams/flashdreams/serving/webrtc/web/mock_ui_server.py +++ b/flashdreams/flashdreams/serving/webrtc/web/mock_ui_server.py @@ -83,13 +83,15 @@ def do_HEAD(self) -> None: def _serve_ui_config(self) -> bool: if urlsplit(self.path).path != "/api/ui/config": return False - adapter_module = ( - "/model-static/adapter.js?v=model-ui-v1" - if self.model_web_dir is not None - and (self.model_web_dir / "adapter.js").is_file() - else None - ) - payload = json.dumps({"adapter_module": adapter_module}).encode("utf-8") + ui_config: dict[str, str | None] = {"adapter_module": None} + if self.model_web_dir is not None: + if (self.model_web_dir / "adapter.js").is_file(): + ui_config["adapter_module"] = "/model-static/adapter.js?v=model-ui-v2" + if (self.model_web_dir / "adapter.css").is_file(): + ui_config["model_stylesheet"] = ( + "/model-static/adapter.css?v=model-ui-v2" + ) + payload = json.dumps(ui_config).encode("utf-8") self.send_response(200) self.send_header("Content-Type", "application/json") self.send_header("Content-Length", str(len(payload))) diff --git a/flashdreams/flashdreams/serving/webrtc/web/request_session.html b/flashdreams/flashdreams/serving/webrtc/web/request_session.html index ee42c82a0..ad158656e 100644 --- a/flashdreams/flashdreams/serving/webrtc/web/request_session.html +++ b/flashdreams/flashdreams/serving/webrtc/web/request_session.html @@ -87,6 +87,6 @@

Client Logs

- + diff --git a/flashdreams/flashdreams/serving/webrtc/web/request_session.js b/flashdreams/flashdreams/serving/webrtc/web/request_session.js index 95c340b2d..a13384bc5 100644 --- a/flashdreams/flashdreams/serving/webrtc/web/request_session.js +++ b/flashdreams/flashdreams/serving/webrtc/web/request_session.js @@ -291,10 +291,14 @@ const modelContext = { async function loadModelAdapter() { let adapter = {} + const stylesheetHrefs = new Set() try { const response = await fetch("/api/ui/config") if (response.ok) { const config = await response.json() + if (typeof config.model_stylesheet === "string" && config.model_stylesheet) { + stylesheetHrefs.add(config.model_stylesheet) + } if (typeof config.adapter_module === "string" && config.adapter_module) { const module = await import(config.adapter_module) if (module.default && typeof module.default === "object") { @@ -308,9 +312,12 @@ async function loadModelAdapter() { modelAdapter = adapter if (typeof adapter.stylesheet === "string" && adapter.stylesheet) { + stylesheetHrefs.add(adapter.stylesheet) + } + for (const href of stylesheetHrefs) { const stylesheet = document.createElement("link") stylesheet.rel = "stylesheet" - stylesheet.href = adapter.stylesheet + stylesheet.href = href document.head.append(stylesheet) } const modelControls = Array.isArray(adapter.controls) ? adapter.controls : [] diff --git a/flashdreams/tests/test_webrtc_serving.py b/flashdreams/tests/test_webrtc_serving.py index 21d789fe0..06e003496 100644 --- a/flashdreams/tests/test_webrtc_serving.py +++ b/flashdreams/tests/test_webrtc_serving.py @@ -332,7 +332,7 @@ def test_shared_viewer_exposes_model_extension_slots() -> None: html = web_dir.joinpath("request_session.html").read_text(encoding="utf-8") javascript = web_dir.joinpath("request_session.js").read_text(encoding="utf-8") - assert "/static/request_session.js?v=shared-webrtc-v3" in html + assert "/static/request_session.js?v=shared-webrtc-v4" in html for slot in ( "modelStageSlot", "modelStatusSlot", @@ -341,6 +341,8 @@ def test_shared_viewer_exposes_model_extension_slots() -> None: ): assert f'id="{slot}"' in html assert 'fetch("/api/ui/config")' in javascript + assert "config.model_stylesheet" in javascript + assert "stylesheetHrefs" in javascript assert "await modelAdapter?.beforeConnect?.(modelContext)" in javascript assert "sendCommand: sendModelCommand" in javascript assert 'id="postprocessField"' in html @@ -399,7 +401,7 @@ async def test_packaged_webrtc_app_serves_model_adapter(tmp_path) -> None: try: config_response = await client.get("/api/ui/config") assert await config_response.json() == { - "adapter_module": "/model-static/adapter.js?v=model-ui-v1" + "adapter_module": "/model-static/adapter.js?v=model-ui-v2" } adapter_response = await client.get("/model-static/adapter.js") assert adapter_response.status == 200 @@ -408,6 +410,39 @@ async def test_packaged_webrtc_app_serves_model_adapter(tmp_path) -> None: await client.close() +@pytest.mark.asyncio +async def test_packaged_webrtc_app_serves_model_stylesheet(tmp_path) -> None: + shared_dir = tmp_path / "shared" + model_dir = tmp_path / "model" + shared_dir.mkdir() + model_dir.mkdir() + (shared_dir / "request_session.html").write_text("session") + (model_dir / "adapter.css").write_text(".stageVideo { object-fit: contain; }") + app = create_packaged_webrtc_app( + web_resource=shared_dir, + model_web_resource=model_dir, + session_manager=_FakeSessionManager(), + request_session_url="http://127.0.0.1:8080/request_session", + preload_name="Test", + as_file_fn=lambda resource: nullcontext(resource), + ) + client = TestClient(TestServer(app)) + await client.start_server() + try: + config_response = await client.get("/api/ui/config") + assert await config_response.json() == { + "adapter_module": None, + "model_stylesheet": "/model-static/adapter.css?v=model-ui-v2", + } + stylesheet_response = await client.get("/model-static/adapter.css") + assert stylesheet_response.status == 200 + assert ( + await stylesheet_response.text() == ".stageVideo { object-fit: contain; }" + ) + finally: + await client.close() + + def test_webrtc_message_helpers_preserve_public_payload_shape() -> None: assert make_error_payload("boom") == {"type": "error", "message": "boom"} assert make_event_ack_payload( diff --git a/integrations/omnidreams/omnidreams/demo/web/adapter.css b/integrations/omnidreams/omnidreams/demo/web/adapter.css index 6fb6d6132..f4cb8868b 100644 --- a/integrations/omnidreams/omnidreams/demo/web/adapter.css +++ b/integrations/omnidreams/omnidreams/demo/web/adapter.css @@ -6,8 +6,9 @@ SPDX-License-Identifier: Apache-2.0 /* Keep the browser from enlarging OmniDreams' native 1280x704 stream. */ .stageVideo { inset: 50% auto auto 50%; - width: min(100vw, 1280px, calc(100vh * 1280 / 704)); - height: min(100vh, 704px, calc(100vw * 704 / 1280)); + width: min(100vw, 1280px, 181.82vh); + height: auto; + max-height: min(100vh, 704px); aspect-ratio: 1280 / 704; transform: translate(-50%, -50%); object-fit: contain; diff --git a/integrations/omnidreams/omnidreams/demo/web/adapter.js b/integrations/omnidreams/omnidreams/demo/web/adapter.js index a34ed03f7..37d19a299 100644 --- a/integrations/omnidreams/omnidreams/demo/web/adapter.js +++ b/integrations/omnidreams/omnidreams/demo/web/adapter.js @@ -3,7 +3,7 @@ export default { modelName: "OmniDreams", - stylesheet: "/model-static/adapter.css?v=omnidreams-ui-v1", + stylesheet: "/model-static/adapter.css?v=model-ui-v2", controls: [ { label: "Drive / Turn", diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index acdb65eed..4a2e89d3a 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -957,10 +957,11 @@ def test_omnidreams_webrtc_adapter_caps_video_display_size() -> None: adapter_js = (web_dir / "adapter.js").read_text(encoding="utf-8") adapter_css = (web_dir / "adapter.css").read_text(encoding="utf-8") - assert 'stylesheet: "/model-static/adapter.css?v=omnidreams-ui-v1"' in adapter_js + assert 'stylesheet: "/model-static/adapter.css?v=model-ui-v2"' in adapter_js assert ".stageVideo" in adapter_css assert "1280px" in adapter_css assert "704px" in adapter_css + assert "calc(" not in adapter_css assert "object-fit: contain" in adapter_css pyproject = Path(__file__).resolve().parents[1] / "pyproject.toml" From ea817a4fc13ce0c6ff9e553c08efa5720a625de2 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 03:14:12 +0000 Subject: [PATCH 46/51] Close async demo provider on pre-driver cancellation --- .../flashdreams/runtime/demo/drivers.py | 16 ++--- .../tests/test_demo_runtime_vertical_slice.py | 70 +++++++++++++++++++ 2 files changed, 78 insertions(+), 8 deletions(-) diff --git a/flashdreams/flashdreams/runtime/demo/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py index 087135c63..9bdd40db8 100644 --- a/flashdreams/flashdreams/runtime/demo/drivers.py +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -453,7 +453,6 @@ async def run_demo_session_async( provider: Any | None = None session_edges: SessionEdges | None = None - driver_started = False try: try: create_provider = getattr(adapter, "create_model_input_provider") @@ -490,7 +489,6 @@ async def run_demo_session_async( "must not be reused." ) driver = run_mode.select_driver() - driver_started = True result = await _run_async_driver( driver=driver, host=context.host, @@ -509,15 +507,13 @@ async def run_demo_session_async( status="cancelled", reason="cancelled during session assembly", error=None, - close_provider=not driver_started, + close_provider=_needs_partial_provider_cleanup(session_edges), ) context.run_metrics.record_session(result) return result except DriverInvariantError as exc: _record_run_session_error(context, exc) - should_record_session = session_edges is not None and ( - driver_started or not session_edges.is_closed - ) + should_record_session = session_edges is not None result = await _close_partial_session_async( context=context, provider=provider, @@ -525,7 +521,7 @@ async def run_demo_session_async( status="failed", reason=str(exc), error=exc, - close_provider=not driver_started, + close_provider=_needs_partial_provider_cleanup(session_edges), ) if should_record_session: context.run_metrics.record_session(result) @@ -539,7 +535,7 @@ async def run_demo_session_async( status="failed", reason=str(exc), error=exc, - close_provider=not driver_started, + close_provider=_needs_partial_provider_cleanup(session_edges), ) context.run_metrics.record_session(result) return result @@ -822,6 +818,10 @@ async def _close_partial_session_async( return RunResult(status=status, reason=reason, error=error) +def _needs_partial_provider_cleanup(session_edges: SessionEdges | None) -> bool: + return session_edges is None or not session_edges.is_closed + + async def _close_provider_async( *, context: RunContext, diff --git a/flashdreams/tests/test_demo_runtime_vertical_slice.py b/flashdreams/tests/test_demo_runtime_vertical_slice.py index 68f887d3a..1f430a459 100644 --- a/flashdreams/tests/test_demo_runtime_vertical_slice.py +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -10,6 +10,7 @@ import pytest +import flashdreams.runtime.demo.drivers as drivers_module from flashdreams.runtime import ( CanonicalInputSchema, IdentityInputMapping, @@ -460,6 +461,75 @@ async def test_run_demo_session_async_invariant_cancellation_finalizes_edges() - assert runtime.start_session_inputs == [] +@pytest.mark.asyncio +async def test_run_demo_session_async_cancels_before_driver_owns_cleanup( + monkeypatch: pytest.MonkeyPatch, +) -> None: + runtime = _FakeVideoRuntime(session=_FakeVideoSession(num_steps=1)) + run_metrics = InMemorySessionMetricsRecorder() + context = _run_context(runtime, run_metrics=run_metrics) + provider = _FakeVideoModelInputProvider() + output = _RecordingOutputSink() + transport = _RecordingTransport() + session_metrics = InMemorySessionMetricsRecorder() + driver_boundary_reached = asyncio.Event() + release_driver = asyncio.Event() + + async def fake_run_async_driver( + *, + driver: object, + host: RuntimeHost, + provider: Any, + session_edges: SessionEdges, + pipeline: StepPipeline, + ) -> RunResult: + del driver, host, provider, session_edges, pipeline + driver_boundary_reached.set() + await release_driver.wait() + return RunResult(status="completed") + + monkeypatch.setattr( + drivers_module, + "_run_async_driver", + fake_run_async_driver, + ) + task = asyncio.create_task( + run_demo_session_async( + context=context, + spec=_spec(), + scenario=_scenario(), + adapter=_FakeDemoAdapter(provider=provider), + run_mode=_FakeRunMode( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=output, + metrics=session_metrics, + transport=transport, + ), + pipeline=StepPipeline(), + ) + ) + + try: + await asyncio.wait_for(driver_boundary_reached.wait(), timeout=2.0) + task.cancel() + result = await task + finally: + release_driver.set() + if not task.done(): + task.cancel() + await asyncio.gather(task, return_exceptions=True) + context.host.close() + + assert result.status == "cancelled" + assert result.reason == "cancelled during session assembly" + assert provider.close_count == 1 + assert output.close_count == 1 + assert transport.close_count == 1 + assert session_metrics.closed + assert run_metrics.sessions == [result] + assert runtime.start_session_inputs == [] + + def test_setup_failure_can_return_skipped_but_not_completed() -> None: skipped = BatchSessionDriver().run_one_session( host=RuntimeHost(_FakeVideoRuntime(session=_FakeVideoSession(num_steps=1))), From 18e4d36c585bd96019de75ebed40c49e86c2c270 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 03:58:46 +0000 Subject: [PATCH 47/51] Route OmniDreams WebRTC through shared realtime driver --- flashdreams/flashdreams/runtime/worker.py | 43 ++++- .../flashdreams/serving/webrtc/manager.py | 166 +++++++++++++--- .../omnidreams/omnidreams/demo/adapter.py | 125 ++++++++++-- .../omnidreams/omnidreams/demo/providers.py | 2 +- .../omnidreams/omnidreams/demo/webrtc.py | 182 +++++++++++++++++- .../omnidreams/tests/test_demo_api.py | 172 +++++++++++++---- 6 files changed, 598 insertions(+), 92 deletions(-) diff --git a/flashdreams/flashdreams/runtime/worker.py b/flashdreams/flashdreams/runtime/worker.py index a545054b0..e523b6e6a 100644 --- a/flashdreams/flashdreams/runtime/worker.py +++ b/flashdreams/flashdreams/runtime/worker.py @@ -13,6 +13,7 @@ import torch _T = TypeVar("_T") +_EXECUTOR_FUTURE_POLL_INTERVAL_S = 0.01 class ModelExecutionWorker: @@ -68,7 +69,7 @@ async def call( self._require_accepting() future = self._submit(func, args, kwargs) try: - return await asyncio.shield(future) + return await _await_executor_future(future) except asyncio.CancelledError: future.add_done_callback(_consume_exception) raise @@ -89,22 +90,24 @@ def call_blocking( async def close(self) -> None: """Drain submitted work and stop accepting lifecycle calls.""" self._require_not_worker_thread() - await asyncio.to_thread(self.close_blocking) + if not self._begin_close(): + return + try: + barrier = asyncio.wrap_future(self._executor.submit(_noop)) + await _await_executor_future(barrier) + finally: + self._finish_close() def close_blocking(self) -> None: """Synchronous close for non-async setup and teardown paths.""" self._require_not_worker_thread() - with self._state_lock: - if self._closed: - return - self._accepting = False + if not self._begin_close(): + return try: barrier = self._executor.submit(_noop) barrier.result() finally: - self._executor.shutdown(wait=True, cancel_futures=False) - with self._state_lock: - self._closed = True + self._finish_close() def _submit( self, @@ -131,6 +134,18 @@ def _require_not_worker_thread(self) -> None: "thread; call the function directly." ) + def _begin_close(self) -> bool: + with self._state_lock: + if self._closed: + return False + self._accepting = False + return True + + def _finish_close(self) -> None: + self._executor.shutdown(wait=True, cancel_futures=False) + with self._state_lock: + self._closed = True + ThreadAffineRuntimeWorker = ModelExecutionWorker @@ -147,6 +162,16 @@ def _noop() -> None: return +async def _await_executor_future(future: asyncio.Future[_T]) -> _T: + """Await an executor future without relying on a single cross-thread wakeup.""" + while not future.done(): + await asyncio.wait( + {future}, + timeout=_EXECUTOR_FUTURE_POLL_INTERVAL_S, + ) + return future.result() + + def _consume_exception(future: asyncio.Future[Any]) -> None: if not future.cancelled(): future.exception() diff --git a/flashdreams/flashdreams/serving/webrtc/manager.py b/flashdreams/flashdreams/serving/webrtc/manager.py index 7e065a72c..4291cd637 100644 --- a/flashdreams/flashdreams/serving/webrtc/manager.py +++ b/flashdreams/flashdreams/serving/webrtc/manager.py @@ -10,10 +10,10 @@ import inspect import json from collections import deque -from collections.abc import Mapping +from collections.abc import Callable, Mapping from collections.abc import Set as AbstractSet from dataclasses import dataclass, field, replace -from typing import Any, Generic, TypeVar +from typing import Any, Generic, TypeVar, cast from aiortc import ( RTCConfiguration, @@ -62,7 +62,9 @@ ) from flashdreams.serving.webrtc.encoders import ( DefaultRTCEncoder, + EncoderBackend, VideoEncoder, + select_encoder, ) from flashdreams.serving.webrtc.media import BufferedVideoTrack, NVENCVideoTrack from flashdreams.serving.webrtc.messages import ( @@ -120,7 +122,7 @@ _SEGMENTS_KEY = "webrtc_segments" _FRAME_TIMES_KEY = "webrtc_frame_times" -_RuntimeT = TypeVar("_RuntimeT", bound=WebRTCSessionRuntime) +_RuntimeT = TypeVar("_RuntimeT") _RuntimeConfigT = TypeVar("_RuntimeConfigT", bound=WebRTCRuntimeConfig) @@ -208,6 +210,27 @@ def _step_request_from_requirements( ) +def _encoder_backend_from_config(value: object) -> EncoderBackend: + backend = str(value) + if backend not in {"auto", "default", "nvenc"}: + raise ValueError( + f"encoder_backend must be 'auto', 'default', or 'nvenc', got {backend!r}." + ) + return cast(EncoderBackend, backend) + + +def _gpu_id_from_device_spec(device_spec: str) -> int: + if not device_spec.startswith("cuda"): + return 0 + _prefix, separator, index = device_spec.partition(":") + if not separator or not index: + return 0 + try: + return int(index) + except ValueError: + return 0 + + class _LegacyWebRTCRuntimeAdapter: """Shared compatibility adapter from old async WebRTC runtimes to RuntimeHost.""" @@ -655,6 +678,11 @@ def __init__( supported_control_keys: AbstractSet[str] | None = None, fatal_generation_errors: bool = False, client_liveness_timeout_s: float = DEFAULT_CLIENT_LIVENESS_TIMEOUT_S, + shared_host: RuntimeHost | None = None, + shared_adapter: Any | None = None, + shared_spec: DemoSpec | None = None, + shared_scenario: PreparedScenario | None = None, + shared_pipeline_factory: Callable[[], StepPipeline] | None = None, ) -> None: if client_liveness_timeout_s <= 0: raise ValueError("client_liveness_timeout_s must be > 0") @@ -678,8 +706,14 @@ def __init__( self._session_lock = asyncio.Lock() self._pending_session_input: Any = None self._shared_runtime_adapter: _LegacyWebRTCRuntimeAdapter | None = None - self._shared_host: RuntimeHost | None = None + self._shared_host: RuntimeHost | None = shared_host + self._owns_shared_host = shared_host is not None self._shared_context: RunContext | None = None + self._shared_adapter = shared_adapter + self._shared_spec = shared_spec + self._shared_scenario = shared_scenario + self._shared_pipeline_factory = shared_pipeline_factory + self._shared_video_encoder: VideoEncoder | None = None @property def pending_session_input(self) -> Any: @@ -742,8 +776,11 @@ def _positive_float_runtime_value(value: Any, *, label: str) -> float: return parsed def _runtime_input_fps(self, runtime: Any) -> float: + peek_input_fps = getattr(runtime, "peek_input_fps", None) + if not callable(peek_input_fps): + return float(self.fps) return self._positive_float_runtime_value( - runtime.peek_input_fps(), + peek_input_fps(), label="peek_input_fps", ) @@ -761,9 +798,22 @@ def _runtime_next_step_request(self, runtime: Any) -> tuple[StepRequest, int]: return request, input_num_frames def _runtime_steady_output_num_frames(self, runtime: Any) -> int: + peek_output_frames = getattr(runtime, "peek_steady_output_num_frames", None) + if callable(peek_output_frames): + return self._positive_int_runtime_value( + peek_output_frames(), + label="peek_steady_output_num_frames", + ) + pipeline = getattr(runtime, "pipeline", None) + get_num_frames = getattr(pipeline, "get_num_frames", None) + if callable(get_num_frames): + return self._positive_int_runtime_value( + get_num_frames(1), + label="pipeline.get_num_frames(1)", + ) return self._positive_int_runtime_value( - runtime.peek_steady_output_num_frames(), - label="peek_steady_output_num_frames", + 1, + label="fallback steady output frame count", ) def _resolve_video_encoder(self) -> VideoEncoder: @@ -776,6 +826,8 @@ def _resolve_video_encoder(self) -> VideoEncoder: transparently get the software path without having to opt in. """ encoder = getattr(self._runtime, "video_encoder", None) + if encoder is None: + encoder = self._shared_video_encoder if encoder is None: encoder = DefaultRTCEncoder(fps=self.fps) return encoder @@ -783,13 +835,15 @@ def _resolve_video_encoder(self) -> VideoEncoder: def _shared_run_context(self, loop: asyncio.AbstractEventLoop) -> RunContext: if self._shared_context is not None: return self._shared_context - runtime_adapter = _LegacyWebRTCRuntimeAdapter( - runtime=self._runtime, - loop=loop, - ) - host = RuntimeHost(runtime_adapter) - self._shared_runtime_adapter = runtime_adapter - self._shared_host = host + host = self._shared_host + if host is None: + runtime_adapter = _LegacyWebRTCRuntimeAdapter( + runtime=self._runtime, + loop=loop, + ) + host = RuntimeHost(runtime_adapter) + self._shared_runtime_adapter = runtime_adapter + self._shared_host = host self._shared_context = RunContext( host=host, run_metrics=InMemorySessionMetricsRecorder(), @@ -821,6 +875,8 @@ async def _reset_runtime_for_session( ) -> None: reset = getattr(context.host.runtime, "reset_for_new_session", None) if not callable(reset): + if self._shared_adapter is not None: + return raise RuntimeError("WebRTC runtime adapter cannot reset sessions.") await context.host.call_async(reset, session_input) @@ -1259,14 +1315,46 @@ def is_runtime_ready(self) -> bool: async def preload_runtime(self) -> None: async with self._preload_lock: if not self._runtime_ready: - await self._runtime.initialize() + initialize = getattr(self._runtime, "initialize", None) + if callable(initialize): + result = initialize() + if inspect.isawaitable(result): + await result + elif self._shared_host is not None: + await asyncio.to_thread(self._shared_host.preload) self._runtime_ready = True + self._initialize_shared_video_encoder() if not self._warmup_complete: await self._run_loopback_warmup_session( num_chunks=self.runtime_config.warmup_chunks ) self._warmup_complete = True + def _initialize_shared_video_encoder(self) -> None: + if self._shared_video_encoder is not None: + return + if getattr(self._runtime, "video_encoder", None) is not None: + return + encoder_backend = getattr(self.runtime_config, "encoder_backend", None) + if encoder_backend is None: + return + backend = _encoder_backend_from_config(encoder_backend) + device_spec = str(getattr(self.runtime_config, "device", "")) + device_type = device_spec.split(":", maxsplit=1)[0] + if device_type != "cuda" and backend == "auto": + backend = "default" + if device_type != "cuda" and backend == "nvenc": + raise RuntimeError("encoder_backend='nvenc' requires a CUDA device.") + self._shared_video_encoder = select_encoder( + backend=backend, + width=self.runtime_config.video_width, + height=self.runtime_config.video_height, + fps=self.fps, + bitrate=int(getattr(self.runtime_config, "encoder_bitrate_bps", 6_000_000)), + gpu_id=_gpu_id_from_device_spec(device_spec), + gop=int(getattr(self.runtime_config, "encoder_gop", self.fps)), + ) + async def create_answer(self, *, offer_sdp: str, offer_type: str) -> dict[str, str]: if not self._runtime_ready or not self._warmup_complete: await self.preload_runtime() @@ -1510,19 +1598,34 @@ async def shutdown(self) -> None: if self._shared_context is not None: await self._shared_context.close_async() if self._shared_host is not None: - self._shared_host.close() + await asyncio.to_thread(self._shared_host.close) self._shared_context = None self._shared_host = None self._shared_runtime_adapter = None - await self._runtime.close() + if self._shared_video_encoder is not None: + self._shared_video_encoder.close() + self._shared_video_encoder = None + if not self._owns_shared_host: + close = getattr(self._runtime, "close", None) + if callable(close): + result = close() + if inspect.isawaitable(result): + await result self._runtime_ready = False self._warmup_complete = False def wait_for_termination(self) -> None: - self._runtime.wait_for_termination() + wait = getattr(self._runtime, "wait_for_termination", None) + if callable(wait): + wait() + return + if self._shared_host is not None: + self._shared_host.run_worker_loop() def send_exit_signal(self) -> None: - self._runtime.send_exit_signal() + send = getattr(self._runtime, "send_exit_signal", None) + if callable(send): + send() async def _handle_datachannel_message( self, @@ -1738,13 +1841,18 @@ async def _run_realtime_driver_session( context: RunContext, session_input: Any, ) -> None: - adapter = _LegacyWebRTCDemoAdapter( - runtime=self._runtime, - identity=self.identity, - session_input=session_input, - ) - spec = self._shared_demo_spec() - scenario = adapter.prepare_scenario(spec) + adapter = self._shared_adapter + spec = self._shared_spec + scenario = self._shared_scenario + if adapter is None or spec is None: + adapter = _LegacyWebRTCDemoAdapter( + runtime=self._runtime, + identity=self.identity, + session_input=session_input, + ) + spec = self._shared_demo_spec() + if scenario is None: + scenario = adapter.prepare_scenario(spec) run_mode = WebRTCRunMode( edge_factory=_ManagedWebRTCSessionEdgeFactory( manager=self, @@ -1759,7 +1867,11 @@ async def _run_realtime_driver_session( scenario=scenario, adapter=adapter, run_mode=run_mode, - pipeline=StepPipeline(), + pipeline=( + self._shared_pipeline_factory() + if self._shared_pipeline_factory is not None + else StepPipeline() + ), reservation=managed_session.reservation, ) if result.status == "completed": diff --git a/integrations/omnidreams/omnidreams/demo/adapter.py b/integrations/omnidreams/omnidreams/demo/adapter.py index 3a7e58dc4..4dc7d1c87 100644 --- a/integrations/omnidreams/omnidreams/demo/adapter.py +++ b/integrations/omnidreams/omnidreams/demo/adapter.py @@ -5,8 +5,9 @@ from __future__ import annotations +import math from collections.abc import Callable -from typing import Any +from typing import Any, cast from omnidreams.config import OMNIDREAMS_CONFIGS, OMNIDREAMS_RUNNERS @@ -25,6 +26,7 @@ ) from flashdreams.runtime.demo.session_inputs import ModelInputProvider from flashdreams.runtime.interfaces import InferenceRuntime +from flashdreams.serving.webrtc.services import WEBRTC_USER_INPUT_SCHEMA from .providers import ( LudusSceneConditioningProvider, @@ -39,13 +41,18 @@ ) from .spec import ( DEFAULT_OMNIDREAMS_PRESET, + DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, OMNIDREAMS_CONDITIONING_LUDUS, OMNIDREAMS_CONDITIONING_MODES, OMNIDREAMS_CONDITIONING_PRECOMPUTED, OMNIDREAMS_MODEL_ID, + LudusBackendName, + OmnidreamsLudusReplayScenario, + OmnidreamsWebRTCScenario, conditioning_mode_from_scenario, resolve_ludus_replay_scenario, resolve_replay_scenario, + resolve_webrtc_scenario, ) RuntimeFactory = Callable[..., InferenceRuntime] @@ -90,25 +97,34 @@ def default_input_mapping(self) -> IdentityInputMapping: return self._mapping def supported_input_modes(self) -> tuple[str, ...]: - return ("replay",) + return ("replay", "keyboard-driving") def supported_output_modes(self) -> tuple[str, ...]: - return ("mp4", "null") + return ("mp4", "null", "webrtc") def supported_conditioning_modes(self) -> tuple[str, ...]: return OMNIDREAMS_CONDITIONING_MODES def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: - if spec.input_mode != "replay": - raise ValueError( - "OmniDreams prepare_scenario currently supports only " - f"input_mode='replay', got {spec.input_mode!r}." - ) if spec.output.mode not in self.supported_output_modes(): raise ValueError( - "OmniDreams replay demo supports output modes " + "OmniDreams demo supports output modes " f"{self.supported_output_modes()}, got {spec.output.mode!r}." ) + if spec.input_mode == "keyboard-driving": + if spec.output.mode != "webrtc": + raise ValueError( + "OmniDreams keyboard-driving input currently requires " + f"output.mode='webrtc', got {spec.output.mode!r}." + ) + return self._prepare_webrtc_scenario(spec) + if spec.input_mode != "replay": + raise ValueError( + "OmniDreams prepare_scenario supports input modes " + f"{self.supported_input_modes()}, got {spec.input_mode!r}." + ) + if spec.output.mode == "webrtc": + raise ValueError("OmniDreams replay input does not support WebRTC output.") conditioning_mode = conditioning_mode_from_scenario(spec.scenario) if conditioning_mode == OMNIDREAMS_CONDITIONING_PRECOMPUTED: scenario = resolve_replay_scenario( @@ -153,6 +169,13 @@ def create_runtime(self, config: InferenceConfig) -> InferenceRuntime: options=OmnidreamsRuntimeOptions( pipeline_config=self._pipeline_config(config), pipeline_factory=self._pipeline_factory, + release_oneshot_encoders_after_cache_init=( + _bool_runtime_option( + config.runtime_options, + "release_oneshot_encoders_after_cache_init", + True, + ) + ), ), ) @@ -161,10 +184,10 @@ def create_model_input_provider( spec: DemoSpec, scenario: PreparedScenario, ) -> ModelInputProvider: - if spec.input_mode != "replay": + if spec.input_mode not in {"replay", "keyboard-driving"}: raise ValueError( - "OmniDreams replay providers currently support only " - f"input_mode='replay', got {spec.input_mode!r}." + "OmniDreams providers support input modes " + f"{self.supported_input_modes()}, got {spec.input_mode!r}." ) if spec.config is None: raise RuntimeError("DemoSpec.config was not initialized.") @@ -188,6 +211,67 @@ def create_model_input_provider( config=spec.config, ) + def _prepare_webrtc_scenario(self, spec: DemoSpec) -> PreparedScenario: + if spec.config is None: + raise RuntimeError("DemoSpec.config was not initialized.") + scenario = self._webrtc_ludus_scenario( + resolve_webrtc_scenario(spec.scenario), + spec=spec, + ) + return PreparedScenario( + initial_inputs=InferenceInput( + global_conditioning={"scenario": scenario}, + ), + source_schema=WEBRTC_USER_INPUT_SCHEMA, + canonicalizer=InputCanonicalizer(), + mapping=self._mapping, + metadata={ + "conditioning_mode": OMNIDREAMS_CONDITIONING_LUDUS, + "model_id": self.model_id, + "preset_id": self._preset_id(spec.config), + "num_views": 1, + }, + ) + + def _webrtc_ludus_scenario( + self, + scenario: OmnidreamsWebRTCScenario, + *, + spec: DemoSpec, + ) -> OmnidreamsLudusReplayScenario: + config = spec.config + if config is None: + raise RuntimeError("DemoSpec.config was not initialized.") + output = spec.output + fps = int(getattr(output, "fps", 30)) + video_height = int(getattr(output, "video_height", 704)) + video_width = int(getattr(output, "video_width", 1280)) + options = config.runtime_options + return OmnidreamsLudusReplayScenario( + keyboard_events=(), + scene_dir=scenario.scene_dir, + scene_uuid=scenario.scene_uuid or DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, + scene_variant=scenario.scene_variant, + camera_name=scenario.camera_name, + total_blocks=int( + options.get( + "total_blocks", + options.get("webrtc_total_blocks", 2_147_483_647), + ) + ), + pixel_height=video_height, + pixel_width=video_width, + fps=fps, + move_speed_per_s=float(options.get("move_speed_per_s", 6.0)), + rotate_speed_rad_per_s=float( + options.get("rotate_speed_rad_per_s", math.radians(35.0)) + ), + ludus_backend=cast( + LudusBackendName, + str(options.get("ludus_backend", "cuda")), + ), + ) + def _preset_id(self, config: InferenceConfig | None) -> str: return ( DEFAULT_OMNIDREAMS_PRESET @@ -214,6 +298,23 @@ def _default_replay_prompt(self, config: InferenceConfig | None) -> str: return "" if runner is None else str(getattr(runner, "prompt", "")) +def _bool_runtime_option( + options: Any, + name: str, + default: bool, +) -> bool: + value = options.get(name, default) + if isinstance(value, bool): + return value + if isinstance(value, str): + lowered = value.strip().lower() + if lowered in {"1", "true", "yes", "on"}: + return True + if lowered in {"0", "false", "no", "off"}: + return False + return bool(value) + + __all__ = [ "OmnidreamsDemoAdapter", "ReplayRuntimeFactory", diff --git a/integrations/omnidreams/omnidreams/demo/providers.py b/integrations/omnidreams/omnidreams/demo/providers.py index 72a174d3b..eff6de760 100644 --- a/integrations/omnidreams/omnidreams/demo/providers.py +++ b/integrations/omnidreams/omnidreams/demo/providers.py @@ -195,7 +195,7 @@ def __init__( supports_recorded_input=True, supports_reset=True, deterministic_given_inputs=True, - user_input_schema=keyboard_driving_user_input_schema(), + user_input_schema=scenario.source_schema, inference_input_schema=precomputed_hdmap_inference_input_schema(), ) diff --git a/integrations/omnidreams/omnidreams/demo/webrtc.py b/integrations/omnidreams/omnidreams/demo/webrtc.py index afcb84f82..d0510fa1d 100644 --- a/integrations/omnidreams/omnidreams/demo/webrtc.py +++ b/integrations/omnidreams/omnidreams/demo/webrtc.py @@ -18,6 +18,7 @@ from __future__ import annotations import math +import os from collections.abc import Callable, Mapping from dataclasses import dataclass, replace from pathlib import Path @@ -46,6 +47,7 @@ from flashdreams.runtime.demo import ( DemoSpec, PreparedScenario, + RuntimeHost, SessionInfo, UserInputWindow, WebRTCAppResources, @@ -68,6 +70,7 @@ from flashdreams.serving.webrtc.server import create_webrtc_app from flashdreams.serving.webrtc.services import WEBRTC_USER_INPUT_SCHEMA +from .adapter import OmnidreamsDemoAdapter, RuntimeFactory from .providers import LudusSceneConditioningProvider from .runtime import OmnidreamsRuntime, OmnidreamsRuntimeOptions, PipelineFactory from .spec import ( @@ -79,6 +82,7 @@ ) WebRTCRuntimeFactory = Callable[..., Any] +SharedRuntimeFactory = RuntimeFactory _WEBRTC_SESSION_TOTAL_BLOCKS = 2_147_483_647 _WEBRTC_STEP_REQUEST_KEY = "omnidreams_webrtc_step_request" @@ -750,11 +754,16 @@ def serve_omnidreams_webrtc_demo( *, spec: DemoSpec, world_rank: int = 0, - runtime_factory: WebRTCRuntimeFactory = OmnidreamsWebRTCModelRuntime, + runtime_factory: WebRTCRuntimeFactory | None = None, + shared_runtime_factory: SharedRuntimeFactory | None = None, create_app_fn: CreateWebRTCApp = create_webrtc_app, server_runner: RunWebRTCServer = run_webrtc_server, ) -> object: """Create OmniDreams' runtime and serve it through the shared WebRTC transport.""" + if runtime_factory is not None and shared_runtime_factory is not None: + raise ValueError( + "Specify either legacy runtime_factory or shared_runtime_factory, not both." + ) if spec.input_mode != "keyboard-driving": raise ValueError( "OmniDreams WebRTC requires input_mode='keyboard-driving', " @@ -771,6 +780,41 @@ def serve_omnidreams_webrtc_demo( f"got {config.model_id!r}." ) scenario = resolve_webrtc_scenario(spec.scenario) + runtime_config = _webrtc_runtime_config( + output=spec.output, + config=config, + scenario=scenario, + ) + if _should_use_legacy_webrtc_path( + scenario=scenario, + runtime_factory=runtime_factory, + ): + return _serve_legacy_omnidreams_webrtc_demo( + spec=spec, + output=spec.output, + runtime_config=runtime_config, + runtime_factory=runtime_factory or OmnidreamsWebRTCModelRuntime, + world_rank=world_rank, + create_app_fn=create_app_fn, + server_runner=server_runner, + ) + return _serve_shared_omnidreams_webrtc_demo( + spec=_shared_webrtc_spec(spec, runtime_config=runtime_config), + output=spec.output, + runtime_config=runtime_config, + shared_runtime_factory=shared_runtime_factory, + world_rank=world_rank, + create_app_fn=create_app_fn, + server_runner=server_runner, + ) + + +def _webrtc_runtime_config( + *, + output: WebRTCOutputSpec, + config: InferenceConfig, + scenario: Any, +) -> OmnidreamsWebRTCModelRuntimeConfig: preset_id = _preset_id(config) seed = _option(config, "seed", 42) runtime_config = OmnidreamsWebRTCModelRuntimeConfig( @@ -781,16 +825,50 @@ def serve_omnidreams_webrtc_demo( scene_variant=scenario.scene_variant, seed=None if seed is None else int(seed), device=config.device or str(_option(config, "device", "cuda:0")), - video_height=spec.output.video_height, - video_width=spec.output.video_width, - fps=spec.output.fps, + video_height=output.video_height, + video_width=output.video_width, + fps=output.fps, camera_name=scenario.camera_name, - warmup_chunks=spec.output.warmup_chunks, - warmup_timeout_s=spec.output.warmup_timeout_s, + warmup_chunks=output.warmup_chunks, + warmup_timeout_s=output.warmup_timeout_s, debug_serve_hdmaps=scenario.debug_serve_hdmaps, encoder_backend="default" if scenario.prefer_sw_encoder else "auto", ) - runtime_config = _apply_runtime_options(runtime_config, config.runtime_options) + return _apply_runtime_options(runtime_config, config.runtime_options) + + +def _should_use_legacy_webrtc_path( + *, + scenario: Any, + runtime_factory: WebRTCRuntimeFactory | None, +) -> bool: + if runtime_factory is not None: + return True + if bool(getattr(scenario, "debug_serve_hdmaps", False)): + logger.info( + "Using the legacy OmniDreams WebRTC path because debug HDMap " + "streaming is still implemented by the compatibility facade." + ) + return True + if _distributed_world_size() > 1: + logger.info( + "Using the legacy OmniDreams WebRTC path for multi-rank serving; " + "shared RuntimeHost distributed fan-out is not yet complete." + ) + return True + return False + + +def _serve_legacy_omnidreams_webrtc_demo( + *, + spec: DemoSpec, + output: WebRTCOutputSpec, + runtime_config: OmnidreamsWebRTCModelRuntimeConfig, + runtime_factory: WebRTCRuntimeFactory, + world_rank: int, + create_app_fn: CreateWebRTCApp, + server_runner: RunWebRTCServer, +) -> object: runtime = runtime_factory(config=runtime_config) manager = BaseWebRTCSessionManager( runtime=runtime, @@ -801,12 +879,12 @@ def serve_omnidreams_webrtc_demo( warmup_label="OmniDreams WebRTC", supported_control_keys=WSAD_SUPPORTED_KEYS, fatal_generation_errors=True, - client_liveness_timeout_s=spec.output.client_liveness_timeout_s, + client_liveness_timeout_s=output.client_liveness_timeout_s, ) from importlib.resources import files return serve_webrtc_demo( - output=spec.output, + output=output, model_id=spec.model_id, session_manager=manager, app_resources=WebRTCAppResources( @@ -819,6 +897,84 @@ def serve_omnidreams_webrtc_demo( ) +def _serve_shared_omnidreams_webrtc_demo( + *, + spec: DemoSpec, + output: WebRTCOutputSpec, + runtime_config: OmnidreamsWebRTCModelRuntimeConfig, + shared_runtime_factory: SharedRuntimeFactory | None, + world_rank: int, + create_app_fn: CreateWebRTCApp, + server_runner: RunWebRTCServer, +) -> object: + adapter = OmnidreamsDemoAdapter(runtime_factory=shared_runtime_factory) + prepared = adapter.prepare_scenario(spec) + config = spec.config + if config is None: + raise RuntimeError("DemoSpec.config was not initialized.") + runtime = adapter.create_runtime(config) + host = RuntimeHost(runtime) + manager = BaseWebRTCSessionManager( + runtime=runtime, + runtime_config=runtime_config, + fps=runtime_config.fps, + identity=runtime_config.pipeline_config_name, + busy_message="An OmniDreams session is already active.", + warmup_label="OmniDreams WebRTC", + supported_control_keys=WSAD_SUPPORTED_KEYS, + fatal_generation_errors=True, + client_liveness_timeout_s=output.client_liveness_timeout_s, + shared_host=host, + shared_adapter=adapter, + shared_spec=spec, + shared_scenario=prepared, + ) + from importlib.resources import files + + return serve_webrtc_demo( + output=output, + model_id=spec.model_id, + session_manager=manager, + app_resources=WebRTCAppResources( + model_web_resource=files("omnidreams.demo").joinpath("web"), + preload_name="OmniDreams", + ), + world_rank=world_rank, + create_app_fn=create_app_fn, + server_runner=server_runner, + ) + + +def _shared_webrtc_spec( + spec: DemoSpec, + *, + runtime_config: OmnidreamsWebRTCModelRuntimeConfig, +) -> DemoSpec: + config = spec.config + if config is None: + raise RuntimeError("DemoSpec.config was not initialized.") + runtime_options = dict(config.runtime_options) + runtime_options.update( + { + "pipeline_config": runtime_config.pipeline_config, + "seed": runtime_config.seed, + "move_speed_per_s": runtime_config.move_speed_per_s, + "rotate_speed_rad_per_s": runtime_config.rotate_speed_rad_per_s, + "release_oneshot_encoders_after_cache_init": False, + } + ) + return replace( + spec, + config=replace( + config, + preset_id=runtime_config.pipeline_config_name, + device=runtime_config.device, + seed=runtime_config.seed, + runtime_options=runtime_options, + ), + ) + + def _preset_id(config: InferenceConfig | None) -> str: return ( DEFAULT_OMNIDREAMS_PRESET @@ -846,6 +1002,13 @@ def _option(config: InferenceConfig, name: str, default: Any) -> Any: return config.runtime_options.get(name, default) +def _distributed_world_size() -> int: + try: + return int(os.environ.get("WORLD_SIZE", "1")) + except ValueError: + return 1 + + def _apply_runtime_options( runtime_config: OmnidreamsWebRTCModelRuntimeConfig, options: Any, @@ -869,6 +1032,7 @@ def _apply_runtime_options( "OmnidreamsWebRTCModelRuntime", "OmnidreamsWebRTCModelRuntimeConfig", "OmnidreamsWebRTCModelRuntimeError", + "SharedRuntimeFactory", "WebRTCRuntimeFactory", "serve_omnidreams_webrtc_demo", ] diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index 4a2e89d3a..59c6f7320 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -62,6 +62,8 @@ Mp4OutputSpec, NullOutputSpec, OutputDecision, + PreparedScenario, + RuntimeHost, SessionInfo, UserInputWindow, WebRTCOutputSpec, @@ -155,12 +157,12 @@ def test_omnidreams_replay_cli_builds_ludus_conditioning_spec( assert spec.config.runtime_options["seed"] == 123 -def test_omnidreams_demo_adapter_declares_replay_modes_only() -> None: +def test_omnidreams_demo_adapter_declares_shared_modes() -> None: adapter = OmnidreamsDemoAdapter() assert adapter.model_id == OMNIDREAMS_MODEL_ID - assert adapter.supported_input_modes() == ("replay",) - assert adapter.supported_output_modes() == ("mp4", "null") + assert adapter.supported_input_modes() == ("replay", "keyboard-driving") + assert adapter.supported_output_modes() == ("mp4", "null", "webrtc") assert adapter.supported_conditioning_modes() == ( OMNIDREAMS_CONDITIONING_PRECOMPUTED, OMNIDREAMS_CONDITIONING_LUDUS, @@ -841,6 +843,7 @@ def test_omnidreams_webrtc_cli_builds_keyboard_driving_spec(tmp_path: Path) -> N def test_omnidreams_webrtc_demo_uses_shared_manager_with_model_config() -> None: pipeline_config = object() + runtime = _FactoryRuntime() spec = DemoSpec( model_id=OMNIDREAMS_MODEL_ID, preset_id=DEFAULT_OMNIDREAMS_PRESET, @@ -849,7 +852,6 @@ def test_omnidreams_webrtc_demo_uses_shared_manager_with_model_config() -> None: scene_uuid="scene-1", scene_variant="rain", camera_name="camera_front_wide_120fov", - debug_serve_hdmaps=True, prefer_sw_encoder=True, ), output=WebRTCOutputSpec( @@ -869,6 +871,83 @@ def test_omnidreams_webrtc_demo_uses_shared_manager_with_model_config() -> None: ), ) + calls: list[dict[str, Any]] = [] + runtime_calls: list[dict[str, Any]] = [] + + def shared_runtime_factory(**kwargs: Any) -> Any: + runtime_calls.append(kwargs) + return runtime + + serve_omnidreams_webrtc_demo( + spec=spec, + world_rank=1, + shared_runtime_factory=shared_runtime_factory, + server_runner=lambda **kwargs: calls.append(kwargs), + ) + + manager = calls[0]["session_manager"] + assert type(manager) is BaseWebRTCSessionManager + assert manager._runtime is runtime + assert isinstance(manager._shared_host, RuntimeHost) + assert isinstance(manager._shared_adapter, OmnidreamsDemoAdapter) + assert isinstance(manager._shared_scenario, PreparedScenario) + assert manager.runtime_config.pipeline_config is pipeline_config + assert manager.runtime_config.pipeline_config_name == DEFAULT_OMNIDREAMS_PRESET + assert manager.runtime_config.scene_uuid == "scene-1" + assert manager.runtime_config.scene_variant == "rain" + assert manager.runtime_config.seed == 123 + assert manager.runtime_config.device == "cuda:7" + assert manager.runtime_config.video_width == 64 + assert manager.runtime_config.video_height == 32 + assert manager.runtime_config.fps == 24 + assert manager.runtime_config.debug_serve_hdmaps is False + assert manager.runtime_config.encoder_backend == "default" + assert manager.identity == DEFAULT_OMNIDREAMS_PRESET + assert len(runtime_calls) == 1 + runtime_config = runtime_calls[0]["config"] + assert runtime_config.seed == 123 + assert runtime_config.runtime_options["pipeline_config"] is pipeline_config + assert ( + runtime_config.runtime_options["release_oneshot_encoders_after_cache_init"] + is False + ) + options = runtime_calls[0]["options"] + assert isinstance(options, OmnidreamsRuntimeOptions) + assert options.release_oneshot_encoders_after_cache_init is False + scenario = manager._shared_scenario.initial_inputs.global_conditioning["scenario"] + assert isinstance(scenario, OmnidreamsLudusReplayScenario) + assert scenario.scene_uuid == "scene-1" + assert scenario.scene_variant == "rain" + assert scenario.pixel_width == 64 + assert scenario.pixel_height == 32 + assert scenario.fps == 24 + assert calls[0]["host"] == "0.0.0.0" + assert calls[0]["port"] == 8082 + + +def test_omnidreams_webrtc_demo_keeps_legacy_runtime_factory_path() -> None: + spec = DemoSpec( + model_id=OMNIDREAMS_MODEL_ID, + preset_id=DEFAULT_OMNIDREAMS_PRESET, + input_mode="keyboard-driving", + scenario=OmnidreamsWebRTCScenario(debug_serve_hdmaps=True), + output=WebRTCOutputSpec( + host="0.0.0.0", + port=8082, + fps=24, + video_width=64, + video_height=32, + warmup_chunks=0, + warmup_timeout_s=1.0, + ), + config=InferenceConfig( + model_id=OMNIDREAMS_MODEL_ID, + preset_id=DEFAULT_OMNIDREAMS_PRESET, + device="cuda:7", + runtime_options={"pipeline_config": object(), "seed": 123}, + ), + ) + calls: list[dict[str, Any]] = [] serve_omnidreams_webrtc_demo( spec=spec, @@ -880,22 +959,8 @@ def test_omnidreams_webrtc_demo_uses_shared_manager_with_model_config() -> None: manager = calls[0]["session_manager"] runtime = manager._runtime assert isinstance(runtime, _FakeWebRTCRuntime) - assert type(manager) is BaseWebRTCSessionManager assert manager.runtime_config is runtime.config - assert runtime.config.pipeline_config is pipeline_config - assert runtime.config.pipeline_config_name == DEFAULT_OMNIDREAMS_PRESET - assert runtime.config.scene_uuid == "scene-1" - assert runtime.config.scene_variant == "rain" - assert runtime.config.seed == 123 - assert runtime.config.device == "cuda:7" - assert runtime.config.video_width == 64 - assert runtime.config.video_height == 32 - assert runtime.config.fps == 24 assert runtime.config.debug_serve_hdmaps is True - assert runtime.config.encoder_backend == "default" - assert manager.identity == DEFAULT_OMNIDREAMS_PRESET - assert calls[0]["host"] == "0.0.0.0" - assert calls[0]["port"] == 8082 def test_omnidreams_webrtc_demo_installs_model_assets_without_routes( @@ -938,7 +1003,7 @@ def fake_create_packaged_webrtc_app(**kwargs: Any) -> web.Application: app = serve_omnidreams_webrtc_demo( spec=spec, - runtime_factory=_FakeWebRTCRuntime, + shared_runtime_factory=lambda **kwargs: _FactoryRuntime(), server_runner=lambda **kwargs: None, ) @@ -1014,7 +1079,7 @@ def fake_server_runner(**kwargs: Any) -> None: app = serve_omnidreams_webrtc_demo( spec=spec, world_rank=0, - runtime_factory=_FakeWebRTCRuntime, + shared_runtime_factory=lambda **kwargs: _FactoryRuntime(), server_runner=fake_server_runner, ) @@ -1155,33 +1220,69 @@ async def test_omnidreams_webrtc_manager_drives_shared_session( scene_path = tmp_path / "scene.usdz" scene_path.write_bytes(b"fake") pipeline = _VariableFrameOmnidreamsPipeline((2, 3)) - config = OmnidreamsWebRTCModelRuntimeConfig( - pipeline_config_name="fake", - pipeline_config=object(), + adapter = OmnidreamsDemoAdapter( pipeline_factory=lambda pipeline_config, device: pipeline, + ) + spec = DemoSpec( + model_id=OMNIDREAMS_MODEL_ID, + preset_id=DEFAULT_OMNIDREAMS_PRESET, + input_mode="keyboard-driving", + scenario=OmnidreamsWebRTCScenario( + scene_dir=scene_path, + scene_uuid="scene-1", + camera_name="camera_front_wide_120fov", + ), + output=WebRTCOutputSpec( + fps=30, + video_width=2, + video_height=2, + warmup_chunks=0, + warmup_timeout_s=1.0, + ), + config=InferenceConfig( + model_id=OMNIDREAMS_MODEL_ID, + preset_id=DEFAULT_OMNIDREAMS_PRESET, + device="cpu", + seed=123, + runtime_options={ + "pipeline_config": object(), + "seed": 123, + "release_oneshot_encoders_after_cache_init": False, + }, + ), + ) + prepared = adapter.prepare_scenario(spec) + assert spec.config is not None + runtime = adapter.create_runtime(spec.config) + runtime_config = OmnidreamsWebRTCModelRuntimeConfig( + pipeline_config_name=DEFAULT_OMNIDREAMS_PRESET, + pipeline_config=object(), scene_dir=scene_path, + scene_uuid="scene-1", device="cpu", fps=30, video_height=2, video_width=2, warmup_chunks=0, ) - runtime = OmnidreamsWebRTCModelRuntime(config=config) + host = RuntimeHost(runtime) manager = BaseWebRTCSessionManager( runtime=runtime, - runtime_config=config, - fps=config.fps, - identity=config.pipeline_config_name, + runtime_config=runtime_config, + fps=runtime_config.fps, + identity=runtime_config.pipeline_config_name, supported_control_keys=frozenset({"w", "a", "s", "d"}), + shared_host=host, + shared_adapter=adapter, + shared_spec=spec, + shared_scenario=prepared, ) - await runtime.initialize() manager._runtime_ready = True - await runtime.reset_for_new_session() loop = asyncio.get_running_loop() context = manager._shared_run_context(loop) reservation = context.admission.try_reserve() assert reservation is not None - resampler = _FakeWebRTCResampler(start_v=loop.time(), fps=config.fps) + resampler = _FakeWebRTCResampler(start_v=loop.time(), fps=runtime_config.fps) input_source = WebRTCInputSource(resampler=resampler) input_source.handle_browser_payload( {"type": "action", "action": {"event": "step"}}, @@ -1191,7 +1292,7 @@ async def test_omnidreams_webrtc_manager_drives_shared_session( channel = _FakeWebRTCChannel() managed_session = ManagedWebRTCSession( runtime=runtime, - video_track=_FakeWebRTCVideoTrack(fps=config.fps), # ty:ignore[invalid-argument-type] + video_track=_FakeWebRTCVideoTrack(fps=runtime_config.fps), # ty:ignore[invalid-argument-type] video_encoder=_FakeWebRTCVideoEncoder(), # ty:ignore[invalid-argument-type] peer_connection=_FakeWebRTCPeerConnection(), resampler=resampler, # ty:ignore[invalid-argument-type] @@ -1213,16 +1314,19 @@ async def test_omnidreams_webrtc_manager_drives_shared_session( chunk = await _wait_for_chunk_done(channel) assert chunk["type"] == "chunk_done" - assert chunk["model"] == "fake" + assert chunk["model"] == DEFAULT_OMNIDREAMS_PRESET assert chunk["num_frames"] == 2 assert [tuple(hdmap.shape) for hdmap in pipeline.generated_hdmaps][:1] == [ (1, 1, 2, 3, 2, 2) ] assert rasterizers[0].calls[0]["timestamps_us"] == (1_000, 34_333) + assert isinstance( + managed_session.input_source, + WebRTCInputSource, + ) finally: transport.close("test complete") - await manager.close_active_session() - await runtime.close() + await manager.shutdown() class _RecordingOutputTarget: From 24c6e07dd7b85030cd203c6b52a22907b4f66ac2 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 04:37:03 +0000 Subject: [PATCH 48/51] ci: add OmniDreams demo runtime GPU workflow Add a standalone GPU workflow for the migrated OmniDreams demo runtime paths. The workflow exercises null output, precomputed HDMap MP4, and Ludus recorded-trace MP4 runs, validates logs and MP4 metadata, and uploads generated artifacts for inspection. --- .github/workflows/omnidreams-demo-runtime.yml | 328 ++++++++++++++++++ 1 file changed, 328 insertions(+) create mode 100644 .github/workflows/omnidreams-demo-runtime.yml diff --git a/.github/workflows/omnidreams-demo-runtime.yml b/.github/workflows/omnidreams-demo-runtime.yml new file mode 100644 index 000000000..a145bfccf --- /dev/null +++ b/.github/workflows/omnidreams-demo-runtime.yml @@ -0,0 +1,328 @@ +name: OmniDreams Demo Runtime + +on: + push: + branches: + - main + - "pull-request/[0-9]+" + paths: + - ".github/workflows/omnidreams-demo-runtime.yml" + - "pyproject.toml" + - "uv.lock" + - "flashdreams/pyproject.toml" + - "flashdreams/flashdreams/core/**" + - "flashdreams/flashdreams/infra/**" + - "flashdreams/flashdreams/runtime/**" + - "flashdreams/flashdreams/serving/**" + - "flashdreams/flashdreams/recipes/taehv/**" + - "flashdreams/flashdreams/recipes/wan/**" + - "integrations/omnidreams/**" + - "integrations/lingbot/**" + workflow_dispatch: + +permissions: + contents: read + +jobs: + demo-runtime: + name: null, precomputed MP4, and Ludus MP4 + runs-on: linux-amd64-gpu-rtxpro6000-latest-2 + timeout-minutes: 180 + defaults: + run: + shell: bash + container: + image: nvidia/cuda:13.2.1-cudnn-devel-ubuntu24.04 + options: --gpus all + env: + UV_PROJECT_ENVIRONMENT: /tmp/flashdreams-venv + UV_LINK_MODE: copy + UV_PYTHON: "3.10" + MAX_JOBS: 8 + ARTIFACT_DIR: artifacts/omnidreams_demo_runtime + NULL_BLOCKS: "10" + PRECOMPUTED_BLOCKS: "225" + LUDUS_BLOCKS: "226" + FPS: "30" + EXAMPLE_DATA_UUID: 239560dc-33d1-11ef-9720-00044bcbccac + LUDUS_SCENE_UUID: 0d404ff7-2b66-498c-b047-1ed8cded60d4 + LUDUS_TRACE: integrations/omnidreams/omnidreams/demo/traces/ludus_forward_sweep_60s.json + EXPECTED_WIDTH: "1280" + EXPECTED_HEIGHT: "704" + MIN_DURATION_SECONDS: "58" + MAX_DURATION_SECONDS: "62" + steps: + - name: Detect GPU architecture + id: gpu-arch + run: | + nvidia-smi + compute_cap=$(nvidia-smi --query-gpu=compute_cap --format=csv,noheader 2>/dev/null | head -1 | tr -d '[:space:]') + arch=$(echo "${compute_cap}" | tr -d '.') + echo "arch=${arch}" >> "$GITHUB_OUTPUT" + echo "Detected GPU compute capability: ${compute_cap} -> sm_${arch}" + + - name: Checkout + uses: actions/checkout@v4 + + - name: Install system dependencies + run: | + apt-get update -qq + DEBIAN_FRONTEND=noninteractive apt-get install -y -qq --no-install-recommends \ + python3 python3-dev python3-venv \ + ffmpeg \ + gcc g++ ninja-build \ + libnccl-dev \ + curl git ca-certificates jq unzip + rm -rf /var/lib/apt/lists/* + + - name: Setup proxy cache + uses: nv-gha-runners/setup-proxy-cache@main + + - name: Setup uv + uses: astral-sh/setup-uv@v6 + with: + enable-cache: true + cache-suffix: "omnidreams-demo-runtime-sm${{ steps.gpu-arch.outputs.arch }}" + prune-cache: false + + - name: Install dependencies + env: + NVTE_CUDA_ARCHS: ${{ steps.gpu-arch.outputs.arch }} + run: | + uv venv --clear + uv sync --locked --extra dev + + - name: Verify GPU availability + run: nvidia-smi + + - name: Run OmniDreams demo modes + env: + HF_TOKEN: ${{ secrets.HF_TOKEN }} + run: | + set -uo pipefail + + log_dir="${ARTIFACT_DIR}/logs" + output_dir="${ARTIFACT_DIR}/outputs" + summary="${ARTIFACT_DIR}/summary.md" + status_file="${ARTIFACT_DIR}/command-status.env" + mkdir -p "${log_dir}" "${output_dir}" + : > "${status_file}" + + odemo() { + uv run --no-sync --package flashdreams-omnidreams omnidreams-demo "$@" + } + + run_demo() { + local name="$1" + shift + local log="${log_dir}/${name}.log" + + { + printf '$' + printf ' %q' "$@" + printf '\n\n' + "$@" + } 2>&1 | tee "${log}" + + local rc="${PIPESTATUS[0]}" + echo "${name}=${rc}" >> "${status_file}" + echo "${name} exit code: ${rc}" | tee -a "${summary}" + return 0 + } + + { + echo "# OmniDreams Demo Runtime CI" + echo + echo "| Mode | Expected blocks | Output |" + echo "| --- | ---: | --- |" + echo "| null | ${NULL_BLOCKS} | none |" + echo "| precomputed MP4 | ${PRECOMPUTED_BLOCKS} | omnidreams-demo-precomputed-1min.mp4 |" + echo "| Ludus MP4 | ${LUDUS_BLOCKS} | omnidreams-demo-ludus-1min.mp4 |" + echo + echo "## Command Status" + } > "${summary}" + + run_demo null \ + odemo replay \ + --output-mode null \ + --device cuda:0 \ + --total-blocks "${NULL_BLOCKS}" + + run_demo precomputed-mp4 \ + odemo replay \ + --device cuda:0 \ + --example-data \ + --example-data-uuid "${EXAMPLE_DATA_UUID}" \ + --total-blocks "${PRECOMPUTED_BLOCKS}" \ + --fps "${FPS}" \ + --output "${output_dir}/omnidreams-demo-precomputed-1min.mp4" + + run_demo ludus-mp4 \ + odemo replay \ + --conditioning-mode ludus-scene-driving \ + --keyboard-trace "${LUDUS_TRACE}" \ + --device cuda:0 \ + --scene-uuid "${LUDUS_SCENE_UUID}" \ + --seed 42 \ + --total-blocks "${LUDUS_BLOCKS}" \ + --output "${output_dir}/omnidreams-demo-ludus-1min.mp4" + + - name: Validate OmniDreams demo artifacts + run: | + set -euo pipefail + + log_dir="${ARTIFACT_DIR}/logs" + output_dir="${ARTIFACT_DIR}/outputs" + probe_dir="${ARTIFACT_DIR}/ffprobe" + summary="${ARTIFACT_DIR}/summary.md" + status_file="${ARTIFACT_DIR}/command-status.env" + mkdir -p "${probe_dir}" + + status_of() { + awk -F= -v name="$1" '$1 == name { print $2 }' "${status_file}" + } + + assert_exit_zero() { + local name="$1" + local rc + rc="$(status_of "${name}")" + if [ "${rc}" != "0" ]; then + echo "${name} command failed with exit code ${rc}" >&2 + exit 1 + fi + } + + assert_clean_log() { + local log="$1" + if grep -En "ERROR|Traceback|Exception|status=failed|Run failed|failed run" "${log}"; then + echo "failure marker found in ${log}" >&2 + exit 1 + fi + } + + assert_log_contains() { + local log="$1" + local pattern="$2" + local label="$3" + if ! grep -Eq "${pattern}" "${log}"; then + echo "expected ${label} in ${log}" >&2 + exit 1 + fi + } + + validate_mp4() { + local mode="$1" + local mp4="$2" + local metadata="${probe_dir}/${mode}.json" + + if [ ! -s "${mp4}" ]; then + echo "expected non-empty MP4 at ${mp4}" >&2 + exit 1 + fi + + ffprobe \ + -v error \ + -select_streams v:0 \ + -show_entries stream=width,height,r_frame_rate,avg_frame_rate,nb_frames,duration:format=duration \ + -of json \ + "${mp4}" > "${metadata}" + + local stream_count width height duration + stream_count="$(jq '.streams | length' "${metadata}")" + width="$(jq -r '.streams[0].width // ""' "${metadata}")" + height="$(jq -r '.streams[0].height // ""' "${metadata}")" + duration="$(jq -r '.streams[0].duration // .format.duration // "0"' "${metadata}")" + + if [ "${stream_count}" -lt 1 ]; then + echo "ffprobe found no video stream in ${mp4}" >&2 + exit 1 + fi + + if [ "${width}" != "${EXPECTED_WIDTH}" ] || [ "${height}" != "${EXPECTED_HEIGHT}" ]; then + echo "unexpected ${mode} resolution ${width}x${height}; expected ${EXPECTED_WIDTH}x${EXPECTED_HEIGHT}" >&2 + exit 1 + fi + + awk \ + -v duration="${duration}" \ + -v min_duration="${MIN_DURATION_SECONDS}" \ + -v max_duration="${MAX_DURATION_SECONDS}" \ + 'BEGIN { + if ((duration + 0) < min_duration || (duration + 0) > max_duration) { + exit 1 + } + }' || { + echo "unexpected ${mode} duration ${duration}s; expected ${MIN_DURATION_SECONDS}-${MAX_DURATION_SECONDS}s" >&2 + exit 1 + } + } + + null_log="${log_dir}/null.log" + precomputed_log="${log_dir}/precomputed-mp4.log" + ludus_log="${log_dir}/ludus-mp4.log" + + assert_exit_zero null + assert_exit_zero precomputed-mp4 + assert_exit_zero ludus-mp4 + + assert_clean_log "${null_log}" + assert_clean_log "${precomputed_log}" + assert_clean_log "${ludus_log}" + + assert_log_contains "${null_log}" "AR 9 encode" "null final AR block" + assert_log_contains "${null_log}" "OmniDreams demo replay step 9 frames=" "null final replay step" + assert_log_contains "${null_log}" "Loaded OmniDreams demo HDMaps shape=.*views=1" "null precomputed HDMaps" + + assert_log_contains "${precomputed_log}" "AR 224 encode" "precomputed final AR block" + assert_log_contains "${precomputed_log}" "OmniDreams demo replay step 224 frames=" "precomputed final replay step" + assert_log_contains "${precomputed_log}" "Loaded OmniDreams demo HDMaps shape=.*views=1" "precomputed HDMaps" + + assert_log_contains "${ludus_log}" "AR 225 encode" "Ludus final AR block" + assert_log_contains "${ludus_log}" "OmniDreams demo replay step 225 frames=" "Ludus final replay step" + assert_log_contains "${ludus_log}" "ludus_backend=cuda" "Ludus CUDA backend" + assert_log_contains "${ludus_log}" "trace_events=[1-9][0-9]*" "nonzero Ludus trace events" + + validate_mp4 precomputed-mp4 "${output_dir}/omnidreams-demo-precomputed-1min.mp4" + validate_mp4 ludus-mp4 "${output_dir}/omnidreams-demo-ludus-1min.mp4" + + { + echo + echo "## Validation" + echo + echo "- All commands exited zero." + echo "- Logs contained expected final AR blocks and provider markers." + echo "- MP4 outputs were non-empty and passed ffprobe stream checks." + } >> "${summary}" + + - name: Trim uv cache for upload + if: always() + run: | + cache_dir="${UV_CACHE_DIR:-/github/home/.cache/uv}" + echo "=== Cache size before trim ===" + du -sh "${cache_dir}" || true + du -sh "${cache_dir}"/*/ 2>/dev/null || true + + rm -rf "${cache_dir}/wheels-v6" + rm -rf "${cache_dir}/archive-v0" + + find "${cache_dir}/git-v0/checkouts" \ + \( -name "build" -o -name "*.egg-info" -o -name "__pycache__" \) \ + -type d -exec rm -rf {} + 2>/dev/null || true + + rm -rf "${cache_dir}/sdists-v9/editable" + + echo "" + echo "=== Cache size after trim ===" + du -sh "${cache_dir}" || true + du -sh "${cache_dir}"/*/ 2>/dev/null || true + echo "" + echo "=== Cached built wheels (sdists-v9) ===" + find "${cache_dir}/sdists-v9" -name "*.whl" -exec ls -lh {} \; 2>/dev/null || true + + - name: Upload demo runtime artifacts + if: always() + uses: actions/upload-artifact@v4 + with: + name: omnidreams-demo-runtime + path: ${{ env.ARTIFACT_DIR }} + if-no-files-found: ignore From 1a9193595bd08002fd7be3cd5c1980332f67d264 Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 04:54:22 +0000 Subject: [PATCH 49/51] ci: shorten OmniDreams demo runtime artifacts Reduce the OmniDreams demo runtime GPU workflow MP4 runs from one minute to roughly 20 seconds. Keep the same null, precomputed HDMap, and Ludus validation coverage while reducing GPU CI runtime and artifact size. --- .github/workflows/omnidreams-demo-runtime.yml | 28 +++++++++---------- 1 file changed, 14 insertions(+), 14 deletions(-) diff --git a/.github/workflows/omnidreams-demo-runtime.yml b/.github/workflows/omnidreams-demo-runtime.yml index a145bfccf..1f494fa83 100644 --- a/.github/workflows/omnidreams-demo-runtime.yml +++ b/.github/workflows/omnidreams-demo-runtime.yml @@ -41,16 +41,16 @@ jobs: MAX_JOBS: 8 ARTIFACT_DIR: artifacts/omnidreams_demo_runtime NULL_BLOCKS: "10" - PRECOMPUTED_BLOCKS: "225" - LUDUS_BLOCKS: "226" + PRECOMPUTED_BLOCKS: "75" + LUDUS_BLOCKS: "76" FPS: "30" EXAMPLE_DATA_UUID: 239560dc-33d1-11ef-9720-00044bcbccac LUDUS_SCENE_UUID: 0d404ff7-2b66-498c-b047-1ed8cded60d4 LUDUS_TRACE: integrations/omnidreams/omnidreams/demo/traces/ludus_forward_sweep_60s.json EXPECTED_WIDTH: "1280" EXPECTED_HEIGHT: "704" - MIN_DURATION_SECONDS: "58" - MAX_DURATION_SECONDS: "62" + MIN_DURATION_SECONDS: "18" + MAX_DURATION_SECONDS: "22" steps: - name: Detect GPU architecture id: gpu-arch @@ -136,8 +136,8 @@ jobs: echo "| Mode | Expected blocks | Output |" echo "| --- | ---: | --- |" echo "| null | ${NULL_BLOCKS} | none |" - echo "| precomputed MP4 | ${PRECOMPUTED_BLOCKS} | omnidreams-demo-precomputed-1min.mp4 |" - echo "| Ludus MP4 | ${LUDUS_BLOCKS} | omnidreams-demo-ludus-1min.mp4 |" + echo "| precomputed MP4 | ${PRECOMPUTED_BLOCKS} | omnidreams-demo-precomputed-20s.mp4 |" + echo "| Ludus MP4 | ${LUDUS_BLOCKS} | omnidreams-demo-ludus-20s.mp4 |" echo echo "## Command Status" } > "${summary}" @@ -155,7 +155,7 @@ jobs: --example-data-uuid "${EXAMPLE_DATA_UUID}" \ --total-blocks "${PRECOMPUTED_BLOCKS}" \ --fps "${FPS}" \ - --output "${output_dir}/omnidreams-demo-precomputed-1min.mp4" + --output "${output_dir}/omnidreams-demo-precomputed-20s.mp4" run_demo ludus-mp4 \ odemo replay \ @@ -165,7 +165,7 @@ jobs: --scene-uuid "${LUDUS_SCENE_UUID}" \ --seed 42 \ --total-blocks "${LUDUS_BLOCKS}" \ - --output "${output_dir}/omnidreams-demo-ludus-1min.mp4" + --output "${output_dir}/omnidreams-demo-ludus-20s.mp4" - name: Validate OmniDreams demo artifacts run: | @@ -273,17 +273,17 @@ jobs: assert_log_contains "${null_log}" "OmniDreams demo replay step 9 frames=" "null final replay step" assert_log_contains "${null_log}" "Loaded OmniDreams demo HDMaps shape=.*views=1" "null precomputed HDMaps" - assert_log_contains "${precomputed_log}" "AR 224 encode" "precomputed final AR block" - assert_log_contains "${precomputed_log}" "OmniDreams demo replay step 224 frames=" "precomputed final replay step" + assert_log_contains "${precomputed_log}" "AR 74 encode" "precomputed final AR block" + assert_log_contains "${precomputed_log}" "OmniDreams demo replay step 74 frames=" "precomputed final replay step" assert_log_contains "${precomputed_log}" "Loaded OmniDreams demo HDMaps shape=.*views=1" "precomputed HDMaps" - assert_log_contains "${ludus_log}" "AR 225 encode" "Ludus final AR block" - assert_log_contains "${ludus_log}" "OmniDreams demo replay step 225 frames=" "Ludus final replay step" + assert_log_contains "${ludus_log}" "AR 75 encode" "Ludus final AR block" + assert_log_contains "${ludus_log}" "OmniDreams demo replay step 75 frames=" "Ludus final replay step" assert_log_contains "${ludus_log}" "ludus_backend=cuda" "Ludus CUDA backend" assert_log_contains "${ludus_log}" "trace_events=[1-9][0-9]*" "nonzero Ludus trace events" - validate_mp4 precomputed-mp4 "${output_dir}/omnidreams-demo-precomputed-1min.mp4" - validate_mp4 ludus-mp4 "${output_dir}/omnidreams-demo-ludus-1min.mp4" + validate_mp4 precomputed-mp4 "${output_dir}/omnidreams-demo-precomputed-20s.mp4" + validate_mp4 ludus-mp4 "${output_dir}/omnidreams-demo-ludus-20s.mp4" { echo From 441e1bbd4466b1158b51611710255b646c75cbca Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 05:53:25 +0000 Subject: [PATCH 50/51] omnidreams: clean up shared demo runtime layout Move the canonical OmniDreams runtime/session into demo/runtime.py while preserving replay compatibility aliases. Split the default shared WebRTC entrypoint from the legacy compatibility facade, share WebRTC config through demo/webrtc_config.py, and cover shared-vs-legacy routing plus lazy legacy imports in CPU tests. --- .../omnidreams/omnidreams/demo/replay.py | 352 +------- .../omnidreams/omnidreams/demo/runtime.py | 345 +++++++- .../omnidreams/omnidreams/demo/webrtc.py | 757 +----------------- .../omnidreams/demo/webrtc_config.py | 86 ++ .../omnidreams/demo/webrtc_legacy.py | 718 +++++++++++++++++ .../omnidreams/tests/test_demo_api.py | 43 +- 6 files changed, 1206 insertions(+), 1095 deletions(-) create mode 100644 integrations/omnidreams/omnidreams/demo/webrtc_config.py create mode 100644 integrations/omnidreams/omnidreams/demo/webrtc_legacy.py diff --git a/integrations/omnidreams/omnidreams/demo/replay.py b/integrations/omnidreams/omnidreams/demo/replay.py index f7eb60cdd..b6c6d91e2 100644 --- a/integrations/omnidreams/omnidreams/demo/replay.py +++ b/integrations/omnidreams/omnidreams/demo/replay.py @@ -1,361 +1,29 @@ # SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. # SPDX-License-Identifier: Apache-2.0 -"""OmniDreams replay runtime for the shared demo runner.""" +"""Compatibility aliases for the OmniDreams runtime module.""" from __future__ import annotations -import os -from collections.abc import Callable, Sequence -from dataclasses import dataclass -from typing import Any - -import torch -import torch.distributed as dist -from loguru import logger -from omnidreams.model_session import OmnidreamsModelSessionCore - -from flashdreams.core.distributed import init as init_distributed -from flashdreams.infra.postprocess import VideoTensorLayout -from flashdreams.infra.runner_io import ( - DEFAULT_RUNNER_INSTALL_HINT, - load_first_frame_tensor, +from .runtime import ( + OmnidreamsRuntime, + OmnidreamsRuntimeOptions, + OmnidreamsSession, + OmnidreamsSessionScenario, + PipelineFactory, ) -from flashdreams.infra.video_output import VideoOutputStream -from flashdreams.runtime.config import InferenceConfig -from flashdreams.runtime.inputs import InferenceInput -from flashdreams.runtime.interfaces import InferenceSession -from flashdreams.runtime.types import StepRequest, StepRequirements, StepResult - -from .spec import OmnidreamsLudusReplayScenario, OmnidreamsReplayScenario - -OmnidreamsSessionScenario = OmnidreamsReplayScenario | OmnidreamsLudusReplayScenario - -PipelineFactory = Callable[[Any, str], Any] - - -@dataclass(frozen=True, kw_only=True, slots=True) -class OmnidreamsRuntimeOptions: - """Construction knobs for the OmniDreams runtime.""" - - pipeline_config: Any - pipeline_factory: PipelineFactory | None = None - output_layout: VideoTensorLayout = "bvtchw" - release_oneshot_encoders_after_cache_init: bool = True - - -class OmnidreamsRuntime: - """Heavyweight OmniDreams runtime consumed by shared demo run modes.""" - - def __init__( - self, - *, - config: InferenceConfig, - options: OmnidreamsRuntimeOptions, - ) -> None: - self.config = config - self.options = options - if _is_torchrun_env() and not dist.is_initialized(): - init_distributed() - - if dist.is_initialized(): - self.local_rank = int(os.environ.get("LOCAL_RANK", "0")) - self.world_size = dist.get_world_size() - self.global_rank = dist.get_rank() - device = f"cuda:{self.local_rank}" - else: - self.local_rank = 0 - self.world_size = 1 - self.global_rank = 0 - device = config.device or "cuda" - - self.is_rank_zero = self.global_rank == 0 - factory = options.pipeline_factory or _default_pipeline_factory - self.pipeline = factory(options.pipeline_config, device) - - def start_session(self, inputs: InferenceInput) -> InferenceSession: - scenario = _scenario_from_inputs(inputs) - return OmnidreamsSession( - pipeline=self.pipeline, - scenario=scenario, - initial_inputs=inputs, - device=torch.device(f"cuda:{self.local_rank}") - if dist.is_initialized() - else torch.device(self.config.device or "cuda"), - is_rank_zero=self.is_rank_zero, - output_layout=self.options.output_layout, - rollout_seed=self.config.seed, - release_oneshot_encoders_after_cache_init=( - self.options.release_oneshot_encoders_after_cache_init - ), - ) - - def close(self) -> None: - pipeline = getattr(self, "pipeline", None) - if pipeline is not None: - close = getattr(pipeline, "close", None) - if callable(close): - close() - del self.pipeline - device = torch.device(self.config.device or "cuda") - if device.type == "cuda" and torch.cuda.is_available(): - torch.cuda.empty_cache() - - -class OmnidreamsSession: - """One OmniDreams rollout over a prepared scenario.""" - - def __init__( - self, - *, - pipeline: Any, - scenario: OmnidreamsSessionScenario, - initial_inputs: InferenceInput, - device: torch.device, - is_rank_zero: bool, - output_layout: VideoTensorLayout, - rollout_seed: int | None, - release_oneshot_encoders_after_cache_init: bool, - ) -> None: - self.pipeline = pipeline - self.scenario = scenario - self._initial_inputs = initial_inputs - self.device = device - self.is_rank_zero = is_rank_zero - self.output_layout = output_layout - self.rollout_seed = rollout_seed - self.release_oneshot_encoders_after_cache_init = ( - release_oneshot_encoders_after_cache_init - ) - self.dtype = torch.bfloat16 - self._closed = False - self._model_session = OmnidreamsModelSessionCore( - pipeline=pipeline, - output_stream_factory=lambda: VideoOutputStream( - postprocess_stream=None, - output_layout=self.output_layout, - ), - ) - self._model_session.reset(self._initialize_cache) - if self.device.type == "cuda" and torch.cuda.is_available(): - torch.cuda.synchronize(device=self.device) - if dist.is_initialized(): - dist.barrier() - - def next_step_requirements(self) -> StepRequirements | None: - if self._closed: - return None - step_index = self._model_session.step_index - if step_index >= self.scenario.total_blocks: - return None - num_frames = self._model_session.next_num_frames() - return StepRequirements( - step_index=step_index, - input_frame_count=num_frames, - ) - - def next_step_request(self) -> StepRequest | None: - requirements = self.next_step_requirements() - if requirements is None: - return None - metadata = dict(requirements.metadata) - metadata["input_frame_count"] = requirements.input_frame_count - if requirements.steady_output_frame_count is not None: - metadata["steady_output_frame_count"] = ( - requirements.steady_output_frame_count - ) - return StepRequest( - step_index=requirements.step_index, - inference_input_schema=requirements.inference_input_schema, - metadata=metadata, - ) - - def step(self, inputs: InferenceInput) -> StepResult: - if self._closed: - raise RuntimeError("OmniDreams replay session is closed.") - - step_index = self._model_session.step_index - num_frames = self._model_session.next_num_frames() - hdmap = _hdmap_from_inputs(inputs) - if hdmap.shape[2] != num_frames: - raise ValueError( - "OmniDreams step HDMap frame count mismatch: " - f"expected {num_frames}, got {hdmap.shape[2]}." - ) - logger.info( - "OmniDreams demo replay step {} frames={}", - step_index, - num_frames, - ) - return self._model_session.step(hdmap) - - def reset(self, inputs: InferenceInput | None = None) -> None: - if inputs is not None: - scenario = _scenario_from_inputs(inputs) - if scenario != self.scenario: - raise ValueError("OmniDreams replay reset cannot swap scenarios.") - self._initial_inputs = inputs - self._model_session.reset(self._initialize_cache) - - def close(self) -> None: - self._closed = True - self._model_session.close() - - def _initialize_cache(self) -> Any: - scenario = self.scenario - _seed_pipeline_for_rollout(self.pipeline, self.rollout_seed) - cache = self.pipeline.initialize_cache( - text=_prompt_from_inputs(self._initial_inputs, scenario), - image=_first_frame_from_inputs( - self._initial_inputs, - scenario=scenario, - device=self.device, - dtype=self.dtype, - ), - view_names=_view_names_from_inputs(self._initial_inputs, scenario), - ) - if self.release_oneshot_encoders_after_cache_init: - release = getattr(self.pipeline, "release_oneshot_encoders", None) - if callable(release): - release() - return cache - - -def _default_pipeline_factory(pipeline_config: Any, device: str) -> Any: - return pipeline_config.setup().to(device=device).eval() - - -def _scenario_from_inputs(inputs: InferenceInput) -> OmnidreamsSessionScenario: - scenario = inputs.global_conditioning.get("scenario") - if not isinstance( - scenario, - (OmnidreamsReplayScenario, OmnidreamsLudusReplayScenario), - ): - raise TypeError( - "OmniDreams replay runtime requires global_conditioning['scenario'] " - "to be an OmnidreamsReplayScenario or OmnidreamsLudusReplayScenario." - ) - return scenario - - -def _prompt_from_inputs( - inputs: InferenceInput, - scenario: OmnidreamsSessionScenario, -) -> list[list[str]]: - prompt = inputs.global_conditioning.get("prompt") - if prompt is None: - if not scenario.prompts: - raise ValueError( - "OmniDreams initial prompt is required when the scenario does " - "not carry fallback prompts." - ) - return [list(scenario.prompts)] - if isinstance(prompt, str): - return [[prompt]] - if isinstance(prompt, Sequence): - values = list(prompt) - if all(isinstance(value, str) for value in values): - return [[str(value) for value in values]] - batches: list[list[str]] = [] - for batch in values: - if not isinstance(batch, Sequence) or isinstance(batch, str): - raise TypeError( - "OmniDreams initial prompt batches must be string sequences." - ) - batches.append([str(item) for item in batch]) - return batches - raise TypeError( - "OmniDreams initial prompt must be a string or sequence of strings." - ) - - -def _first_frame_from_inputs( - inputs: InferenceInput, - *, - scenario: OmnidreamsSessionScenario, - device: torch.device, - dtype: torch.dtype, -) -> torch.Tensor: - first_frame = inputs.global_conditioning.get("first_frame") - if isinstance(first_frame, torch.Tensor): - return first_frame - if first_frame is not None: - raise TypeError("OmniDreams initial first_frame must be a torch.Tensor.") - first_frame_paths = getattr(scenario, "first_frame_paths", ()) - if not first_frame_paths: - raise ValueError( - "OmniDreams initial first_frame tensor is required when the " - "scenario does not carry fallback first_frame_paths." - ) - first_frames = [ - load_first_frame_tensor( - path, - pixel_height=scenario.pixel_height, - pixel_width=scenario.pixel_width, - device=device, - dtype=dtype, - allow_video=True, - install_hint=DEFAULT_RUNNER_INSTALL_HINT, - ) - for path in first_frame_paths - ] - return torch.stack(first_frames, dim=0).unsqueeze(0) - - -def _seed_pipeline_for_rollout(pipeline: Any, seed: int | None) -> None: - if seed is None: - return - diffusion_model = getattr(pipeline, "diffusion_model", None) - rng = getattr(diffusion_model, "rng", None) - if rng is None: - return - rng.manual_seed(int(seed)) - - -def _view_names_from_inputs( - inputs: InferenceInput, - scenario: OmnidreamsSessionScenario, -) -> list[str]: - value = inputs.metadata.get("view_names") or inputs.global_conditioning.get( - "view_names" - ) - if value is None: - return list(scenario.camera_names) - if isinstance(value, str): - return [value] - if isinstance(value, Sequence): - return [str(item) for item in value] - raise TypeError("OmniDreams view_names metadata must be a string sequence.") - - -def _hdmap_from_inputs(inputs: InferenceInput) -> torch.Tensor: - hdmap = inputs.step.get("hdmap") - if not isinstance(hdmap, torch.Tensor): - raise TypeError("OmniDreams session step requires step['hdmap'] tensor.") - if hdmap.ndim != 6: - raise ValueError( - "OmniDreams step['hdmap'] must have shape [B, V, T, C, H, W], " - f"got {tuple(hdmap.shape)}." - ) - return hdmap - - -def _is_torchrun_env() -> bool: - return "RANK" in os.environ and "WORLD_SIZE" in os.environ - OmnidreamsReplayRuntimeOptions = OmnidreamsRuntimeOptions OmnidreamsReplayRuntime = OmnidreamsRuntime OmnidreamsReplaySession = OmnidreamsSession - __all__ = [ + "OmnidreamsReplayRuntime", + "OmnidreamsReplayRuntimeOptions", + "OmnidreamsReplaySession", "OmnidreamsRuntime", "OmnidreamsRuntimeOptions", "OmnidreamsSession", "OmnidreamsSessionScenario", - "OmnidreamsReplayRuntime", - "OmnidreamsReplayRuntimeOptions", - "OmnidreamsReplaySession", "PipelineFactory", ] diff --git a/integrations/omnidreams/omnidreams/demo/runtime.py b/integrations/omnidreams/omnidreams/demo/runtime.py index eae00ff93..1ee99ad31 100644 --- a/integrations/omnidreams/omnidreams/demo/runtime.py +++ b/integrations/omnidreams/omnidreams/demo/runtime.py @@ -3,16 +3,351 @@ """OmniDreams runtime/session contracts for shared demo run modes.""" -from omnidreams.demo.replay import ( - OmnidreamsRuntime, - OmnidreamsRuntimeOptions, - OmnidreamsSession, - PipelineFactory, +from __future__ import annotations + +import os +from collections.abc import Callable, Sequence +from dataclasses import dataclass +from typing import Any + +import torch +import torch.distributed as dist +from loguru import logger +from omnidreams.model_session import OmnidreamsModelSessionCore + +from flashdreams.core.distributed import init as init_distributed +from flashdreams.infra.postprocess import VideoTensorLayout +from flashdreams.infra.runner_io import ( + DEFAULT_RUNNER_INSTALL_HINT, + load_first_frame_tensor, ) +from flashdreams.infra.video_output import VideoOutputStream +from flashdreams.runtime.config import InferenceConfig +from flashdreams.runtime.inputs import InferenceInput +from flashdreams.runtime.interfaces import InferenceSession +from flashdreams.runtime.types import StepRequest, StepRequirements, StepResult + +from .spec import OmnidreamsLudusReplayScenario, OmnidreamsReplayScenario + +OmnidreamsSessionScenario = OmnidreamsReplayScenario | OmnidreamsLudusReplayScenario + +PipelineFactory = Callable[[Any, str], Any] + + +@dataclass(frozen=True, kw_only=True, slots=True) +class OmnidreamsRuntimeOptions: + """Construction knobs for the OmniDreams runtime.""" + + pipeline_config: Any + pipeline_factory: PipelineFactory | None = None + output_layout: VideoTensorLayout = "bvtchw" + release_oneshot_encoders_after_cache_init: bool = True + + +class OmnidreamsRuntime: + """Heavyweight OmniDreams runtime consumed by shared demo run modes.""" + + def __init__( + self, + *, + config: InferenceConfig, + options: OmnidreamsRuntimeOptions, + ) -> None: + self.config = config + self.options = options + if _is_torchrun_env() and not dist.is_initialized(): + init_distributed() + + if dist.is_initialized(): + self.local_rank = int(os.environ.get("LOCAL_RANK", "0")) + self.world_size = dist.get_world_size() + self.global_rank = dist.get_rank() + device = f"cuda:{self.local_rank}" + else: + self.local_rank = 0 + self.world_size = 1 + self.global_rank = 0 + device = config.device or "cuda" + + self.is_rank_zero = self.global_rank == 0 + factory = options.pipeline_factory or _default_pipeline_factory + self.pipeline = factory(options.pipeline_config, device) + + def start_session(self, inputs: InferenceInput) -> InferenceSession: + scenario = _scenario_from_inputs(inputs) + return OmnidreamsSession( + pipeline=self.pipeline, + scenario=scenario, + initial_inputs=inputs, + device=torch.device(f"cuda:{self.local_rank}") + if dist.is_initialized() + else torch.device(self.config.device or "cuda"), + is_rank_zero=self.is_rank_zero, + output_layout=self.options.output_layout, + rollout_seed=self.config.seed, + release_oneshot_encoders_after_cache_init=( + self.options.release_oneshot_encoders_after_cache_init + ), + ) + + def close(self) -> None: + pipeline = getattr(self, "pipeline", None) + if pipeline is not None: + close = getattr(pipeline, "close", None) + if callable(close): + close() + del self.pipeline + device = torch.device(self.config.device or "cuda") + if device.type == "cuda" and torch.cuda.is_available(): + torch.cuda.empty_cache() + + +class OmnidreamsSession: + """One OmniDreams rollout over a prepared scenario.""" + + def __init__( + self, + *, + pipeline: Any, + scenario: OmnidreamsSessionScenario, + initial_inputs: InferenceInput, + device: torch.device, + is_rank_zero: bool, + output_layout: VideoTensorLayout, + rollout_seed: int | None, + release_oneshot_encoders_after_cache_init: bool, + ) -> None: + self.pipeline = pipeline + self.scenario = scenario + self._initial_inputs = initial_inputs + self.device = device + self.is_rank_zero = is_rank_zero + self.output_layout = output_layout + self.rollout_seed = rollout_seed + self.release_oneshot_encoders_after_cache_init = ( + release_oneshot_encoders_after_cache_init + ) + self.dtype = torch.bfloat16 + self._closed = False + self._model_session = OmnidreamsModelSessionCore( + pipeline=pipeline, + output_stream_factory=lambda: VideoOutputStream( + postprocess_stream=None, + output_layout=self.output_layout, + ), + ) + self._model_session.reset(self._initialize_cache) + if self.device.type == "cuda" and torch.cuda.is_available(): + torch.cuda.synchronize(device=self.device) + if dist.is_initialized(): + dist.barrier() + + def next_step_requirements(self) -> StepRequirements | None: + if self._closed: + return None + step_index = self._model_session.step_index + if step_index >= self.scenario.total_blocks: + return None + num_frames = self._model_session.next_num_frames() + return StepRequirements( + step_index=step_index, + input_frame_count=num_frames, + ) + + def next_step_request(self) -> StepRequest | None: + requirements = self.next_step_requirements() + if requirements is None: + return None + metadata = dict(requirements.metadata) + metadata["input_frame_count"] = requirements.input_frame_count + if requirements.steady_output_frame_count is not None: + metadata["steady_output_frame_count"] = ( + requirements.steady_output_frame_count + ) + return StepRequest( + step_index=requirements.step_index, + inference_input_schema=requirements.inference_input_schema, + metadata=metadata, + ) + + def step(self, inputs: InferenceInput) -> StepResult: + if self._closed: + raise RuntimeError("OmniDreams replay session is closed.") + + step_index = self._model_session.step_index + num_frames = self._model_session.next_num_frames() + hdmap = _hdmap_from_inputs(inputs) + if hdmap.shape[2] != num_frames: + raise ValueError( + "OmniDreams step HDMap frame count mismatch: " + f"expected {num_frames}, got {hdmap.shape[2]}." + ) + logger.info( + "OmniDreams demo replay step {} frames={}", + step_index, + num_frames, + ) + return self._model_session.step(hdmap) + + def reset(self, inputs: InferenceInput | None = None) -> None: + if inputs is not None: + scenario = _scenario_from_inputs(inputs) + if scenario != self.scenario: + raise ValueError("OmniDreams replay reset cannot swap scenarios.") + self._initial_inputs = inputs + self._model_session.reset(self._initialize_cache) + + def close(self) -> None: + self._closed = True + self._model_session.close() + + def _initialize_cache(self) -> Any: + scenario = self.scenario + _seed_pipeline_for_rollout(self.pipeline, self.rollout_seed) + cache = self.pipeline.initialize_cache( + text=_prompt_from_inputs(self._initial_inputs, scenario), + image=_first_frame_from_inputs( + self._initial_inputs, + scenario=scenario, + device=self.device, + dtype=self.dtype, + ), + view_names=_view_names_from_inputs(self._initial_inputs, scenario), + ) + if self.release_oneshot_encoders_after_cache_init: + release = getattr(self.pipeline, "release_oneshot_encoders", None) + if callable(release): + release() + return cache + + +def _default_pipeline_factory(pipeline_config: Any, device: str) -> Any: + return pipeline_config.setup().to(device=device).eval() + + +def _scenario_from_inputs(inputs: InferenceInput) -> OmnidreamsSessionScenario: + scenario = inputs.global_conditioning.get("scenario") + if not isinstance( + scenario, + (OmnidreamsReplayScenario, OmnidreamsLudusReplayScenario), + ): + raise TypeError( + "OmniDreams replay runtime requires global_conditioning['scenario'] " + "to be an OmnidreamsReplayScenario or OmnidreamsLudusReplayScenario." + ) + return scenario + + +def _prompt_from_inputs( + inputs: InferenceInput, + scenario: OmnidreamsSessionScenario, +) -> list[list[str]]: + prompt = inputs.global_conditioning.get("prompt") + if prompt is None: + if not scenario.prompts: + raise ValueError( + "OmniDreams initial prompt is required when the scenario does " + "not carry fallback prompts." + ) + return [list(scenario.prompts)] + if isinstance(prompt, str): + return [[prompt]] + if isinstance(prompt, Sequence): + values = list(prompt) + if all(isinstance(value, str) for value in values): + return [[str(value) for value in values]] + batches: list[list[str]] = [] + for batch in values: + if not isinstance(batch, Sequence) or isinstance(batch, str): + raise TypeError( + "OmniDreams initial prompt batches must be string sequences." + ) + batches.append([str(item) for item in batch]) + return batches + raise TypeError( + "OmniDreams initial prompt must be a string or sequence of strings." + ) + + +def _first_frame_from_inputs( + inputs: InferenceInput, + *, + scenario: OmnidreamsSessionScenario, + device: torch.device, + dtype: torch.dtype, +) -> torch.Tensor: + first_frame = inputs.global_conditioning.get("first_frame") + if isinstance(first_frame, torch.Tensor): + return first_frame + if first_frame is not None: + raise TypeError("OmniDreams initial first_frame must be a torch.Tensor.") + first_frame_paths = getattr(scenario, "first_frame_paths", ()) + if not first_frame_paths: + raise ValueError( + "OmniDreams initial first_frame tensor is required when the " + "scenario does not carry fallback first_frame_paths." + ) + first_frames = [ + load_first_frame_tensor( + path, + pixel_height=scenario.pixel_height, + pixel_width=scenario.pixel_width, + device=device, + dtype=dtype, + allow_video=True, + install_hint=DEFAULT_RUNNER_INSTALL_HINT, + ) + for path in first_frame_paths + ] + return torch.stack(first_frames, dim=0).unsqueeze(0) + + +def _seed_pipeline_for_rollout(pipeline: Any, seed: int | None) -> None: + if seed is None: + return + diffusion_model = getattr(pipeline, "diffusion_model", None) + rng = getattr(diffusion_model, "rng", None) + if rng is None: + return + rng.manual_seed(int(seed)) + + +def _view_names_from_inputs( + inputs: InferenceInput, + scenario: OmnidreamsSessionScenario, +) -> list[str]: + value = inputs.metadata.get("view_names") or inputs.global_conditioning.get( + "view_names" + ) + if value is None: + return list(scenario.camera_names) + if isinstance(value, str): + return [value] + if isinstance(value, Sequence): + return [str(item) for item in value] + raise TypeError("OmniDreams view_names metadata must be a string sequence.") + + +def _hdmap_from_inputs(inputs: InferenceInput) -> torch.Tensor: + hdmap = inputs.step.get("hdmap") + if not isinstance(hdmap, torch.Tensor): + raise TypeError("OmniDreams session step requires step['hdmap'] tensor.") + if hdmap.ndim != 6: + raise ValueError( + "OmniDreams step['hdmap'] must have shape [B, V, T, C, H, W], " + f"got {tuple(hdmap.shape)}." + ) + return hdmap + + +def _is_torchrun_env() -> bool: + return "RANK" in os.environ and "WORLD_SIZE" in os.environ + __all__ = [ "OmnidreamsRuntime", "OmnidreamsRuntimeOptions", "OmnidreamsSession", + "OmnidreamsSessionScenario", "PipelineFactory", ] diff --git a/integrations/omnidreams/omnidreams/demo/webrtc.py b/integrations/omnidreams/omnidreams/demo/webrtc.py index d0510fa1d..e5d8e2a84 100644 --- a/integrations/omnidreams/omnidreams/demo/webrtc.py +++ b/integrations/omnidreams/omnidreams/demo/webrtc.py @@ -13,741 +13,45 @@ # See the License for the specific language governing permissions and # limitations under the License. -"""OmniDreams model runtime and browser hooks for the shared WebRTC demo.""" +"""OmniDreams browser hooks for the shared WebRTC demo runtime.""" from __future__ import annotations -import math import os -from collections.abc import Callable, Mapping -from dataclasses import dataclass, replace -from pathlib import Path +from collections.abc import Callable +from dataclasses import replace from typing import Any -import torch from loguru import logger from omnidreams.config import OMNIDREAMS_CONFIGS -from omnidreams.scenes import SCENE_VARIANT_DEFAULT -from omnidreams.transformer import CosmosTransformerConfig -from flashdreams.core.distributed.rank_orchestration import distributed_op -from flashdreams.runtime import ( - CanonicalInputs, - CanonicalInputSchema, - InferenceConfig, - InferenceInput, - InferenceInputSchema, - InputCanonicalizer, - StepRequest, - StepRequirements, - StepResult, - TimeWindow, - step_requirements_from_request, -) +from flashdreams.runtime import InferenceConfig from flashdreams.runtime.demo import ( DemoSpec, - PreparedScenario, RuntimeHost, - SessionInfo, - UserInputWindow, WebRTCAppResources, WebRTCOutputSpec, ) -from flashdreams.runtime.demo.timing import SPARSE_KEY_SEGMENTS_METADATA_KEY from flashdreams.runtime.demo.webrtc import ( CreateWebRTCApp, RunWebRTCServer, serve_webrtc_demo, ) from flashdreams.serving.webrtc.bootstrap import run_webrtc_server -from flashdreams.serving.webrtc.controls import WSAD_SUPPORTED_KEYS, PoseSegment -from flashdreams.serving.webrtc.encoders import EncoderBackend +from flashdreams.serving.webrtc.controls import WSAD_SUPPORTED_KEYS from flashdreams.serving.webrtc.manager import BaseWebRTCSessionManager -from flashdreams.serving.webrtc.runtime import ( - ThreadAffineDistributedWebRTCRuntime, - WebRTCControlSignal, -) from flashdreams.serving.webrtc.server import create_webrtc_app -from flashdreams.serving.webrtc.services import WEBRTC_USER_INPUT_SCHEMA from .adapter import OmnidreamsDemoAdapter, RuntimeFactory -from .providers import LudusSceneConditioningProvider -from .runtime import OmnidreamsRuntime, OmnidreamsRuntimeOptions, PipelineFactory from .spec import ( DEFAULT_OMNIDREAMS_PRESET, - DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, OMNIDREAMS_MODEL_ID, - OmnidreamsLudusReplayScenario, resolve_webrtc_scenario, ) +from .webrtc_config import OmnidreamsWebRTCModelRuntimeConfig WebRTCRuntimeFactory = Callable[..., Any] SharedRuntimeFactory = RuntimeFactory -_WEBRTC_SESSION_TOTAL_BLOCKS = 2_147_483_647 -_WEBRTC_STEP_REQUEST_KEY = "omnidreams_webrtc_step_request" - - -class OmnidreamsWebRTCModelRuntimeError(RuntimeError): - """Raised when the OmniDreams demo runtime is used incorrectly.""" - - -@dataclass(frozen=True, slots=True) -class OmnidreamsWebRTCModelRuntimeConfig: - """Configuration for one scene-driven OmniDreams WebRTC runtime.""" - - pipeline_config_name: str - """User-facing name of the selected OmniDreams pipeline.""" - - pipeline_config: Any - """Resolved single-view OmniDreams pipeline configuration.""" - - scene_dir: Path | None = None - """Local scene root; ``None`` downloads the selected Hugging Face scene.""" - - scene_uuid: str | None = DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID - """Scene UUID used for remote lookup or local archive selection.""" - - scene_variant: str = SCENE_VARIANT_DEFAULT - """Weather variant selected from the scene assets.""" - - seed: int | None = 42 - """Per-rollout seed; ``None`` selects fresh entropy for every session.""" - - device: str = "cuda:0" - """Device used for rendering and model inference.""" - - video_height: int = 704 - """Generated video height in pixels.""" - - video_width: int = 1280 - """Generated video width in pixels.""" - - fps: int = 30 - """Input sampling and output playback frame rate.""" - - camera_name: str = "camera_front_wide_120fov" - """Scene camera controlled by browser keyboard input.""" - - move_speed_per_s: float = 6.0 - """Forward and reverse translation speed in scene units per second.""" - - rotate_speed_rad_per_s: float = math.radians(35.0) - """Left and right rotation speed in radians per second.""" - - warmup_chunks: int = 10 - """Number of synthetic chunks generated before accepting sessions.""" - - warmup_timeout_s: float = 600.0 - """Maximum duration for WebRTC loopback warmup.""" - - debug_serve_hdmaps: bool = False - """Stream rendered conditioning frames without running video generation.""" - - encoder_backend: EncoderBackend = "auto" - """WebRTC video encoder selection policy.""" - - encoder_bitrate_bps: int = 6_000_000 - """Target WebRTC video bitrate in bits per second.""" - - encoder_gop: int = 30 - """WebRTC video encoder group-of-pictures length.""" - - pipeline_factory: PipelineFactory | None = None - """Optional test/runtime override for constructing the shared pipeline.""" - - -class OmnidreamsWebRTCModelRuntime( - ThreadAffineDistributedWebRTCRuntime[ - OmnidreamsWebRTCModelRuntimeConfig, - None, - ] -): - """Compatibility WebRTC facade over the shared OmniDreams runtime/session.""" - - def __init__(self, *, config: OmnidreamsWebRTCModelRuntimeConfig) -> None: - super().__init__( - config=config, - runtime_error_type=OmnidreamsWebRTCModelRuntimeError, - thread_name="omnidreams-demo-runtime", - ) - # The shared WebRTC input source emits normalized runtime events - # (``key_down``/``key_up``). The Ludus provider consumes sparse - # resampler metadata on this transitional path, so keep validation - # aligned with the WebRTC source rather than the replay trace schema. - self.input_source_schema = WEBRTC_USER_INPUT_SCHEMA - self.input_canonicalizer = InputCanonicalizer() - self.input_mapping = _OmnidreamsWebRTCInputMapping() - self._runtime: OmnidreamsRuntime | None = None - self._active_provider: LudusSceneConditioningProvider | None = None - self._active_session: Any | None = None - self._debug_session: _OmnidreamsHDMapDebugSession | None = None - self._steady_output_frame_count_value = 1 - - def _is_runtime_initialized(self) -> bool: - return self._runtime is not None - - def _runtime_step_index(self) -> int: - requirements = self._next_step_requirements_sync() - if requirements is None: - return 0 - return requirements.step_index - - def _next_input_frame_count(self) -> int: - requirements = self._next_step_requirements_sync() - if requirements is None: - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC session is complete." - ) - return requirements.input_frame_count - - def _steady_output_frame_count(self) -> int: - return self._steady_output_frame_count_value - - def _initialize_sync(self) -> None: - if self._runtime is not None: - return - if self._device.type == "cuda" and not torch.cuda.is_available(): - raise RuntimeError("CUDA is required for OmniDreams WebRTC inference.") - _validate_single_view_pipeline_config( - pipeline_config_name=self.config.pipeline_config_name, - pipeline_config=self.config.pipeline_config, - ) - logger.info( - "Setting up shared OmniDreams runtime {} on {} for WebRTC.", - self.config.pipeline_config_name, - self._device, - ) - self._runtime = OmnidreamsRuntime( - config=self._inference_config(), - options=OmnidreamsRuntimeOptions( - pipeline_config=self.config.pipeline_config, - pipeline_factory=self.config.pipeline_factory, - # WebRTC warms the same long-lived runtime before real browser - # sessions. Keep prompt/image encoders available for later - # peer connections until Phase 14 replaces loopback warmup with - # first-class model/runtime warmup. - release_oneshot_encoders_after_cache_init=False, - ), - ) - self._initialize_video_encoder_sync() - - def _reset_rollout_sync(self, session_input: None = None) -> None: - del session_input - self._close_active_session_sync() - runtime = self._require_runtime() - scenario = self._session_scenario() - prepared = PreparedScenario( - initial_inputs=InferenceInput(global_conditioning={"scenario": scenario}), - source_schema=self.input_source_schema, - metadata={ - "conditioning_mode": "ludus-scene-driving", - "model_id": OMNIDREAMS_MODEL_ID, - "preset_id": self.config.pipeline_config_name, - }, - ) - provider = LudusSceneConditioningProvider( - scenario=prepared, - config=self._inference_config(), - ) - try: - initial_input = provider.prepare_initial_input() - session = runtime.start_session(initial_input) - except Exception: - provider.close() - raise - self._active_provider = provider - if self.config.debug_serve_hdmaps: - self._debug_session = _OmnidreamsHDMapDebugSession( - pipeline=runtime.pipeline, - scenario=scenario, - ) - self._active_session = self._debug_session - else: - self._debug_session = None - self._active_session = session - self._steady_output_frame_count_value = _steady_output_frame_count( - self._active_session, - fallback_pipeline=runtime.pipeline, - ) - - def _generate_one_chunk_sync( - self, - *, - segments: list[PoseSegment], - frame_times: list[float], - ) -> StepResult: - request = self._next_step_request_sync() - if request is None: - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC session is complete." - ) - inputs = self.input_mapping.map_step_inputs( - canonical_inputs=CanonicalInputs(), - inference_input=InferenceInput( - metadata={ - SPARSE_KEY_SEGMENTS_METADATA_KEY: tuple(segments), - "frame_times": tuple(frame_times), - "window_start_s": request.step_index / float(self.config.fps), - "window_end_s": (request.step_index + len(frame_times)) - / float(self.config.fps), - } - ), - request=request, - ) - return self._step_active_session_sync(inputs) - - def _close_sync(self) -> None: - self._close_active_session_sync() - runtime = self._runtime - self._runtime = None - if runtime is not None: - runtime.close() - if self._device.type == "cuda" and torch.cuda.is_available(): - torch.cuda.synchronize(device=self._device) - torch.cuda.empty_cache() - - async def start_inference_session(self) -> "_OmnidreamsWebRTCInferenceSession": - self._require_open_and_initialized() - if not await self._worker.call(self._has_active_session_sync): - await self.reset_for_new_session() - return _OmnidreamsWebRTCInferenceSession(self) - - def _next_step_request_sync(self) -> StepRequest | None: - requirements = self._next_step_requirements_sync() - if requirements is None: - return None - metadata = dict(requirements.metadata) - metadata["input_frame_count"] = requirements.input_frame_count - if requirements.steady_output_frame_count is not None: - metadata["steady_output_frame_count"] = ( - requirements.steady_output_frame_count - ) - return StepRequest( - step_index=requirements.step_index, - inference_input_schema=requirements.inference_input_schema, - metadata=metadata, - ) - - def _next_step_requirements_sync(self) -> StepRequirements | None: - session = self._require_active_session() - next_requirements = getattr(session, "next_step_requirements", None) - if callable(next_requirements): - result = next_requirements() - else: - next_request = session.next_step_request() - if next_request is None: - return None - result = step_requirements_from_request(next_request) - if result is None: - return None - if not isinstance(result, StepRequirements): - raise TypeError( - "OmniDreams WebRTC session requirements must be StepRequirements, " - f"got {type(result).__name__}." - ) - return result - - def _session_info_sync(self) -> SessionInfo: - return SessionInfo( - output_layout="bvtchw", - steady_output_frame_count=self._steady_output_frame_count(), - metadata={"model_id": OMNIDREAMS_MODEL_ID}, - ) - - def _step_active_session_sync(self, inputs: InferenceInput) -> StepResult: - provider = self._require_active_provider() - session = self._require_active_session() - request = _request_from_step_inputs(inputs) - requirements = step_requirements_from_request( - request, - allow_user_input_window=True, - ) - window = _user_window_from_step_inputs( - inputs, - request=request, - input_frame_count=requirements.input_frame_count, - ) - prepared = provider.prepare_step(request=requirements, user_window=window) - if prepared.control.close_session: - raise OmnidreamsWebRTCModelRuntimeError( - prepared.control.reason or "OmniDreams WebRTC input is exhausted." - ) - if prepared.control.reset: - reset_input = prepared.control.reset_input - session.reset(reset_input) - if not prepared.control.provider_already_reset: - provider.reset(reset_input) - raise OmnidreamsWebRTCModelRuntimeError( - prepared.control.reason or "OmniDreams WebRTC session reset requested." - ) - if prepared.inference_input is None: - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC provider returned no inference input." - ) - result = session.step(prepared.inference_input) - if not isinstance(result, StepResult): - raise TypeError( - "OmniDreams WebRTC session steps must produce StepResult, got " - f"{type(result).__name__}." - ) - return result - - @distributed_op(WebRTCControlSignal.SESSION_STEP) - def _step_active_session_sync_all_ranks( - self, - inputs: InferenceInput, - ) -> StepResult: - return self._step_active_session_sync(inputs) - - @distributed_op(WebRTCControlSignal.SESSION_CLOSE) - def _close_active_session_sync_all_ranks(self) -> None: - self._close_active_session_sync() - - def _close_active_session_sync(self) -> None: - session = self._active_session - provider = self._active_provider - self._active_session = None - self._debug_session = None - self._active_provider = None - first_error: Exception | None = None - close_session = getattr(session, "close", None) - if callable(close_session): - try: - close_session() - except Exception as exc: - first_error = exc - if provider is not None: - try: - provider.close() - except Exception as exc: - if first_error is None: - first_error = exc - if first_error is not None: - raise first_error - - def _has_active_session_sync(self) -> bool: - return self._active_session is not None and self._active_provider is not None - - def _require_runtime(self) -> OmnidreamsRuntime: - if self._runtime is None: - raise OmnidreamsWebRTCModelRuntimeError("Runtime is not initialized.") - return self._runtime - - def _require_active_session(self) -> Any: - if self._active_session is None: - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC session is not initialized." - ) - return self._active_session - - def _require_active_provider(self) -> LudusSceneConditioningProvider: - if self._active_provider is None: - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC provider is not initialized." - ) - return self._active_provider - - def _inference_config(self) -> InferenceConfig: - return InferenceConfig( - model_id=OMNIDREAMS_MODEL_ID, - preset_id=self.config.pipeline_config_name, - device=str(self.config.device), - seed=self.config.seed, - runtime_options={"seed": self.config.seed}, - ) - - def _session_scenario(self) -> OmnidreamsLudusReplayScenario: - return OmnidreamsLudusReplayScenario( - keyboard_events=(), - scene_dir=self.config.scene_dir, - scene_uuid=self.config.scene_uuid or DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, - scene_variant=self.config.scene_variant, - camera_name=self.config.camera_name, - total_blocks=_WEBRTC_SESSION_TOTAL_BLOCKS, - pixel_height=self.config.video_height, - pixel_width=self.config.video_width, - fps=self.config.fps, - move_speed_per_s=self.config.move_speed_per_s, - rotate_speed_rad_per_s=self.config.rotate_speed_rad_per_s, - ) - - -class _OmnidreamsWebRTCInputMapping: - """Carry shared WebRTC window facts into the OmniDreams session facade.""" - - def validate( - self, - *, - canonical_schema: CanonicalInputSchema | None = None, - inference_input_schema: InferenceInputSchema | None = None, - ) -> None: - del canonical_schema, inference_input_schema - - def map_global_conditioning_inputs( - self, - *, - canonical_inputs: CanonicalInputs, - inference_input: InferenceInput, - ) -> InferenceInput: - del canonical_inputs - return inference_input - - def map_step_inputs( - self, - *, - canonical_inputs: CanonicalInputs, - inference_input: InferenceInput, - request: StepRequest, - ) -> InferenceInput: - del canonical_inputs - step = dict(inference_input.step) - step[_WEBRTC_STEP_REQUEST_KEY] = request - return InferenceInput( - global_conditioning=inference_input.global_conditioning, - step=step, - metadata=inference_input.metadata, - ) - - -class _OmnidreamsWebRTCInferenceSession: - """Synchronous session proxy consumed by the shared WebRTC compatibility path.""" - - def __init__(self, runtime: OmnidreamsWebRTCModelRuntime) -> None: - self._runtime = runtime - self._closed = False - - def session_info(self) -> SessionInfo: - self._require_open() - return self._runtime._worker.call_blocking(self._runtime._session_info_sync) - - def next_step_requirements(self) -> StepRequirements | None: - self._require_open() - return self._runtime._worker.call_blocking( - self._runtime._next_step_requirements_sync - ) - - def next_step_request(self) -> StepRequest | None: - self._require_open() - return self._runtime._worker.call_blocking( - self._runtime._next_step_request_sync - ) - - def step(self, inputs: InferenceInput) -> StepResult: - self._require_open() - return self._runtime._worker.call_blocking( - self._runtime._step_active_session_sync_all_ranks, - inputs, - ) - - def reset(self, inputs: InferenceInput | None = None) -> None: - del inputs - self._require_open() - self._runtime._worker.call_blocking(self._runtime._reset_rollout_sync_all_ranks) - - def close(self) -> None: - if self._closed: - return - self._closed = True - self._runtime._worker.call_blocking( - self._runtime._close_active_session_sync_all_ranks - ) - - def _require_open(self) -> None: - if self._closed: - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC inference session is closed." - ) - - -class _OmnidreamsHDMapDebugSession: - """Session-shaped debug path that streams rendered Ludus HDMaps.""" - - def __init__( - self, *, pipeline: Any, scenario: OmnidreamsLudusReplayScenario - ) -> None: - self._pipeline = pipeline - self._scenario = scenario - self._step_index = 0 - self._closed = False - - def session_info(self) -> SessionInfo: - return SessionInfo( - output_layout="bvtchw", - steady_output_frame_count=self._steady_output_frame_count(), - metadata={"stream": "hdmap"}, - ) - - def next_step_requirements(self) -> StepRequirements | None: - if self._closed or self._step_index >= self._scenario.total_blocks: - return None - return StepRequirements( - step_index=self._step_index, - input_frame_count=self._num_frames(self._step_index), - steady_output_frame_count=self._steady_output_frame_count(), - ) - - def next_step_request(self) -> StepRequest | None: - requirements = self.next_step_requirements() - if requirements is None: - return None - return StepRequest( - step_index=requirements.step_index, - metadata={ - "input_frame_count": requirements.input_frame_count, - "steady_output_frame_count": requirements.steady_output_frame_count, - }, - ) - - def step(self, inputs: InferenceInput) -> StepResult: - requirements = self.next_step_requirements() - if requirements is None: - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC debug session is complete." - ) - hdmap = inputs.step.get("hdmap") - if not isinstance(hdmap, torch.Tensor): - raise TypeError("OmniDreams WebRTC debug session requires step['hdmap'].") - result = StepResult.from_video_chunk( - step_index=requirements.step_index, - video_chunk=hdmap.detach(), - layout="bvtchw", - metadata={"stream": "hdmap"}, - ) - self._step_index += 1 - return result - - def reset(self, inputs: InferenceInput | None = None) -> None: - del inputs - self._step_index = 0 - self._closed = False - - def close(self) -> None: - self._closed = True - - def _steady_output_frame_count(self) -> int: - return self._num_frames(1) - - def _num_frames(self, step_index: int) -> int: - get_num_frames = getattr(self._pipeline, "get_num_frames", None) - if not callable(get_num_frames): - return 1 - return int(get_num_frames(step_index)) - - -def _request_from_step_inputs(inputs: InferenceInput) -> StepRequest: - request = inputs.step.get(_WEBRTC_STEP_REQUEST_KEY) - if not isinstance(request, StepRequest): - raise TypeError( - "OmniDreams WebRTC step input is missing the shared StepRequest." - ) - return request - - -def _user_window_from_step_inputs( - inputs: InferenceInput, - *, - request: StepRequest, - input_frame_count: int, -) -> UserInputWindow: - frame_times = _frame_times_from_metadata(inputs.metadata, input_frame_count) - segments = _segments_from_metadata(inputs.metadata) - window = request.user_input_window or TimeWindow( - start_s=float(inputs.metadata.get("window_start_s", 0.0)), - end_s=float(inputs.metadata.get("window_end_s", frame_times[-1])), - ) - return UserInputWindow( - start_s=window.start_s, - end_s=window.end_s, - frame_times=frame_times, - metadata={SPARSE_KEY_SEGMENTS_METADATA_KEY: segments}, - ) - - -def _frame_times_from_metadata( - metadata: Mapping[str, object], - input_frame_count: int, -) -> tuple[float, ...]: - value = metadata.get("frame_times") - if not isinstance(value, tuple): - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC step input is missing frame_times metadata." - ) - frame_times = tuple( - _float_metadata_value(frame_time, label="frame_times") for frame_time in value - ) - if len(frame_times) != input_frame_count: - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC frame_times length does not match " - f"input_frame_count={input_frame_count}." - ) - return frame_times - - -def _segments_from_metadata(metadata: Mapping[str, object]) -> tuple[PoseSegment, ...]: - value = metadata.get(SPARSE_KEY_SEGMENTS_METADATA_KEY) - if not isinstance(value, tuple): - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC step input is missing resampled key segments." - ) - segments: list[PoseSegment] = [] - for segment in value: - if not isinstance(segment, tuple) or len(segment) != 3: - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC key segments must be 3-tuples." - ) - start_s, end_s, keys = segment - if not isinstance(keys, frozenset | set | tuple | list): - raise OmnidreamsWebRTCModelRuntimeError( - "OmniDreams WebRTC key segment keys must be a sequence." - ) - segments.append( - ( - _float_metadata_value(start_s, label="segment start"), - _float_metadata_value(end_s, label="segment end"), - frozenset(str(key) for key in keys), - ) - ) - return tuple(segments) - - -def _float_metadata_value(value: object, *, label: str) -> float: - if isinstance(value, bool) or not isinstance(value, int | float): - raise OmnidreamsWebRTCModelRuntimeError( - f"OmniDreams WebRTC {label} metadata must be numeric." - ) - return float(value) - - -def _steady_output_frame_count(session: Any, *, fallback_pipeline: Any) -> int: - session_info = getattr(session, "session_info", None) - if callable(session_info): - value = session_info() - if isinstance(value, SessionInfo) and value.steady_output_frame_count: - return int(value.steady_output_frame_count) - get_num_frames = getattr(fallback_pipeline, "get_num_frames", None) - if callable(get_num_frames): - return int(get_num_frames(1)) - return 1 - - -def _validate_single_view_pipeline_config( - *, - pipeline_config_name: str, - pipeline_config: Any, -) -> None: - diffusion_model = getattr(pipeline_config, "diffusion_model", None) - transformer_cfg = getattr(diffusion_model, "transformer", None) - if transformer_cfg is None: - return - if not isinstance(transformer_cfg, CosmosTransformerConfig): - raise TypeError( - "OmniDreams WebRTC requires a CosmosTransformerConfig pipeline." - ) - if transformer_cfg.num_views != 1: - raise ValueError( - "OmniDreams WebRTC supports only single-view configs; " - f"{pipeline_config_name!r} has num_views={transformer_cfg.num_views}." - ) def serve_omnidreams_webrtc_demo( @@ -759,7 +63,7 @@ def serve_omnidreams_webrtc_demo( create_app_fn: CreateWebRTCApp = create_webrtc_app, server_runner: RunWebRTCServer = run_webrtc_server, ) -> object: - """Create OmniDreams' runtime and serve it through the shared WebRTC transport.""" + """Create OmniDreams' runtime and serve it through shared WebRTC transport.""" if runtime_factory is not None and shared_runtime_factory is not None: raise ValueError( "Specify either legacy runtime_factory or shared_runtime_factory, not both." @@ -779,6 +83,7 @@ def serve_omnidreams_webrtc_demo( f"OmniDreams WebRTC requires model_id={OMNIDREAMS_MODEL_ID!r}, " f"got {config.model_id!r}." ) + scenario = resolve_webrtc_scenario(spec.scenario) runtime_config = _webrtc_runtime_config( output=spec.output, @@ -789,6 +94,11 @@ def serve_omnidreams_webrtc_demo( scenario=scenario, runtime_factory=runtime_factory, ): + from .webrtc_legacy import ( # noqa: PLC0415 + OmnidreamsWebRTCModelRuntime, + _serve_legacy_omnidreams_webrtc_demo, + ) + return _serve_legacy_omnidreams_webrtc_demo( spec=spec, output=spec.output, @@ -798,6 +108,7 @@ def serve_omnidreams_webrtc_demo( create_app_fn=create_app_fn, server_runner=server_runner, ) + return _serve_shared_omnidreams_webrtc_demo( spec=_shared_webrtc_spec(spec, runtime_config=runtime_config), output=spec.output, @@ -859,44 +170,6 @@ def _should_use_legacy_webrtc_path( return False -def _serve_legacy_omnidreams_webrtc_demo( - *, - spec: DemoSpec, - output: WebRTCOutputSpec, - runtime_config: OmnidreamsWebRTCModelRuntimeConfig, - runtime_factory: WebRTCRuntimeFactory, - world_rank: int, - create_app_fn: CreateWebRTCApp, - server_runner: RunWebRTCServer, -) -> object: - runtime = runtime_factory(config=runtime_config) - manager = BaseWebRTCSessionManager( - runtime=runtime, - runtime_config=runtime_config, - fps=runtime_config.fps, - identity=runtime_config.pipeline_config_name, - busy_message="An OmniDreams session is already active.", - warmup_label="OmniDreams WebRTC", - supported_control_keys=WSAD_SUPPORTED_KEYS, - fatal_generation_errors=True, - client_liveness_timeout_s=output.client_liveness_timeout_s, - ) - from importlib.resources import files - - return serve_webrtc_demo( - output=output, - model_id=spec.model_id, - session_manager=manager, - app_resources=WebRTCAppResources( - model_web_resource=files("omnidreams.demo").joinpath("web"), - preload_name="OmniDreams", - ), - world_rank=world_rank, - create_app_fn=create_app_fn, - server_runner=server_runner, - ) - - def _serve_shared_omnidreams_webrtc_demo( *, spec: DemoSpec, @@ -1029,9 +302,7 @@ def _apply_runtime_options( __all__ = [ - "OmnidreamsWebRTCModelRuntime", "OmnidreamsWebRTCModelRuntimeConfig", - "OmnidreamsWebRTCModelRuntimeError", "SharedRuntimeFactory", "WebRTCRuntimeFactory", "serve_omnidreams_webrtc_demo", diff --git a/integrations/omnidreams/omnidreams/demo/webrtc_config.py b/integrations/omnidreams/omnidreams/demo/webrtc_config.py new file mode 100644 index 000000000..272123cbd --- /dev/null +++ b/integrations/omnidreams/omnidreams/demo/webrtc_config.py @@ -0,0 +1,86 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Shared OmniDreams WebRTC runtime configuration.""" + +from __future__ import annotations + +import math +from dataclasses import dataclass +from pathlib import Path +from typing import Any + +from omnidreams.scenes import SCENE_VARIANT_DEFAULT + +from flashdreams.serving.webrtc.encoders import EncoderBackend + +from .runtime import PipelineFactory +from .spec import DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID + + +@dataclass(frozen=True, slots=True) +class OmnidreamsWebRTCModelRuntimeConfig: + """Configuration for one scene-driven OmniDreams WebRTC runtime.""" + + pipeline_config_name: str + """User-facing name of the selected OmniDreams pipeline.""" + + pipeline_config: Any + """Resolved single-view OmniDreams pipeline configuration.""" + + scene_dir: Path | None = None + """Local scene root; ``None`` downloads the selected Hugging Face scene.""" + + scene_uuid: str | None = DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID + """Scene UUID used for remote lookup or local archive selection.""" + + scene_variant: str = SCENE_VARIANT_DEFAULT + """Weather variant selected from the scene assets.""" + + seed: int | None = 42 + """Per-rollout seed; ``None`` selects fresh entropy for every session.""" + + device: str = "cuda:0" + """Device used for rendering and model inference.""" + + video_height: int = 704 + """Generated video height in pixels.""" + + video_width: int = 1280 + """Generated video width in pixels.""" + + fps: int = 30 + """Input sampling and output playback frame rate.""" + + camera_name: str = "camera_front_wide_120fov" + """Scene camera controlled by browser keyboard input.""" + + move_speed_per_s: float = 6.0 + """Forward and reverse translation speed in scene units per second.""" + + rotate_speed_rad_per_s: float = math.radians(35.0) + """Left and right rotation speed in radians per second.""" + + warmup_chunks: int = 10 + """Number of synthetic chunks generated before accepting sessions.""" + + warmup_timeout_s: float = 600.0 + """Maximum duration for WebRTC loopback warmup.""" + + debug_serve_hdmaps: bool = False + """Stream rendered conditioning frames without running video generation.""" + + encoder_backend: EncoderBackend = "auto" + """WebRTC video encoder selection policy.""" + + encoder_bitrate_bps: int = 6_000_000 + """Target WebRTC video bitrate in bits per second.""" + + encoder_gop: int = 30 + """WebRTC video encoder group-of-pictures length.""" + + pipeline_factory: PipelineFactory | None = None + """Optional test/runtime override for constructing the shared pipeline.""" + + +__all__ = ["OmnidreamsWebRTCModelRuntimeConfig"] diff --git a/integrations/omnidreams/omnidreams/demo/webrtc_legacy.py b/integrations/omnidreams/omnidreams/demo/webrtc_legacy.py new file mode 100644 index 000000000..0d7768636 --- /dev/null +++ b/integrations/omnidreams/omnidreams/demo/webrtc_legacy.py @@ -0,0 +1,718 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Legacy OmniDreams WebRTC compatibility facade.""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping +from typing import Any + +import torch +from loguru import logger +from omnidreams.transformer import CosmosTransformerConfig + +from flashdreams.core.distributed.rank_orchestration import distributed_op +from flashdreams.runtime import ( + CanonicalInputs, + CanonicalInputSchema, + InferenceConfig, + InferenceInput, + InferenceInputSchema, + InputCanonicalizer, + StepRequest, + StepRequirements, + StepResult, + TimeWindow, + step_requirements_from_request, +) +from flashdreams.runtime.demo import ( + DemoSpec, + PreparedScenario, + SessionInfo, + UserInputWindow, + WebRTCAppResources, + WebRTCOutputSpec, +) +from flashdreams.runtime.demo.timing import SPARSE_KEY_SEGMENTS_METADATA_KEY +from flashdreams.runtime.demo.webrtc import ( + CreateWebRTCApp, + RunWebRTCServer, + serve_webrtc_demo, +) +from flashdreams.serving.webrtc.controls import WSAD_SUPPORTED_KEYS, PoseSegment +from flashdreams.serving.webrtc.manager import BaseWebRTCSessionManager +from flashdreams.serving.webrtc.runtime import ( + ThreadAffineDistributedWebRTCRuntime, + WebRTCControlSignal, +) +from flashdreams.serving.webrtc.services import WEBRTC_USER_INPUT_SCHEMA + +from .providers import LudusSceneConditioningProvider +from .runtime import OmnidreamsRuntime, OmnidreamsRuntimeOptions +from .spec import ( + DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, + OMNIDREAMS_MODEL_ID, + OmnidreamsLudusReplayScenario, +) +from .webrtc_config import OmnidreamsWebRTCModelRuntimeConfig + +WebRTCRuntimeFactory = Callable[..., Any] +_WEBRTC_SESSION_TOTAL_BLOCKS = 2_147_483_647 +_WEBRTC_STEP_REQUEST_KEY = "omnidreams_webrtc_step_request" + + +class OmnidreamsWebRTCModelRuntimeError(RuntimeError): + """Raised when the OmniDreams demo runtime is used incorrectly.""" + + +class OmnidreamsWebRTCModelRuntime( + ThreadAffineDistributedWebRTCRuntime[ + OmnidreamsWebRTCModelRuntimeConfig, + None, + ] +): + """Compatibility WebRTC facade over the shared OmniDreams runtime/session.""" + + def __init__(self, *, config: OmnidreamsWebRTCModelRuntimeConfig) -> None: + super().__init__( + config=config, + runtime_error_type=OmnidreamsWebRTCModelRuntimeError, + thread_name="omnidreams-demo-runtime", + ) + # The shared WebRTC input source emits normalized runtime events + # (``key_down``/``key_up``). The Ludus provider consumes sparse + # resampler metadata on this transitional path, so keep validation + # aligned with the WebRTC source rather than the replay trace schema. + self.input_source_schema = WEBRTC_USER_INPUT_SCHEMA + self.input_canonicalizer = InputCanonicalizer() + self.input_mapping = _OmnidreamsWebRTCInputMapping() + self._runtime: OmnidreamsRuntime | None = None + self._active_provider: LudusSceneConditioningProvider | None = None + self._active_session: Any | None = None + self._debug_session: _OmnidreamsHDMapDebugSession | None = None + self._steady_output_frame_count_value = 1 + + def _is_runtime_initialized(self) -> bool: + return self._runtime is not None + + def _runtime_step_index(self) -> int: + requirements = self._next_step_requirements_sync() + if requirements is None: + return 0 + return requirements.step_index + + def _next_input_frame_count(self) -> int: + requirements = self._next_step_requirements_sync() + if requirements is None: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC session is complete." + ) + return requirements.input_frame_count + + def _steady_output_frame_count(self) -> int: + return self._steady_output_frame_count_value + + def _initialize_sync(self) -> None: + if self._runtime is not None: + return + if self._device.type == "cuda" and not torch.cuda.is_available(): + raise RuntimeError("CUDA is required for OmniDreams WebRTC inference.") + _validate_single_view_pipeline_config( + pipeline_config_name=self.config.pipeline_config_name, + pipeline_config=self.config.pipeline_config, + ) + logger.info( + "Setting up shared OmniDreams runtime {} on {} for WebRTC.", + self.config.pipeline_config_name, + self._device, + ) + self._runtime = OmnidreamsRuntime( + config=self._inference_config(), + options=OmnidreamsRuntimeOptions( + pipeline_config=self.config.pipeline_config, + pipeline_factory=self.config.pipeline_factory, + # WebRTC warms the same long-lived runtime before real browser + # sessions. Keep prompt/image encoders available for later + # peer connections until Phase 14 replaces loopback warmup with + # first-class model/runtime warmup. + release_oneshot_encoders_after_cache_init=False, + ), + ) + self._initialize_video_encoder_sync() + + def _reset_rollout_sync(self, session_input: None = None) -> None: + del session_input + self._close_active_session_sync() + runtime = self._require_runtime() + scenario = self._session_scenario() + prepared = PreparedScenario( + initial_inputs=InferenceInput(global_conditioning={"scenario": scenario}), + source_schema=self.input_source_schema, + metadata={ + "conditioning_mode": "ludus-scene-driving", + "model_id": OMNIDREAMS_MODEL_ID, + "preset_id": self.config.pipeline_config_name, + }, + ) + provider = LudusSceneConditioningProvider( + scenario=prepared, + config=self._inference_config(), + ) + try: + initial_input = provider.prepare_initial_input() + session = runtime.start_session(initial_input) + except Exception: + provider.close() + raise + self._active_provider = provider + if self.config.debug_serve_hdmaps: + self._debug_session = _OmnidreamsHDMapDebugSession( + pipeline=runtime.pipeline, + scenario=scenario, + ) + self._active_session = self._debug_session + else: + self._debug_session = None + self._active_session = session + self._steady_output_frame_count_value = _steady_output_frame_count( + self._active_session, + fallback_pipeline=runtime.pipeline, + ) + + def _generate_one_chunk_sync( + self, + *, + segments: list[PoseSegment], + frame_times: list[float], + ) -> StepResult: + request = self._next_step_request_sync() + if request is None: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC session is complete." + ) + inputs = self.input_mapping.map_step_inputs( + canonical_inputs=CanonicalInputs(), + inference_input=InferenceInput( + metadata={ + SPARSE_KEY_SEGMENTS_METADATA_KEY: tuple(segments), + "frame_times": tuple(frame_times), + "window_start_s": request.step_index / float(self.config.fps), + "window_end_s": (request.step_index + len(frame_times)) + / float(self.config.fps), + } + ), + request=request, + ) + return self._step_active_session_sync(inputs) + + def _close_sync(self) -> None: + self._close_active_session_sync() + runtime = self._runtime + self._runtime = None + if runtime is not None: + runtime.close() + if self._device.type == "cuda" and torch.cuda.is_available(): + torch.cuda.synchronize(device=self._device) + torch.cuda.empty_cache() + + async def start_inference_session(self) -> "_OmnidreamsWebRTCInferenceSession": + self._require_open_and_initialized() + if not await self._worker.call(self._has_active_session_sync): + await self.reset_for_new_session() + return _OmnidreamsWebRTCInferenceSession(self) + + def _next_step_request_sync(self) -> StepRequest | None: + requirements = self._next_step_requirements_sync() + if requirements is None: + return None + metadata = dict(requirements.metadata) + metadata["input_frame_count"] = requirements.input_frame_count + if requirements.steady_output_frame_count is not None: + metadata["steady_output_frame_count"] = ( + requirements.steady_output_frame_count + ) + return StepRequest( + step_index=requirements.step_index, + inference_input_schema=requirements.inference_input_schema, + metadata=metadata, + ) + + def _next_step_requirements_sync(self) -> StepRequirements | None: + session = self._require_active_session() + next_requirements = getattr(session, "next_step_requirements", None) + if callable(next_requirements): + result = next_requirements() + else: + next_request = session.next_step_request() + if next_request is None: + return None + result = step_requirements_from_request(next_request) + if result is None: + return None + if not isinstance(result, StepRequirements): + raise TypeError( + "OmniDreams WebRTC session requirements must be StepRequirements, " + f"got {type(result).__name__}." + ) + return result + + def _session_info_sync(self) -> SessionInfo: + return SessionInfo( + output_layout="bvtchw", + steady_output_frame_count=self._steady_output_frame_count(), + metadata={"model_id": OMNIDREAMS_MODEL_ID}, + ) + + def _step_active_session_sync(self, inputs: InferenceInput) -> StepResult: + provider = self._require_active_provider() + session = self._require_active_session() + request = _request_from_step_inputs(inputs) + requirements = step_requirements_from_request( + request, + allow_user_input_window=True, + ) + window = _user_window_from_step_inputs( + inputs, + request=request, + input_frame_count=requirements.input_frame_count, + ) + prepared = provider.prepare_step(request=requirements, user_window=window) + if prepared.control.close_session: + raise OmnidreamsWebRTCModelRuntimeError( + prepared.control.reason or "OmniDreams WebRTC input is exhausted." + ) + if prepared.control.reset: + reset_input = prepared.control.reset_input + session.reset(reset_input) + if not prepared.control.provider_already_reset: + provider.reset(reset_input) + raise OmnidreamsWebRTCModelRuntimeError( + prepared.control.reason or "OmniDreams WebRTC session reset requested." + ) + if prepared.inference_input is None: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC provider returned no inference input." + ) + result = session.step(prepared.inference_input) + if not isinstance(result, StepResult): + raise TypeError( + "OmniDreams WebRTC session steps must produce StepResult, got " + f"{type(result).__name__}." + ) + return result + + @distributed_op(WebRTCControlSignal.SESSION_STEP) + def _step_active_session_sync_all_ranks( + self, + inputs: InferenceInput, + ) -> StepResult: + return self._step_active_session_sync(inputs) + + @distributed_op(WebRTCControlSignal.SESSION_CLOSE) + def _close_active_session_sync_all_ranks(self) -> None: + self._close_active_session_sync() + + def _close_active_session_sync(self) -> None: + session = self._active_session + provider = self._active_provider + self._active_session = None + self._debug_session = None + self._active_provider = None + first_error: Exception | None = None + close_session = getattr(session, "close", None) + if callable(close_session): + try: + close_session() + except Exception as exc: + first_error = exc + if provider is not None: + try: + provider.close() + except Exception as exc: + if first_error is None: + first_error = exc + if first_error is not None: + raise first_error + + def _has_active_session_sync(self) -> bool: + return self._active_session is not None and self._active_provider is not None + + def _require_runtime(self) -> OmnidreamsRuntime: + if self._runtime is None: + raise OmnidreamsWebRTCModelRuntimeError("Runtime is not initialized.") + return self._runtime + + def _require_active_session(self) -> Any: + if self._active_session is None: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC session is not initialized." + ) + return self._active_session + + def _require_active_provider(self) -> LudusSceneConditioningProvider: + if self._active_provider is None: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC provider is not initialized." + ) + return self._active_provider + + def _inference_config(self) -> InferenceConfig: + return InferenceConfig( + model_id=OMNIDREAMS_MODEL_ID, + preset_id=self.config.pipeline_config_name, + device=str(self.config.device), + seed=self.config.seed, + runtime_options={"seed": self.config.seed}, + ) + + def _session_scenario(self) -> OmnidreamsLudusReplayScenario: + return OmnidreamsLudusReplayScenario( + keyboard_events=(), + scene_dir=self.config.scene_dir, + scene_uuid=self.config.scene_uuid or DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, + scene_variant=self.config.scene_variant, + camera_name=self.config.camera_name, + total_blocks=_WEBRTC_SESSION_TOTAL_BLOCKS, + pixel_height=self.config.video_height, + pixel_width=self.config.video_width, + fps=self.config.fps, + move_speed_per_s=self.config.move_speed_per_s, + rotate_speed_rad_per_s=self.config.rotate_speed_rad_per_s, + ) + + +class _OmnidreamsWebRTCInputMapping: + """Carry shared WebRTC window facts into the OmniDreams session facade.""" + + def validate( + self, + *, + canonical_schema: CanonicalInputSchema | None = None, + inference_input_schema: InferenceInputSchema | None = None, + ) -> None: + del canonical_schema, inference_input_schema + + def map_global_conditioning_inputs( + self, + *, + canonical_inputs: CanonicalInputs, + inference_input: InferenceInput, + ) -> InferenceInput: + del canonical_inputs + return inference_input + + def map_step_inputs( + self, + *, + canonical_inputs: CanonicalInputs, + inference_input: InferenceInput, + request: StepRequest, + ) -> InferenceInput: + del canonical_inputs + step = dict(inference_input.step) + step[_WEBRTC_STEP_REQUEST_KEY] = request + return InferenceInput( + global_conditioning=inference_input.global_conditioning, + step=step, + metadata=inference_input.metadata, + ) + + +class _OmnidreamsWebRTCInferenceSession: + """Synchronous session proxy consumed by the shared WebRTC compatibility path.""" + + def __init__(self, runtime: OmnidreamsWebRTCModelRuntime) -> None: + self._runtime = runtime + self._closed = False + + def session_info(self) -> SessionInfo: + self._require_open() + return self._runtime._worker.call_blocking(self._runtime._session_info_sync) + + def next_step_requirements(self) -> StepRequirements | None: + self._require_open() + return self._runtime._worker.call_blocking( + self._runtime._next_step_requirements_sync + ) + + def next_step_request(self) -> StepRequest | None: + self._require_open() + return self._runtime._worker.call_blocking( + self._runtime._next_step_request_sync + ) + + def step(self, inputs: InferenceInput) -> StepResult: + self._require_open() + return self._runtime._worker.call_blocking( + self._runtime._step_active_session_sync_all_ranks, + inputs, + ) + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self._require_open() + self._runtime._worker.call_blocking(self._runtime._reset_rollout_sync_all_ranks) + + def close(self) -> None: + if self._closed: + return + self._closed = True + self._runtime._worker.call_blocking( + self._runtime._close_active_session_sync_all_ranks + ) + + def _require_open(self) -> None: + if self._closed: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC inference session is closed." + ) + + +class _OmnidreamsHDMapDebugSession: + """Session-shaped debug path that streams rendered Ludus HDMaps.""" + + def __init__( + self, *, pipeline: Any, scenario: OmnidreamsLudusReplayScenario + ) -> None: + self._pipeline = pipeline + self._scenario = scenario + self._step_index = 0 + self._closed = False + + def session_info(self) -> SessionInfo: + return SessionInfo( + output_layout="bvtchw", + steady_output_frame_count=self._steady_output_frame_count(), + metadata={"stream": "hdmap"}, + ) + + def next_step_requirements(self) -> StepRequirements | None: + if self._closed or self._step_index >= self._scenario.total_blocks: + return None + return StepRequirements( + step_index=self._step_index, + input_frame_count=self._num_frames(self._step_index), + steady_output_frame_count=self._steady_output_frame_count(), + ) + + def next_step_request(self) -> StepRequest | None: + requirements = self.next_step_requirements() + if requirements is None: + return None + return StepRequest( + step_index=requirements.step_index, + metadata={ + "input_frame_count": requirements.input_frame_count, + "steady_output_frame_count": requirements.steady_output_frame_count, + }, + ) + + def step(self, inputs: InferenceInput) -> StepResult: + requirements = self.next_step_requirements() + if requirements is None: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC debug session is complete." + ) + hdmap = inputs.step.get("hdmap") + if not isinstance(hdmap, torch.Tensor): + raise TypeError("OmniDreams WebRTC debug session requires step['hdmap'].") + result = StepResult.from_video_chunk( + step_index=requirements.step_index, + video_chunk=hdmap.detach(), + layout="bvtchw", + metadata={"stream": "hdmap"}, + ) + self._step_index += 1 + return result + + def reset(self, inputs: InferenceInput | None = None) -> None: + del inputs + self._step_index = 0 + self._closed = False + + def close(self) -> None: + self._closed = True + + def _steady_output_frame_count(self) -> int: + return self._num_frames(1) + + def _num_frames(self, step_index: int) -> int: + get_num_frames = getattr(self._pipeline, "get_num_frames", None) + if not callable(get_num_frames): + return 1 + return int(get_num_frames(step_index)) + + +def _request_from_step_inputs(inputs: InferenceInput) -> StepRequest: + request = inputs.step.get(_WEBRTC_STEP_REQUEST_KEY) + if not isinstance(request, StepRequest): + raise TypeError( + "OmniDreams WebRTC step input is missing the shared StepRequest." + ) + return request + + +def _user_window_from_step_inputs( + inputs: InferenceInput, + *, + request: StepRequest, + input_frame_count: int, +) -> UserInputWindow: + frame_times = _frame_times_from_metadata(inputs.metadata, input_frame_count) + segments = _segments_from_metadata(inputs.metadata) + window = request.user_input_window or TimeWindow( + start_s=float(inputs.metadata.get("window_start_s", 0.0)), + end_s=float(inputs.metadata.get("window_end_s", frame_times[-1])), + ) + return UserInputWindow( + start_s=window.start_s, + end_s=window.end_s, + frame_times=frame_times, + metadata={SPARSE_KEY_SEGMENTS_METADATA_KEY: segments}, + ) + + +def _frame_times_from_metadata( + metadata: Mapping[str, object], + input_frame_count: int, +) -> tuple[float, ...]: + value = metadata.get("frame_times") + if not isinstance(value, tuple): + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC step input is missing frame_times metadata." + ) + frame_times = tuple( + _float_metadata_value(frame_time, label="frame_times") for frame_time in value + ) + if len(frame_times) != input_frame_count: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC frame_times length does not match " + f"input_frame_count={input_frame_count}." + ) + return frame_times + + +def _segments_from_metadata(metadata: Mapping[str, object]) -> tuple[PoseSegment, ...]: + value = metadata.get(SPARSE_KEY_SEGMENTS_METADATA_KEY) + if not isinstance(value, tuple): + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC step input is missing resampled key segments." + ) + segments: list[PoseSegment] = [] + for segment in value: + if not isinstance(segment, tuple) or len(segment) != 3: + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC key segments must be 3-tuples." + ) + start_s, end_s, keys = segment + if not isinstance(keys, frozenset | set | tuple | list): + raise OmnidreamsWebRTCModelRuntimeError( + "OmniDreams WebRTC key segment keys must be a sequence." + ) + segments.append( + ( + _float_metadata_value(start_s, label="segment start"), + _float_metadata_value(end_s, label="segment end"), + frozenset(str(key) for key in keys), + ) + ) + return tuple(segments) + + +def _float_metadata_value(value: object, *, label: str) -> float: + if isinstance(value, bool) or not isinstance(value, int | float): + raise OmnidreamsWebRTCModelRuntimeError( + f"OmniDreams WebRTC {label} metadata must be numeric." + ) + return float(value) + + +def _steady_output_frame_count(session: Any, *, fallback_pipeline: Any) -> int: + session_info = getattr(session, "session_info", None) + if callable(session_info): + value = session_info() + if isinstance(value, SessionInfo) and value.steady_output_frame_count: + return int(value.steady_output_frame_count) + get_num_frames = getattr(fallback_pipeline, "get_num_frames", None) + if callable(get_num_frames): + return int(get_num_frames(1)) + return 1 + + +def _validate_single_view_pipeline_config( + *, + pipeline_config_name: str, + pipeline_config: Any, +) -> None: + diffusion_model = getattr(pipeline_config, "diffusion_model", None) + transformer_cfg = getattr(diffusion_model, "transformer", None) + if transformer_cfg is None: + return + if not isinstance(transformer_cfg, CosmosTransformerConfig): + raise TypeError( + "OmniDreams WebRTC requires a CosmosTransformerConfig pipeline." + ) + if transformer_cfg.num_views != 1: + raise ValueError( + "OmniDreams WebRTC supports only single-view configs; " + f"{pipeline_config_name!r} has num_views={transformer_cfg.num_views}." + ) + + +def _serve_legacy_omnidreams_webrtc_demo( + *, + spec: DemoSpec, + output: WebRTCOutputSpec, + runtime_config: OmnidreamsWebRTCModelRuntimeConfig, + runtime_factory: WebRTCRuntimeFactory, + world_rank: int, + create_app_fn: CreateWebRTCApp, + server_runner: RunWebRTCServer, +) -> object: + runtime = runtime_factory(config=runtime_config) + manager = BaseWebRTCSessionManager( + runtime=runtime, + runtime_config=runtime_config, + fps=runtime_config.fps, + identity=runtime_config.pipeline_config_name, + busy_message="An OmniDreams session is already active.", + warmup_label="OmniDreams WebRTC", + supported_control_keys=WSAD_SUPPORTED_KEYS, + fatal_generation_errors=True, + client_liveness_timeout_s=output.client_liveness_timeout_s, + ) + from importlib.resources import files + + return serve_webrtc_demo( + output=output, + model_id=spec.model_id, + session_manager=manager, + app_resources=WebRTCAppResources( + model_web_resource=files("omnidreams.demo").joinpath("web"), + preload_name="OmniDreams", + ), + world_rank=world_rank, + create_app_fn=create_app_fn, + server_runner=server_runner, + ) + + +__all__ = [ + "OmnidreamsWebRTCModelRuntime", + "OmnidreamsWebRTCModelRuntimeError", + "WebRTCRuntimeFactory", + "_serve_legacy_omnidreams_webrtc_demo", +] diff --git a/integrations/omnidreams/tests/test_demo_api.py b/integrations/omnidreams/tests/test_demo_api.py index 59c6f7320..3d083535a 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -5,6 +5,7 @@ import asyncio import json +import sys from collections.abc import Sequence from pathlib import Path from types import SimpleNamespace @@ -42,8 +43,8 @@ OmnidreamsSession, ) from omnidreams.demo.webrtc import ( - OmnidreamsWebRTCModelRuntime, OmnidreamsWebRTCModelRuntimeConfig, + _should_use_legacy_webrtc_path, serve_omnidreams_webrtc_demo, ) @@ -715,7 +716,7 @@ def test_omnidreams_replay_runtime_generates_video_step_result( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - import omnidreams.demo.replay as replay_module + import omnidreams.demo.runtime as runtime_module hdmap = tmp_path / "hdmap.mp4" first_frame = tmp_path / "first.png" @@ -723,7 +724,7 @@ def test_omnidreams_replay_runtime_generates_video_step_result( first_frame.write_bytes(b"fake") pipeline = _FakeOmnidreamsPipeline() monkeypatch.setattr( - replay_module, + runtime_module, "load_first_frame_tensor", lambda *args, **kwargs: torch.zeros(1, 3, 2, 2), ) @@ -842,6 +843,8 @@ def test_omnidreams_webrtc_cli_builds_keyboard_driving_spec(tmp_path: Path) -> N def test_omnidreams_webrtc_demo_uses_shared_manager_with_model_config() -> None: + legacy_module_name = "omnidreams.demo.webrtc_legacy" + sys.modules.pop(legacy_module_name, None) pipeline_config = object() runtime = _FactoryRuntime() spec = DemoSpec( @@ -923,6 +926,7 @@ def shared_runtime_factory(**kwargs: Any) -> Any: assert scenario.fps == 24 assert calls[0]["host"] == "0.0.0.0" assert calls[0]["port"] == 8082 + assert legacy_module_name not in sys.modules def test_omnidreams_webrtc_demo_keeps_legacy_runtime_factory_path() -> None: @@ -930,7 +934,7 @@ def test_omnidreams_webrtc_demo_keeps_legacy_runtime_factory_path() -> None: model_id=OMNIDREAMS_MODEL_ID, preset_id=DEFAULT_OMNIDREAMS_PRESET, input_mode="keyboard-driving", - scenario=OmnidreamsWebRTCScenario(debug_serve_hdmaps=True), + scenario=OmnidreamsWebRTCScenario(), output=WebRTCOutputSpec( host="0.0.0.0", port=8082, @@ -960,7 +964,32 @@ def test_omnidreams_webrtc_demo_keeps_legacy_runtime_factory_path() -> None: runtime = manager._runtime assert isinstance(runtime, _FakeWebRTCRuntime) assert manager.runtime_config is runtime.config - assert runtime.config.debug_serve_hdmaps is True + assert runtime.config.debug_serve_hdmaps is False + + +def test_omnidreams_webrtc_demo_keeps_legacy_fallback_gates( + monkeypatch: pytest.MonkeyPatch, +) -> None: + assert _should_use_legacy_webrtc_path( + scenario=OmnidreamsWebRTCScenario(), + runtime_factory=_FakeWebRTCRuntime, + ) + assert _should_use_legacy_webrtc_path( + scenario=OmnidreamsWebRTCScenario(debug_serve_hdmaps=True), + runtime_factory=None, + ) + + monkeypatch.setenv("WORLD_SIZE", "2") + assert _should_use_legacy_webrtc_path( + scenario=OmnidreamsWebRTCScenario(), + runtime_factory=None, + ) + + monkeypatch.setenv("WORLD_SIZE", "not-an-int") + assert not _should_use_legacy_webrtc_path( + scenario=OmnidreamsWebRTCScenario(), + runtime_factory=None, + ) def test_omnidreams_webrtc_demo_installs_model_assets_without_routes( @@ -1096,6 +1125,8 @@ async def test_omnidreams_webrtc_runtime_uses_shared_session( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: + from omnidreams.demo.webrtc_legacy import OmnidreamsWebRTCModelRuntime + _scene, rasterizers = _install_fake_ludus_provider_dependencies(monkeypatch) scene_path = tmp_path / "scene.usdz" scene_path.write_bytes(b"fake") @@ -1181,6 +1212,8 @@ async def test_omnidreams_webrtc_runtime_keeps_encoders_after_warmup_session( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: + from omnidreams.demo.webrtc_legacy import OmnidreamsWebRTCModelRuntime + _install_fake_ludus_provider_dependencies(monkeypatch) scene_path = tmp_path / "scene.usdz" scene_path.write_bytes(b"fake") From 1c12ecf619694f83ec8bed7e0c1d83305f3d0fda Mon Sep 17 00:00:00 2001 From: Jesse Archer Date: Mon, 10 Aug 2026 06:31:30 +0000 Subject: [PATCH 51/51] docs: record validated OmniDreams demo commands Update the OmniDreams demo README with the remote GPU setup and validated null, precomputed MP4, Ludus MP4, and WebRTC commands used during migration testing. Remove the untested explicit benchmark-asset example from the main command list. --- .../omnidreams/omnidreams/demo/README.md | 65 ++++++++++++++----- 1 file changed, 48 insertions(+), 17 deletions(-) diff --git a/integrations/omnidreams/omnidreams/demo/README.md b/integrations/omnidreams/omnidreams/demo/README.md index ce23b48cd..9dab46ed8 100644 --- a/integrations/omnidreams/omnidreams/demo/README.md +++ b/integrations/omnidreams/omnidreams/demo/README.md @@ -8,21 +8,46 @@ SPDX-License-Identifier: Apache-2.0 This folder contains the experimental OmniDreams demo built on `flashdreams.runtime.demo`. -Run commands from the FlashDreams workspace root: +Run commands from the FlashDreams workspace root. The following setup was used +for remote GPU validation on GB300: ```bash cd /path/to/flashdreams export HF_TOKEN= +export CUDA_HOME=/usr/local/cuda-13.1 +export CUDA_PATH="$CUDA_HOME" +export PATH="$CUDA_HOME/bin:$PATH" +export LD_LIBRARY_PATH="$CUDA_HOME/lib64:${LD_LIBRARY_PATH:-}" +hash -r +"$CUDA_HOME/bin/nvcc" --version + +uv sync --python 3.12 --package flashdreams-omnidreams --no-dev ``` -## MP4 Replay +## Null Replay -Generate an MP4 from the bundled single-view sample data: +Run a short replay without writing video output: + +```bash +uv run --python 3.12 --package flashdreams-omnidreams omnidreams-demo replay \ + --output-mode null \ + --device cuda:0 \ + --total-blocks 10 +``` + +## Precomputed MP4 Replay + +Generate an MP4 from bundled single-view sample data and pre-rendered HDMaps: ```bash mkdir -p outputs -uv run --package flashdreams-omnidreams omnidreams-demo replay \ - --output outputs/omnidreams-demo.mp4 +uv run --python 3.12 --package flashdreams-omnidreams omnidreams-demo replay \ + --device cuda:0 \ + --example-data \ + --example-data-uuid 239560dc-33d1-11ef-9720-00044bcbccac \ + --total-blocks 225 \ + --fps 30 \ + --output outputs/omnidreams-demo-precomputed-1min.mp4 ``` This replay path mirrors the benchmark runner path: it uses a prompt, first @@ -30,20 +55,24 @@ frame, and pre-rendered HDMap video. It does not load a Ludus scene or render HDMaps at runtime. The demo defaults to the stable non-perf OmniDreams preset used by the benchmark path. -To provide benchmark-style assets explicitly: +Pass `--example-data-uuid ` to select another bundled single-view sample, +or `--no-example-data` to require explicit asset paths. + +## Ludus MP4 Replay + +Generate an MP4 by rendering HDMap conditioning from a recorded keyboard trace: ```bash -uv run --package flashdreams-omnidreams omnidreams-demo replay \ - --prompt "Driving scene from a front-facing car camera." \ - --hdmap-video-paths /path/to/camera_front_wide_120fov_hdmap.mp4 \ - --first-frame-paths /path/to/first_frame.png \ - --camera-names camera_front_wide_120fov \ - --output outputs/omnidreams-demo.mp4 +uv run --python 3.12 --package flashdreams-omnidreams omnidreams-demo replay \ + --conditioning-mode ludus-scene-driving \ + --keyboard-trace integrations/omnidreams/omnidreams/demo/traces/ludus_forward_sweep_60s.json \ + --device cuda:0 \ + --scene-uuid 0d404ff7-2b66-498c-b047-1ed8cded60d4 \ + --seed 42 \ + --total-blocks 226 \ + --output outputs/omnidreams-demo--ludus-1min.mp4 ``` -Pass `--example-data-uuid ` to select another bundled single-view sample, -or `--no-example-data` to require explicit asset paths. - The `omnidreams-sv-2steps-chunk2-loc6-lightvae-lighttae-perf` preset remains an explicit `--preset-id` opt-in. It should become the default only after the compile/cache behavior is reliable enough for the demo path. @@ -55,9 +84,11 @@ The small model adapter in this package loads one scene, renders HDMap conditioning with Ludus, and runs OmniDreams from browser WASD controls: ```bash -uv run --package flashdreams-omnidreams omnidreams-demo webrtc \ +uv run --python 3.12 --package flashdreams-omnidreams omnidreams-demo webrtc \ --host 0.0.0.0 \ - --port 8082 + --port 8089 \ + --device cuda:0 \ + --scene-uuid 0d404ff7-2b66-498c-b047-1ed8cded60d4 ``` The scene UUID is optional; when omitted, the runtime uses the default