async-gym-agents 0.1.0__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 (22) hide show
  1. async_gym_agents-0.1.0/PKG-INFO +41 -0
  2. async_gym_agents-0.1.0/README.md +23 -0
  3. async_gym_agents-0.1.0/async_gym_agents/agents/__pycache__/async_agent.cpython-310.pyc +0 -0
  4. async_gym_agents-0.1.0/async_gym_agents/agents/__pycache__/async_agent.cpython-311.pyc +0 -0
  5. async_gym_agents-0.1.0/async_gym_agents/agents/__pycache__/injector.cpython-311.pyc +0 -0
  6. async_gym_agents-0.1.0/async_gym_agents/agents/__pycache__/off_policy_injector.cpython-311.pyc +0 -0
  7. async_gym_agents-0.1.0/async_gym_agents/agents/__pycache__/on_policy_injector.cpython-311.pyc +0 -0
  8. async_gym_agents-0.1.0/async_gym_agents/agents/async_agent.py +25 -0
  9. async_gym_agents-0.1.0/async_gym_agents/agents/injector.py +125 -0
  10. async_gym_agents-0.1.0/async_gym_agents/agents/off_policy_injector.py +305 -0
  11. async_gym_agents-0.1.0/async_gym_agents/agents/on_policy_injector.py +215 -0
  12. async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/buggy_lunar_lander.cpython-311.pyc +0 -0
  13. async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/multi_env.cpython-310.pyc +0 -0
  14. async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/multi_env.cpython-311.pyc +0 -0
  15. async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/slow_cartpole.cpython-311.pyc +0 -0
  16. async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/threaded_env.cpython-310.pyc +0 -0
  17. async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/threaded_env.cpython-311.pyc +0 -0
  18. async_gym_agents-0.1.0/async_gym_agents/envs/buggy_lunar_lander.py +46 -0
  19. async_gym_agents-0.1.0/async_gym_agents/envs/multi_env.py +84 -0
  20. async_gym_agents-0.1.0/async_gym_agents/envs/slow_cartpole.py +27 -0
  21. async_gym_agents-0.1.0/async_gym_agents/envs/threaded_env.py +210 -0
  22. async_gym_agents-0.1.0/pyproject.toml +16 -0
@@ -0,0 +1,41 @@
1
+ Metadata-Version: 2.1
2
+ Name: async-gym-agents
3
+ Version: 0.1.0
4
+ Summary: Async agents for Stable Baselines 3
5
+ Author: Jonas Peche
6
+ Author-email: jonas.peche@aon.at
7
+ Requires-Python: >=3.8
8
+ Classifier: Programming Language :: Python :: 3
9
+ Classifier: Programming Language :: Python :: 3.8
10
+ Classifier: Programming Language :: Python :: 3.9
11
+ Classifier: Programming Language :: Python :: 3.10
12
+ Classifier: Programming Language :: Python :: 3.11
13
+ Classifier: Programming Language :: Python :: 3.12
14
+ Requires-Dist: stable-baselines3[extra] (>=2.3.2,<3.0.0)
15
+ Requires-Dist: torchinfo (>=1.8.0,<2.0.0)
16
+ Description-Content-Type: text/markdown
17
+
18
+ # Async Gym Agents
19
+
20
+ Wrapper environments and agent injectors to allow for drop-in async training.
21
+
22
+ ```py
23
+ import gymnasium as gym
24
+ from stable_baselines3 import TD3
25
+
26
+ from async_gym_agents.agents.async_agent import get_injected_agent
27
+ from async_gym_agents.envs.multi_env import IndexableMultiEnv
28
+
29
+ # Create env with 8 parallel envs
30
+ env = IndexableMultiEnv([lambda: gym.make("Pendulum-v1") for i in range(8)])
31
+
32
+ # Create the model, injected with async capabilities
33
+ model = get_injected_agent(TD3)("MlpPolicy", env)
34
+
35
+ # Train the model
36
+ model.learn(total_timesteps=10)
37
+
38
+ # Shut down workers
39
+ model.shutdown()
40
+ ```
41
+
@@ -0,0 +1,23 @@
1
+ # Async Gym Agents
2
+
3
+ Wrapper environments and agent injectors to allow for drop-in async training.
4
+
5
+ ```py
6
+ import gymnasium as gym
7
+ from stable_baselines3 import TD3
8
+
9
+ from async_gym_agents.agents.async_agent import get_injected_agent
10
+ from async_gym_agents.envs.multi_env import IndexableMultiEnv
11
+
12
+ # Create env with 8 parallel envs
13
+ env = IndexableMultiEnv([lambda: gym.make("Pendulum-v1") for i in range(8)])
14
+
15
+ # Create the model, injected with async capabilities
16
+ model = get_injected_agent(TD3)("MlpPolicy", env)
17
+
18
+ # Train the model
19
+ model.learn(total_timesteps=10)
20
+
21
+ # Shut down workers
22
+ model.shutdown()
23
+ ```
@@ -0,0 +1,25 @@
1
+ from stable_baselines3.common.base_class import BaseAlgorithm
2
+ from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm
3
+ from stable_baselines3.common.on_policy_algorithm import OnPolicyAlgorithm
4
+
5
+ from async_gym_agents.agents.off_policy_injector import OffPolicyAlgorithmInjector
6
+ from async_gym_agents.agents.on_policy_injector import OnPolicyAlgorithmInjector
7
+
8
+
9
+ def get_injected_agent(clazz: BaseAlgorithm):
10
+ if issubclass(clazz, OnPolicyAlgorithm):
11
+
12
+ class AsyncAgent(OnPolicyAlgorithmInjector, clazz):
13
+ pass
14
+
15
+ return AsyncAgent
16
+
17
+ elif issubclass(clazz, OffPolicyAlgorithm):
18
+
19
+ class AsyncAgent(OffPolicyAlgorithmInjector, clazz):
20
+ pass
21
+
22
+ return AsyncAgent
23
+
24
+ else:
25
+ raise ValueError(f"Unknown agent class {clazz}!")
@@ -0,0 +1,125 @@
1
+ import queue
2
+ import threading
3
+ from queue import Queue
4
+ from typing import List
5
+
6
+ from async_gym_agents.envs.multi_env import IndexableMultiEnv
7
+
8
+
9
+ class AsyncAgentInjector:
10
+ def __init__(self, max_steps_in_buffer: int = 8, skip_truncated: bool = False):
11
+ self._buffer_utilization = 0.0
12
+ self._buffer_emptiness = 0.0
13
+ self._buffer_stat_count = 0
14
+
15
+ self.running = True
16
+ self.initialized = False
17
+ self.threads = []
18
+
19
+ self.total_episodes = 0
20
+ self.skipped_episodes = 0
21
+ self.skip_truncated = skip_truncated
22
+
23
+ # The larger the queue, the less wait times, but the more outdated the policies training data is
24
+ self.queue = Queue(max_steps_in_buffer)
25
+ self.episode_lock = threading.Lock()
26
+
27
+ # The policy itself is rarely thread-safe
28
+ self.policy_lock = threading.Lock()
29
+
30
+ def _excluded_save_params(self) -> List[str]:
31
+ return [
32
+ "threads",
33
+ "queue",
34
+ "episode_lock",
35
+ "policy_lock",
36
+ ]
37
+
38
+ # noinspection PyUnresolvedReferences
39
+ def get_indexable_env(self) -> IndexableMultiEnv:
40
+ """
41
+ Asserts whether a correct environment is supplied
42
+ """
43
+ assert isinstance(
44
+ self.env, IndexableMultiEnv
45
+ ), "You must pass a IndexableMultiEnv"
46
+ return self.env
47
+
48
+ def _initialize_threads(self):
49
+ self.threads = []
50
+ for index in range(self.get_indexable_env().real_n_envs):
51
+ thread = threading.Thread(
52
+ target=self._collector_loop,
53
+ args=(index,),
54
+ )
55
+ self.threads.append(thread)
56
+ self.threads[index].start()
57
+
58
+ def fetch_transition(self):
59
+ self._buffer_utilization += self.queue.qsize()
60
+ self._buffer_emptiness += 1 if self.queue.empty() else 0
61
+ self._buffer_stat_count += 1
62
+ return self.queue.get()
63
+
64
+ @property
65
+ def buffer_utilization(self) -> float:
66
+ return (
67
+ 0
68
+ if self._buffer_stat_count == 0
69
+ else self._buffer_utilization / self._buffer_stat_count
70
+ )
71
+
72
+ @property
73
+ def buffer_emptyness(self) -> float:
74
+ return (
75
+ 0
76
+ if self._buffer_stat_count == 0
77
+ else self._buffer_emptiness / self._buffer_stat_count
78
+ )
79
+
80
+ @property
81
+ def truncated_episodes_fraction(self) -> float:
82
+ return (
83
+ 0
84
+ if self.total_episodes == 0
85
+ else self.skipped_episodes / self.total_episodes
86
+ )
87
+
88
+ def _episode_generator(self, index: int):
89
+ raise NotImplementedError()
90
+
91
+ def _collector_loop(
92
+ self,
93
+ index: int,
94
+ ):
95
+ """
96
+ Batch-inserts transitions whenever a episode is done.
97
+ """
98
+ for episode in self._episode_generator(index):
99
+ # Keeps track of truncated episodes and optionally removes them
100
+ self.total_episodes += 1
101
+ if episode[-1].infos[0]["TimeLimit.truncated"]:
102
+ self.skipped_episodes += 1
103
+ if self.skip_truncated:
104
+ continue
105
+
106
+ # Feeds the episodes into the queue
107
+ with self.episode_lock:
108
+ for transition in episode:
109
+ while self.running:
110
+ try:
111
+ self.queue.put(transition, block=True, timeout=1)
112
+ break
113
+ except queue.Full:
114
+ pass
115
+
116
+ def shutdown(self):
117
+ """
118
+ Shuts down the workers.
119
+ Shutting down is required to fully release environments.
120
+ Subsequent calls to e.g., train will restart the workers.
121
+ """
122
+ self.running = False
123
+ for thread in self.threads:
124
+ thread.join()
125
+ self.initialized = False
@@ -0,0 +1,305 @@
1
+ from copy import deepcopy
2
+ from dataclasses import dataclass
3
+ from typing import Any, Dict, List, Optional, Tuple, Union
4
+
5
+ import numpy as np
6
+ from gymnasium import spaces
7
+ from stable_baselines3.common.buffers import ReplayBuffer
8
+ from stable_baselines3.common.callbacks import BaseCallback
9
+ from stable_baselines3.common.noise import ActionNoise
10
+ from stable_baselines3.common.off_policy_algorithm import OffPolicyAlgorithm
11
+ from stable_baselines3.common.type_aliases import RolloutReturn, TrainFreq
12
+ from stable_baselines3.common.utils import should_collect_more_steps
13
+ from stable_baselines3.common.vec_env import VecEnv
14
+
15
+ from async_gym_agents.agents.injector import AsyncAgentInjector
16
+ from async_gym_agents.envs.multi_env import IndexableMultiEnv
17
+
18
+
19
+ @dataclass
20
+ class Transition:
21
+ buffer_actions: list[any]
22
+ last_obs: list[any]
23
+ new_obs: list[any]
24
+ rewards: list[float]
25
+ dones: list[bool]
26
+ infos: list[Any]
27
+
28
+
29
+ class OffPolicyAlgorithmInjector(AsyncAgentInjector, OffPolicyAlgorithm):
30
+ def __init__(self, *args, max_steps_in_buffer: int = 8, **kwargs) -> None:
31
+ super().__init__(max_steps_in_buffer)
32
+ super(AsyncAgentInjector, self).__init__(*args, **kwargs)
33
+
34
+ def train(self, *args, **kwargs) -> None:
35
+ with self.policy_lock:
36
+ super().train(*args, **kwargs)
37
+
38
+ def _excluded_save_params(self) -> List[str]:
39
+ return [
40
+ *super()._excluded_save_params(),
41
+ *super(AsyncAgentInjector, self)._excluded_save_params(),
42
+ ]
43
+
44
+ def _store_transition(*args):
45
+ raise NotImplementedError()
46
+
47
+ def _custom_store_transition(
48
+ self,
49
+ replay_buffer: ReplayBuffer,
50
+ buffer_action: np.ndarray,
51
+ last_obs: Union[np.ndarray, Dict[str, np.ndarray]],
52
+ new_obs: Union[np.ndarray, Dict[str, np.ndarray]],
53
+ reward: np.ndarray,
54
+ dones: np.ndarray,
55
+ infos: List[Dict[str, Any]],
56
+ ) -> None:
57
+ """
58
+ Nearly identical to super but stateless (last_obs now passed)
59
+ """
60
+ # Store only the unnormalized version
61
+ if self._vec_normalize_env is not None:
62
+ raise NotImplementedError()
63
+
64
+ # As the VecEnv resets automatically, new_obs is already the
65
+ # first observation of the next episode
66
+ for i, done in enumerate(dones):
67
+ if done and infos[i].get("terminal_observation") is not None:
68
+ if isinstance(new_obs, dict):
69
+ next_obs_ = infos[i]["terminal_observation"]
70
+ # VecNormalize normalizes the terminal observation
71
+ if self._vec_normalize_env is not None:
72
+ next_obs_ = self._vec_normalize_env.unnormalize_obs(next_obs_)
73
+ # Replace next obs for the correct envs
74
+ for key in new_obs.keys():
75
+ new_obs[key][i] = next_obs_[key]
76
+ else:
77
+ new_obs[i] = infos[i]["terminal_observation"]
78
+ # VecNormalize normalizes the terminal observation
79
+ if self._vec_normalize_env is not None:
80
+ new_obs[i] = self._vec_normalize_env.unnormalize_obs(
81
+ new_obs[i, :]
82
+ )
83
+
84
+ replay_buffer.add(
85
+ last_obs,
86
+ new_obs,
87
+ buffer_action,
88
+ reward,
89
+ dones,
90
+ infos,
91
+ )
92
+
93
+ def _sample_action(*args):
94
+ raise NotImplementedError()
95
+
96
+ def _custom_sample_action(
97
+ self,
98
+ learning_starts: int,
99
+ obs,
100
+ action_noise: Optional[ActionNoise] = None,
101
+ ) -> Tuple[np.ndarray, np.ndarray]:
102
+ """
103
+ Very similar but uses passed observation as input
104
+ """
105
+ # Select action randomly or according to policy
106
+ if self.num_timesteps < learning_starts and not (
107
+ self.use_sde and self.use_sde_at_warmup
108
+ ):
109
+ # Warmup phase
110
+ unscaled_action = np.array([self.action_space.sample()])
111
+ else:
112
+ # Note: when using continuous actions,
113
+ # we assume that the policy uses tanh to scale the action
114
+ # We use non-deterministic action in the case of SAC, for TD3, it does not matter
115
+ unscaled_action, _ = self.predict(obs, deterministic=False)
116
+
117
+ # Rescale the action from [low, high] to [-1, 1]
118
+ if isinstance(self.action_space, spaces.Box):
119
+ scaled_action = self.policy.scale_action(unscaled_action)
120
+
121
+ # Add noise to the action (improve exploration)
122
+ if action_noise is not None:
123
+ scaled_action = np.clip(scaled_action + action_noise(), -1, 1)
124
+
125
+ # We store the scaled action in the buffer
126
+ buffer_action = scaled_action
127
+ action = self.policy.unscale_action(scaled_action)
128
+ else:
129
+ # Discrete case, no need to normalize or clip
130
+ buffer_action = unscaled_action
131
+ action = buffer_action
132
+ return action, buffer_action
133
+
134
+ def _episode_generator(self, index: int) -> list[Transition]:
135
+ """
136
+ Continuously plays the game and returns episodes of Transitions
137
+ """
138
+ env = self.get_indexable_env()
139
+ last_obs = env.reset(index=index)
140
+
141
+ episode = []
142
+
143
+ while self.running:
144
+ with self.policy_lock:
145
+ # Select action randomly or according to policy
146
+ actions, buffer_actions = self._custom_sample_action(
147
+ self.learning_starts,
148
+ last_obs,
149
+ self.action_noise,
150
+ )
151
+
152
+ # Rescale and perform action
153
+ new_obs, rewards, dones, infos = env.step(actions, index=index)
154
+
155
+ # Store transition
156
+ episode.append(
157
+ Transition(
158
+ buffer_actions,
159
+ deepcopy(last_obs),
160
+ deepcopy(new_obs),
161
+ rewards,
162
+ dones,
163
+ infos,
164
+ )
165
+ )
166
+ last_obs = new_obs
167
+
168
+ # Start new episode
169
+ if any(dones):
170
+ yield episode
171
+ episode = []
172
+
173
+ def collect_rollouts(
174
+ self,
175
+ env: IndexableMultiEnv,
176
+ callback: BaseCallback,
177
+ train_freq: TrainFreq,
178
+ replay_buffer: ReplayBuffer,
179
+ action_noise: Optional[ActionNoise] = None,
180
+ learning_starts: int = 0,
181
+ log_interval: Optional[int] = None,
182
+ ) -> RolloutReturn:
183
+ """
184
+ Collect experiences and store them into a ``ReplayBuffer``.
185
+
186
+ :param env: The training environment
187
+ :param callback: Callback that will be called at each step
188
+ (and at the beginning and end of the rollout)
189
+ :param train_freq: How much experience to collect
190
+ by doing rollouts of current policy.
191
+ Either ``TrainFreq(<n>, TrainFrequencyUnit.STEP)``
192
+ or ``TrainFreq(<n>, TrainFrequencyUnit.EPISODE)``
193
+ with ``<n>`` being an integer greater than 0.
194
+ :param action_noise: Action noise that will be used for exploration
195
+ Required for deterministic policy (e.g. TD3). This can also be used
196
+ in addition to the stochastic policy for SAC.
197
+ :param learning_starts: Number of steps before learning for the warm-up phase.
198
+ :param replay_buffer:
199
+ :param log_interval: Log data every ``log_interval`` episodes
200
+ :return:
201
+ """
202
+ # Switch to eval mode (this affects batch norm / dropout)
203
+ self.policy.set_training_mode(False)
204
+
205
+ num_collected_steps, num_collected_episodes = 0, 0
206
+
207
+ self.learning_starts = learning_starts
208
+ self.action_noise = action_noise
209
+
210
+ assert isinstance(env, VecEnv), "You must pass a VecEnv"
211
+ assert train_freq.frequency > 0, "Should at least collect one step or episode."
212
+
213
+ if self.use_sde:
214
+ self.actor.reset_noise(1)
215
+
216
+ if not self.initialized:
217
+ self._initialize_threads()
218
+ self.initialized = True
219
+
220
+ callback.on_rollout_start()
221
+ continue_training = True
222
+ while should_collect_more_steps(
223
+ train_freq, num_collected_steps, num_collected_episodes
224
+ ):
225
+ if (
226
+ self.use_sde
227
+ and self.sde_sample_freq > 0
228
+ and num_collected_steps % self.sde_sample_freq == 0
229
+ ):
230
+ # Sample a new noise matrix
231
+ self.actor.reset_noise(1)
232
+
233
+ # Fetch transition (Also the only significant change to super)
234
+ transition: Transition = self.fetch_transition()
235
+
236
+ # Make locals available for callbacks
237
+ buffer_actions = transition.buffer_actions
238
+ self._last_obs = transition.last_obs
239
+ new_obs = transition.new_obs
240
+ rewards = transition.rewards
241
+ dones = transition.dones
242
+ infos = transition.infos
243
+
244
+ # Update stats
245
+ self.num_timesteps += 1
246
+ num_collected_steps += 1
247
+
248
+ # Give access to local variables
249
+ callback.update_locals(locals())
250
+
251
+ # Only stop training if return value is False, not when it is None.
252
+ if not callback.on_step():
253
+ return RolloutReturn(
254
+ num_collected_steps,
255
+ num_collected_episodes,
256
+ continue_training=False,
257
+ )
258
+
259
+ # Retrieve reward and episode length if using Monitor wrapper
260
+ self._update_info_buffer(infos, dones)
261
+
262
+ # Store data in replay buffer (normalized action and unnormalized observation)
263
+ self._custom_store_transition(
264
+ replay_buffer,
265
+ buffer_actions,
266
+ self._last_obs,
267
+ new_obs,
268
+ rewards,
269
+ dones,
270
+ infos,
271
+ )
272
+
273
+ self._update_current_progress_remaining(
274
+ self.num_timesteps, self._total_timesteps
275
+ )
276
+
277
+ # For DQN, check if the target network should be updated
278
+ # and update the exploration schedule
279
+ # For SAC/TD3, the update is dones as the same time as the gradient update
280
+ # see https://github.com/hill-a/stable-baselines/issues/900
281
+ self._on_step()
282
+
283
+ for idx, done in enumerate(dones):
284
+ if done:
285
+ # Update stats
286
+ num_collected_episodes += 1
287
+ self._episode_num += 1
288
+
289
+ if action_noise is not None:
290
+ action_noise.reset()
291
+
292
+ # Log training infos
293
+ if (
294
+ log_interval is not None
295
+ and self._episode_num % log_interval == 0
296
+ ):
297
+ self._dump_logs()
298
+
299
+ callback.on_rollout_end()
300
+
301
+ return RolloutReturn(
302
+ num_collected_steps,
303
+ num_collected_episodes,
304
+ continue_training,
305
+ )
@@ -0,0 +1,215 @@
1
+ from copy import deepcopy
2
+ from dataclasses import dataclass
3
+ from typing import Any, List
4
+
5
+ import numpy as np
6
+ import torch as th
7
+ from gymnasium import spaces
8
+ from stable_baselines3.common.buffers import RolloutBuffer
9
+ from stable_baselines3.common.callbacks import BaseCallback
10
+ from stable_baselines3.common.on_policy_algorithm import OnPolicyAlgorithm
11
+ from stable_baselines3.common.utils import obs_as_tensor
12
+ from stable_baselines3.common.vec_env import VecEnv
13
+
14
+ from async_gym_agents.agents.injector import AsyncAgentInjector
15
+
16
+
17
+ @dataclass
18
+ class Transition:
19
+ actions: list[any]
20
+ values: list[any]
21
+ log_probs: list[any]
22
+ last_obs: list[any]
23
+ new_obs: list[any]
24
+ rewards: list[float]
25
+ dones: list[bool]
26
+ last_dones: list[bool]
27
+ infos: list[Any]
28
+ index: int
29
+
30
+
31
+ class OnPolicyAlgorithmInjector(AsyncAgentInjector, OnPolicyAlgorithm):
32
+ def __init__(self, *args, max_steps_in_buffer: int = 8, **kwargs) -> None:
33
+ super().__init__(max_steps_in_buffer)
34
+ super(AsyncAgentInjector, self).__init__(*args, **kwargs)
35
+
36
+ def train(self, *args, **kwargs) -> None:
37
+ with self.policy_lock:
38
+ super().train()
39
+
40
+ def _excluded_save_params(self) -> List[str]:
41
+ return [
42
+ *super()._excluded_save_params(),
43
+ *super(AsyncAgentInjector, self)._excluded_save_params(),
44
+ ]
45
+
46
+ def _episode_generator(self, index: int) -> list[Transition]:
47
+ """
48
+ Continuously plays the game and returns episodes of Transitions
49
+ """
50
+ env = self.get_indexable_env()
51
+ last_obs = env.reset(index=index)
52
+ last_dones = np.ones((1,), dtype=bool)
53
+
54
+ episode = []
55
+
56
+ while self.running:
57
+ with self.policy_lock:
58
+ with th.no_grad():
59
+ # Convert to pytorch tensor or to TensorDict
60
+ obs_tensor = obs_as_tensor(last_obs, self.device)
61
+ actions, values, log_probs = self.policy(obs_tensor)
62
+ actions = actions.cpu().numpy()
63
+
64
+ # Rescale and perform action
65
+ clipped_actions = actions
66
+
67
+ if isinstance(self.action_space, spaces.Box):
68
+ if self.policy.squash_output:
69
+ # Unscale the actions to match env bounds
70
+ # if they were previously squashed (scaled in [-1, 1])
71
+ clipped_actions = self.policy.unscale_action(clipped_actions)
72
+ else:
73
+ # Otherwise, clip the actions to avoid out of bound error
74
+ # as we are sampling from an unbounded Gaussian distribution
75
+ clipped_actions = np.clip(
76
+ actions, self.action_space.low, self.action_space.high
77
+ )
78
+
79
+ new_obs, rewards, dones, infos = env.step(clipped_actions, index=index)
80
+
81
+ if isinstance(self.action_space, spaces.Discrete):
82
+ # Reshape in case of discrete action
83
+ actions = actions.reshape(-1, 1)
84
+
85
+ # Store transition
86
+ episode.append(
87
+ Transition(
88
+ actions,
89
+ values,
90
+ log_probs,
91
+ deepcopy(last_obs),
92
+ deepcopy(new_obs),
93
+ rewards,
94
+ dones,
95
+ last_dones,
96
+ infos,
97
+ index,
98
+ )
99
+ )
100
+ last_obs = new_obs
101
+ last_dones = dones
102
+
103
+ # Start new episode
104
+ if any(dones):
105
+ yield episode
106
+ episode = []
107
+
108
+ def collect_rollouts(
109
+ self,
110
+ env: VecEnv,
111
+ callback: BaseCallback,
112
+ rollout_buffer: RolloutBuffer,
113
+ n_rollout_steps: int,
114
+ ) -> bool:
115
+ """
116
+ Collect experiences using the current policy and fill a ``RolloutBuffer``.
117
+ The term rollout here refers to the model-free notion and should not
118
+ be used with the concept of rollout used in model-based RL or planning.
119
+
120
+ :param env: The training environment
121
+ :param callback: Callback that will be called at each step
122
+ (and at the beginning and end of the rollout)
123
+ :param rollout_buffer: Buffer to fill with rollouts
124
+ :param n_rollout_steps: Number of experiences to collect per environment
125
+ :return: True if function returned with at least `n_rollout_steps`
126
+ collected, False if callback terminated rollout prematurely.
127
+ """
128
+ assert self._last_obs is not None, "No previous observation was provided"
129
+
130
+ # Switch to eval mode (this affects batch norm / dropout)
131
+ self.policy.set_training_mode(False)
132
+
133
+ n_steps = 0
134
+ rollout_buffer.reset()
135
+
136
+ # Sample new weights for the state dependent exploration
137
+ if self.use_sde:
138
+ self.policy.reset_noise(1)
139
+
140
+ if not self.initialized:
141
+ self._initialize_threads()
142
+ self.initialized = True
143
+
144
+ callback.on_rollout_start()
145
+
146
+ new_obs = None
147
+ dones = None
148
+ while n_steps < n_rollout_steps:
149
+ if (
150
+ self.use_sde
151
+ and self.sde_sample_freq > 0
152
+ and n_steps % self.sde_sample_freq == 0
153
+ ):
154
+ # Sample a new noise matrix
155
+ self.policy.reset_noise(1)
156
+
157
+ # Fetch transitions from workers
158
+ transition: Transition = self.fetch_transition()
159
+
160
+ # Make locals available for callbacks
161
+ new_obs = transition.new_obs
162
+ self._last_obs = transition.last_obs
163
+ actions = transition.actions
164
+ rewards = transition.rewards
165
+ self._last_episode_starts = transition.last_dones
166
+ values = transition.values
167
+ log_probs = transition.log_probs
168
+ dones = transition.dones
169
+ infos = transition.infos
170
+
171
+ self.num_timesteps += 1
172
+
173
+ # Give access to local variables
174
+ callback.update_locals(locals())
175
+ if not callback.on_step():
176
+ return False
177
+
178
+ self._update_info_buffer(infos, dones)
179
+ n_steps += 1
180
+
181
+ # Handle timeout by bootstrapping with value function
182
+ # see GitHub issue #633
183
+ for idx, done in enumerate(dones):
184
+ if (
185
+ done
186
+ and infos[idx].get("terminal_observation") is not None
187
+ and infos[idx].get("TimeLimit.truncated", False)
188
+ ):
189
+ terminal_obs = self.policy.obs_to_tensor(
190
+ infos[idx]["terminal_observation"]
191
+ )[0]
192
+ with th.no_grad():
193
+ terminal_value = self.policy.predict_values(terminal_obs)[0]
194
+ rewards[idx] += self.gamma * terminal_value
195
+
196
+ rollout_buffer.add(
197
+ self._last_obs,
198
+ actions,
199
+ rewards,
200
+ self._last_episode_starts,
201
+ values,
202
+ log_probs,
203
+ )
204
+
205
+ with th.no_grad():
206
+ # Compute value for the last timestep
207
+ values = self.policy.predict_values(obs_as_tensor(new_obs, self.device))
208
+
209
+ rollout_buffer.compute_returns_and_advantage(last_values=values, dones=dones)
210
+
211
+ callback.update_locals(locals())
212
+
213
+ callback.on_rollout_end()
214
+
215
+ return True
@@ -0,0 +1,46 @@
1
+ from typing import Optional
2
+
3
+ import numpy as np
4
+ from gymnasium.envs.box2d import LunarLander
5
+
6
+
7
+ class BuggyLunarLander(LunarLander):
8
+ """
9
+ This environment fakes being buggy by randomly setting truncated to True.
10
+ """
11
+
12
+ def __init__(
13
+ self,
14
+ crash_probability: float = 0.01,
15
+ time_limit: int = 999999,
16
+ render_mode: Optional[str] = None,
17
+ ):
18
+ super().__init__(render_mode)
19
+
20
+ self.tick = 0
21
+ self.tick_until_crash = 0
22
+ self.crash_probability = crash_probability
23
+ self.time_limit = time_limit
24
+
25
+ def reset(self, *, seed: Optional[int] = None, options: Optional[dict] = None):
26
+ self.tick = 0
27
+ self.tick_until_crash = (
28
+ 99999
29
+ if self.crash_probability <= 0
30
+ else np.random.geometric(self.crash_probability)
31
+ )
32
+
33
+ return super().reset(seed=seed, options=options)
34
+
35
+ def step(self, action):
36
+ obs, reward, terminated, truncated, info = super().step(action)
37
+
38
+ self.tick += 1
39
+
40
+ if self.tick >= self.tick_until_crash and not terminated:
41
+ truncated = True
42
+
43
+ if self.tick >= self.time_limit and not truncated:
44
+ terminated = True
45
+
46
+ return obs, reward, terminated, truncated, info
@@ -0,0 +1,84 @@
1
+ from collections import defaultdict
2
+ from typing import Any, Callable, List, Optional, Sequence, Type
3
+
4
+ import numpy as np
5
+ from gymnasium import Env, Wrapper
6
+ from stable_baselines3.common.vec_env import DummyVecEnv, VecEnv
7
+ from stable_baselines3.common.vec_env.base_vec_env import (
8
+ VecEnvIndices,
9
+ VecEnvObs,
10
+ VecEnvStepReturn,
11
+ )
12
+
13
+
14
+ class IndexableMultiEnv(VecEnv):
15
+ """
16
+ Same as multi env but sync
17
+ """
18
+
19
+ def __init__(self, env_fns: List[Callable[[], Env]]):
20
+ self.real_n_envs = len(env_fns)
21
+
22
+ self.envs = [DummyVecEnv([e]) for e in env_fns]
23
+ self.additional = defaultdict(dict)
24
+
25
+ super().__init__(1, self.envs[0].observation_space, self.envs[0].action_space)
26
+
27
+ def step(self, actions: np.ndarray, index: int = 0) -> VecEnvStepReturn:
28
+ self.step_async(actions, index=index)
29
+ return self.step_wait(index=index)
30
+
31
+ def step_async(self, actions: np.ndarray, index: int = 0) -> None:
32
+ self.envs[index].step_async(actions)
33
+
34
+ def step_wait(self, index: int = 0) -> VecEnvStepReturn:
35
+ return self.envs[index].step_wait()
36
+
37
+ def reset(self, index: int = 0, **kwargs) -> VecEnvObs:
38
+ return self.envs[index].reset()
39
+
40
+ def close(self) -> None:
41
+ for env in self.envs:
42
+ env.close()
43
+
44
+ def get_images(self) -> Sequence[Optional[np.ndarray]]:
45
+ raise NotImplementedError
46
+
47
+ def get_attr(self, attr_name: str, indices: VecEnvIndices = None) -> List[Any]:
48
+ return self.envs[self._get_index(indices)].get_attr(attr_name)
49
+
50
+ def set_attr(
51
+ self, attr_name: str, value: Any, indices: VecEnvIndices = None
52
+ ) -> None:
53
+ self.envs[self._get_index(indices)].set_attr(attr_name, value)
54
+
55
+ def env_method(
56
+ self,
57
+ method_name: str,
58
+ *method_args,
59
+ indices: VecEnvIndices = None,
60
+ **method_kwargs,
61
+ ) -> List[Any]:
62
+ return self.envs[self._get_index(indices)].env_method(
63
+ *method_args, *method_kwargs
64
+ )
65
+
66
+ def env_is_wrapped(
67
+ self, wrapper_class: Type[Wrapper], indices: VecEnvIndices = None
68
+ ) -> List[bool]:
69
+ return self.envs[self._get_index(indices)].env_is_wrapped(wrapper_class)
70
+
71
+ def _get_index(self, indices: VecEnvIndices) -> int:
72
+ """
73
+ Convert a flexibly-typed reference to environment indices to an implied list of indices.
74
+
75
+ :param indices: refers to indices of envs.
76
+ :return: the implied list of indices.
77
+ """
78
+ if indices is None:
79
+ return 0
80
+ elif isinstance(indices, int):
81
+ return indices
82
+ raise ValueError(
83
+ f"IndexableMultiEnv only supports a scalar index, not {indices}."
84
+ )
@@ -0,0 +1,27 @@
1
+ import random
2
+ import time
3
+ from typing import Optional
4
+
5
+ from gymnasium.envs.classic_control import CartPoleEnv
6
+
7
+
8
+ class SlowCartPoleEnv(CartPoleEnv):
9
+ """
10
+ This environment fakes being slow.
11
+ """
12
+
13
+ def __init__(
14
+ self,
15
+ min_sleep: float = 0.01,
16
+ max_sleep: float = 0.1,
17
+ render_mode: Optional[str] = None,
18
+ ):
19
+ super().__init__(render_mode)
20
+
21
+ self.min_sleep = min_sleep
22
+ self.max_sleep = max_sleep
23
+
24
+ def step(self, action):
25
+ t = self.min_sleep + random.random() * (self.max_sleep - self.min_sleep)
26
+ time.sleep(t)
27
+ return super().step(action)
@@ -0,0 +1,210 @@
1
+ import threading
2
+ from queue import Queue
3
+ from typing import Any, Callable, Dict, List, Optional, Sequence, Tuple, Type, Union
4
+
5
+ import numpy as np
6
+ from gymnasium import Env, Wrapper, spaces
7
+ from stable_baselines3.common.vec_env import VecEnv
8
+ from stable_baselines3.common.vec_env.base_vec_env import (
9
+ VecEnvIndices,
10
+ VecEnvObs,
11
+ VecEnvStepReturn,
12
+ )
13
+
14
+ # noinspection PyProtectedMember
15
+ from stable_baselines3.common.vec_env.patch_gym import _patch_env
16
+
17
+
18
+ def _stack_observations(
19
+ obs: Union[List[VecEnvObs], Tuple[VecEnvObs]], space: spaces.Space
20
+ ) -> VecEnvObs:
21
+ """
22
+ Stack dict or tuple observations.
23
+ """
24
+ assert len(obs) > 0, "Observations list is empty!"
25
+
26
+ if isinstance(space, spaces.Dict):
27
+ return {
28
+ key: np.stack([single_obs[key] for single_obs in obs])
29
+ for key in space.spaces.keys()
30
+ }
31
+ elif isinstance(space, spaces.Tuple):
32
+ obs_len = len(space.spaces)
33
+ return tuple(
34
+ np.stack([single_obs[i] for single_obs in obs]) for i in range(obs_len)
35
+ )
36
+ else:
37
+ return np.stack(obs)
38
+
39
+
40
+ def _worker(
41
+ task_queue: Queue,
42
+ result_queue: Queue,
43
+ env: Callable,
44
+ ) -> None:
45
+ from stable_baselines3.common.env_util import is_wrapped
46
+
47
+ env: Env = _patch_env(env())
48
+
49
+ reset_info: Optional[Dict[str, Any]] = {}
50
+ while True:
51
+ try:
52
+ cmd, data = task_queue.get()
53
+ if cmd == "step":
54
+ observation, reward, terminated, truncated, info = env.step(data)
55
+ done = terminated or truncated
56
+ info["TimeLimit.truncated"] = truncated and not terminated
57
+ if done:
58
+ # save final observation where user can get it, then reset
59
+ info["terminal_observation"] = observation
60
+ observation, reset_info = env.reset()
61
+ result_queue.put((observation, reward, done, info, reset_info))
62
+ elif cmd == "reset":
63
+ maybe_options = {"options": data[1]} if data[1] else {}
64
+ observation, reset_info = env.reset(seed=data[0], **maybe_options)
65
+ result_queue.put((observation, reset_info))
66
+ elif cmd == "render":
67
+ result_queue.put(env.render())
68
+ elif cmd == "close":
69
+ env.close()
70
+ result_queue.join()
71
+ break
72
+ elif cmd == "get_spaces":
73
+ result_queue.put((env.observation_space, env.action_space))
74
+ elif cmd == "env_method":
75
+ method = getattr(env, data[0])
76
+ result_queue.put(method(*data[1], **data[2]))
77
+ elif cmd == "get_attr":
78
+ result_queue.put(getattr(env, data))
79
+ elif cmd == "set_attr":
80
+ result_queue.put(setattr(env, data[0], data[1]))
81
+ elif cmd == "is_wrapped":
82
+ result_queue.put(is_wrapped(env, data))
83
+ else:
84
+ raise NotImplementedError(f"`{cmd}` is not implemented in the worker")
85
+ except EOFError:
86
+ break
87
+
88
+
89
+ class ThreadedVecEnv(VecEnv):
90
+ def __init__(self, envs: List[Callable]):
91
+ self.waiting = False
92
+ self.closed = False
93
+ n_envs = len(envs)
94
+
95
+ self.task_queues = [Queue() for _ in range(n_envs)]
96
+ self.result_queues = [Queue() for _ in range(n_envs)]
97
+
98
+ self.threads = []
99
+ for task_queue, result_queue, env in zip(
100
+ self.task_queues, self.result_queues, envs
101
+ ):
102
+ args = (task_queue, result_queue, env)
103
+ thread = threading.Thread(target=_worker, args=args, daemon=True)
104
+ thread.start()
105
+ self.threads.append(thread)
106
+
107
+ self.task_queues[0].put(("get_spaces", None))
108
+ observation_space, action_space = self.result_queues[0].get()
109
+
110
+ super().__init__(len(envs), observation_space, action_space)
111
+
112
+ def step_async(self, actions: np.ndarray) -> None:
113
+ for queue, action in zip(self.task_queues, actions):
114
+ queue.put(("step", action))
115
+ self.waiting = True
116
+
117
+ def step_wait(self) -> VecEnvStepReturn:
118
+ results = [queue.get() for queue in self.result_queues]
119
+ self.waiting = False
120
+ obs, rewards, dones, infos, self.reset_infos = zip(*results) # type: ignore[assignment]
121
+ return (
122
+ _stack_observations(obs, self.observation_space),
123
+ np.stack(rewards),
124
+ np.stack(dones),
125
+ infos,
126
+ ) # type: ignore[return-value]
127
+
128
+ def reset(
129
+ self, *, seed: int | None = None, options: dict[str, Any] | None = None
130
+ ) -> VecEnvObs:
131
+ for env_idx, queue in enumerate(self.task_queues):
132
+ queue.put(("reset", (self._seeds[env_idx], self._options[env_idx])))
133
+ results = [queue.get() for queue in self.result_queues]
134
+ obs, self.reset_infos = zip(*results) # type: ignore[assignment]
135
+ # Seeds and options are only used once
136
+ self._reset_seeds()
137
+ self._reset_options()
138
+ return _stack_observations(obs, self.observation_space) # , self.reset_infos
139
+
140
+ def close(self) -> None:
141
+ if self.closed:
142
+ return
143
+ if self.waiting:
144
+ for queue in self.result_queues:
145
+ queue.get()
146
+ for queue in self.task_queues:
147
+ queue.put(("close", None))
148
+ for thread in self.threads:
149
+ thread.join()
150
+ self.closed = True
151
+
152
+ def get_images(self) -> Sequence[Optional[np.ndarray]]:
153
+ raise NotImplementedError
154
+
155
+ def get_attr(self, attr_name: str, indices: VecEnvIndices = None) -> List[Any]:
156
+ """Return attribute from vectorized environment (see base class)."""
157
+ for queue in self._get_target_queues(self.task_queues, indices):
158
+ queue.put(("get_attr", attr_name))
159
+ return [
160
+ queue.get()
161
+ for queue in self._get_target_queues(self.result_queues, indices)
162
+ ]
163
+
164
+ def set_attr(
165
+ self, attr_name: str, value: Any, indices: VecEnvIndices = None
166
+ ) -> None:
167
+ """Set attribute inside vectorized environments (see base class)."""
168
+ for queue in self._get_target_queues(self.task_queues, indices):
169
+ queue.put(("set_attr", (attr_name, value)))
170
+ for queue in self._get_target_queues(self.result_queues, indices):
171
+ queue.get()
172
+
173
+ def env_method(
174
+ self,
175
+ method_name: str,
176
+ *method_args,
177
+ indices: VecEnvIndices = None,
178
+ **method_kwargs,
179
+ ) -> List[Any]:
180
+ """Call instance methods of vectorized environments."""
181
+ for queue in self._get_target_queues(self.task_queues, indices):
182
+ queue.put(("env_method", (method_name, method_args, method_kwargs)))
183
+ return [
184
+ queue.get()
185
+ for queue in self._get_target_queues(self.result_queues, indices)
186
+ ]
187
+
188
+ def env_is_wrapped(
189
+ self, wrapper_class: Type[Wrapper], indices: VecEnvIndices = None
190
+ ) -> List[bool]:
191
+ """Check if worker environments are wrapped with a given wrapper"""
192
+ for queue in self._get_target_queues(self.task_queues, indices):
193
+ queue.put(("is_wrapped", wrapper_class))
194
+ return [
195
+ queue.get()
196
+ for queue in self._get_target_queues(self.result_queues, indices)
197
+ ]
198
+
199
+ def _get_target_queues(
200
+ self, queues: list[Queue], indices: VecEnvIndices
201
+ ) -> List[Any]:
202
+ """
203
+ Get the connection object needed to communicate with the wanted
204
+ envs that are in subprocesses.
205
+
206
+ :param indices: Refers to indices of envs.
207
+ :return: Connection object to communicate between processes.
208
+ """
209
+ indices = self._get_indices(indices)
210
+ return [queues[i] for i in indices]
@@ -0,0 +1,16 @@
1
+ [tool.poetry]
2
+ name = "async-gym-agents"
3
+ version = "0.1.0"
4
+ description = "Async agents for Stable Baselines 3"
5
+ authors = ["Jonas Peche <jonas.peche@aon.at>"]
6
+ readme = "README.md"
7
+
8
+ [tool.poetry.dependencies]
9
+ python = ">=3.8"
10
+ stable-baselines3 = { extras = ["extra"], version = "^2.3.2" }
11
+ torchinfo = "^1.8.0"
12
+
13
+
14
+ [build-system]
15
+ requires = ["poetry-core"]
16
+ build-backend = "poetry.core.masonry.api"