From 6f4e85ca6df2b709fb21c83bc76087f37787327e Mon Sep 17 00:00:00 2001 From: AHMETHAKANBEZIR1 Date: Sun, 4 Oct 2026 09:17:25 +0300 Subject: [PATCH] fix: use static input sizes for BlockDiag products Co-authored-by: OpenAI Codex --- CHANGELOG.md | 5 ++++ gpjax/linalg/custom_operators.py | 6 +++-- tests/test_linalg.py | 44 ++++++++++++++++++++++++++++++++ 3 files changed, 53 insertions(+), 2 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 911fd3019..a5eb2bcfe 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 + +- 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 diff --git a/gpjax/linalg/custom_operators.py b/gpjax/linalg/custom_operators.py index 7cbe70dcf..9d558424a 100644 --- a/gpjax/linalg/custom_operators.py +++ b/gpjax/linalg/custom_operators.py @@ -1,5 +1,7 @@ """Custom Lineax operators for GPJax.""" +from itertools import accumulate + import jax import jax.numpy as jnp import lineax as lx @@ -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) diff --git a/tests/test_linalg.py b/tests/test_linalg.py index 7ce0242a6..7a123e4eb 100644 --- a/tests/test_linalg.py +++ b/tests/test_linalg.py @@ -5,6 +5,7 @@ import jax import jax.numpy as jnp import lineax as lx +import numpy as np import pytest # --- cholesky_factor tests --- @@ -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))