diff --git a/.github/workflows/omnidreams-demo-runtime.yml b/.github/workflows/omnidreams-demo-runtime.yml new file mode 100644 index 000000000..1f494fa83 --- /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: "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: "18" + MAX_DURATION_SECONDS: "22" + 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-20s.mp4 |" + echo "| Ludus MP4 | ${LUDUS_BLOCKS} | omnidreams-demo-ludus-20s.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-20s.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-20s.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 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 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-20s.mp4" + validate_mp4 ludus-mp4 "${output_dir}/omnidreams-demo-ludus-20s.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 diff --git a/flashdreams/flashdreams/runtime/__init__.py b/flashdreams/flashdreams/runtime/__init__.py index fb6eb4b05..6e205823c 100644 --- a/flashdreams/flashdreams/runtime/__init__.py +++ b/flashdreams/flashdreams/runtime/__init__.py @@ -52,14 +52,20 @@ from flashdreams.runtime.metrics import ( InMemoryMetricsRecorder, MetricsRecorder, + MetricsSnapshot, NullMetricsRecorder, RuntimeMetricSample, ) 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 ThreadAffineRuntimeWorker +from flashdreams.runtime.worker import ModelExecutionWorker, ThreadAffineRuntimeWorker __all__ = [ "CanonicalInputs", @@ -90,7 +96,9 @@ "KeyboardToDriverCommand", "MappingCompatibility", "MetricsRecorder", + "MetricsSnapshot", "ModelAdapter", + "ModelExecutionWorker", "Mp4VideoOutputTarget", "NullMetricsRecorder", "NullOutputTarget", @@ -100,10 +108,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/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 7a0535556..60187237c 100644 --- a/flashdreams/flashdreams/runtime/demo/__init__.py +++ b/flashdreams/flashdreams/runtime/demo/__init__.py @@ -3,11 +3,70 @@ """Experimental shared demo API above the inference runtime API.""" -from flashdreams.runtime.demo.outputs import build_output_target -from flashdreams.runtime.demo.replay import run_replay_demo +from flashdreams.runtime.demo.drivers import ( + CLEANUP_TIMEOUT_S, + BatchSessionDriver, + DriverInvariantError, + RealtimeSessionDriver, + run_demo_session, + run_demo_session_async, + shielded_session_cleanup, + uncancel_current_task, +) +from flashdreams.runtime.demo.host import ( + ModelWarmupPlan, + RuntimeHost, + 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 +from flashdreams.runtime.demo.replay import OutputSinkFactory, run_replay_demo +from flashdreams.runtime.demo.run_modes import ( + AsyncSessionDriver, + BenchmarkErrorPolicy, + DefaultErrorPolicy, + ErrorAction, + InMemorySessionMetricsRecorder, + MetricsSnapshot, + Mp4ErrorPolicy, + NativeWindowErrorPolicy, + NoopTransportService, + NullErrorPolicy, + RunContext, + RunMode, + RunModeCapabilities, + RunModeWarmup, + RunResult, + RunSummary, + SessionDriver, + SessionEdges, + SingleSessionAdmissionPolicy, + WebRTCErrorPolicy, + build_model_warmup_plan, + warmup_run_context, +) +from flashdreams.runtime.demo.session_inputs import ( + BatchInputSource, + ControlDecision, + InputSource, + ModelInputProvider, + PreparedStep, + ProviderCapabilities, + RealtimeInputSource, + UserInputWindow, +) from flashdreams.runtime.demo.spec import ( DemoAdapter, DemoSpec, + ModelWarmupAdapter, Mp4OutputSpec, NullOutputSpec, OutputSpec, @@ -15,16 +74,105 @@ 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, +) +from flashdreams.runtime.demo.validation import ( + ResolvedRunCapabilities, + resolve_run_capabilities, + validate_resolved_run, +) __all__ = [ + "BatchInputSource", + "BatchSessionDriver", + "CLEANUP_TIMEOUT_S", + "ControlDecision", + "DefaultErrorPolicy", "DemoAdapter", "DemoSpec", + "DriverInvariantError", + "ErrorAction", + "AsyncSessionDriver", + "ActivationPolicy", + "ActivationResult", + "ActivationSignal", + "AlwaysActiveActivationPolicy", + "BenchmarkErrorPolicy", + "InMemorySessionMetricsRecorder", + "InputSource", + "CatchUpDecision", + "CatchUpPolicy", + "DeterministicClock", + "MetricsSnapshot", + "ModelWarmupAdapter", + "ModelWarmupPlan", + "ModelInputProvider", + "KeyboardRealtimeInputSource", + "Mp4ErrorPolicy", + "Mp4OutputSink", "Mp4OutputSpec", + "NativeWindowErrorPolicy", + "NoopTransportService", "NullOutputSpec", + "NullOutputSink", + "NullErrorPolicy", + "OutputDecision", + "OutputSinkFactory", "OutputSpec", + "OutputSink", "PreparedScenario", + "PreparedStep", + "ProviderCapabilities", + "RealtimeInputSource", + "RealtimeClock", + "RealtimeSessionDriver", + "RealtimeWindowResult", + "ResolvedRunCapabilities", + "ResamplerRealtimeClock", + "RunContext", + "RunMode", + "RunModeCapabilities", + "RunModeWarmup", + "RunResult", + "RunSummary", + "RuntimeHost", + "SessionEdges", + "SessionDriver", + "SessionInfo", + "SignalActivationPolicy", + "SingleSessionAdmissionPolicy", + "SPARSE_KEY_SEGMENTS_METADATA_KEY", + "StepOutcome", + "StepPipeline", + "UserInputWindow", + "WarmupSessionInputs", "WebRTCAppResources", + "WebRTCErrorPolicy", "WebRTCOutputSpec", + "build_output_sink", "build_output_target", + "build_model_warmup_plan", + "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", + "warmup_run_context", ] 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/drivers.py b/flashdreams/flashdreams/runtime/demo/drivers.py new file mode 100644 index 000000000..9bdd40db8 --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/drivers.py @@ -0,0 +1,1011 @@ +# 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 + +import asyncio +import inspect +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 +from .pipeline import StepPipeline +from .run_modes import ( + DriverStatus, + RunContext, + RunMode, + RunResult, + SessionEdges, + SessionReservation, +) +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 +_MODEL_CLEANUP_FAILED_REASON = "model-affine cleanup failed" +_MODEL_CLEANUP_TIMED_OUT_REASON = "model-affine cleanup timed out" + + +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 + invariant_closed = False + 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 + + input_source = cast(BatchInputSource, session_edges.input_source) + 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 = _next_step_requirements(host=host, session=session) + if request is None: + break + user_window = 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 as exc: + if session is not None: + _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), + error=exc, + ) + invariant_closed = True + raise + except Exception as exc: + final_status = "failed" + final_reason = str(exc) + final_error = exc + finally: + if not invariant_closed: + if session is not None: + _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, + reason=final_reason, + error=final_error, + ) + + +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 _next_step_requirements_async( + host=host, + session=session, + ) + if request is None: + break + window_result = await input_source.next_realtime_window( + 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 + ): + 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 DriverInvariantError: + raise + 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 _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, + 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.""" + 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: + 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, + ) + 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 " + "must not be reused." + ) + driver = run_mode.select_driver() + driver_started = True + result = _run_sync_driver( + driver=driver, + host=context.host, + provider=provider, + session_edges=session_edges, + pipeline=pipeline, + ) + 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: + _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 + ): + result = session_edges.close_result( + status="failed", + reason=str(exc), + error=exc, + ) + 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: + _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 + ): + 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() + + +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 + 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, + ) + 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 " + "must not be reused." + ) + driver = run_mode.select_driver() + 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=_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 + result = await _close_partial_session_async( + context=context, + provider=provider, + session_edges=session_edges, + status="failed", + reason=str(exc), + error=exc, + close_provider=_needs_partial_provider_cleanup(session_edges), + ) + if should_record_session: + 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, + session_edges=session_edges, + status="failed", + reason=str(exc), + error=exc, + close_provider=_needs_partial_provider_cleanup(session_edges), + ) + 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 _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) -> bool: + try: + close() + except Exception as exc: + session_edges.record_cleanup_error(exc) + return False + return True + + +def _close_on_host_best_effort( + *, + host: RuntimeHost, + close: Any, + session_edges: SessionEdges, +) -> bool: + try: + 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( + *, + 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: + unhealthy_reason = await _close_model_resources_async( + host=host, + session=session, + provider=provider, + session_edges=session_edges, + timeout_s=timeout_s, + ) + if unhealthy_reason is not None: + host.mark_unhealthy(unhealthy_reason) + 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, + 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 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, + 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) + + +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, + provider: Any, + session_edges: SessionEdges | None, +) -> None: + try: + 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: + _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) + # 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) + + +def _record_run_cleanup_error(context: RunContext, exc: Exception) -> None: + try: + context.run_metrics.record_cleanup_error(exc) + except Exception: + 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_resources_async( + *, + host: RuntimeHost, + session: InferenceSession | None, + provider: ModelInputProvider, + session_edges: SessionEdges, + timeout_s: float, +) -> str | None: + try: + resources_closed = await asyncio.wait_for( + 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: + # 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 _MODEL_CLEANUP_TIMED_OUT_REASON + except Exception as exc: + session_edges.record_cleanup_error(exc) + 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, +) -> 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 + + +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 + 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", + "CLEANUP_TIMEOUT_S", + "DriverInvariantError", + "RealtimeSessionDriver", + "run_demo_session", + "run_demo_session_async", + "shielded_session_cleanup", + "uncancel_current_task", +] diff --git a/flashdreams/flashdreams/runtime/demo/host.py b/flashdreams/flashdreams/runtime/demo/host.py new file mode 100644 index 000000000..6a6e93ce5 --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/host.py @@ -0,0 +1,174 @@ +# SPDX-FileCopyrightText: Copyright (c) 2026 NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# SPDX-License-Identifier: Apache-2.0 + +"""Runtime host and model-execution boundary for shared demos.""" + +from __future__ import annotations + +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") + + +@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) + + def __post_init__(self) -> None: + object.__setattr__(self, "sessions", tuple(self.sessions)) + object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) + + +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 and not self._closed + + @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 on the worker.""" + self._require_open() + return self._worker.call_blocking(func, *args, **kwargs) + + async def call_async( + self, + func: Callable[..., _T], + /, + *args: object, + **kwargs: object, + ) -> _T: + """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 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() + + 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/demo/outputs.py b/flashdreams/flashdreams/runtime/demo/outputs.py index 421ec3bb4..866acbb00 100644 --- a/flashdreams/flashdreams/runtime/demo/outputs.py +++ b/flashdreams/flashdreams/runtime/demo/outputs.py @@ -1,18 +1,270 @@ # 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 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.output import NullOutputTarget, OutputTarget +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 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[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) + 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_record(result)) + return OutputDecision() + + def close(self) -> Sequence[OutputArtifact]: + self.closed = True + 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, *, @@ -42,4 +294,12 @@ def build_output_target( raise TypeError(f"Unsupported demo output spec: {type(output).__name__}.") -__all__ = ["build_output_target"] +__all__ = [ + "Mp4OutputSink", + "NullOutputSink", + "OutputDecision", + "OutputSink", + "SessionInfo", + "build_output_sink", + "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..5a9da08b0 --- /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 StepRequirements, 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: StepRequirements, + 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/replay.py b/flashdreams/flashdreams/runtime/demo/replay.py index 18b873254..fa2f507a9 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 OutputDecision, 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,87 @@ def run_replay_demo( if spec.config is None: raise RuntimeError("DemoSpec.config was not initialized.") + if runner 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(), + ) + + 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, + prepared=prepared, + mapping=mapping, + output_sink_factory=output_sink_factory, + metrics=metrics, + ) + + +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, + 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 +189,422 @@ 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", + ) -> 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, + 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 +623,7 @@ def _require_supported_mode( __all__ = [ "InferenceSessionRunner", + "OutputSinkFactory", "OutputTargetFactory", "run_replay_demo", ] diff --git a/flashdreams/flashdreams/runtime/demo/run_modes.py b/flashdreams/flashdreams/runtime/demo/run_modes.py new file mode 100644 index 000000000..cb93e4e7f --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/run_modes.py @@ -0,0 +1,535 @@ +# 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 + +import asyncio +from collections.abc import Mapping, Sequence +from dataclasses import dataclass, field +from threading import Lock +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 OutputSink + +if TYPE_CHECKING: + from .host import RuntimeHost + 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", + "failed", + "skipped", + "cancelled", + "rejected", + "not_activated", +] + +DriverStatus = Literal[ + "completed", + "failed", + "skipped", + "cancelled", + "not_activated", +] + + +@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: ... + + +class Mp4ErrorPolicy(DefaultErrorPolicy): + """Abort MP4 sessions on setup or step errors.""" + + +class NullErrorPolicy(DefaultErrorPolicy): + """Abort headless/null sessions on setup or step errors.""" + + +class NativeWindowErrorPolicy(DefaultErrorPolicy): + """Abort native-window sessions unless a future UI policy overrides it.""" + + +class BenchmarkErrorPolicy(DefaultErrorPolicy): + """Close failed scenarios while letting benchmark loops continue.""" + + def handle_setup_error(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status="failed", continue_next_scenario=True) + + def handle(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status="failed", continue_next_scenario=True) + + +@dataclass(frozen=True, kw_only=True, slots=True) +class WebRTCErrorPolicy: + """Drop configured recoverable realtime errors, otherwise close the session.""" + + recoverable_exception_types: tuple[type[Exception], ...] = () + + def handle_setup_error(self, exc: Exception) -> ErrorAction: + del exc + return ErrorAction(result_status="failed") + + 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") + + +@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 + + +SessionMetricsRecorder = MetricsRecorder +InMemorySessionMetricsRecorder = InMemoryMetricsRecorder + + +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: ... + + +@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.""" + + 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): + 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", ())), + ) + + 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: "InputSource" + output_sink: OutputSink + cleanup_tasks: set[asyncio.Task[RunResult]] + metrics: SessionMetricsRecorder = field( + default_factory=InMemorySessionMetricsRecorder + ) + error_policy: ErrorPolicy = field(default_factory=DefaultErrorPolicy) + transport: TransportService = field(default_factory=NoopTransportService) + clock: "RealtimeClock | DeterministicClock | None" = None + activation: "ActivationPolicy | None" = None + _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 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 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, + *, + 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.record_cleanup_error(exc) + try: + self.transport.close() + except Exception as exc: + self.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): + """Run/session construction strategy consumed by shared helpers.""" + + name: str + capabilities: RunModeCapabilities + + def validate_run( + self, + *, + spec: "DemoSpec", + adapter: "DemoAdapter", + ) -> None: ... + + def validate_session( + self, + *, + spec: "DemoSpec", + scenario: "PreparedScenario", + adapter: "DemoAdapter", + provider: "ModelInputProvider", + ) -> None: ... + + def create_run_context( + self, + *, + spec: "DemoSpec", + adapter: "DemoAdapter", + host: "RuntimeHost", + model_warmup_plan: ModelWarmupPlan, + ) -> RunContext: ... + + def create_session_edges( + self, + *, + context: RunContext, + spec: "DemoSpec", + scenario: "PreparedScenario", + provider: "ModelInputProvider", + adapter: "DemoAdapter", + ) -> SessionEdges: ... + + 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: ... + + +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", + "BenchmarkErrorPolicy", + "DefaultErrorPolicy", + "DriverStatus", + "ErrorAction", + "ErrorPolicy", + "InMemorySessionMetricsRecorder", + "MetricsSnapshot", + "Mp4ErrorPolicy", + "NativeWindowErrorPolicy", + "NoopTransportService", + "NullErrorPolicy", + "RunContext", + "RunMode", + "RunModeCapabilities", + "RunModeWarmup", + "RunResult", + "RunSummary", + "SessionEdges", + "SessionDriver", + "SessionMetricsRecorder", + "SessionReservation", + "SessionStatus", + "SingleSessionAdmissionPolicy", + "TransportService", + "WebRTCErrorPolicy", + "build_model_warmup_plan", + "warmup_run_context", +] diff --git a/flashdreams/flashdreams/runtime/demo/session_inputs.py b/flashdreams/flashdreams/runtime/demo/session_inputs.py new file mode 100644 index 000000000..dda42995c --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/session_inputs.py @@ -0,0 +1,184 @@ +# 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 TYPE_CHECKING, Protocol, runtime_checkable + +from flashdreams.runtime._utils import freeze_mapping +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.""" + + 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 + user_input_schema: UserInputSchema + + 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: StepRequirements) -> 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: StepRequirements, + clock: "RealtimeClock", + ) -> "RealtimeWindowResult": + """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.""" + + capabilities: ProviderCapabilities + + def prepare_initial_input(self) -> InferenceInput: + """Prepare session-global model inputs.""" + ... + + def prepare_step( + self, + *, + request: StepRequirements, + 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. + + Implementations must be idempotent so driver cleanup and reset control + paths can safely converge after failures. + """ + ... + + def close(self) -> None: + """Release provider-owned resources. + + Implementations must be idempotent and tolerate cleanup after partial + setup or earlier reset failures. + """ + ... + + +__all__ = [ + "BatchInputSource", + "ControlDecision", + "InputSource", + "ModelInputProvider", + "PreparedStep", + "ProviderCapabilities", + "RealtimeInputSource", + "UserInputWindow", +] 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/flashdreams/runtime/demo/timing.py b/flashdreams/flashdreams/runtime/demo/timing.py new file mode 100644 index 000000000..d48e1e213 --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/timing.py @@ -0,0 +1,383 @@ +# 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, UserInputSchema +from flashdreams.runtime.types import StepRequirements + +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: StepRequirements, + 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) + + +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.""" + + 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: StepRequirements, + 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: _SparseInputResampler + max_lag_s: float | None = None + 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 ( + 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: StepRequirements, + 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: StepRequirements) -> int: + """Return the positive input frame count declared by a step requirement.""" + + value = request.input_frame_count + if isinstance(value, bool) or not isinstance(value, int): + raise ValueError("StepRequirements.input_frame_count must be an integer.") + parsed = value + if parsed <= 0: + raise ValueError("StepRequirements.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/flashdreams/runtime/demo/validation.py b/flashdreams/flashdreams/runtime/demo/validation.py new file mode 100644 index 000000000..c5058b148 --- /dev/null +++ b/flashdreams/flashdreams/runtime/demo/validation.py @@ -0,0 +1,200 @@ +# 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( + f"Input source does not satisfy provider raw user input schema: {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/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/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 4130a925f..49ee958a0 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,36 @@ def __post_init__(self) -> None: object.__setattr__(self, "metadata", freeze_mapping(self.metadata)) -__all__ = ["StepRequest", "StepResult"] +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 and not allow_user_input_window: + 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/flashdreams/runtime/worker.py b/flashdreams/flashdreams/runtime/worker.py index af4576c09..e523b6e6a 100644 --- a/flashdreams/flashdreams/runtime/worker.py +++ b/flashdreams/flashdreams/runtime/worker.py @@ -6,15 +6,17 @@ from __future__ import annotations import asyncio +import threading from concurrent.futures import ThreadPoolExecutor from typing import Any, Callable, TypeVar, cast import torch _T = TypeVar("_T") +_EXECUTOR_FUTURE_POLL_INTERVAL_S = 0.01 -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 +40,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,11 +65,11 @@ 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) + return await _await_executor_future(future) except asyncio.CancelledError: future.add_done_callback(_consume_exception) raise @@ -71,21 +82,32 @@ 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: - if self._closed: - return - self._accepting = False - barrier = self._submit(_noop, (), {}) - await asyncio.shield(barrier) - self._executor.shutdown(wait=True, cancel_futures=False) - self._closed = True + self._require_not_worker_thread() + 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() + if not self._begin_close(): + return + try: + barrier = self._executor.submit(_noop) + barrier.result() + finally: + self._finish_close() def _submit( self, @@ -97,9 +119,36 @@ 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." + ) + + 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 + def _invoke( func: Callable[..., _T], @@ -113,9 +162,19 @@ 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() -__all__ = ["ThreadAffineRuntimeWorker"] +__all__ = ["ModelExecutionWorker", "ThreadAffineRuntimeWorker"] 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/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/manager.py b/flashdreams/flashdreams/serving/webrtc/manager.py index a7c6c87a4..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, @@ -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, @@ -37,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 ( @@ -55,6 +82,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,8 +117,12 @@ """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) +_RuntimeT = TypeVar("_RuntimeT") _RuntimeConfigT = TypeVar("_RuntimeConfigT", bound=WebRTCRuntimeConfig) @@ -140,6 +183,421 @@ 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, + ) + + +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.""" + + 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 + 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=inference_input, + 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 +610,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 +650,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() @@ -212,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") @@ -234,6 +705,15 @@ 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 = 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: @@ -296,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", ) @@ -315,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: @@ -330,10 +826,60 @@ 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 + def _shared_run_context(self, loop: asyncio.AbstractEventLoop) -> RunContext: + if self._shared_context is not None: + return self._shared_context + 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(), + 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): + if self._shared_adapter is not None: + return + 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 +1042,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, @@ -769,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() @@ -808,46 +1386,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 +1451,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 +1468,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,15 +1595,37 @@ async def _client_liveness_watchdog( async def shutdown(self) -> None: await self.close_active_session() - await self._runtime.close() + if self._shared_context is not None: + await self._shared_context.close_async() + if self._shared_host is not None: + await asyncio.to_thread(self._shared_host.close) + self._shared_context = None + self._shared_host = None + self._shared_runtime_adapter = None + 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, @@ -1013,6 +1637,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 +1657,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 +1740,202 @@ 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 = 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, + 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=( + self._shared_pipeline_factory() + if self._shared_pipeline_factory is not None + else StepPipeline() + ), + reservation=managed_session.reservation, + ) + 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)) + 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 470569132..25fa09437 100644 --- a/flashdreams/flashdreams/serving/webrtc/media.py +++ b/flashdreams/flashdreams/serving/webrtc/media.py @@ -81,16 +81,37 @@ 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 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 @@ -234,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/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/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/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/services.py b/flashdreams/flashdreams/serving/webrtc/services.py new file mode 100644 index 000000000..3ce38e4be --- /dev/null +++ b/flashdreams/flashdreams/serving/webrtc/services.py @@ -0,0 +1,1161 @@ +# 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 import ( + StepRequirements, + StepResult, + UserInputCapability, + UserInputEvent, + UserInputs, + UserInputSchema, +) +from flashdreams.runtime._utils import freeze_mapping +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"}) +WEBRTC_SKIPPED_INPUTS_METADATA_KEY = "webrtc_skipped_inputs" +WEBRTC_SKIPPED_WINDOW_METADATA_KEY = "webrtc_skipped_window" + + +@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)) + + +@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.""" + + 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 begin_generation(self, generation: int) -> None: ... + + 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) + 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 + 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) + 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=metadata, + ), + 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 + 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: + 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[Any]: + 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_chunk_delivery: Callable[[WebRTCChunkDelivery], 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_chunk_delivery = on_chunk_delivery + self._on_error = on_error + 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, + *, + generation: int, + force_keyframe: bool = False, + ) -> WebRTCOutputBridgeDecision: + 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 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, + dropped=True, + drop_policy="drop_newest", + 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( + accepted=False, + should_stop=True, + dropped=True, + 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, + dropped=True, + drop_policy="drop_newest", + metadata={"reason": "pending queue full"}, + ) + future = asyncio.run_coroutine_threadsafe( + self._deliver( + payload, + chunk=chunk, + generation=generation, + force_keyframe=force_keyframe, + ), + self._loop, + ) + self._pending[future] = generation + 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, + payload: object, + *, + chunk: WebRTCChunkDelivery, + generation: int, + force_keyframe: bool, + ) -> 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[WebRTCChunkDelivery]) -> None: + with self._lock: + self._pending.pop(future, None) + 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.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) + 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) + + 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.""" + + 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", + "WEBRTC_SKIPPED_INPUTS_METADATA_KEY", + "WEBRTC_SKIPPED_WINDOW_METADATA_KEY", + "WebRTCActivationPolicy", + "WebRTCInputSource", + "WebRTCMessageResult", + "WebRTCOfferAnswerer", + "WebRTCOfferRequest", + "WebRTCOutputBridge", + "WebRTCOutputBridgeDecision", + "WebRTCChunkDelivery", + "WebRTCOutputSink", + "WebRTCRunMode", + "WebRTCSessionEdgeFactory", + "WebRTCSessionOfferHandler", + "WebRTCTransportService", +] 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/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_demo_runtime_host.py b/flashdreams/tests/test_demo_runtime_host.py new file mode 100644 index 000000000..b14dafcb6 --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_host.py @@ -0,0 +1,218 @@ +# 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_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_realtime_driver.py b/flashdreams/tests/test_demo_runtime_realtime_driver.py new file mode 100644 index 000000000..5c79b83f7 --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_realtime_driver.py @@ -0,0 +1,916 @@ +# 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, + StepRequirements, + StepResult, + UserInputSchema, +) +from flashdreams.runtime.demo import ( + ActivationResult, + DriverInvariantError, + ErrorAction, + InMemorySessionMetricsRecorder, + OutputDecision, + PreparedStep, + ProviderCapabilities, + RealtimeSessionDriver, + RealtimeWindowResult, + RunContext, + RunResult, + RuntimeHost, + SessionEdges, + SessionInfo, + SingleSessionAdmissionPolicy, + StepPipeline, + UserInputWindow, + 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, + ) + ) + + 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(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() + + 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_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)) + 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"), + ) + assert not host.is_healthy + assert host.unhealthy_reason == "model-affine cleanup failed" + 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() + 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="timeout test", + error=None, + timeout_s=0.001, + ) + + assert result.status == "cancelled" + assert host.unhealthy_reason == "model-affine cleanup timed out" + 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 + + +@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 failed" + 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) + 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, metrics=metrics) + + 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 metrics.catch_up_count == 2 + 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 + user_input_schema = UserInputSchema() + + 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[StepRequirements] = [] + + def is_finished(self) -> bool: + return False + + async def next_realtime_window( + self, + *, + request: StepRequirements, + 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: + capabilities = ProviderCapabilities( + supports_realtime_clock=True, + supports_reset=True, + deterministic_given_inputs=False, + ) + + 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: StepRequirements, + 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_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: + raise AssertionError("demo driver should request StepRequirements") + + 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 _InvariantPipeline(StepPipeline): + def execute_step(self, **kwargs: object) -> Any: + del kwargs + raise DriverInvariantError("step invariant") + + +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 + self.close_targets: list[Any] = [] + + async def call_async( + self, + func: Callable[..., Any], + /, + *args: object, + **kwargs: object, + ) -> Any: + del func, kwargs + 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( + self, + reason: str = "marked unhealthy", + error: Exception | None = None, + ) -> 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 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..8742eb0e9 --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_run_modes.py @@ -0,0 +1,774 @@ +# 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 +import threading +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, + StepRequirements, + UserInputSchema, +) +from flashdreams.runtime.demo import ( + AsyncSessionDriver, + BenchmarkErrorPolicy, + DemoSpec, + DriverInvariantError, + InMemorySessionMetricsRecorder, + ModelWarmupPlan, + Mp4ErrorPolicy, + Mp4OutputSpec, + NativeWindowErrorPolicy, + NullErrorPolicy, + NullOutputSpec, + OutputDecision, + PreparedScenario, + ProviderCapabilities, + RunContext, + RunModeCapabilities, + RunResult, + RuntimeHost, + SessionDriver, + SessionEdges, + SessionInfo, + StepPipeline, + UserInputWindow, + WebRTCErrorPolicy, + 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 + + +@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", + 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 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, + 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 _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, + supports_reset=True, + deterministic_given_inputs=True, + ) + + def __init__(self) -> None: + self.close_count = 0 + + 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, + *, + 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.capabilities = RunModeCapabilities(requires_finite_input=True) + 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 _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, + *, + 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 + 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 + + 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_timing.py b/flashdreams/tests/test_demo_runtime_timing.py new file mode 100644 index 000000000..98e3cf56d --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_timing.py @@ -0,0 +1,237 @@ +# 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 StepRequirements, UserInputSchema +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) -> StepRequirements: + return StepRequirements( + step_index=0, + 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 + user_input_schema = UserInputSchema() + + def is_finished(self) -> bool: + return False + + async def next_realtime_window( + self, + *, + request: StepRequirements, + clock: Any, + ) -> object: + del request, clock + raise AssertionError("Activation timeout must close before requesting input.") diff --git a/flashdreams/tests/test_demo_runtime_validation.py b/flashdreams/tests/test_demo_runtime_validation.py new file mode 100644 index 000000000..c5780bd72 --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_validation.py @@ -0,0 +1,680 @@ +# 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, + CanonicalInputs, + CanonicalInputSchema, + IdentityInputMapping, + InferenceConfig, + InferenceInput, + InferenceInputSchema, + InferenceRuntime, + InputCanonicalizer, + InputField, + InputMapping, + InputMappingSchema, + KeyboardToDriverCommand, + OutputArtifact, + StepRequest, + StepRequirements, + StepResult, + TimeWindow, + UserInputCapability, + UserInputEvent, + UserInputs, + UserInputSchema, +) +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 new file mode 100644 index 000000000..1f430a459 --- /dev/null +++ b/flashdreams/tests/test_demo_runtime_vertical_slice.py @@ -0,0 +1,1328 @@ +# 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 +from collections.abc import Callable, Sequence +from typing import Any, Literal, cast + +import pytest + +import flashdreams.runtime.demo.drivers as drivers_module +from flashdreams.runtime import ( + CanonicalInputSchema, + IdentityInputMapping, + InferenceConfig, + InferenceInput, + InferenceInputSchema, + InferenceRuntime, + InferenceSession, + InputMapping, + OutputArtifact, + StepRequest, + StepRequirements, + StepResult, + UserInputs, + UserInputSchema, +) +from flashdreams.runtime.demo import ( + BatchSessionDriver, + ControlDecision, + DemoSpec, + DriverInvariantError, + ErrorAction, + InMemorySessionMetricsRecorder, + ModelWarmupPlan, + NullOutputSpec, + OutputDecision, + PreparedScenario, + PreparedStep, + ProviderCapabilities, + RunContext, + RunModeCapabilities, + RunResult, + RuntimeHost, + SessionEdges, + SessionInfo, + SingleSessionAdmissionPolicy, + StepPipeline, + UserInputWindow, + run_demo_session, + run_demo_session_async, +) + +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 = StepRequirements(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, + cleanup_tasks=set(), + 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 + assert "prepare_step" not in host.calls + 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) + 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) + 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(), + cleanup_tasks=set(), + 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_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() + 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] + 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: + 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 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 == [] + + +@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 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 == [] + + +@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 == [] + + +@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))), + provider=_FakeVideoModelInputProvider(fail_initial=RuntimeError("skip me")), + session_edges=SessionEdges( + input_source=_FakeBatchInputSource(num_windows=1), + output_sink=_RecordingOutputSink(), + cleanup_tasks=set(), + error_policy=_SetupPolicy(result_status="skipped"), + ), + pipeline=StepPipeline(), + ) + 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, + 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=RuntimeHost(_FakeVideoRuntime(session=_FakeVideoSession(num_steps=1))), + 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_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_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) + 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() + 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 + 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: + 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(), + cleanup_tasks=set(), + 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, + cleanup_tasks=set(), + ), + 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, + cleanup_tasks=set(), + 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_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 + 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 _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", + 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: + capabilities = ProviderCapabilities( + supports_recorded_input=True, + supports_reset=True, + deterministic_given_inputs=True, + ) + + 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"} + ) + 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: StepRequirements, + 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 + if self.fail_close is not 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 + user_input_schema = UserInputSchema() + + 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[StepRequirements] = [] + 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: 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 + user_input_schema = UserInputSchema() + + 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 + 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, + 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] = [] + self.close_count = 0 + + def session_info(self) -> SessionInfo: + return SessionInfo(output_layout="fake-video", steady_output_frame_count=1) + + 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, + 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: + 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 _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 + + def __init__( + self, + *, + 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: + del generation + + def write(self, result: StepResult) -> OutputDecision: + self.results.append(result) + return self.decision + + def close(self) -> Sequence[OutputArtifact]: + self.close_count += 1 + if self.fail_close is not None: + raise self.fail_close + 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 _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, + *, + 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: + name = "fake" + + def __init__( + self, + *, + input_source: _FakeBatchInputSource, + output_sink: _RecordingOutputSink | None = None, + output_sink_factory: ( + Callable[[DemoSpec, PreparedScenario], _RecordingOutputSink] | None + ) = None, + metrics: InMemorySessionMetricsRecorder | None = None, + 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() + 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 + self.select_error = select_error + self.capabilities = RunModeCapabilities( + requires_finite_input=True, + supports_artifacts=True, + ) + + 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, + *, + 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 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, + cleanup_tasks=context.cleanup_tasks, + metrics=self.metrics, + error_policy=self.error_policy or _SetupPolicy(result_status="failed"), + transport=self.transport or _RecordingTransport(), + ) + + def select_driver(self) -> BatchSessionDriver: + if self.select_error is not None: + raise self.select_error + return BatchSessionDriver() 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 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_inference_runtime_api.py b/flashdreams/tests/test_inference_runtime_api.py index 42f75d688..531f2f218 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,15 +17,19 @@ InferenceInputSchema, InMemoryMetricsRecorder, InputField, + MetricsSnapshot, + NullMetricsRecorder, NullOutputTarget, OutputArtifact, RuntimeMetricSample, StepRequest, + StepRequirements, StepResult, TimeWindow, UserInputEvent, UserInputs, UserInputSchema, + step_requirements_from_request, ) pytestmark = pytest.mark.ci_cpu @@ -38,11 +43,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): @@ -54,6 +61,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"), [ @@ -67,6 +81,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 +205,66 @@ 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_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") @@ -233,6 +313,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_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/flashdreams/tests/test_runtime_runner.py b/flashdreams/tests/test_runtime_runner.py index b755ff48a..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, @@ -25,6 +26,7 @@ InputField, InputMapping, InputMappingSchema, + MetricsSnapshot, NullOutputTarget, OutputArtifact, RuntimeMetricSample, @@ -73,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() @@ -656,5 +709,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() 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() diff --git a/flashdreams/tests/test_webrtc_manager.py b/flashdreams/tests/test_webrtc_manager.py index 8e4839e5e..787a8ee7b 100644 --- a/flashdreams/tests/test_webrtc_manager.py +++ b/flashdreams/tests/test_webrtc_manager.py @@ -11,7 +11,16 @@ import pytest import torch -from flashdreams.runtime import StepRequest, StepResult +from flashdreams.runtime import ( + InferenceInput, + StepRequest, + StepRequirements, + StepResult, + 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 from flashdreams.serving.webrtc.encoders import ChunkDeliveryResult @@ -20,6 +29,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 +98,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 +160,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 +483,94 @@ 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 __init__(self) -> None: + self.inference_inputs: list[InferenceInput] = [] + + def map_step_inputs( + self, + *, + canonical_inputs: object, + inference_input: InferenceInput, + request: StepRequest, + ) -> InferenceInput: + 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=mapping, + ) + 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, + 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), + }, + ), + ) + + 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 async def test_action_keydown_reports_error_when_user_event_queue_full( monkeypatch: pytest.MonkeyPatch, @@ -841,6 +989,147 @@ 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_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_services.py b/flashdreams/tests/test_webrtc_services.py new file mode 100644 index 000000000..bda6eee4c --- /dev/null +++ b/flashdreams/tests/test_webrtc_services.py @@ -0,0 +1,751 @@ +# 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()) + step_result = 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 + + encoder.release.set() + await asyncio.wait_for(encoder.done.wait(), timeout=1.0) + 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_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() + 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() + self.prepared_payloads: list[int] = [] + self.delivered_payloads: list[object] = [] + + 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 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 + + def __init__(self) -> None: + self.flush_count = 0 + + def qsize(self) -> int: + return 0 + + async def flush(self) -> None: + self.flush_count += 1 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/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/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/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 diff --git a/integrations/omnidreams/omnidreams/demo/__init__.py b/integrations/omnidreams/omnidreams/demo/__init__.py index 6fa3a9b21..a3d646989 100644 --- a/integrations/omnidreams/omnidreams/demo/__init__.py +++ b/integrations/omnidreams/omnidreams/demo/__init__.py @@ -4,17 +4,33 @@ """Experimental OmniDreams demo adapter built on ``flashdreams.runtime.demo``.""" from omnidreams.demo.adapter import OmnidreamsDemoAdapter +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 a1ca8e8f7..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 @@ -17,28 +18,45 @@ InferenceInput, InferenceInputSchema, InputCanonicalizer, - InputField, UserInputSchema, ) from flashdreams.runtime.demo import ( DemoSpec, - Mp4OutputSpec, PreparedScenario, ) +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 .replay import ( - OmnidreamsReplayRuntime, - OmnidreamsReplayRuntimeOptions, +from .providers import ( + LudusSceneConditioningProvider, + PrecomputedHDMapProvider, + keyboard_driving_user_input_schema, + precomputed_hdmap_inference_input_schema, +) +from .runtime import ( + OmnidreamsRuntime, + OmnidreamsRuntimeOptions, PipelineFactory, ) 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, ) -ReplayRuntimeFactory = Callable[..., InferenceRuntime] +RuntimeFactory = Callable[..., InferenceRuntime] +ReplayRuntimeFactory = RuntimeFactory class OmnidreamsDemoAdapter: @@ -47,10 +65,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() @@ -60,15 +87,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: @@ -78,31 +97,57 @@ 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",) + return ("mp4", "null", "webrtc") + + def supported_conditioning_modes(self) -> tuple[str, ...]: + return OMNIDREAMS_CONDITIONING_MODES def prepare_scenario(self, spec: DemoSpec) -> PreparedScenario: + if spec.output.mode not in self.supported_output_modes(): + raise ValueError( + "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 currently supports only " - f"input_mode='replay', got {spec.input_mode!r}." + "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( + 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}." ) - if not isinstance(spec.output, Mp4OutputSpec): - raise ValueError("OmniDreams replay demo currently requires MP4 output.") - scenario = resolve_replay_scenario( - spec.scenario, - default_prompt=self._default_replay_prompt(spec.config), - ) 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), @@ -119,11 +164,111 @@ 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, + release_oneshot_encoders_after_cache_init=( + _bool_runtime_option( + config.runtime_options, + "release_oneshot_encoders_after_cache_init", + True, + ) + ), + ), + ) + + def create_model_input_provider( + self, + spec: DemoSpec, + scenario: PreparedScenario, + ) -> ModelInputProvider: + if spec.input_mode not in {"replay", "keyboard-driving"}: + raise ValueError( + "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.") + 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, + ) + + 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")), ), ) @@ -153,7 +298,25 @@ 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", + "RuntimeFactory", ] diff --git a/integrations/omnidreams/omnidreams/demo/app.py b/integrations/omnidreams/omnidreams/demo/app.py index 1366643f9..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 @@ -15,6 +16,7 @@ from flashdreams.runtime.demo import ( DemoSpec, Mp4OutputSpec, + NullOutputSpec, WebRTCOutputSpec, ) from flashdreams.runtime.demo.app import DemoApplication @@ -22,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, ) @@ -33,13 +39,32 @@ 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("--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, @@ -54,7 +79,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 +100,21 @@ 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.") + 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 class OmnidreamsDemoApplication(DemoApplication): @@ -108,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, @@ -117,27 +158,56 @@ 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, 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, device=args.device, + seed=args.seed, + runtime_options={"seed": args.seed}, ), ) +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/omnidreams/demo/providers.py b/integrations/omnidreams/omnidreams/demo/providers.py new file mode 100644 index 000000000..eff6de760 --- /dev/null +++ b/integrations/omnidreams/omnidreams/demo/providers.py @@ -0,0 +1,692 @@ +# 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 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 +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.demo import ( + PreparedScenario, + PreparedStep, + ProviderCapabilities, + UserInputWindow, +) +from flashdreams.runtime.demo.session_inputs import ControlDecision +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 ( + DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, + OmnidreamsLudusReplayScenario, + OmnidreamsReplayScenario, +) + + +class PrecomputedHDMapProvider: + """Prepare fixed OmniDreams HDMap conditioning for replay-style runs.""" + + def __init__( + self, + *, + scenario: PreparedScenario, + config: InferenceConfig, + ) -> None: + self._scenario = _precomputed_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.") + + +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=scenario.source_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=( + 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 _precomputed_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 _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'))}") + return torch.device(config.device or "cuda") + + +def _is_rank_zero() -> bool: + return not dist.is_initialized() or dist.get_rank() == 0 + + +__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 8ccb58650..b6c6d91e2 100644 --- a/integrations/omnidreams/omnidreams/demo/replay.py +++ b/integrations/omnidreams/omnidreams/demo/replay.py @@ -1,246 +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 -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 omnidreams.runner import _load_video - -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, StepResult - -from .spec import OmnidreamsReplayScenario - -PipelineFactory = Callable[[Any, str], Any] - - -@dataclass(frozen=True, kw_only=True, slots=True) -class OmnidreamsReplayRuntimeOptions: - """Construction knobs for the replay runtime.""" - - pipeline_config: Any - pipeline_factory: PipelineFactory | None = None - output_layout: VideoTensorLayout = "bvtchw" - - -class OmnidreamsReplayRuntime: - """Heavyweight OmniDreams runtime consumed by ``run_inference_session``.""" - - def __init__( - self, - *, - config: InferenceConfig, - options: OmnidreamsReplayRuntimeOptions, - ) -> 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 OmnidreamsReplaySession( - pipeline=self.pipeline, - scenario=scenario, - 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, - ) - - 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 OmnidreamsReplaySession: - """One MP4 replay rollout over a prepared scenario.""" - - def __init__( - self, - *, - pipeline: Any, - scenario: OmnidreamsReplayScenario, - device: torch.device, - is_rank_zero: bool, - output_layout: VideoTensorLayout, - ) -> None: - self.pipeline = pipeline - self.scenario = scenario - 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( - postprocess_stream=None, - output_layout=self.output_layout, - ), - ) - 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(): - dist.barrier() - - def next_step_request(self) -> StepRequest | 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() - if self._frame_start + num_frames > self._hdmap_videos.shape[2]: - return None - return StepRequest(step_index=step_index) - - 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 - logger.info( - "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] - ) - self._frame_start = frame_end - return result - - 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._model_session.reset(self._initialize_cache) - self._frame_start = 0 - - def close(self) -> None: - self._closed = True - self._model_session.close() - - 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, - 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), - ) - 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() - - -def _scenario_from_inputs(inputs: InferenceInput) -> OmnidreamsReplayScenario: - scenario = inputs.global_conditioning.get("scenario") - if not isinstance(scenario, OmnidreamsReplayScenario): - raise TypeError( - "OmniDreams replay runtime requires global_conditioning['scenario'] " - "to be an OmnidreamsReplayScenario." - ) - return scenario - - -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", "PipelineFactory", ] diff --git a/integrations/omnidreams/omnidreams/demo/runtime.py b/integrations/omnidreams/omnidreams/demo/runtime.py new file mode 100644 index 000000000..1ee99ad31 --- /dev/null +++ b/integrations/omnidreams/omnidreams/demo/runtime.py @@ -0,0 +1,353 @@ +# 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 __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/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/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" + } + ] +} diff --git a/integrations/omnidreams/omnidreams/demo/web/adapter.css b/integrations/omnidreams/omnidreams/demo/web/adapter.css new file mode 100644 index 000000000..f4cb8868b --- /dev/null +++ b/integrations/omnidreams/omnidreams/demo/web/adapter.css @@ -0,0 +1,16 @@ +/* +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, 181.82vh); + height: auto; + max-height: min(100vh, 704px); + 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..37d19a299 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=model-ui-v2", controls: [ { label: "Drive / Turn", diff --git a/integrations/omnidreams/omnidreams/demo/webrtc.py b/integrations/omnidreams/omnidreams/demo/webrtc.py index 00857ea85..e5d8e2a84 100644 --- a/integrations/omnidreams/omnidreams/demo/webrtc.py +++ b/integrations/omnidreams/omnidreams/demo/webrtc.py @@ -13,458 +13,61 @@ # 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 tempfile -import time +import os from collections.abc import Callable -from dataclasses import dataclass, replace -from pathlib import Path +from dataclasses import replace 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.transformer import CosmosTransformerConfig -from flashdreams.runtime import InferenceConfig, StepResult -from flashdreams.runtime.demo import DemoSpec, WebRTCAppResources, WebRTCOutputSpec +from flashdreams.runtime import InferenceConfig +from flashdreams.runtime.demo import ( + DemoSpec, + RuntimeHost, + WebRTCAppResources, + WebRTCOutputSpec, +) 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.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 from flashdreams.serving.webrtc.server import create_webrtc_app +from .adapter import OmnidreamsDemoAdapter, RuntimeFactory from .spec import ( DEFAULT_OMNIDREAMS_PRESET, - DEFAULT_OMNIDREAMS_WEBRTC_SCENE_UUID, OMNIDREAMS_MODEL_ID, resolve_webrtc_scenario, ) +from .webrtc_config import OmnidreamsWebRTCModelRuntimeConfig WebRTCRuntimeFactory = Callable[..., Any] - - -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 = float(np.deg2rad(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.""" - - -class OmnidreamsWebRTCModelRuntime( - ThreadAffineDistributedWebRTCRuntime[ - OmnidreamsWebRTCModelRuntimeConfig, - None, - ] -): - """Run one single-view OmniDreams scene with browser camera controls.""" - - def __init__(self, *, config: OmnidreamsWebRTCModelRuntimeConfig) -> None: - super().__init__( - config=config, - 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", - ) - - def _is_runtime_initialized(self) -> bool: - return self._wrapper is not None and self._renderer is not None - - def _runtime_step_index(self) -> int: - return self._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) - - def _steady_output_frame_count(self) -> int: - return int(self._require_wrapper().frame_chunk_size) - - def _initialize_sync(self) -> None: - if self._wrapper 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) - - logger.info( - "Setting up OmniDreams pipeline {} on {}.", - cfg.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, - ) - - 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}." - ) - - 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." - ) - - 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) - - def _generate_one_chunk_sync( - self, - *, - 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(): - raise OmnidreamsWebRTCModelRuntimeError( - f"Expected {self._next_input_frame_count()} frame times for " - f"step {self._step_index}, got {len(frame_times)}." - ) - if not segments: - raise OmnidreamsWebRTCModelRuntimeError( - f"Step {self._step_index} received no control segments." - ) - - 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, - ) - 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, - ) - self._state = output.state - if self._state.pipeline_cache is not None: - wrapper.finalize_block_generation( - self._state.pipeline_cache, - output.finalization_state, - ) - - 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( - step_index=self._step_index, - video_chunk=video_chunk.detach(), - layout="bvtchw", - metadata=metadata, - ) - expected_frames = len(frame_times) - if result.frame_count != expected_frames: - raise OmnidreamsWebRTCModelRuntimeError( - f"Expected generated chunk to contain {expected_frames} frames, " - f"got {result.frame_count}." - ) - 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 _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 _require_wrapper(self) -> OmnidreamsConditioningWrapper: - if self._wrapper is None: - raise OmnidreamsWebRTCModelRuntimeError("Runtime is not initialized.") - return self._wrapper +SharedRuntimeFactory = RuntimeFactory 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.""" + """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." + ) if spec.input_mode != "keyboard-driving": raise ValueError( "OmniDreams WebRTC requires input_mode='keyboard-driving', " @@ -480,7 +83,49 @@ 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, + config=config, + scenario=scenario, + ) + if _should_use_legacy_webrtc_path( + 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, + 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( @@ -491,17 +136,57 @@ 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) - runtime = runtime_factory(config=runtime_config) + 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_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, @@ -511,12 +196,16 @@ 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, + shared_host=host, + shared_adapter=adapter, + shared_spec=spec, + shared_scenario=prepared, ) 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( @@ -529,6 +218,36 @@ def serve_omnidreams_webrtc_demo( ) +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 @@ -556,6 +275,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, @@ -576,9 +302,8 @@ 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/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 1d2f81ef3..3d083535a 100644 --- a/integrations/omnidreams/tests/test_demo_api.py +++ b/integrations/omnidreams/tests/test_demo_api.py @@ -3,51 +3,83 @@ from __future__ import annotations +import asyncio +import json +import sys 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 +import tomli as tomllib import torch from aiohttp import web 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, ) from omnidreams.demo.app import _replay_spec, _webrtc_spec, parse_args from omnidreams.demo.replay import ( OmnidreamsReplayRuntime, OmnidreamsReplayRuntimeOptions, + OmnidreamsReplaySession, +) +from omnidreams.demo.runtime import ( + OmnidreamsRuntime, + OmnidreamsRuntimeOptions, + OmnidreamsSession, ) from omnidreams.demo.webrtc import ( - OmnidreamsWebRTCModelRuntime, OmnidreamsWebRTCModelRuntimeConfig, + _should_use_legacy_webrtc_path, serve_omnidreams_webrtc_demo, ) from flashdreams.runtime import ( + CanonicalInputs, InferenceConfig, InferenceInput, OutputArtifact, OutputTarget, StepRequest, + StepRequirements, StepResult, ) from flashdreams.runtime.demo import ( DemoSpec, Mp4OutputSpec, + NullOutputSpec, + OutputDecision, + PreparedScenario, + RuntimeHost, + SessionInfo, + UserInputWindow, WebRTCOutputSpec, ) from flashdreams.runtime.demo.replay import run_replay_demo -from flashdreams.serving.webrtc.manager import BaseWebRTCSessionManager +from flashdreams.runtime.demo.timing import SPARSE_KEY_SEGMENTS_METADATA_KEY +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 @@ -59,12 +91,139 @@ def test_omnidreams_demo_defaults_to_stable_non_perf_preset() -> None: assert not args.preset_id.endswith("-perf") -def test_omnidreams_demo_adapter_declares_replay_modes_only() -> None: +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_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_shared_modes() -> 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_input_modes() == ("replay", "keyboard-driving") + assert adapter.supported_output_modes() == ("mp4", "null", "webrtc") + assert adapter.supported_conditioning_modes() == ( + OMNIDREAMS_CONDITIONING_PRECOMPUTED, + OMNIDREAMS_CONDITIONING_LUDUS, + ) + 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: + 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: @@ -107,14 +266,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 @@ -212,31 +374,364 @@ 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.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_ludus_provider_prepares_deterministic_hdmaps( tmp_path: Path, monkeypatch: pytest.MonkeyPatch, ) -> None: - import omnidreams.demo.replay as replay_module + _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, +) -> 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: 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=2, + ) + + result = run_replay_demo( + spec=spec, + adapter=adapter, + output_sink_factory=lambda output_spec: sink, ) - runtime = OmnidreamsReplayRuntime( + 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_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, +) -> 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, +) -> 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: 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.runtime as runtime_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( + runtime_module, + "load_first_frame_tensor", + lambda *args, **kwargs: torch.zeros(1, 3, 2, 2), + ) + + 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, ), @@ -254,11 +749,17 @@ 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 - result = session.step(InferenceInput()) + assert request.metadata["input_frame_count"] == 1 + result = session.step(InferenceInput(step={"hdmap": torch.zeros(1, 1, 1, 3, 2, 2)})) assert result.step_index == 0 assert result.frame_count == 1 @@ -267,6 +768,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"]], @@ -341,7 +843,10 @@ 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( model_id=OMNIDREAMS_MODEL_ID, preset_id=DEFAULT_OMNIDREAMS_PRESET, @@ -350,7 +855,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( @@ -371,32 +875,121 @@ 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, - runtime_factory=_FakeWebRTCRuntime, + shared_runtime_factory=shared_runtime_factory, server_runner=lambda **kwargs: calls.append(kwargs), ) 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._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 + assert legacy_module_name not in sys.modules + + +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(), + 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, + world_rank=1, + runtime_factory=_FakeWebRTCRuntime, + server_runner=lambda **kwargs: calls.append(kwargs), + ) + + manager = calls[0]["session_manager"] + runtime = manager._runtime + assert isinstance(runtime, _FakeWebRTCRuntime) + assert manager.runtime_config is runtime.config + 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( @@ -439,7 +1032,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, ) @@ -453,6 +1046,26 @@ 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=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" + 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: @@ -495,7 +1108,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, ) @@ -508,59 +1121,402 @@ 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: + 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") + 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 wrapper.finalized == [0, 1] + assert [tuple(hdmap.shape) for hdmap in pipeline.generated_hdmaps] == [ + (1, 1, 2, 3, 2, 2), + (1, 1, 3, 3, 2, 2), + ] + 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 + + +@pytest.mark.asyncio +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") + 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, + 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)) + 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, + ) + host = RuntimeHost(runtime) + manager = BaseWebRTCSessionManager( + runtime=runtime, + 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, + ) + manager._runtime_ready = True + 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=runtime_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=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] + 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"] == 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.shutdown() 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, + output: Mp4OutputSpec | NullOutputSpec | None = None, +) -> 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=output or 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()}, + ), + ) + + +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 + + 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 + raise NotImplementedError + + def close(self) -> None: + return 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( @@ -593,7 +1549,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]: @@ -601,61 +1558,171 @@ def finalize(self, *, autoregressive_index: int, cache: object) -> dict[str, flo return {"denoise_s": 0.25} -class _FakeRenderer: +class _VariableFrameOmnidreamsPipeline(_FakeOmnidreamsPipeline): + def __init__(self, frame_counts: tuple[int, ...]) -> None: + super().__init__() + self._frame_counts = frame_counts + + 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 _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 + 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 +class _FakeWebRTCResampler: + def __init__(self, *, start_v: float, fps: int) -> None: + self.dt = 1.0 / fps + self.next_chunk_start_v = start_v - def __init__(self) -> None: - self.calls: list[tuple[str, tuple[int, ...], list[int]]] = [] - self.finalized: list[int] = [] - self.cleaned = False + 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] = [] - def start_generation(self, **kwargs: Any) -> SimpleNamespace: - return self._output("start", kwargs=kwargs, frame_count=2, step_index=0) + async def enqueue_result(self, result: StepResult) -> int: + self.enqueued.append(result) + return result.frame_count - def continue_generation(self, **kwargs: Any) -> SimpleNamespace: - return self._output("continue", kwargs=kwargs, frame_count=3, step_index=1) + def qsize(self) -> int: + return 0 - def _output( + 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, - operation: str, + payload: object, + track: Any, *, - kwargs: dict[str, Any], - frame_count: int, - step_index: int, + force_keyframe: bool = False, ) -> 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()) + del force_keyframe + if not isinstance(payload, StepResult): + raise TypeError("Fake WebRTC encoder expected a StepResult payload.") 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}, + num_frames=await track.enqueue_result(payload), + encode_ms=0.0, ) - 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 _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: