已合并
Update LocalScalarDenseNpu and include ATen/Dispatch_v2.h to support more data types. #27196
Kuteriod创建于 2025年11月28日
Update LocalScalarDenseNpu and include ATen/Dispatch_v2.h to support more data types. #27196
已合并
Kuteriod创建于 2025年11月28日
2 个文件变更+26-29
Mtest/npu/test_tensor.py+22-0
@@ -311,6 +311,28 @@ class TestTensor(TestCase):
311 self.assertTrue(empty_tensor.isnan().all())311 self.assertTrue(empty_tensor.isnan().all())
312 self.assertTrue(empty_strided_tensor.isnan().all())312 self.assertTrue(empty_strided_tensor.isnan().all())
313 313 
314+ def test_print_new_tensor_types(self, device="npu"):
315+ shape = (2, 3)
316+ dtype_factories = {}
317+ 
318+ def _uint_factory(target_dtype):
319+ def _factory():
320+ return torch.randint(low=0, high=1000, size=shape, dtype=target_dtype).npu()
321+ return _factory
322+ 
323+ dtype_factories[torch.uint16] = _uint_factory(torch.uint16)
324+ dtype_factories[torch.uint32] = _uint_factory(torch.uint32)
325+ dtype_factories[torch.uint64] = _uint_factory(torch.uint64)
326+ 
327+ for dtype, factory in dtype_factories.items():
328+ tensor = factory()
329+ scalar = torch.tensor(1, dtype=dtype, device=device)
330+ 
331+ tensor_repr = repr(tensor)
332+ scalar_repr = repr(scalar)
333+ 
334+ self.assertIsInstance(tensor_repr, str)
335+ self.assertIsInstance(scalar_repr, str)
314 336 
315if __name__ == '__main__':337if __name__ == '__main__':
316 run_tests()338 run_tests()
Mtorch_npu/csrc/aten/common/LocalScalarDenseNpu.cpp+4-29
@@ -1,5 +1,6 @@
1#include <ATen/ATen.h>1#include <ATen/ATen.h>
2#include <ATen/NativeFunctions.h>2#include <ATen/NativeFunctions.h>
3+#include <ATen/Dispatch_v2.h>
3 4 
4#include "third_party/acl/inc/acl/acl_base.h"5#include "third_party/acl/inc/acl/acl_base.h"
5#include "third_party/acl/inc/acl/acl_rt.h"6#include "third_party/acl/inc/acl/acl_rt.h"
@@ -10,37 +11,11 @@
10namespace at_npu {11namespace at_npu {
11namespace native {12namespace native {
12 13 
13-#define AT_DISPATCH_CASE_ALL_TYPES_AND5( \
14- SCALARTYPE1, SCALARTYPE2, SCALARTYPE3, SCALARTYPE4, SCALARTYPE5, ...) \
15- AT_DISPATCH_CASE_ALL_TYPES(__VA_ARGS__) \
16- AT_DISPATCH_CASE(SCALARTYPE1, __VA_ARGS__) \
17- AT_DISPATCH_CASE(SCALARTYPE2, __VA_ARGS__) \
18- AT_DISPATCH_CASE(SCALARTYPE3, __VA_ARGS__) \
19- AT_DISPATCH_CASE(SCALARTYPE4, __VA_ARGS__) \
20- AT_DISPATCH_CASE(SCALARTYPE5, __VA_ARGS__)
21- 
22- 
23-#define AT_DISPATCH_ALL_TYPES_AND5( \
24- SCALARTYPE1, SCALARTYPE2, SCALARTYPE3, SCALARTYPE4, SCALARTYPE5, TYPE, NAME, ...) \
25- AT_DISPATCH_SWITCH( \
26- TYPE, \
27- NAME, \
28- AT_DISPATCH_CASE_ALL_TYPES_AND5( \
29- SCALARTYPE1, SCALARTYPE2, SCALARTYPE3, SCALARTYPE4, SCALARTYPE5, __VA_ARGS__))
30- 
31- 
32c10::Scalar NPUNativeFunctions::_local_scalar_dense(const at::Tensor& self)14c10::Scalar NPUNativeFunctions::_local_scalar_dense(const at::Tensor& self)
33{15{
34 c10::Scalar r;16 c10::Scalar r;
35- AT_DISPATCH_ALL_TYPES_AND5(17+ AT_DISPATCH_V2(
36- at::ScalarType::Half,18+ self.scalar_type(), "_local_scalar_dense_npu", AT_WRAP([&] {
37- at::ScalarType::Bool,
38- at::ScalarType::BFloat16,
39- at::ScalarType::Float8_e5m2,
40- at::ScalarType::Float8_e4m3fn,
41- self.scalar_type(),
42- "_local_scalar_dense_npu",
43- [&] {
44 scalar_t value = 0;19 scalar_t value = 0;
45 c10_npu::NPUStream copy_stream = c10_npu::getCurrentNPUStream();20 c10_npu::NPUStream copy_stream = c10_npu::getCurrentNPUStream();
46 // Synchronous copy after stream synchronization21 // Synchronous copy after stream synchronization
@@ -53,7 +28,7 @@ c10::Scalar NPUNativeFunctions::_local_scalar_dense(const at::Tensor& self)
53 sizeof(scalar_t),28 sizeof(scalar_t),
54 ACL_MEMCPY_DEVICE_TO_HOST));29 ACL_MEMCPY_DEVICE_TO_HOST));
55 r = c10::Scalar(value);30 r = c10::Scalar(value);
56- });31+ }), AT_EXPAND(AT_ALL_TYPES_AND_COMPLEX), at::ScalarType::Half, at::ScalarType::Bool, at::ScalarType::BFloat16, AT_EXPAND(AT_FLOAT8_TYPES), AT_EXPAND(AT_BAREBONES_UNSIGNED_TYPES));
57 return r;32 return r;
58}33}
59 34