env-ssl-wrapper 0.4.0__tar.gz → 0.4.2__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.4.0 → env_ssl_wrapper-0.4.2}/PKG-INFO +5 -2
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/README.md +1 -1
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/action_transform_wrapper.py +8 -3
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/adapters.py +13 -41
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/auto_batched_wrapper.py +16 -3
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/done_tracker_wrapper.py +2 -20
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/episode_padding_wrapper.py +36 -24
- env_ssl_wrapper-0.4.2/env_ssl_wrapper/standardize/helpers.py +363 -0
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/image_wrapper.py +39 -9
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/standardize_env_wrapper.py +8 -7
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/standardize_wrapper.py +1 -34
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/tensor_wrapper.py +21 -4
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/time_limit_wrapper.py +18 -3
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/utils.py +20 -20
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/vector.py +15 -34
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/pyproject.toml +4 -1
- env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize/helpers.py +0 -221
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/.gitignore +0 -0
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/LICENSE +0 -0
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/__init__.py +0 -0
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/memory_trace.py +0 -0
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/__init__.py +0 -0
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/flatten_obs_wrapper.py +0 -0
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/mocks.py +0 -0
- {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/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.4.
|
|
3
|
+
Version: 0.4.2
|
|
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
|
|
@@ -47,14 +47,17 @@ Requires-Dist: fire; extra == 'examples'
|
|
|
47
47
|
Provides-Extra: test
|
|
48
48
|
Requires-Dist: box2d-py>=2.3.5; extra == 'test'
|
|
49
49
|
Requires-Dist: dm-control; extra == 'test'
|
|
50
|
+
Requires-Dist: gym-pusht; extra == 'test'
|
|
50
51
|
Requires-Dist: gymnasium-robotics; extra == 'test'
|
|
51
52
|
Requires-Dist: gymnasium>=0.29.1; extra == 'test'
|
|
52
53
|
Requires-Dist: imageio; extra == 'test'
|
|
54
|
+
Requires-Dist: mean-conc-beta; extra == 'test'
|
|
53
55
|
Requires-Dist: mujoco; extra == 'test'
|
|
54
56
|
Requires-Dist: pillow; extra == 'test'
|
|
55
57
|
Requires-Dist: populora; extra == 'test'
|
|
56
58
|
Requires-Dist: pybullet; extra == 'test'
|
|
57
59
|
Requires-Dist: pygame; extra == 'test'
|
|
60
|
+
Requires-Dist: pymunk<7; extra == 'test'
|
|
58
61
|
Requires-Dist: pytest; extra == 'test'
|
|
59
62
|
Description-Content-Type: text/markdown
|
|
60
63
|
|
|
@@ -81,7 +84,7 @@ env = StandardizeEnvWrapper(any_env)
|
|
|
81
84
|
obs, info = env.reset() # torch.float32, batched
|
|
82
85
|
|
|
83
86
|
while not env.all_done:
|
|
84
|
-
actions = torch.randint(0, 2, (
|
|
87
|
+
actions = torch.randint(0, 2, (1,))
|
|
85
88
|
obs, reward, terminated, truncated, info = env.step(actions)
|
|
86
89
|
```
|
|
87
90
|
|
|
@@ -21,7 +21,7 @@ env = StandardizeEnvWrapper(any_env)
|
|
|
21
21
|
obs, info = env.reset() # torch.float32, batched
|
|
22
22
|
|
|
23
23
|
while not env.all_done:
|
|
24
|
-
actions = torch.randint(0, 2, (
|
|
24
|
+
actions = torch.randint(0, 2, (1,))
|
|
25
25
|
obs, reward, terminated, truncated, info = env.step(actions)
|
|
26
26
|
```
|
|
27
27
|
|
|
@@ -97,9 +97,14 @@ class ActionTransformWrapper(EnvWrapper):
|
|
|
97
97
|
if was_scalar:
|
|
98
98
|
t = rearrange(t, '-> 1')
|
|
99
99
|
|
|
100
|
-
|
|
101
|
-
|
|
102
|
-
|
|
100
|
+
# bounds may be unbatched (single_action_space) or already batched
|
|
101
|
+
# (vector envs exposing only a batched action_space)
|
|
102
|
+
|
|
103
|
+
if low.shape != t.shape and not (low.ndim > 0 and t.shape[-low.ndim:] == low.shape):
|
|
104
|
+
dim = t.shape[-1]
|
|
105
|
+
low = np.broadcast_to(low, (dim,))
|
|
106
|
+
high = np.broadcast_to(high, (dim,))
|
|
107
|
+
|
|
103
108
|
valid = np.isfinite(low) & np.isfinite(high)
|
|
104
109
|
|
|
105
110
|
if is_tensor(t):
|
|
@@ -1,32 +1,20 @@
|
|
|
1
1
|
from __future__ import annotations
|
|
2
2
|
|
|
3
|
-
import numpy as np
|
|
4
|
-
import torch
|
|
5
|
-
from torch import is_tensor
|
|
6
|
-
|
|
7
3
|
from .helpers import (
|
|
8
4
|
EnvWrapper,
|
|
9
5
|
default,
|
|
10
6
|
exists,
|
|
11
7
|
first_existing,
|
|
12
8
|
get_attr,
|
|
9
|
+
is_time_step,
|
|
10
|
+
normalize_reset_out,
|
|
11
|
+
normalize_step_out,
|
|
13
12
|
safe_close,
|
|
14
13
|
truthy_attr,
|
|
14
|
+
zero_like,
|
|
15
15
|
)
|
|
16
16
|
from .spaces import space_from_action_spec
|
|
17
17
|
|
|
18
|
-
# helpers
|
|
19
|
-
|
|
20
|
-
def zero_like(x):
|
|
21
|
-
if is_tensor(x):
|
|
22
|
-
return torch.zeros_like(x, dtype = torch.bool)
|
|
23
|
-
|
|
24
|
-
arr = np.asarray(x)
|
|
25
|
-
return np.zeros_like(arr, dtype = bool) if arr.ndim > 0 else False
|
|
26
|
-
|
|
27
|
-
def is_time_step(out):
|
|
28
|
-
return exists(get_attr(out, 'step_type')) and exists(get_attr(out, 'observation'))
|
|
29
|
-
|
|
30
18
|
# base adapter
|
|
31
19
|
|
|
32
20
|
class BaseEnvAdapter:
|
|
@@ -38,27 +26,10 @@ class BaseEnvAdapter:
|
|
|
38
26
|
self.env = env
|
|
39
27
|
|
|
40
28
|
def step(self, action) -> tuple:
|
|
41
|
-
|
|
42
|
-
if is_time_step(out):
|
|
43
|
-
last = out.last() if callable(get_attr(out, 'last')) else out.step_type == 2
|
|
44
|
-
return out.observation, out.reward, last, False, dict(discount = out.discount)
|
|
45
|
-
if len(out) == 5:
|
|
46
|
-
obs, reward, term, trunc, info = out
|
|
47
|
-
return obs, reward, term, trunc, info if isinstance(info, dict) else {}
|
|
48
|
-
if len(out) in (3, 4):
|
|
49
|
-
obs, reward, done, *rest = out
|
|
50
|
-
info = rest[0] if rest and isinstance(rest[0], dict) else {}
|
|
51
|
-
return obs, reward, done, zero_like(done), info
|
|
52
|
-
raise ValueError(f'cannot standardize step output of length {len(out)}')
|
|
29
|
+
return normalize_step_out(self.env.step(action))
|
|
53
30
|
|
|
54
31
|
def reset(self, **kwargs) -> tuple:
|
|
55
|
-
|
|
56
|
-
if is_time_step(out):
|
|
57
|
-
return out.observation, {}
|
|
58
|
-
if isinstance(out, tuple) and len(out) == 2:
|
|
59
|
-
obs, info = out
|
|
60
|
-
return obs, {} if info is None else (info if isinstance(info, dict) else {})
|
|
61
|
-
return out, {}
|
|
32
|
+
return normalize_reset_out(self.env.reset(**kwargs))
|
|
62
33
|
|
|
63
34
|
def seed(self, seed: int):
|
|
64
35
|
if callable(get_attr(self.env, 'seed')):
|
|
@@ -381,15 +352,16 @@ class GymnasiumAdapter(BaseEnvAdapter):
|
|
|
381
352
|
|
|
382
353
|
@property
|
|
383
354
|
def autoresets(self) -> bool:
|
|
384
|
-
|
|
385
|
-
|
|
355
|
+
mode = first_existing(self.env, 'autoreset', 'autoresets', 'autoreset_mode')
|
|
356
|
+
|
|
357
|
+
if exists(mode):
|
|
358
|
+
return truthy_attr(mode)
|
|
359
|
+
|
|
386
360
|
try:
|
|
387
361
|
from gymnasium.vector import VectorEnv
|
|
388
|
-
|
|
389
|
-
return True
|
|
362
|
+
return isinstance(self.env, VectorEnv)
|
|
390
363
|
except ImportError:
|
|
391
|
-
|
|
392
|
-
return False
|
|
364
|
+
return False
|
|
393
365
|
|
|
394
366
|
def seed(self, seed: int):
|
|
395
367
|
try:
|
{env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/auto_batched_wrapper.py
RENAMED
|
@@ -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
|
|
@@ -202,7 +213,9 @@ class AutoBatchedWrapper(EnvWrapper):
|
|
|
202
213
|
|
|
203
214
|
obs, reward, terminated, truncated, info = *maybe_expand_dim(out[:4]), out[4]
|
|
204
215
|
|
|
205
|
-
if isinstance(info, dict)
|
|
206
|
-
|
|
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])
|
|
207
220
|
|
|
208
221
|
return obs, reward, terminated, truncated, info
|
{env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/done_tracker_wrapper.py
RENAMED
|
@@ -2,8 +2,6 @@ from __future__ import annotations
|
|
|
2
2
|
|
|
3
3
|
import numpy as np
|
|
4
4
|
|
|
5
|
-
from torch.utils._pytree import tree_flatten
|
|
6
|
-
|
|
7
5
|
from .auto_batched_wrapper import AutoBatchedWrapper
|
|
8
6
|
from .helpers import (
|
|
9
7
|
EnvWrapper,
|
|
@@ -11,27 +9,11 @@ from .helpers import (
|
|
|
11
9
|
env_autoresets,
|
|
12
10
|
env_num_envs,
|
|
13
11
|
exists,
|
|
14
|
-
|
|
15
|
-
is_array,
|
|
12
|
+
get_batch_size,
|
|
16
13
|
is_vectorized,
|
|
17
14
|
to_numpy,
|
|
18
15
|
)
|
|
19
16
|
|
|
20
|
-
# helper functions
|
|
21
|
-
|
|
22
|
-
def get_batch_size(tree) -> int | None:
|
|
23
|
-
leaves, _ = tree_flatten(tree)
|
|
24
|
-
|
|
25
|
-
if not leaves:
|
|
26
|
-
return None
|
|
27
|
-
|
|
28
|
-
first = leaves[0]
|
|
29
|
-
|
|
30
|
-
if is_array(first):
|
|
31
|
-
return len(first) if first.ndim > 0 else None
|
|
32
|
-
|
|
33
|
-
return len(first) if exists(get_attr(first, '__len__')) else None
|
|
34
|
-
|
|
35
17
|
# classes
|
|
36
18
|
|
|
37
19
|
class DoneTrackerWrapper(EnvWrapper):
|
|
@@ -113,6 +95,6 @@ class DoneTrackerWrapper(EnvWrapper):
|
|
|
113
95
|
info['episode_lengths'] = self.episode_lengths.copy()
|
|
114
96
|
|
|
115
97
|
if self.all_done:
|
|
116
|
-
info.update(needs_reset =
|
|
98
|
+
info.update(needs_reset = self.needs_reset, all_done = True)
|
|
117
99
|
|
|
118
100
|
return obs, reward, terminated, truncated, info
|
|
@@ -8,19 +8,31 @@ 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
|
|
11
|
+
from .helpers import (
|
|
12
|
+
EnvWrapper,
|
|
13
|
+
copy_leaf,
|
|
14
|
+
dones_of,
|
|
15
|
+
env_autoresets,
|
|
16
|
+
exists,
|
|
17
|
+
first_existing,
|
|
18
|
+
is_vectorized,
|
|
19
|
+
to_numpy
|
|
20
|
+
)
|
|
12
21
|
|
|
13
22
|
# helpers
|
|
14
23
|
|
|
15
|
-
def
|
|
16
|
-
if is_tensor(x):
|
|
17
|
-
|
|
18
|
-
|
|
24
|
+
def broadcast_mask(mask, x):
|
|
25
|
+
if not is_tensor(x):
|
|
26
|
+
return mask
|
|
27
|
+
|
|
28
|
+
m = torch.as_tensor(mask, device = x.device, dtype = torch.bool)
|
|
29
|
+
diff = x.ndim - m.ndim
|
|
19
30
|
|
|
20
|
-
|
|
21
|
-
m = rearrange(m, f'... -> ... {" ".join(["1"] * diff)}')
|
|
31
|
+
return rearrange(m, f'... -> ... {" ".join(["1"] * diff)}') if diff > 0 else m
|
|
22
32
|
|
|
23
|
-
|
|
33
|
+
def zero_mask(x, mask, fill_scalar = None):
|
|
34
|
+
if is_tensor(x):
|
|
35
|
+
return torch.where(broadcast_mask(mask, x), torch.zeros_like(x), x)
|
|
24
36
|
|
|
25
37
|
arr = np.asarray(x)
|
|
26
38
|
|
|
@@ -40,13 +52,7 @@ def back_to_mask_type(dones, newly):
|
|
|
40
52
|
|
|
41
53
|
def merge_final(current, value, mask):
|
|
42
54
|
if is_tensor(current):
|
|
43
|
-
|
|
44
|
-
diff = current.ndim - m.ndim
|
|
45
|
-
|
|
46
|
-
if diff > 0:
|
|
47
|
-
m = rearrange(m, f'... -> ... {" ".join(["1"] * diff)}')
|
|
48
|
-
|
|
49
|
-
return torch.where(m, value, current)
|
|
55
|
+
return torch.where(broadcast_mask(mask, current), value, current)
|
|
50
56
|
|
|
51
57
|
curr = np.asarray(current)
|
|
52
58
|
|
|
@@ -60,10 +66,11 @@ def merge_final(current, value, mask):
|
|
|
60
66
|
# class
|
|
61
67
|
|
|
62
68
|
class EpisodePaddingWrapper(EnvWrapper):
|
|
63
|
-
def __init__(self, env):
|
|
69
|
+
def __init__(self, env, pad_autoreset: bool = True):
|
|
64
70
|
super().__init__(env)
|
|
65
71
|
self.is_vector = is_vectorized(env)
|
|
66
72
|
self.autoreset = env_autoresets(env)
|
|
73
|
+
self.pad_autoreset = pad_autoreset
|
|
67
74
|
self._last_obs = None
|
|
68
75
|
self._is_done = None
|
|
69
76
|
self._final_obs = None
|
|
@@ -87,25 +94,30 @@ class EpisodePaddingWrapper(EnvWrapper):
|
|
|
87
94
|
if self._is_done is None or len(self._is_done) != len(mask):
|
|
88
95
|
self._is_done = np.zeros(len(mask), dtype = bool)
|
|
89
96
|
|
|
90
|
-
|
|
91
|
-
|
|
97
|
+
# autoreset envs revive every done slot, so each done is a new
|
|
98
|
+
# terminal transition; non-autoreset slots stay done until reset
|
|
99
|
+
|
|
100
|
+
newly = mask if self.autoreset else mask & ~self._is_done
|
|
101
|
+
self._is_done |= mask
|
|
92
102
|
|
|
93
103
|
if mask.any():
|
|
94
104
|
assert exists(self._last_obs), 'environment needs reset before calling step. call env.reset() first'
|
|
95
105
|
|
|
96
|
-
newly = mask & ~self._is_done
|
|
97
|
-
self._is_done |= mask
|
|
98
|
-
|
|
99
106
|
if newly.any():
|
|
100
|
-
|
|
107
|
+
final_val = first_existing(info, 'final_observation', 'final_obs')
|
|
108
|
+
value = final_val if exists(final_val) else self._last_obs
|
|
101
109
|
|
|
102
110
|
if self._final_obs is None:
|
|
103
111
|
self._final_obs = tree_map(copy_leaf, value)
|
|
104
112
|
else:
|
|
105
113
|
self._final_obs = tree_map(partial(merge_final, mask = newly), self._final_obs, value)
|
|
106
114
|
|
|
107
|
-
|
|
108
|
-
|
|
115
|
+
# zero pad done slots — non-autoresetting envs stay frozen on the terminal transition,
|
|
116
|
+
# while autoresetting envs already re-emit the true terminal obs (unless told to pad anyway)
|
|
117
|
+
|
|
118
|
+
if not self.autoreset or self.pad_autoreset:
|
|
119
|
+
obs = tree_map(partial(zero_mask, mask = mask), obs)
|
|
120
|
+
reward = zero_mask(reward, mask & ~newly, fill_scalar = 0.0)
|
|
109
121
|
|
|
110
122
|
info['final_observation'] = self._final_obs
|
|
111
123
|
info['_final_observation'] = back_to_mask_type(dones, mask)
|
|
@@ -0,0 +1,363 @@
|
|
|
1
|
+
from __future__ import annotations
|
|
2
|
+
|
|
3
|
+
import inspect
|
|
4
|
+
import numpy as np
|
|
5
|
+
import torch
|
|
6
|
+
from torch import is_tensor
|
|
7
|
+
from torch.utils._pytree import tree_flatten, tree_map, tree_structure, tree_unflatten
|
|
8
|
+
|
|
9
|
+
# helpers
|
|
10
|
+
|
|
11
|
+
def exists(v):
|
|
12
|
+
return v is not None
|
|
13
|
+
|
|
14
|
+
def default(v, d):
|
|
15
|
+
return v if exists(v) else d
|
|
16
|
+
|
|
17
|
+
def get_attr(obj, name, default = None):
|
|
18
|
+
# properties that raise count as missing
|
|
19
|
+
try:
|
|
20
|
+
return getattr(obj, name, default)
|
|
21
|
+
except Exception:
|
|
22
|
+
return default
|
|
23
|
+
|
|
24
|
+
def truthy_attr(value):
|
|
25
|
+
# flags arrive as None, methods, numpy scalars — only honest truths count.
|
|
26
|
+
# gymnasium's AutoresetMode.DISABLED is a truthy enum that means "off"
|
|
27
|
+
|
|
28
|
+
if not exists(value) or callable(value):
|
|
29
|
+
return False
|
|
30
|
+
|
|
31
|
+
if getattr(value, 'name', None) == 'DISABLED':
|
|
32
|
+
return False
|
|
33
|
+
|
|
34
|
+
try:
|
|
35
|
+
return bool(value)
|
|
36
|
+
except Exception:
|
|
37
|
+
return False
|
|
38
|
+
|
|
39
|
+
def first_existing(obj, *names):
|
|
40
|
+
for name in names:
|
|
41
|
+
if isinstance(obj, dict):
|
|
42
|
+
if name in obj and exists(obj[name]):
|
|
43
|
+
return obj[name]
|
|
44
|
+
else:
|
|
45
|
+
value = get_attr(obj, name)
|
|
46
|
+
if exists(value):
|
|
47
|
+
return value
|
|
48
|
+
|
|
49
|
+
return None
|
|
50
|
+
|
|
51
|
+
def is_scalar(v):
|
|
52
|
+
return isinstance(v, (int, float, bool, np.number, np.bool_))
|
|
53
|
+
|
|
54
|
+
def get_batch_size(tree) -> int | None:
|
|
55
|
+
leaves, _ = tree_flatten(tree)
|
|
56
|
+
|
|
57
|
+
if not leaves:
|
|
58
|
+
return None
|
|
59
|
+
|
|
60
|
+
first = leaves[0]
|
|
61
|
+
|
|
62
|
+
if is_array(first):
|
|
63
|
+
return len(first) if first.ndim > 0 else None
|
|
64
|
+
|
|
65
|
+
return len(first) if exists(get_attr(first, '__len__')) else None
|
|
66
|
+
|
|
67
|
+
def is_array(v):
|
|
68
|
+
return is_tensor(v) or isinstance(v, np.ndarray)
|
|
69
|
+
|
|
70
|
+
def to_numpy(t):
|
|
71
|
+
return t.detach().cpu().numpy() if is_tensor(t) else np.asarray(t)
|
|
72
|
+
|
|
73
|
+
def any_true(x):
|
|
74
|
+
if is_tensor(x):
|
|
75
|
+
return bool(x.any())
|
|
76
|
+
return bool(np.asarray(x).any())
|
|
77
|
+
|
|
78
|
+
def copy_leaf(x):
|
|
79
|
+
if is_tensor(x):
|
|
80
|
+
return x.clone()
|
|
81
|
+
|
|
82
|
+
if isinstance(x, np.ndarray):
|
|
83
|
+
return x.copy()
|
|
84
|
+
|
|
85
|
+
return x
|
|
86
|
+
|
|
87
|
+
def dones_of(terminated, truncated):
|
|
88
|
+
if not isinstance(terminated, (dict, list, tuple)):
|
|
89
|
+
return terminated | truncated
|
|
90
|
+
return tree_map(lambda a, b: a | b, terminated, truncated)
|
|
91
|
+
|
|
92
|
+
# sim step / reset normalization
|
|
93
|
+
|
|
94
|
+
def is_time_step(out):
|
|
95
|
+
return exists(get_attr(out, 'step_type')) and exists(get_attr(out, 'observation'))
|
|
96
|
+
|
|
97
|
+
def zero_like(x):
|
|
98
|
+
if is_tensor(x):
|
|
99
|
+
return torch.zeros_like(x, dtype = torch.bool)
|
|
100
|
+
|
|
101
|
+
arr = np.asarray(x)
|
|
102
|
+
return np.zeros_like(arr, dtype = bool) if arr.ndim > 0 else False
|
|
103
|
+
|
|
104
|
+
def normalize_reset_out(out):
|
|
105
|
+
if is_time_step(out):
|
|
106
|
+
return out.observation, {}
|
|
107
|
+
|
|
108
|
+
if isinstance(out, tuple) and len(out) == 2:
|
|
109
|
+
obs, info = out
|
|
110
|
+
return obs, {} if info is None else (info if isinstance(info, dict) else {})
|
|
111
|
+
|
|
112
|
+
return out, {}
|
|
113
|
+
|
|
114
|
+
def normalize_step_out(out):
|
|
115
|
+
if is_time_step(out):
|
|
116
|
+
last = out.last() if callable(get_attr(out, 'last')) else out.step_type == 2
|
|
117
|
+
return out.observation, out.reward, last, False, dict(discount = out.discount)
|
|
118
|
+
|
|
119
|
+
if len(out) == 5:
|
|
120
|
+
return out
|
|
121
|
+
|
|
122
|
+
if len(out) in (3, 4):
|
|
123
|
+
obs, reward, done, *rest = out
|
|
124
|
+
info = rest[0] if rest and isinstance(rest[0], dict) else {}
|
|
125
|
+
return obs, reward, done, zero_like(done), info
|
|
126
|
+
|
|
127
|
+
raise ValueError(f'could not standardize step output of length {len(out)}')
|
|
128
|
+
|
|
129
|
+
def _zero_leaf(x):
|
|
130
|
+
if is_tensor(x):
|
|
131
|
+
return torch.zeros_like(x)
|
|
132
|
+
if isinstance(x, np.ndarray):
|
|
133
|
+
return np.zeros_like(x)
|
|
134
|
+
if isinstance(x, bool):
|
|
135
|
+
return False
|
|
136
|
+
if isinstance(x, (int, float, np.number)):
|
|
137
|
+
return np.zeros_like(x)
|
|
138
|
+
return 0
|
|
139
|
+
|
|
140
|
+
def _stack_leaves(leaves):
|
|
141
|
+
if all(map(is_tensor, leaves)):
|
|
142
|
+
return torch.stack(leaves)
|
|
143
|
+
return np.stack(leaves)
|
|
144
|
+
|
|
145
|
+
def stack_trees(trees):
|
|
146
|
+
first = trees[0]
|
|
147
|
+
|
|
148
|
+
if is_tensor(first):
|
|
149
|
+
return torch.stack(trees)
|
|
150
|
+
|
|
151
|
+
if isinstance(first, np.ndarray):
|
|
152
|
+
return np.stack(trees)
|
|
153
|
+
|
|
154
|
+
if isinstance(first, dict):
|
|
155
|
+
return {key: stack_trees([t[key] for t in trees]) for key in first}
|
|
156
|
+
|
|
157
|
+
if isinstance(first, tuple):
|
|
158
|
+
return tuple(stack_trees([t[i] for t in trees]) for i in range(len(first)))
|
|
159
|
+
|
|
160
|
+
leaves = [tree_flatten(tree)[0] for tree in trees]
|
|
161
|
+
stacked = [_stack_leaves(col) for col in zip(*leaves)]
|
|
162
|
+
return tree_unflatten(stacked, tree_structure(trees[0]))
|
|
163
|
+
|
|
164
|
+
def unpack_vector_observations(arr):
|
|
165
|
+
# unpacks a 1D sequence / numpy object array of unbatched single-env observations
|
|
166
|
+
# (or None for un-terminated slots) into the canonical batched pytree format matching obs
|
|
167
|
+
|
|
168
|
+
if not isinstance(arr, (np.ndarray, list, tuple)):
|
|
169
|
+
return arr
|
|
170
|
+
|
|
171
|
+
if isinstance(arr, np.ndarray) and arr.dtype != object:
|
|
172
|
+
return arr
|
|
173
|
+
|
|
174
|
+
sample = next((x for x in arr if x is not None), None)
|
|
175
|
+
if sample is None:
|
|
176
|
+
return arr
|
|
177
|
+
|
|
178
|
+
trees = [tree_map(_zero_leaf, sample) if x is None else x for x in arr]
|
|
179
|
+
return stack_trees(trees)
|
|
180
|
+
|
|
181
|
+
# environment probes
|
|
182
|
+
|
|
183
|
+
def get_adapter(env):
|
|
184
|
+
from .adapters import get_adapter as _get_adapter
|
|
185
|
+
return _get_adapter(env)
|
|
186
|
+
|
|
187
|
+
def env_num_envs(env) -> int:
|
|
188
|
+
return get_adapter(env).num_envs
|
|
189
|
+
|
|
190
|
+
def env_autoresets(env) -> bool:
|
|
191
|
+
return get_adapter(env).autoresets
|
|
192
|
+
|
|
193
|
+
def env_render_mode(env):
|
|
194
|
+
return get_attr(env, 'render_mode', 'custom')
|
|
195
|
+
|
|
196
|
+
def env_render(env, height, width, camera = None):
|
|
197
|
+
return get_adapter(env).render(height, width, camera)
|
|
198
|
+
|
|
199
|
+
def is_vectorized(env) -> bool:
|
|
200
|
+
return get_adapter(env).is_vectorized
|
|
201
|
+
|
|
202
|
+
# gymnasium 1.x surfaces final observations as 'final_obs' / '_final_obs', everything here uses 'final_observation' / '_final_observation'
|
|
203
|
+
|
|
204
|
+
FINAL_OBSERVATION_KEYS = ('final_observation', 'final_obs')
|
|
205
|
+
FINAL_OBSERVATION_MASK_KEYS = ('_final_observation', '_final_obs')
|
|
206
|
+
|
|
207
|
+
# (gymnasium 1.x name, standard name)
|
|
208
|
+
|
|
209
|
+
FINAL_OBS_ALIASES = (('final_obs', 'final_observation'), ('_final_obs', '_final_observation'))
|
|
210
|
+
|
|
211
|
+
def has_final_observation(info):
|
|
212
|
+
return isinstance(info, dict) and any(key in info for key in FINAL_OBSERVATION_KEYS)
|
|
213
|
+
|
|
214
|
+
def maybe_get_final_observation(info):
|
|
215
|
+
if isinstance(info, dict):
|
|
216
|
+
for key in FINAL_OBSERVATION_KEYS:
|
|
217
|
+
if key in info:
|
|
218
|
+
return info[key]
|
|
219
|
+
|
|
220
|
+
return None
|
|
221
|
+
|
|
222
|
+
def get_final_observation(info, default = None):
|
|
223
|
+
final_obs = maybe_get_final_observation(info)
|
|
224
|
+
|
|
225
|
+
if exists(final_obs):
|
|
226
|
+
return final_obs
|
|
227
|
+
|
|
228
|
+
if exists(default):
|
|
229
|
+
return default
|
|
230
|
+
|
|
231
|
+
raise KeyError("no 'final_observation' found in info")
|
|
232
|
+
|
|
233
|
+
def maybe_transform_final_observation(info, fn):
|
|
234
|
+
if not has_final_observation(info):
|
|
235
|
+
return info
|
|
236
|
+
|
|
237
|
+
for key in FINAL_OBSERVATION_KEYS:
|
|
238
|
+
if key in info:
|
|
239
|
+
info[key] = fn(info[key])
|
|
240
|
+
|
|
241
|
+
return info
|
|
242
|
+
|
|
243
|
+
def mark_terminal_obs(info, obs, dones, is_vector):
|
|
244
|
+
if not isinstance(info, dict) or not any_true(dones):
|
|
245
|
+
return
|
|
246
|
+
|
|
247
|
+
# vector envs: gymnasium provides final_obs as a 1D object array of single-env obs (or None).
|
|
248
|
+
# unpack into the standardized batched pytree format matching obs.
|
|
249
|
+
|
|
250
|
+
if is_vector:
|
|
251
|
+
for key in FINAL_OBSERVATION_KEYS:
|
|
252
|
+
if key in info:
|
|
253
|
+
info[key] = unpack_vector_observations(info[key])
|
|
254
|
+
|
|
255
|
+
# gymnasium 1.x names these 'final_obs' / '_final_obs' — alias to the standard names when provided
|
|
256
|
+
|
|
257
|
+
for src, dst in FINAL_OBS_ALIASES:
|
|
258
|
+
if src in info and dst not in info:
|
|
259
|
+
info[dst] = info[src]
|
|
260
|
+
|
|
261
|
+
# single envs get the final observation synthesized on termination if the sim provided none
|
|
262
|
+
|
|
263
|
+
if not is_vector and 'final_observation' not in info:
|
|
264
|
+
info['final_observation'] = obs
|
|
265
|
+
info['_final_observation'] = True
|
|
266
|
+
|
|
267
|
+
def instantiate_env(env):
|
|
268
|
+
if isinstance(env, str):
|
|
269
|
+
import gymnasium as gym
|
|
270
|
+
return gym.make(env)
|
|
271
|
+
|
|
272
|
+
if isinstance(env, type) or (callable(env) and not exists(get_attr(env, 'reset'))):
|
|
273
|
+
return env()
|
|
274
|
+
|
|
275
|
+
return env
|
|
276
|
+
|
|
277
|
+
def safe_close(env):
|
|
278
|
+
if not exists(env):
|
|
279
|
+
return
|
|
280
|
+
|
|
281
|
+
close_fn = get_attr(env, 'close')
|
|
282
|
+
|
|
283
|
+
if callable(close_fn):
|
|
284
|
+
try:
|
|
285
|
+
close_fn()
|
|
286
|
+
except Exception:
|
|
287
|
+
pass
|
|
288
|
+
|
|
289
|
+
# base wrapper
|
|
290
|
+
|
|
291
|
+
class EnvWrapper:
|
|
292
|
+
def __init__(self, env):
|
|
293
|
+
self.env = env
|
|
294
|
+
|
|
295
|
+
def close(self):
|
|
296
|
+
safe_close(self.env)
|
|
297
|
+
|
|
298
|
+
def __enter__(self):
|
|
299
|
+
return self
|
|
300
|
+
|
|
301
|
+
def __exit__(self, exc_type, exc_val, exc_tb):
|
|
302
|
+
self.close()
|
|
303
|
+
|
|
304
|
+
def __getattr__(self, name):
|
|
305
|
+
if name.startswith('_'):
|
|
306
|
+
raise AttributeError(f"attempted to get missing private attribute '{name}'")
|
|
307
|
+
return getattr(self.env, name)
|
|
308
|
+
|
|
309
|
+
def accepts_done_param(fn):
|
|
310
|
+
try:
|
|
311
|
+
params = inspect.signature(fn).parameters
|
|
312
|
+
return 'done' in params or any(p.kind == inspect.Parameter.VAR_KEYWORD for p in params.values())
|
|
313
|
+
except (ValueError, TypeError):
|
|
314
|
+
return False
|
|
315
|
+
|
|
316
|
+
class TransformObservationWrapper(EnvWrapper):
|
|
317
|
+
"""
|
|
318
|
+
Base observation wrapper that automatically handles:
|
|
319
|
+
- Calling transform_obs on observations in reset() and step()
|
|
320
|
+
- Detecting environment autoreset and passing `done` to stateful transforms
|
|
321
|
+
- Propagating transformed observations to info['final_observation']
|
|
322
|
+
"""
|
|
323
|
+
|
|
324
|
+
def __init__(self, env):
|
|
325
|
+
super().__init__(env)
|
|
326
|
+
self.autoresets = env_autoresets(env)
|
|
327
|
+
self.takes_done = accepts_done_param(self.transform_obs)
|
|
328
|
+
|
|
329
|
+
def transform_obs(self, obs, done = None):
|
|
330
|
+
return obs
|
|
331
|
+
|
|
332
|
+
def transform(self, obs, done = None):
|
|
333
|
+
return self.transform_obs(obs, done = done) if self.takes_done else self.transform_obs(obs)
|
|
334
|
+
|
|
335
|
+
def reset(self, **kwargs):
|
|
336
|
+
obs, info = self.env.reset(**kwargs)
|
|
337
|
+
obs = self.transform_obs(obs)
|
|
338
|
+
|
|
339
|
+
if isinstance(info, dict):
|
|
340
|
+
for key in FINAL_OBSERVATION_KEYS:
|
|
341
|
+
if key in info:
|
|
342
|
+
info[key] = self.transform_obs(info[key])
|
|
343
|
+
|
|
344
|
+
return obs, info
|
|
345
|
+
|
|
346
|
+
def step(self, action):
|
|
347
|
+
obs, reward, terminated, truncated, info = self.env.step(action)
|
|
348
|
+
done = dones_of(terminated, truncated) if self.autoresets else None
|
|
349
|
+
|
|
350
|
+
out = self.transform_obs(obs, done = done) if self.takes_done else self.transform_obs(obs)
|
|
351
|
+
|
|
352
|
+
if isinstance(info, dict):
|
|
353
|
+
for key in FINAL_OBSERVATION_KEYS:
|
|
354
|
+
if key in info:
|
|
355
|
+
if not self.takes_done:
|
|
356
|
+
info[key] = self.transform_obs(info[key])
|
|
357
|
+
elif not self.autoresets:
|
|
358
|
+
info[key] = out
|
|
359
|
+
|
|
360
|
+
return out, reward, terminated, truncated, info
|
|
361
|
+
|
|
362
|
+
ObservationWrapper = TransformObservationWrapper
|
|
363
|
+
|