Making MultivariateNormal Policies work (and porting Normal to
Independent)
This commit is contained in:
@@ -6,7 +6,7 @@ from ..distributions import UniversalGaussianDistribution, AnyDistribution
|
||||
|
||||
|
||||
def get_mean_and_chol(p: AnyDistribution, expand=False):
|
||||
if isinstance(p, th.distributions.Normal):
|
||||
if isinstance(p, th.distributions.Normal) or isinstance(p, th.distributions.Independent):
|
||||
if expand:
|
||||
return p.mean, th.diag_embed(p.stddev)
|
||||
else:
|
||||
@@ -32,7 +32,7 @@ def get_mean_and_sqrt(p: UniversalGaussianDistribution):
|
||||
|
||||
|
||||
def get_cov(p: AnyDistribution):
|
||||
if isinstance(p, th.distributions.Normal):
|
||||
if isinstance(p, th.distributions.Normal) or isinstance(p, th.distributions.Independent):
|
||||
return th.diag_embed(p.variance)
|
||||
elif isinstance(p, th.distributions.MultivariateNormal):
|
||||
return p.covariance_matrix
|
||||
@@ -45,7 +45,7 @@ def get_cov(p: AnyDistribution):
|
||||
def has_diag_cov(p: AnyDistribution, numerical_check=True):
|
||||
if isinstance(p, SB3_Distribution):
|
||||
return has_diag_cov(p.distribution, numerical_check=numerical_check)
|
||||
if isinstance(p, th.distributions.Normal):
|
||||
if isinstance(p, th.distributions.Normal) or isinstance(p, th.distributions.Independent):
|
||||
return True
|
||||
if not numerical_check:
|
||||
return False
|
||||
@@ -67,11 +67,15 @@ def get_diag_cov_vec(p: AnyDistribution, check_diag=True, numerical_check=True):
|
||||
|
||||
def new_dist_like(orig_p: AnyDistribution, mean: th.Tensor, chol: th.Tensor):
|
||||
if isinstance(orig_p, UniversalGaussianDistribution):
|
||||
return orig_p.new_list_like_me(mean, chol)
|
||||
return orig_p.new_dist_like_me(mean, chol)
|
||||
elif isinstance(orig_p, th.distributions.Normal):
|
||||
if orig_p.stddev.shape != chol.shape:
|
||||
chol = th.diagonal(chol, dim1=1, dim2=2)
|
||||
return th.distributions.Normal(mean, chol)
|
||||
elif isinstance(orig_p, th.distributions.Independent):
|
||||
if orig_p.stddev.shape != chol.shape:
|
||||
chol = th.diagonal(chol, dim1=1, dim2=2)
|
||||
return th.distributions.Independent(th.distributions.Normal(mean, chol), 1)
|
||||
elif isinstance(orig_p, th.distributions.MultivariateNormal):
|
||||
return th.distributions.MultivariateNormal(mean, scale_tril=chol)
|
||||
elif isinstance(orig_p, SB3_Distribution):
|
||||
@@ -79,6 +83,10 @@ def new_dist_like(orig_p: AnyDistribution, mean: th.Tensor, chol: th.Tensor):
|
||||
if isinstance(p, th.distributions.Normal):
|
||||
p_out = orig_p.__class__(orig_p.action_dim)
|
||||
p_out.distribution = th.distributions.Normal(mean, chol)
|
||||
elif isinstance(p, th.distributions.Independent):
|
||||
p_out = orig_p.__class__(orig_p.action_dim)
|
||||
p_out.distribution = th.distributions.Independent(
|
||||
th.distributions.Normal(mean, chol), 1)
|
||||
elif isinstance(p, th.distributions.MultivariateNormal):
|
||||
p_out = orig_p.__class__(orig_p.action_dim)
|
||||
p_out.distribution = th.distributions.MultivariateNormal(
|
||||
|
||||
@@ -1,4 +1,5 @@
|
||||
from typing import Any, Dict, Optional, Type, Union, NamedTuple
|
||||
from more_itertools import distribute
|
||||
|
||||
import numpy as np
|
||||
import torch as th
|
||||
@@ -10,6 +11,8 @@ from stable_baselines3.common.vec_env import VecEnv
|
||||
from stable_baselines3.common.callbacks import BaseCallback
|
||||
from stable_baselines3.common.utils import obs_as_tensor
|
||||
|
||||
from ..misc.distTools import get_mean_and_chol
|
||||
from ..distributions.distributions import Strength, UniversalGaussianDistribution
|
||||
|
||||
# TRL requires the origina mean and covariance from the policy when the datapoint was created.
|
||||
# GaussianRolloutBuffer extends the RolloutBuffer by these two fields
|
||||
@@ -120,6 +123,12 @@ class GaussianRolloutCollectorAuxclass():
|
||||
def _setup_model(self) -> None:
|
||||
super()._setup_model()
|
||||
|
||||
cov_shape = self.action_space.shape
|
||||
|
||||
if isinstance(self.policy.action_dist, UniversalGaussianDistribution):
|
||||
if self.policy.action_dist.cov_strength == Strength.FULL:
|
||||
cov_shape = cov_shape + cov_shape
|
||||
|
||||
self.rollout_buffer = GaussianRolloutBuffer(
|
||||
self.n_steps,
|
||||
self.observation_space,
|
||||
@@ -128,6 +137,7 @@ class GaussianRolloutCollectorAuxclass():
|
||||
gamma=self.gamma,
|
||||
gae_lambda=self.gae_lambda,
|
||||
n_envs=self.n_envs,
|
||||
cov_shape=cov_shape,
|
||||
)
|
||||
|
||||
def collect_rollouts(
|
||||
|
||||
Reference in New Issue
Block a user