From 1264bdec00e5551ea76cdcee91d497a9ca28b817 Mon Sep 17 00:00:00 2001 From: AHMETHAKANBEZIR1 Date: Sun, 4 Oct 2026 08:43:42 +0300 Subject: [PATCH 1/2] fix: support multidimensional Gaussian sample shapes Co-authored-by: OpenAI Codex --- CHANGELOG.md | 5 +++++ gpjax/distributions.py | 3 ++- tests/test_gaussian_distribution.py | 22 ++++++++++++++++++++++ 3 files changed, 29 insertions(+), 1 deletion(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 911fd3019..402cdd6e7 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -13,6 +13,11 @@ and this project adheres to [Semantic Versioning](https://semver.org/spec/v2.0.0 - Allow JAX and JAXlib 0.11 in downstream environments by removing the `<0.11` dependency bounds ([#801](https://github.com/QuantClimate/GPJax/issues/801)). +### Fixed + +- Apply the Gaussian covariance transform to each sample for multidimensional + `GaussianDistribution.sample` shapes, including shapes with empty axes. + ## [1.0.0] — 2026-09-28 ### Added diff --git a/gpjax/distributions.py b/gpjax/distributions.py index a69550b8f..1e169f25c 100644 --- a/gpjax/distributions.py +++ b/gpjax/distributions.py @@ -122,7 +122,8 @@ def affine_transformation(_x): if not sample_shape: return affine_transformation(white_noise) - return vmap(affine_transformation)(white_noise) + flat_noise = white_noise.reshape((-1, self.event_shape[0])) + return vmap(affine_transformation)(flat_noise).reshape(white_noise.shape) @property def mean(self) -> Float[Array, " N"]: diff --git a/tests/test_gaussian_distribution.py b/tests/test_gaussian_distribution.py index b1ec1811f..a4f892aec 100644 --- a/tests/test_gaussian_distribution.py +++ b/tests/test_gaussian_distribution.py @@ -6,6 +6,8 @@ import jax import jax.numpy as jnp import lineax as lx +import numpy as np +import pytest def _load_distributions(): @@ -60,6 +62,26 @@ def test_sample_shape(): assert samples.shape == (10, 2) +@pytest.mark.parametrize( + "sample_shape", [(), (3,), (2, 3), (2, 2), (2, 1, 3), (0, 3), (2, 0)] +) +@pytest.mark.parametrize("dtype", [jnp.float32, jnp.float64]) +def test_sample_matches_affine_normal_for_all_sample_axes(sample_shape, dtype): + mu = jnp.array([1.0, -2.0], dtype=dtype) + covariance = jnp.array([[2.0, 0.5], [0.5, 1.0]], dtype=dtype) + distribution = GaussianDistribution( + loc=mu, scale=lx.MatrixLinearOperator(covariance) + ) + key = jax.random.key(17) + white_noise = jax.random.normal(key, shape=(*sample_shape, 2)) + expected = mu + white_noise @ jnp.linalg.cholesky(covariance).T + + for sample in [distribution.sample, jax.jit(distribution.sample, static_argnums=1)]: + actual = sample(key, sample_shape) + assert actual.shape == (*sample_shape, 2) + np.testing.assert_allclose(actual, expected, rtol=1e-6, atol=1e-6) + + def test_log_prob_standard_normal(): mu = jnp.zeros(2) cov = lx.MatrixLinearOperator(jnp.eye(2)) From 8718c295cf000e9fe5cdc424c7f2a47163771dc6 Mon Sep 17 00:00:00 2001 From: AHMETHAKANBEZIR1 Date: Sun, 4 Oct 2026 08:55:42 +0300 Subject: [PATCH 2/2] fix: preserve empty Gaussian event dimensions Co-authored-by: OpenAI Codex --- gpjax/distributions.py | 4 +++- tests/test_gaussian_distribution.py | 9 +++++++++ 2 files changed, 12 insertions(+), 1 deletion(-) diff --git a/gpjax/distributions.py b/gpjax/distributions.py index 1e169f25c..d19258669 100644 --- a/gpjax/distributions.py +++ b/gpjax/distributions.py @@ -14,6 +14,8 @@ # ============================================================================== +import math + from beartype.typing import ( Optional, ) @@ -122,7 +124,7 @@ def affine_transformation(_x): if not sample_shape: return affine_transformation(white_noise) - flat_noise = white_noise.reshape((-1, self.event_shape[0])) + flat_noise = white_noise.reshape((math.prod(sample_shape), self.event_shape[0])) return vmap(affine_transformation)(flat_noise).reshape(white_noise.shape) @property diff --git a/tests/test_gaussian_distribution.py b/tests/test_gaussian_distribution.py index a4f892aec..9fcf8a34b 100644 --- a/tests/test_gaussian_distribution.py +++ b/tests/test_gaussian_distribution.py @@ -91,6 +91,15 @@ def test_log_prob_standard_normal(): assert jnp.allclose(lp, expected, atol=1e-5) +@pytest.mark.parametrize("sample_shape", [(), (3,), (2, 3), (0,)]) +def test_sample_zero_dimensional_event(sample_shape): + distribution = GaussianDistribution( + loc=jnp.zeros(0), scale=lx.MatrixLinearOperator(jnp.eye(0)) + ) + for sample in [distribution.sample, jax.jit(distribution.sample, static_argnums=1)]: + assert sample(jax.random.key(0), sample_shape).shape == (*sample_shape, 0) + + def test_covariance_returns_dense(): mu = jnp.zeros(2) A = jnp.array([[2.0, 1.0], [1.0, 3.0]])