Compare commits
10
Commits
44eb3335ff
...
master
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
1096dbd848 | ||
|
|
404320c5cc | ||
|
|
4d6ed9b3ac | ||
|
|
7fca6186d5 | ||
|
|
2e0ca977bc | ||
|
|
e83cb9a8a5 | ||
|
|
3e2b988a2f | ||
|
|
de2b9a10d6 | ||
|
|
9fb0014a99 | ||
|
|
8e991ae05b |
@@ -12,7 +12,7 @@ JAX bindings and native implementations of differentiable trust region projectio
|
|||||||
- Multiple projection types:
|
- Multiple projection types:
|
||||||
- KL (Kullback-Leibler divergence)
|
- KL (Kullback-Leibler divergence)
|
||||||
- Wasserstein (only diagonal covariance)
|
- Wasserstein (only diagonal covariance)
|
||||||
- Frobenius (wip, not tested)
|
- Frobenius (wip, problem with cov projections)
|
||||||
- Identity (no projection)
|
- Identity (no projection)
|
||||||
- Support for both diagonal and full covariance Gaussians (induced from cholesky decomposition)
|
- Support for both diagonal and full covariance Gaussians (induced from cholesky decomposition)
|
||||||
- Contextual and non-contextual standard deviations (non-contextual means all standard deviations in batch are expected to be the same)
|
- Contextual and non-contextual standard deviations (non-contextual means all standard deviations in batch are expected to be the same)
|
||||||
@@ -65,9 +65,7 @@ pytest tests/test_projections.py
|
|||||||
*Note*: The test suite verifies:
|
*Note*: The test suite verifies:
|
||||||
|
|
||||||
1. All projections run without errors and maintain basic properties (shapes, positive definiteness)
|
1. All projections run without errors and maintain basic properties (shapes, positive definiteness)
|
||||||
2. KL bounds are actually (approximately) met for:
|
2. KL bounds are actually (approximately) met for true KL projection (both diagonal and full covariance)
|
||||||
- KL projection (both diagonal and full covariance)
|
|
||||||
- Wasserstein projection (diagonal covariance only)
|
|
||||||
3. Gradients can be computed through all projections:
|
3. Gradients can be computed through all projections:
|
||||||
- Both through projection operation and trust region loss
|
- Both through projection operation and trust region loss
|
||||||
- Gradients have correct shapes and are finite
|
- Gradients have correct shapes and are finite
|
||||||
@@ -1,6 +1,8 @@
|
|||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
from typing import Dict
|
from typing import Dict
|
||||||
|
import jax
|
||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
|
from functools import partial
|
||||||
|
|
||||||
class BaseProjection(ABC):
|
class BaseProjection(ABC):
|
||||||
def __init__(self, trust_region_coeff: float = 1.0, mean_bound: float = 0.01,
|
def __init__(self, trust_region_coeff: float = 1.0, mean_bound: float = 0.01,
|
||||||
@@ -8,22 +10,48 @@ class BaseProjection(ABC):
|
|||||||
self.trust_region_coeff = trust_region_coeff
|
self.trust_region_coeff = trust_region_coeff
|
||||||
self.mean_bound = mean_bound
|
self.mean_bound = mean_bound
|
||||||
self.cov_bound = cov_bound
|
self.cov_bound = cov_bound
|
||||||
self.full_cov = full_cov
|
|
||||||
self.contextual_std = contextual_std
|
self.contextual_std = contextual_std
|
||||||
|
self.full_cov = full_cov
|
||||||
|
|
||||||
@abstractmethod
|
@abstractmethod
|
||||||
def project(self, policy_params: Dict[str, jnp.ndarray],
|
def project(self, policy_params: Dict[str, jnp.ndarray],
|
||||||
old_policy_params: Dict[str, jnp.ndarray]) -> Dict[str, jnp.ndarray]:
|
old_policy_params: Dict[str, jnp.ndarray]) -> Dict[str, jnp.ndarray]:
|
||||||
"""Project policy parameters.
|
"""Project parameters to satisfy trust region constraints."""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@abstractmethod
|
||||||
|
def get_trust_region_loss(self, policy_params: Dict[str, jnp.ndarray],
|
||||||
|
proj_policy_params: Dict[str, jnp.ndarray]) -> jnp.ndarray:
|
||||||
|
"""Compute trust region loss between original and projected parameters."""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@partial(jax.jit, static_argnames=('self'))
|
||||||
|
def _mean_projection(self, mean: jnp.ndarray, old_mean: jnp.ndarray,
|
||||||
|
mean_part: jnp.ndarray) -> jnp.ndarray:
|
||||||
|
"""Project mean based on the Mahalanobis objective and trust region.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
policy_params: Dictionary with:
|
mean: Current mean vectors
|
||||||
- 'loc': mean parameters (batch_size, dim)
|
old_mean: Old mean vectors
|
||||||
- 'scale': standard deviations (batch_size, dim) if full_cov=False
|
mean_part: Mahalanobis/Euclidean distance between the two mean vectors
|
||||||
- 'scale_tril': Cholesky factor (batch_size, dim, dim) if full_cov=True
|
|
||||||
old_policy_params: Same format as policy_params
|
Returns:
|
||||||
|
Projected mean that satisfies the trust region
|
||||||
"""
|
"""
|
||||||
pass
|
mask = mean_part > self.mean_bound
|
||||||
|
omega = jnp.ones_like(mean_part)
|
||||||
|
omega = jnp.where(mask, jnp.sqrt(mean_part / self.mean_bound) - 1., omega)
|
||||||
|
omega = jnp.maximum(-omega, omega)[..., None]
|
||||||
|
|
||||||
|
# Use matrix operations instead of boolean indexing
|
||||||
|
m = (mean + omega * old_mean) / (1. + omega + 1e-16)
|
||||||
|
mask_matrix = mask[..., None].astype(mean.dtype)
|
||||||
|
return mask_matrix * m + (1 - mask_matrix) * mean
|
||||||
|
|
||||||
|
def _cov_projection(self, scale_or_tril: jnp.ndarray, old_scale_or_tril: jnp.ndarray,
|
||||||
|
cov_part: jnp.ndarray) -> jnp.ndarray:
|
||||||
|
"""Project covariance parameters."""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
def _calc_covariance(self, params: Dict[str, jnp.ndarray]) -> jnp.ndarray:
|
def _calc_covariance(self, params: Dict[str, jnp.ndarray]) -> jnp.ndarray:
|
||||||
"""Convert scale representation to covariance matrix."""
|
"""Convert scale representation to covariance matrix."""
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
from .base_projection import BaseProjection
|
from .base_projection import BaseProjection
|
||||||
from typing import Dict
|
from typing import Dict
|
||||||
|
import jax
|
||||||
|
from functools import partial
|
||||||
|
|
||||||
class FrobeniusProjection(BaseProjection):
|
class FrobeniusProjection(BaseProjection):
|
||||||
def __init__(self, trust_region_coeff: float = 1.0, mean_bound: float = 0.01,
|
def __init__(self, trust_region_coeff: float = 1.0, mean_bound: float = 0.01,
|
||||||
@@ -46,6 +48,7 @@ class FrobeniusProjection(BaseProjection):
|
|||||||
else:
|
else:
|
||||||
return {"loc": proj_mean, "scale": scale_or_tril}
|
return {"loc": proj_mean, "scale": scale_or_tril}
|
||||||
|
|
||||||
|
@partial(jax.jit, static_argnames=('self'))
|
||||||
def get_trust_region_loss(self, policy_params: Dict[str, jnp.ndarray],
|
def get_trust_region_loss(self, policy_params: Dict[str, jnp.ndarray],
|
||||||
proj_policy_params: Dict[str, jnp.ndarray]) -> jnp.ndarray:
|
proj_policy_params: Dict[str, jnp.ndarray]) -> jnp.ndarray:
|
||||||
mean = policy_params["loc"]
|
mean = policy_params["loc"]
|
||||||
@@ -59,46 +62,53 @@ class FrobeniusProjection(BaseProjection):
|
|||||||
|
|
||||||
return (mean_diff + cov_diff).mean() * self.trust_region_coeff
|
return (mean_diff + cov_diff).mean() * self.trust_region_coeff
|
||||||
|
|
||||||
|
@partial(jax.jit, static_argnames=('self'))
|
||||||
def _gaussian_frobenius(self, p, q):
|
def _gaussian_frobenius(self, p, q):
|
||||||
mean, cov = p
|
mean, cov = p
|
||||||
old_mean, old_cov = q
|
old_mean, old_cov = q
|
||||||
|
|
||||||
if self.scale_prec:
|
if self.scale_prec:
|
||||||
prec_old = jnp.linalg.inv(old_cov)
|
# Mahalanobis distance for mean
|
||||||
mean_part = jnp.sum(jnp.matmul(mean - old_mean, prec_old) * (mean - old_mean), axis=-1)
|
diff = mean - old_mean
|
||||||
cov_part = jnp.sum(prec_old * cov, axis=(-2, -1)) - jnp.log(jnp.linalg.det(jnp.matmul(prec_old, cov))) - mean.shape[-1]
|
if old_cov.ndim == mean.ndim: # diagonal case
|
||||||
|
mean_part = jnp.sum(jnp.square(diff / old_cov), axis=-1)
|
||||||
|
else:
|
||||||
|
solved = jax.scipy.linalg.solve_triangular(
|
||||||
|
old_cov, diff[..., None], lower=True
|
||||||
|
)
|
||||||
|
mean_part = jnp.sum(jnp.square(solved.squeeze(-1)), axis=-1)
|
||||||
else:
|
else:
|
||||||
mean_part = jnp.sum(jnp.square(mean - old_mean), axis=-1)
|
mean_part = jnp.sum(jnp.square(mean - old_mean), axis=-1)
|
||||||
cov_part = jnp.sum(jnp.square(cov - old_cov), axis=(-2, -1))
|
|
||||||
|
# Frobenius norm for covariance
|
||||||
|
if cov.ndim == mean.ndim: # diagonal case
|
||||||
|
diff = old_cov - cov
|
||||||
|
cov_part = jnp.sum(jnp.square(diff), axis=-1)
|
||||||
|
else:
|
||||||
|
diff = jnp.matmul(old_cov, jnp.swapaxes(old_cov, -1, -2)) - \
|
||||||
|
jnp.matmul(cov, jnp.swapaxes(cov, -1, -2))
|
||||||
|
cov_part = jnp.sum(jnp.square(diff), axis=(-2, -1))
|
||||||
|
|
||||||
return mean_part, cov_part
|
return mean_part, cov_part
|
||||||
|
|
||||||
def _mean_projection(self, mean: jnp.ndarray, old_mean: jnp.ndarray,
|
|
||||||
mean_part: jnp.ndarray) -> jnp.ndarray:
|
|
||||||
diff = mean - old_mean
|
|
||||||
norm = jnp.sqrt(mean_part)
|
|
||||||
return jnp.where(norm > self.mean_bound,
|
|
||||||
old_mean + diff * self.mean_bound / norm[..., None],
|
|
||||||
mean)
|
|
||||||
|
|
||||||
def _cov_projection(self, cov: jnp.ndarray, old_cov: jnp.ndarray,
|
def _cov_projection(self, cov: jnp.ndarray, old_cov: jnp.ndarray,
|
||||||
cov_part: jnp.ndarray) -> jnp.ndarray:
|
cov_part: jnp.ndarray) -> jnp.ndarray:
|
||||||
batch_shape = cov.shape[:-2]
|
batch_shape = cov.shape[:-2] if cov.ndim > 2 else cov.shape[:-1]
|
||||||
cov_mask = cov_part > self.cov_bound
|
cov_mask = cov_part > self.cov_bound
|
||||||
|
|
||||||
eta = jnp.ones(batch_shape, dtype=cov.dtype)
|
eta = jnp.ones(batch_shape, dtype=cov.dtype)
|
||||||
eta = jnp.where(cov_mask,
|
eta = jnp.where(cov_mask,
|
||||||
jnp.sqrt(cov_part / self.cov_bound) - 1.,
|
jnp.sqrt(cov_part / self.cov_bound) - 1.,
|
||||||
eta)
|
eta)
|
||||||
eta = jnp.maximum(-eta, eta)
|
eta = jnp.maximum(-eta, eta)
|
||||||
|
|
||||||
if self.full_cov:
|
if self.full_cov:
|
||||||
new_cov = (cov + jnp.einsum('...,...ij->...ij', eta, old_cov)) / (1. + eta + 1e-16)[..., None, None]
|
new_cov = (cov + jnp.einsum('...,...ij->...ij', eta, old_cov)) / \
|
||||||
|
(1. + eta + 1e-16)[..., None, None]
|
||||||
|
mask_matrix = cov_mask[..., None, None].astype(cov.dtype)
|
||||||
|
proj_cov = mask_matrix * new_cov + (1 - mask_matrix) * cov
|
||||||
|
return jnp.linalg.cholesky(proj_cov)
|
||||||
else:
|
else:
|
||||||
# For diagonal case, simple broadcasting
|
|
||||||
new_cov = (cov + eta[..., None] * old_cov) / (1. + eta + 1e-16)[..., None]
|
new_cov = (cov + eta[..., None] * old_cov) / (1. + eta + 1e-16)[..., None]
|
||||||
|
mask_matrix = cov_mask[..., None].astype(cov.dtype)
|
||||||
proj_cov = jnp.where(cov_mask[..., None] if not self.full_cov else cov_mask[..., None, None],
|
return mask_matrix * jnp.sqrt(new_cov) + (1 - mask_matrix) * cov
|
||||||
new_cov, cov)
|
|
||||||
|
|
||||||
return proj_cov
|
|
||||||
@@ -1,6 +1,8 @@
|
|||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
from .base_projection import BaseProjection
|
from .base_projection import BaseProjection
|
||||||
from typing import Dict
|
from typing import Dict
|
||||||
|
import jax
|
||||||
|
from functools import partial
|
||||||
|
|
||||||
class IdentityProjection(BaseProjection):
|
class IdentityProjection(BaseProjection):
|
||||||
def __init__(self, trust_region_coeff: float = 1.0, mean_bound: float = 0.01,
|
def __init__(self, trust_region_coeff: float = 1.0, mean_bound: float = 0.01,
|
||||||
@@ -8,10 +10,12 @@ class IdentityProjection(BaseProjection):
|
|||||||
super().__init__(trust_region_coeff=trust_region_coeff, mean_bound=mean_bound,
|
super().__init__(trust_region_coeff=trust_region_coeff, mean_bound=mean_bound,
|
||||||
cov_bound=cov_bound, contextual_std=contextual_std, full_cov=full_cov)
|
cov_bound=cov_bound, contextual_std=contextual_std, full_cov=full_cov)
|
||||||
|
|
||||||
|
@partial(jax.jit, static_argnames=('self'))
|
||||||
def project(self, policy_params: Dict[str, jnp.ndarray],
|
def project(self, policy_params: Dict[str, jnp.ndarray],
|
||||||
old_policy_params: Dict[str, jnp.ndarray]) -> Dict[str, jnp.ndarray]:
|
old_policy_params: Dict[str, jnp.ndarray]) -> Dict[str, jnp.ndarray]:
|
||||||
return policy_params
|
return policy_params
|
||||||
|
|
||||||
|
@partial(jax.jit, static_argnames=('self'))
|
||||||
def get_trust_region_loss(self, policy_params: Dict[str, jnp.ndarray],
|
def get_trust_region_loss(self, policy_params: Dict[str, jnp.ndarray],
|
||||||
proj_policy_params: Dict[str, jnp.ndarray]) -> jnp.ndarray:
|
proj_policy_params: Dict[str, jnp.ndarray]) -> jnp.ndarray:
|
||||||
return jnp.array(0.0)
|
return jnp.array(0.0)
|
||||||
+89
-61
@@ -13,6 +13,24 @@ from .exception_projection import makeExceptionProjection
|
|||||||
|
|
||||||
MAX_EVAL = 1000
|
MAX_EVAL = 1000
|
||||||
|
|
||||||
|
# Cache for projection operators
|
||||||
|
_diag_proj_op = None
|
||||||
|
_full_proj_op = None
|
||||||
|
|
||||||
|
def _get_diag_proj_op(batch_shape, dim):
|
||||||
|
global _diag_proj_op
|
||||||
|
if _diag_proj_op is None:
|
||||||
|
_diag_proj_op = cpp_projection.BatchedDiagCovOnlyProjection(
|
||||||
|
batch_shape, dim, max_eval=MAX_EVAL)
|
||||||
|
return _diag_proj_op
|
||||||
|
|
||||||
|
def _get_full_proj_op(batch_shape, dim):
|
||||||
|
global _full_proj_op
|
||||||
|
if _full_proj_op is None:
|
||||||
|
_full_proj_op = cpp_projection.BatchedCovOnlyProjection(
|
||||||
|
batch_shape, dim, max_eval=MAX_EVAL)
|
||||||
|
return _full_proj_op
|
||||||
|
|
||||||
class KLProjection(BaseProjection):
|
class KLProjection(BaseProjection):
|
||||||
"""KL divergence-based projection for Gaussian policies.
|
"""KL divergence-based projection for Gaussian policies.
|
||||||
|
|
||||||
@@ -68,13 +86,25 @@ class KLProjection(BaseProjection):
|
|||||||
else:
|
else:
|
||||||
return {"loc": proj_mean, "scale": proj_scale_or_tril}
|
return {"loc": proj_mean, "scale": proj_scale_or_tril}
|
||||||
|
|
||||||
|
@partial(jax.jit, static_argnames=('self'))
|
||||||
def get_trust_region_loss(self, policy_params: Dict[str, jnp.ndarray],
|
def get_trust_region_loss(self, policy_params: Dict[str, jnp.ndarray],
|
||||||
proj_policy_params: Dict[str, jnp.ndarray]) -> jnp.ndarray:
|
proj_policy_params: Dict[str, jnp.ndarray]) -> jnp.ndarray:
|
||||||
mean, scale_or_tril = policy_params["loc"], policy_params["scale"]
|
"""Compute trust region loss between original and projected parameters."""
|
||||||
proj_mean, proj_scale_or_tril = proj_policy_params["loc"], proj_policy_params["scale"]
|
# Get the right scale parameter based on full_cov
|
||||||
|
mean = policy_params["loc"]
|
||||||
|
proj_mean = proj_policy_params["loc"]
|
||||||
|
|
||||||
|
if self.full_cov:
|
||||||
|
scale_or_tril = policy_params["scale_tril"]
|
||||||
|
proj_scale_or_tril = proj_policy_params["scale_tril"]
|
||||||
|
else:
|
||||||
|
scale_or_tril = policy_params["scale"]
|
||||||
|
proj_scale_or_tril = proj_policy_params["scale"]
|
||||||
|
|
||||||
kl = sum(self._gaussian_kl((mean, scale_or_tril), (proj_mean, proj_scale_or_tril)))
|
kl = sum(self._gaussian_kl((mean, scale_or_tril), (proj_mean, proj_scale_or_tril)))
|
||||||
return jnp.mean(kl) * self.trust_region_coeff
|
return jnp.mean(kl) * self.trust_region_coeff
|
||||||
|
|
||||||
|
@partial(jax.jit, static_argnames=('self'))
|
||||||
def _gaussian_kl(self, p: Tuple[jnp.ndarray, jnp.ndarray],
|
def _gaussian_kl(self, p: Tuple[jnp.ndarray, jnp.ndarray],
|
||||||
q: Tuple[jnp.ndarray, jnp.ndarray]) -> Tuple[jnp.ndarray, jnp.ndarray]:
|
q: Tuple[jnp.ndarray, jnp.ndarray]) -> Tuple[jnp.ndarray, jnp.ndarray]:
|
||||||
mean, scale_or_tril = p
|
mean, scale_or_tril = p
|
||||||
@@ -99,6 +129,7 @@ class KLProjection(BaseProjection):
|
|||||||
|
|
||||||
return maha_part, cov_part
|
return maha_part, cov_part
|
||||||
|
|
||||||
|
@partial(jax.jit, static_argnames=('self'))
|
||||||
def _maha(self, x: jnp.ndarray, y: jnp.ndarray, scale_or_tril: jnp.ndarray) -> jnp.ndarray:
|
def _maha(self, x: jnp.ndarray, y: jnp.ndarray, scale_or_tril: jnp.ndarray) -> jnp.ndarray:
|
||||||
diff = x - y
|
diff = x - y
|
||||||
if self.full_cov:
|
if self.full_cov:
|
||||||
@@ -109,21 +140,17 @@ class KLProjection(BaseProjection):
|
|||||||
else:
|
else:
|
||||||
return jnp.sum(jnp.square(diff / scale_or_tril), axis=-1)
|
return jnp.sum(jnp.square(diff / scale_or_tril), axis=-1)
|
||||||
|
|
||||||
|
@partial(jax.jit, static_argnames=('self'))
|
||||||
def _log_determinant(self, scale_or_tril: jnp.ndarray) -> jnp.ndarray:
|
def _log_determinant(self, scale_or_tril: jnp.ndarray) -> jnp.ndarray:
|
||||||
if self.full_cov:
|
if self.full_cov:
|
||||||
return 2 * jnp.sum(jnp.log(jnp.diagonal(scale_or_tril, axis1=-2, axis2=-1)), axis=-1)
|
return 2 * jnp.sum(jnp.log(jnp.diagonal(scale_or_tril, axis1=-2, axis2=-1)), axis=-1)
|
||||||
else:
|
else:
|
||||||
return 2 * jnp.sum(jnp.log(scale_or_tril), axis=-1)
|
return 2 * jnp.sum(jnp.log(scale_or_tril), axis=-1)
|
||||||
|
|
||||||
|
@partial(jax.jit, static_argnames=('self'))
|
||||||
def _batched_trace_square(self, x: jnp.ndarray) -> jnp.ndarray:
|
def _batched_trace_square(self, x: jnp.ndarray) -> jnp.ndarray:
|
||||||
return jnp.sum(x ** 2, axis=(-2, -1))
|
return jnp.sum(x ** 2, axis=(-2, -1))
|
||||||
|
|
||||||
def _mean_projection(self, mean: jnp.ndarray, old_mean: jnp.ndarray,
|
|
||||||
mean_part: jnp.ndarray) -> jnp.ndarray:
|
|
||||||
return old_mean + (mean - old_mean) * jnp.sqrt(
|
|
||||||
self.mean_bound / (mean_part + 1e-8)
|
|
||||||
)[..., None]
|
|
||||||
|
|
||||||
def _cov_projection(self, scale_or_tril: jnp.ndarray, old_scale_or_tril: jnp.ndarray, cov_part: jnp.ndarray) -> jnp.ndarray:
|
def _cov_projection(self, scale_or_tril: jnp.ndarray, old_scale_or_tril: jnp.ndarray, cov_part: jnp.ndarray) -> jnp.ndarray:
|
||||||
if self.full_cov:
|
if self.full_cov:
|
||||||
cov = jnp.matmul(scale_or_tril, jnp.swapaxes(scale_or_tril, -1, -2))
|
cov = jnp.matmul(scale_or_tril, jnp.swapaxes(scale_or_tril, -1, -2))
|
||||||
@@ -133,14 +160,13 @@ class KLProjection(BaseProjection):
|
|||||||
old_cov = old_scale_or_tril ** 2
|
old_cov = old_scale_or_tril ** 2
|
||||||
|
|
||||||
mask = cov_part > self.cov_bound
|
mask = cov_part > self.cov_bound
|
||||||
proj_scale_or_tril = jnp.zeros_like(scale_or_tril)
|
proj_scale_or_tril = scale_or_tril # Start with original scale
|
||||||
proj_scale_or_tril = jnp.where(~mask, scale_or_tril, proj_scale_or_tril)
|
|
||||||
|
|
||||||
if mask.any():
|
if mask.any():
|
||||||
if self.full_cov:
|
if self.full_cov:
|
||||||
proj_cov = project_full_covariance(cov, scale_or_tril, old_scale_or_tril, self.cov_bound)
|
proj_cov = project_full_covariance(cov, scale_or_tril, old_scale_or_tril, self.cov_bound)
|
||||||
is_invalid = jnp.isnan(proj_cov.mean(axis=(-2, -1))) & mask
|
is_invalid = jnp.isnan(proj_cov.mean(axis=(-2, -1)))
|
||||||
proj_scale_or_tril = jnp.where(is_invalid, old_scale_or_tril, proj_scale_or_tril)
|
proj_scale_or_tril = jnp.where(is_invalid[..., None, None], old_scale_or_tril, scale_or_tril)
|
||||||
mask = mask & ~is_invalid
|
mask = mask & ~is_invalid
|
||||||
chol = jnp.linalg.cholesky(proj_cov)
|
chol = jnp.linalg.cholesky(proj_cov)
|
||||||
proj_scale_or_tril = jnp.where(mask[..., None, None], chol, proj_scale_or_tril)
|
proj_scale_or_tril = jnp.where(mask[..., None, None], chol, proj_scale_or_tril)
|
||||||
@@ -148,10 +174,12 @@ class KLProjection(BaseProjection):
|
|||||||
proj_cov = project_diag_covariance(cov, old_cov, self.cov_bound)
|
proj_cov = project_diag_covariance(cov, old_cov, self.cov_bound)
|
||||||
is_invalid = (jnp.isnan(proj_cov.mean(axis=-1)) |
|
is_invalid = (jnp.isnan(proj_cov.mean(axis=-1)) |
|
||||||
jnp.isinf(proj_cov.mean(axis=-1)) |
|
jnp.isinf(proj_cov.mean(axis=-1)) |
|
||||||
(proj_cov.min(axis=-1) < 0)) & mask
|
(proj_cov.min(axis=-1) < 0))
|
||||||
proj_scale_or_tril = jnp.where(is_invalid, old_scale_or_tril, proj_scale_or_tril)
|
proj_scale_or_tril = jnp.where(is_invalid[..., None], old_scale_or_tril, scale_or_tril)
|
||||||
mask = mask & ~is_invalid
|
mask = mask & ~is_invalid
|
||||||
proj_scale_or_tril = jnp.where(mask[..., None], jnp.sqrt(proj_cov), proj_scale_or_tril)
|
proj_scale_or_tril = jnp.where(mask[..., None], jnp.sqrt(proj_cov), scale_or_tril)
|
||||||
|
else:
|
||||||
|
proj_scale_or_tril = scale_or_tril
|
||||||
|
|
||||||
return proj_scale_or_tril
|
return proj_scale_or_tril
|
||||||
|
|
||||||
@@ -166,6 +194,49 @@ class KLProjection(BaseProjection):
|
|||||||
if key not in policy_params or key not in old_policy_params:
|
if key not in policy_params or key not in old_policy_params:
|
||||||
raise KeyError(f"Missing required key '{key}' in policy parameters")
|
raise KeyError(f"Missing required key '{key}' in policy parameters")
|
||||||
|
|
||||||
|
@partial(jax.custom_vjp, nondiff_argnums=(2,))
|
||||||
|
def project_diag_covariance(cov, old_cov, eps_cov):
|
||||||
|
"""JAX wrapper for C++ diagonal covariance projection"""
|
||||||
|
batch_shape = cov.shape[0]
|
||||||
|
dim = cov.shape[-1]
|
||||||
|
|
||||||
|
cov_np = np.asarray(cov)
|
||||||
|
old_cov_np = np.asarray(old_cov)
|
||||||
|
eps = eps_cov * np.ones(batch_shape, dtype=old_cov_np.dtype)
|
||||||
|
|
||||||
|
p_op = _get_diag_proj_op(batch_shape, dim)
|
||||||
|
|
||||||
|
try:
|
||||||
|
proj_cov = p_op.forward(eps, old_cov_np, cov_np)
|
||||||
|
except:
|
||||||
|
proj_cov = cov_np # Return input on failure
|
||||||
|
|
||||||
|
return jnp.array(proj_cov)
|
||||||
|
|
||||||
|
def project_diag_covariance_fwd(cov, old_cov, eps_cov):
|
||||||
|
y = project_diag_covariance(cov, old_cov, eps_cov)
|
||||||
|
return y, (cov, old_cov)
|
||||||
|
|
||||||
|
def project_diag_covariance_bwd(eps_cov, res, g):
|
||||||
|
cov, old_cov = res
|
||||||
|
|
||||||
|
# Convert to numpy for C++ backward pass
|
||||||
|
g_np = np.asarray(g)
|
||||||
|
batch_shape = g_np.shape[0]
|
||||||
|
dim = g_np.shape[-1]
|
||||||
|
|
||||||
|
# Get C++ projection operator
|
||||||
|
p_op = _get_diag_proj_op(batch_shape, dim)
|
||||||
|
|
||||||
|
# Run C++ backward pass
|
||||||
|
grad_cov = p_op.backward(g_np)
|
||||||
|
|
||||||
|
# Convert back to JAX array
|
||||||
|
return jnp.array(grad_cov), None
|
||||||
|
|
||||||
|
# Register VJP rule for diagonal covariance projection
|
||||||
|
project_diag_covariance.defvjp(project_diag_covariance_fwd, project_diag_covariance_bwd)
|
||||||
|
|
||||||
@partial(jax.custom_vjp, nondiff_argnums=(3,))
|
@partial(jax.custom_vjp, nondiff_argnums=(3,))
|
||||||
def project_full_covariance(cov, chol, old_chol, eps_cov):
|
def project_full_covariance(cov, chol, old_chol, eps_cov):
|
||||||
"""JAX wrapper for C++ full covariance projection"""
|
"""JAX wrapper for C++ full covariance projection"""
|
||||||
@@ -179,7 +250,7 @@ def project_full_covariance(cov, chol, old_chol, eps_cov):
|
|||||||
eps = eps_cov * np.ones(batch_shape)
|
eps = eps_cov * np.ones(batch_shape)
|
||||||
|
|
||||||
# Create C++ projection operator directly
|
# Create C++ projection operator directly
|
||||||
p_op = cpp_projection.BatchedCovOnlyProjection(batch_shape, dim, max_eval=MAX_EVAL)
|
p_op = _get_full_proj_op(batch_shape, dim)
|
||||||
|
|
||||||
# Run C++ projection
|
# Run C++ projection
|
||||||
proj_cov = p_op.forward(eps, old_chol_np, chol_np, cov_np)
|
proj_cov = p_op.forward(eps, old_chol_np, chol_np, cov_np)
|
||||||
@@ -203,7 +274,7 @@ def project_full_covariance_bwd(eps_cov, res, g):
|
|||||||
dim = g_np.shape[-1]
|
dim = g_np.shape[-1]
|
||||||
|
|
||||||
# Get C++ projection operator
|
# Get C++ projection operator
|
||||||
p_op = cpp_projection.BatchedCovOnlyProjection(batch_shape, dim, max_eval=MAX_EVAL)
|
p_op = _get_full_proj_op(batch_shape, dim)
|
||||||
|
|
||||||
# Run C++ backward pass
|
# Run C++ backward pass
|
||||||
grad_cov = p_op.backward(g_np)
|
grad_cov = p_op.backward(g_np)
|
||||||
@@ -214,48 +285,5 @@ def project_full_covariance_bwd(eps_cov, res, g):
|
|||||||
# Register VJP rule for full covariance projection
|
# Register VJP rule for full covariance projection
|
||||||
project_full_covariance.defvjp(project_full_covariance_fwd, project_full_covariance_bwd)
|
project_full_covariance.defvjp(project_full_covariance_fwd, project_full_covariance_bwd)
|
||||||
|
|
||||||
@partial(jax.custom_vjp, nondiff_argnums=(2,))
|
|
||||||
def project_diag_covariance(cov, old_cov, eps_cov):
|
|
||||||
"""JAX wrapper for C++ diagonal covariance projection"""
|
|
||||||
# Convert JAX arrays to numpy for C++ function
|
|
||||||
cov_np = np.asarray(cov)
|
|
||||||
old_cov_np = np.asarray(old_cov)
|
|
||||||
batch_shape = cov_np.shape[0]
|
|
||||||
dim = cov_np.shape[-1]
|
|
||||||
eps = eps_cov * np.ones(batch_shape)
|
|
||||||
|
|
||||||
# Create C++ projection operator directly
|
|
||||||
p_op = cpp_projection.BatchedDiagCovOnlyProjection(batch_shape, dim, max_eval=MAX_EVAL)
|
|
||||||
|
|
||||||
# Run C++ projection
|
|
||||||
proj_cov = p_op.forward(eps, old_cov_np, cov_np)
|
|
||||||
|
|
||||||
# Convert back to JAX array
|
|
||||||
return jnp.array(proj_cov)
|
|
||||||
|
|
||||||
def project_diag_covariance_fwd(cov, old_cov, eps_cov):
|
|
||||||
y = project_diag_covariance(cov, old_cov, eps_cov)
|
|
||||||
return y, (cov, old_cov)
|
|
||||||
|
|
||||||
def project_diag_covariance_bwd(eps_cov, res, g):
|
|
||||||
cov, old_cov = res
|
|
||||||
|
|
||||||
# Convert to numpy for C++ backward pass
|
|
||||||
g_np = np.asarray(g)
|
|
||||||
batch_shape = g_np.shape[0]
|
|
||||||
dim = g_np.shape[-1]
|
|
||||||
|
|
||||||
# Get C++ projection operator
|
|
||||||
p_op = cpp_projection.BatchedDiagCovOnlyProjection(batch_shape, dim, max_eval=MAX_EVAL)
|
|
||||||
|
|
||||||
# Run C++ backward pass
|
|
||||||
grad_cov = p_op.backward(g_np)
|
|
||||||
|
|
||||||
# Convert back to JAX array
|
|
||||||
return jnp.array(grad_cov), None
|
|
||||||
|
|
||||||
# Register VJP rule for diagonal covariance projection
|
|
||||||
project_diag_covariance.defvjp(project_diag_covariance_fwd, project_diag_covariance_bwd)
|
|
||||||
|
|
||||||
if not cpp_projection_available:
|
if not cpp_projection_available:
|
||||||
KLProjection = makeExceptionProjection("ITPAL (C++ library) is not available. Please install the C++ library to use this projection.")
|
KLProjection = makeExceptionProjection("ITPAL (C++ library) is not available. Please install the C++ library to use this projection.")
|
||||||
@@ -1,7 +1,10 @@
|
|||||||
import jax.numpy as jnp
|
import jax.numpy as jnp
|
||||||
from .base_projection import BaseProjection
|
from .base_projection import BaseProjection
|
||||||
from typing import Dict, Tuple
|
from typing import Dict, Tuple
|
||||||
|
import jax
|
||||||
|
from functools import partial
|
||||||
|
|
||||||
|
@jax.jit
|
||||||
def scale_tril_to_sqrt(scale_tril: jnp.ndarray) -> jnp.ndarray:
|
def scale_tril_to_sqrt(scale_tril: jnp.ndarray) -> jnp.ndarray:
|
||||||
"""
|
"""
|
||||||
'Converts' scale_tril to scale_sqrt.
|
'Converts' scale_tril to scale_sqrt.
|
||||||
@@ -22,7 +25,8 @@ class WassersteinProjection(BaseProjection):
|
|||||||
|
|
||||||
def project(self, policy_params: Dict[str, jnp.ndarray],
|
def project(self, policy_params: Dict[str, jnp.ndarray],
|
||||||
old_policy_params: Dict[str, jnp.ndarray]) -> Dict[str, jnp.ndarray]:
|
old_policy_params: Dict[str, jnp.ndarray]) -> Dict[str, jnp.ndarray]:
|
||||||
assert not self.full_cov, "Wasserstein projection only supports diagonal covariance"
|
if self.full_cov:
|
||||||
|
print("Warning: Wasserstein projection with full covariance is wip, we recommend using diagonal covariance instead.")
|
||||||
|
|
||||||
mean = policy_params["loc"] # shape: (batch_size, dim)
|
mean = policy_params["loc"] # shape: (batch_size, dim)
|
||||||
old_mean = old_policy_params["loc"]
|
old_mean = old_policy_params["loc"]
|
||||||
@@ -50,6 +54,7 @@ class WassersteinProjection(BaseProjection):
|
|||||||
|
|
||||||
return {"loc": proj_mean, "scale": proj_scale}
|
return {"loc": proj_mean, "scale": proj_scale}
|
||||||
|
|
||||||
|
@partial(jax.jit, static_argnames=('self'))
|
||||||
def get_trust_region_loss(self, policy_params: Dict[str, jnp.ndarray],
|
def get_trust_region_loss(self, policy_params: Dict[str, jnp.ndarray],
|
||||||
proj_policy_params: Dict[str, jnp.ndarray]) -> jnp.ndarray:
|
proj_policy_params: Dict[str, jnp.ndarray]) -> jnp.ndarray:
|
||||||
mean = policy_params["loc"]
|
mean = policy_params["loc"]
|
||||||
@@ -63,37 +68,51 @@ class WassersteinProjection(BaseProjection):
|
|||||||
w2 = mean_part + cov_part
|
w2 = mean_part + cov_part
|
||||||
return w2.mean() * self.trust_region_coeff
|
return w2.mean() * self.trust_region_coeff
|
||||||
|
|
||||||
def _mean_projection(self, mean: jnp.ndarray, old_mean: jnp.ndarray,
|
|
||||||
mean_part: jnp.ndarray) -> jnp.ndarray:
|
|
||||||
diff = mean - old_mean
|
|
||||||
norm = jnp.sqrt(mean_part)
|
|
||||||
return jnp.where(norm > self.mean_bound,
|
|
||||||
old_mean + diff * self.mean_bound / norm[..., None],
|
|
||||||
mean)
|
|
||||||
|
|
||||||
def _scale_projection(self, scale: jnp.ndarray, old_scale: jnp.ndarray,
|
def _scale_projection(self, scale: jnp.ndarray, old_scale: jnp.ndarray,
|
||||||
scale_part: jnp.ndarray) -> jnp.ndarray:
|
scale_part: jnp.ndarray) -> jnp.ndarray:
|
||||||
"""Project scale parameters (standard deviations for diagonal case)"""
|
"""Project scale parameters using multiplicative update.
|
||||||
diff = scale - old_scale
|
|
||||||
norm = jnp.sqrt(scale_part)
|
|
||||||
|
|
||||||
if scale.ndim == 2: # Batched scale
|
Args:
|
||||||
norm = norm[..., None]
|
scale: Current scale/sqrt of covariance
|
||||||
|
old_scale: Previous scale/sqrt of covariance
|
||||||
|
scale_part: W2 distance between scales
|
||||||
|
|
||||||
return jnp.where(norm > self.cov_bound,
|
Returns:
|
||||||
old_scale + diff * self.cov_bound / norm,
|
Projected scale that satisfies the trust region constraint
|
||||||
scale)
|
"""
|
||||||
|
# Check if projection needed
|
||||||
|
cov_mask = scale_part > self.cov_bound
|
||||||
|
|
||||||
|
# Compute eta (multiplier for the update)
|
||||||
|
batch_shape = scale.shape[:-2] if scale.ndim > 2 else scale.shape[:-1]
|
||||||
|
eta = jnp.ones(batch_shape, dtype=scale.dtype)
|
||||||
|
eta = jnp.where(cov_mask,
|
||||||
|
jnp.sqrt(scale_part / self.cov_bound) - 1.,
|
||||||
|
eta)
|
||||||
|
eta = jnp.maximum(-eta, eta)
|
||||||
|
|
||||||
|
# Multiplicative update with matrix operations
|
||||||
|
if scale.ndim > 2: # Full covariance case
|
||||||
|
new_scale = (scale + jnp.einsum('...,...ij->...ij', eta, old_scale)) / \
|
||||||
|
(1. + eta + 1e-16)[..., None, None]
|
||||||
|
mask_matrix = cov_mask[..., None, None].astype(scale.dtype)
|
||||||
|
return mask_matrix * new_scale + (1 - mask_matrix) * scale
|
||||||
|
else: # Diagonal case
|
||||||
|
new_scale = (scale + eta[..., None] * old_scale) / \
|
||||||
|
(1. + eta + 1e-16)[..., None]
|
||||||
|
mask_matrix = cov_mask[..., None].astype(scale.dtype)
|
||||||
|
return mask_matrix * new_scale + (1 - mask_matrix) * scale
|
||||||
|
|
||||||
def _gaussian_wasserstein(self, p, q):
|
@staticmethod
|
||||||
|
@jax.jit
|
||||||
|
def _gaussian_wasserstein(p, q):
|
||||||
mean, scale = p
|
mean, scale = p
|
||||||
mean_other, scale_other = q
|
mean_other, scale_other = q
|
||||||
|
|
||||||
# Keep batch dimension by only summing over feature dimension
|
# Euclidean distance for mean part (we're in diagonal case)
|
||||||
mean_part = jnp.sum(jnp.square(mean - mean_other), axis=-1) # -> (batch_size,)
|
mean_part = jnp.sum(jnp.square(mean - mean_other), axis=-1)
|
||||||
|
|
||||||
if scale.ndim == mean.ndim: # Batched scale
|
# Standard W2 objective for covariance (diagonal case)
|
||||||
cov_part = jnp.sum(scale_other**2 + scale**2 - 2 * scale_other * scale, axis=-1)
|
cov_part = jnp.sum(scale_other**2 + scale**2 - 2 * scale_other * scale, axis=-1)
|
||||||
else: # Non-contextual scale (single scale for all batches)
|
|
||||||
cov_part = jnp.sum(scale_other**2 + scale**2 - 2 * scale_other * scale)
|
|
||||||
|
|
||||||
return mean_part, cov_part
|
return mean_part, cov_part
|
||||||
@@ -0,0 +1,50 @@
|
|||||||
|
import jax
|
||||||
|
import jax.numpy as jnp
|
||||||
|
import time
|
||||||
|
from itpal_jax import FrobeniusProjection
|
||||||
|
|
||||||
|
def generate_params(key, batch_size, dim):
|
||||||
|
keys = jax.random.split(key, 2)
|
||||||
|
return {
|
||||||
|
"loc": jax.random.normal(keys[0], (batch_size, dim)),
|
||||||
|
"scale": jax.nn.softplus(jax.random.normal(keys[1], (batch_size, dim)))
|
||||||
|
}
|
||||||
|
|
||||||
|
def main():
|
||||||
|
# Test parameters
|
||||||
|
batch_size = 32
|
||||||
|
dim = 8
|
||||||
|
n_iterations = 1000
|
||||||
|
|
||||||
|
# Initialize projector
|
||||||
|
proj = FrobeniusProjection(mean_bound=0.1, cov_bound=0.1, contextual_std=True)
|
||||||
|
|
||||||
|
# Compile function
|
||||||
|
proj_fn = lambda p, op: proj.project(p, op)
|
||||||
|
proj_fn = jax.jit(proj_fn)
|
||||||
|
|
||||||
|
# Generate initial key
|
||||||
|
key = jax.random.PRNGKey(0)
|
||||||
|
|
||||||
|
# Warmup
|
||||||
|
for _ in range(10):
|
||||||
|
key, subkey1, subkey2 = jax.random.split(key, 3)
|
||||||
|
params = generate_params(subkey1, batch_size, dim)
|
||||||
|
old_params = generate_params(subkey2, batch_size, dim)
|
||||||
|
proj_fn(params, old_params)
|
||||||
|
|
||||||
|
# Time projections
|
||||||
|
start_time = time.time()
|
||||||
|
for _ in range(n_iterations):
|
||||||
|
key, subkey1, subkey2 = jax.random.split(key, 3)
|
||||||
|
params = generate_params(subkey1, batch_size, dim)
|
||||||
|
old_params = generate_params(subkey2, batch_size, dim)
|
||||||
|
proj_fn(params, old_params)
|
||||||
|
end_time = time.time()
|
||||||
|
|
||||||
|
print(f"Frobenius Projection:")
|
||||||
|
print(f"Average time per projection: {(end_time - start_time) / n_iterations * 1000:.3f} ms")
|
||||||
|
print(f"Total time for {n_iterations} iterations: {end_time - start_time:.3f} s")
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -0,0 +1,49 @@
|
|||||||
|
import jax
|
||||||
|
import jax.numpy as jnp
|
||||||
|
import time
|
||||||
|
from itpal_jax import KLProjection
|
||||||
|
|
||||||
|
def generate_params(key, batch_size, dim):
|
||||||
|
keys = jax.random.split(key, 2)
|
||||||
|
return {
|
||||||
|
"loc": jax.random.normal(keys[0], (batch_size, dim)),
|
||||||
|
"scale": jax.nn.softplus(jax.random.normal(keys[1], (batch_size, dim)))
|
||||||
|
}
|
||||||
|
|
||||||
|
def main():
|
||||||
|
# Test parameters
|
||||||
|
batch_size = 32
|
||||||
|
dim = 8
|
||||||
|
n_iterations = 1000
|
||||||
|
|
||||||
|
# Initialize projector
|
||||||
|
proj = KLProjection(mean_bound=0.1, cov_bound=0.1, contextual_std=True)
|
||||||
|
|
||||||
|
# No JIT for KL projection since it uses C++ backend
|
||||||
|
proj_fn = lambda p, op: proj.project(p, op)
|
||||||
|
|
||||||
|
# Generate initial key
|
||||||
|
key = jax.random.PRNGKey(0)
|
||||||
|
|
||||||
|
# Warmup
|
||||||
|
for _ in range(10):
|
||||||
|
key, subkey1, subkey2 = jax.random.split(key, 3)
|
||||||
|
params = generate_params(subkey1, batch_size, dim)
|
||||||
|
old_params = generate_params(subkey2, batch_size, dim)
|
||||||
|
proj_fn(params, old_params)
|
||||||
|
|
||||||
|
# Time projections
|
||||||
|
start_time = time.time()
|
||||||
|
for _ in range(n_iterations):
|
||||||
|
key, subkey1, subkey2 = jax.random.split(key, 3)
|
||||||
|
params = generate_params(subkey1, batch_size, dim)
|
||||||
|
old_params = generate_params(subkey2, batch_size, dim)
|
||||||
|
proj_fn(params, old_params)
|
||||||
|
end_time = time.time()
|
||||||
|
|
||||||
|
print(f"KL Projection:")
|
||||||
|
print(f"Average time per projection: {(end_time - start_time) / n_iterations * 1000:.3f} ms")
|
||||||
|
print(f"Total time for {n_iterations} iterations: {end_time - start_time:.3f} s")
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -0,0 +1,50 @@
|
|||||||
|
import jax
|
||||||
|
import jax.numpy as jnp
|
||||||
|
import time
|
||||||
|
from itpal_jax import WassersteinProjection
|
||||||
|
|
||||||
|
def generate_params(key, batch_size, dim):
|
||||||
|
keys = jax.random.split(key, 2)
|
||||||
|
return {
|
||||||
|
"loc": jax.random.normal(keys[0], (batch_size, dim)),
|
||||||
|
"scale": jax.nn.softplus(jax.random.normal(keys[1], (batch_size, dim)))
|
||||||
|
}
|
||||||
|
|
||||||
|
def main():
|
||||||
|
# Test parameters
|
||||||
|
batch_size = 32
|
||||||
|
dim = 8
|
||||||
|
n_iterations = 1000
|
||||||
|
|
||||||
|
# Initialize projector
|
||||||
|
proj = WassersteinProjection(mean_bound=0.1, cov_bound=0.1, contextual_std=True)
|
||||||
|
|
||||||
|
# Compile function
|
||||||
|
proj_fn = lambda p, op: proj.project(p, op)
|
||||||
|
proj_fn = jax.jit(proj_fn)
|
||||||
|
|
||||||
|
# Generate initial key
|
||||||
|
key = jax.random.PRNGKey(0)
|
||||||
|
|
||||||
|
# Warmup
|
||||||
|
for _ in range(10):
|
||||||
|
key, subkey1, subkey2 = jax.random.split(key, 3)
|
||||||
|
params = generate_params(subkey1, batch_size, dim)
|
||||||
|
old_params = generate_params(subkey2, batch_size, dim)
|
||||||
|
proj_fn(params, old_params)
|
||||||
|
|
||||||
|
# Time projections
|
||||||
|
start_time = time.time()
|
||||||
|
for _ in range(n_iterations):
|
||||||
|
key, subkey1, subkey2 = jax.random.split(key, 3)
|
||||||
|
params = generate_params(subkey1, batch_size, dim)
|
||||||
|
old_params = generate_params(subkey2, batch_size, dim)
|
||||||
|
proj_fn(params, old_params)
|
||||||
|
end_time = time.time()
|
||||||
|
|
||||||
|
print(f"Wasserstein Projection:")
|
||||||
|
print(f"Average time per projection: {(end_time - start_time) / n_iterations * 1000:.3f} ms")
|
||||||
|
print(f"Total time for {n_iterations} iterations: {end_time - start_time:.3f} s")
|
||||||
|
|
||||||
|
if __name__ == "__main__":
|
||||||
|
main()
|
||||||
@@ -94,8 +94,8 @@ def test_diagonal_projection(ProjectionClass, needs_cpp, gaussian_params):
|
|||||||
assert jnp.all(jnp.isfinite(proj_params["scale"]))
|
assert jnp.all(jnp.isfinite(proj_params["scale"]))
|
||||||
assert jnp.all(proj_params["scale"] > 0)
|
assert jnp.all(proj_params["scale"] > 0)
|
||||||
|
|
||||||
# Only check KL bounds for KL projection (and W2, which should approx hold as well)
|
# Only check KL bounds for KL projection
|
||||||
if ProjectionClass in [KLProjection, WassersteinProjection]:
|
if ProjectionClass in [KLProjection]:
|
||||||
kl = compute_gaussian_kl(proj_params, gaussian_params["old_params"])
|
kl = compute_gaussian_kl(proj_params, gaussian_params["old_params"])
|
||||||
max_kl = (mean_bound + cov_bound) * 1.1 # Allow 10% margin
|
max_kl = (mean_bound + cov_bound) * 1.1 # Allow 10% margin
|
||||||
|
|
||||||
@@ -151,6 +151,10 @@ def test_full_covariance_projection(ProjectionClass):
|
|||||||
eigvals = jnp.linalg.eigvalsh(cov)
|
eigvals = jnp.linalg.eigvalsh(cov)
|
||||||
assert jnp.all(eigvals > 0)
|
assert jnp.all(eigvals > 0)
|
||||||
|
|
||||||
|
# Check trust region loss computation works
|
||||||
|
tr_loss = proj.get_trust_region_loss(params, proj_params)
|
||||||
|
assert jnp.isfinite(tr_loss)
|
||||||
|
|
||||||
# Only check KL bounds for KL projection
|
# Only check KL bounds for KL projection
|
||||||
if ProjectionClass in [KLProjection]:
|
if ProjectionClass in [KLProjection]:
|
||||||
kl = compute_gaussian_kl(proj_params, old_params, full_cov=True)
|
kl = compute_gaussian_kl(proj_params, old_params, full_cov=True)
|
||||||
|
|||||||
Reference in New Issue
Block a user