Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
5 changes: 5 additions & 0 deletions CHANGELOG.md
Original file line number Diff line number Diff line change
Expand Up @@ -26,6 +26,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
Expand Down
5 changes: 4 additions & 1 deletion gpjax/distributions.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,6 +14,8 @@
# ==============================================================================


import math

from beartype.typing import (
Optional,
)
Expand Down Expand Up @@ -122,7 +124,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((math.prod(sample_shape), self.event_shape[0]))
return vmap(affine_transformation)(flat_noise).reshape(white_noise.shape)

@property
def mean(self) -> Float[Array, " N"]:
Expand Down
31 changes: 31 additions & 0 deletions tests/test_gaussian_distribution.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,8 @@
import jax
import jax.numpy as jnp
import lineax as lx
import numpy as np
import pytest


def _load_distributions():
Expand Down Expand Up @@ -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))
Expand All @@ -69,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]])
Expand Down
Loading