Skip to content

[Feature Request] AsyncEnvPool v2 #4061

Description

@vmoens

[Feature] AsyncEnvPool v2: shared-memory slot exchange, deadline batch former, and a SEED-style inference path

Motivation

AsyncEnvPool already has the right semantics for high-throughput acting: envs advance independently, min_get allows partial-batch harvesting, and step_and_maybe_reset inside the worker hides reset latency from the trainer. What limits it today is the data plane:

  • Every step ships a tensordict through an mp.Queue (pickling + Python object churn),
  • results carry env_index as NonTensorData and are re-assembled with _sort_results + torch.stack on every recv (a fresh allocation + copy per batch),
  • there is no batch-forming policy (min_get is "at least N", with no max/deadline) and no instrumentation to tell whether a run is env-bound or inference-bound,
  • the pool ends at one host, and the consumer half (gather → policy → scatter, overlapped with batch forming) is left entirely to user Python.

None of this matters for cheap Python envs on a single node — queue overhead vanishes under env step time. It matters a lot in the regime of many slow envs (sim gangs, game engines, browsers), pixel observations, and a policy pinned to a large accelerator. That is the SEED RL regime ([Espeholt et al., 2019](https://arxiv.org/abs/1910.06591)), and nothing in the PyTorch ecosystem currently serves it out of the box. I have seen an internal SEED-style stack built along the lines below sustain fleets of heavyweight simulator envs feeding a central learner; every item in this proposal is a generalization of something that measurably mattered there.

Proposal

Ordered by expected payoff. Items 1–3 are the core; 4–6 build on them.

1. Slot-based shared-memory exchange (replace queues on the hot path)

Pre-allocate one contiguous shared (and optionally pinned) tensordict with a num_envs leading dimension — the request buffer — plus a matching response buffer for actions. Each worker owns slot i:

  • worker writes its (obs, step_type, reward) record in place (update_) and raises a ready flag (semaphore / eventfd / lock-free ready ring);
  • recv collects ready indices and returns a view into the request buffer plus the index list — no pickle, no sort, no stack;
  • send scatters actions into the response buffer and signals the corresponding slots — no message back either.

tensordict already has the storage machinery (shared/memmap allocation, in-place updates); the new pieces are the readiness signaling and the index-batch former. This deletes the three hot-path costs (queue serialization, _sort_results, per-recv torch.stack) in one change.

pool = AsyncEnvPool(env_fns, backend="multiprocessing", exchange="shm")
ready = pool.recv(min_get=1, max_get=64, timeout=0.002)   # TensorDict view + ready.env_indices
actions = policy(ready)
pool.send(actions, ready.env_indices)

A discount field need not cross the exchange at all: encode the termination kind in the step type (FIRST / MID / TERMINATION / TRUNCATION) and derive discount on the consumer side.

2. Deadline-based batch former + health metrics

Extend recv to a real batching policy: return when min_get is satisfied and (max_get reached or timeout expired). Export counters:

  • batch fill ratio and partial-batch fraction,
  • obs→batch dwell (how long a ready slot waits to be included),
  • batch→action dwell (how long a formed batch waits on the policy),
  • consumer busy fraction.

These four numbers are the difference between knowing a run is env-bound vs inference-bound and guessing. They are cheap to collect at the exchange and should be first-class, not left to user-side timers.

3. Fixed-shape inference path

Partial batches break torch.compile / CUDA graphs. Provide an opt-in consumer utility that:

  • pads formed batches to a fixed size (or a small set of bucket sizes),
  • gathers ready slots into a pinned staging buffer and issues a single async H2D on a dedicated stream,
  • double-buffers so batch N+1 forms while batch N computes.

This can live as a helper around the exchange rather than inside the pool, but it should ship with torchrl — it is the part every user currently rebuilds, badly, in a Python loop.

4. Worker-owned autoreset as the steady-state protocol

In the async hot loop the trainer should never send "reset": workers loop step/auto-reset forever (gymnasium NEXT_STEP semantics — terminal step delivered, action into a terminal state discarded, next record is FIRST) and only ever ship records. step_and_maybe_reset already does the env-side work; this makes it the only steady-state mode of the async pool, which simplifies the wire protocol of item 5 and guarantees reset latency is always hidden.

5. Remote feeder: a small env-host protocol

The exchange from item 1 doesn't care who writes a slot. Add a thin cross-host feeder:

  • env-host client: connect, handshake for the obs/action spec, then a bidi gRPC stream of records — the host owns its env and its resets; the client is ~200 lines of Python;
  • server side: the stream ingress claims a slot in the same exchange the local workers use.

This takes AsyncEnvPool cross-host with zero changes to the consumer loop, and lets envs be "anything that speaks the stream" — containers, other languages, other machines. Note this is deliberately a different regime from the Ray/RPC collectors, which replicate the policy onto actors (IMPALA-style) and ship trajectories; the SEED shape wins when the policy is too big or too hot to replicate onto env nodes.

6. Optional: a native (Rust or C++) dispatch engine

With items 1–5 in Python, the remaining ceiling is the coordinator itself: a Python thread polling ready flags and forming batches steals GIL time from the policy thread between CUDA launches, and gRPC ingress in Python adds per-message overhead. The fix is small and well-bounded — the slot exchange, ready-ring, batch former, and stream ingress are a few hundred lines of Rust (PyO3) or C++ behind exactly the API above:

  • Python keeps ownership of the policy, the training loop, and the buffers' tensor views;
  • the native engine owns readiness signaling, batch forming, and (for item 5) the gRPC ingress, and hands Python (batch_view, indices) pairs;
  • GIL contention stops being a scaling wall, and the engine can expose the item-2 metrics with nanosecond-resolution dwell tracking essentially for free.

I've seen this exact split (Python/JAX policy over native-owned slot buffers) hold up well in production; I'd scope it as an optional extension backend (exchange="native") rather than a rewrite, so the pure-Python path remains the default and the reference implementation.

Non-goals

  • Replacing ParallelEnv / the sync batched API — lockstep vectorization stays the right tool for evaluation and cheap envs.
  • Replacing the distributed/Ray collectors — actor-side inference remains the right shape when the policy is small.

Rollout sketch

  1. exchange="shm" + recv(min_get, max_get, timeout) + metrics (items 1–2) behind the existing AsyncEnvPool API.
  2. Fixed-shape consumer helper (item 3) + autoreset-only steady state (item 4).
  3. Env-host protocol (item 5).
  4. Native engine backend (item 6), gated on profiling results from 1–3.

Happy to iterate on the API sketch — opening this to collect feedback on the overall direction first.

## Building a v2: see VecNormV2 example

https://docs.pytorch.org/rl/main/_modules/torchrl/envs/transforms/vecnorm.html#VecNormV2

cc @theap06

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

enhancementNew feature or request

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions