已合并
test(reductions): note NPU/XLA skip searchsorted non-contiguous UserWarning check. #37097
test(reductions): note NPU/XLA skip searchsorted non-contiguous UserWarning check. #37097
已合并
Margaret_wangrui创建于 5月29日
1 个文件变更+4-3
Mtest/test_reductions.py+4-3
@@ -1559,16 +1559,17 @@ class TestReductions(TestCase):
1559 self.assertEqual(torch.searchsorted(boundaries, values_nan, right=True), expected_result)1559 self.assertEqual(torch.searchsorted(boundaries, values_nan, right=True), expected_result)
1560 self.assertEqual(torch.searchsorted(boundaries, values_nan, side='right'), expected_result)1560 self.assertEqual(torch.searchsorted(boundaries, values_nan, side='right'), expected_result)
1561 1561 
1562- # type promotion and non contiguous tensors1562+ # type promotion and non-contiguous inputs.
1563+ # CPU/CUDA: expect a non-contiguous UserWarning. XLA/NPU: permute().to() can still be contiguous,
1564+ # so only assertEqual (no warn check).
1563 values_3d_permute = values_3d.permute(2, 1, 0).to(torch.int32)1565 values_3d_permute = values_3d.permute(2, 1, 0).to(torch.int32)
1564 boundaries_permute = values_3d.permute(2, 1, 0).to(torch.float64)1566 boundaries_permute = values_3d.permute(2, 1, 0).to(torch.float64)
1565 expected_result = torch.tensor([[[0, 0], [0, 1]], [[2, 0], [0, 1]], [[2, 0], [0, 0]]], device=device)1567 expected_result = torch.tensor([[[0, 0], [0, 1]], [[2, 0], [0, 1]], [[2, 0], [0, 0]]], device=device)
1566- if self.device_type != 'xla':1568+ if self.device_type != 'xla' and self.device_type != 'npu':
1567 self.assertWarnsRegex(1569 self.assertWarnsRegex(
1568 UserWarning, "tensor is non-contiguous",1570 UserWarning, "tensor is non-contiguous",
1569 lambda: self.assertEqual(torch.searchsorted(boundaries_permute, values_3d_permute), expected_result))1571 lambda: self.assertEqual(torch.searchsorted(boundaries_permute, values_3d_permute), expected_result))
1570 else:1572 else:
1571- # All tensors in XLA is contiguous even doing permute, no warning msg will be generate in XLA
1572 self.assertEqual(torch.searchsorted(boundaries_permute, values_3d_permute), expected_result)1573 self.assertEqual(torch.searchsorted(boundaries_permute, values_3d_permute), expected_result)
1573 1574 
1574 # scalar type1575 # scalar type