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

- Split `BlockDiag` input vectors with static input block sizes, so matrix-vector
products work under JIT and with rectangular blocks.

## [1.0.0] — 2026-09-28

### Added
Expand Down
6 changes: 4 additions & 2 deletions gpjax/linalg/custom_operators.py
Original file line number Diff line number Diff line change
@@ -1,5 +1,7 @@
"""Custom Lineax operators for GPJax."""

from itertools import accumulate

import jax
import jax.numpy as jnp
import lineax as lx
Expand All @@ -14,8 +16,8 @@ def __init__(self, blocks):
self.blocks = tuple(blocks)

def mv(self, x):
sizes = [b.out_structure().shape[0] for b in self.blocks]
splits = jnp.cumsum(jnp.array(sizes[:-1]))
sizes = [b.in_structure().shape[0] for b in self.blocks]
splits = tuple(accumulate(sizes[:-1]))
xs = jnp.split(x, splits)
ys = [b.mv(xi) for b, xi in zip(self.blocks, xs, strict=False)]
return jnp.concatenate(ys)
Expand Down
44 changes: 44 additions & 0 deletions tests/test_linalg.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
import jax
import jax.numpy as jnp
import lineax as lx
import numpy as np
import pytest

# --- cholesky_factor tests ---
Expand Down Expand Up @@ -208,6 +209,49 @@ def test_block_diag_mv():
assert jnp.allclose(result, expected)


@pytest.mark.parametrize(
"block_shapes",
[((2, 2), (3, 3)), ((2, 3), (3, 2)), ((2, 3),), ((1, 2), (3, 1), (2, 4))],
)
@pytest.mark.parametrize("dtype", [jnp.float32, jnp.float64])
@pytest.mark.parametrize("use_jit", [False, True])
def test_block_diag_mv_matches_dense_value_and_gradients(block_shapes, dtype, use_jit):
matrices = tuple(
jnp.arange(rows * cols, dtype=dtype).reshape(rows, cols) + index + 1
for index, (rows, cols) in enumerate(block_shapes)
)
x = jnp.arange(sum(cols for _, cols in block_shapes), dtype=dtype) + 0.5
weights = jnp.arange(sum(rows for rows, _ in block_shapes), dtype=dtype) + 1

def apply(blocks, vector):
operator = BlockDiag(tuple(lx.MatrixLinearOperator(block) for block in blocks))
return operator.mv(vector)

def loss(blocks, vector):
return jnp.dot(weights, apply(blocks, vector))

def reference_loss(blocks, vector):
return jnp.dot(weights, jax.scipy.linalg.block_diag(*blocks) @ vector)

evaluate = jax.jit(apply) if use_jit else apply
result = evaluate(matrices, x)
expected = jax.scipy.linalg.block_diag(*matrices) @ x
assert result.shape == expected.shape
np.testing.assert_allclose(result, expected, rtol=1e-6, atol=1e-6)

differentiate = jax.grad(loss, argnums=(0, 1))
if use_jit:
differentiate = jax.jit(differentiate)
actual_gradients = differentiate(matrices, x)
expected_gradients = jax.grad(reference_loss, argnums=(0, 1))(matrices, x)
for actual, expected in zip(
jax.tree.leaves(actual_gradients),
jax.tree.leaves(expected_gradients),
strict=True,
):
np.testing.assert_allclose(actual, expected, rtol=1e-6, atol=1e-6)


def test_block_diag_as_matrix():
A = lx.MatrixLinearOperator(jnp.eye(2))
B = lx.MatrixLinearOperator(2.0 * jnp.eye(3))
Expand Down
Loading