From fa8d131fe27bc7e402b3b66db3cdcd50f76ffa2a Mon Sep 17 00:00:00 2001 From: Vincent Moens Date: Wed, 24 Jun 2026 11:57:25 -0700 Subject: [PATCH] Update [ghstack-poisoned] --- sota-implementations/vla_grpo/README.md | 26 ++ .../vla_grpo/config/vla_grpo_libero.yaml | 17 + .../vla_grpo/config/vla_grpo_toy.yaml | 18 + sota-implementations/vla_grpo/test_openvla.py | 181 ++++++++ sota-implementations/vla_grpo/utils.py | 410 +++++++++++++++++- sota-implementations/vla_grpo/vla-grpo.py | 47 +- test/test_custom_envs.py | 5 +- torchrl/envs/custom/vla.py | 9 +- 8 files changed, 680 insertions(+), 33 deletions(-) diff --git a/sota-implementations/vla_grpo/README.md b/sota-implementations/vla_grpo/README.md index bda1af31c4e..bc06d071f2e 100644 --- a/sota-implementations/vla_grpo/README.md +++ b/sota-implementations/vla_grpo/README.md @@ -204,6 +204,32 @@ throughput split into collection and optimization: - `throughput/train_decisions_per_s` - `throughput/optim_steps_per_s` +The collector can also be switched between synchronous and async execution +paths for throughput experiments: + +```bash +# fully synchronous baseline +python sota-implementations/vla_grpo/vla-grpo.py \ + collector.async_env=false collector.async_policy=false + +# asynchronous env slots, but no policy auto-batching +python sota-implementations/vla_grpo/vla-grpo.py \ + collector.async_env=true collector.async_policy=false env.num_envs=8 + +# asynchronous env slots plus auto-batched policy inference +python sota-implementations/vla_grpo/vla-grpo.py \ + collector.async_env=true collector.async_policy=true env.num_envs=8 \ + collector.server_max_batch_size=8 collector.server_timeout=0.01 +``` + +`collector.async_env=true` uses `AsyncBatchedCollector` so faster environment +slots do not wait at a global step barrier. `collector.async_policy=true` routes +policy calls through an inference server; with multiple async env slots this +enables auto-batching and logs `policy_server/*` counters such as average batch +size, request rate, and queue/forward latency. The `false/true` combination is +available as a policy-server plumbing ablation, but policy auto-batching is most +meaningful when several env slots submit requests concurrently. + Eval rollouts can also be rendered to video (`logger.record_video=true`, on by default). A dedicated single-environment recorder is built with `from_pixels=True`: `ToyVLAEnv` renders the tracking scene, while `LiberoEnv` diff --git a/sota-implementations/vla_grpo/config/vla_grpo_libero.yaml b/sota-implementations/vla_grpo/config/vla_grpo_libero.yaml index d9c90e8881e..35ab8586f0b 100644 --- a/sota-implementations/vla_grpo/config/vla_grpo_libero.yaml +++ b/sota-implementations/vla_grpo/config/vla_grpo_libero.yaml @@ -91,6 +91,23 @@ collector: min_replay_decisions: null # null/0 = collect one target group wave total_iters: 100 # the paper's total_epochs policy_device: null # null = policy.device; set e.g. cuda:1 for rollout inference + # Execution-mode switches for throughput ablations: + # - false/false: regular synchronous TorchRL Collector. + # - true/false: async env slots with one request per policy forward. + # - true/true: async env slots plus auto-batched policy inference. + # - false/true: sync env stepping through the policy server path. + async_env: false + async_policy: false + env_backend: threading # AsyncBatchedCollector env backend: threading | multiprocessing + policy_backend: threading # inference transport: threading | multiprocessing | ray | monarch + server_backend: thread # process server needs a policy_factory and is not used here + server_max_batch_size: null # null = env.num_envs when async_policy=true + server_min_batch_size: 1 + server_timeout: 0.01 + server_collect_stats: true + server_stats_window_size: 1024 + max_inflight_per_env: 1 + storing_device: null advantage: trajectory_return: sum # binary success return per trajectory diff --git a/sota-implementations/vla_grpo/config/vla_grpo_toy.yaml b/sota-implementations/vla_grpo/config/vla_grpo_toy.yaml index e9a85d1d920..4f84fe14779 100644 --- a/sota-implementations/vla_grpo/config/vla_grpo_toy.yaml +++ b/sota-implementations/vla_grpo/config/vla_grpo_toy.yaml @@ -11,6 +11,7 @@ env: success_tol: 0.35 # sized so a random policy succeeds sometimes (cold-start signal) max_outer_steps: 6 # episode truncation, in chunk decisions render_size: 64 # side length of the from_pixels eval-video frame + num_envs: 1 # async-env workers; ToyVLAEnv grouped rollouts are per worker seed: 0 tokenizer: @@ -35,6 +36,23 @@ collector: max_same_policy_collect_attempts: 2 min_replay_decisions: null # null/0 = collect one target group wave total_iters: 200 + # Execution-mode switches for throughput ablations: + # - false/false: regular synchronous TorchRL Collector. + # - true/false: async env slots with one request per policy forward. + # - true/true: async env slots plus auto-batched policy inference. + # - false/true: sync env stepping through the policy server path. + async_env: false + async_policy: false + env_backend: threading # AsyncBatchedCollector env backend: threading | multiprocessing + policy_backend: threading # inference transport: threading | multiprocessing | ray | monarch + server_backend: thread # process server needs a policy_factory and is not used here + server_max_batch_size: null # null = number of async envs when async_policy=true + server_min_batch_size: 1 + server_timeout: 0.01 + server_collect_stats: true + server_stats_window_size: 1024 + max_inflight_per_env: 1 + storing_device: null advantage: trajectory_return: sum # binary success return per trajectory diff --git a/sota-implementations/vla_grpo/test_openvla.py b/sota-implementations/vla_grpo/test_openvla.py index c4ccc960002..7de11f08aa8 100644 --- a/sota-implementations/vla_grpo/test_openvla.py +++ b/sota-implementations/vla_grpo/test_openvla.py @@ -356,6 +356,187 @@ def hook(_): assert captured["kwargs"]["frames_per_batch"] == 4 +def test_make_collector_async_env_uses_async_batched_collector(monkeypatch): + captured = {} + + class _FakeAsyncCollector: + def __init__(self, *args, **kwargs): + captured["args"] = args + captured["kwargs"] = kwargs + + def __iter__(self): + return self + + def __next__(self): + raise StopIteration + + def server_stats(self, *, reset=False): + return {"requests": 0} + + def shutdown(self): + captured["shutdown"] = True + + class _FakeEnv: + batch_size = torch.Size([1]) + device = torch.device("cpu") + + cfg = SimpleNamespace( + collector=SimpleNamespace( + groups_per_iter=4, + group_size=2, + async_env=True, + async_policy=True, + server_min_batch_size=2, + ), + env=SimpleNamespace( + backend="toy", + action_dim=2, + state_dim=4, + image_shape=(3, 8, 8), + render_size=16, + success_steps=2, + success_tol=0.25, + max_outer_steps=3, + num_envs=4, + seed=0, + ), + ) + monkeypatch.setattr(utils, "AsyncBatchedCollector", _FakeAsyncCollector) + + collector = utils.make_collector( + cfg, + _FakeEnv(), + object(), + torch.device("cpu"), + tokenizer=object(), + replay_buffer=object(), + ) + collector._ensure_collector() + + assert len(captured["kwargs"]["create_env_fn"]) == 4 + assert captured["kwargs"]["yield_completed_trajectories"] + server_config = captured["kwargs"]["server_config"] + assert server_config.max_batch_size == 4 + assert server_config.min_batch_size == 2 + + +def test_make_collector_async_env_without_policy_batching(monkeypatch): + captured = {} + + class _FakeAsyncCollector: + def __init__(self, *args, **kwargs): + captured["kwargs"] = kwargs + + def __iter__(self): + return self + + def __next__(self): + raise StopIteration + + def server_stats(self, *, reset=False): + return {} + + def shutdown(self): + pass + + class _FakeEnv: + batch_size = torch.Size([1]) + device = torch.device("cpu") + + cfg = SimpleNamespace( + collector=SimpleNamespace( + groups_per_iter=2, + group_size=2, + async_env=True, + async_policy=False, + ), + env=SimpleNamespace( + backend="toy", + action_dim=2, + state_dim=4, + image_shape=(3, 8, 8), + render_size=16, + success_steps=2, + success_tol=0.25, + max_outer_steps=3, + num_envs=2, + seed=0, + ), + ) + monkeypatch.setattr(utils, "AsyncBatchedCollector", _FakeAsyncCollector) + + collector = utils.make_collector( + cfg, + _FakeEnv(), + object(), + torch.device("cpu"), + tokenizer=object(), + ) + collector._ensure_collector() + + server_config = captured["kwargs"]["server_config"] + assert server_config.max_batch_size == 1 + assert server_config.timeout == 0.0 + + +def test_make_collector_sync_env_can_use_policy_server(monkeypatch): + captured = {} + + class _FakeCollector: + def __init__(self, *args, **kwargs): + captured["collector_args"] = args + captured["collector_kwargs"] = kwargs + self.requested_frames_per_batch = kwargs["frames_per_batch"] + + def shutdown(self, *args, **kwargs): + captured["collector_shutdown"] = True + + def reset(self, *args, **kwargs): + captured["collector_reset"] = True + + class _FakeServer: + def __init__(self, *args, **kwargs): + captured["server_args"] = args + captured["server_kwargs"] = kwargs + + def start(self): + return self + + def shutdown(self): + captured["server_shutdown"] = True + + def stats(self, *, reset=False): + return {"requests": 0} + + class _FakeEnv: + batch_size = torch.Size([2]) + device = None + + policy = SimpleNamespace( + in_keys=["observation"], out_keys=[("vla_action", "tokens")] + ) + cfg = SimpleNamespace( + collector=SimpleNamespace( + groups_per_iter=2, + group_size=1, + async_policy=True, + ), + env=SimpleNamespace(max_outer_steps=3), + ) + monkeypatch.setattr(utils, "Collector", _FakeCollector) + monkeypatch.setattr(utils, "InferenceServer", _FakeServer) + + collector = utils.make_collector(cfg, _FakeEnv(), policy, torch.device("cpu")) + + assert isinstance(collector, utils._ServerBackedCollector) + assert isinstance(captured["collector_args"][1], utils.PolicyClientModule) + assert captured["server_kwargs"]["server_config"].max_batch_size == 2 + assert captured["collector_kwargs"]["policy_device"] == torch.device("cpu") + assert captured["collector_kwargs"]["trust_policy"] is True + collector.shutdown() + assert captured["server_shutdown"] + + def test_make_replay_buffer_scales_capacity_with_overcollection(): cfg = SimpleNamespace( collector=SimpleNamespace( diff --git a/sota-implementations/vla_grpo/utils.py b/sota-implementations/vla_grpo/utils.py index bdef29a624b..a3523118dd2 100644 --- a/sota-implementations/vla_grpo/utils.py +++ b/sota-implementations/vla_grpo/utils.py @@ -28,8 +28,8 @@ import torch -from tensordict import TensorDictBase -from torchrl.collectors import Collector +from tensordict import lazy_stack, TensorDictBase +from torchrl.collectors import AsyncBatchedCollector, Collector from torchrl.data import LazyTensorStorage, TensorDictReplayBuffer from torchrl.data.replay_buffers.samplers import SamplerWithoutReplacement from torchrl.data.vla import ( @@ -50,6 +50,12 @@ TransformedEnv, ) from torchrl.envs.utils import ExplorationType, set_exploration_type +from torchrl.modules.inference_server import ( + InferenceServer, + InferenceServerConfig, + PolicyClientModule, + ThreadingTransport, +) from torchrl.modules.vla import TinyVLA, VLAWrapperBase from torchrl.objectives import ClipPPOLoss from torchrl.objectives.llm import MCAdvantage, MCAdvantageSelector @@ -61,6 +67,13 @@ LOG_PROBS_KEY = ("vla_action", "log_probs") +def _cfg_get(section, key: str, default=None): + get = getattr(section, "get", None) + if get is not None: + return get(key, default) + return getattr(section, key, default) + + def candidate_group_size(cfg) -> int: """Number of rollout candidates collected for each GRPO group.""" collector_get = getattr( @@ -247,6 +260,146 @@ def _make_libero_worker( ) +def _num_envs_from_cfg(cfg, *, eval_mode: bool = False) -> int: + if cfg.env.backend == "toy": + default = 1 + else: + default = cfg.env.eval_num_envs if eval_mode else cfg.env.num_envs + key = "eval_num_envs" if eval_mode else "num_envs" + return int(_cfg_get(cfg.env, key, default)) + + +def _validate_libero_env_count( + cfg, num_envs: int, *, group_repeats=None, eval_mode: bool = False, override=False +) -> None: + task_ids = list(cfg.env.task_ids) + parallel_group_repeats = ( + not eval_mode + and group_repeats is not None + and bool(_cfg_get(cfg.env, "parallel_group_repeats", False)) + ) + task_coverage_envs = num_envs + if parallel_group_repeats: + group_repeats = int(group_repeats) + candidate_repeats = candidate_group_size(cfg) + if candidate_repeats % group_repeats: + raise ValueError( + "collector.candidate_group_size must be a multiple of " + "collector.group_size when env.parallel_group_repeats=true " + f"({candidate_repeats=} and {group_repeats=})." + ) + if num_envs % group_repeats: + raise ValueError( + "env.num_envs must be a multiple of collector.group_size " + "when env.parallel_group_repeats=true so every parallel " + f"group has exactly {group_repeats} workers ({num_envs=})." + ) + if _cfg_get(cfg.env, "train_init_state_mode", "random") == "random": + raise ValueError( + "env.parallel_group_repeats=true requires " + "env.train_init_state_mode='cycle' or 'fixed'. Random " + "init-state sampling is local to each worker, so workers " + "sharing a group id would not necessarily share the same " + "initial state." + ) + task_coverage_envs = num_envs // group_repeats + if not override and task_coverage_envs < len(task_ids): + raise ValueError( + f"{'eval_num_envs' if eval_mode else 'num_envs'} ({num_envs}) " + f"must cover task_ids ({len(task_ids)} tasks): each worker is " + "bound to one task; fewer workers would silently drop tasks. " + "With env.parallel_group_repeats=true, task coverage is " + f"num_envs / collector.group_size ({task_coverage_envs})." + ) + if not override and task_coverage_envs % len(task_ids): + warnings.warn( + f"effective task workers ({task_coverage_envs}) is not a " + f"multiple of the number of tasks ({len(task_ids)}): tasks " + "will be sampled unevenly." + ) + + +def _make_env_worker( + cfg, + tokenizer: ActionTokenizerBase, + worker_idx: int, + *, + group_repeats: int | None = None, + seed: int | None = None, + device: torch.device | None = None, + eval_mode: bool = False, + from_pixels: bool = False, +) -> TransformedEnv: + worker_seed = None if seed is None else int(seed) + int(worker_idx) + if cfg.env.backend == "toy": + base = ToyVLAEnv( + action_dim=cfg.env.action_dim, + state_dim=cfg.env.state_dim, + image_shape=tuple(cfg.env.image_shape), + from_pixels=from_pixels, + render_size=_cfg_get(cfg.env, "render_size", 64), + success_steps=cfg.env.success_steps, + success_tol=cfg.env.success_tol, + group_repeats=group_repeats, + group_id_offset=worker_idx * GROUP_ID_OFFSET, + batch_size=[], + seed=worker_seed, + device=device, + ) + elif cfg.env.backend == "libero": + base = _make_libero_worker( + cfg, + worker_idx, + group_repeats=group_repeats, + eval_mode=eval_mode, + from_pixels=from_pixels, + ) + if worker_seed is not None: + base.set_seed(worker_seed) + else: + raise ValueError(f"Unknown env backend {cfg.env.backend!r}.") + return TransformedEnv(base, _chunk_transform(cfg, tokenizer)) + + +def make_async_env_factories( + cfg, + tokenizer: ActionTokenizerBase, + *, + group_repeats: int | None = None, + seed: int | None = None, + device: torch.device | None = None, + eval_mode: bool = False, + from_pixels: bool = False, + num_envs: int | None = None, +) -> list[Callable[[], TransformedEnv]]: + """Build one transformed VLA env factory per async collection slot.""" + override = num_envs is not None + if num_envs is None: + num_envs = _num_envs_from_cfg(cfg, eval_mode=eval_mode) + if cfg.env.backend == "libero": + _validate_libero_env_count( + cfg, + num_envs, + group_repeats=group_repeats, + eval_mode=eval_mode, + override=override, + ) + return [ + partial( + _make_env_worker, + cfg, + tokenizer, + worker_idx, + group_repeats=group_repeats, + seed=seed, + device=device, + eval_mode=eval_mode, + from_pixels=from_pixels, + ) + for worker_idx in range(num_envs) + ] + + def make_env( cfg, tokenizer: ActionTokenizerBase, @@ -439,30 +592,176 @@ def make_optimizer(cfg, loss_module: ClipPPOLoss): return optim, scheduler +class _ServerBackedCollector: + """Collector wrapper that owns a thread-backed inference server.""" + + def __init__(self, collector: Collector, server: InferenceServer) -> None: + self.collector = collector + self.server = server + self.requested_frames_per_batch = collector.requested_frames_per_batch + + def __iter__(self): + return iter(self.collector) + + def reset(self, *args, **kwargs) -> None: + self.collector.reset(*args, **kwargs) + + def shutdown(self, *args, **kwargs) -> None: + try: + self.collector.shutdown(*args, **kwargs) + finally: + self.server.shutdown() + + def server_stats(self, *, reset: bool = False) -> dict[str, float | int]: + return self.server.stats(reset=reset) + + def __getattr__(self, name): + return getattr(self.collector, name) + + +class _BatchedPolicyClientModule(PolicyClientModule): + """Split a synchronous batched env observation into server requests.""" + + def forward(self, tensordict: TensorDictBase) -> TensorDictBase: + if tensordict.ndim == 0: + return super().forward(tensordict) + batch_size = tensordict.batch_size + flat_tensordict = tensordict.reshape(-1) + futures = [self.submit(td) for td in flat_tensordict.unbind(0)] + result = lazy_stack([future.result() for future in futures], 0) + result = result.reshape(batch_size) + self._check_policy_lag(result) + return result + + +class _AsyncReplayCollector: + """AsyncBatchedCollector wrapper that writes complete trajectories to replay.""" + + yields_complete_trajectories = True + requested_frames_per_batch = 1 + + def __init__( + self, + *, + create_env_fn: list[Callable[[], EnvBase]], + policy: VLAWrapperBase, + replay_buffer: TensorDictReplayBuffer | None, + collector_kwargs: dict, + ) -> None: + self._create_env_fn = create_env_fn + self._policy = policy + self._replay_buffer = replay_buffer + self._collector_kwargs = collector_kwargs + self._collector = None + self._iterator = None + self._last_server_stats: dict[str, float | int] = {} + self.num_envs = len(create_env_fn) + + def _ensure_collector(self): + if self._collector is None: + self._collector = AsyncBatchedCollector( + create_env_fn=self._create_env_fn, + policy=self._policy, + **self._collector_kwargs, + ) + self._iterator = iter(self._collector) + return self._collector + + def __iter__(self): + return self + + def __next__(self): + self._ensure_collector() + traj = next(self._iterator) + if self._replay_buffer is not None: + self._replay_buffer.extend(traj) + return traj + + def pause_collection(self) -> None: + if self._collector is None: + return + self._last_server_stats = self._collector.server_stats(reset=True) + self._collector.shutdown() + self._collector = None + self._iterator = None + + def reset(self, *args, **kwargs) -> None: + self.pause_collection() + + def shutdown(self, *args, **kwargs) -> None: + self.pause_collection() + + def server_stats(self, *, reset: bool = False) -> dict[str, float | int]: + if self._collector is not None: + return self._collector.server_stats(reset=reset) + result = dict(self._last_server_stats) + if reset: + self._last_server_stats = {} + return result + + +def _training_group_repeats(cfg) -> int: + env_get = getattr( + cfg.env, + "get", + lambda key, default=None: getattr(cfg.env, key, default), + ) + return ( + cfg.collector.group_size + if env_get("parallel_group_repeats", False) + else candidate_group_size(cfg) + ) + + +def _server_config_from_collector(cfg, *, num_envs: int) -> InferenceServerConfig: + collector_get = getattr( + cfg.collector, + "get", + lambda key, default=None: getattr(cfg.collector, key, default), + ) + async_policy = bool(collector_get("async_policy", False)) + if async_policy: + max_batch_size = int(collector_get("server_max_batch_size", None) or num_envs) + min_batch_size = int(collector_get("server_min_batch_size", 1)) + timeout = float(collector_get("server_timeout", 0.01)) + else: + max_batch_size = 1 + min_batch_size = 1 + timeout = 0.0 + return InferenceServerConfig( + max_batch_size=max_batch_size, + min_batch_size=min_batch_size, + timeout=timeout, + collect_stats=bool(collector_get("server_collect_stats", True)), + stats_window_size=int(collector_get("server_stats_window_size", 1024)), + ) + + def make_collector( cfg, env: EnvBase, policy: VLAWrapperBase, device: torch.device, *, + tokenizer: ActionTokenizerBase | None = None, replay_buffer: TensorDictReplayBuffer | None = None, post_collect_hook: Callable[[TensorDictBase], None] | None = None, ) -> Collector: - """Endless synchronous collector assembling complete trajectories. - - With no replay buffer, each yielded batch holds exactly one iteration's - worth of complete, done-terminated trajectories, concatenated along time - (``trajs_per_batch`` with ``traj_format="cat"``: flat and unpadded, - episodes delimited by the done flags -- no padding frames for the - image-heavy VLA observations; episodes spanning internal collection - steps are reassembled by the collector, in-flight episodes are held - back). - - With a replay buffer, TorchRL's collector writer path is used instead: - complete trajectories are pushed to the buffer as each internal rollout - batch finishes, and the iterator yields ``None``. This lets the replay - buffer transform keep incomplete GRPO groups across same-policy collection - polls until enough useful decisions have reached storage. + """Build a VLA rollout collector. + + The default path is the synchronous TorchRL ``Collector`` used by the + original recipe. Set ``collector.async_env=true`` to use + ``AsyncBatchedCollector`` env slots; set ``collector.async_policy=true`` to + route policy calls through an inference server with configurable + auto-batching. The async-env path yields complete trajectories and writes + them to the replay buffer directly, so the training loop can use the same + replay/advantage machinery across execution modes. + + With no replay buffer on the synchronous path, each yielded batch holds + complete, done-terminated trajectories, concatenated along time + (``trajs_per_batch`` with ``traj_format="cat"``). With a replay buffer, + TorchRL's collector writer path pushes complete trajectories to storage as + each internal rollout batch finishes. The policy is held by reference (in-place optimizer updates apply immediately) and observations/actions are cast between the env's and the @@ -471,7 +770,17 @@ def make_collector( via :func:`~torchrl.envs.utils.exploration_type` -- no policy mutation, so this works with any collector (including multi-process workers). """ - num_envs = env.batch_size[0] if env.batch_size else 1 + collector_get = getattr( + cfg.collector, + "get", + lambda key, default=None: getattr(cfg.collector, key, default), + ) + async_env = bool(collector_get("async_env", False)) + async_policy = bool(collector_get("async_policy", False)) + if async_env: + num_envs = _num_envs_from_cfg(cfg) + else: + num_envs = env.batch_size[0] if env.batch_size else 1 groups_per_iter = int(cfg.collector.groups_per_iter) group_size = int(cfg.collector.group_size) candidate_size = candidate_group_size(cfg) @@ -548,20 +857,79 @@ def make_collector( frames_per_batch = ( num_envs if replay_buffer is not None else (num_envs * cfg.env.max_outer_steps) ) - return Collector( + server_config = _server_config_from_collector(cfg, num_envs=num_envs) + if async_env: + if tokenizer is None: + raise ValueError( + "tokenizer is required when collector.async_env=true so async " + "environment factories can decode action tokens." + ) + create_env_fn = make_async_env_factories( + cfg, + tokenizer, + group_repeats=_training_group_repeats(cfg), + seed=cfg.env.seed, + device=env_device if cfg.env.backend == "toy" else None, + ) + return _AsyncReplayCollector( + create_env_fn=create_env_fn, + policy=policy, + replay_buffer=replay_buffer, + collector_kwargs={ + "frames_per_batch": 1, + "total_frames": -1, + "yield_completed_trajectories": True, + "env_backend": collector_get("env_backend", "threading"), + "policy_backend": collector_get("policy_backend", "threading"), + "server_backend": collector_get("server_backend", "thread"), + "server_config": server_config, + "policy_device": device, + "output_device": env_device, + "env_device": env_device, + "storing_device": collector_get("storing_device", None), + "max_inflight_per_env": collector_get("max_inflight_per_env", 1), + "verbose": bool(collector_get("verbose", False)), + }, + ) + + collector_policy = policy + collector_policy_device = device + server = None + if async_policy: + transport = ThreadingTransport() + server = InferenceServer( + policy, + transport, + server_config=server_config, + policy_device=device, + output_device=env_device, + ).start() + collector_policy = _BatchedPolicyClientModule( + transport, + in_keys=getattr(policy, "in_keys", None), + out_keys=getattr(policy, "out_keys", None), + max_inflight=None, + ) + collector_policy_device = env_device + + collector = Collector( env, - policy, + collector_policy, frames_per_batch=frames_per_batch, total_frames=-1, trajs_per_batch=groups_per_iter * candidate_size, traj_format="cat", exploration_type=ExplorationType.RANDOM, - policy_device=device, + policy_device=collector_policy_device, env_device=env_device, reset_at_each_iter=False, replay_buffer=replay_buffer, post_collect_hook=post_collect_hook, + trust_policy=True if async_policy else None, ) + if server is not None: + return _ServerBackedCollector(collector, server) + return collector def evaluate(env: TransformedEnv, policy: VLAWrapperBase, cfg) -> float: diff --git a/sota-implementations/vla_grpo/vla-grpo.py b/sota-implementations/vla_grpo/vla-grpo.py index 61668122940..d72fd015ef6 100644 --- a/sota-implementations/vla_grpo/vla-grpo.py +++ b/sota-implementations/vla_grpo/vla-grpo.py @@ -319,6 +319,13 @@ def main(cfg): # noqa: F821 "get", lambda key, default=None: getattr(cfg.env, key, default), ) + collector_get = getattr( + cfg.collector, + "get", + lambda key, default=None: getattr(cfg.collector, key, default), + ) + async_rollout_env = bool(collector_get("async_env", False)) + async_rollout_policy = bool(collector_get("async_policy", False)) train_group_repeats = ( cfg.collector.group_size if env_get("parallel_group_repeats", False) @@ -330,6 +337,7 @@ def main(cfg): # noqa: F821 group_repeats=train_group_repeats, seed=cfg.env.seed, device=device if cfg.env.backend == "toy" else None, + num_envs=1 if async_rollout_env else None, ) eval_process = bool(cfg.logger.get("eval_process", False)) eval_env = None @@ -383,6 +391,7 @@ def main(cfg): # noqa: F821 train_env, rollout_policy, rollout_device, + tokenizer=tokenizer, replay_buffer=replay_buffer, ) collector_iter = iter(collector) @@ -406,16 +415,23 @@ def main(cfg): # noqa: F821 int(cfg.collector.get("max_same_policy_collect_attempts", 1)), 1 ) min_replay_decisions = int(cfg.collector.get("min_replay_decisions", 0) or 0) - num_envs = train_env.batch_size[0] if train_env.batch_size else 1 - collector_frames_per_poll = max( - int(getattr(collector, "requested_frames_per_batch", num_envs)), 1 - ) - collect_polls_per_group_wave = max( - math.ceil( - episodes_per_iter * int(cfg.env.max_outer_steps) / collector_frames_per_poll - ), - 1, + num_envs = int(getattr(collector, "num_envs", 0)) or ( + train_env.batch_size[0] if train_env.batch_size else 1 ) + if getattr(collector, "yields_complete_trajectories", False): + collect_polls_per_group_wave = max(episodes_per_iter, 1) + else: + collector_frames_per_poll = max( + int(getattr(collector, "requested_frames_per_batch", num_envs)), 1 + ) + collect_polls_per_group_wave = max( + math.ceil( + episodes_per_iter + * int(cfg.env.max_outer_steps) + / collector_frames_per_poll + ), + 1, + ) max_collect_polls_per_iter = ( collect_polls_per_group_wave * max_collect_batches_per_iter ) @@ -440,6 +456,7 @@ def main(cfg): # noqa: F821 # runs under exploration_type=RANDOM (set in make_collector); the policy # reads that context, so the script never mutates it. with timeit("collect") as collect_timer: + collector_server_stats = {} if not retry_same_policy: advantage_transform.reset_stats() same_policy_collect_polls = 0 @@ -497,6 +514,13 @@ def main(cfg): # noqa: F821 else: safety_cap_hit = True + pause_collection = getattr(collector, "pause_collection", None) + if pause_collection is not None: + pause_collection() + server_stats = getattr(collector, "server_stats", None) + if server_stats is not None: + collector_server_stats = server_stats(reset=True) + if min_replay_decisions > 0 and len(replay_buffer) < min_replay_decisions: safety_cap_hit = True @@ -566,7 +590,12 @@ def main(cfg): # noqa: F821 if completed_trajectories else 0.0 ), + "collector/async_env": float(async_rollout_env), + "collector/async_policy": float(async_rollout_policy), } + for key, value in collector_server_stats.items(): + if isinstance(value, (float, int)): + group_metrics[f"policy_server/{key}"] = value # PPO update over the decisions that survived dynamic sampling, with # gradient accumulation (micro-batches of mini_batch_size decisions) num_decisions = len(replay_buffer) diff --git a/test/test_custom_envs.py b/test/test_custom_envs.py index 8558a6f6c1f..7de89c33cd0 100644 --- a/test/test_custom_envs.py +++ b/test/test_custom_envs.py @@ -397,6 +397,7 @@ def test_tracking_grouped_inits(self, batch_size): state_dim=4, success_steps=2, group_repeats=3, + group_id_offset=10, batch_size=batch_size, seed=0, ) @@ -405,7 +406,7 @@ def test_tracking_grouped_inits(self, batch_size): td = env.reset() group_ids.append(td["group_id"].reshape(()).item()) targets.append(td["observation", "state"][..., 2:4].clone()) - assert group_ids == [0, 0, 0, 1, 1, 1] + assert group_ids == [10, 10, 10, 11, 11, 11] torch.testing.assert_close(targets[1], targets[0]) torch.testing.assert_close(targets[2], targets[0]) torch.testing.assert_close(targets[4], targets[3]) @@ -413,7 +414,7 @@ def test_tracking_grouped_inits(self, batch_size): assert not torch.allclose(targets[3], targets[0]) # the group id rides every step of the episode td["action"] = td["observation", "state"][..., 2:4] - assert env.step(td)["next", "group_id"].reshape(()).item() == 1 + assert env.step(td)["next", "group_id"].reshape(()).item() == 11 def test_tracking_group_repeats_validation(self): with pytest.raises(ValueError, match="success_steps"): diff --git a/torchrl/envs/custom/vla.py b/torchrl/envs/custom/vla.py index 451a6991615..2b368afb9ff 100644 --- a/torchrl/envs/custom/vla.py +++ b/torchrl/envs/custom/vla.py @@ -95,6 +95,10 @@ class ToyVLAEnv(EnvBase): advantages require (n rollouts per initial state, e.g. grouped by :class:`~torchrl.objectives.llm.MCAdvantage`). Defaults to ``None`` (a fresh target every episode, no ``group_id`` entry). + group_id_offset (int, optional): offset added to grouped rollout ids. + This lets several single ToyVLAEnv instances collect grouped + rollouts in parallel without mixing unrelated initial states in the + downstream group-advantage transform. Defaults to ``0``. batch_size (torch.Size, optional): number of vectorized copies. Defaults to ``torch.Size([])`` (a single environment). device (torch.device, optional): device of the specs. @@ -150,6 +154,7 @@ def __init__( success_steps: int | None = None, success_tol: float = 0.25, group_repeats: int | None = None, + group_id_offset: int = 0, batch_size: torch.Size | None = None, device: torch.device | None = None, seed: int | None = None, @@ -195,6 +200,7 @@ def __init__( self.success_steps = int(success_steps) if success_steps is not None else None self.success_tol = float(success_tol) self.group_repeats = int(group_repeats) if group_repeats is not None else None + self.group_id_offset = int(group_id_offset) if self.group_repeats is not None and self.batch_size.numel() > 1: raise ValueError( "group_repeats only supports a single environment " @@ -368,7 +374,8 @@ def _reset(self, tensordict: TensorDictBase | None = None, **kwargs) -> TensorDi if self._episode_count % self.group_repeats == 0: self._target = self._sample_target() self._group_id = torch.full_like( - self._group_id, self._episode_count // self.group_repeats + self._group_id, + self.group_id_offset + self._episode_count // self.group_repeats, ) self._episode_count += 1 else: