Skip to content
Merged
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
394 changes: 233 additions & 161 deletions .claude/skills/remote-training/SKILL.md

Large diffs are not rendered by default.

43 changes: 41 additions & 2 deletions positronic/offboard/vendor_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import asyncio
import logging
import time
from abc import ABC, abstractmethod
from collections.abc import Callable
from typing import Any
Expand Down Expand Up @@ -113,13 +114,24 @@ class VendorServer(ABC):
resolve_model(None) → create_policy → reset → warmup
"""

def __init__(self, codec: Codec | None, host: str = '0.0.0.0', port: int = 8000, recording_dir: str | None = None):
def __init__(
self,
codec: Codec | None,
host: str = '0.0.0.0',
port: int = 8000,
recording_dir: str | None = None,
idle_timeout_min: float | None = None,
):
self.codec = codec
self.host = host
self.port = port
if recording_dir:
self.codec = RecordingCodec(self.codec, pos3.sync(recording_dir))

self.idle_timeout_min = idle_timeout_min
self._active_sessions = 0
self._last_activity = time.monotonic()

self.metadata: dict[str, Any] = {}

self.app = FastAPI()
Expand Down Expand Up @@ -168,6 +180,8 @@ async def websocket_endpoint(self, websocket: WebSocket, model_id: str | None =
await websocket.accept()
logger.info(f'Connected to {websocket.client} requesting {model_id or "default"}')

self._active_sessions += 1
self._last_activity = time.monotonic()
model_handle = None
try:
model_handle, extra_meta = await self.resolve_model(model_id, websocket)
Expand All @@ -180,6 +194,7 @@ async def websocket_endpoint(self, websocket: WebSocket, model_id: str | None =
try:
while True:
message = await websocket.receive_bytes()
self._last_activity = time.monotonic()
try:
raw_obs = deserialise(message)
actions = policy.select_action(raw_obs)
Expand All @@ -198,6 +213,8 @@ async def websocket_endpoint(self, websocket: WebSocket, model_id: str | None =
except Exception:
logger.debug('Failed to send error to client', exc_info=True)
finally:
self._active_sessions = max(0, self._active_sessions - 1)
self._last_activity = time.monotonic()
if model_handle is not None:
await self.release_policy(model_handle)

Expand All @@ -207,11 +224,33 @@ async def _startup(self):
policy.reset()
await self.warmup(policy)

async def _idle_watchdog(self, server: uvicorn.Server):
timeout_s = self.idle_timeout_min * 60
poll = min(timeout_s, 30)
while not server.should_exit:
await asyncio.sleep(poll)
if self._active_sessions > 0:
continue
idle = time.monotonic() - self._last_activity
if idle >= timeout_s:
logger.warning(f'No activity for {idle:.0f}s (idle timeout {timeout_s:.0f}s); shutting down server')
server.should_exit = True
return

def serve(self):
async def _run():
await self._startup()
config = uvicorn.Config(self.app, host=self.host, port=self.port, log_level='info')
await uvicorn.Server(config).serve()
server = uvicorn.Server(config)
self._last_activity = time.monotonic()
watchdog = None
if self.idle_timeout_min and self.idle_timeout_min > 0:
watchdog = asyncio.create_task(self._idle_watchdog(server))
try:
await server.serve()
finally:
if watchdog is not None:
watchdog.cancel()

try:
asyncio.run(_run())
Expand Down
8 changes: 7 additions & 1 deletion positronic/vendors/dreamzero/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -202,8 +202,11 @@ def __init__(
roboarena_port: int = 1234,
enable_dit_cache: bool = True,
recording_dir: str | None = None,
idle_timeout_min: float | None = None,
):
super().__init__(codec=codec, host=host, port=port, recording_dir=recording_dir)
super().__init__(
codec=codec, host=host, port=port, recording_dir=recording_dir, idle_timeout_min=idle_timeout_min
)
self.model_path = model_path
self.dreamzero_venv = Path(dreamzero_venv)
self.backbone = backbone
Expand Down Expand Up @@ -272,6 +275,7 @@ def shutdown_model(self):
port=8000,
enable_dit_cache=True,
recording_dir=None,
idle_timeout_min=None,
)
def server(
codec: Codec | None,
Expand All @@ -282,6 +286,7 @@ def server(
port: int,
enable_dit_cache: bool,
recording_dir: str | None,
idle_timeout_min: float | None,
):
"""Starts the DreamZero inference server."""
with pos3.mirror():
Expand All @@ -294,6 +299,7 @@ def server(
port=port,
enable_dit_cache=enable_dit_cache,
recording_dir=recording_dir,
idle_timeout_min=idle_timeout_min,
).serve()


Expand Down
8 changes: 7 additions & 1 deletion positronic/vendors/gr00t/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -278,8 +278,11 @@ def __init__(
zmq_port: int = 5555,
recording_dir: str | None = None,
ready_timeout: float = 120.0,
idle_timeout_min: float | None = None,
):
super().__init__(codec=codec, host=host, port=port, recording_dir=recording_dir)
super().__init__(
codec=codec, host=host, port=port, recording_dir=recording_dir, idle_timeout_min=idle_timeout_min
)
self.checkpoints_dir = checkpoints_dir.rstrip('/')
self.checkpoint = checkpoint
self.modality_config = modality_config
Expand Down Expand Up @@ -376,6 +379,7 @@ def shutdown_model(self):
modality_config='ee',
recording_dir=None,
ready_timeout=120.0,
idle_timeout_min=None,
)
def server(
codec: Codec,
Expand All @@ -386,6 +390,7 @@ def server(
modality_config: str,
recording_dir: str | None,
ready_timeout: float,
idle_timeout_min: float | None,
):
"""Starts the GR00T inference server with encoding/decoding."""

Expand All @@ -399,6 +404,7 @@ def server(
port=port,
recording_dir=recording_dir,
ready_timeout=ready_timeout,
idle_timeout_min=idle_timeout_min,
).serve()


Expand Down
29 changes: 25 additions & 4 deletions positronic/vendors/lerobot/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -33,8 +33,11 @@ def __init__(
port: int = 8000,
device: str | None = None,
recording_dir: str | None = None,
idle_timeout_min: float | None = None,
):
super().__init__(codec=codec, host=host, port=port, recording_dir=recording_dir)
super().__init__(
codec=codec, host=host, port=port, recording_dir=recording_dir, idle_timeout_min=idle_timeout_min
)
self.checkpoints_dir = str(checkpoints_dir).rstrip('/') + '/checkpoints'
self.checkpoint = checkpoint
self.device = device or _detect_device()
Expand Down Expand Up @@ -73,10 +76,28 @@ async def release_policy(self, model_handle):
await self.policy_manager.release_session()


@cfn.config(codec=lerobot_codecs.ee, checkpoint=None, port=8000, host='0.0.0.0', recording_dir=None)
def main(checkpoints_dir: str, checkpoint: str | None, codec, port: int, host: str, recording_dir: str | None):
@cfn.config(
codec=lerobot_codecs.ee, checkpoint=None, port=8000, host='0.0.0.0', recording_dir=None, idle_timeout_min=None
)
def main(
checkpoints_dir: str,
checkpoint: str | None,
codec,
port: int,
host: str,
recording_dir: str | None,
idle_timeout_min: float | None,
):
checkpoints_dir = str(pos3.download(checkpoints_dir))
InferenceServer(codec, checkpoints_dir, checkpoint, host=host, port=port, recording_dir=recording_dir).serve()
InferenceServer(
codec,
checkpoints_dir,
checkpoint,
host=host,
port=port,
recording_dir=recording_dir,
idle_timeout_min=idle_timeout_min,
).serve()


phail = main.override(
Expand Down
35 changes: 29 additions & 6 deletions positronic/vendors/lerobot_0_3_3/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,8 +43,11 @@ def __init__(
metadata: dict[str, Any] | None = None,
device: str | None = None,
recording_dir: str | None = None,
idle_timeout_min: float | None = None,
):
super().__init__(codec=codec, host=host, port=port, recording_dir=recording_dir)
super().__init__(
codec=codec, host=host, port=port, recording_dir=recording_dir, idle_timeout_min=idle_timeout_min
)
self.policy_factory = policy_factory
self.checkpoints_dir = str(checkpoints_dir).rstrip('/') + '/checkpoints'
self.checkpoint = checkpoint
Expand Down Expand Up @@ -96,7 +99,15 @@ def act(checkpoint_path: str) -> PreTrainedPolicy:
return policy


@cfn.config(policy_factory=act, codec=lerobot_codecs.ee, checkpoint=None, port=8000, host='0.0.0.0', recording_dir=None)
@cfn.config(
policy_factory=act,
codec=lerobot_codecs.ee,
checkpoint=None,
port=8000,
host='0.0.0.0',
recording_dir=None,
idle_timeout_min=None,
)
def main(
policy_factory: Callable[[str], PreTrainedPolicy],
checkpoints_dir: str,
Expand All @@ -105,10 +116,18 @@ def main(
port: int,
host: str,
recording_dir: str | None,
idle_timeout_min: float | None,
):
checkpoints_dir = str(pos3.download(checkpoints_dir))
InferenceServer(
policy_factory, codec, checkpoints_dir, checkpoint, host=host, port=port, recording_dir=recording_dir
policy_factory,
codec,
checkpoints_dir,
checkpoint,
host=host,
port=port,
recording_dir=recording_dir,
idle_timeout_min=idle_timeout_min,
).serve()


Expand All @@ -123,12 +142,16 @@ def main(
_DEMO_CHECKPOINT = 's3://positronic-public/checkpoints/sim_stack_cubes/act/'


@cfn.config(policy_factory=act, codec=lerobot_codecs.ee, checkpoint=None, port=8000, host='0.0.0.0')
def demo(policy_factory, checkpoint, codec, port, host):
@cfn.config(
policy_factory=act, codec=lerobot_codecs.ee, checkpoint=None, port=8000, host='0.0.0.0', idle_timeout_min=None
)
def demo(policy_factory, checkpoint, codec, port, host, idle_timeout_min):
from positronic.cfg.ds import PUBLIC

checkpoints_dir = str(pos3.download(_DEMO_CHECKPOINT, profile=PUBLIC))
InferenceServer(policy_factory, codec, checkpoints_dir, checkpoint, host=host, port=port).serve()
InferenceServer(
policy_factory, codec, checkpoints_dir, checkpoint, host=host, port=port, idle_timeout_min=idle_timeout_min
).serve()


if __name__ == '__main__':
Expand Down
8 changes: 7 additions & 1 deletion positronic/vendors/openpi/server.py
Original file line number Diff line number Diff line change
Expand Up @@ -190,8 +190,11 @@ def __init__(
openpi_ws_port: int = 8001,
metadata: dict[str, Any] | None = None,
recording_dir: str | None = None,
idle_timeout_min: float | None = None,
):
super().__init__(codec=codec, host=host, port=port, recording_dir=recording_dir)
super().__init__(
codec=codec, host=host, port=port, recording_dir=recording_dir, idle_timeout_min=idle_timeout_min
)
self.checkpoints_dir = str(checkpoints_dir).rstrip('/')
self.config_name = config_name
self.checkpoint = checkpoint
Expand Down Expand Up @@ -280,6 +283,7 @@ def shutdown_model(self):
port=8000,
openpi_ws_port=8001,
recording_dir=None,
idle_timeout_min=None,
)
def server(
codec,
Expand All @@ -290,6 +294,7 @@ def server(
port: int,
openpi_ws_port: int,
recording_dir: str | None,
idle_timeout_min: float | None,
):
"""OpenPI inference server.

Expand All @@ -316,6 +321,7 @@ def server(
port=port,
openpi_ws_port=openpi_ws_port,
recording_dir=recording_dir,
idle_timeout_min=idle_timeout_min,
).serve()


Expand Down
Loading
Loading