Purge history for public version

This commit is contained in:
2022-10-27 21:12:58 +02:00
commit ed9ce2bbad
12 changed files with 2029 additions and 0 deletions
+134
View File
@@ -0,0 +1,134 @@
from typing import Any, Dict, List, Optional, Tuple, Union
import torch as th
from stable_baselines3.common.distributions import Distribution as SB3_Distribution
class UniversalGaussianDistribution(SB3_Distribution):
pass
AnyDistribution = Union[SB3_Distribution, UniversalGaussianDistribution]
def get_mean_and_chol(p: AnyDistribution, expand=False):
if isinstance(p, th.distributions.Normal) or isinstance(p, th.distributions.Independent):
if expand:
return p.mean, th.diag_embed(p.stddev)
else:
return p.mean, p.stddev
elif isinstance(p, th.distributions.MultivariateNormal):
return p.mean, p.scale_tril
elif isinstance(p, SB3_Distribution):
return get_mean_and_chol(p.distribution, expand=expand)
else:
raise Exception('Dist-Type not implemented')
def get_mean_and_sqrt(p: UniversalGaussianDistribution, expand=False):
if not hasattr(p, 'cov_sqrt'):
raise Exception(
'Distribution was not induced from sqrt. On-demand calculation is not supported.')
else:
mean, chol = get_mean_and_chol(p, expand=False)
sqrt_cov = p.cov_sqrt
if mean.shape[0] != sqrt_cov.shape[0]:
shape = list(sqrt_cov.shape)
shape[0] = mean.shape[0]
shape = tuple(shape)
sqrt_cov = sqrt_cov.expand(shape)
if expand and len(sqrt_cov.shape) <= 2:
sqrt_cov = th.diag_embed(sqrt_cov)
return mean, sqrt_cov
def get_cov(p: AnyDistribution):
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
elif isinstance(p, SB3_Distribution):
return get_cov(p.distribution)
else:
raise Exception('Dist-Type not implemented')
def has_diag_cov(p: AnyDistribution, numerical_check=False):
if isinstance(p, SB3_Distribution):
return has_diag_cov(p.distribution, numerical_check=numerical_check)
if isinstance(p, th.distributions.Normal) or isinstance(p, th.distributions.Independent):
return True
if not numerical_check:
return False
# Check if matrix is diag
cov = get_cov(p)
return th.equal(cov - th.diag_embed(th.diagonal(cov, dim1=-2, dim2=-1)), th.zeros_like(cov))
def is_contextual(p: AnyDistribution):
# TODO: Implement for UniveralGaussianDist
return False
def get_diag_cov_vec(p: AnyDistribution, check_diag=True, numerical_check=False):
if check_diag and not has_diag_cov(p, numerical_check=numerical_check):
raise Exception('Cannot reduce cov-mat to diag-vec: Is not diagonal')
return th.diagonal(get_cov(p), dim1=-2, dim2=-1)
def new_dist_like(orig_p: AnyDistribution, mean: th.Tensor, chol: th.Tensor):
if isinstance(orig_p, UniversalGaussianDistribution):
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):
p = orig_p.distribution
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(
mean, scale_tril=chol)
else:
raise Exception('Dist-Type not implemented (of sb3 dist)')
return p_out
else:
raise Exception('Dist-Type not implemented')
def new_dist_like_from_sqrt(orig_p: AnyDistribution, mean: th.Tensor, cov_sqrt: th.Tensor):
chol = _sqrt_to_chol(cov_sqrt, only_diag=has_diag_cov(orig_p))
new = new_dist_like(orig_p, mean, chol)
new.cov_sqrt = cov_sqrt
if hasattr(new, 'distribution'):
new.distribution.cov_sqrt = cov_sqrt
return new
def _sqrt_to_chol(cov_sqrt, only_diag=False):
cov = th.bmm(cov_sqrt.mT, cov_sqrt)
cov += th.eye(cov.shape[-1]).expand(cov.shape)*(1e-6)
chol = th.linalg.cholesky(cov)
if only_diag:
chol = th.diagonal(chol, dim1=-2, dim2=-1)
return chol
+31
View File
@@ -0,0 +1,31 @@
import torch as th
from torch.distributions.multivariate_normal import _batch_mahalanobis
def mahalanobis_alt(u, v, std):
"""
Stolen from Fabian's Code (Public Version)
"""
delta = u - v
return th.triangular_solve(delta, std, upper=False)[0].pow(2).sum([-2, -1])
def mahalanobis(u, v, chol):
delta = u - v
return _batch_mahalanobis(chol, delta)
def frob_sq(diff, is_spd=False):
# If diff is spd, we can use a (probably) more performant algorithm
if is_spd:
return _frob_sq_spd(diff)
return th.norm(diff, p='fro', dim=tuple(range(1, diff.dim()))).pow(2)
def _frob_sq_spd(diff):
return _batch_trace(diff @ diff)
def _batch_trace(x):
return th.diagonal(x, dim1=-2, dim2=-1).sum(-1)
@@ -0,0 +1,5 @@
#TODO: License or such
from .base_projection_layer import BaseProjectionLayer
from .frob_projection_layer import FrobeniusProjectionLayer
from .kl_projection_layer import KLProjectionLayer
from .w2_projection_layer import WassersteinProjectionLayer
@@ -0,0 +1,198 @@
from typing import Any, Dict, Optional, Type, Union, Tuple, final
import torch as th
from stable_baselines3.common.distributions import kl_divergence
from stable_baselines3.common.distributions import Distribution
from ..misc.distTools import *
class BaseProjectionLayer(object):
def __init__(self,
mean_bound: float = 0.03,
cov_bound: float = 1e-3,
trust_region_coeff: float = 1.0,
scale_prec: bool = True,
do_entropy_proj: bool = False,
entropy_eq: bool = False,
entropy_first: bool = False,
):
self.mean_bound = mean_bound
self.cov_bound = cov_bound
self.trust_region_coeff = trust_region_coeff
self.do_entropy_proj = do_entropy_proj
self.entropy_first = scale_prec
self.scale_prec = scale_prec
self.mean_eq = False
self.entropy_first = entropy_first
self.entropy_proj = entropy_equality_projection if entropy_eq else entropy_inequality_projection
def __call__(self, p, q, step, *args, **kwargs):
# TODO: self.entropy_schedule(self.initial_entropy, self.target_entropy, self.temperature, step) * p[0].new_ones(p[0].shape[0])
entropy_bound = 'lol'
return self._projection(p, q, eps=self.mean_bound, eps_cov=self.cov_bound, beta=entropy_bound, **kwargs)
@final
def _projection(self, p, q, eps: th.Tensor, eps_cov: th.Tensor, beta: th.Tensor, **kwargs):
"""
Template method with hook _trust_region_projection() to encode specific functionality.
(Optional) entropy projection is executed before or after as specified by entropy_first.
Do not override this. For Python >= 3.8 you can use the @final decorator to enforce not overwriting.
Args:
policy: policy instance
p: current distribution
q: old distribution
eps: mean trust region bound
eps_cov: covariance trust region bound
beta: entropy bound
**kwargs:
Returns:
projected mean, projected std
"""
####################################################################################################################
# entropy projection in the beginning
if self.do_entropy_proj and self.entropy_first:
p = self.entropy_proj(p, beta)
####################################################################################################################
# trust region projection for mean and cov bounds
new_p = self._trust_region_projection(
p, q, eps, eps_cov, **kwargs)
####################################################################################################################
# entropy projection in the end
if not self.do_entropy_proj or self.entropy_first:
return new_p
return self.entropy_proj(new_p, beta)
def _trust_region_projection(self, p, q, eps: th.Tensor, eps_cov: th.Tensor, **kwargs):
"""
Hook for implementing the specific trust region projection
Args:
p: current distribution
q: old distribution
eps: mean trust region bound
eps_cov: covariance trust region bound
**kwargs:
Returns:
projected
"""
return p
def get_trust_region_loss(self, p, proj_p):
# p:
# predicted distribution from network output
# proj_p:
# projected distribution
proj_mean, proj_chol = get_mean_and_chol(proj_p)
p_target = new_dist_like(p, proj_mean, proj_chol)
kl_diff = self.trust_region_value(p, p_target)
kl_loss = kl_diff.mean()
return kl_loss * self.trust_region_coeff
def trust_region_value(self, p, q):
"""
Computes the KL divergence between two Gaussian distributions p and q_values.
Returns:
full kl divergence
"""
return kl_divergence(p, q)
def entropy_inequality_projection(p: th.distributions.Normal,
beta: Union[float, th.Tensor]):
"""
Stolen and adapted from Fabian's Code (Public Version)
Projects std to satisfy an entropy INEQUALITY constraint.
Args:
p: current distribution
beta: target entropy for EACH std or general bound for all stds
Returns:
projected std that satisfies the entropy bound
"""
mean, std = p.mean, p.stddev
k = std.shape[-1]
batch_shape = std.shape[:-2]
ent = p.entropy()
mask = ent < beta
# if nothing has to be projected skip computation
if (~mask).all():
return p
alpha = th.ones(batch_shape, dtype=std.dtype, device=std.device)
alpha[mask] = th.exp((beta[mask] - ent[mask]) / k)
proj_std = th.einsum('ijk,i->ijk', std, alpha)
new_mean, new_std = mean, th.where(mask[..., None, None], proj_std, std)
return th.distributions.Normal(new_mean, new_std)
def entropy_equality_projection(p: th.distributions.Normal,
beta: Union[float, th.Tensor]):
"""
Stolen and adapted from Fabian's Code (Public Version)
Projects std to satisfy an entropy EQUALITY constraint.
Args:
p: current distribution
beta: target entropy for EACH std or general bound for all stds
Returns:
projected std that satisfies the entropy bound
"""
mean, std = p.mean, p.stddev
k = std.shape[-1]
ent = p.entropy()
alpha = th.exp((beta - ent) / k)
proj_std = th.einsum('ijk,i->ijk', std, alpha)
new_mean, new_std = mean, proj_std
return th.distributions.Normal(new_mean, new_std)
def mean_projection(mean: th.Tensor, old_mean: th.Tensor, maha: th.Tensor, eps: th.Tensor):
"""
Stolen from Fabian's Code (Public Version)
Projects the mean based on the Mahalanobis objective and trust region.
Args:
mean: current mean vectors
old_mean: old mean vectors
maha: Mahalanobis distance between the two mean vectors
eps: trust region bound
Returns:
projected mean that satisfies the trust region
"""
batch_shape = mean.shape[:-1]
mask = maha > eps
################################################################################################################
# mean projection maha
# if nothing has to be projected skip computation
if mask.any():
omega = th.ones(batch_shape, dtype=mean.dtype, device=mean.device)
omega[mask] = th.sqrt(maha[mask] / eps) - 1.
omega = th.max(-omega, omega)[..., None]
m = (mean + omega * old_mean) / (1 + omega + 1e-16)
proj_mean = th.where(mask[..., None], m, mean)
else:
proj_mean = mean
return proj_mean
@@ -0,0 +1,130 @@
import torch as th
from typing import Tuple
from .base_projection_layer import BaseProjectionLayer, mean_projection
from ..misc.norm import mahalanobis, frob_sq
from ..misc.distTools import get_mean_and_chol, get_cov, new_dist_like
class FrobeniusProjectionLayer(BaseProjectionLayer):
def _trust_region_projection(self, p, q, eps: th.Tensor, eps_cov: th.Tensor, **kwargs):
"""
Stolen from Fabian's Code (Public Version)
Runs Frobenius projection layer and constructs cholesky of covariance
Args:
policy: policy instance
p: current distribution
q: old distribution
eps: (modified) kl bound/ kl bound for mean part
eps_cov: (modified) kl bound for cov part
beta: (modified) entropy bound
**kwargs:
Returns: mean, cov cholesky
"""
mean, chol = get_mean_and_chol(p, expand=True)
old_mean, old_chol = get_mean_and_chol(q, expand=True)
batch_shape = mean.shape[:-1]
####################################################################################################################
# precompute mean and cov part of frob projection, which are used for the projection.
mean_part, cov_part, cov, cov_old = gaussian_frobenius(
p, q, self.scale_prec, True)
################################################################################################################
# mean projection maha/euclidean
proj_mean = mean_projection(mean, old_mean, mean_part, eps)
################################################################################################################
# cov projection frobenius
cov_mask = cov_part > eps_cov
if cov_mask.any():
eta = th.ones(batch_shape, dtype=chol.dtype, device=chol.device)
eta[cov_mask] = th.sqrt(cov_part[cov_mask] / eps_cov) - 1.
eta = th.max(-eta, eta)
new_cov = (cov + th.einsum('i,ijk->ijk', eta, cov_old)
) / (1. + eta + 1e-16)[..., None, None]
proj_chol = th.where(
cov_mask[..., None, None], th.linalg.cholesky(new_cov), chol)
else:
proj_chol = chol
proj_p = new_dist_like(p, proj_mean, proj_chol)
return proj_p
def trust_region_value(self, p, q):
"""
Stolen from Fabian's Code (Public Version)
Computes the Frobenius metric between two Gaussian distributions p and q.
Args:
policy: policy instance
p: current distribution
q: old distribution
Returns:
mean and covariance part of Frobenius metric
"""
return gaussian_frobenius(p, q, self.scale_prec)
def get_trust_region_loss(self, p, proj_p):
"""
Stolen from Fabian's Code (Public Version)
"""
mean_diff, _ = self.trust_region_value(p, proj_p)
if False and policy.contextual_std:
# Compute MSE here, because we found the Frobenius norm tends to generate values that explode for the cov
p_mean, proj_p_mean = p.mean, proj_p.mean
cov_diff = (p_mean - proj_p_mean).pow(2).sum([-1, -2])
delta_loss = (mean_diff + cov_diff).mean()
else:
delta_loss = mean_diff.mean()
return delta_loss * self.trust_region_coeff
def gaussian_frobenius(p, q, scale_prec: bool = False, return_cov: bool = False):
"""
Stolen from Fabian' Code (Public Version)
Compute (p - q_values) (L_oL_o^T)^-1 (p - 1)^T + |LL^T - L_oL_o^T|_F^2 with p,q_values ~ N(y, LL^T)
Args:
policy: current policy
p: mean and chol of gaussian p
q: mean and chol of gaussian q_values
return_cov: return cov matrices for further computations
scale_prec: scale objective with precision matrix
Returns: mahalanobis distance, squared frobenius norm
"""
mean, chol = get_mean_and_chol(p)
mean_other, chol_other = get_mean_and_chol(q)
if scale_prec:
# maha objective for mean
mean_part = mahalanobis(mean, mean_other, chol_other)
else:
# euclidean distance for mean
# mean_part = ch.norm(mean_other - mean, ord=2, axis=1) ** 2
mean_part = ((mean_other - mean) ** 2).sum(1)
# frob objective for cov
cov = get_cov(p)
cov_other = get_cov(q)
diff = cov_other - cov
# Matrix is real symmetric PSD, therefore |A @ A^H|^2_F = tr{A @ A^H} = tr{A @ A}
#cov_part = torch_batched_trace(diff @ diff)
cov_part = frob_sq(diff, is_spd=True)
if return_cov:
return mean_part, cov_part, cov, cov_other
return mean_part, cov_part
@@ -0,0 +1,7 @@
from .base_projection_layer import BaseProjectionLayer
class KLProjectionLayer(BaseProjectionLayer):
def __init__(self, *args, **kwargs):
raise Exception(
"KL Projections are not avaible in the public release. You would need to have access to the internal version of ALR's Project ITPAL to use it anyway...")
@@ -0,0 +1,143 @@
import numpy as np
import torch as th
from typing import Tuple, Any
from ..misc.norm import mahalanobis
from .base_projection_layer import BaseProjectionLayer, mean_projection, mean_equality_projection
from ..misc.norm import mahalanobis, _batch_trace
from ..misc.distTools import get_diag_cov_vec, get_mean_and_chol, get_mean_and_sqrt, get_cov, new_dist_like_from_sqrt, has_diag_cov
class WassersteinProjectionLayer(BaseProjectionLayer):
"""
Stolen from Fabian's Code (Public Version)
"""
def _trust_region_projection(self, p, q, eps: th.Tensor, eps_cov: th.Tensor, **kwargs):
"""
Runs commutative Wasserstein projection layer and constructs sqrt of covariance
Args:
policy: policy instance
p: current distribution
q: old distribution
eps: (modified) kl bound/ kl bound for mean part
eps_cov: (modified) kl bound for cov part
**kwargs:
Returns:
mean, cov sqrt
"""
mean, sqrt = get_mean_and_sqrt(p, expand=True)
old_mean, old_sqrt = get_mean_and_sqrt(q, expand=True)
batch_shape = mean.shape[:-1]
####################################################################################################################
# precompute mean and cov part of W2, which are used for the projection.
# Both parts differ based on precision scaling.
# If activated, the mean part is the maha distance and the cov has a more complex term in the inner parenthesis.
mean_part, cov_part = gaussian_wasserstein_commutative(
p, q, self.scale_prec)
####################################################################################################################
# project mean (w/ or w/o precision scaling)
proj_mean = mean_projection(mean, old_mean, mean_part, eps)
####################################################################################################################
# project covariance (w/ or w/o precision scaling)
cov_mask = cov_part > eps_cov
if cov_mask.any():
# gradient issue with ch.where, it executes both paths and gives NaN gradient.
eta = th.ones(batch_shape, dtype=sqrt.dtype, device=sqrt.device)
eta[cov_mask] = th.sqrt(cov_part[cov_mask] / eps_cov) - 1.
eta = th.max(-eta, eta)
new_sqrt = (sqrt + th.einsum('i,ijk->ijk', eta, old_sqrt)
) / (1. + eta + 1e-16)[..., None, None]
proj_sqrt = th.where(cov_mask[..., None, None], new_sqrt, sqrt)
else:
proj_sqrt = sqrt
proj_p = new_dist_like_from_sqrt(p, proj_mean, proj_sqrt)
return proj_p
def trust_region_value(self, p, q):
"""
Computes the Wasserstein distance between two Gaussian distributions p and q.
Args:
policy: policy instance
p: current distribution
q: old distribution
Returns:
mean and covariance part of Wasserstein distance
"""
mean_part, cov_part = gaussian_wasserstein_commutative(
p, q, scale_prec=self.scale_prec)
return mean_part + cov_part
def get_trust_region_loss(self, p, proj_p):
# p:
# predicted distribution from network output
# proj_p:
# projected distribution
proj_mean, proj_sqrt = get_mean_and_sqrt(proj_p)
p_target = new_dist_like_from_sqrt(p, proj_mean, proj_sqrt)
kl_diff = self.trust_region_value(p, p_target)
kl_loss = kl_diff.mean()
return kl_loss * self.trust_region_coeff
def gaussian_wasserstein_commutative(p, q, scale_prec=False) -> Tuple[th.Tensor, th.Tensor]:
"""
Compute mean part and cov part of W_2(p || q_values) with p,q_values ~ N(y, SS).
This version DOES assume commutativity of both distributions, i.e. covariance matrices.
This is less general and assumes both distributions are somewhat close together.
When scale_prec is true scale both distributions with old precision matrix.
Args:
policy: current policy
p: mean and sqrt of gaussian p
q: mean and sqrt of gaussian q_values
scale_prec: scale objective by old precision matrix.
This penalizes directions based on old uncertainty/covariance.
Returns: mean part of W2, cov part of W2
"""
mean, sqrt = get_mean_and_sqrt(p, expand=True)
mean_other, sqrt_other = get_mean_and_sqrt(q, expand=True)
if scale_prec:
# maha objective for mean
mean_part = mahalanobis(mean, mean_other, sqrt_other)
else:
# euclidean distance for mean
# mean_part = ch.norm(mean_other - mean, ord=2, axis=1) ** 2
mean_part = ((mean_other - mean) ** 2).sum(1)
cov = get_cov(p)
if scale_prec and False:
# cov constraint scaled with precision of old dist
batch_dim, dim = mean.shape
identity = th.eye(dim, dtype=sqrt.dtype, device=sqrt.device)
sqrt_inv_other = th.linalg.solve(sqrt_other, identity)
c = sqrt_inv_other @ cov @ sqrt_inv_other
cov_part = _batch_trace(
identity + c - 2 * sqrt_inv_other @ sqrt)
else:
# W2 objective for cov assuming normal W2 objective for mean
cov_other = get_cov(q)
try:
cov_part = _batch_trace(
cov_other + cov - 2 * th.bmm(sqrt_other, sqrt))
except:
import pdb
pdb.set_trace()
return mean_part, cov_part