mirror of
https://github.com/saymrwulf/stable-baselines3.git
synced 2026-09-17 22:30:59 +00:00
Merge pull request #65 from Antonin-Raffin/feat/policy-save-load
Policy save/load - Action dist refactor
This commit is contained in:
commit
cf840ed928
12 changed files with 430 additions and 288 deletions
|
|
@ -10,10 +10,12 @@ Pre-Release 0.4.0a0 (WIP)
|
||||||
Breaking Changes:
|
Breaking Changes:
|
||||||
^^^^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^^^^
|
||||||
- Removed CEMRL
|
- Removed CEMRL
|
||||||
|
- Model saved with previous versions cannot be loaded (because of the pre-preprocessing)
|
||||||
|
|
||||||
New Features:
|
New Features:
|
||||||
^^^^^^^^^^^^^
|
^^^^^^^^^^^^^
|
||||||
- Add support for Discrete observation spaces
|
- Add support for Discrete observation spaces
|
||||||
|
- Add saving/loading for policy weights, so the policy can be used without the model
|
||||||
|
|
||||||
Bug Fixes:
|
Bug Fixes:
|
||||||
^^^^^^^^^^
|
^^^^^^^^^^
|
||||||
|
|
@ -26,6 +28,8 @@ Others:
|
||||||
^^^^^^^
|
^^^^^^^
|
||||||
- Refactor handling of observation and action spaces
|
- Refactor handling of observation and action spaces
|
||||||
- Refactored features extraction to have proper preprocessing
|
- Refactored features extraction to have proper preprocessing
|
||||||
|
- Refactored action distributions
|
||||||
|
|
||||||
|
|
||||||
Documentation:
|
Documentation:
|
||||||
^^^^^^^^^^^^^^
|
^^^^^^^^^^^^^^
|
||||||
|
|
|
||||||
|
|
@ -1,2 +1,2 @@
|
||||||
#!/bin/bash
|
#!/bin/bash
|
||||||
python -m pytest --cov-config .coveragerc --cov-report html --cov-report term --cov=. -v
|
python3 -m pytest --cov-config .coveragerc --cov-report html --cov-report term --cov=. -v
|
||||||
|
|
|
||||||
|
|
@ -38,7 +38,8 @@ def test_squashed_gaussian(model_class):
|
||||||
gaussian_mean = th.rand(N_SAMPLES, N_ACTIONS)
|
gaussian_mean = th.rand(N_SAMPLES, N_ACTIONS)
|
||||||
dist = SquashedDiagGaussianDistribution(N_ACTIONS)
|
dist = SquashedDiagGaussianDistribution(N_ACTIONS)
|
||||||
_, log_std = dist.proba_distribution_net(N_FEATURES)
|
_, log_std = dist.proba_distribution_net(N_FEATURES)
|
||||||
actions, _ = dist.proba_distribution(gaussian_mean, log_std)
|
dist = dist.proba_distribution(gaussian_mean, log_std)
|
||||||
|
actions = dist.get_actions()
|
||||||
assert th.max(th.abs(actions)) <= 1.0
|
assert th.max(th.abs(actions)) <= 1.0
|
||||||
|
|
||||||
def test_sde_distribution():
|
def test_sde_distribution():
|
||||||
|
|
@ -51,7 +52,8 @@ def test_sde_distribution():
|
||||||
_, log_std = dist.proba_distribution_net(N_FEATURES)
|
_, log_std = dist.proba_distribution_net(N_FEATURES)
|
||||||
dist.sample_weights(log_std, batch_size=N_SAMPLES)
|
dist.sample_weights(log_std, batch_size=N_SAMPLES)
|
||||||
|
|
||||||
actions, _ = dist.proba_distribution(deterministic_actions, log_std, state)
|
dist = dist.proba_distribution(deterministic_actions, log_std, state)
|
||||||
|
actions = dist.get_actions()
|
||||||
|
|
||||||
assert th.allclose(actions.mean(), dist.distribution.mean.mean(), rtol=1e-3)
|
assert th.allclose(actions.mean(), dist.distribution.mean.mean(), rtol=1e-3)
|
||||||
assert th.allclose(actions.std(), dist.distribution.scale.mean(), rtol=1e-3)
|
assert th.allclose(actions.std(), dist.distribution.scale.mean(), rtol=1e-3)
|
||||||
|
|
@ -71,11 +73,12 @@ def test_entropy(dist):
|
||||||
_, log_std = dist.proba_distribution_net(N_FEATURES, log_std_init=th.log(th.tensor(0.2)))
|
_, log_std = dist.proba_distribution_net(N_FEATURES, log_std_init=th.log(th.tensor(0.2)))
|
||||||
|
|
||||||
if isinstance(dist, DiagGaussianDistribution):
|
if isinstance(dist, DiagGaussianDistribution):
|
||||||
actions, dist = dist.proba_distribution(deterministic_actions, log_std)
|
dist = dist.proba_distribution(deterministic_actions, log_std)
|
||||||
else:
|
else:
|
||||||
dist.sample_weights(log_std, batch_size=N_SAMPLES)
|
dist.sample_weights(log_std, batch_size=N_SAMPLES)
|
||||||
actions, dist = dist.proba_distribution(deterministic_actions, log_std, state)
|
dist = dist.proba_distribution(deterministic_actions, log_std, state)
|
||||||
|
|
||||||
|
actions = dist.get_actions()
|
||||||
entropy = dist.entropy()
|
entropy = dist.entropy()
|
||||||
log_prob = dist.log_prob(actions)
|
log_prob = dist.log_prob(actions)
|
||||||
assert th.allclose(entropy.mean(), -log_prob.mean(), rtol=5e-3)
|
assert th.allclose(entropy.mean(), -log_prob.mean(), rtol=5e-3)
|
||||||
|
|
@ -88,8 +91,9 @@ def test_categorical():
|
||||||
set_random_seed(1)
|
set_random_seed(1)
|
||||||
state = th.rand(N_SAMPLES, N_FEATURES)
|
state = th.rand(N_SAMPLES, N_FEATURES)
|
||||||
action_logits = th.rand(N_SAMPLES, N_ACTIONS)
|
action_logits = th.rand(N_SAMPLES, N_ACTIONS)
|
||||||
actions, dist = dist.proba_distribution(action_logits)
|
dist = dist.proba_distribution(action_logits)
|
||||||
|
|
||||||
|
actions = dist.get_actions()
|
||||||
entropy = dist.entropy()
|
entropy = dist.entropy()
|
||||||
log_prob = dist.log_prob(actions)
|
log_prob = dist.log_prob(actions)
|
||||||
assert th.allclose(entropy.mean(), -log_prob.mean(), rtol=1e-4)
|
assert th.allclose(entropy.mean(), -log_prob.mean(), rtol=1e-4)
|
||||||
|
|
|
||||||
|
|
@ -20,9 +20,9 @@ def test_continuous(model_class):
|
||||||
env = IdentityEnvBox(eps=0.5)
|
env = IdentityEnvBox(eps=0.5)
|
||||||
|
|
||||||
n_steps = {
|
n_steps = {
|
||||||
A2C: 3000,
|
A2C: 3500,
|
||||||
PPO: 3000,
|
PPO: 3000,
|
||||||
SAC: 500,
|
SAC: 700,
|
||||||
TD3: 500
|
TD3: 500
|
||||||
}[model_class]
|
}[model_class]
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -160,3 +160,79 @@ def test_save_load_replay_buffer(model_class):
|
||||||
|
|
||||||
# clear file from os
|
# clear file from os
|
||||||
os.remove(replay_path)
|
os.remove(replay_path)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.parametrize("model_class", MODEL_LIST)
|
||||||
|
def test_save_load_policy(model_class):
|
||||||
|
"""
|
||||||
|
Test saving and loading policy only.
|
||||||
|
|
||||||
|
:param model_class: (BaseRLModel) A RL model
|
||||||
|
"""
|
||||||
|
env = DummyVecEnv([lambda: IdentityEnvBox(10)])
|
||||||
|
|
||||||
|
# create model
|
||||||
|
model = model_class('MlpPolicy', env, policy_kwargs=dict(net_arch=[16]), verbose=1, create_eval_env=True)
|
||||||
|
model.learn(total_timesteps=500, eval_freq=250)
|
||||||
|
|
||||||
|
env.reset()
|
||||||
|
observations = np.array([env.step(env.action_space.sample())[0] for _ in range(10)])
|
||||||
|
observations = observations.reshape(10, -1)
|
||||||
|
|
||||||
|
policy = model.policy
|
||||||
|
actor = None
|
||||||
|
if model_class in [SAC, TD3]:
|
||||||
|
actor = policy.actor
|
||||||
|
|
||||||
|
# Get dictionary of current parameters
|
||||||
|
params = deepcopy(policy.state_dict())
|
||||||
|
|
||||||
|
# Modify all parameters to be random values
|
||||||
|
random_params = dict((param_name, th.rand_like(param)) for param_name, param in params.items())
|
||||||
|
|
||||||
|
# Update model parameters with the new random values
|
||||||
|
policy.load_state_dict(random_params)
|
||||||
|
|
||||||
|
new_params = policy.state_dict()
|
||||||
|
# Check that all params are different now
|
||||||
|
for k in params:
|
||||||
|
assert not th.allclose(params[k], new_params[k]), "Parameters did not change as expected."
|
||||||
|
|
||||||
|
params = new_params
|
||||||
|
|
||||||
|
# get selected actions
|
||||||
|
selected_actions, _ = policy.predict(observations, deterministic=True)
|
||||||
|
# Should also work with the actor only
|
||||||
|
if actor is not None:
|
||||||
|
selected_actions_actor, _ = actor.predict(observations, deterministic=True)
|
||||||
|
|
||||||
|
# Save and load policy
|
||||||
|
policy.save("./logs/policy_weights.pkl")
|
||||||
|
# Save and load actor
|
||||||
|
if actor is not None:
|
||||||
|
actor.save("./logs/actor_weights.pkl")
|
||||||
|
|
||||||
|
policy.load("./logs/policy_weights.pkl")
|
||||||
|
if actor is not None:
|
||||||
|
actor.load("./logs/actor_weights.pkl")
|
||||||
|
|
||||||
|
# check if params are still the same after load
|
||||||
|
new_params = policy.state_dict()
|
||||||
|
|
||||||
|
# Check that all params are the same as before save load procedure now
|
||||||
|
for key in params:
|
||||||
|
assert th.allclose(params[key], new_params[key]), "Policy parameters not the same after save and load."
|
||||||
|
|
||||||
|
# check if model still selects the same actions
|
||||||
|
new_selected_actions, _ = policy.predict(observations, deterministic=True)
|
||||||
|
assert np.allclose(selected_actions, new_selected_actions, 1e-4)
|
||||||
|
|
||||||
|
if actor is not None:
|
||||||
|
new_selected_actions_actor, _ = actor.predict(observations, deterministic=True)
|
||||||
|
assert np.allclose(selected_actions_actor, new_selected_actions_actor, 1e-4)
|
||||||
|
assert np.allclose(selected_actions_actor, new_selected_actions, 1e-4)
|
||||||
|
|
||||||
|
# clear file from os
|
||||||
|
os.remove("./logs/policy_weights.pkl")
|
||||||
|
if actor is not None:
|
||||||
|
os.remove("./logs/actor_weights.pkl")
|
||||||
|
|
|
||||||
|
|
@ -158,27 +158,6 @@ class BaseRLModel(ABC):
|
||||||
assert eval_env.num_envs == 1
|
assert eval_env.num_envs == 1
|
||||||
return eval_env
|
return eval_env
|
||||||
|
|
||||||
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: (np.ndarray) Action to scale
|
|
||||||
:return: (np.ndarray) Scaled action
|
|
||||||
"""
|
|
||||||
low, high = self.action_space.low, self.action_space.high
|
|
||||||
return 2.0 * ((action - low) / (high - low)) - 1.0
|
|
||||||
|
|
||||||
def unscale_action(self, scaled_action: np.ndarray) -> np.ndarray:
|
|
||||||
"""
|
|
||||||
Rescale the action from [-1, 1] to [low, high]
|
|
||||||
(no need for symmetric action space)
|
|
||||||
|
|
||||||
:param scaled_action: Action to un-scale
|
|
||||||
"""
|
|
||||||
low, high = self.action_space.low, self.action_space.high
|
|
||||||
return low + (0.5 * (scaled_action + 1.0) * (high - low))
|
|
||||||
|
|
||||||
def _setup_lr_schedule(self) -> None:
|
def _setup_lr_schedule(self) -> None:
|
||||||
"""Transform to callable if needed."""
|
"""Transform to callable if needed."""
|
||||||
self.lr_schedule = get_schedule_fn(self.learning_rate)
|
self.lr_schedule = get_schedule_fn(self.learning_rate)
|
||||||
|
|
@ -318,57 +297,6 @@ class BaseRLModel(ABC):
|
||||||
"""
|
"""
|
||||||
raise NotImplementedError()
|
raise NotImplementedError()
|
||||||
|
|
||||||
@staticmethod
|
|
||||||
def _is_vectorized_observation(observation: np.ndarray, observation_space: gym.spaces.Space) -> bool:
|
|
||||||
"""
|
|
||||||
For every observation type, detects and validates the shape,
|
|
||||||
then returns whether or not the observation is vectorized.
|
|
||||||
|
|
||||||
:param observation: (np.ndarray) the input observation to validate
|
|
||||||
:param observation_space: (gym.spaces) the observation space
|
|
||||||
:return: (bool) whether the given observation is vectorized or not
|
|
||||||
"""
|
|
||||||
if isinstance(observation_space, gym.spaces.Box):
|
|
||||||
if observation.shape == observation_space.shape:
|
|
||||||
return False
|
|
||||||
elif observation.shape[1:] == observation_space.shape:
|
|
||||||
return True
|
|
||||||
else:
|
|
||||||
raise ValueError("Error: Unexpected observation shape {} for ".format(observation.shape) +
|
|
||||||
"Box environment, please use {} ".format(observation_space.shape) +
|
|
||||||
"or (n_env, {}) for the observation shape."
|
|
||||||
.format(", ".join(map(str, observation_space.shape))))
|
|
||||||
elif isinstance(observation_space, gym.spaces.Discrete):
|
|
||||||
if observation.shape == (): # A numpy array of a number, has shape empty tuple '()'
|
|
||||||
return False
|
|
||||||
elif len(observation.shape) == 1:
|
|
||||||
return True
|
|
||||||
else:
|
|
||||||
raise ValueError("Error: Unexpected observation shape {} for ".format(observation.shape) +
|
|
||||||
"Discrete environment, please use (1,) or (n_env, 1) for the observation shape.")
|
|
||||||
# TODO: add support for MultiDiscrete and MultiBinary observation spaces
|
|
||||||
# elif isinstance(observation_space, gym.spaces.MultiDiscrete):
|
|
||||||
# if observation.shape == (len(observation_space.nvec),):
|
|
||||||
# return False
|
|
||||||
# elif len(observation.shape) == 2 and observation.shape[1] == len(observation_space.nvec):
|
|
||||||
# return True
|
|
||||||
# else:
|
|
||||||
# raise ValueError("Error: Unexpected observation shape {} for MultiDiscrete ".format(observation.shape) +
|
|
||||||
# "environment, please use ({},) or ".format(len(observation_space.nvec)) +
|
|
||||||
# "(n_env, {}) for the observation shape.".format(len(observation_space.nvec)))
|
|
||||||
# elif isinstance(observation_space, gym.spaces.MultiBinary):
|
|
||||||
# if observation.shape == (observation_space.n,):
|
|
||||||
# return False
|
|
||||||
# elif len(observation.shape) == 2 and observation.shape[1] == observation_space.n:
|
|
||||||
# return True
|
|
||||||
# else:
|
|
||||||
# raise ValueError("Error: Unexpected observation shape {} for MultiBinary ".format(observation.shape) +
|
|
||||||
# "environment, please use ({},) or ".format(observation_space.n) +
|
|
||||||
# "(n_env, {}) for the observation shape.".format(observation_space.n))
|
|
||||||
else:
|
|
||||||
raise ValueError("Error: Cannot determine if the observation is vectorized with the space type {}."
|
|
||||||
.format(observation_space))
|
|
||||||
|
|
||||||
def predict(self, observation: np.ndarray,
|
def predict(self, observation: np.ndarray,
|
||||||
state: Optional[np.ndarray] = None,
|
state: Optional[np.ndarray] = None,
|
||||||
mask: Optional[np.ndarray] = None,
|
mask: Optional[np.ndarray] = None,
|
||||||
|
|
@ -383,36 +311,7 @@ class BaseRLModel(ABC):
|
||||||
:return: (Tuple[np.ndarray, Optional[np.ndarray]]) the model's action and the next state
|
:return: (Tuple[np.ndarray, Optional[np.ndarray]]) the model's action and the next state
|
||||||
(used in recurrent policies)
|
(used in recurrent policies)
|
||||||
"""
|
"""
|
||||||
# TODO: move this block to BasePolicy
|
return self.policy.predict(observation, state, mask, deterministic)
|
||||||
# if state is None:
|
|
||||||
# state = self.initial_state
|
|
||||||
# if mask is None:
|
|
||||||
# mask = [False for _ in range(self.n_envs)]
|
|
||||||
observation = np.array(observation)
|
|
||||||
vectorized_env = self._is_vectorized_observation(observation, self.observation_space)
|
|
||||||
|
|
||||||
observation = observation.reshape((-1,) + self.observation_space.shape)
|
|
||||||
observation = th.as_tensor(observation).to(self.device)
|
|
||||||
with th.no_grad():
|
|
||||||
actions = self.policy.predict(observation, deterministic=deterministic)
|
|
||||||
# Convert to numpy
|
|
||||||
actions = actions.cpu().numpy()
|
|
||||||
|
|
||||||
# Rescale to proper domain when using squashing
|
|
||||||
if isinstance(self.action_space, gym.spaces.Box) and self.policy.squash_output:
|
|
||||||
actions = self.unscale_action(actions)
|
|
||||||
|
|
||||||
clipped_actions = actions
|
|
||||||
# Clip the actions to avoid out of bound error when using gaussian distribution
|
|
||||||
if isinstance(self.action_space, gym.spaces.Box) and not self.policy.squash_output:
|
|
||||||
clipped_actions = np.clip(actions, self.action_space.low, self.action_space.high)
|
|
||||||
|
|
||||||
if not vectorized_env:
|
|
||||||
if state is not None:
|
|
||||||
raise ValueError("Error: The environment must be vectorized when using recurrent policies.")
|
|
||||||
clipped_actions = clipped_actions[0]
|
|
||||||
|
|
||||||
return clipped_actions, state
|
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def load(cls, load_path: str, env: Optional[GymEnv] = None, **kwargs):
|
def load(cls, load_path: str, env: Optional[GymEnv] = None, **kwargs):
|
||||||
|
|
@ -484,10 +383,7 @@ class BaseRLModel(ABC):
|
||||||
raise ValueError(f"Error: the file {load_path} could not be found")
|
raise ValueError(f"Error: the file {load_path} could not be found")
|
||||||
|
|
||||||
# set device to cpu if cuda is not available
|
# set device to cpu if cuda is not available
|
||||||
if th.cuda.is_available():
|
device = th.device('cuda') if th.cuda.is_available() else th.device('cpu')
|
||||||
device = th.device('cuda')
|
|
||||||
else:
|
|
||||||
device = th.device('cpu')
|
|
||||||
|
|
||||||
# Open the zip archive and load data
|
# Open the zip archive and load data
|
||||||
try:
|
try:
|
||||||
|
|
@ -534,20 +430,6 @@ class BaseRLModel(ABC):
|
||||||
# load the parameters with the right `map_location`
|
# load the parameters with the right `map_location`
|
||||||
params[os.path.splitext(file_path)[0]] = th.load(file_content, map_location=device)
|
params[os.path.splitext(file_path)[0]] = th.load(file_content, map_location=device)
|
||||||
|
|
||||||
# for backward compatibility
|
|
||||||
if params.get('params') is not None:
|
|
||||||
params_copy = {}
|
|
||||||
for name in params:
|
|
||||||
if name == 'params':
|
|
||||||
params_copy['policy'] = params[name]
|
|
||||||
elif name == 'opt':
|
|
||||||
params_copy['policy.optimizer'] = params[name]
|
|
||||||
# Special case for SAC
|
|
||||||
elif name == 'ent_coef_optimizer':
|
|
||||||
params_copy[name] = params[name]
|
|
||||||
else:
|
|
||||||
params_copy[name + '.optimizer'] = params[name]
|
|
||||||
params = params_copy
|
|
||||||
except zipfile.BadZipFile:
|
except zipfile.BadZipFile:
|
||||||
# load_path wasn't a zip file
|
# load_path wasn't a zip file
|
||||||
raise ValueError(f"Error: the file {load_path} wasn't a zip-file")
|
raise ValueError(f"Error: the file {load_path} wasn't a zip-file")
|
||||||
|
|
@ -925,7 +807,7 @@ class OffPolicyRLModel(BaseRLModel):
|
||||||
unscaled_action, _ = self.predict(obs, deterministic=False)
|
unscaled_action, _ = self.predict(obs, deterministic=False)
|
||||||
|
|
||||||
# Rescale the action from [low, high] to [-1, 1]
|
# Rescale the action from [low, high] to [-1, 1]
|
||||||
scaled_action = self.scale_action(unscaled_action)
|
scaled_action = self.policy.scale_action(unscaled_action)
|
||||||
|
|
||||||
if self.use_sde:
|
if self.use_sde:
|
||||||
# When using SDE, the action can be out of bounds
|
# When using SDE, the action can be out of bounds
|
||||||
|
|
@ -941,7 +823,7 @@ class OffPolicyRLModel(BaseRLModel):
|
||||||
clipped_action = np.clip(clipped_action + action_noise(), -1, 1)
|
clipped_action = np.clip(clipped_action + action_noise(), -1, 1)
|
||||||
|
|
||||||
# Rescale and perform action
|
# Rescale and perform action
|
||||||
new_obs, reward, done, infos = env.step(self.unscale_action(clipped_action))
|
new_obs, reward, done, infos = env.step(self.policy.unscale_action(clipped_action))
|
||||||
|
|
||||||
# Only stop training if return value is False, not when it is None.
|
# Only stop training if return value is False, not when it is None.
|
||||||
if callback.on_step() is False:
|
if callback.on_step() is False:
|
||||||
|
|
|
||||||
|
|
@ -33,12 +33,50 @@ class Distribution(object):
|
||||||
|
|
||||||
def sample(self) -> th.Tensor:
|
def sample(self) -> th.Tensor:
|
||||||
"""
|
"""
|
||||||
returns a sample from the probabilty distribution
|
Returns a sample from the probabilty distribution
|
||||||
|
|
||||||
:return: (th.Tensor) the stochastic action
|
:return: (th.Tensor) the stochastic action
|
||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def mode(self) -> th.Tensor:
|
||||||
|
"""
|
||||||
|
Returns the most likely action (deterministic output)
|
||||||
|
from the probabilty distribution
|
||||||
|
|
||||||
|
:return: (th.Tensor) the stochastic action
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def get_actions(self, deterministic: bool = False) -> th.Tensor:
|
||||||
|
"""
|
||||||
|
Return actions according to the probabilty distribution.
|
||||||
|
|
||||||
|
:param deterministic: (bool)
|
||||||
|
:return: (th.Tensor)
|
||||||
|
"""
|
||||||
|
if deterministic:
|
||||||
|
return self.mode()
|
||||||
|
return self.sample()
|
||||||
|
|
||||||
|
def actions_from_params(self, *args, **kwargs) -> th.Tensor:
|
||||||
|
"""
|
||||||
|
Returns samples from the probabilty distribution
|
||||||
|
given its parameters.
|
||||||
|
|
||||||
|
:return: (th.Tensor) actions
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def log_prob_from_params(self, *args, **kwargs) -> Tuple[th.Tensor, th.Tensor]:
|
||||||
|
"""
|
||||||
|
Returns samples and the associated log probabilties
|
||||||
|
from the probabilty distribution given its parameters.
|
||||||
|
|
||||||
|
:return: (th.Tuple[th.Tensor, th.Tensor]) actions and log prob
|
||||||
|
"""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
|
||||||
def sum_independent_dims(tensor: th.Tensor) -> th.Tensor:
|
def sum_independent_dims(tensor: th.Tensor) -> th.Tensor:
|
||||||
"""
|
"""
|
||||||
|
|
@ -88,34 +126,37 @@ class DiagGaussianDistribution(Distribution):
|
||||||
return mean_actions, log_std
|
return mean_actions, log_std
|
||||||
|
|
||||||
def proba_distribution(self, mean_actions: th.Tensor,
|
def proba_distribution(self, mean_actions: th.Tensor,
|
||||||
log_std: th.Tensor,
|
log_std: th.Tensor) -> 'DiagGaussianDistribution':
|
||||||
deterministic: bool = False) -> Tuple[th.Tensor, 'DiagGaussianDistribution']:
|
|
||||||
"""
|
"""
|
||||||
Create and sample for the distribution given its parameters (mean, std)
|
Create the distribution given its parameters (mean, std)
|
||||||
|
|
||||||
:param mean_actions: (th.Tensor)
|
:param mean_actions: (th.Tensor)
|
||||||
:param log_std: (th.Tensor)
|
:param log_std: (th.Tensor)
|
||||||
:param deterministic: (bool)
|
:return: (DiagGaussianDistribution)
|
||||||
:return: (th.Tensor)
|
|
||||||
"""
|
"""
|
||||||
action_std = th.ones_like(mean_actions) * log_std.exp()
|
action_std = th.ones_like(mean_actions) * log_std.exp()
|
||||||
self.distribution = Normal(mean_actions, action_std)
|
self.distribution = Normal(mean_actions, action_std)
|
||||||
if deterministic:
|
return self
|
||||||
action = self.mode()
|
|
||||||
else:
|
|
||||||
action = self.sample()
|
|
||||||
return action, self
|
|
||||||
|
|
||||||
def mode(self) -> th.Tensor:
|
def mode(self) -> th.Tensor:
|
||||||
return self.distribution.mean
|
return self.distribution.mean
|
||||||
|
|
||||||
def sample(self) -> th.Tensor:
|
def sample(self) -> th.Tensor:
|
||||||
|
# Reparametrization trick to pass gradients
|
||||||
return self.distribution.rsample()
|
return self.distribution.rsample()
|
||||||
|
|
||||||
def entropy(self) -> th.Tensor:
|
def entropy(self) -> th.Tensor:
|
||||||
return sum_independent_dims(self.distribution.entropy())
|
return sum_independent_dims(self.distribution.entropy())
|
||||||
|
|
||||||
def log_prob_from_params(self, mean_actions: th.Tensor, log_std: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
def actions_from_params(self, mean_actions: th.Tensor,
|
||||||
|
log_std: th.Tensor,
|
||||||
|
deterministic: bool = False) -> th.Tensor:
|
||||||
|
# Update the proba distribution
|
||||||
|
self.proba_distribution(mean_actions, log_std)
|
||||||
|
return self.get_actions(deterministic=deterministic)
|
||||||
|
|
||||||
|
def log_prob_from_params(self, mean_actions: th.Tensor,
|
||||||
|
log_std: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
||||||
"""
|
"""
|
||||||
Compute the log probabilty of taking an action
|
Compute the log probabilty of taking an action
|
||||||
given the distribution parameters.
|
given the distribution parameters.
|
||||||
|
|
@ -124,20 +165,19 @@ class DiagGaussianDistribution(Distribution):
|
||||||
:param log_std: (th.Tensor)
|
:param log_std: (th.Tensor)
|
||||||
:return: (Tuple[th.Tensor, th.Tensor])
|
:return: (Tuple[th.Tensor, th.Tensor])
|
||||||
"""
|
"""
|
||||||
action, _ = self.proba_distribution(mean_actions, log_std)
|
actions = self.actions_from_params(mean_actions, log_std)
|
||||||
log_prob = self.log_prob(action)
|
log_prob = self.log_prob(actions)
|
||||||
return action, log_prob
|
return actions, log_prob
|
||||||
|
|
||||||
def log_prob(self, action: th.Tensor) -> th.Tensor:
|
def log_prob(self, actions: th.Tensor) -> th.Tensor:
|
||||||
"""
|
"""
|
||||||
Get the log probabilty of an action given a distribution.
|
Get the log probabilties of actions according to the distribution.
|
||||||
Note that you must call ``proba_distribution()`` method
|
Note that you must call ``proba_distribution()`` method before.
|
||||||
before.
|
|
||||||
|
|
||||||
:param action: (th.Tensor)
|
:param actions: (th.Tensor)
|
||||||
:return: (th.Tensor)
|
:return: (th.Tensor)
|
||||||
"""
|
"""
|
||||||
log_prob = self.distribution.log_prob(action)
|
log_prob = self.distribution.log_prob(actions)
|
||||||
return sum_independent_dims(log_prob)
|
return sum_independent_dims(log_prob)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -154,17 +194,17 @@ class SquashedDiagGaussianDistribution(DiagGaussianDistribution):
|
||||||
super(SquashedDiagGaussianDistribution, self).__init__(action_dim)
|
super(SquashedDiagGaussianDistribution, self).__init__(action_dim)
|
||||||
# Avoid NaN (prevents division by zero or log of zero)
|
# Avoid NaN (prevents division by zero or log of zero)
|
||||||
self.epsilon = epsilon
|
self.epsilon = epsilon
|
||||||
self.gaussian_action = None
|
self.gaussian_actions = None
|
||||||
|
|
||||||
def proba_distribution(self, mean_actions, log_std, deterministic=False):
|
def proba_distribution(self, mean_actions: th.Tensor,
|
||||||
action, _ = super(SquashedDiagGaussianDistribution, self).proba_distribution(mean_actions, log_std,
|
log_std: th.Tensor) -> 'SquashedDiagGaussianDistribution':
|
||||||
deterministic)
|
super(SquashedDiagGaussianDistribution, self).proba_distribution(mean_actions, log_std)
|
||||||
return action, self
|
return self
|
||||||
|
|
||||||
def mode(self) -> th.Tensor:
|
def mode(self) -> th.Tensor:
|
||||||
self.gaussian_action = self.distribution.mean
|
self.gaussian_actions = self.distribution.mean
|
||||||
# Squash the output
|
# Squash the output
|
||||||
return th.tanh(self.gaussian_action)
|
return th.tanh(self.gaussian_actions)
|
||||||
|
|
||||||
def entropy(self) -> Optional[th.Tensor]:
|
def entropy(self) -> Optional[th.Tensor]:
|
||||||
# No analytical form,
|
# No analytical form,
|
||||||
|
|
@ -172,27 +212,30 @@ class SquashedDiagGaussianDistribution(DiagGaussianDistribution):
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def sample(self) -> th.Tensor:
|
def sample(self) -> th.Tensor:
|
||||||
self.gaussian_action = self.distribution.rsample()
|
# Reparametrization trick to pass gradients
|
||||||
return th.tanh(self.gaussian_action)
|
self.gaussian_actions = self.distribution.rsample()
|
||||||
|
return th.tanh(self.gaussian_actions)
|
||||||
|
|
||||||
def log_prob_from_params(self, mean_actions, log_std) -> Tuple[th.Tensor, th.Tensor]:
|
def log_prob_from_params(self, mean_actions: th.Tensor,
|
||||||
action, _ = self.proba_distribution(mean_actions, log_std)
|
log_std: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
||||||
log_prob = self.log_prob(action, self.gaussian_action)
|
action = self.actions_from_params(mean_actions, log_std)
|
||||||
|
log_prob = self.log_prob(action, self.gaussian_actions)
|
||||||
return action, log_prob
|
return action, log_prob
|
||||||
|
|
||||||
def log_prob(self, action: th.Tensor, gaussian_action: Optional[th.Tensor] = None) -> th.Tensor:
|
def log_prob(self, actions: th.Tensor,
|
||||||
|
gaussian_actions: Optional[th.Tensor] = None) -> th.Tensor:
|
||||||
# Inverse tanh
|
# Inverse tanh
|
||||||
# Naive implementation (not stable): 0.5 * torch.log((1 + x) / (1 - x))
|
# Naive implementation (not stable): 0.5 * torch.log((1 + x) / (1 - x))
|
||||||
# We use numpy to avoid numerical instability
|
# We use numpy to avoid numerical instability
|
||||||
if gaussian_action is None:
|
if gaussian_actions is None:
|
||||||
# It will be clipped to avoid NaN when inversing tanh
|
# It will be clipped to avoid NaN when inversing tanh
|
||||||
gaussian_action = TanhBijector.inverse(action)
|
gaussian_actions = TanhBijector.inverse(actions)
|
||||||
|
|
||||||
# Log likelihood for a Gaussian distribution
|
# Log likelihood for a Gaussian distribution
|
||||||
log_prob = super(SquashedDiagGaussianDistribution, self).log_prob(gaussian_action)
|
log_prob = super(SquashedDiagGaussianDistribution, self).log_prob(gaussian_actions)
|
||||||
# Squash correction (from original SAC implementation)
|
# Squash correction (from original SAC implementation)
|
||||||
# this comes from the fact that tanh is bijective and differentiable
|
# this comes from the fact that tanh is bijective and differentiable
|
||||||
log_prob -= th.sum(th.log(1 - action ** 2 + self.epsilon), dim=1)
|
log_prob -= th.sum(th.log(1 - actions ** 2 + self.epsilon), dim=1)
|
||||||
return log_prob
|
return log_prob
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -214,20 +257,16 @@ class CategoricalDistribution(Distribution):
|
||||||
it will be the logits of the Categorical distribution.
|
it will be the logits of the Categorical distribution.
|
||||||
You can then get probabilties using a softmax.
|
You can then get probabilties using a softmax.
|
||||||
|
|
||||||
:param latent_dim: (int) Dimension og the last layer of the policy (before the action layer)
|
:param latent_dim: (int) Dimension of the last layer
|
||||||
|
of the policy network (before the action layer)
|
||||||
:return: (nn.Linear)
|
:return: (nn.Linear)
|
||||||
"""
|
"""
|
||||||
action_logits = nn.Linear(latent_dim, self.action_dim)
|
action_logits = nn.Linear(latent_dim, self.action_dim)
|
||||||
return action_logits
|
return action_logits
|
||||||
|
|
||||||
def proba_distribution(self, action_logits: th.Tensor,
|
def proba_distribution(self, action_logits: th.Tensor) -> 'CategoricalDistribution':
|
||||||
deterministic: bool = False) -> Tuple[th.Tensor, 'CategoricalDistribution']:
|
|
||||||
self.distribution = Categorical(logits=action_logits)
|
self.distribution = Categorical(logits=action_logits)
|
||||||
if deterministic:
|
return self
|
||||||
action = self.mode()
|
|
||||||
else:
|
|
||||||
action = self.sample()
|
|
||||||
return action, self
|
|
||||||
|
|
||||||
def mode(self) -> th.Tensor:
|
def mode(self) -> th.Tensor:
|
||||||
return th.argmax(self.distribution.probs, dim=1)
|
return th.argmax(self.distribution.probs, dim=1)
|
||||||
|
|
@ -238,14 +277,19 @@ class CategoricalDistribution(Distribution):
|
||||||
def entropy(self) -> th.Tensor:
|
def entropy(self) -> th.Tensor:
|
||||||
return self.distribution.entropy()
|
return self.distribution.entropy()
|
||||||
|
|
||||||
def log_prob_from_params(self, action_logits: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
def actions_from_params(self, action_logits: th.Tensor,
|
||||||
action, _ = self.proba_distribution(action_logits)
|
deterministic: bool = False) -> th.Tensor:
|
||||||
log_prob = self.log_prob(action)
|
# Update the proba distribution
|
||||||
return action, log_prob
|
self.proba_distribution(action_logits)
|
||||||
|
return self.get_actions(deterministic=deterministic)
|
||||||
|
|
||||||
def log_prob(self, action: th.Tensor) -> th.Tensor:
|
def log_prob_from_params(self, action_logits: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
||||||
log_prob = self.distribution.log_prob(action)
|
actions = self.actions_from_params(action_logits)
|
||||||
return log_prob
|
log_prob = self.log_prob(actions)
|
||||||
|
return actions, log_prob
|
||||||
|
|
||||||
|
def log_prob(self, actions: th.Tensor) -> th.Tensor:
|
||||||
|
return self.distribution.log_prob(actions)
|
||||||
|
|
||||||
|
|
||||||
class StateDependentNoiseDistribution(Distribution):
|
class StateDependentNoiseDistribution(Distribution):
|
||||||
|
|
@ -283,6 +327,7 @@ class StateDependentNoiseDistribution(Distribution):
|
||||||
self.weights_dist = None
|
self.weights_dist = None
|
||||||
self.exploration_mat = None
|
self.exploration_mat = None
|
||||||
self.exploration_matrices = None
|
self.exploration_matrices = None
|
||||||
|
self._latent_sde = None
|
||||||
self.use_expln = use_expln
|
self.use_expln = use_expln
|
||||||
self.full_std = full_std
|
self.full_std = full_std
|
||||||
self.epsilon = epsilon
|
self.epsilon = epsilon
|
||||||
|
|
@ -327,7 +372,9 @@ class StateDependentNoiseDistribution(Distribution):
|
||||||
"""
|
"""
|
||||||
std = self.get_std(log_std)
|
std = self.get_std(log_std)
|
||||||
self.weights_dist = Normal(th.zeros_like(std), std)
|
self.weights_dist = Normal(th.zeros_like(std), std)
|
||||||
|
# Reparametrization trick to pass gradients
|
||||||
self.exploration_mat = self.weights_dist.rsample()
|
self.exploration_mat = self.weights_dist.rsample()
|
||||||
|
# Pre-compute matrices in case of parallel exploration
|
||||||
self.exploration_matrices = self.weights_dist.rsample((batch_size,))
|
self.exploration_matrices = self.weights_dist.rsample((batch_size,))
|
||||||
|
|
||||||
def proba_distribution_net(self, latent_dim: int, log_std_init: float = -2.0,
|
def proba_distribution_net(self, latent_dim: int, log_std_init: float = -2.0,
|
||||||
|
|
@ -358,33 +405,26 @@ class StateDependentNoiseDistribution(Distribution):
|
||||||
|
|
||||||
def proba_distribution(self, mean_actions: th.Tensor,
|
def proba_distribution(self, mean_actions: th.Tensor,
|
||||||
log_std: th.Tensor,
|
log_std: th.Tensor,
|
||||||
latent_sde: th.Tensor,
|
latent_sde: th.Tensor) -> 'StateDependentNoiseDistribution':
|
||||||
deterministic: bool = False) -> Tuple[th.Tensor, 'StateDependentNoiseDistribution']:
|
|
||||||
"""
|
"""
|
||||||
Create and sample for the distribution given its parameters (mean, std)
|
Create the distribution given its parameters (mean, std)
|
||||||
|
|
||||||
:param mean_actions: (th.Tensor)
|
:param mean_actions: (th.Tensor)
|
||||||
:param log_std: (th.Tensor)
|
:param log_std: (th.Tensor)
|
||||||
:param latent_sde: (th.Tensor)
|
:param latent_sde: (th.Tensor)
|
||||||
:param deterministic: (bool)
|
:return: (StateDependentNoiseDistribution)
|
||||||
:return: (Tuple[th.Tensor, Distribution])
|
|
||||||
"""
|
"""
|
||||||
# Stop gradient if we don't want to influence the features
|
# Stop gradient if we don't want to influence the features
|
||||||
latent_sde = latent_sde if self.learn_features else latent_sde.detach()
|
self._latent_sde = latent_sde if self.learn_features else latent_sde.detach()
|
||||||
variance = th.mm(latent_sde ** 2, self.get_std(log_std) ** 2)
|
variance = th.mm(latent_sde ** 2, self.get_std(log_std) ** 2)
|
||||||
self.distribution = Normal(mean_actions, th.sqrt(variance + self.epsilon))
|
self.distribution = Normal(mean_actions, th.sqrt(variance + self.epsilon))
|
||||||
|
return self
|
||||||
if deterministic:
|
|
||||||
action = self.mode()
|
|
||||||
else:
|
|
||||||
action = self.sample(latent_sde)
|
|
||||||
return action, self
|
|
||||||
|
|
||||||
def mode(self) -> th.Tensor:
|
def mode(self) -> th.Tensor:
|
||||||
action = self.distribution.mean
|
actions = self.distribution.mean
|
||||||
if self.bijector is not None:
|
if self.bijector is not None:
|
||||||
return self.bijector.forward(action)
|
return self.bijector.forward(actions)
|
||||||
return action
|
return actions
|
||||||
|
|
||||||
def get_noise(self, latent_sde: th.Tensor) -> th.Tensor:
|
def get_noise(self, latent_sde: th.Tensor) -> th.Tensor:
|
||||||
latent_sde = latent_sde if self.learn_features else latent_sde.detach()
|
latent_sde = latent_sde if self.learn_features else latent_sde.detach()
|
||||||
|
|
@ -398,12 +438,12 @@ class StateDependentNoiseDistribution(Distribution):
|
||||||
noise = th.bmm(latent_sde, self.exploration_matrices)
|
noise = th.bmm(latent_sde, self.exploration_matrices)
|
||||||
return noise.squeeze(1)
|
return noise.squeeze(1)
|
||||||
|
|
||||||
def sample(self, latent_sde: th.Tensor) -> th.Tensor:
|
def sample(self) -> th.Tensor:
|
||||||
noise = self.get_noise(latent_sde)
|
noise = self.get_noise(self._latent_sde)
|
||||||
action = self.distribution.mean + noise
|
actions = self.distribution.mean + noise
|
||||||
if self.bijector is not None:
|
if self.bijector is not None:
|
||||||
return self.bijector.forward(action)
|
return self.bijector.forward(actions)
|
||||||
return action
|
return actions
|
||||||
|
|
||||||
def entropy(self) -> Optional[th.Tensor]:
|
def entropy(self) -> Optional[th.Tensor]:
|
||||||
# No analytical form,
|
# No analytical form,
|
||||||
|
|
@ -412,26 +452,34 @@ class StateDependentNoiseDistribution(Distribution):
|
||||||
return None
|
return None
|
||||||
return sum_independent_dims(self.distribution.entropy())
|
return sum_independent_dims(self.distribution.entropy())
|
||||||
|
|
||||||
|
def actions_from_params(self, mean_actions: th.Tensor,
|
||||||
|
log_std: th.Tensor,
|
||||||
|
latent_sde: th.Tensor,
|
||||||
|
deterministic: bool = False) -> th.Tensor:
|
||||||
|
# Update the proba distribution
|
||||||
|
self.proba_distribution(mean_actions, log_std, latent_sde)
|
||||||
|
return self.get_actions(deterministic=deterministic)
|
||||||
|
|
||||||
def log_prob_from_params(self, mean_actions: th.Tensor,
|
def log_prob_from_params(self, mean_actions: th.Tensor,
|
||||||
log_std: th.Tensor,
|
log_std: th.Tensor,
|
||||||
latent_sde: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
latent_sde: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
||||||
action, _ = self.proba_distribution(mean_actions, log_std, latent_sde)
|
actions = self.actions_from_params(mean_actions, log_std, latent_sde)
|
||||||
log_prob = self.log_prob(action)
|
log_prob = self.log_prob(actions)
|
||||||
return action, log_prob
|
return actions, log_prob
|
||||||
|
|
||||||
def log_prob(self, action: th.Tensor) -> th.Tensor:
|
def log_prob(self, actions: th.Tensor) -> th.Tensor:
|
||||||
if self.bijector is not None:
|
if self.bijector is not None:
|
||||||
gaussian_action = self.bijector.inverse(action)
|
gaussian_actions = self.bijector.inverse(actions)
|
||||||
else:
|
else:
|
||||||
gaussian_action = action
|
gaussian_actions = actions
|
||||||
# log likelihood for a gaussian
|
# log likelihood for a gaussian
|
||||||
log_prob = self.distribution.log_prob(gaussian_action)
|
log_prob = self.distribution.log_prob(gaussian_actions)
|
||||||
# Sum along action dim
|
# Sum along action dim
|
||||||
log_prob = sum_independent_dims(log_prob)
|
log_prob = sum_independent_dims(log_prob)
|
||||||
|
|
||||||
if self.bijector is not None:
|
if self.bijector is not None:
|
||||||
# Squash correction (from original SAC implementation)
|
# Squash correction (from original SAC implementation)
|
||||||
log_prob -= th.sum(self.bijector.log_prob_correction(gaussian_action), dim=1)
|
log_prob -= th.sum(self.bijector.log_prob_correction(gaussian_actions), dim=1)
|
||||||
return log_prob
|
return log_prob
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -63,7 +63,7 @@ class BasePolicy(nn.Module):
|
||||||
def forward(self, *_args, **kwargs):
|
def forward(self, *_args, **kwargs):
|
||||||
raise NotImplementedError()
|
raise NotImplementedError()
|
||||||
|
|
||||||
def predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
def _predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
||||||
"""
|
"""
|
||||||
Get the action according to the policy for a given observation.
|
Get the action according to the policy for a given observation.
|
||||||
|
|
||||||
|
|
@ -73,9 +73,127 @@ class BasePolicy(nn.Module):
|
||||||
"""
|
"""
|
||||||
raise NotImplementedError()
|
raise NotImplementedError()
|
||||||
|
|
||||||
|
def predict(self, observation: np.ndarray,
|
||||||
|
state: Optional[np.ndarray] = None,
|
||||||
|
mask: Optional[np.ndarray] = None,
|
||||||
|
deterministic: bool = False) -> Tuple[np.ndarray, Optional[np.ndarray]]:
|
||||||
|
"""
|
||||||
|
Get the policy action and state from an observation (and optional state).
|
||||||
|
|
||||||
|
:param observation: (np.ndarray) the input observation
|
||||||
|
:param state: (Optional[np.ndarray]) The last states (can be None, used in recurrent policies)
|
||||||
|
:param mask: (Optional[np.ndarray]) The last masks (can be None, used in recurrent policies)
|
||||||
|
:param deterministic: (bool) Whether or not to return deterministic actions.
|
||||||
|
:return: (Tuple[np.ndarray, Optional[np.ndarray]]) the model's action and the next state
|
||||||
|
(used in recurrent policies)
|
||||||
|
"""
|
||||||
|
# if state is None:
|
||||||
|
# state = self.initial_state
|
||||||
|
# if mask is None:
|
||||||
|
# mask = [False for _ in range(self.n_envs)]
|
||||||
|
observation = np.array(observation)
|
||||||
|
vectorized_env = self._is_vectorized_observation(observation, self.observation_space)
|
||||||
|
|
||||||
|
observation = observation.reshape((-1,) + self.observation_space.shape)
|
||||||
|
observation = th.as_tensor(observation).to(self.device)
|
||||||
|
with th.no_grad():
|
||||||
|
actions = self._predict(observation, deterministic=deterministic)
|
||||||
|
# Convert to numpy
|
||||||
|
actions = actions.cpu().numpy()
|
||||||
|
|
||||||
|
# Rescale to proper domain when using squashing
|
||||||
|
if isinstance(self.action_space, gym.spaces.Box) and self.squash_output:
|
||||||
|
actions = self.unscale_action(actions)
|
||||||
|
|
||||||
|
clipped_actions = actions
|
||||||
|
# Clip the actions to avoid out of bound error when using gaussian distribution
|
||||||
|
if isinstance(self.action_space, gym.spaces.Box) and not self.squash_output:
|
||||||
|
clipped_actions = np.clip(actions, self.action_space.low, self.action_space.high)
|
||||||
|
|
||||||
|
if not vectorized_env:
|
||||||
|
if state is not None:
|
||||||
|
raise ValueError("Error: The environment must be vectorized when using recurrent policies.")
|
||||||
|
clipped_actions = clipped_actions[0]
|
||||||
|
|
||||||
|
return clipped_actions, state
|
||||||
|
|
||||||
|
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: (np.ndarray) Action to scale
|
||||||
|
:return: (np.ndarray) Scaled action
|
||||||
|
"""
|
||||||
|
low, high = self.action_space.low, self.action_space.high
|
||||||
|
return 2.0 * ((action - low) / (high - low)) - 1.0
|
||||||
|
|
||||||
|
def unscale_action(self, scaled_action: np.ndarray) -> np.ndarray:
|
||||||
|
"""
|
||||||
|
Rescale the action from [-1, 1] to [low, high]
|
||||||
|
(no need for symmetric action space)
|
||||||
|
|
||||||
|
:param scaled_action: Action to un-scale
|
||||||
|
"""
|
||||||
|
low, high = self.action_space.low, self.action_space.high
|
||||||
|
return low + (0.5 * (scaled_action + 1.0) * (high - low))
|
||||||
|
|
||||||
|
@staticmethod
|
||||||
|
def _is_vectorized_observation(observation: np.ndarray, observation_space: gym.spaces.Space) -> bool:
|
||||||
|
"""
|
||||||
|
For every observation type, detects and validates the shape,
|
||||||
|
then returns whether or not the observation is vectorized.
|
||||||
|
|
||||||
|
:param observation: (np.ndarray) the input observation to validate
|
||||||
|
:param observation_space: (gym.spaces) the observation space
|
||||||
|
:return: (bool) whether the given observation is vectorized or not
|
||||||
|
"""
|
||||||
|
if isinstance(observation_space, gym.spaces.Box):
|
||||||
|
if observation.shape == observation_space.shape:
|
||||||
|
return False
|
||||||
|
elif observation.shape[1:] == observation_space.shape:
|
||||||
|
return True
|
||||||
|
else:
|
||||||
|
raise ValueError("Error: Unexpected observation shape {} for ".format(observation.shape) +
|
||||||
|
"Box environment, please use {} ".format(observation_space.shape) +
|
||||||
|
"or (n_env, {}) for the observation shape."
|
||||||
|
.format(", ".join(map(str, observation_space.shape))))
|
||||||
|
elif isinstance(observation_space, gym.spaces.Discrete):
|
||||||
|
if observation.shape == (): # A numpy array of a number, has shape empty tuple '()'
|
||||||
|
return False
|
||||||
|
elif len(observation.shape) == 1:
|
||||||
|
return True
|
||||||
|
else:
|
||||||
|
raise ValueError("Error: Unexpected observation shape {} for ".format(observation.shape) +
|
||||||
|
"Discrete environment, please use (1,) or (n_env, 1) for the observation shape.")
|
||||||
|
# TODO: add support for MultiDiscrete and MultiBinary observation spaces
|
||||||
|
# elif isinstance(observation_space, gym.spaces.MultiDiscrete):
|
||||||
|
# if observation.shape == (len(observation_space.nvec),):
|
||||||
|
# return False
|
||||||
|
# elif len(observation.shape) == 2 and observation.shape[1] == len(observation_space.nvec):
|
||||||
|
# return True
|
||||||
|
# else:
|
||||||
|
# raise ValueError("Error: Unexpected observation shape {} for MultiDiscrete ".format(observation.shape) +
|
||||||
|
# "environment, please use ({},) or ".format(len(observation_space.nvec)) +
|
||||||
|
# "(n_env, {}) for the observation shape.".format(len(observation_space.nvec)))
|
||||||
|
# elif isinstance(observation_space, gym.spaces.MultiBinary):
|
||||||
|
# if observation.shape == (observation_space.n,):
|
||||||
|
# return False
|
||||||
|
# elif len(observation.shape) == 2 and observation.shape[1] == observation_space.n:
|
||||||
|
# return True
|
||||||
|
# else:
|
||||||
|
# raise ValueError("Error: Unexpected observation shape {} for MultiBinary ".format(observation.shape) +
|
||||||
|
# "environment, please use ({},) or ".format(observation_space.n) +
|
||||||
|
# "(n_env, {}) for the observation shape.".format(observation_space.n))
|
||||||
|
else:
|
||||||
|
raise ValueError("Error: Cannot determine if the observation is vectorized with the space type {}."
|
||||||
|
.format(observation_space))
|
||||||
|
|
||||||
|
|
||||||
def save(self, path: str) -> None:
|
def save(self, path: str) -> None:
|
||||||
"""
|
"""
|
||||||
Save model to a given location.
|
Save policy weights to a given location.
|
||||||
|
NOTE: we don't save policy parameters
|
||||||
|
|
||||||
:param path: (str)
|
:param path: (str)
|
||||||
"""
|
"""
|
||||||
|
|
@ -83,7 +201,8 @@ class BasePolicy(nn.Module):
|
||||||
|
|
||||||
def load(self, path: str) -> None:
|
def load(self, path: str) -> None:
|
||||||
"""
|
"""
|
||||||
Load saved model from path.
|
Load policy weights from path.
|
||||||
|
NOTE: we don't load policy parameters
|
||||||
|
|
||||||
:param path: (str)
|
:param path: (str)
|
||||||
"""
|
"""
|
||||||
|
|
|
||||||
|
|
@ -155,11 +155,11 @@ class PPOPolicy(BasePolicy):
|
||||||
"""
|
"""
|
||||||
latent_pi, latent_vf, latent_sde = self._get_latent(obs)
|
latent_pi, latent_vf, latent_sde = self._get_latent(obs)
|
||||||
# Evaluate the values for the given observations
|
# Evaluate the values for the given observations
|
||||||
value = self.value_net(latent_vf)
|
values = self.value_net(latent_vf)
|
||||||
action, action_distribution = self._get_action_dist_from_latent(latent_pi, latent_sde=latent_sde,
|
distribution = self._get_action_dist_from_latent(latent_pi, latent_sde=latent_sde)
|
||||||
deterministic=deterministic)
|
actions = distribution.get_actions(deterministic=deterministic)
|
||||||
log_prob = action_distribution.log_prob(action)
|
log_prob = distribution.log_prob(actions)
|
||||||
return action, value, log_prob
|
return actions, values, log_prob
|
||||||
|
|
||||||
def _get_latent(self, obs: th.Tensor) -> Tuple[th.Tensor, th.Tensor, th.Tensor]:
|
def _get_latent(self, obs: th.Tensor) -> Tuple[th.Tensor, th.Tensor, th.Tensor]:
|
||||||
"""
|
"""
|
||||||
|
|
@ -180,33 +180,29 @@ class PPOPolicy(BasePolicy):
|
||||||
return latent_pi, latent_vf, latent_sde
|
return latent_pi, latent_vf, latent_sde
|
||||||
|
|
||||||
def _get_action_dist_from_latent(self, latent_pi: th.Tensor,
|
def _get_action_dist_from_latent(self, latent_pi: th.Tensor,
|
||||||
latent_sde: Optional[th.Tensor] = None,
|
latent_sde: Optional[th.Tensor] = None) -> Distribution:
|
||||||
deterministic: bool = False) -> Tuple[th.Tensor, Distribution]:
|
|
||||||
"""
|
"""
|
||||||
Retrieve action and associated action distribution
|
Retrieve action distribution given the latent codes.
|
||||||
given the latent codes.
|
|
||||||
|
|
||||||
:param latent_pi: (th.Tensor) Latent code for the actor
|
:param latent_pi: (th.Tensor) Latent code for the actor
|
||||||
:param latent_sde: (Optional[th.Tensor]) Latent code for the SDE exploration function
|
:param latent_sde: (Optional[th.Tensor]) Latent code for the SDE exploration function
|
||||||
:param deterministic: (bool) Whether to sample or use deterministic actions
|
:return: (Distribution) Action distribution
|
||||||
:return: (Tuple[th.Tensor, Distribution]) Action and action distribution
|
|
||||||
"""
|
"""
|
||||||
mean_actions = self.action_net(latent_pi)
|
mean_actions = self.action_net(latent_pi)
|
||||||
|
|
||||||
if isinstance(self.action_dist, DiagGaussianDistribution):
|
if isinstance(self.action_dist, DiagGaussianDistribution):
|
||||||
return self.action_dist.proba_distribution(mean_actions, self.log_std, deterministic=deterministic)
|
return self.action_dist.proba_distribution(mean_actions, self.log_std)
|
||||||
|
|
||||||
elif isinstance(self.action_dist, CategoricalDistribution):
|
elif isinstance(self.action_dist, CategoricalDistribution):
|
||||||
# Here mean_actions are the logits before the softmax
|
# Here mean_actions are the logits before the softmax
|
||||||
return self.action_dist.proba_distribution(mean_actions, deterministic=deterministic)
|
return self.action_dist.proba_distribution(action_logits=mean_actions)
|
||||||
|
|
||||||
elif isinstance(self.action_dist, StateDependentNoiseDistribution):
|
elif isinstance(self.action_dist, StateDependentNoiseDistribution):
|
||||||
return self.action_dist.proba_distribution(mean_actions, self.log_std, latent_sde,
|
return self.action_dist.proba_distribution(mean_actions, self.log_std, latent_sde)
|
||||||
deterministic=deterministic)
|
|
||||||
else:
|
else:
|
||||||
raise ValueError('Invalid action distribution')
|
raise ValueError('Invalid action distribution')
|
||||||
|
|
||||||
def predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
def _predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
||||||
"""
|
"""
|
||||||
Get the action according to the policy for a given observation.
|
Get the action according to the policy for a given observation.
|
||||||
|
|
||||||
|
|
@ -215,27 +211,25 @@ class PPOPolicy(BasePolicy):
|
||||||
:return: (th.Tensor) Taken action according to the policy
|
:return: (th.Tensor) Taken action according to the policy
|
||||||
"""
|
"""
|
||||||
latent_pi, _, latent_sde = self._get_latent(observation)
|
latent_pi, _, latent_sde = self._get_latent(observation)
|
||||||
action, _ = self._get_action_dist_from_latent(latent_pi, latent_sde, deterministic=deterministic)
|
distribution = self._get_action_dist_from_latent(latent_pi, latent_sde)
|
||||||
return action
|
return distribution.get_actions(deterministic=deterministic)
|
||||||
|
|
||||||
def evaluate_actions(self, obs: th.Tensor,
|
def evaluate_actions(self, obs: th.Tensor,
|
||||||
actions: th.Tensor,
|
actions: th.Tensor) -> Tuple[th.Tensor, th.Tensor, th.Tensor]:
|
||||||
deterministic: bool = False) -> Tuple[th.Tensor, th.Tensor, th.Tensor]:
|
|
||||||
"""
|
"""
|
||||||
Evaluate actions according to the current policy,
|
Evaluate actions according to the current policy,
|
||||||
given the observations.
|
given the observations.
|
||||||
|
|
||||||
:param obs: (th.Tensor)
|
:param obs: (th.Tensor)
|
||||||
:param actions: (th.Tensor)
|
:param actions: (th.Tensor)
|
||||||
:param deterministic: (bool)
|
|
||||||
:return: (th.Tensor, th.Tensor, th.Tensor) estimated value, log likelihood of taking those actions
|
:return: (th.Tensor, th.Tensor, th.Tensor) estimated value, log likelihood of taking those actions
|
||||||
and entropy of the action distribution.
|
and entropy of the action distribution.
|
||||||
"""
|
"""
|
||||||
latent_pi, latent_vf, latent_sde = self._get_latent(obs)
|
latent_pi, latent_vf, latent_sde = self._get_latent(obs)
|
||||||
_, action_distribution = self._get_action_dist_from_latent(latent_pi, latent_sde, deterministic=deterministic)
|
distribution = self._get_action_dist_from_latent(latent_pi, latent_sde)
|
||||||
log_prob = action_distribution.log_prob(actions)
|
log_prob = distribution.log_prob(actions)
|
||||||
values = self.value_net(latent_vf)
|
values = self.value_net(latent_vf)
|
||||||
return values, log_prob, action_distribution.entropy()
|
return values, log_prob, distribution.entropy()
|
||||||
|
|
||||||
|
|
||||||
MlpPolicy = PPOPolicy
|
MlpPolicy = PPOPolicy
|
||||||
|
|
|
||||||
|
|
@ -1,4 +1,4 @@
|
||||||
from typing import Optional, List, Tuple, Callable, Union, Type
|
from typing import Optional, List, Tuple, Callable, Union, Type, Dict
|
||||||
|
|
||||||
import gym
|
import gym
|
||||||
import torch as th
|
import torch as th
|
||||||
|
|
@ -38,6 +38,7 @@ class Actor(BasePolicy):
|
||||||
:param clip_mean: (float) Clip the mean output when using SDE to avoid numerical instability.
|
:param clip_mean: (float) Clip the mean output when using SDE to avoid numerical instability.
|
||||||
:param normalize_images: (bool) Whether to normalize images or not,
|
:param normalize_images: (bool) Whether to normalize images or not,
|
||||||
dividing by 255.0 (True by default)
|
dividing by 255.0 (True by default)
|
||||||
|
:param device: (Union[th.device, str]) Device on which the code should run.
|
||||||
"""
|
"""
|
||||||
def __init__(self, observation_space: gym.spaces.Space,
|
def __init__(self, observation_space: gym.spaces.Space,
|
||||||
action_space: gym.spaces.Space,
|
action_space: gym.spaces.Space,
|
||||||
|
|
@ -51,10 +52,12 @@ class Actor(BasePolicy):
|
||||||
sde_net_arch: Optional[List[int]] = None,
|
sde_net_arch: Optional[List[int]] = None,
|
||||||
use_expln: bool = False,
|
use_expln: bool = False,
|
||||||
clip_mean: float = 2.0,
|
clip_mean: float = 2.0,
|
||||||
normalize_images: bool = True):
|
normalize_images: bool = True,
|
||||||
|
device: Union[th.device, str] = 'cpu'):
|
||||||
super(Actor, self).__init__(observation_space, action_space,
|
super(Actor, self).__init__(observation_space, action_space,
|
||||||
features_extractor=features_extractor,
|
features_extractor=features_extractor,
|
||||||
normalize_images=normalize_images)
|
normalize_images=normalize_images,
|
||||||
|
device=device)
|
||||||
|
|
||||||
action_dim = get_action_dim(self.action_space)
|
action_dim = get_action_dim(self.action_space)
|
||||||
|
|
||||||
|
|
@ -108,38 +111,43 @@ class Actor(BasePolicy):
|
||||||
'reset_noise() is only available when using SDE'
|
'reset_noise() is only available when using SDE'
|
||||||
self.action_dist.sample_weights(self.log_std, batch_size=batch_size)
|
self.action_dist.sample_weights(self.log_std, batch_size=batch_size)
|
||||||
|
|
||||||
def _get_latent(self, obs: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
def get_action_dist_params(self, obs: th.Tensor) -> Tuple[th.Tensor, th.Tensor, Dict[str, th.Tensor]]:
|
||||||
|
"""
|
||||||
|
Get the parameters for the action distribution.
|
||||||
|
|
||||||
|
:param obs: (th.Tensor)
|
||||||
|
:return: (Tuple[th.Tensor, th.Tensor, Dict[str, th.Tensor]])
|
||||||
|
Mean, standard deviation and optional keyword arguments.
|
||||||
|
"""
|
||||||
features = self.extract_features(obs)
|
features = self.extract_features(obs)
|
||||||
latent_pi = self.latent_pi(features)
|
latent_pi = self.latent_pi(features)
|
||||||
latent_sde = self.sde_features_extractor(features) if self.sde_features_extractor is not None else latent_pi
|
|
||||||
return latent_pi, latent_sde
|
|
||||||
|
|
||||||
def get_action_dist_params(self, obs: th.Tensor) -> Tuple[th.Tensor, th.Tensor, th.Tensor]:
|
|
||||||
latent_pi, latent_sde = self._get_latent(obs)
|
|
||||||
mean_actions = self.mu(latent_pi)
|
mean_actions = self.mu(latent_pi)
|
||||||
|
|
||||||
if self.use_sde:
|
if self.use_sde:
|
||||||
log_std = self.log_std
|
latent_sde = latent_pi
|
||||||
else:
|
if self.sde_features_extractor is not None:
|
||||||
|
latent_sde = self.sde_features_extractor(features)
|
||||||
|
return mean_actions, self.log_std, dict(latent_sde=latent_sde)
|
||||||
|
# Unstructured exploration (Original implementation)
|
||||||
log_std = self.log_std(latent_pi)
|
log_std = self.log_std(latent_pi)
|
||||||
# Original Implementation to cap the standard deviation
|
# Original Implementation to cap the standard deviation
|
||||||
log_std = th.clamp(log_std, LOG_STD_MIN, LOG_STD_MAX)
|
log_std = th.clamp(log_std, LOG_STD_MIN, LOG_STD_MAX)
|
||||||
return mean_actions, log_std, latent_sde
|
return mean_actions, log_std, {}
|
||||||
|
|
||||||
def forward(self, obs: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
def forward(self, obs: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
||||||
mean_actions, log_std, latent_sde = self.get_action_dist_params(obs)
|
mean_actions, log_std, kwargs = self.get_action_dist_params(obs)
|
||||||
kwargs = dict(latent_sde=latent_sde) if self.use_sde else {}
|
|
||||||
# Note: the action is squashed
|
# Note: the action is squashed
|
||||||
action, _ = self.action_dist.proba_distribution(mean_actions, log_std,
|
return self.action_dist.actions_from_params(mean_actions, log_std,
|
||||||
deterministic=deterministic, **kwargs)
|
deterministic=deterministic, **kwargs)
|
||||||
return action
|
|
||||||
|
|
||||||
def action_log_prob(self, obs: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
def action_log_prob(self, obs: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
||||||
mean_actions, log_std, latent_sde = self.get_action_dist_params(obs)
|
mean_actions, log_std, kwargs = self.get_action_dist_params(obs)
|
||||||
kwargs = dict(latent_sde=latent_sde) if self.use_sde else {}
|
|
||||||
# return action and associated log prob
|
# return action and associated log prob
|
||||||
return self.action_dist.log_prob_from_params(mean_actions, log_std, **kwargs)
|
return self.action_dist.log_prob_from_params(mean_actions, log_std, **kwargs)
|
||||||
|
|
||||||
|
def _predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
||||||
|
return self.forward(observation, deterministic)
|
||||||
|
|
||||||
|
|
||||||
class Critic(BasePolicy):
|
class Critic(BasePolicy):
|
||||||
"""
|
"""
|
||||||
|
|
@ -154,6 +162,7 @@ class Critic(BasePolicy):
|
||||||
:param activation_fn: (Type[nn.Module]) Activation function
|
:param activation_fn: (Type[nn.Module]) Activation function
|
||||||
:param normalize_images: (bool) Whether to normalize images or not,
|
:param normalize_images: (bool) Whether to normalize images or not,
|
||||||
dividing by 255.0 (True by default)
|
dividing by 255.0 (True by default)
|
||||||
|
:param device: (Union[th.device, str]) Device on which the code should run.
|
||||||
"""
|
"""
|
||||||
def __init__(self, observation_space: gym.spaces.Space,
|
def __init__(self, observation_space: gym.spaces.Space,
|
||||||
action_space: gym.spaces.Space,
|
action_space: gym.spaces.Space,
|
||||||
|
|
@ -161,10 +170,12 @@ class Critic(BasePolicy):
|
||||||
features_extractor: nn.Module,
|
features_extractor: nn.Module,
|
||||||
features_dim: int,
|
features_dim: int,
|
||||||
activation_fn: Type[nn.Module] = nn.ReLU,
|
activation_fn: Type[nn.Module] = nn.ReLU,
|
||||||
normalize_images: bool = True):
|
normalize_images: bool = True,
|
||||||
|
device: Union[th.device, str] = 'cpu'):
|
||||||
super(Critic, self).__init__(observation_space, action_space,
|
super(Critic, self).__init__(observation_space, action_space,
|
||||||
features_extractor=features_extractor,
|
features_extractor=features_extractor,
|
||||||
normalize_images=normalize_images)
|
normalize_images=normalize_images,
|
||||||
|
device=device)
|
||||||
|
|
||||||
action_dim = get_action_dim(self.action_space)
|
action_dim = get_action_dim(self.action_space)
|
||||||
|
|
||||||
|
|
@ -234,7 +245,8 @@ class SACPolicy(BasePolicy):
|
||||||
'features_dim': self.features_dim,
|
'features_dim': self.features_dim,
|
||||||
'net_arch': self.net_arch,
|
'net_arch': self.net_arch,
|
||||||
'activation_fn': self.activation_fn,
|
'activation_fn': self.activation_fn,
|
||||||
'normalize_images': normalize_images
|
'normalize_images': normalize_images,
|
||||||
|
'device': device
|
||||||
}
|
}
|
||||||
self.actor_kwargs = self.net_args.copy()
|
self.actor_kwargs = self.net_args.copy()
|
||||||
sde_kwargs = {
|
sde_kwargs = {
|
||||||
|
|
@ -268,7 +280,7 @@ class SACPolicy(BasePolicy):
|
||||||
def forward(self, obs: th.Tensor) -> th.Tensor:
|
def forward(self, obs: th.Tensor) -> th.Tensor:
|
||||||
return self.predict(obs, deterministic=False)
|
return self.predict(obs, deterministic=False)
|
||||||
|
|
||||||
def predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
def _predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
||||||
return self.actor(observation, deterministic)
|
return self.actor(observation, deterministic)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -35,6 +35,7 @@ class Actor(BasePolicy):
|
||||||
above zero and prevent it from growing too fast. In practice, ``exp()`` is usually enough.
|
above zero and prevent it from growing too fast. In practice, ``exp()`` is usually enough.
|
||||||
:param normalize_images: (bool) Whether to normalize images or not,
|
:param normalize_images: (bool) Whether to normalize images or not,
|
||||||
dividing by 255.0 (True by default)
|
dividing by 255.0 (True by default)
|
||||||
|
:param device: (Union[th.device, str]) Device on which the code should run.
|
||||||
"""
|
"""
|
||||||
def __init__(self,
|
def __init__(self,
|
||||||
observation_space: gym.spaces.Space,
|
observation_space: gym.spaces.Space,
|
||||||
|
|
@ -50,10 +51,12 @@ class Actor(BasePolicy):
|
||||||
full_std: bool = False,
|
full_std: bool = False,
|
||||||
sde_net_arch: Optional[List[int]] = None,
|
sde_net_arch: Optional[List[int]] = None,
|
||||||
use_expln: bool = False,
|
use_expln: bool = False,
|
||||||
normalize_images: bool = True):
|
normalize_images: bool = True,
|
||||||
|
device: Union[th.device, str] = 'cpu'):
|
||||||
super(Actor, self).__init__(observation_space, action_space,
|
super(Actor, self).__init__(observation_space, action_space,
|
||||||
features_extractor=features_extractor,
|
features_extractor=features_extractor,
|
||||||
normalize_images=normalize_images)
|
normalize_images=normalize_images,
|
||||||
|
device=device)
|
||||||
|
|
||||||
self.latent_pi, self.log_std = None, None
|
self.latent_pi, self.log_std = None, None
|
||||||
self.weights_dist, self.exploration_mat = None, None
|
self.weights_dist, self.exploration_mat = None, None
|
||||||
|
|
@ -104,18 +107,13 @@ class Actor(BasePolicy):
|
||||||
"""
|
"""
|
||||||
return self.action_dist.get_std(self.log_std)
|
return self.action_dist.get_std(self.log_std)
|
||||||
|
|
||||||
def _get_action_dist_from_latent(self, latent_pi: th.Tensor,
|
|
||||||
latent_sde: th.Tensor) -> Tuple[th.Tensor, Distribution]:
|
|
||||||
mean_actions = self.mu(latent_pi)
|
|
||||||
return self.action_dist.proba_distribution(mean_actions, self.log_std, latent_sde)
|
|
||||||
|
|
||||||
def _get_latent(self, obs: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
def _get_latent(self, obs: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
||||||
features = self.extract_features(obs)
|
features = self.extract_features(obs)
|
||||||
latent_pi = self.latent_pi(features)
|
latent_pi = self.latent_pi(features)
|
||||||
latent_sde = self.sde_features_extractor(features) if self.sde_features_extractor is not None else latent_pi
|
latent_sde = self.sde_features_extractor(features) if self.sde_features_extractor is not None else latent_pi
|
||||||
return latent_pi, latent_sde
|
return latent_pi, latent_sde
|
||||||
|
|
||||||
def evaluate_actions(self, obs: th.Tensor, action: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
def evaluate_actions(self, obs: th.Tensor, actions: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
||||||
"""
|
"""
|
||||||
Evaluate actions according to the current policy,
|
Evaluate actions according to the current policy,
|
||||||
given the observations. Only useful when using SDE.
|
given the observations. Only useful when using SDE.
|
||||||
|
|
@ -126,9 +124,9 @@ class Actor(BasePolicy):
|
||||||
and entropy of the action distribution.
|
and entropy of the action distribution.
|
||||||
"""
|
"""
|
||||||
latent_pi, latent_sde = self._get_latent(obs)
|
latent_pi, latent_sde = self._get_latent(obs)
|
||||||
_, distribution = self._get_action_dist_from_latent(latent_pi, latent_sde)
|
mean_actions = self.mu(latent_pi)
|
||||||
log_prob = distribution.log_prob(action)
|
distribution = self.action_dist.proba_distribution(mean_actions, self.log_std, latent_sde)
|
||||||
# value = self.value_net(latent_vf)
|
log_prob = distribution.log_prob(actions)
|
||||||
return log_prob, distribution.entropy()
|
return log_prob, distribution.entropy()
|
||||||
|
|
||||||
def reset_noise(self) -> None:
|
def reset_noise(self) -> None:
|
||||||
|
|
@ -150,12 +148,13 @@ class Actor(BasePolicy):
|
||||||
# -> set squash_output=True in the action_dist?
|
# -> set squash_output=True in the action_dist?
|
||||||
# NOTE: the clipping is done in the rollout for now
|
# NOTE: the clipping is done in the rollout for now
|
||||||
return self.mu(latent_pi) + noise
|
return self.mu(latent_pi) + noise
|
||||||
# action, _ = self._get_action_dist_from_latent(latent_pi)
|
|
||||||
# return action
|
|
||||||
else:
|
else:
|
||||||
features = self.extract_features(obs)
|
features = self.extract_features(obs)
|
||||||
return self.mu(features)
|
return self.mu(features)
|
||||||
|
|
||||||
|
def _predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
||||||
|
return self.forward(observation, deterministic=deterministic)
|
||||||
|
|
||||||
|
|
||||||
class Critic(BasePolicy):
|
class Critic(BasePolicy):
|
||||||
"""
|
"""
|
||||||
|
|
@ -171,6 +170,7 @@ class Critic(BasePolicy):
|
||||||
:param activation_fn: (Type[nn.Module]) Activation function
|
:param activation_fn: (Type[nn.Module]) Activation function
|
||||||
:param normalize_images: (bool) Whether to normalize images or not,
|
:param normalize_images: (bool) Whether to normalize images or not,
|
||||||
dividing by 255.0 (True by default)
|
dividing by 255.0 (True by default)
|
||||||
|
:param device: (Union[th.device, str]) Device on which the code should run.
|
||||||
"""
|
"""
|
||||||
def __init__(self, observation_space: gym.spaces.Space,
|
def __init__(self, observation_space: gym.spaces.Space,
|
||||||
action_space: gym.spaces.Space,
|
action_space: gym.spaces.Space,
|
||||||
|
|
@ -178,10 +178,12 @@ class Critic(BasePolicy):
|
||||||
features_extractor: nn.Module,
|
features_extractor: nn.Module,
|
||||||
features_dim: int,
|
features_dim: int,
|
||||||
activation_fn: Type[nn.Module] = nn.ReLU,
|
activation_fn: Type[nn.Module] = nn.ReLU,
|
||||||
normalize_images: bool = True):
|
normalize_images: bool = True,
|
||||||
|
device: Union[th.device, str] = 'cpu'):
|
||||||
super(Critic, self).__init__(observation_space, action_space,
|
super(Critic, self).__init__(observation_space, action_space,
|
||||||
features_extractor=features_extractor,
|
features_extractor=features_extractor,
|
||||||
normalize_images=normalize_images)
|
normalize_images=normalize_images,
|
||||||
|
device=device)
|
||||||
|
|
||||||
action_dim = get_action_dim(self.action_space)
|
action_dim = get_action_dim(self.action_space)
|
||||||
|
|
||||||
|
|
@ -191,14 +193,14 @@ class Critic(BasePolicy):
|
||||||
q2_net = create_mlp(features_dim + action_dim, 1, net_arch, activation_fn)
|
q2_net = create_mlp(features_dim + action_dim, 1, net_arch, activation_fn)
|
||||||
self.q2_net = nn.Sequential(*q2_net)
|
self.q2_net = nn.Sequential(*q2_net)
|
||||||
|
|
||||||
def forward(self, obs: th.Tensor, action: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
def forward(self, obs: th.Tensor, actions: th.Tensor) -> Tuple[th.Tensor, th.Tensor]:
|
||||||
features = self.extract_features(obs)
|
features = self.extract_features(obs)
|
||||||
qvalue_input = th.cat([features, action], dim=1)
|
qvalue_input = th.cat([features, actions], dim=1)
|
||||||
return self.q1_net(qvalue_input), self.q2_net(qvalue_input)
|
return self.q1_net(qvalue_input), self.q2_net(qvalue_input)
|
||||||
|
|
||||||
def q1_forward(self, obs: th.Tensor, action: th.Tensor) -> th.Tensor:
|
def q1_forward(self, obs: th.Tensor, actions: th.Tensor) -> th.Tensor:
|
||||||
features = self.extract_features(obs)
|
features = self.extract_features(obs)
|
||||||
return self.q1_net(th.cat([features, action], dim=1))
|
return self.q1_net(th.cat([features, actions], dim=1))
|
||||||
|
|
||||||
|
|
||||||
class ValueFunction(BasePolicy):
|
class ValueFunction(BasePolicy):
|
||||||
|
|
@ -245,7 +247,7 @@ class TD3Policy(BasePolicy):
|
||||||
:param action_space: (gym.spaces.Space) Action space
|
:param action_space: (gym.spaces.Space) Action space
|
||||||
:param lr_schedule: (Callable) Learning rate schedule (could be constant)
|
:param lr_schedule: (Callable) Learning rate schedule (could be constant)
|
||||||
:param net_arch: (Optional[List[int]]) The specification of the policy and value networks.
|
:param net_arch: (Optional[List[int]]) The specification of the policy and value networks.
|
||||||
:param device: (str or th.device) Device on which the code should run.
|
:param device: (Union[th.device, str]) Device on which the code should run.
|
||||||
:param activation_fn: (Type[nn.Module]) Activation function
|
:param activation_fn: (Type[nn.Module]) Activation function
|
||||||
:param use_sde: (bool) Whether to use State Dependent Exploration or not
|
:param use_sde: (bool) Whether to use State Dependent Exploration or not
|
||||||
:param log_std_init: (float) Initial value for the log standard deviation
|
:param log_std_init: (float) Initial value for the log standard deviation
|
||||||
|
|
@ -290,7 +292,8 @@ class TD3Policy(BasePolicy):
|
||||||
'features_dim': self.features_dim,
|
'features_dim': self.features_dim,
|
||||||
'net_arch': self.net_arch,
|
'net_arch': self.net_arch,
|
||||||
'activation_fn': self.activation_fn,
|
'activation_fn': self.activation_fn,
|
||||||
'normalize_images': normalize_images
|
'normalize_images': normalize_images,
|
||||||
|
'device': device
|
||||||
}
|
}
|
||||||
self.actor_kwargs = self.net_args.copy()
|
self.actor_kwargs = self.net_args.copy()
|
||||||
sde_kwargs = {
|
sde_kwargs = {
|
||||||
|
|
@ -338,9 +341,9 @@ class TD3Policy(BasePolicy):
|
||||||
return Critic(**self.net_args).to(self.device)
|
return Critic(**self.net_args).to(self.device)
|
||||||
|
|
||||||
def forward(self, observation: th.Tensor, deterministic: bool = False):
|
def forward(self, observation: th.Tensor, deterministic: bool = False):
|
||||||
return self.predict(observation, deterministic=deterministic)
|
return self._predict(observation, deterministic=deterministic)
|
||||||
|
|
||||||
def predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
def _predict(self, observation: th.Tensor, deterministic: bool = False) -> th.Tensor:
|
||||||
return self.actor(observation, deterministic=deterministic)
|
return self.actor(observation, deterministic=deterministic)
|
||||||
|
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -1 +1 @@
|
||||||
0.4.0a2
|
0.4.0a3
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue