From 0ad743c85d3c2dbdfab8f18b6c519a45751b4a08 Mon Sep 17 00:00:00 2001 From: Antonin Raffin Date: Fri, 25 Oct 2019 10:59:15 +0200 Subject: [PATCH] Add A2C --- README.md | 8 +-- tests/test_run.py | 7 +- torchy_baselines/__init__.py | 3 +- torchy_baselines/a2c/__init__.py | 2 + torchy_baselines/a2c/a2c.py | 109 +++++++++++++++++++++++++++++++ 5 files changed, 120 insertions(+), 9 deletions(-) create mode 100644 torchy_baselines/a2c/__init__.py create mode 100644 torchy_baselines/a2c/a2c.py diff --git a/README.md b/README.md index a73f690..b5624ec 100644 --- a/README.md +++ b/README.md @@ -8,6 +8,7 @@ PyTorch version of [Stable Baselines](https://github.com/hill-a/stable-baselines ## Implemented Algorithms +- A2C - CEM-RL (with TD3) - PPO - SAC @@ -18,11 +19,8 @@ PyTorch version of [Stable Baselines](https://github.com/hill-a/stable-baselines TODO: - save/load -- predict -- flexible mlp -- logger -- better monitor wrapper? -- A2C +- better predict +- complete logger Later: - get_parameters / set_parameters diff --git a/tests/test_run.py b/tests/test_run.py index 9740921..32a4b30 100644 --- a/tests/test_run.py +++ b/tests/test_run.py @@ -3,7 +3,7 @@ import os import pytest import numpy as np -from torchy_baselines import TD3, CEMRL, PPO, SAC +from torchy_baselines import A2C, CEMRL, PPO, SAC, TD3 from torchy_baselines.common.noise import NormalActionNoise @@ -28,9 +28,10 @@ def test_cemrl(): os.remove("test_save.pth") +@pytest.mark.parametrize("model_class", [A2C, PPO]) @pytest.mark.parametrize("env_id", ['CartPole-v1', 'Pendulum-v0']) -def test_ppo(env_id): - model = PPO('MlpPolicy', env_id, policy_kwargs=dict(net_arch=[16]), verbose=1, create_eval_env=True) +def test_onpolicy(model_class, env_id): + model = model_class('MlpPolicy', env_id, policy_kwargs=dict(net_arch=[16]), verbose=1, create_eval_env=True) model.learn(total_timesteps=1000, eval_freq=500) # model.save("test_save") # model.load("test_save") diff --git a/torchy_baselines/__init__.py b/torchy_baselines/__init__.py index b9dabaa..a5896e6 100644 --- a/torchy_baselines/__init__.py +++ b/torchy_baselines/__init__.py @@ -1,6 +1,7 @@ +from torchy_baselines.a2c import A2C from torchy_baselines.cem_rl import CEMRL from torchy_baselines.ppo import PPO from torchy_baselines.sac import SAC from torchy_baselines.td3 import TD3 -__version__ = "0.0.4" +__version__ = "0.0.5a" diff --git a/torchy_baselines/a2c/__init__.py b/torchy_baselines/a2c/__init__.py new file mode 100644 index 0000000..0cc4be0 --- /dev/null +++ b/torchy_baselines/a2c/__init__.py @@ -0,0 +1,2 @@ +from torchy_baselines.a2c.a2c import A2C +from torchy_baselines.ppo.policies import MlpPolicy diff --git a/torchy_baselines/a2c/a2c.py b/torchy_baselines/a2c/a2c.py new file mode 100644 index 0000000..4de140b --- /dev/null +++ b/torchy_baselines/a2c/a2c.py @@ -0,0 +1,109 @@ +from gym import spaces +import torch as th +import torch.nn.functional as F + +from torchy_baselines.common.utils import explained_variance +from torchy_baselines.ppo.ppo import PPO +from torchy_baselines.ppo.policies import PPOPolicy + + +class A2C(PPO): + """ + Advantage Actor Critic (A2C) + + Paper: https://arxiv.org/abs/1602.01783 + Code: This implementation borrows code from https://github.com/ikostrikov/pytorch-a2c-ppo-acktr-gail and + and Stable Baselines (https://github.com/hill-a/stable-baselines) + + Introduction to A2C: https://hackernoon.com/intuitive-rl-intro-to-advantage-actor-critic-a2c-4ff545978752 + + :param policy: (PPOPolicy or str) The policy model to use (MlpPolicy, CnnPolicy, ...) + :param env: (Gym environment or str) The environment to learn from (if registered in Gym, can be str) + :param learning_rate: (float or callable) The learning rate, it can be a function + :param n_steps: (int) The number of steps to run for each environment per update + (i.e. batch size is n_steps * n_env where n_env is number of environment copies running in parallel) + :param batch_size: (int) Minibatch size + :param n_epochs: (int) Number of epoch when optimizing the surrogate loss + :param gamma: (float) Discount factor + :param gae_lambda: (float) Factor for trade-off of bias vs variance for Generalized Advantage Estimator + :param ent_coef: (float) Entropy coefficient for the loss calculation + :param vf_coef: (float) Value function coefficient for the loss calculation + :param max_grad_norm: (float) The maximum value for the gradient clipping + :param tensorboard_log: (str) the log location for tensorboard (if None, no logging) + :param create_eval_env: (bool) Whether to create a second environment that will be + used for evaluating the agent periodically. (Only available when passing string for the environment) + :param policy_kwargs: (dict) additional arguments to be passed to the policy on creation + :param verbose: (int) the verbosity level: 0 none, 1 training information, 2 tensorflow debug + :param seed: (int) Seed for the pseudo random generators + :param device: (str or th.device) Device (cpu, cuda, ...) on which the code should be run. + Setting it to auto, the code will be run on the GPU if possible. + :param _init_setup_model: (bool) Whether or not to build the network at the creation of the instance + """ + + def __init__(self, policy, env, learning_rate=3e-4, + n_steps=2048, batch_size=64, n_epochs=1, + gamma=0.99, gae_lambda=0.95, + ent_coef=0.0, vf_coef=0.5, max_grad_norm=0.5, + tensorboard_log=None, create_eval_env=False, + policy_kwargs=None, verbose=0, seed=0, device='auto', + _init_setup_model=True): + + super(A2C, self).__init__(policy, env, learning_rate=learning_rate, + n_steps=n_steps, batch_size=batch_size, n_epochs=n_epochs, + gamma=gamma, gae_lambda=gae_lambda, ent_coef=ent_coef, + vf_coef=vf_coef, max_grad_norm=max_grad_norm, + tensorboard_log=tensorboard_log, policy_kwargs=policy_kwargs, + verbose=verbose, device=device, create_eval_env=create_eval_env, + seed=seed, _init_setup_model=False) + + self.batch_size = n_steps + + if _init_setup_model: + self._setup_model() + + def train(self, gradient_steps, batch_size=64): + + for gradient_step in range(gradient_steps): + # approx_kl_divs = [] + # Sample replay buffer + for replay_data in self.rollout_buffer.get(batch_size): + # Unpack + obs, action, _, _, advantage, return_batch = replay_data + + if isinstance(self.action_space, spaces.Discrete): + # Convert discrete action for float to long + action = action.long().flatten() + + values, log_prob, entropy = self.policy.get_policy_stats(obs, action) + values = values.flatten() + # Normalize advantage + # TODO: check without + advantage = (advantage - advantage.mean()) / (advantage.std() + 1e-8) + + policy_loss = -(advantage * log_prob).mean() + + # Value loss using the TD(gae_lambda) target + value_loss = F.mse_loss(return_batch, values) + + # Entropy loss favor exploration + entropy_loss = th.mean(entropy) + + loss = policy_loss + self.ent_coef * entropy_loss + self.vf_coef * value_loss + + # Optimization step + self.policy.optimizer.zero_grad() + loss.backward() + # Clip grad norm + th.nn.utils.clip_grad_norm_(self.policy.parameters(), self.max_grad_norm) + self.policy.optimizer.step() + # approx_kl_divs.append(th.mean(old_log_prob - log_prob).detach().cpu().numpy()) + + # print(explained_variance(self.rollout_buffer.returns.flatten().cpu().numpy(), + # self.rollout_buffer.values.flatten().cpu().numpy())) + + def learn(self, total_timesteps, callback=None, log_interval=100, + eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="A2C", reset_num_timesteps=True): + + return super(A2C, self).learn(total_timesteps=total_timesteps, callback=callback, log_interval=log_interval, + eval_env=eval_env, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes, + tb_log_name=tb_log_name, reset_num_timesteps=reset_num_timesteps)