已合并
fix the circular dependency issue & add ut #26174
梁松伟创建于 2025年10月31日
fix the circular dependency issue & add ut #26174
已合并
梁松伟创建于 2025年10月31日
2 个文件变更+82-11
Atest/_inductor/test_npu_kernel_features.py+82-0
@@ -0,0 +1,82 @@
1+import sympy
2+import torch
3+import torch_npu
4+ 
5+from torch_npu.testing.testcase import TestCase, run_tests
6+from torch_npu._inductor.codegen.npu_kernel_features import NumelList
7+ 
8+ 
9+class TestNumeList(TestCase):
10+ def test_numels(self):
11+ numel_list = NumelList([2, 3, 4])
12+ self.assertEqual(numel_list.numels(), 24)
13+
14+ def test_equality(self):
15+ numel_list1 = NumelList([2, 3, 4])
16+ numel_list2 = NumelList([2, 3, 4])
17+ self.assertTrue(numel_list1 == numel_list2)
18+ 
19+ self.assertTrue(numel_list1 == 24)
20+ self.assertFalse(numel_list1 == 25)
21+ 
22+ def test_less_than(self):
23+ numel_list1 = NumelList([2, 3, 4])
24+ numel_list2 = NumelList([3, 4, 5])
25+ self.assertTrue(numel_list1 < numel_list2)
26+ 
27+ self.assertTrue(numel_list1 < 25)
28+ self.assertFalse(numel_list1 < 24)
29+ 
30+ def test_greater_than(self):
31+ numel_list1 = NumelList([2, 3, 5])
32+ numel_list2 = NumelList([2, 3, 4])
33+ self.assertTrue(numel_list1 > numel_list2)
34+ 
35+ def test_less_than_or_equal(self):
36+ numel_list1 = NumelList([2, 3, 4])
37+ numel_list2 = NumelList([3, 4, 5])
38+ self.assertTrue(numel_list1 <= numel_list2)
39+ 
40+ self.assertTrue(numel_list1 <= 25)
41+ self.assertTrue(numel_list1 <= 24)
42+ self.assertFalse(numel_list1 <= 23)
43+ 
44+ def test_greater_than_or_equal(self):
45+ numel_list1 = NumelList([2, 3, 5])
46+ numel_list2 = NumelList([2, 3, 4])
47+ self.assertTrue(numel_list1 >= numel_list2)
48+
49+ def test_modulo(self):
50+ numel_list = NumelList([2, 3, 4])
51+ self.assertEqual(numel_list % 5, 4)
52+
53+ def test_division(self):
54+ numel_list = NumelList([2, 3, 4])
55+ self.assertEqual(numel_list / 2, 12.0)
56+ self.assertEqual(numel_list // 2, 12)
57+ 
58+ def test_multiplication(self):
59+ numel_list = NumelList([2, 3, 4])
60+ self.assertEqual(numel_list * 2, 48)
61+ self.assertEqual(2 * numel_list, 48)
62+ 
63+ def test_addition(self):
64+ numel_list = NumelList([2, 3, 4])
65+ self.assertEqual(numel_list + 2, 26)
66+ self.assertEqual(2 + numel_list, 26)
67+
68+ def test_hash(self):
69+ # 测试相同内容的hash值相同
70+ numel_list1 = NumelList([2, 3, 4])
71+ numel_list2 = NumelList([2, 3, 4])
72+ self.assertEqual(hash(numel_list1), hash(numel_list2))
73+ 
74+ # 测试不同内容的hash值不同
75+ numel_list3 = NumelList([2, 3, 5])
76+ self.assertNotEqual(hash(numel_list1), hash(numel_list3))
77+ 
78+ # 测试NumelList对象的hash值与整数的hash值不同
79+ self.assertNotEqual(hash(numel_list1), hash(24))
80+ 
81+if __name__ == "__main__":
82+ run_tests()
Mtorch_npu/_inductor/codegen/npu_kernel_features.py+0-11
@@ -12,17 +12,6 @@ from torch.utils._ordered_set import OrderedSet
12 12 
13 13 
14class NumelList(Tuple):14class NumelList(Tuple):
15- 
16- @staticmethod
17- def calc_numels(other):
18- if isinstance(other, Iterable):
19- numel = NumelList.calc_numels(other)
20- return numel
21- elif isinstance(other, NumelList):
22- return other.numels()
23- else:
24- return other
25- 
26 def numels(self):15 def numels(self):
27 numel = functools.reduce(lambda a, b: a * b, self, 1)16 numel = functools.reduce(lambda a, b: a * b, self, 1)
28 return numel17 return numel