@@ -400,7 +400,6 @@ def test_inv(x):
400400def _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