From 2a660e9a41586f930b7921efb6f1131c87161224 Mon Sep 17 00:00:00 2001 From: Antonin Raffin Date: Thu, 12 Sep 2019 15:38:15 +0200 Subject: [PATCH] Update closer to original implementation for CEMRL --- torchy_baselines/cem_rl/cem.py | 2 +- torchy_baselines/cem_rl/cem_rl.py | 9 ++++++--- torchy_baselines/td3/td3.py | 18 +++++++++++++----- 3 files changed, 20 insertions(+), 9 deletions(-) diff --git a/torchy_baselines/cem_rl/cem.py b/torchy_baselines/cem_rl/cem.py index 0ec81a2..57513c5 100644 --- a/torchy_baselines/cem_rl/cem.py +++ b/torchy_baselines/cem_rl/cem.py @@ -7,7 +7,7 @@ import numpy as np class CEM(object): """ - Cross-entropy methods. + Cross-entropy method with diagonal covariance (separable CEM) """ def __init__(self, num_params, diff --git a/torchy_baselines/cem_rl/cem_rl.py b/torchy_baselines/cem_rl/cem_rl.py index bd02d0f..f3d4336 100644 --- a/torchy_baselines/cem_rl/cem_rl.py +++ b/torchy_baselines/cem_rl/cem_rl.py @@ -80,12 +80,15 @@ class CEMRL(TD3): # In the paper: 2 * actor_steps // self.n_grad # In the original implementation: actor_steps // self.n_grad - # Difference with current implementation: + # Difference with TD3 implementation: # the target critic is updated in the train_critic() - # instead of the train_actor() + # instead of the train_actor() and no policy delay # Issue with this update style: the bigger the population, the slower the code if self.update_style == 'original': - self.train_critic(actor_steps // self.n_grad) + self.train_critic(actor_steps // self.n_grad, tau=0.005) + self.train_actor(actor_steps, tau_critic=0.0) + elif self.update_style == 'original_td3': + self.train_critic(actor_steps // self.n_grad, tau=0.0) self.train_actor(actor_steps) else: # Closer to td3: with policy delay diff --git a/torchy_baselines/td3/td3.py b/torchy_baselines/td3/td3.py index 3bb1bae..50bb83f 100644 --- a/torchy_baselines/td3/td3.py +++ b/torchy_baselines/td3/td3.py @@ -76,7 +76,7 @@ class TD3(BaseRLModel): return self.max_action * self.select_action(observation) def train_critic(self, n_iterations=1, batch_size=100, discount=0.99, - policy_noise=0.2, noise_clip=0.5, replay_data=None): + policy_noise=0.2, noise_clip=0.5, replay_data=None, tau=0.0): for it in range(n_iterations): # Sample replay buffer @@ -106,7 +106,14 @@ class TD3(BaseRLModel): critic_loss.backward() self.critic.optimizer.step() - def train_actor(self, n_iterations=1, batch_size=100, tau=0.005, replay_data=None): + # Update the frozen target models + # Note: by default, for TD3, this update is done in train_actor + # however, for CEMRL it is done here + if tau > 0: + 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 train_actor(self, n_iterations=1, batch_size=100, tau_actor=0.005, tau_critic=0.005, replay_data=None): for it in range(n_iterations): # Sample replay buffer @@ -124,11 +131,12 @@ class TD3(BaseRLModel): self.actor.optimizer.step() # Update the frozen target models - 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) + if tau_critic > 0: + for param, target_param in zip(self.critic.parameters(), self.critic_target.parameters()): + target_param.data.copy_(tau_critic * param.data + (1 - tau_critic) * target_param.data) for param, target_param in zip(self.actor.parameters(), self.actor_target.parameters()): - target_param.data.copy_(tau * param.data + (1 - tau) * target_param.data) + target_param.data.copy_(tau_actor * param.data + (1 - tau_actor) * target_param.data) def train(self, n_iterations, batch_size=100, discount=0.99, tau=0.005, policy_noise=0.2, noise_clip=0.5, policy_freq=2):