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.
- {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.0}/PKG-INFO +65 -19
- {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.0}/README.md +63 -18
- env_ssl_wrapper-0.4.0/env_ssl_wrapper/__init__.py +47 -0
- env_ssl_wrapper-0.4.0/env_ssl_wrapper/memory_trace.py +93 -0
- env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize/__init__.py +166 -0
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/adapters.py +1 -14
- {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
- {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
- {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
- env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize/flatten_obs_wrapper.py +75 -0
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/helpers.py +87 -0
- env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize/standardize_env_wrapper.py +66 -0
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/tensor_wrapper.py +44 -32
- {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
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/utils.py +29 -5
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/vector.py +41 -5
- {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.0}/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-0.4.0}/.gitignore +0 -0
- {env_ssl_wrapper-0.3.0 → env_ssl_wrapper-0.4.0}/LICENSE +0 -0
- {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
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/image_wrapper.py +0 -0
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/mocks.py +0 -0
- {env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/spaces.py +0 -0
- {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
|
+
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
|
-
##
|
|
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')):
|
|
@@ -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
|
-
|
|
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
|
|
33
|
-
|
|
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
|
-
|
|
36
|
-
|
|
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
|
-
|
|
44
|
-
|
|
46
|
+
if t.dtype == np.float64:
|
|
47
|
+
t = t.astype(np.float32)
|
|
45
48
|
|
|
46
|
-
|
|
47
|
-
t = t.astype(np.float32)
|
|
49
|
+
return t
|
|
48
50
|
|
|
49
|
-
|
|
50
|
-
|
|
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 =
|
|
108
|
-
terminated =
|
|
109
|
-
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
|
-
|
|
30
|
-
|
|
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)
|
{env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/utils.py
RENAMED
|
@@ -27,14 +27,27 @@ WRAPPERS = dict(
|
|
|
27
27
|
|
|
28
28
|
def parse_wrapper(wrapper):
|
|
29
29
|
if isinstance(wrapper, str):
|
|
30
|
-
if wrapper
|
|
31
|
-
|
|
32
|
-
|
|
33
|
-
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
|
-
|
|
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
|
-
|
|
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.
|
|
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
|
-
|
|
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
|
+
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
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{env_ssl_wrapper-0.3.0/env_ssl_wrapper → env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize}/mocks.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|