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
15 changes: 14 additions & 1 deletion test/envs/test_env_base.py
Original file line number Diff line number Diff line change
Expand Up @@ -25,7 +25,7 @@
from torch import nn

from torchrl.data.tensor_specs import Binary, Composite, NonTensor, Unbounded
from torchrl.envs import EnvBase, ParallelEnv, SerialEnv
from torchrl.envs import AsyncEnvPool, EnvBase, ParallelEnv, SerialEnv
from torchrl.envs.libs.gym import GymEnv
from torchrl.envs.transforms import StepCounter, Transform, TransformedEnv
from torchrl.envs.utils import check_env_specs, make_composite_from_td, step_mdp
Expand Down Expand Up @@ -1102,6 +1102,19 @@ def test_transformed_env_delegates_implicit_state_detection(self):
out = env.reset(td)
assert (out["observation"] == 0).all()

def test_async_pool_kwargs_forwarding(self):
# Regression test: verify that kwargs (like set_state=True) are properly
# forwarded across the AsyncEnvPool boundary to the sub-environments.
# _KwargOnlySetStateEnv is explicitly hardcoded to raise a RuntimeError
# if set_state=True reaches its _reset method.
env = AsyncEnvPool([_KwargOnlySetStateEnv, _KwargOnlySetStateEnv])
try:
with pytest.raises(RuntimeError, match="unexpected implicit set_state"):
# If kwargs are silently dropped, this will NOT raise an error and the test will fail!
env.reset(set_state=True)
finally:
env.close()


_MOCK_ENV_KEY_CASES = [
# (class name, out_key, pixel shape or None for vector obs)
Expand Down
52 changes: 38 additions & 14 deletions torchrl/envs/async_envs.py
Original file line number Diff line number Diff line change
Expand Up @@ -527,7 +527,7 @@ def reset(
if indices.shape != tensordict.shape:
indices = expand_as_right(indices, tensordict)
tensordict[self._env_idx_key] = indices
self.async_reset_send(tensordict)
self.async_reset_send(tensordict, **kwargs)
tensordict = self.async_reset_recv(min_get=self.num_envs)
return tensordict

Expand Down Expand Up @@ -995,13 +995,14 @@ def _send_worker_batches(
*,
per_env: bool,
record_action: bool,
**kwargs,
) -> None:
requests = [[] for _ in range(self.num_workers)]
for env_index, local_td in _zip_strict(env_idx, local_tds):
data = self._prepare_worker_data(
env_index, local_td, record_action=record_action
)
requests[self._env_to_worker[env_index]].append((env_index, data))
requests[self._env_to_worker[env_index]].append((env_index, data, kwargs))
for worker_index, worker_requests in enumerate(requests):
if worker_requests:
self.input_queue[worker_index].put((msg, worker_requests, per_env))
Expand Down Expand Up @@ -1239,6 +1240,7 @@ def async_reset_send(
self,
tensordict: TensorDictBase | None = None,
env_index: int | list[int] | None = None,
**kwargs,
) -> None:
tensordict, env_idx = self._maybe_make_tensordict(tensordict, env_index, True)
_per_env = isinstance(env_index, int)
Expand All @@ -1257,6 +1259,7 @@ def async_reset_send(
tensordict.unbind(0),
per_env=_per_env,
record_action=False,
**kwargs,
)

def async_reset_recv(
Expand Down Expand Up @@ -1309,6 +1312,7 @@ def _async_private_reset_send(
self,
tensordict: TensorDictBase | None = None,
env_index: int | list[int] | None = None,
**kwargs,
) -> None:
tensordict, env_idx = self._maybe_make_tensordict(tensordict, env_index, True)

Expand All @@ -1324,6 +1328,7 @@ def _async_private_reset_send(
tensordict.unbind(0),
per_env=False,
record_action=False,
**kwargs,
)

_async_private_reset_recv = async_reset_recv
Expand Down Expand Up @@ -1476,8 +1481,13 @@ def _worker_exec(
next_slots.unbind(0),
)
]
for env_index, data in requests:
local_input_queues[env_index].put((msg, data, per_env))
for request in requests:
if len(request) == 3:
env_index, data, kwargs = request
else:
env_index, data = request
kwargs = {}
local_input_queues[env_index].put((msg, data, per_env, kwargs))

@classmethod
def _env_exec(
Expand Down Expand Up @@ -1513,11 +1523,15 @@ def _env_exec(

while True:
msg_data = input_queue.get()
if len(msg_data) == 3:
if len(msg_data) == 4:
msg, data, per_env, kwargs = msg_data
elif len(msg_data) == 3:
msg, data, per_env = msg_data
kwargs = {}
else:
msg, data = msg_data
per_env = False
kwargs = {}
if grouped_input and msg != "shutdown":
if msg == "init_shm":
input_slots, result_slots, next_slots, clock = data
Expand All @@ -1539,7 +1553,7 @@ def _env_exec(
elif msg == "reset":
if shared_slots is not None:
data = shared_slots[0].select(*data, strict=True)
data = env.reset(data)
data = env.reset(data, **kwargs)
target = per_env_reset_queue if per_env else reset_queue
if shared_slots is None:
data.set(cls._env_idx_key, NonTensorData(data=i))
Expand Down Expand Up @@ -1694,14 +1708,22 @@ def _private_step_func(cls, env_td: tuple[EnvBase, TensorDictBase, int]):
return env._step(td).set(cls._env_idx_key, NonTensorData(data=idx))

@classmethod
def _reset_func(cls, env_td: tuple[EnvBase, TensorDictBase]):
env, td, idx = env_td
return env.reset(td).set(cls._env_idx_key, NonTensorData(data=idx))
def _reset_func(cls, env_td: tuple):
if len(env_td) == 4:
env, td, idx, kwargs = env_td
else:
env, td, idx = env_td
kwargs = {}
return env.reset(td, **kwargs).set(cls._env_idx_key, NonTensorData(data=idx))

@classmethod
def _private_reset_func(cls, env_td: tuple[EnvBase, TensorDictBase]):
env, td, idx = env_td
return env._reset(td).set(cls._env_idx_key, NonTensorData(data=idx))
def _private_reset_func(cls, env_td: tuple):
if len(env_td) == 4:
env, td, idx, kwargs = env_td
else:
env, td, idx = env_td
kwargs = {}
return env._reset(td, **kwargs).set(cls._env_idx_key, NonTensorData(data=idx))

@classmethod
def _step_and_maybe_reset_func(cls, env_td: tuple[EnvBase, TensorDictBase]):
Expand Down Expand Up @@ -1922,6 +1944,7 @@ def async_reset_send(
self,
tensordict: TensorDictBase | None = None,
env_index: int | list[int] | None = None,
**kwargs,
) -> None:
tensordict, env_idx = self._maybe_make_tensordict(tensordict, env_index, True)
_per_env = isinstance(env_index, int)
Expand All @@ -1936,7 +1959,7 @@ def async_reset_send(
tds = tensordict.unbind(0)
envs = [self.envs[idx] for idx in env_idx]
futures = [
self._pool.submit(self._reset_func, (env, td, idx))
self._pool.submit(self._reset_func, (env, td, idx, kwargs))
for env, td, idx in zip(envs, tds, env_idx)
]
if _per_env:
Expand Down Expand Up @@ -1977,6 +2000,7 @@ def _async_private_reset_send(
self,
tensordict: TensorDictBase | None = None,
env_index: int | list[int] | None = None,
**kwargs,
) -> None:
tensordict, env_idx = self._maybe_make_tensordict(tensordict, env_index, True)

Expand All @@ -1989,7 +2013,7 @@ def _async_private_reset_send(
tds = tensordict.unbind(0)
envs = [self.envs[idx] for idx in env_idx]
futures = [
self._pool.submit(self._private_reset_func, (env, td, idx))
self._pool.submit(self._private_reset_func, (env, td, idx, kwargs))
for env, td, idx in zip(envs, tds, env_idx)
]
self._current_reset = self._current_reset + len(futures)
Expand Down
Loading