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.
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/PKG-INFO +1 -1
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/action_chunk.py +17 -2
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/auto_batched_wrapper.py +8 -12
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/helpers.py +3 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/pyproject.toml +1 -1
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/.gitignore +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/LICENSE +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/README.md +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/__init__.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/memory_trace.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/__init__.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/action_transform_wrapper.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/adapters.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/done_tracker_wrapper.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/episode_padding_wrapper.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/flatten_obs_wrapper.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/image_wrapper.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/mocks.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/spaces.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/standardize_env_wrapper.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/standardize_wrapper.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/tensor_wrapper.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/time_limit_wrapper.py +0 -0
- {env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/utils.py +0 -0
- {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.
|
|
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
|
-
|
|
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
|
{env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/auto_batched_wrapper.py
RENAMED
|
@@ -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,
|
|
157
|
-
|
|
158
|
-
|
|
159
|
-
|
|
160
|
-
|
|
161
|
-
|
|
162
|
-
|
|
163
|
-
|
|
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
|
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/done_tracker_wrapper.py
RENAMED
|
File without changes
|
|
File without changes
|
{env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/flatten_obs_wrapper.py
RENAMED
|
File without changes
|
{env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/image_wrapper.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|
|
File without changes
|
{env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/standardize_wrapper.py
RENAMED
|
File without changes
|
{env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/tensor_wrapper.py
RENAMED
|
File without changes
|
{env_ssl_wrapper-0.4.4 → env_ssl_wrapper-0.4.6}/env_ssl_wrapper/standardize/time_limit_wrapper.py
RENAMED
|
File without changes
|
|
File without changes
|
|
File without changes
|