diff --git a/README.md b/README.md index 094901b..64dfbb9 100644 --- a/README.md +++ b/README.md @@ -7,7 +7,6 @@ PyTorch version of [Stable Baselines](https://github.com/hill-a/stable-baselines), a set of improved implementations of reinforcement learning algorithms. TODO: -- SAC - save/load - automatic choice for action distribution - predict diff --git a/setup.py b/setup.py index 994d949..89a3b57 100644 --- a/setup.py +++ b/setup.py @@ -34,7 +34,7 @@ setup(name='torchy_baselines', license="MIT", long_description="", long_description_content_type='text/markdown', - version="0.0.3", + version="0.0.4", ) # python setup.py sdist diff --git a/tests/test_run.py b/tests/test_run.py index 546a3c0..439eb86 100644 --- a/tests/test_run.py +++ b/tests/test_run.py @@ -1,12 +1,12 @@ import os -from torchy_baselines import TD3, CEMRL, PPO +from torchy_baselines import TD3, CEMRL, PPO, SAC def test_td3(): model = TD3('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[64, 64]), start_timesteps=100, verbose=1, create_eval_env=True) - model.learn(total_timesteps=20000, eval_freq=1000) + model.learn(total_timesteps=1000, eval_freq=500) model.save("test_save") model.load("test_save") os.remove("test_save.pth") @@ -15,7 +15,7 @@ def test_td3(): def test_cemrl(): model = CEMRL('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[16]), pop_size=2, n_grad=1, start_timesteps=100, verbose=1, create_eval_env=True) - model.learn(total_timesteps=20000, eval_freq=1000) + model.learn(total_timesteps=1000, eval_freq=500) model.save("test_save") model.load("test_save") os.remove("test_save.pth") @@ -27,3 +27,8 @@ def test_ppo(): # model.save("test_save") # model.load("test_save") # os.remove("test_save.pth") + +def test_sac(): + model = SAC('MlpPolicy', 'Pendulum-v0', policy_kwargs=dict(net_arch=[64, 64]), + start_timesteps=100, verbose=1, create_eval_env=True, ent_coef='auto') + model.learn(total_timesteps=1000, eval_freq=500) diff --git a/torchy_baselines/__init__.py b/torchy_baselines/__init__.py index e95418f..b9dabaa 100644 --- a/torchy_baselines/__init__.py +++ b/torchy_baselines/__init__.py @@ -1,5 +1,6 @@ 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.2" +__version__ = "0.0.4" diff --git a/torchy_baselines/common/distributions.py b/torchy_baselines/common/distributions.py index e188667..ec19087 100644 --- a/torchy_baselines/common/distributions.py +++ b/torchy_baselines/common/distributions.py @@ -97,7 +97,13 @@ class SquashedDiagGaussianDistribution(DiagGaussianDistribution): return th.tanh(self.distribution.mean) def sample(self): - return th.tanh(self.distribution.rsample()) + self.gaussian_action = self.distribution.rsample() + return th.tanh(self.gaussian_action) + + def log_prob_from_params(self, mean_actions, log_std): + action, _ = self.proba_distribution(mean_actions, log_std) + log_prob = self.log_prob(action, self.gaussian_action) + return action, log_prob def log_prob(self, action, gaussian_action=None): # Inverse tanh diff --git a/torchy_baselines/sac/__init__.py b/torchy_baselines/sac/__init__.py new file mode 100644 index 0000000..6f70061 --- /dev/null +++ b/torchy_baselines/sac/__init__.py @@ -0,0 +1 @@ +from torchy_baselines.sac.sac import SAC diff --git a/torchy_baselines/sac/policies.py b/torchy_baselines/sac/policies.py new file mode 100644 index 0000000..c89004e --- /dev/null +++ b/torchy_baselines/sac/policies.py @@ -0,0 +1,111 @@ +import torch as th +import torch.nn as nn + +from torchy_baselines.common.policies import BasePolicy, register_policy, create_mlp, BaseNetwork +from torchy_baselines.common.distributions import SquashedDiagGaussianDistribution + +# CAP the standard deviation of the actor +LOG_STD_MAX = 2 +LOG_STD_MIN = -20 + + +class Actor(BaseNetwork): + def __init__(self, state_dim, action_dim, net_arch=None, activation_fn=nn.ReLU): + super(Actor, self).__init__() + + if net_arch is None: + net_arch = [256, 256] + + # TODO: orthogonal initialization? + actor_net = create_mlp(state_dim, -1, net_arch, activation_fn) + self.actor_net = nn.Sequential(*actor_net) + + self.action_dist = SquashedDiagGaussianDistribution(action_dim) + self.mu = nn.Linear(net_arch[-1], action_dim) + self.log_std = nn.Linear(net_arch[-1], action_dim) + + def get_action_dist_params(self, state): + latent = self.actor_net(state) + mean_actions, log_std = self.mu(latent), self.log_std(latent) + # Original Implementation to cap the standard deviation + log_std = th.clamp(log_std, LOG_STD_MIN, LOG_STD_MAX) + return mean_actions, log_std + + def forward(self, state, deterministic=False): + mean_actions, log_std = self.get_action_dist_params(state) + # Note the action is squashed + action, _ = self.action_dist.proba_distribution(mean_actions, log_std, deterministic=deterministic) + return action + + def action_log_prob(self, state): + mean_actions, log_std = self.get_action_dist_params(state) + action, log_prob = self.action_dist.log_prob_from_params(mean_actions, log_std) + return action, log_prob + + +class Critic(BaseNetwork): + def __init__(self, state_dim, action_dim, + net_arch=None, activation_fn=nn.ReLU): + super(Critic, self).__init__() + + if net_arch is None: + net_arch = [256, 256] + + q1_net = create_mlp(state_dim + action_dim, 1, net_arch, activation_fn) + self.q1_net = nn.Sequential(*q1_net) + + q2_net = create_mlp(state_dim + action_dim, 1, net_arch, activation_fn) + self.q2_net = nn.Sequential(*q2_net) + + self.q_networks = [self.q1_net, self.q2_net] + + def forward(self, obs, action): + qvalue_input = th.cat([obs, action], dim=1) + return [q_net(qvalue_input) for q_net in self.q_networks] + + def q1_forward(self, obs, action): + return self.q_networks[0](th.cat([obs, action], dim=1)) + + +class SACPolicy(BasePolicy): + def __init__(self, observation_space, action_space, + learning_rate=1e-3, net_arch=None, device='cpu', + activation_fn=nn.ReLU): + super(SACPolicy, self).__init__(observation_space, action_space, device) + self.state_dim = self.observation_space.shape[0] + self.action_dim = self.action_space.shape[0] + self.net_arch = net_arch + self.activation_fn = activation_fn + self.net_args = { + 'state_dim': self.state_dim, + 'action_dim': self.action_dim, + 'net_arch': self.net_arch, + 'activation_fn': self.activation_fn + } + self.actor, self.actor_target = None, None + self.critic, self.critic_target = None, None + + self._build(learning_rate) + + def _build(self, learning_rate): + self.actor = self.make_actor() + self.actor.optimizer = th.optim.Adam(self.actor.parameters(), lr=learning_rate) + + self.critic = self.make_critic() + self.critic_target = self.make_critic() + self.critic_target.load_state_dict(self.critic.state_dict()) + self.critic.optimizer = th.optim.Adam(self.critic.parameters(), lr=learning_rate) + + def actor_forward(self, state, deterministic=False): + pass + + def make_actor(self): + return Actor(**self.net_args).to(self.device) + + def make_critic(self): + return Critic(**self.net_args).to(self.device) + + +MlpPolicy = SACPolicy + +register_policy("MlpPolicy", MlpPolicy) diff --git a/torchy_baselines/sac/sac.py b/torchy_baselines/sac/sac.py new file mode 100644 index 0000000..01a9d09 --- /dev/null +++ b/torchy_baselines/sac/sac.py @@ -0,0 +1,235 @@ +import time + +import torch as th +import torch.nn.functional as F +import numpy as np + +from torchy_baselines.common.base_class import BaseRLModel +from torchy_baselines.common.buffers import ReplayBuffer +from torchy_baselines.common.evaluation import evaluate_policy +from torchy_baselines.sac.policies import SACPolicy + + +class SAC(BaseRLModel): + """ + Implementation of Soft Actor-Critic (SAC) + Off-Policy Maximum Entropy Deep Reinforcement Learning with a Stochastic Actor, + Paper: https://arxiv.org/abs/1801.01290 + Code: This implementation borrows code from original implementation (https://github.com/haarnoja/sac) + from OpenAI Spinning Up (https://github.com/openai/spinningup) and from the Softlearning repo + (https://github.com/rail-berkeley/softlearning/) + + Note: we use double q target and not value target as discussed + in https://github.com/hill-a/stable-baselines/issues/270 + """ + + def __init__(self, policy, env, policy_kwargs=None, verbose=0, + buffer_size=int(1e6), learning_rate=3e-4, seed=0, device='auto', + ent_coef='auto', target_entropy='auto', gamma=0.99, + action_noise_std=0.0, start_timesteps=100, + batch_size=64, create_eval_env=False, + _init_setup_model=True): + + super(SAC, self).__init__(policy, env, SACPolicy, policy_kwargs, verbose, device, + create_eval_env=create_eval_env) + + self.max_action = np.abs(self.action_space.high) + self.action_noise_std = action_noise_std + self.learning_rate = learning_rate + self.buffer_size = buffer_size + self.start_timesteps = start_timesteps + self._seed = seed + self.batch_size = batch_size + + self.ent_coef = ent_coef + self.target_entropy = target_entropy + self.log_ent_coef = None + # self.target_update_interval = target_update_interval + # self.gradient_steps = gradient_steps + self.gamma = gamma + + if _init_setup_model: + self._setup_model() + + def _setup_model(self): + state_dim, action_dim = self.observation_space.shape[0], self.action_space.shape[0] + self.seed(self._seed) + + # Target entropy is used when learning the entropy coefficient + if self.target_entropy == 'auto': + # automatically set target entropy if needed + self.target_entropy = -np.prod(self.env.action_space.shape).astype(np.float32) + else: + # Force conversion + # this will also throw an error for unexpected string + self.target_entropy = float(self.target_entropy) + + # The entropy coefficient or entropy can be learned automatically + # see Automating Entropy Adjustment for Maximum Entropy RL section + # of https://arxiv.org/abs/1812.05905 + if isinstance(self.ent_coef, str) and self.ent_coef.startswith('auto'): + # Default initial value of ent_coef when learned + init_value = 1.0 + if '_' in self.ent_coef: + init_value = float(self.ent_coef.split('_')[1]) + assert init_value > 0., "The initial value of ent_coef must be greater than 0" + + # Note: we optimize the log of the entropy coeff which is slightly different from the paper + # as discussed in https://github.com/rail-berkeley/softlearning/issues/37 + self.log_ent_coef = th.log(th.ones(1, device=self.device) * init_value).requires_grad_(True) + # Important: detach the variable from the graph + # so we don't change it with other losses + # see https://github.com/rail-berkeley/softlearning/issues/60 + self.ent_coef = th.exp(self.log_ent_coef.detach()) + self.ent_coef_optimizer = th.optim.Adam([self.log_ent_coef], lr=self.learning_rate) + else: + # Force conversion to float + # this will throw an error if a malformed string (different from 'auto') + # is passed + self.ent_coef = float(self.ent_coef) + + self.replay_buffer = ReplayBuffer(self.buffer_size, state_dim, action_dim, self.device) + self.policy = self.policy(self.observation_space, self.action_space, + self.learning_rate, device=self.device, **self.policy_kwargs) + self.policy = self.policy.to(self.device) + self._create_aliases() + + def _create_aliases(self): + self.actor = self.policy.actor + self.critic = self.policy.critic + self.critic_target = self.policy.critic_target + + def select_action(self, observation): + # Normally not needed + observation = np.array(observation) + with th.no_grad(): + observation = th.FloatTensor(observation.reshape(1, -1)).to(self.device) + return self.actor(observation).cpu().data.numpy() + + def predict(self, observation, state=None, mask=None, deterministic=True): + """ + Get the model's action from an observation + + :param observation: (np.ndarray) the input observation + :param state: (np.ndarray) The last states (can be None, used in recurrent policies) + :param mask: (np.ndarray) The last masks (can be None, used in recurrent policies) + :param deterministic: (bool) Whether or not to return deterministic actions. + :return: (np.ndarray, np.ndarray) the model's action and the next state (used in recurrent policies) + """ + return self.max_action * self.select_action(observation) + + def train(self, n_iterations, batch_size=64, tau=0.005): + + for it in range(n_iterations): + + # Sample replay buffer + replay_data = self.replay_buffer.sample(batch_size) + + state, action_batch, next_state, done, reward = replay_data + + # Action by the current actor for the sampled state + action_pi, log_prob = self.actor.action_log_prob(state) + log_prob = log_prob.reshape(-1, 1) + + ent_coef_loss = None + if not isinstance(self.ent_coef, float): + ent_coef_loss = -(self.log_ent_coef * (log_prob + self.target_entropy).detach()).mean() + + # Optimize entropy coefficient, also called + # entropy temperature or alpha in the paper + if ent_coef_loss is not None: + self.ent_coef_optimizer.zero_grad() + ent_coef_loss.backward() + self.ent_coef_optimizer.step() + + # Select action according to policy + next_action, next_log_prob = self.actor.action_log_prob(next_state) + + # Compute the target Q value + target_q1, target_q2 = self.critic_target(next_state, next_action) + target_q = th.min(target_q1, target_q2) + target_q = reward + ((1 - done) * self.gamma * target_q).detach() + + # td error + entropy term + q_backup = (target_q - self.ent_coef * next_log_prob.reshape(-1, 1)).detach() + + # Get current Q estimates + # using action from the replay buffer + current_q1, current_q2 = self.critic(state, action_batch) + + # Compute critic loss + critic_loss = 0.5 * (F.mse_loss(current_q1, q_backup) + F.mse_loss(current_q2, q_backup)) + + # Optimize the critic + self.critic.optimizer.zero_grad() + critic_loss.backward() + self.critic.optimizer.step() + + # Compute actor loss + # Alternative: actor_loss = th.mean(log_prob - min_qf_pi) + actor_loss = (self.ent_coef * log_prob - self.critic.q1_forward(state, action_pi)).mean() + + # Optimize the actor + self.actor.optimizer.zero_grad() + actor_loss.backward() + self.actor.optimizer.step() + + # Update target networks + for param, target_param in zip(self.critic.parameters(), self.critic_target.parameters()): + target_param.data.copy_(tau * param.data + (1 - tau) * target_param.data) + + def learn(self, total_timesteps, callback=None, log_interval=100, + eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="TD3", reset_num_timesteps=True): + + timesteps_since_eval = 0 + episode_num = 0 + evaluations = [] + start_time = time.time() + eval_env = self._get_eval_env(eval_env) + + while self.num_timesteps < total_timesteps: + + if callback is not None: + # Only stop training if return value is False, not when it is None. + if callback(locals(), globals()) is False: + break + + episode_reward, episode_timesteps = self.collect_rollouts(self.env, n_episodes=1, + action_noise_std=self.action_noise_std, + deterministic=False, callback=None, + start_timesteps=self.start_timesteps, + num_timesteps=self.num_timesteps, + replay_buffer=self.replay_buffer) + episode_num += 1 + self.num_timesteps += episode_timesteps + timesteps_since_eval += episode_timesteps + + if self.num_timesteps > 0: + if self.verbose > 1: + print("Total T: {} Episode Num: {} Episode T: {} Reward: {}".format( + self.num_timesteps, episode_num, episode_timesteps, episode_reward)) + self.train(episode_timesteps, batch_size=self.batch_size) + + # Evaluate episode + if 0 < eval_freq <= timesteps_since_eval and eval_env is not None: + timesteps_since_eval %= eval_freq + mean_reward, _ = evaluate_policy(self, eval_env, n_eval_episodes) + evaluations.append(mean_reward) + if self.verbose > 0: + print("Eval num_timesteps={}, mean_reward={:.2f}".format(self.num_timesteps, evaluations[-1])) + print("FPS: {:.2f}".format(self.num_timesteps / (time.time() - start_time))) + + return self + + def save(self, path): + if not path.endswith('.pth'): + path += '.pth' + th.save(self.policy.state_dict(), path) + + def load(self, path, env=None, **_kwargs): + if not path.endswith('.pth'): + path += '.pth' + if env is not None: + pass + self.policy.load_state_dict(th.load(path)) + self._create_aliases() diff --git a/torchy_baselines/td3/td3.py b/torchy_baselines/td3/td3.py index 7a527e8..76aa9f1 100644 --- a/torchy_baselines/td3/td3.py +++ b/torchy_baselines/td3/td3.py @@ -115,9 +115,9 @@ class TD3(BaseRLModel): for it in range(n_iterations): # Sample replay buffer if replay_data is None: - state, action, next_state, done, reward = self.replay_buffer.sample(batch_size) + state, _, next_state, done, reward = self.replay_buffer.sample(batch_size) else: - state, action, next_state, done, reward = replay_data + state, _, next_state, done, reward = replay_data # Compute actor loss actor_loss = -self.critic.q1_forward(state, self.actor(state)).mean()