From cf0b13fb100966f54425daec1d184fa632eb9077 Mon Sep 17 00:00:00 2001 From: Vivek Goel Date: Tue, 8 Sep 2026 15:51:19 +0530 Subject: [PATCH 1/4] Enable two-rank CFG-parallel RoboLab serving Signed-off-by: Vivek Goel --- .../scripts/action_policy_server_robolab.py | 112 +++++++++++++++--- .../action_policy_server_robolab_test.py | 27 +++++ 2 files changed, 124 insertions(+), 15 deletions(-) diff --git a/cosmos_framework/scripts/action_policy_server_robolab.py b/cosmos_framework/scripts/action_policy_server_robolab.py index 71bb2c80f..aefacf883 100644 --- a/cosmos_framework/scripts/action_policy_server_robolab.py +++ b/cosmos_framework/scripts/action_policy_server_robolab.py @@ -28,8 +28,8 @@ import json import os import socket -import time import threading +import time from dataclasses import dataclass from pathlib import Path from typing import Any, Literal @@ -37,6 +37,7 @@ import numpy as np import pydantic import torch +import torch.distributed as dist import torch.nn.functional as F import tyro @@ -363,6 +364,22 @@ 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) -> dict[str, int]: + if cfg_parallel: + if world_size != 2: + raise ValueError(f"--cfg-parallel requires exactly 2 ranks, got world size {world_size}") + return {"dp_shard_size": 1, "cfgp_size": 2, "cp_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, "cfgp_size": 1, "cp_size": 1} + class RobolabPolicyService: def __init__(self, args: RobolabServerArgs) -> None: @@ -429,10 +446,21 @@ def __init__(self, args: RobolabServerArgs) -> None: ) def _build_setup_args(self, args: RobolabServerArgs) -> OmniSetupArgs: + 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, + ) + + # 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 @@ -580,25 +608,23 @@ 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]: + def _infer_impl(self, obs: dict[str, Any], seed: int) -> dict[str, Any]: start_time = time.monotonic() sample = self._build_sample(obs) 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(): + 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, + ) action = samples["action"][0][:, : self.cfg.action_dim] # [T,D] action = action[self.cfg.history_length :] # [T2,D] @@ -634,11 +660,67 @@ 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 + + @staticmethod + def _broadcast_request(request: dict[str, Any] | None) -> dict[str, Any]: + payload: list[Any] = [request] + dist.broadcast_object_list( + payload, + src=0, + device=torch.device("cuda", torch.cuda.current_device()), + ) + received = payload[0] + if not isinstance(received, dict): + raise TypeError(f"Expected a distributed request dict, got {type(received).__name__}") + return received + + def infer(self, obs: dict[str, Any]) -> dict[str, Any]: + # The WebSocket server may dispatch requests from multiple threads. Keep + # the broadcast and distributed model call in one globally ordered + # critical section so every rank executes collectives in the same order. + with self._lock: + seed = self._next_seed() + if self._distributed_enabled(): + self._broadcast_request({"obs": obs, "seed": seed}) + return self._infer_impl(obs, seed) + + 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 = self._broadcast_request(None) + obs = request.get("obs") + seed = request.get("seed") + 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__}") + # The result is intentionally discarded. This rank participates in + # the CFGP/FSDP collectives; rank 0 returns the response to RoboLab. + self._infer_impl(obs, seed) + 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") diff --git a/cosmos_framework/scripts/action_policy_server_robolab_test.py b/cosmos_framework/scripts/action_policy_server_robolab_test.py index f2b66aca5..7a100ad24 100644 --- a/cosmos_framework/scripts/action_policy_server_robolab_test.py +++ b/cosmos_framework/scripts/action_policy_server_robolab_test.py @@ -95,6 +95,7 @@ def test_server_args_default_to_released_droid_serving_config() -> None: assert args.num_steps == 4 assert args.shift == 5.0 assert args.deterministic_seed is False + assert args.cfg_parallel is False def test_server_args_accept_guidance_interval() -> None: @@ -103,6 +104,32 @@ def test_server_args_accept_guidance_interval() -> None: assert args.guidance_interval == (960.0, 1001.0) +@pytest.mark.parametrize( + ("cfg_parallel", "world_size", "expected"), + [ + (False, 1, {"dp_shard_size": 1, "cfgp_size": 1, "cp_size": 1}), + (True, 2, {"dp_shard_size": 1, "cfgp_size": 2, "cp_size": 1}), + ], +) +def test_resolve_parallelism_overrides( + cfg_parallel: bool, + world_size: int, + expected: dict[str, int], +) -> None: + assert ( + robolab_server._resolve_parallelism_overrides(cfg_parallel=cfg_parallel, world_size=world_size) == expected + ) + + +@pytest.mark.parametrize(("cfg_parallel", "world_size"), [(False, 2), (True, 1), (True, 4)]) +def test_resolve_parallelism_overrides_rejects_unsupported_launches( + cfg_parallel: bool, + world_size: int, +) -> None: + with pytest.raises(ValueError): + robolab_server._resolve_parallelism_overrides(cfg_parallel=cfg_parallel, world_size=world_size) + + def test_joint_pos_observation_preprocessing_matches_internal_layout() -> None: service = object.__new__(robolab_server.RobolabPolicyService) service.cfg = robolab_server.RobolabPolicyConfig( From f35d206220947a32381c89064a9766d03de72c8a Mon Sep 17 00:00:00 2001 From: Vivek Goel Date: Fri, 18 Sep 2026 16:32:50 +0530 Subject: [PATCH 2/4] Fix two-rank RoboLab control communication --- .../scripts/action_policy_server_robolab.py | 191 ++++++++++++++---- .../action_policy_server_robolab_test.py | 164 ++++++++++++++- 2 files changed, 306 insertions(+), 49 deletions(-) diff --git a/cosmos_framework/scripts/action_policy_server_robolab.py b/cosmos_framework/scripts/action_policy_server_robolab.py index aefacf883..4c8474f1c 100644 --- a/cosmos_framework/scripts/action_policy_server_robolab.py +++ b/cosmos_framework/scripts/action_policy_server_robolab.py @@ -31,6 +31,7 @@ import threading import time from dataclasses import dataclass +from datetime import timedelta from pathlib import Path from typing import Any, Literal @@ -53,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 ( @@ -78,6 +80,7 @@ "sides, with the robot visible." ) _DEFAULT_HF_REVISION = "main" +_CONTROL_GROUP_TIMEOUT = timedelta(days=365) _ROBOLAB_POLICY_HF_REPOSITORIES = { "Cosmos3-Nano-Policy-DROID": "nvidia/Cosmos3-Nano-Policy-DROID", "nvidia/Cosmos3-Nano-Policy-DROID": "nvidia/Cosmos3-Nano-Policy-DROID", @@ -148,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: @@ -177,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}") @@ -368,29 +377,54 @@ class RobolabServerArgs(pydantic.BaseModel): """Use exactly two ranks for CFG parallelism without FSDP or context parallelism.""" -def _resolve_parallelism_overrides(*, cfg_parallel: bool, world_size: int) -> dict[str, int]: +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}") - return {"dp_shard_size": 1, "cfgp_size": 2, "cp_size": 1} + 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, "cfgp_size": 1, "cp_size": 1} + 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}" @@ -435,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} " @@ -445,13 +481,7 @@ 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: - 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, - ) - + 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. @@ -608,14 +638,12 @@ def _build_sample(self, obs: dict[str, Any]) -> dict[str, Any]: sample["ai_caption"] = json.dumps(sample["ai_caption"]) return sample - def _infer_impl(self, obs: dict[str, Any], seed: int) -> 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) log.info(f"[robolab-policy-server] prompt={data_batch['ai_caption'][0]!r} seed={seed}") with torch.inference_mode(): - samples = self.model.generate_samples_from_batch( + return self.model.generate_samples_from_batch( data_batch, guidance=self.cfg.guidance, guidance_interval=( @@ -626,6 +654,13 @@ def _infer_impl(self, obs: dict[str, Any], seed: int) -> dict[str, Any]: 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] @@ -664,28 +699,54 @@ def _infer_impl(self, obs: dict[str, Any], seed: int) -> dict[str, Any]: def _distributed_enabled() -> bool: return dist.is_available() and dist.is_initialized() and dist.get_world_size() > 1 - @staticmethod - def _broadcast_request(request: dict[str, Any] | None) -> dict[str, Any]: - payload: list[Any] = [request] - dist.broadcast_object_list( - payload, - src=0, - device=torch.device("cuda", torch.cuda.current_device()), - ) + 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 request dict, got {type(received).__name__}") + 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 infer(self, obs: dict[str, Any]) -> dict[str, Any]: - # The WebSocket server may dispatch requests from multiple threads. Keep - # the broadcast and distributed model call in one globally ordered - # critical section so every rank executes collectives in the same order. + # Serialize request dispatch and CFGP generation so every rank enters + # the model collectives in the same order. with self._lock: + start_time = time.monotonic() + # Reject malformed client input before dispatching anything to the + # worker. The worker therefore remains ready for the next request. + sample = self._build_sample(obs) seed = self._next_seed() if self._distributed_enabled(): - self._broadcast_request({"obs": obs, "seed": seed}) - return self._infer_impl(obs, seed) + request_id = self._control_request_id + self._send_control_request(request_id, {"kind": "infer", "obs": obs, "seed": seed}) + self._control_request_id += 1 + self._wait_worker_ready(request_id) + 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.""" @@ -698,16 +759,54 @@ def worker_loop(self) -> None: rank0_only=False, ) while True: - request = self._broadcast_request(None) + 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) + continue + obs = request.get("obs") seed = request.get("seed") - 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__}") - # The result is intentionally discarded. This rank participates in - # the CFGP/FSDP collectives; rank 0 returns the response to RoboLab. - self._infer_impl(obs, seed) + 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: + error = f"{type(exc).__name__}: {exc}" + self._send_worker_ready(request_id, error=error) + log.exception( + f"[robolab-policy-server] rank {rank} rejected distributed request {request_id}", + rank0_only=False, + ) + continue + + self._send_worker_ready(request_id) + # Only CFGP generation runs on this rank. Rank 0 formats and returns + # the response to RoboLab. + 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: @@ -724,8 +823,14 @@ def serve(args: RobolabServerArgs) -> None: 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: diff --git a/cosmos_framework/scripts/action_policy_server_robolab_test.py b/cosmos_framework/scripts/action_policy_server_robolab_test.py index 7a100ad24..237caecaa 100644 --- a/cosmos_framework/scripts/action_policy_server_robolab_test.py +++ b/cosmos_framework/scripts/action_policy_server_robolab_test.py @@ -6,9 +6,10 @@ from __future__ import annotations import sys +import threading from pathlib import Path from typing import Any -from unittest.mock import patch +from unittest.mock import Mock, patch import numpy as np import pytest @@ -35,12 +36,20 @@ def fake_download(checkpoint: Any) -> str: calls.append((checkpoint.repository, checkpoint.revision)) return str(downloaded_path) + rank0_downloads: list[Any] = [] + + def fake_download_on_rank0(download: Any) -> Path: + rank0_downloads.append(download) + return Path(download()) + monkeypatch.setattr(robolab_server.CheckpointDirHf, "download", fake_download) + monkeypatch.setattr(robolab_server, "_download_on_rank0", fake_download_on_rank0) resolved = robolab_server._resolve_checkpoint_path("Cosmos3-Nano-Policy-DROID", hf_revision="test-revision") assert resolved == str(downloaded_path) assert calls == [("nvidia/Cosmos3-Nano-Policy-DROID", "test-revision")] + assert len(rank0_downloads) == 1 def test_resolve_checkpoint_keeps_existing_local_path(tmp_path: Path) -> None: @@ -105,19 +114,25 @@ def test_server_args_accept_guidance_interval() -> None: @pytest.mark.parametrize( - ("cfg_parallel", "world_size", "expected"), + ("cfg_parallel", "world_size", "guidance", "expected"), [ - (False, 1, {"dp_shard_size": 1, "cfgp_size": 1, "cp_size": 1}), - (True, 2, {"dp_shard_size": 1, "cfgp_size": 2, "cp_size": 1}), + (False, 1, 1.0, {"dp_shard_size": 1}), + (True, 2, 3.0, {"dp_shard_size": 1}), ], ) def test_resolve_parallelism_overrides( cfg_parallel: bool, world_size: int, + guidance: float, expected: dict[str, int], ) -> None: assert ( - robolab_server._resolve_parallelism_overrides(cfg_parallel=cfg_parallel, world_size=world_size) == expected + robolab_server._resolve_parallelism_overrides( + cfg_parallel=cfg_parallel, + world_size=world_size, + guidance=guidance, + ) + == expected ) @@ -127,7 +142,144 @@ def test_resolve_parallelism_overrides_rejects_unsupported_launches( world_size: int, ) -> None: with pytest.raises(ValueError): - robolab_server._resolve_parallelism_overrides(cfg_parallel=cfg_parallel, world_size=world_size) + robolab_server._resolve_parallelism_overrides( + cfg_parallel=cfg_parallel, + world_size=world_size, + guidance=3.0, + ) + + +def test_resolve_parallelism_overrides_rejects_cfg_parallel_without_cfg() -> None: + with pytest.raises(ValueError, match="guidance"): + robolab_server._resolve_parallelism_overrides( + cfg_parallel=True, + world_size=2, + guidance=1.0, + ) + + +def test_build_control_group_uses_gloo_with_long_idle_timeout() -> None: + control_group = Mock() + with ( + patch.object(robolab_server.dist, "is_available", return_value=True), + patch.object(robolab_server.dist, "is_initialized", return_value=True), + patch.object(robolab_server.dist, "get_world_size", return_value=2), + patch.object(robolab_server.dist, "new_group", return_value=control_group) as new_group, + ): + assert robolab_server._build_control_group() is control_group + + new_group.assert_called_once_with( + ranks=[0, 1], + backend="gloo", + timeout=robolab_server._CONTROL_GROUP_TIMEOUT, + ) + + +def test_control_messages_use_the_gloo_group() -> None: + service = object.__new__(robolab_server.RobolabPolicyService) + service._control_group = Mock() + message = {"request_id": 3, "kind": "infer"} + + with patch.object(robolab_server.dist, "broadcast_object_list") as broadcast: + assert service._broadcast_control_message(message, src=0) is message + + broadcast.assert_called_once_with([message], src=0, group=service._control_group) + + +def test_invalid_launch_is_rejected_before_checkpoint_resolution() -> None: + args = robolab_server.RobolabServerArgs(cfg_parallel=True) + with ( + patch.object(robolab_server.torch.cuda, "is_available", return_value=True), + patch.object(robolab_server, "maybe_init_distributed"), + patch.object(robolab_server.dist, "is_available", return_value=True), + patch.object(robolab_server.dist, "is_initialized", return_value=True), + patch.object(robolab_server.dist, "get_world_size", return_value=4), + patch.object(robolab_server, "_resolve_checkpoint_path") as resolve_checkpoint, + pytest.raises(ValueError, match="exactly 2 ranks"), + ): + robolab_server.RobolabPolicyService(args) + + resolve_checkpoint.assert_not_called() + + +def test_infer_rejects_invalid_observation_before_distributed_dispatch() -> None: + service = object.__new__(robolab_server.RobolabPolicyService) + service._lock = threading.Lock() + service._build_sample = Mock(side_effect=ValueError("bad observation")) + service._next_seed = Mock() + service._send_control_request = Mock() + + with pytest.raises(ValueError, match="bad observation"): + service.infer({"prompt": "missing image and state"}) + + service._next_seed.assert_not_called() + service._send_control_request.assert_not_called() + + +def test_infer_waits_for_worker_then_formats_rank0_output() -> None: + service = object.__new__(robolab_server.RobolabPolicyService) + service._lock = threading.Lock() + service._control_request_id = 0 + service._build_sample = Mock(return_value={"sample": "prepared"}) + service._next_seed = Mock(return_value=17) + service._distributed_enabled = Mock(return_value=True) + service._send_control_request = Mock() + service._wait_worker_ready = Mock() + service._generate = Mock(return_value={"samples": "generated"}) + service._format_outputs = Mock(return_value={"action": "formatted"}) + obs = {"prompt": "move"} + + assert service.infer(obs) == {"action": "formatted"} + + service._send_control_request.assert_called_once_with( + 0, + {"kind": "infer", "obs": obs, "seed": 17}, + ) + service._wait_worker_ready.assert_called_once_with(0) + service._generate.assert_called_once_with({"sample": "prepared"}, 17) + service._format_outputs.assert_called_once() + + +def test_worker_reports_preparation_error_then_processes_next_request() -> None: + service = object.__new__(robolab_server.RobolabPolicyService) + service._control_request_id = 0 + service._distributed_enabled = Mock(return_value=True) + service._receive_control_request = Mock( + side_effect=[ + {"kind": "infer", "obs": {"bad": True}, "seed": 7}, + {"kind": "infer", "obs": {"good": True}, "seed": 8}, + {"kind": "shutdown"}, + ] + ) + service._build_sample = Mock(side_effect=[ValueError("bad sample"), {"sample": "valid"}]) + service._generate = Mock(return_value={}) + service._send_worker_ready = Mock() + + with patch.object(robolab_server.dist, "get_rank", return_value=1): + service.worker_loop() + + assert service._send_worker_ready.call_args_list[0].args == (0,) + assert service._send_worker_ready.call_args_list[0].kwargs == {"error": "ValueError: bad sample"} + assert service._send_worker_ready.call_args_list[1].args == (1,) + assert service._send_worker_ready.call_args_list[1].kwargs == {} + assert service._send_worker_ready.call_args_list[2].args == (2,) + assert service._send_worker_ready.call_args_list[2].kwargs == {} + service._generate.assert_called_once_with({"sample": "valid"}, 8) + + +def test_shutdown_worker_uses_next_control_request_and_waits_for_ack() -> None: + service = object.__new__(robolab_server.RobolabPolicyService) + service._lock = threading.Lock() + service._control_request_id = 4 + service._distributed_enabled = Mock(return_value=True) + service._send_control_request = Mock() + service._wait_worker_ready = Mock() + + with patch.object(robolab_server.dist, "get_rank", return_value=0): + service.shutdown_worker() + + service._send_control_request.assert_called_once_with(4, {"kind": "shutdown"}) + service._wait_worker_ready.assert_called_once_with(4) def test_joint_pos_observation_preprocessing_matches_internal_layout() -> None: From b3bbf95e8ae88f6dce7e78acadba8e8327750d86 Mon Sep 17 00:00:00 2001 From: Vivek Goel Date: Tue, 22 Sep 2026 18:45:47 +0530 Subject: [PATCH 3/4] Optimize two-rank RoboLab request preparation --- .../scripts/action_policy_server_robolab.py | 77 +++++++++-- .../action_policy_server_robolab_test.py | 120 +++++++++++++++--- docs/action_policy_droid_server.md | 22 ++++ 3 files changed, 193 insertions(+), 26 deletions(-) diff --git a/cosmos_framework/scripts/action_policy_server_robolab.py b/cosmos_framework/scripts/action_policy_server_robolab.py index 4c8474f1c..b2a557c28 100644 --- a/cosmos_framework/scripts/action_policy_server_robolab.py +++ b/cosmos_framework/scripts/action_policy_server_robolab.py @@ -80,7 +80,7 @@ "sides, with the robot visible." ) _DEFAULT_HF_REVISION = "main" -_CONTROL_GROUP_TIMEOUT = timedelta(days=365) +_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", @@ -731,20 +731,66 @@ def _wait_worker_ready(self, request_id: int) -> None: 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() - # Reject malformed client input before dispatching anything to the - # worker. The worker therefore remains ready for the next request. - sample = self._build_sample(obs) 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 - self._wait_worker_ready(request_id) + + 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) @@ -774,10 +820,12 @@ def worker_loop(self) -> None: 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) - continue + 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__}") @@ -785,17 +833,24 @@ def worker_loop(self) -> None: raise TypeError(f"Distributed request 'seed' must be an int, got {type(seed).__name__}") sample = self._build_sample(obs) except Exception as exc: - error = f"{type(exc).__name__}: {exc}" - self._send_worker_ready(request_id, error=error) - log.exception( - f"[robolab-policy-server] rank {rank} rejected distributed request {request_id}", + 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 - self._send_worker_ready(request_id) # 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: diff --git a/cosmos_framework/scripts/action_policy_server_robolab_test.py b/cosmos_framework/scripts/action_policy_server_robolab_test.py index 237caecaa..e2f178604 100644 --- a/cosmos_framework/scripts/action_policy_server_robolab_test.py +++ b/cosmos_framework/scripts/action_policy_server_robolab_test.py @@ -158,7 +158,7 @@ def test_resolve_parallelism_overrides_rejects_cfg_parallel_without_cfg() -> Non ) -def test_build_control_group_uses_gloo_with_long_idle_timeout() -> None: +def test_build_control_group_uses_gloo_with_three_hour_timeout() -> None: control_group = Mock() with ( patch.object(robolab_server.dist, "is_available", return_value=True), @@ -173,6 +173,7 @@ def test_build_control_group_uses_gloo_with_long_idle_timeout() -> None: backend="gloo", timeout=robolab_server._CONTROL_GROUP_TIMEOUT, ) + assert robolab_server._CONTROL_GROUP_TIMEOUT.total_seconds() == 3 * 60 * 60 def test_control_messages_use_the_gloo_group() -> None: @@ -202,21 +203,47 @@ def test_invalid_launch_is_rejected_before_checkpoint_resolution() -> None: resolve_checkpoint.assert_not_called() -def test_infer_rejects_invalid_observation_before_distributed_dispatch() -> None: +def test_preparation_status_exchange_uses_the_gloo_group() -> None: + service = object.__new__(robolab_server.RobolabPolicyService) + service._control_group = Mock() + + def gather_statuses(output: list[Any], status: dict[str, Any], *, group: Any) -> None: + assert group is service._control_group + output[:] = [status, {"request_id": 3, "rank": 1, "error": "ValueError: worker failed"}] + + with ( + patch.object(robolab_server.dist, "get_rank", return_value=0), + patch.object(robolab_server.dist, "get_world_size", return_value=2), + patch.object(robolab_server.dist, "all_gather_object", side_effect=gather_statuses) as all_gather, + ): + errors = service._exchange_preparation_status(3, error=None) + + assert errors == {1: "ValueError: worker failed"} + all_gather.assert_called_once() + + +def test_infer_dispatches_invalid_observation_then_exchanges_error_status() -> None: service = object.__new__(robolab_server.RobolabPolicyService) service._lock = threading.Lock() + service._control_request_id = 0 service._build_sample = Mock(side_effect=ValueError("bad observation")) - service._next_seed = Mock() + service._next_seed = Mock(return_value=17) + service._distributed_enabled = Mock(return_value=True) service._send_control_request = Mock() + service._exchange_preparation_status = Mock( + return_value={0: "ValueError: bad observation", 1: "ValueError: bad observation"} + ) + obs = {"prompt": "missing image and state"} with pytest.raises(ValueError, match="bad observation"): - service.infer({"prompt": "missing image and state"}) + service.infer(obs) - service._next_seed.assert_not_called() - service._send_control_request.assert_not_called() + service._send_control_request.assert_called_once_with(0, {"kind": "infer", "obs": obs, "seed": 17}) + service._exchange_preparation_status.assert_called_once_with(0, error="ValueError: bad observation") + assert service._control_request_id == 1 -def test_infer_waits_for_worker_then_formats_rank0_output() -> None: +def test_infer_exchanges_preparation_status_then_formats_rank0_output() -> None: service = object.__new__(robolab_server.RobolabPolicyService) service._lock = threading.Lock() service._control_request_id = 0 @@ -224,7 +251,7 @@ def test_infer_waits_for_worker_then_formats_rank0_output() -> None: service._next_seed = Mock(return_value=17) service._distributed_enabled = Mock(return_value=True) service._send_control_request = Mock() - service._wait_worker_ready = Mock() + service._exchange_preparation_status = Mock(return_value={}) service._generate = Mock(return_value={"samples": "generated"}) service._format_outputs = Mock(return_value={"action": "formatted"}) obs = {"prompt": "move"} @@ -235,11 +262,51 @@ def test_infer_waits_for_worker_then_formats_rank0_output() -> None: 0, {"kind": "infer", "obs": obs, "seed": 17}, ) - service._wait_worker_ready.assert_called_once_with(0) + service._exchange_preparation_status.assert_called_once_with(0, error=None) service._generate.assert_called_once_with({"sample": "prepared"}, 17) service._format_outputs.assert_called_once() +def test_infer_skips_generation_when_worker_preparation_fails() -> None: + service = object.__new__(robolab_server.RobolabPolicyService) + service._lock = threading.Lock() + service._control_request_id = 0 + service._build_sample = Mock(return_value={"sample": "prepared"}) + service._next_seed = Mock(return_value=17) + service._distributed_enabled = Mock(return_value=True) + service._send_control_request = Mock() + service._exchange_preparation_status = Mock(return_value={1: "ValueError: worker failed"}) + service._generate = Mock() + + with pytest.raises(RuntimeError, match="worker failed"): + service.infer({"prompt": "move"}) + + service._generate.assert_not_called() + + +def test_worker_acknowledges_unsupported_request_then_raises() -> None: + service = object.__new__(robolab_server.RobolabPolicyService) + service._control_request_id = 0 + service._distributed_enabled = Mock(return_value=True) + service._receive_control_request = Mock(return_value={"kind": "unsupported"}) + service._send_worker_ready = Mock() + service._build_sample = Mock() + + with ( + patch.object(robolab_server.dist, "get_rank", return_value=1), + pytest.raises(RuntimeError, match="Unsupported distributed request kind: 'unsupported'"), + ): + service.worker_loop() + + service._send_worker_ready.assert_called_once_with( + 0, + error="Unsupported distributed request kind: 'unsupported'", + ) + service._receive_control_request.assert_called_once_with(0) + assert service._control_request_id == 1 + service._build_sample.assert_not_called() + + def test_worker_reports_preparation_error_then_processes_next_request() -> None: service = object.__new__(robolab_server.RobolabPolicyService) service._control_request_id = 0 @@ -254,19 +321,42 @@ def test_worker_reports_preparation_error_then_processes_next_request() -> None: service._build_sample = Mock(side_effect=[ValueError("bad sample"), {"sample": "valid"}]) service._generate = Mock(return_value={}) service._send_worker_ready = Mock() + service._exchange_preparation_status = Mock(side_effect=[{1: "ValueError: bad sample"}, {}]) with patch.object(robolab_server.dist, "get_rank", return_value=1): service.worker_loop() - assert service._send_worker_ready.call_args_list[0].args == (0,) - assert service._send_worker_ready.call_args_list[0].kwargs == {"error": "ValueError: bad sample"} - assert service._send_worker_ready.call_args_list[1].args == (1,) - assert service._send_worker_ready.call_args_list[1].kwargs == {} - assert service._send_worker_ready.call_args_list[2].args == (2,) - assert service._send_worker_ready.call_args_list[2].kwargs == {} + assert service._exchange_preparation_status.call_args_list[0].args == (0,) + assert service._exchange_preparation_status.call_args_list[0].kwargs == {"error": "ValueError: bad sample"} + assert service._exchange_preparation_status.call_args_list[1].args == (1,) + assert service._exchange_preparation_status.call_args_list[1].kwargs == {"error": None} + service._send_worker_ready.assert_called_once_with(2) service._generate.assert_called_once_with({"sample": "valid"}, 8) +def test_worker_skips_generation_when_rank0_preparation_fails() -> None: + service = object.__new__(robolab_server.RobolabPolicyService) + service._control_request_id = 0 + service._distributed_enabled = Mock(return_value=True) + service._receive_control_request = Mock( + side_effect=[ + {"kind": "infer", "obs": {"valid": True}, "seed": 7}, + {"kind": "shutdown"}, + ] + ) + service._build_sample = Mock(return_value={"sample": "valid"}) + service._generate = Mock() + service._send_worker_ready = Mock() + service._exchange_preparation_status = Mock(return_value={0: "ValueError: server failed"}) + + with patch.object(robolab_server.dist, "get_rank", return_value=1): + service.worker_loop() + + service._exchange_preparation_status.assert_called_once_with(0, error=None) + service._generate.assert_not_called() + service._send_worker_ready.assert_called_once_with(1) + + def test_shutdown_worker_uses_next_control_request_and_waits_for_ack() -> None: service = object.__new__(robolab_server.RobolabPolicyService) service._lock = threading.Lock() diff --git a/docs/action_policy_droid_server.md b/docs/action_policy_droid_server.md index fc1c3e915..a4c8a5da5 100644 --- a/docs/action_policy_droid_server.md +++ b/docs/action_policy_droid_server.md @@ -87,6 +87,28 @@ Inside the container, start the policy server: timesteps in the inclusive range `[960, 1001]`. Omit `--guidance-interval` to apply guidance at every denoising step. +### Two-rank CFG parallelism + +To serve with classifier-free guidance parallelized across two local GPUs, launch +exactly two processes and pass `--cfg-parallel`. Set `OMP_NUM_THREADS` to the +number of physical CPU cores available to the job divided by the two local +ranks. Without an explicit value, `torchrun` defaults each process to one OpenMP +thread, which can make request preprocessing slower. + +For example, on a host where the job has 32 physical CPU cores available: + +```bash +OMP_NUM_THREADS=16 torchrun --nproc-per-node=2 \ + -m cosmos_framework.scripts.action_policy_server_robolab \ + --cfg-parallel \ + --port 8000 +``` + +Use cores assigned to the job rather than the host-wide CPU count when running +inside a container, CPU set, or scheduler allocation. The ratio is a starting +point; benchmark a few nearby values if CPU preprocessing is important to +end-to-end latency. + ## Simulation Client Clone [`RoboLab`](https://github.com/NVlabs/RoboLab): From 849ed776dc990f3398b75d180f44e6a5571346c5 Mon Sep 17 00:00:00 2001 From: Vivek Goel Date: Tue, 22 Sep 2026 19:00:36 +0530 Subject: [PATCH 4/4] Fix generated policy server documentation TOC --- docs/action_policy_droid_server.md | 1 + 1 file changed, 1 insertion(+) diff --git a/docs/action_policy_droid_server.md b/docs/action_policy_droid_server.md index a4c8a5da5..9997b3315 100644 --- a/docs/action_policy_droid_server.md +++ b/docs/action_policy_droid_server.md @@ -14,6 +14,7 @@ ______________________________________________________________________ **Table of Contents** - [Policy Server](#policy-server) + - [Two-rank CFG parallelism](#two-rank-cfg-parallelism) - [Simulation Client](#simulation-client) ______________________________________________________________________