Refactor evaluation

This commit is contained in:
Antonin Raffin 2020-01-27 15:53:27 +01:00
parent d514cd9126
commit a628354721
8 changed files with 125 additions and 143 deletions

View file

@ -130,8 +130,10 @@ class A2C(PPO):
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
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):
eval_env=None, eval_freq=-1, n_eval_episodes=5,
tb_log_name="A2C", eval_log_path=None, 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)
tb_log_name=tb_log_name, eval_log_path=eval_log_path,
reset_num_timesteps=reset_num_timesteps)

View file

@ -100,9 +100,10 @@ class CEMRL(TD3):
elitism=self.elitism)
def learn(self, total_timesteps, callback=None, log_interval=4,
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="CEMRL", reset_num_timesteps=True):
eval_env=None, eval_freq=-1, n_eval_episodes=5,
tb_log_name="CEMRL", eval_log_path=None, reset_num_timesteps=True):
timesteps_since_eval, episode_num, evaluations, obs, eval_env, callback = self._setup_learn(eval_env, callback)
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq, n_eval_episodes, eval_log_path)
actor_steps = 0
continue_training = True
@ -155,21 +156,6 @@ class CEMRL(TD3):
# Get the params back in the population
self.es_params[i] = self.actor.parameters_to_vector()
# Evaluate agent
if 0 < eval_freq <= timesteps_since_eval and eval_env is not None:
timesteps_since_eval %= eval_freq
self.actor.load_from_vector(self.es.mu)
sync_envs_normalization(self.env, eval_env)
mean_reward, std_reward = evaluate_policy(self, eval_env, n_eval_episodes)
evaluations.append(mean_reward)
if self.verbose > 0:
print("Eval num_timesteps={}, "
"episode_reward={:.2f} +/- {:.2f}".format(self.num_timesteps, mean_reward, std_reward))
print("FPS: {:.2f}".format(self.num_timesteps / (time.time() - self.start_time)))
actor_steps = 0
# evaluate all actors
for params in self.es_params:
@ -180,7 +166,6 @@ class CEMRL(TD3):
n_steps=-1, action_noise=self.action_noise,
deterministic=False, callback=callback,
learning_starts=self.learning_starts,
num_timesteps=self.num_timesteps,
replay_buffer=self.replay_buffer,
obs=obs, episode_num=episode_num,
log_interval=log_interval)
@ -192,8 +177,6 @@ class CEMRL(TD3):
break
episode_num += n_episodes
self.num_timesteps += episode_timesteps
timesteps_since_eval += episode_timesteps
actor_steps += episode_timesteps
self.fitnesses.append(episode_reward)
@ -202,7 +185,6 @@ class CEMRL(TD3):
self._update_current_progress(self.num_timesteps, total_timesteps)
self.es.tell(self.es_params, self.fitnesses)
timesteps_since_eval += actor_steps
callback.on_training_end()

View file

@ -2,8 +2,7 @@ import time
import os
import io
import zipfile
import typing
from typing import Union, Type, Optional, Dict, Any, List, Tuple
from typing import Union, Type, Optional, Dict, Any, List, Tuple, Callable
from abc import ABC, abstractmethod
from collections import deque
@ -14,16 +13,13 @@ import numpy as np
from torchy_baselines.common import logger
from torchy_baselines.common.policies import BasePolicy, get_policy_from_name
from torchy_baselines.common.utils import set_random_seed, get_schedule_fn, update_learning_rate
from torchy_baselines.common.vec_env import DummyVecEnv, VecEnv, unwrap_vec_normalize, sync_envs_normalization
from torchy_baselines.common.vec_env import DummyVecEnv, VecEnv, unwrap_vec_normalize
from torchy_baselines.common.monitor import Monitor
from torchy_baselines.common.evaluation import evaluate_policy
from torchy_baselines.common.save_util import data_to_json, json_to_data
from torchy_baselines.common.type_aliases import GymEnv, TensorDict, OptimizerStateDict
from torchy_baselines.common.callbacks import BaseCallback, CallbackList, ConvertCallback, EvalCallback
from torchy_baselines.common.noise import ActionNoise
if typing.TYPE_CHECKING:
from torchy_baselines.common.callbacks import BaseCallback
class BaseRLModel(ABC):
"""
@ -50,6 +46,7 @@ class BaseRLModel(ABC):
:param sde_sample_freq: Sample a new noise matrix every n steps when using SDE
Default: -1 (only sample at the beginning of the rollout)
"""
def __init__(self,
policy: Type[BasePolicy],
env: Union[GymEnv, str],
@ -133,9 +130,6 @@ class BaseRLModel(ABC):
def _get_eval_env(self, eval_env: Optional[GymEnv]) -> Optional[GymEnv]:
"""
Return the environment that will be used for evaluation.
:param eval_env:
:return:
"""
if eval_env is None:
eval_env = self.eval_env
@ -146,34 +140,10 @@ class BaseRLModel(ABC):
assert eval_env.num_envs == 1
return eval_env
# Type hint as string to avoid circular import
def _init_callback(self, callback) -> 'BaseCallback':
"""
Note: we cannot use type hint here because of circular import.
:param callback: (Union[callable, [BaseCallback], BaseCallback, None])
:return: (BaseCallback)
"""
# Avoid circular import
from torchy_baselines.common.callbacks import BaseCallback, CallbackList, ConvertCallback
# Convert a list of callbacks into a callback
if isinstance(callback, list):
callback = CallbackList(callback)
# Convert functional callback to object
if not isinstance(callback, BaseCallback):
callback = ConvertCallback(callback)
callback.init_callback(self)
return callback
def scale_action(self, action: np.ndarray) -> np.ndarray:
"""
Rescale the action from [low, high] to [-1, 1]
(no need for symmetric action space)
:param action:
:return:
"""
low, high = self.action_space.low, self.action_space.high
return 2.0 * ((action - low) / (high - low)) - 1.0
@ -182,9 +152,6 @@ class BaseRLModel(ABC):
"""
Rescale the action from [-1, 1] to [low, high]
(no need for symmetric action space)
:param scaled_action:
:return:
"""
low, high = self.action_space.low, self.action_space.high
return low + (0.5 * (scaled_action + 1.0) * (high - low))
@ -300,11 +267,13 @@ class BaseRLModel(ABC):
@abstractmethod
def learn(self, total_timesteps: int,
callback=None, log_interval: int = 100,
callback: Union[None, Callable, List[BaseCallback], BaseCallback] = None,
log_interval: int = 100,
tb_log_name: str = "run",
eval_env: Optional[GymEnv] = None,
eval_freq: int = -1,
n_eval_episodes: int = 5,
eval_log_path: Optional[str] = None,
reset_num_timesteps: bool = True):
"""
Return a trained model.
@ -318,6 +287,8 @@ class BaseRLModel(ABC):
:param eval_env: (gym.Env) Environment that will be used to evaluate the agent
:param eval_freq: (int) Evaluate the agent every `eval_freq` timesteps (this may vary a little)
:param n_eval_episodes: (int) Number of episode to evaluate the agent
:param eval_log_path: (Optional[str]) Path to a folder where the evaluations will be saved
:param reset_num_timesteps: (bool)
:return: (BaseRLModel) the trained model
"""
raise NotImplementedError()
@ -391,7 +362,8 @@ class BaseRLModel(ABC):
@staticmethod
def _load_from_file(load_path: str, load_data: bool = True) -> (Tuple[Optional[Dict[str, Any]],
Optional[TensorDict], Optional[OptimizerStateDict]]):
Optional[TensorDict],
Optional[OptimizerStateDict]]):
""" Load model data from a .zip archive
:param load_path: Where to load the model from
@ -473,14 +445,53 @@ class BaseRLModel(ABC):
if self.eval_env is not None:
self.eval_env.seed(seed)
def _setup_learn(self, eval_env: Optional[GymEnv], callback=None) -> (Tuple[int, int,
List[Any], np.ndarray, Optional[VecEnv], Any]):
def _init_callback(self,
callback: Union[None, Callable, List[BaseCallback], BaseCallback],
eval_env: Optional[VecEnv] = None,
eval_freq: int = 10000,
n_eval_episodes: int = 5,
log_path: Optional[str] = None) -> BaseCallback:
"""
:param callback: (Union[callable, [BaseCallback], BaseCallback, None])
:return: (BaseCallback)
"""
# Convert a list of callbacks into a callback
if isinstance(callback, list):
callback = CallbackList(callback)
# Convert functional callback to object
if not isinstance(callback, BaseCallback):
callback = ConvertCallback(callback)
# Create eval callback in charge of the evaluation
if eval_env is not None:
# Same folder as the rest
best_model_save_path = os.path.dirname(log_path) if log_path is not None else None
eval_callback = EvalCallback(eval_env,
best_model_save_path=best_model_save_path,
log_path=log_path, eval_freq=eval_freq, n_eval_episodes=n_eval_episodes)
callback = CallbackList([callback, eval_callback])
callback.init_callback(self)
return callback
def _setup_learn(self,
eval_env: Optional[GymEnv],
callback: Union[None, Callable, List[BaseCallback], BaseCallback] = None,
eval_freq: int = 10000,
n_eval_episodes: int = 5,
log_path: Optional[str] = None
) -> Tuple[int, np.ndarray, BaseCallback]:
"""
Initialize different variables needed for training.
:param eval_env: (Optional[GymEnv])
:param callback: (Union[None, BaseCallback, List[BaseCallback, Callable]])
:return: (int, int, [float], np.ndarray, VecEnv, BaseCallback)
:param eval_freq: (int)
:param n_eval_episodes: (int)
:param log_path (Optional[str]):
:return: (Tuple[int, np.ndarray, BaseCallback])
"""
self.start_time = time.time()
self.ep_info_buffer = deque(maxlen=100)
@ -488,18 +499,18 @@ class BaseRLModel(ABC):
if self.action_noise is not None:
self.action_noise.reset()
callback = self._init_callback(callback)
timesteps_since_eval, episode_num = 0, 0
evaluations = []
if eval_env is not None and self.seed is not None:
eval_env.seed(self.seed)
eval_env = self._get_eval_env(eval_env)
obs = self.env.reset() # type: GymEnv
obs = self.env.reset()
return timesteps_since_eval, episode_num, evaluations, obs, eval_env, callback
# Create eval callback if needed
callback = self._init_callback(callback, eval_env, eval_freq, n_eval_episodes, log_path)
return episode_num, obs, callback
def _update_info_buffer(self, infos: List[Dict[str, Any]]) -> None:
"""
@ -521,7 +532,6 @@ class BaseRLModel(ABC):
action_noise: Optional[ActionNoise] = None,
deterministic: bool = False,
learning_starts: int = 0,
num_timesteps: int = 0,
replay_buffer=None,
obs: Optional[np.ndarray] = None,
episode_num: int = 0,
@ -537,7 +547,6 @@ class BaseRLModel(ABC):
:param deterministic: (bool)
:param callback: (BaseCallback)
:param learning_starts: (int)
:param num_timesteps: (int)
:param replay_buffer: (ReplayBuffer)
:param obs: (np.ndarray)
:param episode_num: (int)
@ -583,7 +592,7 @@ class BaseRLModel(ABC):
# Select action randomly or according to policy
# TODO: use action from policy when using SDE during the warmup phase?
# if num_timesteps < learning_starts and not self.use_sde:
if num_timesteps < learning_starts:
if self.num_timesteps < learning_starts:
# Warmup phase
unscaled_action = np.array([self.action_space.sample()])
else:
@ -642,7 +651,7 @@ class BaseRLModel(ABC):
if self._vec_normalize_env is not None:
obs_ = new_obs_
num_timesteps += 1
self.num_timesteps += 1
episode_timesteps += 1
total_steps += 1
if 0 < n_steps <= total_steps:
@ -659,7 +668,7 @@ class BaseRLModel(ABC):
# Display training infos
if self.verbose >= 1 and log_interval is not None and (
episode_num + total_episodes) % log_interval == 0:
fps = int(num_timesteps / (time.time() - self.start_time))
fps = int(self.num_timesteps / (time.time() - self.start_time))
logger.logkv("episodes", episode_num + total_episodes)
if len(self.ep_info_buffer) > 0 and len(self.ep_info_buffer[0]) > 0:
logger.logkv('ep_rew_mean', self.safe_mean([ep_info['r'] for ep_info in self.ep_info_buffer]))
@ -667,7 +676,7 @@ class BaseRLModel(ABC):
# logger.logkv("n_updates", n_updates)
logger.logkv("fps", fps)
logger.logkv('time_elapsed', int(time.time() - self.start_time))
logger.logkv("total timesteps", num_timesteps)
logger.logkv("total timesteps", self.num_timesteps)
if self.use_sde:
logger.logkv("std", (self.actor.get_std()).mean().item())
logger.dumpkvs()
@ -776,27 +785,3 @@ class BaseRLModel(ABC):
params_to_save = self.get_policy_parameters()
opt_params_to_save = self.get_opt_parameters()
self._save_to_file_zip(path, data=data, params=params_to_save, opt_params=opt_params_to_save)
def _eval_policy(self, eval_freq: int, eval_env: GymEnv, n_eval_episodes: int,
timesteps_since_eval: int, render: bool = False, deterministic: bool = True) -> int:
"""
Evaluate the current policy on a test environment.
:param eval_freq: Evaluate the agent every `eval_freq` timesteps (this may vary a little)
:param n_eval_episodes: Number of episode to evaluate the agent
:parma timesteps_since_eval: Number of timesteps since last evaluation
:param deterministic: Whether to use deterministic or stochastic actions
:param render: Whether to render the eval env or not
:return: Number of timesteps since last evaluation
"""
if 0 < eval_freq <= timesteps_since_eval and eval_env is not None:
timesteps_since_eval %= eval_freq
# Synchronise the normalization stats if needed
sync_envs_normalization(self.env, eval_env)
mean_reward, std_reward = evaluate_policy(self, eval_env, n_eval_episodes,
render=render, deterministic=deterministic)
if self.verbose > 0:
print(f"Eval num_timesteps={self.num_timesteps}, "
f"episode_reward={mean_reward:.2f} +/- {std_reward:.2f}")
print(f"FPS: {self.num_timesteps / (time.time() - self.start_time):.2f}")
return timesteps_since_eval

View file

@ -1,15 +1,18 @@
import os
from abc import ABC, abstractmethod
import typing
from typing import Union, List, Dict, Any, Optional
import gym
import numpy as np
from torchy_baselines.common.base_class import BaseRLModel # pytype: disable=pyi-error
from torchy_baselines.common.vec_env import VecEnv, sync_envs_normalization
from torchy_baselines.common.evaluation import evaluate_policy
from torchy_baselines.common.logger import Logger
if typing.TYPE_CHECKING:
from torchy_baselines.common.base_class import BaseRLModel # pytype: disable=pyi-error
class BaseCallback(ABC):
"""
@ -31,7 +34,8 @@ class BaseCallback(ABC):
# to have access to the parent object
self.parent = None # type: Optional[BaseCallback]
def init_callback(self, model: BaseRLModel) -> None:
# Type hint as string to avoid circular import
def init_callback(self, model: 'BaseRLModel') -> None:
"""
Initialize the callback by saving references to the
RL model and the training environment for convenience.
@ -105,11 +109,13 @@ class EventCallback(BaseCallback):
if callback is not None:
self.callback.parent = self
def init_callback(self, model: BaseRLModel) -> None:
def init_callback(self, model: 'BaseRLModel') -> None:
super(EventCallback, self).init_callback(model)
if self.callback is not None:
self.callback.init_callback(self.model)
def _on_training_start(self) -> None:
if self.callback is not None:
self.callback.on_training_start(self.locals, self.globals)
def _on_event(self) -> bool:
@ -117,6 +123,9 @@ class EventCallback(BaseCallback):
return self.callback()
return True
def _on_step(self) -> bool:
return True
class CallbackList(BaseCallback):
def __init__(self, callbacks: List[BaseCallback]):
@ -179,7 +188,7 @@ class ConvertCallback(BaseCallback):
"""
Convert functional callback (old-style) to object.
:param on_step: (callable)
:param callback: (callable)
:param verbose: (int)
"""
def __init__(self, callback, verbose=0):
@ -207,6 +216,7 @@ class EvalCallback(EventCallback):
according to performance on the eval env will be saved.
:param deterministic: (bool) Whether the evaluation should
use a stochastic or deterministic actions.
:param deterministic: (bool) Whether to render or not the environment during evaluation
:param verbose: (int)
"""
def __init__(self, eval_env: Union[gym.Env, VecEnv],
@ -216,12 +226,15 @@ class EvalCallback(EventCallback):
log_path: str = None,
best_model_save_path: str = None,
deterministic: bool = True,
render: bool = False,
verbose: int = 1):
super(EvalCallback, self).__init__(callback_on_new_best, verbose=verbose)
self.n_eval_episodes = n_eval_episodes
self.eval_freq = eval_freq
self.best_mean_reward = -np.inf
self.deterministic = deterministic
self.render = render
if isinstance(eval_env, VecEnv):
assert eval_env.num_envs == 1, "You must pass only one environment for evaluation"
@ -230,6 +243,7 @@ class EvalCallback(EventCallback):
self.log_path = log_path
self.evaluations_results = []
self.evaluations_timesteps = []
self.evaluations_length = []
def _init_callback(self):
# Does not work when eval_env is a gym.Env and training_env is a VecEnv
@ -244,22 +258,30 @@ class EvalCallback(EventCallback):
def _on_step(self) -> bool:
if self.n_calls % self.eval_freq == 0:
if self.eval_freq > 0 and self.n_calls % self.eval_freq == 0:
# Sync training and eval env if there is VecNormalize
sync_envs_normalization(self.training_env, self.eval_env)
episode_rewards, _ = evaluate_policy(self.model, self.eval_env, n_eval_episodes=self.n_eval_episodes,
deterministic=self.deterministic, return_episode_rewards=True)
episode_rewards, episode_lengths = evaluate_policy(self.model, self.eval_env,
n_eval_episodes=self.n_eval_episodes,
render=self.render,
deterministic=self.deterministic,
return_episode_rewards=True)
if self.log_path is not None:
self.evaluations_timesteps.append(self.num_timesteps)
self.evaluations_results.append(episode_rewards)
np.savez(self.log_path, timesteps=self.evaluations_timesteps, results=self.evaluations_results)
self.evaluations_length.append(episode_lengths)
np.savez(self.log_path, timesteps=self.evaluations_timesteps,
results=self.evaluations_results, ep_lengths=self.evaluations_length)
mean_reward, std_reward = np.mean(episode_rewards), np.std(episode_rewards)
mean_ep_length, std_ep_length = np.mean(episode_lengths), np.std(episode_lengths)
if self.verbose > 0:
print(f"Eval num_timesteps={self.num_timesteps}, "
f"episode_reward={mean_reward:.2f} +/- {std_reward:.2f}")
print(f"Episode length: {mean_ep_length:.2f} +/- {std_ep_length:.2f}")
if mean_reward > self.best_mean_reward:
if self.verbose > 0:

View file

@ -24,31 +24,33 @@ def evaluate_policy(model, env, n_eval_episodes=10, deterministic=True,
:param return_episode_rewards: (bool) If True, a list of reward per episode
will be returned instead of the mean.
:return: (float, float) Mean reward per episode, std of reward per episode
returns ([float], int) when `return_episode_rewards` is True
returns ([float], [int]) when `return_episode_rewards` is True
"""
if isinstance(env, VecEnv):
assert env.num_envs == 1, "You must pass only one environment when using this function"
episode_rewards, n_steps = [], 0
episode_rewards, episode_lengths = [], []
for _ in range(n_eval_episodes):
obs = env.reset()
done = False
episode_reward = 0.0
episode_length = 0
while not done:
action = model.predict(obs, deterministic=deterministic)
obs, reward, done, _info = env.step(action)
episode_reward += reward
if callback is not None:
callback(locals(), globals())
n_steps += 1
episode_length += 1
if render:
env.render()
episode_rewards.append(episode_reward)
episode_lengths.append(episode_length)
mean_reward = np.mean(episode_rewards)
std_reward = np.std(episode_rewards)
if reward_threshold is not None:
assert mean_reward > reward_threshold, (f'Mean reward below threshold: '
'{mean_reward:.2f} < {reward_threshold:.2f}')
if return_episode_rewards:
return episode_rewards, n_steps
return episode_rewards, episode_lengths
return mean_reward, std_reward

View file

@ -193,6 +193,8 @@ class PPO(BaseRLModel):
self._update_info_buffer(infos)
n_steps += 1
self.num_timesteps += env.num_envs
if isinstance(self.action_space, gym.spaces.Discrete):
# Reshape in case of discrete action
actions = actions.reshape(-1, 1)
@ -284,9 +286,11 @@ class PPO(BaseRLModel):
logger.logkv("std", th.exp(self.policy.log_std).mean().item())
def learn(self, total_timesteps, callback=None, log_interval=1,
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="PPO", reset_num_timesteps=True):
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="PPO",
eval_log_path=None, reset_num_timesteps=True):
timesteps_since_eval, iteration, evaluations, obs, eval_env, callback = self._setup_learn(eval_env, callback)
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq, n_eval_episodes, eval_log_path)
iteration = 0
if self.tensorboard_log is not None and SummaryWriter is not None:
self.tb_writer = SummaryWriter(log_dir=os.path.join(self.tensorboard_log, tb_log_name))
@ -295,15 +299,15 @@ class PPO(BaseRLModel):
while self.num_timesteps < total_timesteps:
obs, continue_training = self.collect_rollouts(self.env, callback, self.rollout_buffer, n_rollout_steps=self.n_steps,
obs, continue_training = self.collect_rollouts(self.env, callback,
self.rollout_buffer,
n_rollout_steps=self.n_steps,
obs=obs)
if continue_training is False:
break
iteration += 1
self.num_timesteps += self.n_steps * self.n_envs
timesteps_since_eval += self.n_steps * self.n_envs
self._update_current_progress(self.num_timesteps, total_timesteps)
# Display training infos
@ -320,9 +324,6 @@ class PPO(BaseRLModel):
self.train(self.n_epochs, batch_size=self.batch_size)
# Evaluate the agent
timesteps_since_eval = self._eval_policy(eval_freq, eval_env, n_eval_episodes,
timesteps_since_eval, deterministic=True)
# For tensorboard integration
# if self.tb_writer is not None:
# self.tb_writer.add_scalar('Eval/reward', mean_reward, self.num_timesteps)

View file

@ -257,9 +257,9 @@ class SAC(BaseRLModel):
def learn(self, total_timesteps, callback=None, log_interval=4,
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="SAC",
reset_num_timesteps=True):
eval_log_path=None, reset_num_timesteps=True):
timesteps_since_eval, episode_num, evaluations, obs, eval_env, callback = self._setup_learn(eval_env, callback)
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq, n_eval_episodes, eval_log_path)
callback.on_training_start(locals(), globals())
@ -268,7 +268,6 @@ class SAC(BaseRLModel):
n_steps=self.train_freq, action_noise=self.action_noise,
deterministic=False, callback=callback,
learning_starts=self.learning_starts,
num_timesteps=self.num_timesteps,
replay_buffer=self.replay_buffer,
obs=obs, episode_num=episode_num,
log_interval=log_interval)
@ -278,9 +277,7 @@ class SAC(BaseRLModel):
if continue_training is False:
break
self.num_timesteps += episode_timesteps
episode_num += n_episodes
timesteps_since_eval += episode_timesteps
self._update_current_progress(self.num_timesteps, total_timesteps)
if self.num_timesteps > 0 and self.num_timesteps > self.learning_starts:
@ -288,9 +285,6 @@ class SAC(BaseRLModel):
self.train(gradient_steps, batch_size=self.batch_size)
timesteps_since_eval = self._eval_policy(eval_freq, eval_env, n_eval_episodes,
timesteps_since_eval, deterministic=True)
callback.on_training_end()
return self

View file

@ -249,9 +249,10 @@ class TD3(BaseRLModel):
del self.rollout_data
def learn(self, total_timesteps, callback=None, log_interval=4,
eval_env=None, eval_freq=-1, n_eval_episodes=5, tb_log_name="TD3", reset_num_timesteps=True):
eval_env=None, eval_freq=-1, n_eval_episodes=5,
tb_log_name="TD3", eval_log_path=None, reset_num_timesteps=True):
timesteps_since_eval, episode_num, evaluations, obs, eval_env, callback = self._setup_learn(eval_env, callback)
episode_num, obs, callback = self._setup_learn(eval_env, callback, eval_freq, n_eval_episodes, eval_log_path)
callback.on_training_start(locals(), globals())
@ -261,7 +262,6 @@ class TD3(BaseRLModel):
n_steps=self.train_freq, action_noise=self.action_noise,
deterministic=False, callback=callback,
learning_starts=self.learning_starts,
num_timesteps=self.num_timesteps,
replay_buffer=self.replay_buffer,
obs=obs, episode_num=episode_num,
log_interval=log_interval)
@ -272,8 +272,6 @@ class TD3(BaseRLModel):
break
episode_num += n_episodes
self.num_timesteps += episode_timesteps
timesteps_since_eval += episode_timesteps
self._update_current_progress(self.num_timesteps, total_timesteps)
if self.num_timesteps > 0 and self.num_timesteps > self.learning_starts:
@ -290,10 +288,6 @@ class TD3(BaseRLModel):
gradient_steps = self.gradient_steps if self.gradient_steps > 0 else episode_timesteps
self.train(gradient_steps, batch_size=self.batch_size, policy_delay=self.policy_delay)
# Evaluate the agent
timesteps_since_eval = self._eval_policy(eval_freq, eval_env, n_eval_episodes,
timesteps_since_eval, deterministic=True)
callback.on_training_end()
return self