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.
Files changed (25) hide show
  1. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/PKG-INFO +5 -2
  2. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/README.md +1 -1
  3. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/action_transform_wrapper.py +8 -3
  4. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/adapters.py +13 -41
  5. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/auto_batched_wrapper.py +16 -3
  6. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/done_tracker_wrapper.py +2 -20
  7. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/episode_padding_wrapper.py +36 -24
  8. env_ssl_wrapper-0.4.2/env_ssl_wrapper/standardize/helpers.py +363 -0
  9. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/image_wrapper.py +39 -9
  10. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/standardize_env_wrapper.py +8 -7
  11. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/standardize_wrapper.py +1 -34
  12. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/tensor_wrapper.py +21 -4
  13. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/time_limit_wrapper.py +18 -3
  14. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/utils.py +20 -20
  15. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/vector.py +15 -34
  16. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/pyproject.toml +4 -1
  17. env_ssl_wrapper-0.4.0/env_ssl_wrapper/standardize/helpers.py +0 -221
  18. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/.gitignore +0 -0
  19. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/LICENSE +0 -0
  20. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/__init__.py +0 -0
  21. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/memory_trace.py +0 -0
  22. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/__init__.py +0 -0
  23. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/flatten_obs_wrapper.py +0 -0
  24. {env_ssl_wrapper-0.4.0 → env_ssl_wrapper-0.4.2}/env_ssl_wrapper/standardize/mocks.py +0 -0
  25. {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.0
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, (8,))
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, (8,))
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
- dim = t.shape[-1]
101
- low = np.broadcast_to(low, (dim,))
102
- high = np.broadcast_to(high, (dim,))
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
- out = self.env.step(action)
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
- out = self.env.reset(**kwargs)
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
- if truthy_attr(first_existing(self.env, 'autoreset', 'autoresets', 'autoreset_mode')):
385
- return True
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
- if isinstance(self.env, VectorEnv):
389
- return True
362
+ return isinstance(self.env, VectorEnv)
390
363
  except ImportError:
391
- pass
392
- return False
364
+ return False
393
365
 
394
366
  def seed(self, seed: int):
395
367
  try:
@@ -7,7 +7,18 @@ import torch
7
7
  from torch.utils._pytree import tree_map
8
8
  from einops import rearrange
9
9
 
10
- from .helpers import EnvWrapper, default, exists, first_existing, get_attr, is_array, is_scalar, is_tensor, is_vectorized
10
+ from .helpers import (
11
+ FINAL_OBSERVATION_KEYS,
12
+ EnvWrapper,
13
+ default,
14
+ exists,
15
+ first_existing,
16
+ get_attr,
17
+ is_array,
18
+ is_scalar,
19
+ is_tensor,
20
+ is_vectorized,
21
+ )
11
22
  from .spaces import space_from_action_spec
12
23
 
13
24
  # helper functions
@@ -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) and 'final_observation' in info:
206
- info['final_observation'] = maybe_expand_dim(info['final_observation'])
216
+ if isinstance(info, dict):
217
+ for key in FINAL_OBSERVATION_KEYS:
218
+ if key in info:
219
+ info[key] = maybe_expand_dim(info[key])
207
220
 
208
221
  return obs, reward, terminated, truncated, info
@@ -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
- get_attr,
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 = True, all_done = True)
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 EnvWrapper, copy_leaf, dones_of, env_autoresets, exists, is_vectorized, to_numpy
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 zero_mask(x, mask, fill_scalar = None):
16
- if is_tensor(x):
17
- m = torch.as_tensor(mask, device = x.device, dtype = torch.bool)
18
- diff = x.ndim - m.ndim
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
- if diff > 0:
21
- m = rearrange(m, f'... -> ... {" ".join(["1"] * diff)}')
31
+ return rearrange(m, f'... -> ... {" ".join(["1"] * diff)}') if diff > 0 else m
22
32
 
23
- return torch.where(m, torch.zeros_like(x), x)
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
- m = torch.as_tensor(mask, device = current.device, dtype = torch.bool)
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
- if self.autoreset:
91
- self._is_done &= mask
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
- value = info['final_observation'] if 'final_observation' in info else self._last_obs
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
- obs = tree_map(partial(zero_mask, mask = mask), obs)
108
- reward = zero_mask(reward, mask & ~newly, fill_scalar = 0.0)
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
+