diff --git a/stable_baselines3/cmaes/cmaes.py b/stable_baselines3/cmaes/cmaes.py index 027aa80..2f20999 100644 --- a/stable_baselines3/cmaes/cmaes.py +++ b/stable_baselines3/cmaes/cmaes.py @@ -11,6 +11,7 @@ from stable_baselines3.common.base_class import BaseAlgorithm from stable_baselines3.common.policies import BasePolicy from stable_baselines3.common.type_aliases import GymEnv, MaybeCallback from stable_baselines3.common.utils import safe_mean +from stable_baselines3.sac import SAC class CMAES(BaseAlgorithm): @@ -21,7 +22,9 @@ class CMAES(BaseAlgorithm): n_steps: int = 200, n_individuals: int = -1, std_init: float = 0.5, - best_individual: Optional[np.ndarray] = None, + best_individual: Union[np.ndarray, None, str] = None, + diagonal_cov: bool = False, + pop_size: Optional[int] = None, policy_kwargs: Dict[str, Any] = None, tensorboard_log: Optional[str] = None, verbose: int = 0, @@ -49,11 +52,15 @@ class CMAES(BaseAlgorithm): ) self.policy_kwargs["device"] = self.device + if isinstance(best_individual, str): + best_individual = SAC.load(best_individual).actor.parameters_to_vector() # pytype:disable=attribute-error self.best_individual = best_individual self.best_ever = None self.std_init = std_init self.n_steps = n_steps self.es = None + self.diagonal_cov = diagonal_cov + self.pop_size = pop_size if _init_setup_model: self._setup_model() @@ -95,11 +102,17 @@ class CMAES(BaseAlgorithm): options = {"seed": self.seed} if self.env.num_envs > 1: options["popsize"] = self.env.num_envs + if self.pop_size is not None: + options["popsize"] = self.pop_size + if self.diagonal_cov: + options["CMA_diagonal"] = True self.es = cma.CMAEvolutionStrategy(self.best_individual, self.std_init, options) continue_training = True while self.num_timesteps < total_timesteps and not self.es.stop() and continue_training: candidates = self.es.ask() + # Add best + candidates.append(self.best_individual) returns = np.zeros((len(candidates),)) candidate_idx = 0 candidate_steps = 0 @@ -131,6 +144,10 @@ class CMAES(BaseAlgorithm): self._update_current_progress_remaining(self.num_timesteps, self._total_timesteps) if candidate_steps > self.n_steps: + if self.verbose > 0: + print(f"Candidate {candidate_idx + 1}, return={returns[candidate_idx]:.2f}") + # force reset + self._last_obs = self.env.reset() candidate_idx += 1 candidate_steps = 0 if candidate_idx < len(candidates): @@ -149,6 +166,8 @@ class CMAES(BaseAlgorithm): # TODO: inject best solution from time to time? self.es.tell(candidates, -1 * returns) + if self.verbose > 0: + print(f"Mean return={np.mean(returns):.2f} +/- {np.std(returns):.2f}") # TODO: load best individual when using predict self.best_ever.update(self.es.best) self.best_individual = self.best_ever.x diff --git a/stable_baselines3/cmaes/policies.py b/stable_baselines3/cmaes/policies.py index 3a5b988..09ca548 100644 --- a/stable_baselines3/cmaes/policies.py +++ b/stable_baselines3/cmaes/policies.py @@ -7,9 +7,10 @@ from torch import nn as nn from stable_baselines3.common.policies import BasePolicy, register_policy from stable_baselines3.common.preprocessing import get_action_dim from stable_baselines3.common.torch_layers import BaseFeaturesExtractor, FlattenExtractor, NatureCNN, create_mlp +from stable_baselines3.sac.policies import Actor -class Actor(BasePolicy): +class ActorCMA(BasePolicy): """ Actor network (policy) for TD3. @@ -36,7 +37,7 @@ class Actor(BasePolicy): normalize_images: bool = True, device: Union[th.device, str] = "auto", ): - super(Actor, self).__init__( + super(ActorCMA, self).__init__( observation_space, action_space, features_extractor=features_extractor, @@ -124,6 +125,7 @@ class CMAESPolicy(BasePolicy): "activation_fn": self.activation_fn, "normalize_images": normalize_images, "device": device, + "use_sde": True, } self.actor = Actor(**self.net_args).to(self.device) self.best_actor = Actor(**self.net_args).to(self.device)