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
292 changes: 267 additions & 25 deletions cosmos_framework/scripts/action_policy_server_robolab.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,15 +28,17 @@
import json
import os
import socket
import time
import threading
import time
from dataclasses import dataclass
from datetime import timedelta
from pathlib import Path
from typing import Any, Literal

import numpy as np
import pydantic
import torch
import torch.distributed as dist
import torch.nn.functional as F
import tyro

Expand All @@ -52,6 +54,7 @@
from cosmos_framework.inference.args import GuidanceInterval, OmniSetupArgs, OmniSetupOverrides
from cosmos_framework.inference.common.args import ConfigFileType, ConfigOverrides, tyro_cli
from cosmos_framework.inference.common.config import deserialize_config, deserialize_config_dict, load_config
from cosmos_framework.inference.common.inference import _download_on_rank0
from cosmos_framework.inference.common.init import init_output_dir
from cosmos_framework.inference.inference import OmniInference
from cosmos_framework.scripts.action_policy_server_utils import (
Expand All @@ -77,6 +80,7 @@
"sides, with the robot visible."
)
_DEFAULT_HF_REVISION = "main"
_CONTROL_GROUP_TIMEOUT = timedelta(hours=3)
_ROBOLAB_POLICY_HF_REPOSITORIES = {
"Cosmos3-Nano-Policy-DROID": "nvidia/Cosmos3-Nano-Policy-DROID",
"nvidia/Cosmos3-Nano-Policy-DROID": "nvidia/Cosmos3-Nano-Policy-DROID",
Expand Down Expand Up @@ -147,7 +151,11 @@ def _resolve_checkpoint_path(checkpoint_path: str, *, hf_revision: str) -> str:
f"[robolab-policy-server] downloading consolidated checkpoint from Hugging Face: "
f"repository={repository!r} revision={hf_revision!r}"
)
return CheckpointDirHf(repository=repository, revision=hf_revision).download()
return str(
_download_on_rank0(
CheckpointDirHf(repository=repository, revision=hf_revision).download,
)
)


def _validate_checkpoint(checkpoint_path: str, *, allow_dcp_checkpoint: bool) -> None:
Expand Down Expand Up @@ -176,8 +184,10 @@ def _validate_checkpoint(checkpoint_path: str, *, allow_dcp_checkpoint: bool) ->
has_config = (checkpoint_dir / "config.json").exists()
has_consolidated_safetensors = any(checkpoint_dir.glob("*.safetensors"))
has_diffusers_safetensors_index = (checkpoint_dir / "model.safetensors.index.json").exists()
if not checkpoint_dir.is_dir() or not has_config or not (
has_consolidated_safetensors or has_diffusers_safetensors_index
if (
not checkpoint_dir.is_dir()
or not has_config
or not (has_consolidated_safetensors or has_diffusers_safetensors_index)
):
raise ValueError(f"Invalid safetensors checkpoint directory: {checkpoint_dir}")

Expand Down Expand Up @@ -363,17 +373,58 @@ class RobolabServerArgs(pydantic.BaseModel):
format_prompt_as_json: bool | None = None
"""Serve prompts as structured JSON (matching training ``format_prompt_as_json``)."""

cfg_parallel: bool = False
"""Use exactly two ranks for CFG parallelism without FSDP or context parallelism."""


def _resolve_parallelism_overrides(*, cfg_parallel: bool, world_size: int, guidance: float) -> dict[str, int]:
if cfg_parallel:
if world_size != 2:
raise ValueError(f"--cfg-parallel requires exactly 2 ranks, got world size {world_size}")
if guidance == 1.0:
raise ValueError("--cfg-parallel requires --guidance to differ from 1.0 so the CFG branch is active")
# The latency preset derives cfgp_size=2, cp_size=1, and
# dp_replicate_size=2 for a two-rank launch. Only disable FSDP.
return {"dp_shard_size": 1}

if world_size != 1:
raise ValueError(
f"A {world_size}-rank RoboLab server requires --cfg-parallel; non-CFG multi-rank serving is unsupported"
)
return {"dp_shard_size": 1}


def _build_control_group() -> dist.ProcessGroup:
"""Build a long-lived Gloo control group separate from NCCL generation."""
if not (dist.is_available() and dist.is_initialized() and dist.get_world_size() == 2):
raise RuntimeError("The RoboLab control group requires an initialized two-rank process group")

# Both ranks create this group during synchronized startup. The long timeout
# allows rank 1 to wait for requests without occupying the NCCL model group.
return dist.new_group(
ranks=[0, 1],
backend="gloo",
timeout=_CONTROL_GROUP_TIMEOUT,
)


class RobolabPolicyService:
def __init__(self, args: RobolabServerArgs) -> None:
if not torch.cuda.is_available():
raise RuntimeError("CUDA is required for OmniMoTModel inference in this repo.")
maybe_init_distributed()
world_size = dist.get_world_size() if dist.is_available() and dist.is_initialized() else 1
parallelism_overrides = _resolve_parallelism_overrides(
cfg_parallel=args.cfg_parallel,
world_size=world_size,
guidance=float(args.guidance),
)

resolved_checkpoint_path = _resolve_checkpoint_path(args.checkpoint_path, hf_revision=args.hf_revision)
args = args.model_copy(update={"checkpoint_path": resolved_checkpoint_path})
_validate_checkpoint(args.checkpoint_path, allow_dcp_checkpoint=args.allow_dcp_checkpoint)
maybe_init_distributed()

setup_args = self._build_setup_args(args)
setup_args = self._build_setup_args(args, parallelism_overrides)
log.info(
f"[robolab-policy-server] loading model: checkpoint_path={setup_args.checkpoint_path!r} "
f"config_file={setup_args.config_file!r} experiment={setup_args.experiment!r}"
Expand Down Expand Up @@ -418,6 +469,8 @@ def __init__(self, args: RobolabServerArgs) -> None:

self._lock = threading.Lock()
self._rng = np.random.default_rng(self.cfg.seed)
self._control_group = _build_control_group() if self._distributed_enabled() else None
self._control_request_id = 0
log.info(
f"[robolab-policy-server] ready domain={self.cfg.domain_name!r} resolution={self.cfg.resolution!r} "
f"action_space={self.cfg.action_space} action_dim={self.cfg.action_dim} "
Expand All @@ -428,11 +481,16 @@ def __init__(self, args: RobolabServerArgs) -> None:
f"seed={self.cfg.seed} deterministic_seed={self.cfg.deterministic_seed}"
)

def _build_setup_args(self, args: RobolabServerArgs) -> OmniSetupArgs:
def _build_setup_args(self, args: RobolabServerArgs, parallelism_overrides: dict[str, int]) -> OmniSetupArgs:
# RoboLab calls the model directly and never runs OmniInference's output
# guardrails. Avoid constructing and downloading unused guardrail models
# independently on every distributed rank.
setup_overrides: dict[str, Any] = {
"checkpoint_path": args.checkpoint_path,
"output_dir": args.output_dir or _DEFAULT_ROBOLAB_OUTPUT_DIR,
"sampler": args.sampler,
"guardrails": False,
**parallelism_overrides,
}
if args.experiment is not None:
setup_overrides["experiment"] = args.experiment
Expand Down Expand Up @@ -580,26 +638,29 @@ def _build_sample(self, obs: dict[str, Any]) -> dict[str, Any]:
sample["ai_caption"] = json.dumps(sample["ai_caption"])
return sample

def infer(self, obs: dict[str, Any]) -> dict[str, Any]:
start_time = time.monotonic()
sample = self._build_sample(obs)
def _generate(self, sample: dict[str, Any], seed: int) -> dict[str, Any]:
data_batch = _build_data_batch_from_sample(sample)
seed = self._next_seed()
log.info(f"[robolab-policy-server] prompt={data_batch['ai_caption'][0]!r} seed={seed}")

with self._lock:
with torch.inference_mode():
samples = self.model.generate_samples_from_batch(
data_batch,
guidance=self.cfg.guidance,
guidance_interval=(
list(self.cfg.guidance_interval) if self.cfg.guidance_interval is not None else None
),
seed=[seed],
num_steps=self.cfg.num_steps,
shift=self.cfg.shift,
)
with torch.inference_mode():
return self.model.generate_samples_from_batch(
data_batch,
guidance=self.cfg.guidance,
guidance_interval=(
list(self.cfg.guidance_interval) if self.cfg.guidance_interval is not None else None
),
seed=[seed],
num_steps=self.cfg.num_steps,
shift=self.cfg.shift,
)

def _format_outputs(
self,
obs: dict[str, Any],
samples: dict[str, Any],
*,
start_time: float,
) -> dict[str, Any]:
action = samples["action"][0][:, : self.cfg.action_dim] # [T,D]
action = action[self.cfg.history_length :] # [T2,D]
action_np = action.detach().cpu().numpy() # [T2,D]
Expand Down Expand Up @@ -634,16 +695,197 @@ def infer(self, obs: dict[str, Any]) -> dict[str, Any]:
print(f"infer_ms: {infer_ms:.1f}")
return outputs

@staticmethod
def _distributed_enabled() -> bool:
return dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1

def _broadcast_control_message(self, message: dict[str, Any] | None, *, src: int) -> dict[str, Any]:
if self._control_group is None:
raise RuntimeError("Distributed control group is not initialized")
payload: list[Any] = [message]
dist.broadcast_object_list(payload, src=src, group=self._control_group)
received = payload[0]
if not isinstance(received, dict):
raise TypeError(f"Expected a distributed control dict, got {type(received).__name__}")
return received

def _send_control_request(self, request_id: int, request: dict[str, Any]) -> None:
self._broadcast_control_message({"request_id": request_id, **request}, src=0)

def _receive_control_request(self, request_id: int) -> dict[str, Any]:
request = self._broadcast_control_message(None, src=0)
if request.get("request_id") != request_id:
raise RuntimeError(f"Expected distributed request {request_id}, got {request.get('request_id')!r}")
return request

def _send_worker_ready(self, request_id: int, *, error: str | None = None) -> None:
self._broadcast_control_message({"request_id": request_id, "error": error}, src=1)

def _wait_worker_ready(self, request_id: int) -> None:
response = self._broadcast_control_message(None, src=1)
if response.get("request_id") != request_id:
raise RuntimeError(
f"Expected readiness for distributed request {request_id}, got {response.get('request_id')!r}"
)
error = response.get("error")
if error is not None:
raise RuntimeError(f"Distributed worker rejected request {request_id}: {error}")

def _exchange_preparation_status(self, request_id: int, *, error: str | None) -> dict[int, str]:
if self._control_group is None:
raise RuntimeError("Distributed control group is not initialized")

status = {"request_id": request_id, "rank": dist.get_rank(), "error": error}
gathered_statuses: list[Any] = [None] * dist.get_world_size(group=self._control_group)
dist.all_gather_object(gathered_statuses, status, group=self._control_group)

errors: dict[int, str] = {}
for gathered_status in gathered_statuses:
if not isinstance(gathered_status, dict):
raise TypeError(f"Expected a distributed preparation status dict, got {type(gathered_status).__name__}")
if gathered_status.get("request_id") != request_id:
raise RuntimeError(
f"Expected preparation status for request {request_id}, got {gathered_status.get('request_id')!r}"
)
rank = gathered_status.get("rank")
if not isinstance(rank, int):
raise TypeError(f"Expected preparation status rank to be an int, got {type(rank).__name__}")
gathered_error = gathered_status.get("error")
if gathered_error is not None:
if not isinstance(gathered_error, str):
raise TypeError(
f"Expected preparation status error to be a string, got {type(gathered_error).__name__}"
)
errors[rank] = gathered_error
return errors

def infer(self, obs: dict[str, Any]) -> dict[str, Any]:
# Serialize request dispatch and CFGP generation so every rank enters
# the model collectives in the same order.
with self._lock:
start_time = time.monotonic()
seed = self._next_seed()
request_id: int | None = None
if self._distributed_enabled():
request_id = self._control_request_id
self._send_control_request(request_id, {"kind": "infer", "obs": obs, "seed": seed})
self._control_request_id += 1

sample: dict[str, Any] | None = None
preparation_exception: Exception | None = None
try:
sample = self._build_sample(obs)
except Exception as exc:
preparation_exception = exc

if request_id is not None:
error = None
if preparation_exception is not None:
error = f"{type(preparation_exception).__name__}: {preparation_exception}"
preparation_errors = self._exchange_preparation_status(request_id, error=error)
if preparation_errors:
if preparation_exception is not None:
raise preparation_exception
raise RuntimeError(f"Distributed request {request_id} preparation failed: {preparation_errors}")
elif preparation_exception is not None:
raise preparation_exception

assert sample is not None
samples = self._generate(sample, seed)
return self._format_outputs(obs, samples, start_time=start_time)

def worker_loop(self) -> None:
"""Run distributed inference work on non-server ranks."""
if not self._distributed_enabled():
raise RuntimeError("worker_loop requires an initialized multi-rank process group")

rank = dist.get_rank()
log.info(
f"[robolab-policy-server] rank {rank} ready as a distributed inference worker",
rank0_only=False,
)
while True:
request_id = self._control_request_id
request = self._receive_control_request(request_id)
self._control_request_id += 1
kind = request.get("kind")
if kind == "shutdown":
self._send_worker_ready(request_id)
log.info(
f"[robolab-policy-server] rank {rank} received shutdown",
rank0_only=False,
)
return
if kind != "infer":
error = f"Unsupported distributed request kind: {kind!r}"
self._send_worker_ready(request_id, error=error)
log.error(f"[robolab-policy-server] rank {rank}: {error}", rank0_only=False)
raise RuntimeError(error)

obs = request.get("obs")
seed = request.get("seed")
sample: dict[str, Any] | None = None
preparation_exception: Exception | None = None
try:
if not isinstance(obs, dict):
raise TypeError(f"Distributed request 'obs' must be a dict, got {type(obs).__name__}")
if not isinstance(seed, int):
raise TypeError(f"Distributed request 'seed' must be an int, got {type(seed).__name__}")
sample = self._build_sample(obs)
except Exception as exc:
preparation_exception = exc

error = None
if preparation_exception is not None:
error = f"{type(preparation_exception).__name__}: {preparation_exception}"
preparation_errors = self._exchange_preparation_status(request_id, error=error)
if preparation_errors:
log.error(
f"[robolab-policy-server] rank {rank} rejected distributed request {request_id}: "
f"{preparation_errors}",
rank0_only=False,
)
continue

# Only CFGP generation runs on this rank. Rank 0 formats and returns
# the response to RoboLab.
assert sample is not None
assert isinstance(seed, int)
self._generate(sample, seed)

def shutdown_worker(self) -> None:
"""Ask the non-server rank to leave its control loop cleanly."""
if not self._distributed_enabled() or dist.get_rank() != 0:
return
with self._lock:
request_id = self._control_request_id
self._send_control_request(request_id, {"kind": "shutdown"})
self._control_request_id += 1
self._wait_worker_ready(request_id)


def serve(args: RobolabServerArgs) -> None:
hostname = socket.gethostname()
log.info(f"[robolab-policy-server] starting host={hostname} bind={args.host}:{int(args.port)}")
service = RobolabPolicyService(args)

# Only global rank 0 owns the public socket. All other ranks stay alive and
# enter the worker loop so every request executes on the full CFGP group.
if service._distributed_enabled() and dist.get_rank() != 0:
service.worker_loop()
return

local_ip = get_local_ip()
log.info(f"[robolab-policy-server] Server accessible at: ws://{local_ip}:{int(args.port)}/")
log.info(f"[robolab-policy-server] Health check: http://{local_ip}:{int(args.port)}/healthz")
server_cls = _load_openpi_websocket_policy_server()
server_cls(policy=service, host=args.host, port=int(args.port), metadata={}).serve_forever()
try:
server_cls = _load_openpi_websocket_policy_server()
server_cls(policy=service, host=args.host, port=int(args.port), metadata={}).serve_forever()
finally:
try:
service.shutdown_worker()
except Exception as exc:
log.warning(f"[robolab-policy-server] failed to stop distributed worker cleanly: {exc}")


def main() -> None:
Expand Down
Loading
Loading