已合并
[fix] keep npu tensor strides for empty_like #31959
zhangqiongwen创建于 3月18日
[fix] keep npu tensor strides for empty_like #31959
已合并
共 2 个文件变更+44-1
| @@ -596,6 +596,41 @@ class TestNpu(TestCase): | |||
| 596 | x_like = torch.empty_like(x) | 596 | x_like = torch.empty_like(x) |
| 597 | res = x_like + 1 | 597 | res = x_like + 1 |
| 598 | 598 | ||
| 599 | + def test_function_torch_empty_like_with_stride(self): | ||
| 600 | + # if a is contiguous, stride of b should be same as a | ||
| 601 | + a = torch.empty([16, 32], device='npu') | ||
| 602 | + self.assertTrue(a.is_contiguous()) | ||
| 603 | + self.assertEqual(a.stride(), (32, 1)) | ||
| 604 | + | ||
| 605 | + b = torch.empty_like(a) | ||
| 606 | + self.assertTrue(b.is_contiguous()) | ||
| 607 | + self.assertEqual(b.stride(), (32, 1)) | ||
| 608 | + | ||
| 609 | + b = torch.empty_like(a, memory_format=torch.preserve_format) | ||
| 610 | + self.assertTrue(b.is_contiguous()) | ||
| 611 | + self.assertEqual(b.stride(), (32, 1)) | ||
| 612 | + | ||
| 613 | + b = torch.empty_like(a, memory_format=torch.contiguous_format) | ||
| 614 | + self.assertTrue(b.is_contiguous()) | ||
| 615 | + self.assertEqual(b.stride(), (32, 1)) | ||
| 616 | + | ||
| 617 | + # if a is Not contiguous, stride of b should be same as a when memory_format=torch.preserve_format(default) | ||
| 618 | + a = torch.empty([16, 32], device='npu').T | ||
| 619 | + self.assertFalse(a.is_contiguous()) | ||
| 620 | + self.assertEqual(a.stride(), (1, 32)) | ||
| 621 | + | ||
| 622 | + b = torch.empty_like(a) | ||
| 623 | + self.assertFalse(b.is_contiguous()) | ||
| 624 | + self.assertEqual(b.stride(), (1, 32)) | ||
| 625 | + | ||
| 626 | + b = torch.empty_like(a, memory_format=torch.preserve_format) | ||
| 627 | + self.assertFalse(b.is_contiguous()) | ||
| 628 | + self.assertEqual(b.stride(), (1, 32)) | ||
| 629 | + | ||
| 630 | + b = torch.empty_like(a, memory_format=torch.contiguous_format) | ||
| 631 | + self.assertTrue(b.is_contiguous()) | ||
| 632 | + self.assertEqual(b.stride(), (16, 1)) | ||
| 633 | + | ||
| 599 | def test_function_torch_empty_like_in_fake_tensor_mode(self): | 634 | def test_function_torch_empty_like_in_fake_tensor_mode(self): |
| 600 | with torch._subclasses.fake_tensor.FakeTensorMode(): | 635 | with torch._subclasses.fake_tensor.FakeTensorMode(): |
| 601 | x = torch.rand(3, 3).npu() | 636 | x = torch.rand(3, 3).npu() |
| @@ -237,7 +237,15 @@ at::Tensor empty_like_npu( | |||
| 237 | if ((typeid(*self.storage().unsafeGetStorageImpl()) != typeid(torch_npu::NPUStorageImpl))) { | 237 | if ((typeid(*self.storage().unsafeGetStorageImpl()) != typeid(torch_npu::NPUStorageImpl))) { |
| 238 | npu_format = ACL_FORMAT_ND; | 238 | npu_format = ACL_FORMAT_ND; |
| 239 | } | 239 | } |
| 240 | - result = OpPreparation::ApplyTensorWithFormat(self.sizes(), options, npu_format); | 240 | + if (FormatHelper::IsBaseFormatType(npu_format) && self.unsafeGetTensorImpl()->support_as_strided() && |
| 241 | + self.layout() == c10::kStrided && | ||
| 242 | + (!optional_memory_format.has_value() || optional_memory_format.value() == c10::MemoryFormat::Preserve)) { | ||
| 243 | + // keep strides | ||
| 244 | + std::vector<int64_t> strides = at::infer_dense_strides(self.sizes(), self.strides()); | ||
| 245 | + result = at::empty_strided(self.sizes(), strides, options.memory_format(std::nullopt)); | ||
| 246 | + } else { | ||
| 247 | + result = OpPreparation::ApplyTensorWithFormat(self.sizes(), options, npu_format); | ||
| 248 | + } | ||
| 241 | } | 249 | } |
| 242 | } | 250 | } |
| 243 | 251 | ||