From d465c7003966d8a3f1395230fcd0cfc7f7d3358b Mon Sep 17 00:00:00 2001 From: coder-jayp Date: Sat, 3 Oct 2026 21:32:14 +0530 Subject: [PATCH] Fix kwargs dropping in AsyncEnvPool.reset - Passed **kwargs through AsyncEnvPool.reset to async_reset_send. - Updated _send_worker_batches to append kwargs to the worker queue requests. - Updated _worker_exec and _env_exec to gracefully unpack tuple lengths (falling back to {} for old messages). - Updated ThreadingAsyncEnvPool reset functions to pass kwargs. - Added regression test test_async_pool_kwargs_forwarding to guarantee kwargs are not dropped. --- test/envs/test_env_base.py | 15 +++++++++- torchrl/envs/async_envs.py | 57 +++++++++++++++++++++++++------------- 2 files changed, 52 insertions(+), 20 deletions(-) diff --git a/test/envs/test_env_base.py b/test/envs/test_env_base.py index 48b36641075..29e7e71abf2 100644 --- a/test/envs/test_env_base.py +++ b/test/envs/test_env_base.py @@ -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 @@ -1100,6 +1100,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) diff --git a/torchrl/envs/async_envs.py b/torchrl/envs/async_envs.py index f6380514437..6993eb6e45e 100644 --- a/torchrl/envs/async_envs.py +++ b/torchrl/envs/async_envs.py @@ -509,11 +509,6 @@ def reset( tensordict: TensorDictBase | None = None, **kwargs, ) -> TensorDictBase: - if kwargs.get("set_state", False): - raise NotImplementedError( - "EnvPool (AsyncEnvPool / ParallelEnv) currently drops kwargs and " - "cannot forward set_state=True to sub-environments." - ) if self._current_step > 0: raise RuntimeError("Some envs are still processing a step.") if tensordict is None: @@ -532,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 @@ -1000,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)) @@ -1244,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) @@ -1262,6 +1259,7 @@ def async_reset_send( tensordict.unbind(0), per_env=_per_env, record_action=False, + **kwargs, ) def async_reset_recv( @@ -1314,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) @@ -1329,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 @@ -1481,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( @@ -1518,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 @@ -1544,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)) @@ -1699,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]): @@ -1927,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) @@ -1941,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: @@ -1982,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) @@ -1994,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)