Skip to content

Commit 7376b8d

Browse files
committed
ENH: test @ and @= operators in matmul tests
1 parent ef5b39f commit 7376b8d

1 file changed

Lines changed: 17 additions & 1 deletion

File tree

‎array_api_tests/test_linalg.py‎

Lines changed: 17 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -400,7 +400,6 @@ def test_inv(x):
400400
def _test_matmul(namespace, x1, x2):
401401
matmul = namespace.matmul
402402

403-
# TODO: Make this also test the @ operator
404403
if (x1.shape == () or x2.shape == ()
405404
or len(x1.shape) == len(x2.shape) == 1 and x1.shape != x2.shape
406405
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):
410409
# libraries will use a custom exception class.
411410
ph.raises(Exception, lambda: xp.matmul(x1, x2),
412411
"matmul did not raise an exception for invalid shapes")
412+
ph.raises(Exception, lambda: x1 @ x2,
413+
"@ did not raise an exception for invalid shapes")
413414
return
414415
else:
415416
res = matmul(x1, x2)
@@ -437,6 +438,21 @@ def _test_matmul(namespace, x1, x2):
437438
expected=stack_shape + (x1.shape[-2], x2.shape[-1]))
438439
_test_stacks(matmul, x1, x2, res=res)
439440

441+
# Test the @ operator against the matmul() result
442+
res_op = x1 @ x2
443+
ph.assert_dtype("@", in_dtype=[x1.dtype, x2.dtype], out_dtype=res_op.dtype)
444+
assert_equal(res, res_op, "@ gives a different result from matmul()")
445+
446+
# Test @= where the result fits into x1 (same shape and dtype). Only
447+
# values are checked: libraries may implement @= as rebinding
448+
# (x1 = x1 @ x2) rather than true in-place mutation (numpy itself
449+
# does not define __imatmul__), and in-place mutation keeps x1's
450+
# dtype while matmul() promotes (e.g. uint8 @= uint16 stays uint8).
451+
if res.shape == x1.shape and res.dtype == x1.dtype:
452+
x1_inplace = xp.asarray(x1, copy=True)
453+
x1_inplace @= x2
454+
assert_equal(res, x1_inplace, "@= gives a different result from matmul()")
455+
440456
@pytest.mark.unvectorized
441457
@pytest.mark.xp_extension('linalg')
442458
@given(

0 commit comments

Comments
 (0)