Compare commits

...
7 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
10 changed files with 238 additions and 49 deletions
+1 -1
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)
+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."""
+7 -10
View File
@@ -2,6 +2,7 @@ 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 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,
@@ -47,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"]
@@ -60,6 +62,7 @@ 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
@@ -88,14 +91,6 @@ class FrobeniusProjection(BaseProjection):
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] if cov.ndim > 2 else cov.shape[:-1] batch_shape = cov.shape[:-2] if cov.ndim > 2 else cov.shape[:-1]
@@ -110,8 +105,10 @@ class FrobeniusProjection(BaseProjection):
if self.full_cov: if self.full_cov:
new_cov = (cov + jnp.einsum('...,...ij->...ij', eta, old_cov)) / \ new_cov = (cov + jnp.einsum('...,...ij->...ij', eta, old_cov)) / \
(1. + eta + 1e-16)[..., None, None] (1. + eta + 1e-16)[..., None, None]
proj_cov = jnp.where(cov_mask[..., None, None], new_cov, cov) 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) return jnp.linalg.cholesky(proj_cov)
else: else:
new_cov = (cov + eta[..., None] * old_cov) / (1. + eta + 1e-16)[..., None] new_cov = (cov + eta[..., None] * old_cov) / (1. + eta + 1e-16)[..., None]
return jnp.where(cov_mask[..., None], jnp.sqrt(new_cov), cov) mask_matrix = cov_mask[..., None].astype(cov.dtype)
return mask_matrix * jnp.sqrt(new_cov) + (1 - mask_matrix) * 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)
+26 -16
View File
@@ -86,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
@@ -117,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:
@@ -127,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))
@@ -151,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)
@@ -166,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
+11 -14
View File
@@ -2,7 +2,9 @@ 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 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.
@@ -52,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"]
@@ -65,14 +68,6 @@ 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 using multiplicative update. """Project scale parameters using multiplicative update.
@@ -96,19 +91,21 @@ class WassersteinProjection(BaseProjection):
eta) eta)
eta = jnp.maximum(-eta, eta) eta = jnp.maximum(-eta, eta)
# Multiplicative update with correct broadcasting # Multiplicative update with matrix operations
if scale.ndim > 2: # Full covariance case if scale.ndim > 2: # Full covariance case
new_scale = (scale + jnp.einsum('...,...ij->...ij', eta, old_scale)) / \ new_scale = (scale + jnp.einsum('...,...ij->...ij', eta, old_scale)) / \
(1. + eta + 1e-16)[..., None, None] (1. + eta + 1e-16)[..., None, None]
mask = cov_mask[..., None, None] mask_matrix = cov_mask[..., None, None].astype(scale.dtype)
return mask_matrix * new_scale + (1 - mask_matrix) * scale
else: # Diagonal case else: # Diagonal case
new_scale = (scale + eta[..., None] * old_scale) / \ new_scale = (scale + eta[..., None] * old_scale) / \
(1. + eta + 1e-16)[..., None] (1. + eta + 1e-16)[..., None]
mask = cov_mask[..., None] mask_matrix = cov_mask[..., None].astype(scale.dtype)
return mask_matrix * new_scale + (1 - mask_matrix) * scale
return jnp.where(mask, new_scale, 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
+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()
+4
View File
@@ -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)