Skip to content
Open
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
158 changes: 158 additions & 0 deletions benchmarks/benchmark_parallel_env_startup.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,158 @@
# Copyright (c) Meta Platforms, Inc. and affiliates.
#
# This source code is licensed under the MIT license found in the
# LICENSE file in the root directory of this source tree.

"""Benchmark ``ParallelEnv`` startup metadata strategies.

The benchmark records constructor calls in separate marker files so parent-side
temporary environments and long-lived worker environments can be counted
independently. For example::

python benchmarks/benchmark_parallel_env_startup.py --workers 80 --delay 0.1

Use ``--mode`` to run one strategy, and ``--start-method forkserver`` to compare
the no-shadow-environment path with the default ``spawn`` worker startup.
"""

from __future__ import annotations

import argparse
import json
import os
import tempfile
import time
from functools import partial
from pathlib import Path
from typing import Literal

from torchrl import timeit
from torchrl._utils import logger as torchrl_logger
from torchrl.envs import ParallelEnv
from torchrl.testing.mocking_classes import CountingEnv


class _StartupEnv(CountingEnv):
def __init__(
self,
worker_idx: int,
*,
marker_dir: str,
parent_pid: int,
delay: float,
) -> None:
if delay:
time.sleep(delay)
role = "parent" if os.getpid() == parent_pid else "worker"
self._marker_path = Path(marker_dir) / (
f"constructed-{role}-{worker_idx}-{os.getpid()}"
)
self._marker_path.touch()
super().__init__()

def close(self, *, raise_if_closed: bool = True) -> None:
self._marker_path.with_name(f"closed-{self._marker_path.name}").touch()
super().close(raise_if_closed=raise_if_closed)


def _run_mode(
mode: Literal["legacy", "homogeneous", "workers"],
*,
num_workers: int,
delay: float,
start_method: Literal["spawn", "forkserver", "fork"],
marker_dir: Path,
) -> dict[str, float | int | str]:
parent_pid = os.getpid()
common_factory = partial(
_StartupEnv,
marker_dir=str(marker_dir),
parent_pid=parent_pid,
delay=delay,
)
create_env_kwargs = [
{"worker_idx": worker_idx} for worker_idx in range(num_workers)
]
if mode == "legacy":
create_env_fn = [
partial(common_factory, worker_idx=worker_idx)
for worker_idx in range(num_workers)
]
create_env_kwargs = None
else:
create_env_fn = common_factory

with timeit(f"{mode}/construct") as construct_timer:
env = ParallelEnv(
num_workers,
create_env_fn,
create_env_kwargs=create_env_kwargs,
metadata_from_workers=mode == "workers",
use_buffers=False,
mp_start_method=start_method,
)
construct_seconds = construct_timer.elapsed()
try:
with timeit(f"{mode}/reset") as reset_timer:
tensordict = env.reset()
reset_seconds = reset_timer.elapsed()
with timeit(f"{mode}/first_step") as step_timer:
env.rand_step(tensordict)
first_step_seconds = step_timer.elapsed()
finally:
env.close(raise_if_closed=False)

parent_constructions = len(list(marker_dir.glob("constructed-parent-*")))
worker_constructions = len(list(marker_dir.glob("constructed-worker-*")))
closed_constructions = len(list(marker_dir.glob("closed-constructed-*")))
return {
"mode": mode,
"start_method": start_method,
"num_workers": num_workers,
"parent_constructions": parent_constructions,
"worker_constructions": worker_constructions,
"closed_constructions": closed_constructions,
"construct_seconds": construct_seconds,
"reset_seconds": reset_seconds,
"first_step_seconds": first_step_seconds,
}


def main() -> None:
parser = argparse.ArgumentParser()
parser.add_argument("--workers", type=int, default=80)
parser.add_argument("--delay", type=float, default=0.05)
parser.add_argument(
"--mode",
choices=("all", "legacy", "homogeneous", "workers"),
default="all",
)
parser.add_argument(
"--start-method",
choices=("spawn", "forkserver", "fork"),
default="spawn",
)
parser.add_argument("--output", type=Path)
args = parser.parse_args()

modes = ("legacy", "homogeneous", "workers") if args.mode == "all" else (args.mode,)
results = []
for mode in modes:
with tempfile.TemporaryDirectory() as marker_dir:
result = _run_mode(
mode,
num_workers=args.workers,
delay=args.delay,
start_method=args.start_method,
marker_dir=Path(marker_dir),
)
results.append(result)
torchrl_logger.info("ParallelEnv startup benchmark: %s", result)

if args.output is not None:
args.output.parent.mkdir(parents=True, exist_ok=True)
args.output.write_text(json.dumps(results, indent=2) + "\n")


if __name__ == "__main__":
main()
28 changes: 28 additions & 0 deletions docs/source/reference/envs_vectorized.rst
Original file line number Diff line number Diff line change
Expand Up @@ -63,6 +63,34 @@ needed.
parallel env can be a bottleneck. This is why, for instance, TorchRL tests are so slow.
Once the environment is launched, a great speedup should be observed.

Environments with especially expensive constructors can avoid a second,
parent-side construction pass by setting ``metadata_from_workers=True``.
In this opt-in mode, the real worker environments report their metadata before
normal initialization. The workers start eagerly, their tensor schemas are
validated for compatibility, and no temporary environment is created in the
parent. This mode currently requires pipe-based communication with
``use_buffers=False``. All workers must expose the same tensor schema: specs
and example tensors may only differ in non-tensor payload values (such as
language instructions). Environments with genuinely heterogeneous specs
should keep the default metadata path. At shutdown, workers started in this
mode are closed one at a time to bound teardown resource spikes; the
per-worker grace period is controlled by the ``shutdown_timeout`` argument.
Use one common factory with ``create_env_kwargs`` for
worker-specific arguments:

.. code-block:: python

from functools import partial

make_env = partial(GymEnv, "Pendulum-v1")
env = ParallelEnv(
4,
make_env,
create_env_kwargs=[{"g": 9.0 + worker_idx} for worker_idx in range(4)],
metadata_from_workers=True,
use_buffers=False,
)

.. note::

*TorchRL requires precise specs*: Another thing to take in consideration is
Expand Down
122 changes: 122 additions & 0 deletions sota-implementations/vla_grpo/test_openvla.py
Original file line number Diff line number Diff line change
Expand Up @@ -818,6 +818,128 @@ def test_chunk_transform_openvla_tokens_decodes_then_postprocesses(self):


class TestLiberoWorkers:
def test_standalone_env_factory_uses_worker_metadata(self, monkeypatch):
captured = {}

def fake_parallel_env(num_envs, env_factory, **kwargs):
captured["num_envs"] = num_envs
captured["env_factory"] = env_factory
captured["kwargs"] = kwargs
return object()

cfg = SimpleNamespace(
env=SimpleNamespace(
backend="libero",
task_ids=[0, 1],
num_envs=2,
parallel_group_repeats=False,
),
collector=SimpleNamespace(group_size=8, candidate_group_size=None),
)
monkeypatch.setattr(utils, "ParallelEnv", fake_parallel_env)
monkeypatch.setattr(utils, "_chunk_transform", lambda *args, **kwargs: object())
monkeypatch.setattr(utils, "TransformedEnv", lambda base, transform: base)

result = utils.make_env(cfg, tokenizer=None)

assert result is not None
assert captured["num_envs"] == 2
assert captured["kwargs"]["create_env_kwargs"] == [
{"worker_idx": 0},
{"worker_idx": 1},
]
assert captured["kwargs"]["metadata_from_workers"]
assert captured["kwargs"]["use_buffers"] is False
assert captured["env_factory"].func is utils._make_libero_worker

def test_collector_factory_uses_common_env_factory_and_worker_kwargs(
self, monkeypatch
):
captured = {}

def fake_make_env_worker(*args, **kwargs):
return args, kwargs

def fake_parallel_env(num_envs, env_factory, **kwargs):
captured["num_envs"] = num_envs
captured["env_factory"] = env_factory
captured["kwargs"] = kwargs
return object()

cfg = SimpleNamespace(env=SimpleNamespace(backend="libero"))
tokenizer = object()
monkeypatch.setattr(utils, "_make_env_worker", fake_make_env_worker)
monkeypatch.setattr(utils, "ParallelEnv", fake_parallel_env)

result = utils._make_collector_env(
cfg,
tokenizer=tokenizer,
num_envs=3,
group_repeats=8,
seed=100,
device=torch.device("cpu"),
worker_idx_offset=16,
render_gpu_device_id=2,
)

assert result is not None
assert captured["num_envs"] == 3
assert captured["kwargs"]["create_env_kwargs"] == [
{"worker_idx": 0},
{"worker_idx": 1},
{"worker_idx": 2},
]
assert captured["kwargs"]["metadata_from_workers"]
assert captured["kwargs"]["use_buffers"] is False
assert captured["kwargs"]["mp_start_method"] == "spawn"
args, kwargs = captured["env_factory"](worker_idx=2)
assert args == (cfg, tokenizer)
assert kwargs == {
"worker_idx": 2,
"group_repeats": 8,
"seed": 100,
"device": None,
"worker_idx_offset": 16,
"render_gpu_device_id": 2,
}

def test_collector_factory_preserves_group_offset_and_seed(self, monkeypatch):
captured = {}

class FakeLiberoEnv:
def set_seed(self, seed):
captured["seed"] = seed

def fake_make_libero_worker(cfg, worker_idx, **kwargs):
captured["worker_idx"] = worker_idx
captured["worker_kwargs"] = kwargs
return FakeLiberoEnv()

monkeypatch.setattr(utils, "_make_libero_worker", fake_make_libero_worker)
monkeypatch.setattr(utils, "_chunk_transform", lambda *args, **kwargs: object())
monkeypatch.setattr(utils, "TransformedEnv", lambda base, transform: base)
cfg = SimpleNamespace(env=SimpleNamespace(backend="libero"))

utils._make_env_worker(
cfg,
tokenizer=None,
worker_idx=3,
group_repeats=8,
seed=100,
worker_idx_offset=16,
render_gpu_device_id=2,
)

assert captured["worker_idx"] == 3
assert captured["worker_kwargs"] == {
"group_repeats": 8,
"eval_mode": False,
"from_pixels": False,
"worker_idx_offset": 16,
"render_gpu_device_id": 2,
}
assert captured["seed"] == 119

def test_libero_worker_assignment_serial_and_parallel_groups(self):
class _EnvCfg(dict):
def __getattr__(self, name):
Expand Down
Loading
Loading