env-ssl-wrapper 0.3.0__tar.gz → 0.4.1__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 (27) hide show
  1. {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.1}/PKG-INFO +65 -19
  2. {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.1}/README.md +63 -18
  3. env_ssl_wrapper-0.4.1/env_ssl_wrapper/__init__.py +47 -0
  4. env_ssl_wrapper-0.4.1/env_ssl_wrapper/memory_trace.py +93 -0
  5. env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize/__init__.py +166 -0
  6. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/adapters.py +1 -14
  7. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/auto_batched_wrapper.py +22 -3
  8. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/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.1/env_ssl_wrapper/standardize}/episode_padding_wrapper.py +27 -8
  10. env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize/flatten_obs_wrapper.py +75 -0
  11. env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize/helpers.py +313 -0
  12. env_ssl_wrapper-0.4.1/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.1/env_ssl_wrapper/standardize}/standardize_wrapper.py +0 -1
  14. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/tensor_wrapper.py +65 -36
  15. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/time_limit_wrapper.py +9 -2
  16. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/utils.py +33 -7
  17. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/vector.py +38 -19
  18. {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.1}/pyproject.toml +4 -6
  19. env_ssl_wrapper-0.3.0/env_ssl_wrapper/__init__.py +0 -59
  20. env_ssl_wrapper-0.3.0/env_ssl_wrapper/flatten_obs_wrapper.py +0 -79
  21. env_ssl_wrapper-0.3.0/env_ssl_wrapper/helpers.py +0 -134
  22. {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.1}/.gitignore +0 -0
  23. {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.1}/LICENSE +0 -0
  24. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/action_transform_wrapper.py +0 -0
  25. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/image_wrapper.py +0 -0
  26. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/mocks.py +0 -0
  27. {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/spaces.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.1
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')):
@@ -7,7 +7,18 @@ import torch
7
7
  from torch.utils._pytree import tree_map
8
8
  from einops import rearrange
9
9
 
10
- from .helpers import EnvWrapper, default, exists, first_existing, get_attr, is_array, is_scalar, is_tensor, is_vectorized
10
+ from .helpers import (
11
+ FINAL_OBSERVATION_KEYS,
12
+ EnvWrapper,
13
+ default,
14
+ exists,
15
+ first_existing,
16
+ get_attr,
17
+ is_array,
18
+ is_scalar,
19
+ is_tensor,
20
+ is_vectorized,
21
+ )
11
22
  from .spaces import space_from_action_spec
12
23
 
13
24
  # helper functions
@@ -39,6 +50,12 @@ def maybe_expand_dim(x):
39
50
 
40
51
  return rearrange(arr, '... -> 1 ...')
41
52
 
53
+ if is_tensor(x):
54
+ return rearrange(x, '-> 1') if x.ndim == 0 else rearrange(x, '... -> 1 ...')
55
+
56
+ if isinstance(x, np.ndarray) and x.dtype.kind in 'biufc':
57
+ return rearrange(x, '-> 1') if x.ndim == 0 else rearrange(x, '... -> 1 ...')
58
+
42
59
  return tree_map(_expand, x)
43
60
 
44
61
  def is_integer_dtype(t):
@@ -196,7 +213,9 @@ class AutoBatchedWrapper(EnvWrapper):
196
213
 
197
214
  obs, reward, terminated, truncated, info = *maybe_expand_dim(out[:4]), out[4]
198
215
 
199
- if isinstance(info, dict) and 'final_observation' in info:
200
- info['final_observation'] = maybe_expand_dim(info['final_observation'])
216
+ if isinstance(info, dict):
217
+ for key in FINAL_OBSERVATION_KEYS:
218
+ if key in info:
219
+ info[key] = maybe_expand_dim(info[key])
201
220
 
202
221
  return obs, reward, terminated, truncated, info
@@ -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