Skip to content
Closed
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
4 changes: 2 additions & 2 deletions README.md
Original file line number Diff line number Diff line change
Expand Up @@ -20,7 +20,7 @@ Reactor Runtime turns an inference pipeline into a real-time, interactive media
- ✅ **Typed, validated commands.** Declare the commands your model accepts with standard Python types and constraints. The runtime validates every payload before your handler runs and compiles the surface into an OpenAPI schema that drives typed client SDKs.
- 🔎 **Traceable logs.** `get_logger()` writes structured records — readable `key=value` in a terminal, JSON for a log pipeline. Every record a session writes carries that session's id automatically, so one filter recovers everything a single run logged.
- 📦 **One container, anywhere.** The `reactor` CLI scaffolds a workspace, builds a small image, and runs it locally. The same image deploys to [Reactor](https://reactor.inc)'s GPU cloud unchanged.
- 🧩 **Experimental multi-GPU workers.** Coordinate local GPU ranks with `WorkerGroup`, using the GPU count declared in your model manifest.
- 🧩 **Experimental multi-GPU video.** Declare a worker and typed output with `DistributedVideoModel`; the runtime manages local GPU workers and session cleanup.

## How it works

Expand Down Expand Up @@ -81,7 +81,7 @@ logger.info("scene changed", prompt=self.prompt)

Records render as `key=value` text by default, or as one JSON object per line under `REACTOR_LOG_FORMAT=json`. While a session is live, its id is stamped on every record, so tracing one run's logs never requires threading an id through your call sites. Every record also carries the lifecycle phase it was written in, at both granularities: `state`, the session state machine's word, and `runtime_state`, the coarse word the health endpoint serves — so the logs of one phase — loading weights, a live session, teardown — are filterable by whichever vocabulary you are reading off another surface. The stamp is applied where records are written rather than where they are made, so a plain `logging.getLogger(__name__)` and the libraries your model imports are covered too.

For single-node multi-GPU video, the experimental `reactor_runtime.distributed` package provides `DistributedWorker` rank hooks and `WorkerGroup` orchestration. Construct the group during model loading with the manifest-derived `self.world_size`; it owns spawning, frame transport, liveness checks, and teardown. These low-level primitives stay off the package root and may change between minor releases. They are not a generic action/tensor API or a multi-session scheduler.
For single-node multi-GPU video, the experimental `DistributedVideoModel` adapter keeps the same commands and typed output: declare a worker, snapshot controls, and map its frames to an `Output`. The runtime handles worker sessions, pause/restart, and cleanup using the manifest's GPU count. Raw `WorkerGroup` orchestration is reserved for advanced use. This video-only surface lives in `reactor_runtime.distributed`, stays off the package root, and may change between minor releases. It is not a generic action/tensor API or a multi-session scheduler.

## Install

Expand Down
54 changes: 37 additions & 17 deletions src/reactor_runtime/distributed/__init__.py
Original file line number Diff line number Diff line change
@@ -1,28 +1,48 @@
"""Experimental single-node, single-session multi-GPU video primitives.

Signatures may change between minor releases. Import from
``reactor_runtime.distributed``; these names stay off the package root.

Subclass :class:`DistributedWorker` to define per-rank setup, warmup, session,
and generation hooks. :class:`WorkerGroup` owns process spawning, commands and
acknowledgements, bounded uint8 frame transport, liveness checks, and teardown.
A one-worker group runs inline without child processes or a process group.
Use the manifest-derived ``self.world_size`` when constructing the group
during model loading.

Torch is imported lazily in workers when the model image provides it. The
controller stays torch-free. The framework creates the process group before
``setup``, and the model owns its compute collectives. Clean shutdown includes
a final barrier before destruction. This package does not provide arbitrary
structured outputs or multi-session scheduling.
"""Multi-GPU worker abstraction for real-time streaming models.

.. warning:: **Experimental.** This package supports single-node, single-session
uint8 video, not arbitrary structured outputs or multi-session scheduling.
Signatures may change between minor releases. Import from
``reactor_runtime.distributed``; these names stay off the package root.

Use :class:`DistributedVideoModel` for the video authoring path: declare a
worker and map its frames to an Output. The adapter owns threading, session
bookkeeping, pause/restart, and cleanup. It is a :class:`ReactorModel`, not an
engine host or multi-session scheduler. Use :class:`WorkerGroup` directly only
when custom orchestration requires the raw primitives.

A model that needs several GPUs cannot simply run several copies of
itself: it is a server, with one event loop, one session, and one output
stream. So the process splits in two roles, vended here as two classes:

- :class:`DistributedWorker` — the per-GPU class a model author
subclasses. Hooks: ``setup`` / ``warmup`` / ``start_session`` /
``generate_chunk`` / ``end_session``.
- :class:`WorkerGroup` — the controller handle a model holds, created in
``load()``. Owns process spawning, the command/ack protocol, the
frame transport, liveness detection, and teardown. A one-worker group
runs inline without child processes, shared memory, or a process group.

``torch`` is imported lazily and only where available: the controller
side of a multi-worker group runs torch-free, and workers use it (NCCL process group,
CUDA device binding, host-memory pinning) when the image provides it.

The process group belongs entirely to the model. The framework creates
it before ``setup`` runs and issues no compute collectives of its own, so a
model is free to carve its own sub-groups out of it — tensor-,
sequence-, or context-parallel — by calling ``new_group`` or
``init_device_mesh`` inside ``setup``, where every rank reaches the call
in the same order. Clean shutdown includes a final barrier before destruction.
"""

from reactor_runtime.distributed.errors import WorkerCrashed, WorkerError
from reactor_runtime.distributed.frames import SharedFrameBuffer
from reactor_runtime.distributed.group import WorkerGroup
from reactor_runtime.distributed.model import DistributedVideoModel
from reactor_runtime.distributed.worker import DistributedWorker

__all__ = [
"DistributedVideoModel",
"DistributedWorker",
"SharedFrameBuffer",
"WorkerCrashed",
Expand Down
271 changes: 271 additions & 0 deletions src/reactor_runtime/distributed/model.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,271 @@
# Copyright (c) 2026 Reactor Technologies, Inc. All rights reserved.
"""Experimental single-session video adapter on the ordinary model lifecycle."""

from __future__ import annotations

import asyncio
import copy
import math
import pickle
import time
from collections.abc import Callable
from concurrent.futures import ThreadPoolExecutor
from functools import partial
from pathlib import Path
from typing import Any, ClassVar

import numpy as np

from reactor_runtime.core.model import (
ClientDisconnected,
ReactorEvent,
SessionEnded,
SessionStarted,
)
from reactor_runtime.distributed.group import WorkerGroup
from reactor_runtime.distributed.worker import DistributedWorker
from reactor_runtime.interface.model.reactor_model import ReactorModel
from reactor_runtime.interface.tracks import Output


class DistributedVideoModel(ReactorModel):
"""Declare per-rank computation without managing worker sessions or threads.

Experimental, single-node, single-session uint8 video only. Declare
``worker`` and ``frame_shape``, implement :meth:`to_output`, and optionally
:meth:`controls`. Commands and lifecycle hooks use the ordinary decorators.
The manifest supplies ``world_size``; one worker runs in-process.

The adapter owns ``load`` and ``run``. Use :meth:`worker_setup` to read model
configuration and pass weight paths to workers; use :meth:`session_params`
for initialization data. Every hook returning a dictionary is snapshotted
before dispatch, and must contain picklable CPU data. GPU state belongs to
the worker, not the controller.

``paused`` holds the current worker session, including one completed
in-flight chunk. :meth:`restart_generation` instead drops that chunk and
reinitializes the workers at the next safe boundary. Neither operation
interrupts a collective. Last-viewer disconnects also restart generation;
additional viewers share the existing sequence.

Attributes:
worker: Importable per-rank worker class.
frame_shape: Maximum ``(frames, height, width, channels)`` buffer shape.
session_seed: Initial seed, in NumPy's uint32 range.
adaptive_fps: Measure chunk throughput for playout; otherwise use ``fps``.
startup_timeout: Multi-worker setup/warmup deadline in seconds.
command_timeout: Multi-worker session/chunk deadline in seconds.
shutdown_timeout: Grace period before terminating worker processes.
init_process_group: Initialize NCCL/gloo; disable for CPU protocol tests.
"""

worker: ClassVar[type[DistributedWorker] | None] = None
frame_shape: ClassVar[tuple[int, ...] | None] = None
session_seed: int = 0
adaptive_fps: ClassVar[bool] = False
startup_timeout: ClassVar[float] = 3600.0
command_timeout: ClassVar[float] = 300.0
shutdown_timeout: ClassVar[float] = 60.0
init_process_group: ClassVar[bool] = True

def __init__(self) -> None:
super().__init__()
self._workers: WorkerGroup | None = None
self._executor: ThreadPoolExecutor | None = None
self._epoch = 0
self._active = False
self._paused = False
self._wakeup: asyncio.Event | None = None

def load(self, config_path: Path | None) -> None:
"""Validate declarations, then load every rank on its owning thread."""
if self._workers is not None:
raise RuntimeError("load() must be called only once")
if not isinstance(self.worker, type) or not issubclass(self.worker, DistributedWorker):
raise TypeError("declare worker = YourDistributedWorker on the model")
if self.frame_shape is None:
raise TypeError("declare frame_shape = (max_frames, height, width, channels)")
if type(self).to_output is DistributedVideoModel.to_output:
raise TypeError("implement to_output(self, frames) to return your typed Output")
if type(self.session_seed) is not int or not 0 <= self.session_seed < 2**32:
raise ValueError("session_seed must be an integer in [0, 2**32)")
for name in ("startup_timeout", "command_timeout", "shutdown_timeout"):
value = getattr(self, name)
if type(value) not in (int, float) or not math.isfinite(value) or value <= 0:
raise ValueError(f"{name} must be a finite positive number of seconds")
setup = _snapshot(self.worker_setup(config_path), "worker_setup()")
self._workers = WorkerGroup(
self.worker,
frame_shape=self.frame_shape,
world_size=self.world_size,
setup_kwargs=setup,
init_process_group=self.init_process_group,
)
self._executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="video-workers")
try:
self._executor.submit(self._workers.start, timeout=self.startup_timeout).result()
except BaseException:
self._executor.shutdown(wait=True)
raise

def worker_setup(self, config_path: Path | None) -> dict[str, Any]:
"""Read model configuration and return kwargs for every worker's setup."""
return {}

def session_params(self) -> dict[str, Any]:
"""Return worker initialization data for a new sequence or restart."""
return {}

def controls(self) -> dict[str, Any]:
"""Return the current conditioning, snapshotted at each chunk boundary."""
return {}

def to_output(self, frames: np.ndarray) -> Output:
"""Map a completed video chunk into the model's declared output track."""
raise NotImplementedError

@property
def paused(self) -> bool:
"""Whether generation is paused, preserving worker state for resume."""
return self._paused

@paused.setter
def paused(self, value: bool) -> None:
if type(value) is not bool:
raise TypeError("paused must be a bool")
if value and not self._paused:
self.output.flush()
self._paused = value
self._wake()

def restart_generation(self, *, seed: int | None = None) -> None:
"""Discard pending output and begin a fresh sequence after in-flight work.

Call from a command or lifecycle hook. Controls are preserved, and a
paused model remains paused. An optional seed also applies to subsequent
sequences; it must fit NumPy's uint32 range.
"""
if seed is not None:
if type(seed) is not int or not 0 <= seed < 2**32:
raise ValueError("seed must be an integer in [0, 2**32)")
self.session_seed = seed
self._epoch += 1
self.output.flush()
self._wake()

def _on_loop_ready(self) -> None:
super()._on_loop_ready()
self._wakeup = asyncio.Event()

async def _dispatch_reactor_event(self, event: ReactorEvent) -> None:
# Bookkeeping is independent of user-decorated hooks: overriding a hook
# cannot accidentally bypass stale-result protection or worker teardown.
if isinstance(event, SessionStarted):
self._active = False
self._paused = False
self.restart_generation()
elif isinstance(event, SessionEnded):
self._active = False
self.restart_generation()
elif isinstance(event, ClientDisconnected) and event.total == 0:
self.restart_generation()
await super()._dispatch_reactor_event(event)
if isinstance(event, SessionStarted):
self._active = True
self._wake()

async def run(self) -> None:
"""Drive serial worker calls, holding or discarding output at safe boundaries."""
workers = self._workers
if workers is None:
raise RuntimeError("load() must run first; override worker_setup(), not load()")
try:
while True:
await self._wait_ready()
epoch = self._epoch
params = _snapshot(self.session_params(), "session_params()")
await self._call(
workers.start_session,
params,
seed=self.session_seed,
timeout=self.command_timeout,
)
index = 0
while self._current(epoch):
await self._wait_ready(epoch)
if not self._current(epoch):
break
controls = _snapshot(self.controls(), "controls()")
started = time.perf_counter()
frames = await self._call(
workers.generate,
index,
controls,
timeout=self.command_timeout,
)
compute_time = time.perf_counter() - started
await self._wait_ready(epoch)
if not self._current(epoch):
break
output = self.to_output(frames)
if not isinstance(output, Output):
raise TypeError("to_output() must return an Output instance")
await self.emit(
output,
compute_time=compute_time if self.adaptive_fps else None,
)
index += 1
await self._call(workers.end_session, timeout=self.command_timeout)
finally:
try:
await self._call(workers.shutdown, timeout=self.shutdown_timeout)
finally:
if self._executor is not None:
self._executor.shutdown(wait=False)

def _current(self, epoch: int) -> bool:
return epoch == self._epoch and self._active and self.connected.is_set()

def _wake(self) -> None:
if self._wakeup is not None:
self._wakeup.set()

async def _wait_ready(self, epoch: int | None = None) -> None:
if self._wakeup is None:
raise RuntimeError("run() must be started by the runtime's model loop")
while True:
if epoch is not None and not self._current(epoch):
return
if self._active and self.connected.is_set() and not self._paused:
return
self._wakeup.clear()
await self._wakeup.wait()

async def _call[T](self, fn: Callable[..., T], *args: Any, **kwargs: Any) -> T:
assert self._executor is not None
future = asyncio.get_running_loop().run_in_executor(
self._executor,
partial(fn, *args, **kwargs),
)
try:
return await asyncio.shield(future)
except asyncio.CancelledError:
# Cancellation cannot stop a running GPU hook. Drain the command
# before shutdown releases its buffers, even in the inline backend.
try:
await future
finally:
raise


def _snapshot(value: dict[str, Any], hook: str) -> dict[str, Any]:
if not isinstance(value, dict):
raise TypeError(f"{hook} must return a dict, got {type(value).__name__}")
try:
result = copy.deepcopy(value)
pickle.dumps(result)
except Exception as exc:
raise TypeError(
f"{hook} must return picklable CPU data; pass paths and scalars, not live resources"
) from exc
return result
Loading