env-ssl-wrapper 0.3.0__tar.gz → 0.4.0__tar.gz

This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
Files changed (26) hide show
  1. {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.0}/PKG-INFO +65 -19
  2. {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.0}/README.md +63 -18
  3. env_ssl_wrapper-0.4.0/env_ssl_wrapper/__init__.py +47 -0
  4. env_ssl_wrapper-0.4.0/env_ssl_wrapper/memory_trace.py +93 -0
  5. env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize/__init__.py +166 -0
  6. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/adapters.py +1 -14
  7. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/auto_batched_wrapper.py +6 -0
  8. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/done_tracker_wrapper.py +5 -0
  9. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/episode_padding_wrapper.py +8 -4
  10. env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize/flatten_obs_wrapper.py +75 -0
  11. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/helpers.py +87 -0
  12. env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize/standardize_env_wrapper.py +66 -0
  13. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/tensor_wrapper.py +44 -32
  14. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/time_limit_wrapper.py +9 -2
  15. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/utils.py +29 -5
  16. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/vector.py +41 -5
  17. {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.0}/pyproject.toml +4 -6
  18. env_ssl_wrapper-0.3.0/env_ssl_wrapper/__init__.py +0 -59
  19. env_ssl_wrapper-0.3.0/env_ssl_wrapper/flatten_obs_wrapper.py +0 -79
  20. {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.0}/.gitignore +0 -0
  21. {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.0}/LICENSE +0 -0
  22. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/action_transform_wrapper.py +0 -0
  23. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/image_wrapper.py +0 -0
  24. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/mocks.py +0 -0
  25. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/spaces.py +0 -0
  26. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/standardize_wrapper.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: env-ssl-wrapper
3
- Version: 0.3.0
3
+ Version: 0.4.0
4
4
  Summary: One torch-native interface for any MDP environment
5
5
  Project-URL: Homepage, https://pypi.org/project/env-ssl-wrapper/
6
6
  Project-URL: Repository, https://codeberg.org/lucidrains/env-ssl-wrapper
@@ -43,6 +43,7 @@ Requires-Dist: torch>=2.5
43
43
  Requires-Dist: x-mlps-pytorch
44
44
  Requires-Dist: x-transformers
45
45
  Provides-Extra: examples
46
+ Requires-Dist: fire; extra == 'examples'
46
47
  Provides-Extra: test
47
48
  Requires-Dist: box2d-py>=2.3.5; extra == 'test'
48
49
  Requires-Dist: dm-control; extra == 'test'
@@ -67,19 +68,18 @@ One line turns any simulator's environment — mujoco warp, isaac sim, pybullet,
67
68
  pip install env-ssl-wrapper
68
69
  ```
69
70
 
70
- ## Usage
71
+ ## Standardize
71
72
 
72
73
  ```python
73
74
  import torch
74
- from env_ssl_wrapper import compose_env
75
+ from env_ssl_wrapper import StandardizeEnvWrapper
75
76
 
76
- env = compose_env(
77
- any_env, # any env from any sim
78
- ('tensor', dict(device='cpu')), # wrap with whatever you need
79
- 'done_tracker',
80
- )
77
+ # one master wrapper to standardize any simulator environment
78
+
79
+ env = StandardizeEnvWrapper(any_env)
80
+
81
+ obs, info = env.reset() # torch.float32, batched
81
82
 
82
- obs, info = env.reset() # torch.float32, batched
83
83
  while not env.all_done:
84
84
  actions = torch.randint(0, 2, (8,))
85
85
  obs, reward, terminated, truncated, info = env.step(actions)
@@ -87,6 +87,50 @@ while not env.all_done:
87
87
 
88
88
  Works identically for every simulator.
89
89
 
90
+ Or compose individual wrappers piecemeal:
91
+
92
+ ```python
93
+ from env_ssl_wrapper import compose_env
94
+
95
+ env = compose_env(
96
+ any_env,
97
+ ('tensor', dict(device = 'cpu')),
98
+ 'done_tracker',
99
+ )
100
+ ```
101
+
102
+ ## Research Wrappers
103
+
104
+ ### Memory Trace
105
+
106
+ A cheap memory that keeps track of an exponential moving average (EMA) of a state, or subset ([Eberhard et al., 2025](https://arxiv.org/abs/2503.15200)).
107
+
108
+ ```python
109
+ import gymnasium as gym
110
+ from env_ssl_wrapper import StandardizeEnvWrapper, MemoryTraceWrapper
111
+
112
+ env = StandardizeEnvWrapper(gym.make('LunarLander-v3'))
113
+ env = MemoryTraceWrapper(env, lambdas = (0.9, 0.99))
114
+
115
+ obs, info = env.reset()
116
+ # obs['obs'] -> (1, 8)
117
+ # obs['trace_0.9'] -> (1, 8)
118
+ # obs['trace_0.99'] -> (1, 8)
119
+ ```
120
+
121
+ Or on a subset of dictionary observations:
122
+
123
+ ```python
124
+ env = MemoryTraceWrapper(env, lambdas = (0.9, 0.99), keys = 'proprio')
125
+ # obs['proprio_trace_0.9'] -> (1, 2)
126
+ ```
127
+
128
+ Run the PPO benchmark on POMDP LunarLander:
129
+
130
+ ```bash
131
+ uv run test_memory_trace.py
132
+ ```
133
+
90
134
  ## Wrappers
91
135
 
92
136
  Pass wrappers as strings (default config) or `(name, dict)` tuples (custom config), in any order.
@@ -166,15 +210,6 @@ obs, info = env.reset()
166
210
  obs, reward, terminated, truncated, info = env.step(torch.randn(1, 1))
167
211
  ```
168
212
 
169
- ## Mock sims
170
-
171
- `env_ssl_wrapper.mocks` ships dependency-free stand-ins emulating each simulator's quirks (`GymnasiumMockEnv`, `IsaacMockEnv`, `DMControlMockEnv`, ...) for testing your code without installing the real sims.
172
-
173
- ```python
174
- from env_ssl_wrapper.mocks import IsaacMockEnv
175
- env = compose_env(IsaacMockEnv(), 'tensor', 'done_tracker')
176
- ```
177
-
178
213
  ## Multiprocessing
179
214
 
180
215
  Parallelize any single environment or factory into an autoresetting vector env:
@@ -191,5 +226,16 @@ with MultiprocessingVecEnv('CartPole-v1', num_envs = 8) as env:
191
226
 
192
227
  ```bash
193
228
  uv sync --extra test
194
- uv run pytest tests/test_real_envs.py
229
+ uv run pytest
230
+ ```
231
+
232
+ ## Citations
233
+
234
+ ```bibtex
235
+ @article{eberhard2025partially,
236
+ title = {Partially Observable Reinforcement Learning with Memory Traces},
237
+ author = {Onno Eberhard and Michael Muehlebach and Claire Vernade},
238
+ journal = {arXiv preprint arXiv:2503.15200},
239
+ year = {2025}
240
+ }
195
241
  ```
@@ -8,19 +8,18 @@ One line turns any simulator's environment — mujoco warp, isaac sim, pybullet,
8
8
  pip install env-ssl-wrapper
9
9
  ```
10
10
 
11
- ## Usage
11
+ ## Standardize
12
12
 
13
13
  ```python
14
14
  import torch
15
- from env_ssl_wrapper import compose_env
15
+ from env_ssl_wrapper import StandardizeEnvWrapper
16
16
 
17
- env = compose_env(
18
- any_env, # any env from any sim
19
- ('tensor', dict(device='cpu')), # wrap with whatever you need
20
- 'done_tracker',
21
- )
17
+ # one master wrapper to standardize any simulator environment
18
+
19
+ env = StandardizeEnvWrapper(any_env)
20
+
21
+ obs, info = env.reset() # torch.float32, batched
22
22
 
23
- obs, info = env.reset() # torch.float32, batched
24
23
  while not env.all_done:
25
24
  actions = torch.randint(0, 2, (8,))
26
25
  obs, reward, terminated, truncated, info = env.step(actions)
@@ -28,6 +27,50 @@ while not env.all_done:
28
27
 
29
28
  Works identically for every simulator.
30
29
 
30
+ Or compose individual wrappers piecemeal:
31
+
32
+ ```python
33
+ from env_ssl_wrapper import compose_env
34
+
35
+ env = compose_env(
36
+ any_env,
37
+ ('tensor', dict(device = 'cpu')),
38
+ 'done_tracker',
39
+ )
40
+ ```
41
+
42
+ ## Research Wrappers
43
+
44
+ ### Memory Trace
45
+
46
+ A cheap memory that keeps track of an exponential moving average (EMA) of a state, or subset ([Eberhard et al., 2025](https://arxiv.org/abs/2503.15200)).
47
+
48
+ ```python
49
+ import gymnasium as gym
50
+ from env_ssl_wrapper import StandardizeEnvWrapper, MemoryTraceWrapper
51
+
52
+ env = StandardizeEnvWrapper(gym.make('LunarLander-v3'))
53
+ env = MemoryTraceWrapper(env, lambdas = (0.9, 0.99))
54
+
55
+ obs, info = env.reset()
56
+ # obs['obs'] -> (1, 8)
57
+ # obs['trace_0.9'] -> (1, 8)
58
+ # obs['trace_0.99'] -> (1, 8)
59
+ ```
60
+
61
+ Or on a subset of dictionary observations:
62
+
63
+ ```python
64
+ env = MemoryTraceWrapper(env, lambdas = (0.9, 0.99), keys = 'proprio')
65
+ # obs['proprio_trace_0.9'] -> (1, 2)
66
+ ```
67
+
68
+ Run the PPO benchmark on POMDP LunarLander:
69
+
70
+ ```bash
71
+ uv run test_memory_trace.py
72
+ ```
73
+
31
74
  ## Wrappers
32
75
 
33
76
  Pass wrappers as strings (default config) or `(name, dict)` tuples (custom config), in any order.
@@ -107,15 +150,6 @@ obs, info = env.reset()
107
150
  obs, reward, terminated, truncated, info = env.step(torch.randn(1, 1))
108
151
  ```
109
152
 
110
- ## Mock sims
111
-
112
- `env_ssl_wrapper.mocks` ships dependency-free stand-ins emulating each simulator's quirks (`GymnasiumMockEnv`, `IsaacMockEnv`, `DMControlMockEnv`, ...) for testing your code without installing the real sims.
113
-
114
- ```python
115
- from env_ssl_wrapper.mocks import IsaacMockEnv
116
- env = compose_env(IsaacMockEnv(), 'tensor', 'done_tracker')
117
- ```
118
-
119
153
  ## Multiprocessing
120
154
 
121
155
  Parallelize any single environment or factory into an autoresetting vector env:
@@ -132,5 +166,16 @@ with MultiprocessingVecEnv('CartPole-v1', num_envs = 8) as env:
132
166
 
133
167
  ```bash
134
168
  uv sync --extra test
135
- uv run pytest tests/test_real_envs.py
169
+ uv run pytest
170
+ ```
171
+
172
+ ## Citations
173
+
174
+ ```bibtex
175
+ @article{eberhard2025partially,
176
+ title = {Partially Observable Reinforcement Learning with Memory Traces},
177
+ author = {Onno Eberhard and Michael Muehlebach and Claire Vernade},
178
+ journal = {arXiv preprint arXiv:2503.15200},
179
+ year = {2025}
180
+ }
136
181
  ```
@@ -0,0 +1,47 @@
1
+ from __future__ import annotations
2
+
3
+ import sys
4
+ from . import standardize
5
+ from .standardize import *
6
+
7
+ from .memory_trace import MemoryTraceWrapper
8
+
9
+ # Wire backwards-compatibility aliases in sys.modules and module globals
10
+ # so imports like `from env_ssl_wrapper.done_tracker_wrapper import DoneTrackerWrapper`
11
+ # or `import env_ssl_wrapper.mocks` continue to work without breaking.
12
+
13
+ _STANDARDIZE_SUBMODULES = (
14
+ 'adapters',
15
+ 'auto_batched_wrapper',
16
+ 'action_transform_wrapper',
17
+ 'done_tracker_wrapper',
18
+ 'episode_padding_wrapper',
19
+ 'flatten_obs_wrapper',
20
+ 'helpers',
21
+ 'image_wrapper',
22
+ 'mocks',
23
+ 'spaces',
24
+ 'standardize_wrapper',
25
+ 'tensor_wrapper',
26
+ 'time_limit_wrapper',
27
+ 'standardize_env_wrapper',
28
+ 'utils',
29
+ 'vector',
30
+ )
31
+
32
+ for _name in _STANDARDIZE_SUBMODULES:
33
+ _mod = getattr(standardize, _name, None)
34
+ if _mod is not None:
35
+ sys.modules[f'{__name__}.{_name}'] = _mod
36
+ globals()[_name] = _mod
37
+
38
+ __all__ = [
39
+ *standardize.__all__,
40
+ 'MemoryTraceWrapper',
41
+ ]
42
+
43
+ def __getattr__(name):
44
+ if hasattr(standardize, name):
45
+ return getattr(standardize, name)
46
+ raise AttributeError(f"module '{__name__}' has no attribute '{name}'")
47
+
@@ -0,0 +1,93 @@
1
+ from __future__ import annotations
2
+ from typing import Sequence
3
+
4
+ import torch
5
+ from torch import is_tensor
6
+ from torch_einops_utils import pad_right_ndim_to
7
+
8
+ from .standardize.helpers import (
9
+ TransformObservationWrapper,
10
+ exists,
11
+ default,
12
+ any_true,
13
+ )
14
+
15
+ # helper functions
16
+
17
+ def cast_tuple(val):
18
+ return val if isinstance(val, (tuple, list)) else (val,)
19
+
20
+ def calc_lerp_weight(lam, done, x):
21
+ if not exists(done) or not any_true(done):
22
+ return 1. - lam
23
+
24
+ # reset trace on done, else decay
25
+
26
+ if not is_tensor(done):
27
+ done = torch.as_tensor(done, device = x.device)
28
+
29
+ weight = torch.where(done, 1., 1. - lam).to(x)
30
+ return pad_right_ndim_to(weight, x.ndim)
31
+
32
+ # Partially Observable Reinforcement Learning with Memory Traces - Eberhard et al.
33
+ # https://arxiv.org/abs/2503.15200
34
+
35
+ class MemoryTraceWrapper(TransformObservationWrapper):
36
+
37
+ def __init__(
38
+ self,
39
+ env,
40
+ lambdas: float | Sequence[float] = (0.9, 0.99),
41
+ trace_key: str = 'trace',
42
+ obs_key: str = 'obs',
43
+ keys: str | Sequence[str] | None = None,
44
+ ):
45
+ super().__init__(env)
46
+
47
+ self.lambdas = tuple(float(l) for l in cast_tuple(lambdas))
48
+ assert all(0. <= l <= 1. for l in self.lambdas), f'lambdas must be within [0, 1], got {self.lambdas}'
49
+
50
+ self.trace_key = trace_key
51
+ self.obs_key = obs_key
52
+ self.keys = set(cast_tuple(keys)) if exists(keys) else None
53
+
54
+ self.traces = dict()
55
+
56
+ def trace_key_for(self, key, lam):
57
+ prefix = self.trace_key if key == self.obs_key else f'{key}_{self.trace_key}'
58
+ return f'{prefix}_{lam}' if len(self.lambdas) > 1 else prefix
59
+
60
+ def transform_obs(self, obs, done = None):
61
+ out = dict(obs) if isinstance(obs, dict) else {self.obs_key: obs}
62
+ target_keys = tuple(k for k in default(self.keys, tuple(out.keys())) if k in out)
63
+
64
+ for key in target_keys:
65
+ val = out.get(key)
66
+
67
+ if not exists(val):
68
+ continue
69
+
70
+ if not is_tensor(val):
71
+ val = torch.as_tensor(val)
72
+ out[key] = val
73
+
74
+ val_float = val.float() if not val.is_floating_point() else val
75
+
76
+ # init or lerp update traces
77
+ # z_t = λ * z_{t-1} + (1 - λ) * y_t
78
+
79
+ if key not in self.traces:
80
+ traces = [val_float.clone() for _ in self.lambdas]
81
+ else:
82
+ traces = [prev.lerp(val_float, calc_lerp_weight(lam, done, val_float)) for lam, prev in zip(self.lambdas, self.traces[key])]
83
+
84
+ self.traces[key] = traces
85
+
86
+ for lam, trace in zip(self.lambdas, traces):
87
+ out[self.trace_key_for(key, lam)] = trace
88
+
89
+ return out
90
+
91
+ def reset(self, **kwargs):
92
+ self.traces = dict()
93
+ return super().reset(**kwargs)
@@ -0,0 +1,166 @@
1
+ from __future__ import annotations
2
+
3
+ from . import (
4
+ adapters,
5
+ auto_batched_wrapper,
6
+ action_transform_wrapper,
7
+ done_tracker_wrapper,
8
+ episode_padding_wrapper,
9
+ flatten_obs_wrapper,
10
+ helpers,
11
+ image_wrapper,
12
+ mocks,
13
+ spaces,
14
+ standardize_wrapper,
15
+ tensor_wrapper,
16
+ time_limit_wrapper,
17
+ utils,
18
+ vector,
19
+ )
20
+
21
+ from .standardize_wrapper import StandardizeWrapper
22
+ from .image_wrapper import ImageObservationWrapper
23
+ from .auto_batched_wrapper import AutoBatchedWrapper
24
+ from .tensor_wrapper import TensorWrapper
25
+ from .action_transform_wrapper import ActionTransformWrapper
26
+ from .done_tracker_wrapper import DoneTrackerWrapper
27
+ from .flatten_obs_wrapper import FlattenObsWrapper
28
+ from .episode_padding_wrapper import EpisodePaddingWrapper
29
+ from .time_limit_wrapper import TimeLimitWrapper
30
+ from .vector import MultiprocessingVecEnv
31
+ from .standardize_env_wrapper import StandardizeEnvWrapper, StandardizeEnv, StandardizedEnv
32
+
33
+ from .adapters import (
34
+ BaseEnvAdapter,
35
+ get_adapter,
36
+ register_adapter,
37
+ MujocoWarpAdapter,
38
+ IsaacAdapter,
39
+ PyBulletAdapter,
40
+ DMControlAdapter,
41
+ PufferLibAdapter,
42
+ RoboticsAdapter,
43
+ GymnasiumAdapter,
44
+ LegacyGymAdapter,
45
+ DefaultAdapter,
46
+ )
47
+
48
+ from .spaces import (
49
+ InferredSpace,
50
+ infer_observation_space,
51
+ space_from_action_spec,
52
+ action_space_dim,
53
+ action_space_is_discrete,
54
+ action_space_is_box,
55
+ action_space_bounds,
56
+ action_dim_of,
57
+ obs_dim_of,
58
+ )
59
+
60
+ from .mocks import (
61
+ MockEnv,
62
+ GymnasiumMockEnv,
63
+ GymnasiumDiscreteMockEnv,
64
+ LegacyGymMockEnv,
65
+ PyBulletMockEnv,
66
+ DMControlMockEnv,
67
+ IsaacMockEnv,
68
+ AutoresetVectorMockEnv,
69
+ PufferVectorMockEnv,
70
+ PufferTensorMockEnv,
71
+ ManiSkillMockEnv,
72
+ BraxMockEnv,
73
+ MjxMockEnv,
74
+ MetaWorldMockEnv,
75
+ TrifingerMockEnv,
76
+ HabitatMockEnv,
77
+ TupleObsMockEnv,
78
+ JaxArray,
79
+ )
80
+
81
+ from .utils import wrap_env, compose_env
82
+ from .helpers import (
83
+ TransformObservationWrapper,
84
+ ObservationWrapper,
85
+ has_final_observation,
86
+ maybe_get_final_observation,
87
+ get_final_observation,
88
+ maybe_transform_final_observation,
89
+ )
90
+
91
+ __all__ = [
92
+ 'TransformObservationWrapper',
93
+ 'ObservationWrapper',
94
+ 'adapters',
95
+ 'auto_batched_wrapper',
96
+ 'action_transform_wrapper',
97
+ 'done_tracker_wrapper',
98
+ 'episode_padding_wrapper',
99
+ 'flatten_obs_wrapper',
100
+ 'helpers',
101
+ 'image_wrapper',
102
+ 'mocks',
103
+ 'spaces',
104
+ 'standardize_wrapper',
105
+ 'tensor_wrapper',
106
+ 'time_limit_wrapper',
107
+ 'utils',
108
+ 'vector',
109
+ 'StandardizeWrapper',
110
+ 'ImageObservationWrapper',
111
+ 'AutoBatchedWrapper',
112
+ 'TensorWrapper',
113
+ 'ActionTransformWrapper',
114
+ 'DoneTrackerWrapper',
115
+ 'FlattenObsWrapper',
116
+ 'EpisodePaddingWrapper',
117
+ 'TimeLimitWrapper',
118
+ 'MultiprocessingVecEnv',
119
+ 'StandardizeEnvWrapper',
120
+ 'StandardizeEnv',
121
+ 'BaseEnvAdapter',
122
+ 'get_adapter',
123
+ 'register_adapter',
124
+ 'MujocoWarpAdapter',
125
+ 'IsaacAdapter',
126
+ 'PyBulletAdapter',
127
+ 'DMControlAdapter',
128
+ 'PufferLibAdapter',
129
+ 'RoboticsAdapter',
130
+ 'GymnasiumAdapter',
131
+ 'LegacyGymAdapter',
132
+ 'DefaultAdapter',
133
+ 'InferredSpace',
134
+ 'infer_observation_space',
135
+ 'space_from_action_spec',
136
+ 'action_space_dim',
137
+ 'action_space_is_discrete',
138
+ 'action_space_is_box',
139
+ 'action_space_bounds',
140
+ 'action_dim_of',
141
+ 'obs_dim_of',
142
+ 'MockEnv',
143
+ 'GymnasiumMockEnv',
144
+ 'GymnasiumDiscreteMockEnv',
145
+ 'LegacyGymMockEnv',
146
+ 'PyBulletMockEnv',
147
+ 'DMControlMockEnv',
148
+ 'IsaacMockEnv',
149
+ 'AutoresetVectorMockEnv',
150
+ 'PufferVectorMockEnv',
151
+ 'PufferTensorMockEnv',
152
+ 'ManiSkillMockEnv',
153
+ 'BraxMockEnv',
154
+ 'MjxMockEnv',
155
+ 'MetaWorldMockEnv',
156
+ 'TrifingerMockEnv',
157
+ 'HabitatMockEnv',
158
+ 'TupleObsMockEnv',
159
+ 'JaxArray',
160
+ 'wrap_env',
161
+ 'compose_env',
162
+ 'has_final_observation',
163
+ 'maybe_get_final_observation',
164
+ 'get_final_observation',
165
+ 'maybe_transform_final_observation',
166
+ ]
@@ -175,7 +175,7 @@ class WrapperAdapter(BaseEnvAdapter):
175
175
  @property
176
176
  def autoresets(self) -> bool:
177
177
  val = first_existing(self.env, 'autoreset', 'autoresets', 'autoreset_mode')
178
- return bool(val) if exists(val) else self.inner_adapter.autoresets
178
+ return truthy_attr(val) if exists(val) else self.inner_adapter.autoresets
179
179
 
180
180
  @property
181
181
  def action_space(self):
@@ -202,19 +202,6 @@ class DMControlAdapter(BaseEnvAdapter):
202
202
  has_physics = exists(get_attr(env, 'physics')) and callable(get_attr(get_attr(env, 'physics'), 'render'))
203
203
  return is_dm_type or has_physics
204
204
 
205
- def step(self, action):
206
- out = self.env.step(action)
207
- if is_time_step(out):
208
- last = out.last() if callable(get_attr(out, 'last')) else out.step_type == 2
209
- return out.observation, out.reward, last, False, dict(discount = out.discount)
210
- return super().step(action)
211
-
212
- def reset(self, **kwargs):
213
- out = self.env.reset(**kwargs)
214
- if is_time_step(out):
215
- return out.observation, {}
216
- return super().reset(**kwargs)
217
-
218
205
  def seed(self, seed: int):
219
206
  random_state = get_attr(get_attr(self.env, 'task'), '_random')
220
207
  if exists(random_state) and callable(get_attr(random_state, 'seed')):
@@ -39,6 +39,12 @@ def maybe_expand_dim(x):
39
39
 
40
40
  return rearrange(arr, '... -> 1 ...')
41
41
 
42
+ if is_tensor(x):
43
+ return rearrange(x, '-> 1') if x.ndim == 0 else rearrange(x, '... -> 1 ...')
44
+
45
+ if isinstance(x, np.ndarray) and x.dtype.kind in 'biufc':
46
+ return rearrange(x, '-> 1') if x.ndim == 0 else rearrange(x, '... -> 1 ...')
47
+
42
48
  return tree_map(_expand, x)
43
49
 
44
50
  def is_integer_dtype(t):
@@ -12,6 +12,7 @@ from .helpers import (
12
12
  env_num_envs,
13
13
  exists,
14
14
  get_attr,
15
+ is_array,
15
16
  is_vectorized,
16
17
  to_numpy,
17
18
  )
@@ -25,6 +26,10 @@ def get_batch_size(tree) -> int | None:
25
26
  return None
26
27
 
27
28
  first = leaves[0]
29
+
30
+ if is_array(first):
31
+ return len(first) if first.ndim > 0 else None
32
+
28
33
  return len(first) if exists(get_attr(first, '__len__')) else None
29
34
 
30
35
  # classes
@@ -8,7 +8,7 @@ from torch import is_tensor
8
8
  from torch.utils._pytree import tree_map
9
9
  from einops import rearrange
10
10
 
11
- from .helpers import EnvWrapper, copy_leaf, dones_of, exists, is_vectorized, to_numpy
11
+ from .helpers import EnvWrapper, copy_leaf, dones_of, env_autoresets, exists, is_vectorized, to_numpy
12
12
 
13
13
  # helpers
14
14
 
@@ -63,6 +63,7 @@ class EpisodePaddingWrapper(EnvWrapper):
63
63
  def __init__(self, env):
64
64
  super().__init__(env)
65
65
  self.is_vector = is_vectorized(env)
66
+ self.autoreset = env_autoresets(env)
66
67
  self._last_obs = None
67
68
  self._is_done = None
68
69
  self._final_obs = None
@@ -83,12 +84,15 @@ class EpisodePaddingWrapper(EnvWrapper):
83
84
  dones = dones_of(terminated, truncated)
84
85
  mask = to_numpy(dones).astype(bool)
85
86
 
87
+ if self._is_done is None or len(self._is_done) != len(mask):
88
+ self._is_done = np.zeros(len(mask), dtype = bool)
89
+
90
+ if self.autoreset:
91
+ self._is_done &= mask
92
+
86
93
  if mask.any():
87
94
  assert exists(self._last_obs), 'environment needs reset before calling step. call env.reset() first'
88
95
 
89
- if self._is_done is None or len(self._is_done) != len(mask):
90
- self._is_done = np.zeros(len(mask), dtype = bool)
91
-
92
96
  newly = mask & ~self._is_done
93
97
  self._is_done |= mask
94
98
 
@@ -0,0 +1,75 @@
1
+ from __future__ import annotations
2
+
3
+ import numpy as np
4
+ import torch
5
+ from torch import is_tensor
6
+ from torch.utils._pytree import tree_flatten
7
+ from einops import rearrange
8
+
9
+ from .helpers import (
10
+ TransformObservationWrapper,
11
+ default,
12
+ is_array,
13
+ is_scalar,
14
+ is_vectorized,
15
+ )
16
+
17
+ # helpers
18
+
19
+ def flattenable(t):
20
+ if is_array(t) or is_scalar(t):
21
+ return True
22
+
23
+ try:
24
+ return np.asarray(t).dtype.kind not in 'USO'
25
+ except Exception:
26
+ return False
27
+
28
+ def flatten_leaf(t, is_vector = False):
29
+ if not is_array(t):
30
+ t = np.asarray(t)
31
+
32
+ if not is_vector:
33
+ if t.ndim == 0:
34
+ return rearrange(t, '-> 1')
35
+ return rearrange(t, '... -> (...)')
36
+
37
+ if t.ndim == 0:
38
+ return rearrange(t, '-> 1 1')
39
+
40
+ if t.ndim == 1:
41
+ return rearrange(t, 'b -> b 1')
42
+
43
+ return rearrange(t, 'b ... -> b (...)')
44
+
45
+ def concat_leaves(leaves, is_vector = False):
46
+ axis = -1 if is_vector else 0
47
+ leaves = [flatten_leaf(t, is_vector = is_vector) for t in leaves]
48
+
49
+ if is_tensor(leaves[0]):
50
+ return torch.cat(leaves, dim = axis)
51
+
52
+ return np.concatenate(leaves, axis = axis)
53
+
54
+ # class
55
+
56
+ class FlattenObsWrapper(TransformObservationWrapper):
57
+ def __init__(self, env, is_vector: bool | None = None):
58
+ super().__init__(env)
59
+ self.is_vector = default(is_vector, is_vectorized(env))
60
+
61
+ def transform_obs(self, obs):
62
+ if is_array(obs):
63
+ return obs
64
+
65
+ leaves, _ = tree_flatten(obs)
66
+ leaves = [t for t in leaves if flattenable(t)]
67
+
68
+ if len(leaves) == 0:
69
+ return obs
70
+
71
+ if len(leaves) == 1:
72
+ return flatten_leaf(leaves[0], is_vector = self.is_vector)
73
+
74
+ return concat_leaves(leaves, is_vector = self.is_vector)
75
+
@@ -1,5 +1,6 @@
1
1
  from __future__ import annotations
2
2
 
3
+ import inspect
3
4
  import numpy as np
4
5
  from torch import is_tensor
5
6
  from torch.utils._pytree import tree_map
@@ -62,6 +63,8 @@ def copy_leaf(x):
62
63
  return x
63
64
 
64
65
  def dones_of(terminated, truncated):
66
+ if not isinstance(terminated, (dict, list, tuple)):
67
+ return terminated | truncated
65
68
  return tree_map(lambda a, b: a | b, terminated, truncated)
66
69
 
67
70
  # environment probes
@@ -85,6 +88,33 @@ def env_render(env, height, width, camera = None):
85
88
  def is_vectorized(env) -> bool:
86
89
  return get_adapter(env).is_vectorized
87
90
 
91
+ def has_final_observation(info):
92
+ return isinstance(info, dict) and 'final_observation' in info
93
+
94
+ def maybe_get_final_observation(info):
95
+ if isinstance(info, dict):
96
+ return info.get('final_observation')
97
+
98
+ return None
99
+
100
+ def get_final_observation(info, default = None):
101
+ final_obs = maybe_get_final_observation(info)
102
+
103
+ if exists(final_obs):
104
+ return final_obs
105
+
106
+ if exists(default):
107
+ return default
108
+
109
+ raise KeyError("no 'final_observation' found in info")
110
+
111
+ def maybe_transform_final_observation(info, fn):
112
+ if not has_final_observation(info):
113
+ return info
114
+
115
+ info['final_observation'] = fn(info['final_observation'])
116
+ return info
117
+
88
118
  def mark_terminal_obs(info, obs, dones, is_vector):
89
119
  # single-env terminal contract — vector envs handled by EpisodePaddingWrapper
90
120
  if not is_vector and isinstance(info, dict) and 'final_observation' not in info and any_true(dones):
@@ -132,3 +162,60 @@ class EnvWrapper:
132
162
  if name.startswith('_'):
133
163
  raise AttributeError(f"attempted to get missing private attribute '{name}'")
134
164
  return getattr(self.env, name)
165
+
166
+ def accepts_done_param(fn):
167
+ try:
168
+ params = inspect.signature(fn).parameters
169
+ return 'done' in params or any(p.kind == inspect.Parameter.VAR_KEYWORD for p in params.values())
170
+ except (ValueError, TypeError):
171
+ return False
172
+
173
+ class TransformObservationWrapper(EnvWrapper):
174
+ """
175
+ Base observation wrapper that automatically handles:
176
+ - Calling transform_obs on observations in reset() and step()
177
+ - Detecting environment autoreset and passing `done` to stateful transforms
178
+ - Propagating transformed observations to info['final_observation']
179
+ """
180
+
181
+ def __init__(self, env):
182
+ super().__init__(env)
183
+ self.autoresets = env_autoresets(env)
184
+ self.takes_done = accepts_done_param(self.transform_obs)
185
+
186
+ def transform_obs(self, obs, done = None):
187
+ if hasattr(self, 'observation'):
188
+ return self.observation(obs)
189
+ return obs
190
+
191
+ def transform(self, obs, done = None):
192
+ return self.transform_obs(obs, done = done) if self.takes_done else self.transform_obs(obs)
193
+
194
+ def observation(self, obs):
195
+ return self.transform_obs(obs)
196
+
197
+ def reset(self, **kwargs):
198
+ obs, info = self.env.reset(**kwargs)
199
+ obs = self.transform_obs(obs)
200
+
201
+ if isinstance(info, dict) and 'final_observation' in info:
202
+ info['final_observation'] = self.transform_obs(info['final_observation'])
203
+
204
+ return obs, info
205
+
206
+ def step(self, action):
207
+ obs, reward, terminated, truncated, info = self.env.step(action)
208
+ done = dones_of(terminated, truncated) if self.autoresets else None
209
+
210
+ out = self.transform_obs(obs, done = done) if self.takes_done else self.transform_obs(obs)
211
+
212
+ if isinstance(info, dict) and 'final_observation' in info:
213
+ if not self.takes_done:
214
+ info['final_observation'] = self.transform_obs(info['final_observation'])
215
+ elif not self.autoresets:
216
+ info['final_observation'] = out
217
+
218
+ return out, reward, terminated, truncated, info
219
+
220
+ ObservationWrapper = TransformObservationWrapper
221
+
@@ -0,0 +1,66 @@
1
+ from __future__ import annotations
2
+
3
+ import torch
4
+ from .helpers import EnvWrapper, exists, default
5
+ from .utils import compose_env
6
+
7
+ class StandardizeEnvWrapper(EnvWrapper):
8
+ """
9
+ Master environment wrapper that turns any simulator environment
10
+ into a torch-native, batched, standardized RL environment.
11
+ """
12
+
13
+ def __init__(
14
+ self,
15
+ env,
16
+ device: str | torch.device = 'cpu',
17
+ auto_batch: bool = True,
18
+ pad_episodes: bool = True,
19
+ done_tracker: bool = True,
20
+ flatten_obs: bool = False,
21
+ action_transform: bool = False,
22
+ max_timesteps: int | None = None,
23
+ image_size: tuple[int, int] | None = None,
24
+ lambdas: tuple[float, ...] | list[float] | None = None,
25
+ **kwargs
26
+ ):
27
+ wrappers = []
28
+
29
+ if exists(image_size):
30
+ wrappers.append(('image', dict(image_size = image_size)))
31
+
32
+ if action_transform:
33
+ wrappers.append('action_transform')
34
+
35
+ if exists(max_timesteps):
36
+ wrappers.append(('time_limit', dict(max_timesteps = max_timesteps)))
37
+
38
+ if auto_batch:
39
+ wrappers.append('auto_batch')
40
+
41
+ if exists(device):
42
+ wrappers.append(('tensor', dict(device = device)))
43
+
44
+ if done_tracker:
45
+ wrappers.append('done_tracker')
46
+
47
+ if exists(lambdas):
48
+ wrappers.append(('memory_trace', dict(lambdas = lambdas, **kwargs)))
49
+
50
+ if flatten_obs:
51
+ wrappers.append('flatten_obs')
52
+
53
+ wrapped = compose_env(env, *wrappers, pad_episodes = pad_episodes)
54
+
55
+ super().__init__(wrapped)
56
+
57
+ def reset(self, **kwargs):
58
+ return self.env.reset(**kwargs)
59
+
60
+ def step(self, action):
61
+ return self.env.step(action)
62
+
63
+ # aliases
64
+
65
+ StandardizeEnv = StandardizeEnvWrapper
66
+ StandardizedEnv = StandardizeEnvWrapper
@@ -11,43 +11,49 @@ from .helpers import EnvWrapper, exists, get_attr, is_scalar
11
11
 
12
12
  # helpers
13
13
 
14
+ def to_torch_leaf(t, device, cast_obs_to_float = True):
15
+ if not is_tensor(t):
16
+ if isinstance(t, np.ndarray):
17
+ t = from_numpy(t.copy())
18
+ elif is_scalar(t):
19
+ t = tensor(t)
20
+ elif exists(get_attr(t, '__array__')):
21
+ t = from_numpy(np.asarray(t))
22
+ else:
23
+ return t
24
+
25
+ dtype = torch.float32 if cast_obs_to_float and t.dtype != torch.bool else t.dtype
26
+ return t.to(device = device, dtype = dtype)
27
+
14
28
  def numpy_to_torch(x, device, cast_obs_to_float = True):
15
29
  # numpy / scalars / foreign array-likes to torch; float32 unless disabled
30
+ if not isinstance(x, (dict, list, tuple)):
31
+ return to_torch_leaf(x, device, cast_obs_to_float)
16
32
 
17
- def _to_torch(t):
18
- if not is_tensor(t):
19
- if isinstance(t, np.ndarray):
20
- t = from_numpy(np.array(t))
21
- elif is_scalar(t):
22
- t = tensor(t)
23
- elif exists(get_attr(t, '__array__')):
24
- t = from_numpy(np.array(t))
25
- else:
26
- return t
27
-
28
- dtype = torch.float32 if cast_obs_to_float and t.dtype != torch.bool else t.dtype
29
- return t.to(device = device, dtype = dtype)
30
- return tree_map(_to_torch, x)
33
+ return tree_map(partial(to_torch_leaf, device = device, cast_obs_to_float = cast_obs_to_float), x)
31
34
 
32
- def torch_to_numpy(x):
33
- # torch to numpy; 0-dim collapses to scalar, float64 → float32
35
+ def to_numpy_leaf(t):
36
+ if is_tensor(t):
37
+ t = t.detach().cpu().numpy()
38
+ elif is_scalar(t):
39
+ t = np.asarray(t)
40
+ else:
41
+ return t
34
42
 
35
- def _to_numpy(t):
36
- if is_tensor(t):
37
- t = t.detach().cpu().numpy()
38
- elif is_scalar(t):
39
- t = np.asarray(t)
40
- else:
41
- return t
43
+ if t.ndim == 0:
44
+ return t.item()
42
45
 
43
- if t.ndim == 0:
44
- return t.item()
46
+ if t.dtype == np.float64:
47
+ t = t.astype(np.float32)
45
48
 
46
- if t.dtype == np.float64:
47
- t = t.astype(np.float32)
49
+ return t
48
50
 
49
- return t
50
- return tree_map(_to_numpy, x)
51
+ def torch_to_numpy(x):
52
+ # torch to numpy; 0-dim collapses to scalar, float64 → float32
53
+ if not isinstance(x, (dict, list, tuple)):
54
+ return to_numpy_leaf(x)
55
+
56
+ return tree_map(to_numpy_leaf, x)
51
57
 
52
58
  # rewards float32, dones bool
53
59
 
@@ -96,6 +102,12 @@ class TensorWrapper(EnvWrapper):
96
102
 
97
103
  return obs, info
98
104
 
105
+ def to_contract(self, t, to_float = False):
106
+ if not isinstance(t, (dict, list, tuple)):
107
+ leaf = to_torch_leaf(t, self.device, cast_obs_to_float = False)
108
+ return contract(leaf, to_float = to_float)
109
+ return contract(self.cast(t), to_float = to_float)
110
+
99
111
  def step(self, action):
100
112
  action = torch_to_numpy(action) if self.convert_in else action
101
113
  obs, reward, terminated, truncated, info = self.env.step(action)
@@ -104,9 +116,9 @@ class TensorWrapper(EnvWrapper):
104
116
  return obs, reward, terminated, truncated, info
105
117
 
106
118
  obs = self.cast(obs)
107
- reward = contract(self.cast(reward), to_float = True)
108
- terminated = contract(self.cast(terminated))
109
- truncated = contract(self.cast(truncated))
119
+ reward = self.to_contract(reward, to_float = True)
120
+ terminated = self.to_contract(terminated)
121
+ truncated = self.to_contract(truncated)
110
122
  self.cast_info(info)
111
123
 
112
124
  return obs, reward, terminated, truncated, info
@@ -26,8 +26,15 @@ def back_to_like(t, numpy_arr):
26
26
  return numpy_arr
27
27
 
28
28
  class TimeLimitWrapper(EnvWrapper):
29
- # caps episode length at max_timesteps, setting truncated = True for
30
- # capped envs (vectorized and single alike); timers reset per episode
29
+ """
30
+ caps episode length at max_timesteps, setting truncated = True for
31
+ capped envs (vectorized and single alike); timers reset per episode.
32
+
33
+ Note on vector environments:
34
+ Marking a slot as truncated signals the time limit to downstream consumers,
35
+ but does not stop the underlying simulator slot if it runs continuously or
36
+ autoresets. Downstream loops should check `truncated` to reset consumer state.
37
+ """
31
38
 
32
39
  def __init__(self, env, max_timesteps):
33
40
  super().__init__(env)
@@ -27,14 +27,27 @@ WRAPPERS = dict(
27
27
 
28
28
  def parse_wrapper(wrapper):
29
29
  if isinstance(wrapper, str):
30
- if wrapper not in WRAPPERS:
31
- raise ValueError(f'unknown wrapper {wrapper!r} — choose from {sorted(WRAPPERS)}')
32
-
33
- wrapper = WRAPPERS[wrapper]
30
+ if wrapper == 'memory_trace':
31
+ from ..memory_trace import MemoryTraceWrapper
32
+ wrapper = MemoryTraceWrapper
33
+ elif wrapper in ('standardize_env', 'master'):
34
+ from .standardize_env_wrapper import StandardizeEnvWrapper
35
+ wrapper = StandardizeEnvWrapper
36
+ elif wrapper not in WRAPPERS:
37
+ raise ValueError(f'unknown wrapper {wrapper!r} — choose from {sorted([*WRAPPERS, "memory_trace", "standardize_env"])}')
38
+ else:
39
+ wrapper = WRAPPERS[wrapper]
34
40
 
35
41
  if isinstance(wrapper, tuple):
36
42
  name, kwargs = wrapper
37
- wrapper = partial(WRAPPERS.get(name, name), **kwargs)
43
+ if name == 'memory_trace':
44
+ from ..memory_trace import MemoryTraceWrapper
45
+ wrapper = partial(MemoryTraceWrapper, **kwargs)
46
+ elif name in ('standardize_env', 'master'):
47
+ from .standardize_env_wrapper import StandardizeEnvWrapper
48
+ wrapper = partial(StandardizeEnvWrapper, **kwargs)
49
+ else:
50
+ wrapper = partial(WRAPPERS.get(name, name), **kwargs)
38
51
 
39
52
  elif isinstance(wrapper, dict):
40
53
  raise ValueError("wrapper kwargs must be passed as (name, kwargs), e.g. ('tensor', dict(device = 'cpu'))")
@@ -63,6 +76,17 @@ def compose_env(env, *wrappers, pad_episodes: bool = True):
63
76
  funcs.insert(1, EpisodePaddingWrapper)
64
77
  classes.insert(1, EpisodePaddingWrapper)
65
78
 
79
+ from ..memory_trace import MemoryTraceWrapper
80
+ if MemoryTraceWrapper in classes and TensorWrapper in classes:
81
+ idx_mem = classes.index(MemoryTraceWrapper)
82
+ idx_ten = classes.index(TensorWrapper)
83
+ if idx_mem < idx_ten:
84
+ f = funcs.pop(idx_mem)
85
+ c = classes.pop(idx_mem)
86
+ idx_ten = classes.index(TensorWrapper)
87
+ funcs.insert(idx_ten + 1, f)
88
+ classes.insert(idx_ten + 1, c)
89
+
66
90
  assert len(set(classes)) == len(classes), 'duplicate wrappers found'
67
91
 
68
92
  for func in funcs:
@@ -6,7 +6,7 @@ import torch
6
6
  from torch import is_tensor
7
7
  from torch.utils._pytree import tree_flatten, tree_map, tree_structure, tree_unflatten
8
8
 
9
- from .helpers import any_true, dones_of, exists, get_attr, instantiate_env, safe_close
9
+ from .helpers import any_true, dones_of, exists, get_attr, instantiate_env, safe_close, truthy_attr
10
10
  from .spaces import action_dim_of
11
11
  from .standardize_wrapper import StandardizeWrapper
12
12
 
@@ -23,6 +23,20 @@ def _stack_leaves(leaves):
23
23
  return np.stack(leaves)
24
24
 
25
25
  def _stack_trees(trees):
26
+ first = trees[0]
27
+
28
+ if is_tensor(first):
29
+ return torch.stack(trees)
30
+
31
+ if isinstance(first, np.ndarray):
32
+ return np.stack(trees)
33
+
34
+ if isinstance(first, dict):
35
+ return {key: _stack_trees([t[key] for t in trees]) for key in first}
36
+
37
+ if isinstance(first, tuple):
38
+ return tuple(_stack_trees([t[i] for t in trees]) for i in range(len(first)))
39
+
26
40
  leaves = [tree_flatten(tree)[0] for tree in trees]
27
41
  stacked = [_stack_leaves(col) for col in zip(*leaves)]
28
42
  return tree_unflatten(stacked, tree_structure(trees[0]))
@@ -57,7 +71,10 @@ def _exec(env, cmd, payload):
57
71
  obs, reward, terminated, truncated, info = env.step(payload)
58
72
 
59
73
  if any_true(dones_of(terminated, truncated)):
60
- obs, final_obs = env.reset()[0], obs
74
+ if truthy_attr(get_attr(env.adapter, 'autoresets')):
75
+ final_obs = info.get('final_observation', obs)
76
+ else:
77
+ obs, final_obs = env.reset()[0], obs
61
78
  else:
62
79
  final_obs = None
63
80
 
@@ -165,15 +182,34 @@ def _shutdown(conns, procs):
165
182
  for conn in conns:
166
183
  try:
167
184
  conn.send(('close', None))
168
- except (EOFError, BrokenPipeError):
185
+ except (EOFError, BrokenPipeError, OSError):
169
186
  pass
170
187
 
171
188
  for proc in procs:
172
- proc.join(timeout = 2)
189
+ if proc.is_alive():
190
+ proc.join(timeout = 0.2)
173
191
 
174
192
  if proc.is_alive():
175
193
  proc.terminate()
176
194
 
195
+ for conn in conns:
196
+ try:
197
+ conn.close()
198
+ except Exception:
199
+ pass
200
+
201
+ def _split_actions(actions, num_envs):
202
+ if isinstance(actions, dict):
203
+ assert all(len(v) == num_envs for v in actions.values()), f'expected {num_envs} actions per key'
204
+ return [{k: v[i] for k, v in actions.items()} for i in range(num_envs)]
205
+
206
+ if isinstance(actions, tuple) and not is_tensor(actions) and not isinstance(actions, np.ndarray):
207
+ assert all(len(v) == num_envs for v in actions), f'expected {num_envs} actions per element'
208
+ return [tuple(elem[i] for elem in actions) for i in range(num_envs)]
209
+
210
+ assert len(actions) == num_envs, f'expected {num_envs} actions, but got {len(actions)}'
211
+ return actions
212
+
177
213
  # class
178
214
 
179
215
  class MultiprocessingVecEnv:
@@ -260,7 +296,7 @@ class MultiprocessingVecEnv:
260
296
  return _stack_trees([obs for obs, _ in results]), _merge_infos([info for _, info in results])
261
297
 
262
298
  def step(self, actions):
263
- assert len(actions) == self.num_envs, f'expected {self.num_envs} actions, but got {len(actions)}'
299
+ actions = _split_actions(actions, self.num_envs)
264
300
 
265
301
  for conn, action in zip(self._conns, actions):
266
302
  _safe_send(conn, ('step', action))
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "env-ssl-wrapper"
3
- version = "0.3.0"
3
+ version = "0.4.0"
4
4
  description = "One torch-native interface for any MDP environment"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
@@ -39,7 +39,9 @@ Homepage = "https://pypi.org/project/env-ssl-wrapper/"
39
39
  Repository = "https://codeberg.org/lucidrains/env-ssl-wrapper"
40
40
 
41
41
  [project.optional-dependencies]
42
- examples = []
42
+ examples = [
43
+ "fire",
44
+ ]
43
45
  test = [
44
46
  "pytest",
45
47
  "gymnasium>=0.29.1",
@@ -66,10 +68,6 @@ testpaths = [
66
68
  requires = ["hatchling"]
67
69
  build-backend = "hatchling.build"
68
70
 
69
- [tool.rye]
70
- managed = true
71
- dev-dependencies = []
72
-
73
71
  [tool.hatch.metadata]
74
72
  allow-direct-references = true
75
73
 
@@ -1,59 +0,0 @@
1
- from .standardize_wrapper import StandardizeWrapper
2
- from .image_wrapper import ImageObservationWrapper
3
- from .auto_batched_wrapper import AutoBatchedWrapper
4
- from .tensor_wrapper import TensorWrapper
5
- from .action_transform_wrapper import ActionTransformWrapper
6
- from .done_tracker_wrapper import DoneTrackerWrapper
7
- from .flatten_obs_wrapper import FlattenObsWrapper
8
- from .episode_padding_wrapper import EpisodePaddingWrapper
9
- from .time_limit_wrapper import TimeLimitWrapper
10
- from .vector import MultiprocessingVecEnv
11
-
12
- from .adapters import (
13
- BaseEnvAdapter,
14
- get_adapter,
15
- register_adapter,
16
- MujocoWarpAdapter,
17
- IsaacAdapter,
18
- PyBulletAdapter,
19
- DMControlAdapter,
20
- PufferLibAdapter,
21
- RoboticsAdapter,
22
- GymnasiumAdapter,
23
- LegacyGymAdapter,
24
- DefaultAdapter,
25
- )
26
- from .spaces import (
27
- InferredSpace,
28
- infer_observation_space,
29
- space_from_action_spec,
30
- action_space_dim,
31
- action_space_is_discrete,
32
- action_space_is_box,
33
- action_space_bounds,
34
- action_dim_of,
35
- obs_dim_of,
36
- )
37
-
38
- from .mocks import (
39
- MockEnv,
40
- GymnasiumMockEnv,
41
- GymnasiumDiscreteMockEnv,
42
- LegacyGymMockEnv,
43
- PyBulletMockEnv,
44
- DMControlMockEnv,
45
- IsaacMockEnv,
46
- AutoresetVectorMockEnv,
47
- PufferVectorMockEnv,
48
- PufferTensorMockEnv,
49
- ManiSkillMockEnv,
50
- BraxMockEnv,
51
- MjxMockEnv,
52
- MetaWorldMockEnv,
53
- TrifingerMockEnv,
54
- HabitatMockEnv,
55
- TupleObsMockEnv,
56
- JaxArray
57
- )
58
-
59
- from .utils import wrap_env, compose_env
@@ -1,79 +0,0 @@
1
- from __future__ import annotations
2
-
3
- import numpy as np
4
- import torch
5
- from torch import is_tensor
6
- from torch.utils._pytree import tree_flatten
7
- from einops import rearrange
8
-
9
- from .helpers import EnvWrapper, is_array, is_scalar
10
-
11
- # helpers
12
-
13
- def flattenable(t):
14
- if is_array(t) or is_scalar(t):
15
- return True
16
-
17
- try:
18
- return np.asarray(t).dtype.kind not in 'USO'
19
- except Exception:
20
- return False
21
-
22
- def flatten_leaf(t):
23
- if not is_array(t):
24
- t = np.asarray(t)
25
-
26
- if t.ndim == 0:
27
- return rearrange(t, '-> 1 1')
28
-
29
- if t.ndim == 1:
30
- return rearrange(t, 'b -> b 1')
31
-
32
- return rearrange(t, 'b ... -> b (...)')
33
-
34
- def concat_leaves(leaves):
35
- axis, flatten = -1, True
36
-
37
- if leaves[0].ndim == 1 and not all(len(t) == len(leaves[0]) for t in leaves):
38
- axis, flatten = 0, False
39
-
40
- leaves = [flatten_leaf(t) for t in leaves] if flatten else leaves
41
-
42
- if is_tensor(leaves[0]):
43
- return torch.cat(leaves, dim = axis)
44
-
45
- return np.concatenate(leaves, axis = axis)
46
-
47
- # class
48
-
49
- class FlattenObsWrapper(EnvWrapper):
50
- def __init__(self, env):
51
- super().__init__(env)
52
-
53
- def observation(self, obs):
54
- if is_array(obs):
55
- return obs
56
-
57
- leaves, _ = tree_flatten(obs)
58
- leaves = [t for t in leaves if flattenable(t)]
59
-
60
- if len(leaves) == 0:
61
- return obs
62
-
63
- if len(leaves) == 1:
64
- return flatten_leaf(leaves[0])
65
-
66
- return concat_leaves(leaves)
67
-
68
- def final_observation(self, info):
69
- if isinstance(info, dict) and 'final_observation' in info:
70
- info['final_observation'] = self.observation(info['final_observation'])
71
- return info
72
-
73
- def reset(self, **kwargs):
74
- obs, info = self.env.reset(**kwargs)
75
- return self.observation(obs), self.final_observation(info)
76
-
77
- def step(self, action):
78
- obs, reward, terminated, truncated, info = self.env.step(action)
79
- return self.observation(obs), reward, terminated, truncated, self.final_observation(info)
File without changes