Communicative_CRN / src /trainer.py
Sanni Henry
Initial deploy: Gradio landmark detection demo
377b913
Raw
History Blame Contribute Delete
10.6 kB
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()