From fdd7a4b3a007d21e28396dbbea55bf13f19d5ba1 Mon Sep 17 00:00:00 2001 From: Sylvester Kaczmarek <16242628+sylvesterkaczmarek@users.noreply.github.com> Date: Thu, 10 Sep 2026 17:55:46 +0100 Subject: [PATCH] Fix concatenation along negative batch axes --- dataclass_array/array_dataclass_test.py | 53 +++++++++++++++++++++---- dataclass_array/ops.py | 8 +++- 2 files changed, 52 insertions(+), 9 deletions(-) diff --git a/dataclass_array/array_dataclass_test.py b/dataclass_array/array_dataclass_test.py index 9d14379..e2136f0 100644 --- a/dataclass_array/array_dataclass_test.py +++ b/dataclass_array/array_dataclass_test.py @@ -748,20 +748,57 @@ class PointDynamicShape(dca.DataclassArray): @enp.testing.parametrize_xnp() -@pytest.mark.parametrize('batch_shape', [(1,), (3,)]) -def test_concatenate(xnp: enp.NpModule, batch_shape: Shape): +@pytest.mark.parametrize( + 'batch_shape,axis', + [ + ((1,), 0), + ((3,), 0), + ((3,), -1), + ((2, 3), 0), + ((2, 3), 1), + ((2, 3), -1), + ((2, 3), -2), + ], +) +def test_concatenate(xnp: enp.NpModule, batch_shape: Shape, axis: int): class TestConcatenateClass(dca.DataclassArray): x: FloatArray['*shape 3'] # pyrefly: ignore[not-a-type] y: FloatArray['*shape'] # pyrefly: ignore[not-a-type] - p = TestConcatenateClass( - x=xnp.zeros(batch_shape + (3,), dtype=xnp.float32), - y=xnp.zeros(batch_shape, dtype=xnp.float32), + x = np.arange(np.prod(batch_shape) * 3, dtype=np.float32).reshape( + batch_shape + (3,) + ) + y = np.arange(np.prod(batch_shape), dtype=np.float32).reshape(batch_shape) + p = TestConcatenateClass(x=xnp.asarray(x), y=xnp.asarray(y)) + p2 = p.replace(x=p.x + 100, y=p.y + 100) + result = dca.concat([p, p2], axis=axis) + batch_axis = axis % len(batch_shape) + expected_shape = list(batch_shape) + expected_shape[batch_axis] *= 2 + assert result.shape == tuple(expected_shape) + assert result.xnp is xnp + np.testing.assert_array_equal( + result.x, np.concatenate([x, x + 100], axis=batch_axis) ) + np.testing.assert_array_equal( + result.y, np.concatenate([y, y + 100], axis=batch_axis) + ) + + +@enp.testing.parametrize_xnp() +@pytest.mark.parametrize('axis', [0, 1, -1, -2]) +def test_concatenate_nested_batch_axes(xnp: enp.NpModule, axis: int): + p = Nested.make((2, 3), xnp) + expected_shape = (4, 3) if axis % 2 == 0 else (2, 6) + result = dca.concat([p, p], axis=axis) + dca.testing.assert_array_equal(result, Nested.make(expected_shape, xnp)) + - p_concatenated = dca.concat([p, p, p]) - assert p_concatenated.x.shape == tuple(x * 3 for x in batch_shape) + (3,) - assert p_concatenated.y.shape == tuple(x * 3 for x in batch_shape) +@pytest.mark.parametrize('axis', [-3, 2, 3]) +def test_concatenate_rejects_non_batch_axes(axis: int): + p = Isometrie.make((2, 3), np) + with pytest.raises(np.exceptions.AxisError): + dca.concat([p, p], axis=axis) def test_class_getitem(): diff --git a/dataclass_array/ops.py b/dataclass_array/ops.py index fb13e22..554cba3 100644 --- a/dataclass_array/ops.py +++ b/dataclass_array/ops.py @@ -30,6 +30,7 @@ def _ops_base( arrays: Iterable[DcT], *, axis: int, + new_axis: bool, array_fn: Callable[ [ enp.NpModule, @@ -73,7 +74,10 @@ def _ops_base( xnp = first_arr.xnp # If axis < 0, normalize the axis such as the last axis is before the inner # shape - axis = np_utils.to_absolute_axis(axis, ndim=first_arr.ndim + 1) # pyrefly: ignore[bad-assignment] + ndim = first_arr.ndim + int(new_axis) + axis = np_utils.to_absolute_axis( + axis, ndim=ndim + ) # pyrefly: ignore[bad-assignment] # Iterating over only the fields of the `first_arr` will skip optional fields # if those are not set in `first_arr`, even if they are present in others. @@ -96,6 +100,7 @@ def stack( return _ops_base( arrays, axis=axis, + new_axis=True, array_fn=lambda xnp, axis, f: xnp.stack( # pylint: disable=g-long-lambda [getattr(arr, f.name) for arr in arrays], axis=axis ), @@ -111,6 +116,7 @@ def concat(arrays: Iterable[DcT], *, axis: int = 0) -> DcT: return _ops_base( arrays, axis=axis, + new_axis=False, array_fn=lambda xnp, axis, f: xnp.concatenate( # pylint: disable=g-long-lambda [getattr(arr, f.name) for arr in arrays], axis=axis ),