已合并
support_list_for_register_meta_npu #39250
cuiduo创建于 6月25日
support_list_for_register_meta_npu #39250
已合并
cuiduo创建于 6月25日
4 个文件变更+59-8
@@ -0,0 +1,33 @@
1+import torch
atomgit-bot
atomgit-botatomgit-bot6月25日

🔵 Low Priority

测试文件名 test_nanquantile_defaul.py 中 "default" 拼写为 "defaul"(缺少末尾字母 't')。该测试旨在覆盖 torch.ops.aten.nanquantile.default,正确文件名应为 test_nanquantile_default.py。文件名拼写错误可能导致他人按正确名称搜索时找不到该文件。

建议:将文件重命名为 test_nanquantile_default.py(补全缺失的 't')。

likedislike
2+from torch.testing._internal.common_utils import run_tests, parametrize, instantiate_parametrized_tests
3+from testutils import TestUtils
4+import torch_npu
5+import torch_npu._inductor
6+from torch._inductor.utils import run_and_get_code
7+ 
8+ 
9+class TestNanquantileDefault(TestUtils):
10+ 
11+ @parametrize("dynamic", [False])
12+ @parametrize('shape', [(3, 4)])
13+ @parametrize('dtype', ['float32'])
14+ def test_nanquantile(self, shape, dtype, dynamic):
15+ a = torch.randn(shape, dtype=torch.float32, device="npu")
atomgit-bot
atomgit-botatomgit-bot6月25日

🔵 Low Priority

测试函数 test_nanquantiledtype 参数通过 @parametrize('dtype', ['float32']) 进行参数化注入(第 14 行),但在第 16 行的函数体内,张量创建使用了硬编码的 dtype=torch.float32,完全忽略了传入的 dtype 参数值。

这意味着:

  1. dtype 参数化形同虚设——无论传什么 dtype,测试都只测 float32
  2. 如果将来扩展 dtype 参数值(如添加 'float16'),测试不会覆盖新 dtype,导致遗漏

对比同项目其他测试文件(如 test_nanquantile_defaul.py 所在目录的其他测试),它们使用 self._generate_tensor(shape, dtype) 来正确使用 dtype 参数。

建议:将 dtype=torch.float32 改为使用传入的 dtype 参数,例如:torch.randn(shape, dtype=getattr(torch, dtype), device="npu"),或者参照其他测试使用 self._generate_tensor(shape, dtype)

likedislike
16+ a[0, 0] = float('nan')
17+ q = torch.tensor([0.25, 0.5, 0.75], dtype=torch.float32, device="npu")
18+ 
19+ def fn(a, q):
20+ return torch.ops.aten.nanquantile.default(a, q, dim=0, keepdim=False)
21+ 
22+ r1 = fn(a, q)
23+ func = torch.compile(fn, backend="inductor", dynamic=dynamic)
24+ r, codes = run_and_get_code(func, a, q)
25+ self.assertEqual(r, r1, atol=1e-3, rtol=1e-3)
26+ self.assertTrue('nanquantile' in codes[0])
27+ 
28+ 
29+instantiate_parametrized_tests(TestNanquantileDefault)
30+ 
31+ 
32+if __name__ == "__main__":
33+ run_tests()
@@ -68,7 +68,9 @@ def _load_triton_backend():
68 disable_foreach,68 disable_foreach,
69 patch_fx_node_is_input_dependent_cudagraph_unsafe,69 patch_fx_node_is_input_dependent_cudagraph_unsafe,
70 )70 )
71+ from ._npu_meta_registration import npu_patch_meta
71 72 
73+ npu_patch_meta()
72 def _inductor_register_backend_for_device():74 def _inductor_register_backend_for_device():
73 from .codegen.cpp_wrapper import CppWrapperNpu75 from .codegen.cpp_wrapper import CppWrapperNpu
74 from .codegen.scheduling import NPUTritonScheduling76 from .codegen.scheduling import NPUTritonScheduling
Rtorch_npu/utils/_npu_meta_registration.pytorch_npu/_inductor/_npu_meta_registration.py+24-1
@@ -8,6 +8,7 @@ from torch import Tensor
8from torch._C import DispatchKey8from torch._C import DispatchKey
9from torch._decomp import decomposition_table, meta_table9from torch._decomp import decomposition_table, meta_table
10from torch._inductor import decomposition as inductor_decompo10from torch._inductor import decomposition as inductor_decompo
11+from torch._meta_registrations import device_hint
11from torch._ops import OpOverload, OpOverloadPacket12from torch._ops import OpOverload, OpOverloadPacket
12from torch._prims_common.wrappers import out_wrapper13from torch._prims_common.wrappers import out_wrapper
13from torch._subclasses import fake_tensor as _subclasses_fake_tensor14from torch._subclasses import fake_tensor as _subclasses_fake_tensor
@@ -106,7 +107,9 @@ def patch_torch_decomp_decompositions():
106 107 
107def register_meta_npu(op, avoid_fallback_flag=False, inductor_decomp=False):108def register_meta_npu(op, avoid_fallback_flag=False, inductor_decomp=False):
108 def meta_decorator(fn: Callable):109 def meta_decorator(fn: Callable):
109- _add_op_to_meta_table(op, fn, avoid_fallback_flag, inductor_decomp)110+ ops = op if isinstance(op, list) else [op]
111+ for single_op in ops:
112+ _add_op_to_meta_table(single_op, fn, avoid_fallback_flag, inductor_decomp)
110 return fn113 return fn
111 114 
112 return meta_decorator115 return meta_decorator
@@ -148,6 +151,26 @@ def npu_patch_meta():
148 patch_torch_inductor_decompositions()151 patch_torch_inductor_decompositions()
149 152 
150 153 
154+@register_meta_npu(
155+ [
156+ aten.sort.default,
157+ aten.sort.stable,
158+ aten.sort.values,
159+ aten.sort.values_stable,
160+ ]
161+ , avoid_fallback_flag=True
162+)
163+def meta_sort(self, stable=None, dim=-1, descending=False, values=None, indices=None):
164+ if device_hint(self) == "npu":
165+ v = torch.empty(self.shape, dtype=self.dtype, device=self.device)
166+ i = torch.empty(self.shape, dtype=torch.int64, device=self.device)
167+ 
168+ return v, i
169+ else:
170+ from torch._meta_registrations import meta_sort
171+ meta_sort(self, stable=stable, dim=dim, descending=descending, values=values, indices=indices)
172+ 
173+ 
151@register_meta_npu(aten.index_put.default)174@register_meta_npu(aten.index_put.default)
152def meta_index_put_patch(self, indices, values, accumulate=False):175def meta_index_put_patch(self, indices, values, accumulate=False):
153 return self.new_empty(self.shape)176 return self.new_empty(self.shape)
@@ -13,10 +13,3 @@ def apply_npu_format_patch():
13 from torch_npu.npu._format import _apply_npu_format_patch13 from torch_npu.npu._format import _apply_npu_format_patch
14 14 
15 _apply_npu_format_patch()15 _apply_npu_format_patch()
16- 
17- 
18-@PatchManager.register_patch("npu")
19-def apply_npu_meta_patch():
20- from torch_npu.utils._npu_meta_registration import npu_patch_meta
21- 
22- npu_patch_meta()