已合并
fix the circular dependency issue & add ut #26174
梁松伟创建于 2025年10月31日
fix the circular dependency issue & add ut #26174
已合并
共 2 个文件变更+82-11
| @@ -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() | ||
| @@ -12,17 +12,6 @@ from torch.utils._ordered_set import OrderedSet | |||
| 12 | 12 | ||
| 13 | 13 | ||
| 14 | class NumelList(Tuple): | 14 | class NumelList(Tuple): |
| 15 | - | ||
| 16 | - | ||
| 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 numel | 17 | return numel |