Spaces:
Sleeping
Sleeping
File size: 10,563 Bytes
377b913 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 | import torch
import numpy as np
from expreplay import ReplayMemory
from DQNModel import DQN
from evaluator import Evaluator
from tqdm import tqdm
class Trainer(object):
def __init__(self,
env,
eval_env=None,
image_size=(45, 45, 45),
update_frequency=4,
replay_buffer_size=1e6,
init_memory_size=5e4,
max_episodes=100,
steps_per_episode=50,
eps=1,
min_eps=0.1,
delta=0.001,
batch_size=4,
gamma=0.9,
number_actions=6,
frame_history=4,
model_name="CommNet",
logger=None,
train_freq=1,
team_reward=False,
attention=False,
lr=1e-3,
scheduler_gamma=0.5,
scheduler_step_size=100,
checkpoint=None
):
self.env = env
self.eval_env = eval_env
self.agents = env.agents
self.image_size = image_size
self.update_frequency = update_frequency
self.replay_buffer_size = replay_buffer_size
self.init_memory_size = init_memory_size
self.max_episodes = max_episodes
self.steps_per_episode = steps_per_episode
self.eps = eps
self.min_eps = min_eps
self.delta = delta
self.batch_size = batch_size
self.gamma = gamma
self.number_actions = number_actions
self.frame_history = frame_history
self.epoch_length = self.env.files.num_files
self.best_val_distance = float('inf')
self.buffer = ReplayMemory(
self.replay_buffer_size,
self.image_size,
self.frame_history,
self.agents)
self.dqn = DQN(
self.agents,
self.frame_history,
logger=logger,
type=model_name,
collective_rewards=team_reward,
attention=attention,
lr=lr,
scheduler_gamma=scheduler_gamma,
scheduler_step_size=scheduler_step_size)
self.dqn.q_network.train(True)
self.logger = logger
self.train_freq = train_freq
self.start_episode = 1
self.start_acc_steps = 0
self.start_eps = eps
if checkpoint is not None:
if self.logger:
self.logger.log("Restoring from checkpoint...")
self.dqn.q_network.load_state_dict(checkpoint['q_network_state_dict'])
if 'target_network_state_dict' in checkpoint:
self.dqn.target_network.load_state_dict(checkpoint['target_network_state_dict'])
if 'optimiser_state_dict' in checkpoint:
self.dqn.optimiser.load_state_dict(checkpoint['optimiser_state_dict'])
if 'scheduler_state_dict' in checkpoint:
self.dqn.scheduler.load_state_dict(checkpoint['scheduler_state_dict'])
self.start_episode = checkpoint.get('episode', 0) + 1
self.start_acc_steps = checkpoint.get('acc_steps', 0)
self.start_eps = checkpoint.get('eps', self.eps)
if self.logger:
self.logger.log(f"Resumed from episode {checkpoint.get('episode', 0)}, step {checkpoint.get('acc_steps', 0)}")
self.evaluator = Evaluator(eval_env,
self.dqn.q_network,
logger,
self.agents,
steps_per_episode)
def train(self):
self.logger.log(self.dqn.q_network)
self.init_memory()
episode = self.start_episode
acc_steps = self.start_acc_steps
self.eps = self.start_eps
epoch_distances = []
while episode <= self.max_episodes:
# Reset the environment for the start of the episode.
obs = self.env.reset()
self.buffer._hist.clear()
terminal = [False for _ in range(self.agents)]
losses = []
score = [0] * self.agents
for step_num in range(self.steps_per_episode):
acc_steps += 1
# Step the agent once, and get the transition tuple
index = self.buffer.append_obs(obs)
acts, q_values = self.get_next_actions(
self.buffer.recent_state())
next_obs, reward, terminal, info = self.env.step(
np.copy(acts), q_values, terminal)
self.buffer.append_effect((index, obs, acts, reward, terminal))
score = [sum(x) for x in zip(score, reward)]
obs = next_obs
if acc_steps % self.train_freq == 0:
mini_batch = self.buffer.sample(self.batch_size)
loss = self.dqn.train_q_network(mini_batch, self.gamma)
losses.append(loss)
if all(t for t in terminal):
break
epoch_distances.append([info['distError_' + str(i)]
for i in range(self.agents)])
self.append_episode_board(info, score, "train", episode)
if (episode * self.epoch_length) % self.update_frequency == 0:
self.dqn.copy_to_target_network()
self.eps = max(self.min_eps, self.eps - self.delta)
# Every epoch
if episode % self.epoch_length == 0:
self.append_epoch_board(epoch_distances, self.eps, losses,
"train", episode)
self.validation_epoch(episode)
self.dqn.save_model(name="latest_dqn.pt", forced=True)
self.dqn.save_checkpoint(
name="latest_checkpoint.pt",
episode=episode,
eps=self.eps,
acc_steps=acc_steps,
forced=True)
self.dqn.scheduler.step()
epoch_distances = []
episode += 1
def init_memory(self):
self.logger.log("Initialising memory buffer...")
pbar = tqdm(desc="Memory buffer", total=self.init_memory_size)
while len(self.buffer) < self.init_memory_size:
# Reset the environment for the start of the episode.
obs = self.env.reset()
self.buffer._hist.clear()
terminal = [False for _ in range(self.agents)]
steps = 0
for _ in range(self.steps_per_episode):
steps += 1
index = self.buffer.append_obs(obs)
acts, q_values = self.get_next_actions(obs)
next_obs, reward, terminal, info = self.env.step(
acts, q_values, terminal)
self.buffer.append_effect((index, obs, acts, reward, terminal))
obs = next_obs
if all(t for t in terminal):
break
pbar.update(steps)
pbar.close()
self.logger.log("Memory buffer filled")
def validation_epoch(self, episode):
if self.eval_env is None:
return
self.dqn.q_network.train(False)
epoch_distances = []
for k in range(self.eval_env.files.num_files):
self.logger.log(f"eval episode {k}")
(score, start_dists, q_values,
info) = self.evaluator.play_one_episode()
epoch_distances.append([info['distError_' + str(i)]
for i in range(self.agents)])
val_dists = self.append_epoch_board(epoch_distances, name="eval",
episode=episode)
if (val_dists < self.best_val_distance):
self.logger.log("Improved new best mean validation distances")
self.best_val_distance = val_dists
self.dqn.save_model(name="best_dqn.pt", forced=True)
self.dqn.q_network.train(True)
def append_episode_board(self, info, score, name="train", episode=0):
dists = {str(i):
info['distError_' + str(i)] for i in range(self.agents)}
self.logger.write_to_board(f"{name}/dist", dists, episode)
scores = {str(i): score[i] for i in range(self.agents)}
self.logger.write_to_board(f"{name}/score", scores, episode)
def append_epoch_board(self, epoch_dists, eps=0, losses=[],
name="train", episode=0):
epoch_dists = np.array(epoch_dists)
if name == "train":
lr = self.dqn.scheduler.state_dict()["_last_lr"]
if isinstance(lr, list):
lr = lr[0]
self.logger.write_to_board(name, {"eps": eps, "lr": lr}, episode)
if len(losses) > 0:
loss_dict = {"loss": sum(losses) / len(losses)}
self.logger.write_to_board(name, loss_dict, episode)
for i in range(self.agents):
mean_dist = sum(epoch_dists[:, i]) / len(epoch_dists[:, i])
mean_dist_dict = {str(i): mean_dist}
self.logger.write_to_board(
f"{name}/mean_dist", mean_dist_dict, episode)
min_dist_dict = {str(i): min(epoch_dists[:, i])}
self.logger.write_to_board(
f"{name}/min_dist", min_dist_dict, episode)
max_dist_dict = {str(i): max(epoch_dists[:, i])}
self.logger.write_to_board(
f"{name}/max_dist", max_dist_dict, episode)
return np.array(list(mean_dist_dict.values())).mean()
def get_next_actions(self, obs_stack):
# epsilon-greedy policy
if np.random.random() < self.eps:
q_values = np.zeros((self.agents, self.number_actions))
actions = np.random.randint(self.number_actions, size=self.agents)
else:
actions, q_values = self.get_greedy_actions(
obs_stack, doubleLearning=True)
return actions, q_values
def get_greedy_actions(self, obs_stack, doubleLearning=True):
inputs = torch.tensor(obs_stack).unsqueeze(0)
if doubleLearning:
q_vals = self.dqn.q_network.forward(inputs).detach().squeeze(0)
else:
q_vals = self.dqn.target_network.forward(
inputs).detach().squeeze(0)
idx = torch.max(q_vals, -1)[1]
greedy_steps = np.array(idx, dtype=np.int32).flatten()
return greedy_steps, q_vals.data.numpy()
|