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

🔵 Low Priority

新文件 test_nanquantile_defaul.py 的文件名中 "defaul" 缺少末尾字母 "t",应为 "default"。该测试测试的是 torch.ops.aten.nanquantile.default,文件名意图明显是 test_nanquantile_default。同目录下没有其他类似拼写错误的文件。虽然这只是文件名问题,不影响 Python 运行,但会误导其他开发者,且与项目命名规范不一致。

建议:将文件名改为 test_nanquantile_default.py

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
atomgit-bot
atomgit-botatomgit-bot6月25日

🔵 Low Priority

test_nanquantile_defaul.py 同时导入了 torch_npu(第 5 行)和 torch_npu._inductor(第 6 行)。第 6 行的 torch_npu._inductor 导入仅用于触发 inductor 插件的副作用初始化(注册 lowering 等),但该模块的返回值未被使用。如果上游 PyTorch 版本变更导致 torch_npu._inductor 的自动初始化失败,而 import torch_npu._inductor 是唯一触发初始化的路径,这本身是合理的。但观察第 24 行 torch.compile(fn, backend="inductor", dynamic=dynamic) 通常会自动触发 inductor 后端初始化,因此该导入可能是不必要的。

建议:确认 import torch_npu._inductor 是否确实为测试所必需。如果 torch.compile 已自动触发初始化,可移除该导入以减少隐式副作用依赖。

likedislike
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")
16+ a[0, 0] = float('nan')
17+ q = torch.tensor([0.25, 0.5, 0.75], dtype=torch.float32, device="npu")
atomgit-bot
atomgit-botatomgit-bot6月25日

🔵 Low Priority

第 14 行通过 @parametrize('dtype', ['float32'])dtype 进行了参数化,但第 16 行创建张量时硬编码了 dtype=torch.float32,完全忽略了 dtype 参数。这使得参数化形同虚设——如果未来有人将 ['float32'] 扩展为 ['float32', 'float16'],后者的测试实际仍在测 float32。

同时第 18 行的 q 张量也硬编码为 dtype=torch.float32

建议:将硬编码的 dtype=torch.float32 改为使用参数化传入的 dtype,例如 dtype=getattr(torch, dtype)dtype=torch.__dict__[dtype],与同目录其他测试(如 test_ge.py)的 _generate_tensor 用法一致。

likedislike
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()
@@ -85,6 +85,9 @@ def _load_triton_backend():
85 disable_foreach,85 disable_foreach,
86 get_current_raw_stream,86 get_current_raw_stream,
87 )87 )
88+ from ._npu_meta_registration import npu_patch_meta
89+ 
90+ npu_patch_meta()
88 91 
89 def _inductor_register_backend_for_device():92 def _inductor_register_backend_for_device():
90 from .codegen.cpp_wrapper import CppWrapperNpu93 from .codegen.cpp_wrapper import CppWrapperNpu
Rtorch_npu/utils/_npu_meta_registration.pytorch_npu/_inductor/_npu_meta_registration.py+24-1
@@ -19,6 +19,7 @@ from torch._dynamo.exc import Unsupported
19from torch._dynamo.variables.lists import TupleVariable19from torch._dynamo.variables.lists import TupleVariable
20from torch._dynamo.variables.nn_module import NNModuleVariable20from torch._dynamo.variables.nn_module import NNModuleVariable
21from torch._decomp import meta_table21from torch._decomp import meta_table
22+from torch._meta_registrations import device_hint
22import torch_npu23import torch_npu
23 24 
24aten = torch.ops.aten25aten = torch.ops.aten
@@ -109,7 +110,9 @@ def patch_torch_decomp_decompositions():
109 110 
110def register_meta_npu(op, avoid_fallback_flag=False, inductor_decomp=False):111def register_meta_npu(op, avoid_fallback_flag=False, inductor_decomp=False):
111 def meta_decorator(fn: Callable):112 def meta_decorator(fn: Callable):
112- _add_op_to_meta_table(op, fn, avoid_fallback_flag, inductor_decomp)113+ ops = op if isinstance(op, list) else [op]
114+ for single_op in ops:
115+ _add_op_to_meta_table(single_op, fn, avoid_fallback_flag, inductor_decomp)
113 return fn116 return fn
114 117 
115 return meta_decorator118 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()