Compare commits

...
10 Commits
Author SHA1 Message Date
dodox 1096dbd848 Perf tests 2025-01-07 18:24:41 +01:00
dodox 404320c5cc revert kl, cxant kit compile c-binding 2025-01-07 18:23:50 +01:00
dodox 4d6ed9b3ac Better jit (bool mask via matmul) 2025-01-07 16:54:20 +01:00
dodox 7fca6186d5 jit wherever possible 2024-12-21 19:21:24 +01:00
dodox 2e0ca977bc Update README 2024-12-21 18:53:44 +01:00
dodox e83cb9a8a5 Also check loss calc works for full cov case 2024-12-21 18:53:27 +01:00
dodox 3e2b988a2f Fixes for contextual KL 2024-12-21 18:53:11 +01:00
dodox de2b9a10d6 Updated README 2024-12-21 18:31:26 +01:00
dodox 9fb0014a99 Updated tests (no check kl for w2) 2024-12-21 18:31:07 +01:00
dodox 8e991ae05b Fixes 2024-12-21 18:31:01 +01:00
10 changed files with 361 additions and 121 deletions
+2 -4
View File
@@ -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
+36 -8
View File
@@ -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."""
+32 -22
View File
@@ -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
+4
View File
@@ -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)
+88 -60
View File
@@ -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.")
+43 -24
View File
@@ -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
def _gaussian_wasserstein(self, p, q): # 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
@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
+50
View File
@@ -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()
+49
View File
@@ -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()
+50
View File
@@ -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()
+6 -2
View File
@@ -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)