env-ssl-wrapper 0.4.2__tar.gz → 0.4.3__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.2 → env_ssl_wrapper-0.4.3}/PKG-INFO +49 -1
  2. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/README.md +48 -0
  3. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/__init__.py +2 -0
  4. env_ssl_wrapper-0.4.3/env_ssl_wrapper/action_chunk.py +117 -0
  5. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/standardize_env_wrapper.py +6 -0
  6. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/utils.py +13 -1
  7. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/pyproject.toml +1 -1
  8. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/.gitignore +0 -0
  9. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/LICENSE +0 -0
  10. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/memory_trace.py +0 -0
  11. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/__init__.py +0 -0
  12. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/action_transform_wrapper.py +0 -0
  13. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/adapters.py +0 -0
  14. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/auto_batched_wrapper.py +0 -0
  15. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/done_tracker_wrapper.py +0 -0
  16. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/episode_padding_wrapper.py +0 -0
  17. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/flatten_obs_wrapper.py +0 -0
  18. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/helpers.py +0 -0
  19. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/image_wrapper.py +0 -0
  20. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/mocks.py +0 -0
  21. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/spaces.py +0 -0
  22. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/standardize_wrapper.py +0 -0
  23. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/tensor_wrapper.py +0 -0
  24. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/env_ssl_wrapper/standardize/time_limit_wrapper.py +0 -0
  25. {env_ssl_wrapper-0.4.2 → env_ssl_wrapper-0.4.3}/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.2
3
+ Version: 0.4.3
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
@@ -134,6 +134,53 @@ Run the PPO benchmark on POMDP LunarLander:
134
134
  uv run test_memory_trace.py
135
135
  ```
136
136
 
137
+ ### Action Chunk
138
+
139
+ Open-loop action chunking: one `step` executes a whole chunk of actions and returns only the final state, temporally compressing the environment by the chunk length.
140
+
141
+ ```python
142
+ import gymnasium as gym
143
+ import torch
144
+ from env_ssl_wrapper import StandardizeEnvWrapper, ActionChunkWrapper
145
+
146
+ env = StandardizeEnvWrapper(gym.make('CartPole-v1'))
147
+ env = ActionChunkWrapper(env, chunk_len = 2, gamma = 0.99)
148
+
149
+ obs, info = env.reset()
150
+
151
+ while True:
152
+ actions = torch.randint(0, 2, (1, 2)) # (num_envs, chunk_len)
153
+ obs, reward, terminated, truncated, info = env.step(actions)
154
+
155
+ # reward -> discounted sum of substeps (r0 + gamma * r1 + ...)
156
+ # info['discount'] -> gamma ** chunk_length (macro-step discount for GAE / Bellman target)
157
+ # info['chunk_length'] -> env steps actually executed (drops below chunk_len on terminal chunk)
158
+ # info['chunk_rewards'] -> (1, chunk_length) raw per-step rewards
159
+
160
+ if bool((terminated | truncated).item()):
161
+ obs, info = env.reset()
162
+ ```
163
+
164
+ Chunks are shaped `(num_envs, chunk_len, *action_shape)` (`(num_envs, chunk_len)` for discrete actions).
165
+
166
+ Execution stops early the moment any env terminates or truncates mid-chunk — the terminal state and done flags of that substep are returned and the rest of the chunk is dropped, so a new episode is never silently advanced.
167
+
168
+ Pass `gamma` (default `1.`) to discount intra-chunk rewards $r = \sum_{i=0}^{L-1} \gamma^i r_i$. The macro-transition discount factor to the next state is provided as `info['discount'] = gamma ** chunk_length`. `reward_mode` can also be `'mean'` or `'last'`.
169
+
170
+ Can also be passed directly to `StandardizeEnvWrapper` or `compose_env`:
171
+
172
+ ```python
173
+ env = StandardizeEnvWrapper(gym.make('CartPole-v1'), chunk_len = 2, chunk_gamma = 0.99)
174
+ ```
175
+
176
+ Run the chunked PPO benchmark on CartPole, or check chunked rollouts against a per-step reference:
177
+
178
+ ```bash
179
+ uv run test_action_chunk.py --chunk_len=2 # PPO on CartPole
180
+ uv run test_action_chunk.py --sweep=True # PPO over several chunk lengths
181
+ uv run test_action_chunk.py --verify=True --sweep=True # chunked rollouts match per-step reference
182
+ ```
183
+
137
184
  ## Wrappers
138
185
 
139
186
  Pass wrappers as strings (default config) or `(name, dict)` tuples (custom config), in any order.
@@ -148,6 +195,7 @@ Pass wrappers as strings (default config) or `(name, dict)` tuples (custom confi
148
195
  | `action_transform` | Rescales actions from a canonical `(0, 1)` range to the env's bounds. |
149
196
  | `tensor` | NumPy → torch on a device, torch actions → numpy for the sim. |
150
197
  | `flatten_obs` | Flattens dict/tuple observations into a single vector. |
198
+ | `action_chunk` | Executes actions in open-loop chunks of length k. `('action_chunk', dict(chunk_len = 2, gamma = 0.99))` |
151
199
 
152
200
  Every env emits the same contract: obs `torch.float32`, rewards `torch.float32`, `terminated`/`truncated` `torch.bool`. `env.seed(n)` works on every sim.
153
201
 
@@ -71,6 +71,53 @@ Run the PPO benchmark on POMDP LunarLander:
71
71
  uv run test_memory_trace.py
72
72
  ```
73
73
 
74
+ ### Action Chunk
75
+
76
+ Open-loop action chunking: one `step` executes a whole chunk of actions and returns only the final state, temporally compressing the environment by the chunk length.
77
+
78
+ ```python
79
+ import gymnasium as gym
80
+ import torch
81
+ from env_ssl_wrapper import StandardizeEnvWrapper, ActionChunkWrapper
82
+
83
+ env = StandardizeEnvWrapper(gym.make('CartPole-v1'))
84
+ env = ActionChunkWrapper(env, chunk_len = 2, gamma = 0.99)
85
+
86
+ obs, info = env.reset()
87
+
88
+ while True:
89
+ actions = torch.randint(0, 2, (1, 2)) # (num_envs, chunk_len)
90
+ obs, reward, terminated, truncated, info = env.step(actions)
91
+
92
+ # reward -> discounted sum of substeps (r0 + gamma * r1 + ...)
93
+ # info['discount'] -> gamma ** chunk_length (macro-step discount for GAE / Bellman target)
94
+ # info['chunk_length'] -> env steps actually executed (drops below chunk_len on terminal chunk)
95
+ # info['chunk_rewards'] -> (1, chunk_length) raw per-step rewards
96
+
97
+ if bool((terminated | truncated).item()):
98
+ obs, info = env.reset()
99
+ ```
100
+
101
+ Chunks are shaped `(num_envs, chunk_len, *action_shape)` (`(num_envs, chunk_len)` for discrete actions).
102
+
103
+ Execution stops early the moment any env terminates or truncates mid-chunk — the terminal state and done flags of that substep are returned and the rest of the chunk is dropped, so a new episode is never silently advanced.
104
+
105
+ Pass `gamma` (default `1.`) to discount intra-chunk rewards $r = \sum_{i=0}^{L-1} \gamma^i r_i$. The macro-transition discount factor to the next state is provided as `info['discount'] = gamma ** chunk_length`. `reward_mode` can also be `'mean'` or `'last'`.
106
+
107
+ Can also be passed directly to `StandardizeEnvWrapper` or `compose_env`:
108
+
109
+ ```python
110
+ env = StandardizeEnvWrapper(gym.make('CartPole-v1'), chunk_len = 2, chunk_gamma = 0.99)
111
+ ```
112
+
113
+ Run the chunked PPO benchmark on CartPole, or check chunked rollouts against a per-step reference:
114
+
115
+ ```bash
116
+ uv run test_action_chunk.py --chunk_len=2 # PPO on CartPole
117
+ uv run test_action_chunk.py --sweep=True # PPO over several chunk lengths
118
+ uv run test_action_chunk.py --verify=True --sweep=True # chunked rollouts match per-step reference
119
+ ```
120
+
74
121
  ## Wrappers
75
122
 
76
123
  Pass wrappers as strings (default config) or `(name, dict)` tuples (custom config), in any order.
@@ -85,6 +132,7 @@ Pass wrappers as strings (default config) or `(name, dict)` tuples (custom confi
85
132
  | `action_transform` | Rescales actions from a canonical `(0, 1)` range to the env's bounds. |
86
133
  | `tensor` | NumPy → torch on a device, torch actions → numpy for the sim. |
87
134
  | `flatten_obs` | Flattens dict/tuple observations into a single vector. |
135
+ | `action_chunk` | Executes actions in open-loop chunks of length k. `('action_chunk', dict(chunk_len = 2, gamma = 0.99))` |
88
136
 
89
137
  Every env emits the same contract: obs `torch.float32`, rewards `torch.float32`, `terminated`/`truncated` `torch.bool`. `env.seed(n)` works on every sim.
90
138
 
@@ -5,6 +5,7 @@ from . import standardize
5
5
  from .standardize import *
6
6
 
7
7
  from .memory_trace import MemoryTraceWrapper
8
+ from .action_chunk import ActionChunkWrapper
8
9
 
9
10
  # Wire backwards-compatibility aliases in sys.modules and module globals
10
11
  # so imports like `from env_ssl_wrapper.done_tracker_wrapper import DoneTrackerWrapper`
@@ -38,6 +39,7 @@ for _name in _STANDARDIZE_SUBMODULES:
38
39
  __all__ = [
39
40
  *standardize.__all__,
40
41
  'MemoryTraceWrapper',
42
+ 'ActionChunkWrapper',
41
43
  ]
42
44
 
43
45
  def __getattr__(name):
@@ -0,0 +1,117 @@
1
+ from __future__ import annotations
2
+
3
+ import numpy as np
4
+ import torch
5
+ from einops import reduce
6
+ from torch import is_tensor
7
+
8
+ from .standardize.helpers import (
9
+ EnvWrapper,
10
+ exists,
11
+ default,
12
+ any_true,
13
+ dones_of,
14
+ get_attr,
15
+ )
16
+
17
+ # helpers
18
+
19
+ def stack_steps(steps):
20
+ return torch.stack(steps, dim = -1) if is_tensor(steps[0]) else np.stack(steps, axis = -1)
21
+
22
+ def chunk_steps(actions, axis):
23
+ return actions.unbind(dim = axis) if is_tensor(actions) else np.moveaxis(actions, axis, 0)
24
+
25
+ # wrapper
26
+
27
+ class ActionChunkWrapper(EnvWrapper):
28
+ """
29
+ Open-loop action chunking - executes a chunk of actions in one step,
30
+ temporally compressing the environment by the chunk length.
31
+
32
+ action chunks are shaped (num_envs, chunk_len, *action_shape)
33
+ """
34
+
35
+ def __init__(
36
+ self,
37
+ env,
38
+ chunk_len: int,
39
+ gamma: float = 1.,
40
+ discount: float | None = None,
41
+ reward_mode: str = 'sum'
42
+ ):
43
+ super().__init__(env)
44
+
45
+ gamma = default(discount, gamma)
46
+ assert chunk_len >= 1, f'chunk_len must be at least 1, got {chunk_len}'
47
+ assert 0. <= gamma <= 1., f'gamma must be between 0 and 1, got {gamma}'
48
+ assert reward_mode in ('sum', 'mean', 'last'), f'unknown reward_mode {reward_mode!r}'
49
+
50
+ self.chunk_len = chunk_len
51
+ self.gamma = float(gamma)
52
+ self.reward_mode = reward_mode
53
+
54
+ action_space = get_attr(env, 'action_space')
55
+ self.action_shape = tuple(get_attr(action_space, 'shape', ()) or ())
56
+
57
+ @property
58
+ def chunk_action_shape(self):
59
+ return (self.chunk_len, *self.action_shape)
60
+
61
+ def reset(self, **kwargs):
62
+ return self.env.reset(**kwargs)
63
+
64
+ def step(self, actions):
65
+ if isinstance(actions, dict):
66
+ raise NotImplementedError('action chunking does not support dict actions yet')
67
+
68
+ if not (is_tensor(actions) or isinstance(actions, np.ndarray)):
69
+ actions = np.asarray(actions)
70
+
71
+ chunk_axis = actions.ndim - 1 - len(self.action_shape)
72
+
73
+ assert chunk_axis >= 1, (
74
+ f'action chunks must be shaped (num_envs, chunk_len, *action_shape), got {tuple(actions.shape)}'
75
+ )
76
+
77
+ assert actions.shape[chunk_axis] == self.chunk_len, (
78
+ f'expected chunk length {self.chunk_len}, got {actions.shape[chunk_axis]}'
79
+ )
80
+
81
+ rewards = []
82
+ out = None
83
+
84
+ for action in chunk_steps(actions, chunk_axis):
85
+ out = self.env.step(action)
86
+ rewards.append(out[1])
87
+
88
+ if any_true(dones_of(out[2], out[3])):
89
+ break
90
+
91
+ obs, _, terminated, truncated, info = out
92
+
93
+ rewards = stack_steps(rewards)
94
+ executed_len = rewards.shape[-1]
95
+
96
+ if self.reward_mode == 'last':
97
+ reward = rewards[..., -1]
98
+ elif self.reward_mode == 'mean':
99
+ reward = reduce(rewards, '... k -> ...', 'mean')
100
+ elif self.gamma != 1.:
101
+ if is_tensor(rewards):
102
+ dtype = rewards.dtype if rewards.is_floating_point() else torch.float32
103
+ discounts = (self.gamma ** torch.arange(executed_len, device = rewards.device)).to(dtype = dtype)
104
+ else:
105
+ dtype = rewards.dtype if np.issubdtype(rewards.dtype, np.floating) else np.float32
106
+ discounts = (self.gamma ** np.arange(executed_len)).astype(dtype)
107
+
108
+ reward = reduce(rewards * discounts, '... k -> ...', 'sum')
109
+ else:
110
+ reward = reduce(rewards, '... k -> ...', 'sum')
111
+
112
+ info = dict(info) if isinstance(info, dict) else {}
113
+ info['chunk_length'] = executed_len
114
+ info['chunk_rewards'] = rewards
115
+ info['discount'] = self.gamma ** executed_len
116
+
117
+ return obs, reward, terminated, truncated, info
@@ -23,6 +23,9 @@ class StandardizeEnvWrapper(EnvWrapper):
23
23
  image_size: int | tuple[int, int] | None = None,
24
24
  lambdas: tuple[float, ...] | list[float] | None = None,
25
25
  keys: str | tuple[str, ...] | None = None,
26
+ chunk_len: int | None = None,
27
+ chunk_gamma: float = 1.,
28
+ chunk_reward_mode: str = 'sum',
26
29
  ):
27
30
  wrappers = []
28
31
 
@@ -48,6 +51,9 @@ class StandardizeEnvWrapper(EnvWrapper):
48
51
  if exists(lambdas):
49
52
  wrappers.append(('memory_trace', dict(lambdas = lambdas, keys = keys)))
50
53
 
54
+ if exists(chunk_len):
55
+ wrappers.append(('action_chunk', dict(chunk_len = chunk_len, gamma = chunk_gamma, reward_mode = chunk_reward_mode)))
56
+
51
57
  if flatten_obs:
52
58
  wrappers.append('flatten_obs')
53
59
 
@@ -34,10 +34,14 @@ def get_wrapper(name):
34
34
  from ..memory_trace import MemoryTraceWrapper
35
35
  return MemoryTraceWrapper
36
36
 
37
+ if name in ('action_chunk', 'chunk'):
38
+ from ..action_chunk import ActionChunkWrapper
39
+ return ActionChunkWrapper
40
+
37
41
  if name in WRAPPERS:
38
42
  return WRAPPERS[name]
39
43
 
40
- raise ValueError(f'unknown wrapper {name!r} — choose from {sorted([*WRAPPERS, "memory_trace", "standardize_env"])}')
44
+ raise ValueError(f'unknown wrapper {name!r} — choose from {sorted([*WRAPPERS, "memory_trace", "action_chunk", "standardize_env"])}')
41
45
 
42
46
  def parse_wrapper(wrapper):
43
47
  if isinstance(wrapper, str):
@@ -87,6 +91,14 @@ def compose_env(env, *wrappers, pad_episodes: bool = True):
87
91
  funcs.insert(idx_ten + 1, f)
88
92
  classes.insert(idx_ten + 1, c)
89
93
 
94
+ # action chunking changes the step signature, so it always goes outermost
95
+
96
+ from ..action_chunk import ActionChunkWrapper
97
+ if ActionChunkWrapper in classes:
98
+ idx_chunk = classes.index(ActionChunkWrapper)
99
+ funcs.append(funcs.pop(idx_chunk))
100
+ classes.append(classes.pop(idx_chunk))
101
+
90
102
  assert len(set(classes)) == len(classes), 'duplicate wrappers found'
91
103
 
92
104
  for func in funcs:
@@ -1,6 +1,6 @@
1
1
  [project]
2
2
  name = "env-ssl-wrapper"
3
- version = "0.4.2"
3
+ version = "0.4.3"
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