已合并
[fix] keep npu tensor strides for empty_like #31959
[fix] keep npu tensor strides for empty_like #31959
已合并
zhangqiongwen创建于 3月18日
共 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 + 1597 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