diff --git a/agent/async_agent.py b/agent/async_agent.py index fe2d738..e4bf2a7 100644 --- a/agent/async_agent.py +++ b/agent/async_agent.py @@ -25,18 +25,18 @@ def train(id, config, learning_network, target_network): config.logger.debug('worker %d, episode %d, return %f, avg return %f, episode steps %d, total steps %d' % ( id, episode, rewards[-1], np.mean(rewards[-100:]), steps, config.total_steps.value)) -def evaluate(config, task, learning_network): +def evaluate(config, task, actor, critic): test_rewards = [] test_points = [] - worker = config.worker(config, learning_network, None) + worker = config.worker(config, actor, critic) # config.logger = Logger('./evaluation_log', gym.logger) while True: steps = config.total_steps.value if steps % config.test_interval == 0: - worker.worker_network.load_state_dict(learning_network.state_dict()) - with open('data/%s-%s-model-%s.bin' % ( - config.tag, config.worker.__name__, task.name), 'wb') as f: - pickle.dump(learning_network.state_dict(), f) + # worker.worker_network.load_state_dict(learning_network.state_dict()) + # with open('data/%s-%s-model-%s.bin' % ( + # config.tag, config.worker.__name__, task.name), 'wb') as f: + # pickle.dump(learning_network.state_dict(), f) rewards = np.zeros(config.test_repetitions) for i in range(config.test_repetitions): rewards[i] = worker.episode(deterministic=True)[1] @@ -62,15 +62,17 @@ class AsyncAgent: def run(self): config = self.config task = config.task_fn() - learning_network = config.network_fn() - learning_network.share_memory() - target_network = config.network_fn() - target_network.share_memory() - target_network.load_state_dict(learning_network.state_dict()) + actor = config.actor_fn() + actor.share_memory() + critic = config.critic_fn() + critic.share_memory() + # target_network = config.network_fn() + # target_network.share_memory() + # target_network.load_state_dict(learning_network.state_dict()) os.environ['OMP_NUM_THREADS'] = '1' - args = [(i, config, learning_network, target_network) for i in range(config.num_workers)] - args.append((config, task, learning_network)) + args = [(i, config, actor, critic) for i in range(config.num_workers)] + args.append((config, task, actor, critic)) procs = [mp.Process(target=train, args=args[i]) for i in range(config.num_workers)] procs.append(mp.Process(target=evaluate, args=args[-1])) for p in procs: p.start() diff --git a/async_worker/actor_critic.py b/async_worker/actor_critic.py index 93ffa6f..4971a92 100644 --- a/async_worker/actor_critic.py +++ b/async_worker/actor_critic.py @@ -9,14 +9,19 @@ from torch.autograd import Variable import torch.nn as nn class AdvantageActorCritic: - def __init__(self, config, learning_network, target_network): + def __init__(self, config, actor, critic): self.config = config - self.optimizer = config.optimizer_fn(learning_network.parameters()) - self.worker_network = config.network_fn() - self.worker_network.load_state_dict(learning_network.state_dict()) + self.actor_opt = torch.optim.SGD(actor.parameters(), lr=0.0001) + self.critic_opt = torch.optim.SGD(critic.parameters(), lr=0.0001) + # self.optimizer = config.optimizer_fn(learning_network.parameters()) + self.worker_actor = config.actor_fn() + self.worker_actor.load_state_dict(actor.state_dict()) + self.worker_critic = config.critic_fn() + self.worker_critic.load_state_dict(critic.state_dict()) self.task = config.task_fn() self.policy = config.policy_fn() - self.learning_network = learning_network + self.actor = actor + self.critic = critic def episode(self, deterministic=False): config = self.config @@ -26,7 +31,9 @@ class AdvantageActorCritic: pending = [] while not config.stop_signal.value and \ (not config.max_episode_length or steps < config.max_episode_length): - prob, log_prob, value = self.worker_network.predict(np.stack([state])) + prob, log_prob = self.actor.predict(np.stack([state])) + value = self.critic.predict(np.stack([state])) + # prob, log_prob, value = self.worker_network.predict(np.stack([state])) action = self.policy.sample(prob.data.numpy().flatten(), deterministic) next_state, reward, terminal, _ = self.task.step(action) @@ -44,38 +51,53 @@ class AdvantageActorCritic: config.total_steps.value += 1 if terminal or len(pending) >= config.update_interval: - loss = 0 + policy_loss = 0 + value_loss = 0 if terminal: R = torch.FloatTensor([[0]]) else: - R = self.worker_network.critic(np.stack([next_state])).data + R = self.worker_critic.predict(np.stack([next_state])).data GAE = torch.FloatTensor([[0]]) for i in reversed(range(len(pending))): prob, log_prob, value, action, reward = pending[i] - if i == len(pending) - 1: - delta = reward + config.discount * R - value.data - else: - delta = reward + pending[i + 1][2].data - value.data - GAE = config.discount * config.gae_tau * GAE + delta - loss += -log_prob.gather(1, Variable(torch.LongTensor([[action]]))) * Variable(GAE) - loss += config.entropy_weight * torch.sum(torch.mul(prob, log_prob)) + R = config.discount * R + reward + GAE = R - value.data + # if i == len(pending) - 1: + # delta = reward + config.discount * R - value.data + # else: + # delta = reward + pending[i + 1][2].data - value.data + # GAE = config.discount * config.gae_tau * GAE + delta + policy_loss += -log_prob.gather(1, Variable(torch.LongTensor([[action]]))) * Variable(GAE) + policy_loss += config.entropy_weight * torch.sum(torch.mul(prob, log_prob)) R = reward + config.discount * R - loss += 0.5 * (Variable(R) - value).pow(2) + value_loss += 0.5 * (Variable(R) - value).pow(2) pending = [] - self.worker_network.zero_grad() - self.optimizer.zero_grad() - loss.backward() - nn.utils.clip_grad_norm(self.worker_network.parameters(), config.gradient_clip) + self.worker_actor.zero_grad() + self.actor_opt.zero_grad() + policy_loss.backward() + nn.utils.clip_grad_norm(self.worker_actor.parameters(), config.gradient_clip) for param, worker_param in zip( - self.learning_network.parameters(), self.worker_network.parameters()): + self.actor.parameters(), self.worker_actor.parameters()): if param.grad is not None: break param._grad = worker_param.grad - self.optimizer.step() - self.worker_network.load_state_dict(self.learning_network.state_dict()) - self.worker_network.reset(terminal) + self.actor_opt.step() + self.worker_actor.load_state_dict(self.actor.state_dict()) + # self.worker_network.reset(terminal) + + self.worker_critic.zero_grad() + self.critic_opt.zero_grad() + value_loss.backward() + nn.utils.clip_grad_norm(self.worker_critic.parameters(), config.gradient_clip) + for param, worker_param in zip( + self.critic.parameters(), self.worker_critic.parameters()): + if param.grad is not None: + break + param._grad = worker_param.grad + self.critic_opt.step() + self.worker_critic.load_state_dict(self.critic.state_dict()) if terminal: break diff --git a/main.py b/main.py index c21b0a7..b6b7ff9 100644 --- a/main.py +++ b/main.py @@ -44,11 +44,41 @@ def async_cart_pole(): agent = AsyncAgent(config) agent.run() +HIDDEN = 100 +class ActorNet(nn.Module, BasicNet): + def __init__(self): + super(ActorNet, self).__init__() + self.fc = nn.Linear(4, HIDDEN) + self.actor = nn.Linear(HIDDEN, 2) + BasicNet.__init__(self, None, False, False) + + def predict(self, x): + x = self.to_torch_variable(x) + x = F.relu(self.fc(x)) + x = self.actor(x) + prob = F.softmax(x) + log_prob = F.log_softmax(x) + return prob, log_prob + + +class CriticNet(nn.Module, BasicNet): + def __init__(self): + super(CriticNet, self).__init__() + self.fc = nn.Linear(4, HIDDEN) + self.critic = nn.Linear(HIDDEN, 1) + BasicNet.__init__(self, None, False, False) + + def predict(self, x): + x = self.to_torch_variable(x) + x = F.relu(self.fc(x)) + value = self.critic(x) + return value + def a3c_cart_pole(): config = Config() config.task_fn = lambda: CartPole() - config.optimizer_fn = lambda params: torch.optim.Adam(params, 0.001) - config.network_fn = lambda: ActorCriticFCNet(4, 2) + config.actor_fn = lambda: ActorNet() + config.critic_fn = lambda: CriticNet() config.policy_fn = SamplePolicy config.worker = AdvantageActorCritic config.discount = 0.99 @@ -233,10 +263,10 @@ if __name__ == '__main__': # dqn_cart_pole() # async_cart_pole() - # a3c_cart_pole() + a3c_cart_pole() # a3c_pendulum() # a3c_walker() - ddpg_pendulum() + # ddpg_pendulum() # ddpg_walker() # dqn_pixel_atari('PongNoFrameskip-v3')