env-ssl-wrapper 0.4.4__tar.gz → 0.4.6__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.4 → env_ssl_wrapper-0.4.6}/PKG-INFO +1 -1
  2. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/action_chunk.py +17 -2
  3. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/auto_batched_wrapper.py +8 -12
  4. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/helpers.py +3 -0
  5. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/pyproject.toml +1 -1
  6. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/.gitignore +0 -0
  7. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/LICENSE +0 -0
  8. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/README.md +0 -0
  9. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/__init__.py +0 -0
  10. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/memory_trace.py +0 -0
  11. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/__init__.py +0 -0
  12. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/action_transform_wrapper.py +0 -0
  13. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/adapters.py +0 -0
  14. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/done_tracker_wrapper.py +0 -0
  15. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/episode_padding_wrapper.py +0 -0
  16. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/flatten_obs_wrapper.py +0 -0
  17. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/image_wrapper.py +0 -0
  18. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/mocks.py +0 -0
  19. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/spaces.py +0 -0
  20. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/standardize_env_wrapper.py +0 -0
  21. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/standardize_wrapper.py +0 -0
  22. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/tensor_wrapper.py +0 -0
  23. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/time_limit_wrapper.py +0 -0
  24. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/utils.py +0 -0
  25. {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/vector.py +0 -0
@@ -1,6 +1,6 @@
1
1
  Metadata-Version: 2.5
2
2
  Name: env-ssl-wrapper
3
- Version: 0.4.4
3
+ Version: 0.4.6
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
@@ -7,11 +7,12 @@ from torch import is_tensor
7
7
 
8
8
  from .standardize.helpers import (
9
9
  EnvWrapper,
10
- exists,
10
+ copy_leaf,
11
11
  default,
12
12
  any_true,
13
13
  dones_of,
14
14
  get_attr,
15
+ is_vectorized,
15
16
  )
16
17
 
17
18
  # helpers
@@ -22,6 +23,16 @@ def stack_steps(steps):
22
23
  def chunk_steps(actions, axis):
23
24
  return actions.unbind(dim = axis) if is_tensor(actions) else np.moveaxis(actions, axis, 0)
24
25
 
26
+ def unbatch_step(action):
27
+ # chunks always carry a num_envs axis - drop it for plain single envs
28
+ if is_tensor(action):
29
+ return action.squeeze(0)
30
+
31
+ if isinstance(action, np.ndarray):
32
+ return np.squeeze(action, axis = 0) if action.ndim > 0 else action
33
+
34
+ return action
35
+
25
36
  # wrapper
26
37
 
27
38
  class ActionChunkWrapper(EnvWrapper):
@@ -53,6 +64,7 @@ class ActionChunkWrapper(EnvWrapper):
53
64
 
54
65
  action_space = get_attr(env, 'action_space')
55
66
  self.action_shape = tuple(get_attr(action_space, 'shape', ()) or ())
67
+ self.expects_batch = is_vectorized(env)
56
68
 
57
69
  @property
58
70
  def chunk_action_shape(self):
@@ -82,8 +94,11 @@ class ActionChunkWrapper(EnvWrapper):
82
94
  out = None
83
95
 
84
96
  for action in chunk_steps(actions, chunk_axis):
97
+ if not self.expects_batch:
98
+ action = unbatch_step(action)
99
+
85
100
  out = self.env.step(action)
86
- rewards.append(out[1])
101
+ rewards.append(copy_leaf(out[1]))
87
102
 
88
103
  if any_true(dones_of(out[2], out[3])):
89
104
  break
@@ -153,18 +153,14 @@ def is_numeric_container(x):
153
153
  def maybe_squeeze_dim(x, shape_tree = None, is_vector = False, prepend_batch = False):
154
154
  # reshape actions to match the env's space, falling back to heuristics
155
155
 
156
- if isinstance(shape_tree, (list, dict)):
157
- keyed = isinstance(shape_tree, dict)
158
-
159
- ok = x.keys() == shape_tree.keys() if keyed else isinstance(x, (list, tuple)) and len(x) == len(shape_tree)
160
- assert ok, f'action structure does not match its {"dict" if keyed else "tuple"} action space'
161
-
162
- children = list(x.values()) if keyed else list(x)
163
- subtrees = list(shape_tree.values()) if keyed else shape_tree
164
-
165
- leaves = [maybe_squeeze_dim(child, subtree, is_vector, prepend_batch) for child, subtree in zip(children, subtrees)]
166
-
167
- return dict(zip(x.keys(), leaves)) if keyed else rebuild_container(x, leaves)
156
+ if isinstance(shape_tree, dict):
157
+ assert isinstance(x, dict) and x.keys() == shape_tree.keys(), 'action structure does not match its dict action space'
158
+ return {key: maybe_squeeze_dim(child, shape_tree[key], is_vector, prepend_batch) for key, child in x.items()}
159
+
160
+ if isinstance(shape_tree, list):
161
+ assert isinstance(x, (list, tuple)) and len(x) == len(shape_tree), 'action structure does not match its tuple action space'
162
+ leaves = [maybe_squeeze_dim(child, subtree, is_vector, prepend_batch) for child, subtree in zip(x, shape_tree)]
163
+ return rebuild_container(x, leaves)
168
164
 
169
165
  # leaf-shaped tree — claims the whole input
170
166
 
@@ -84,6 +84,9 @@ def copy_leaf(x):
84
84
 
85
85
  return x
86
86
 
87
+ def copy_tree(tree):
88
+ return tree_map(copy_leaf, tree)
89
+
87
90
  def dones_of(terminated, truncated):
88
91
  if not isinstance(terminated, (dict, list, tuple)):
89
92
  return terminated | truncated
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "env-ssl-wrapper"
3
- version = "0.4.4"
3
+ version = "0.4.6"
4
4
  description = "One torch-native interface for any MDP environment"
5
5
  authors = [
6
6
  { name = "Phil Wang", email = "lucidrains@gmail.com" }
File without changes