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.
- {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.1}/PKG-INFO +65 -19
- {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.1}/README.md +63 -18
- env_ssl_wrapper-0.4.1/env_ssl_wrapper/__init__.py +47 -0
- env_ssl_wrapper-0.4.1/env_ssl_wrapper/memory_trace.py +93 -0
- env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize/__init__.py +166 -0
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/adapters.py +1 -14
- {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
- {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
- {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
- env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize/flatten_obs_wrapper.py +75 -0
- env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize/helpers.py +313 -0
- env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize/standardize_env_wrapper.py +66 -0
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/standardize_wrapper.py +0 -1
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/tensor_wrapper.py +65 -36
- {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
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/utils.py +33 -7
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/vector.py +38 -19
- {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.1}/pyproject.toml +4 -6
- env_ssl_wrapper-0.3.0/env_ssl_wrapper/__init__.py +0 -59
- env_ssl_wrapper-0.3.0/env_ssl_wrapper/flatten_obs_wrapper.py +0 -79
- env_ssl_wrapper-0.3.0/env_ssl_wrapper/helpers.py +0 -134
- {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.1}/.gitignore +0 -0
- {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.1}/LICENSE +0 -0
- {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
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/image_wrapper.py +0 -0
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.1/env_ssl_wrapper/standardize}/mocks.py +0 -0
- {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
|
+
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
|
-
##
|
|
71
|
+
## Standardize
|
|
71
72
|
|
|
72
73
|
```python
|
|
73
74
|
import torch
|
|
74
|
-
from env_ssl_wrapper import
|
|
75
|
+
from env_ssl_wrapper import StandardizeEnvWrapper
|
|
75
76
|
|
|
76
|
-
|
|
77
|
-
|
|
78
|
-
|
|
79
|
-
|
|
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
|
|
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
|
-
##
|
|
11
|
+
## Standardize
|
|
12
12
|
|
|
13
13
|
```python
|
|
14
14
|
import torch
|
|
15
|
-
from env_ssl_wrapper import
|
|
15
|
+
from env_ssl_wrapper import StandardizeEnvWrapper
|
|
16
16
|
|
|
17
|
-
|
|
18
|
-
|
|
19
|
-
|
|
20
|
-
|
|
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
|
|
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
|
|
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
|
|
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)
|
|
200
|
-
|
|
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
|