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.
- async_gym_agents-0.1.0/PKG-INFO +41 -0
- async_gym_agents-0.1.0/README.md +23 -0
- async_gym_agents-0.1.0/async_gym_agents/agents/__pycache__/async_agent.cpython-310.pyc +0 -0
- async_gym_agents-0.1.0/async_gym_agents/agents/__pycache__/async_agent.cpython-311.pyc +0 -0
- async_gym_agents-0.1.0/async_gym_agents/agents/__pycache__/injector.cpython-311.pyc +0 -0
- async_gym_agents-0.1.0/async_gym_agents/agents/__pycache__/off_policy_injector.cpython-311.pyc +0 -0
- async_gym_agents-0.1.0/async_gym_agents/agents/__pycache__/on_policy_injector.cpython-311.pyc +0 -0
- async_gym_agents-0.1.0/async_gym_agents/agents/async_agent.py +25 -0
- async_gym_agents-0.1.0/async_gym_agents/agents/injector.py +125 -0
- async_gym_agents-0.1.0/async_gym_agents/agents/off_policy_injector.py +305 -0
- async_gym_agents-0.1.0/async_gym_agents/agents/on_policy_injector.py +215 -0
- async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/buggy_lunar_lander.cpython-311.pyc +0 -0
- async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/multi_env.cpython-310.pyc +0 -0
- async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/multi_env.cpython-311.pyc +0 -0
- async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/slow_cartpole.cpython-311.pyc +0 -0
- async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/threaded_env.cpython-310.pyc +0 -0
- async_gym_agents-0.1.0/async_gym_agents/envs/__pycache__/threaded_env.cpython-311.pyc +0 -0
- async_gym_agents-0.1.0/async_gym_agents/envs/buggy_lunar_lander.py +46 -0
- async_gym_agents-0.1.0/async_gym_agents/envs/multi_env.py +84 -0
- async_gym_agents-0.1.0/async_gym_agents/envs/slow_cartpole.py +27 -0
- async_gym_agents-0.1.0/async_gym_agents/envs/threaded_env.py +210 -0
- 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
|
+
```
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
async_gym_agents-0.1.0/async_gym_agents/agents/__pycache__/off_policy_injector.cpython-311.pyc
ADDED
|
Binary file
|
|
Binary file
|
|
@@ -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
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
Binary file
|
|
@@ -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"
|