已合并
test(reductions): note NPU/XLA skip searchsorted non-contiguous UserWarning check. #37097
Margaret_wangrui创建于 5月29日
test(reductions): note NPU/XLA skip searchsorted non-contiguous UserWarning check. #37097
已合并
共 1 个文件变更+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 tensors | 1562 | + # 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 type | 1575 | # scalar type |