From 8c9ec8614be830b353aa744bbfdc8835d007fc10 Mon Sep 17 00:00:00 2001 From: haroune-dev Date: Mon, 14 Sep 2026 15:44:28 +0100 Subject: [PATCH 1/2] ENH: test @ and @= operators in matmul tests --- array_api_tests/test_linalg.py | 15 ++++++++++++++- 1 file changed, 14 insertions(+), 1 deletion(-) diff --git a/array_api_tests/test_linalg.py b/array_api_tests/test_linalg.py index 87a7652f..a652fca3 100644 --- a/array_api_tests/test_linalg.py +++ b/array_api_tests/test_linalg.py @@ -400,7 +400,6 @@ def test_inv(x): def _test_matmul(namespace, x1, x2): matmul = namespace.matmul - # TODO: Make this also test the @ operator if (x1.shape == () or x2.shape == () or len(x1.shape) == len(x2.shape) == 1 and x1.shape != x2.shape or len(x1.shape) == 1 and len(x2.shape) >= 2 and x1.shape[0] != x2.shape[-2] @@ -410,6 +409,8 @@ def _test_matmul(namespace, x1, x2): # libraries will use a custom exception class. ph.raises(Exception, lambda: xp.matmul(x1, x2), "matmul did not raise an exception for invalid shapes") + ph.raises(Exception, lambda: x1 @ x2, + "@ did not raise an exception for invalid shapes") return else: res = matmul(x1, x2) @@ -437,6 +438,18 @@ def _test_matmul(namespace, x1, x2): expected=stack_shape + (x1.shape[-2], x2.shape[-1])) _test_stacks(matmul, x1, x2, res=res) + # Test @ matches matmul() + res_op = x1 @ x2 + ph.assert_dtype("@", in_dtype=[x1.dtype, x2.dtype], out_dtype=res_op.dtype) + assert_equal(res, res_op, "@ gives a different result from matmul()") + + # Test @= where result fits in x1; compare values only since + # in-place keeps x1's dtype and may fall back to x1 = x1 @ x2. + if res.shape == x1.shape and res.dtype == x1.dtype: + x1_inplace = xp.asarray(x1, copy=True) + x1_inplace @= x2 + assert_equal(res, x1_inplace, "@= gives a different result from matmul()") + @pytest.mark.unvectorized @pytest.mark.xp_extension('linalg') @given( From d7e00218d8a97707e9bba746053e0237f0cc79d5 Mon Sep 17 00:00:00 2001 From: haroune-dev Date: Sun, 27 Sep 2026 11:01:49 +0100 Subject: [PATCH 2/2] test: account for torch matmul dtype limitation --- array_api_tests/test_linalg.py | 29 ++++++++++++++++------------- 1 file changed, 16 insertions(+), 13 deletions(-) diff --git a/array_api_tests/test_linalg.py b/array_api_tests/test_linalg.py index a652fca3..6e9baf6e 100644 --- a/array_api_tests/test_linalg.py +++ b/array_api_tests/test_linalg.py @@ -409,8 +409,9 @@ def _test_matmul(namespace, x1, x2): # libraries will use a custom exception class. ph.raises(Exception, lambda: xp.matmul(x1, x2), "matmul did not raise an exception for invalid shapes") - ph.raises(Exception, lambda: x1 @ x2, - "@ did not raise an exception for invalid shapes") + if x1.dtype == x2.dtype: + ph.raises(Exception, lambda: x1 @ x2, + "@ did not raise an exception for invalid shapes") return else: res = matmul(x1, x2) @@ -438,17 +439,19 @@ def _test_matmul(namespace, x1, x2): expected=stack_shape + (x1.shape[-2], x2.shape[-1])) _test_stacks(matmul, x1, x2, res=res) - # Test @ matches matmul() - res_op = x1 @ x2 - ph.assert_dtype("@", in_dtype=[x1.dtype, x2.dtype], out_dtype=res_op.dtype) - assert_equal(res, res_op, "@ gives a different result from matmul()") - - # Test @= where result fits in x1; compare values only since - # in-place keeps x1's dtype and may fall back to x1 = x1 @ x2. - if res.shape == x1.shape and res.dtype == x1.dtype: - x1_inplace = xp.asarray(x1, copy=True) - x1_inplace @= x2 - assert_equal(res, x1_inplace, "@= gives a different result from matmul()") + # @ bypasses the compat wrapper, so mixed-dtype promotion cannot be + # tested portably here. See data-apis/array-api-compat#245. + if x1.dtype == x2.dtype: + res_op = x1 @ x2 + ph.assert_dtype("@", in_dtype=[x1.dtype, x2.dtype], out_dtype=res_op.dtype) + assert_equal(res, res_op, "@ gives a different result from matmul()") + + # Test @= where result fits in x1; compare values only since + # in-place keeps x1's dtype and may fall back to x1 = x1 @ x2. + if res.shape == x1.shape and res.dtype == x1.dtype: + x1_inplace = xp.asarray(x1, copy=True) + x1_inplace @= x2 + assert_equal(res, x1_inplace, "@= gives a different result from matmul()") @pytest.mark.unvectorized @pytest.mark.xp_extension('linalg')