已合并
test(package): add testcase for PackageExporter additional APIs #37842
PAGEMRW创建于 6月8日
test(package): add testcase for PackageExporter additional APIs #37842
已合并
PAGEMRW创建于 6月8日
已删除 :test-package-exporter-additional-api-v2.10.0合入到Ascend/pytorchv2.10.0
1 个文件变更+124-0
Atest/package/test_package_exporter_additional_api.py+124-0
@@ -0,0 +1,124 @@
1+"""
2+Add validation cases for torch.package APIs in torch-npu CI:
3+ 
4+1. PyTorch community lacks sufficient and direct API validations for some torch.package APIs, so this file is added.
5+2. This file validates torch.package.PackageExporter.add_dependency,
6+ torch.package.PackageExporter.all_paths,
7+ torch.package.PackageExporter.dependency_graph_string,
8+ torch.package.PackageExporter.get_unique_id,
9+ torch.package.PackageExporter.register_intern_hook,
10+ and torch.package.PackageExporter.close.
11+"""
12+ 
13+# Owner(s): ["oncall: package/deploy"]
14+ 
15+from io import BytesIO
16+ 
17+from torch.package import PackageExporter, PackageImporter, PackagingError
18+from torch.testing._internal.common_utils import TestCase, run_tests
19+ 
20+ 
21+class TestPackageExporterAdditionalAPI(TestCase):
22+ """Direct tests for PackageExporter APIs not covered by existing package tests."""
23+ 
24+ def test_add_dependency(self):
25+ buffer = BytesIO()
26+ exporter = PackageExporter(buffer)
27+ 
28+ exporter.add_dependency("math")
29+ 
30+ self.assertIn("math", exporter.dependency_graph.nodes)
31+ 
32+ exporter.close()
33+ buffer.seek(0)
34+ importer = PackageImporter(buffer)
35+ 
36+ import math
37+ 
38+ self.assertIs(importer.import_module("math"), math)
39+ 
40+ def test_add_dependency_nonexistent_module_raises(self):
41+ buffer = BytesIO()
42+ exporter = PackageExporter(buffer)
43+ 
44+ exporter.add_dependency("nonexistent_module_for_package_exporter_test")
45+ 
46+ with self.assertRaises(PackagingError):
47+ exporter.close()
48+ 
49+ def test_all_paths(self):
50+ exporter = PackageExporter(BytesIO())
51+ exporter.dependency_graph.add_edge("a", "b")
52+ exporter.dependency_graph.add_edge("b", "c")
53+ exporter.dependency_graph.add_edge("a", "d")
54+ 
55+ paths = exporter.all_paths("a", "c")
56+ 
57+ self.assertIn('"a" -> "b"', paths)
58+ self.assertIn('"b" -> "c"', paths)
59+ self.assertNotIn('"a" -> "d"', paths)
60+ 
61+ def test_dependency_graph_string(self):
62+ exporter = PackageExporter(BytesIO())
63+ exporter.dependency_graph.add_edge("a", "b")
64+ 
65+ graph = exporter.dependency_graph_string()
66+ 
67+ self.assertIn("digraph G", graph)
68+ self.assertIn('"a" -> "b"', graph)
69+ 
70+ def test_get_unique_id(self):
71+ exporter = PackageExporter(BytesIO())
72+ 
73+ self.assertEqual(exporter.get_unique_id(), "0")
74+ self.assertEqual(exporter.get_unique_id(), "1")
75+ self.assertEqual(exporter.get_unique_id(), "2")
76+ 
77+ def test_register_intern_hook(self):
78+ buffer = BytesIO()
79+ interned_modules = []
80+ 
81+ def intern_hook(package_exporter, module_name):
82+ interned_modules.append(module_name)
83+ 
84+ with PackageExporter(buffer) as exporter:
85+ exporter.register_intern_hook(intern_hook)
86+ exporter.save_source_string("foo", "VALUE = 1", dependencies=False)
87+ 
88+ self.assertEqual(interned_modules, ["foo"])
89+ 
90+ def test_register_intern_hook_remove(self):
91+ buffer = BytesIO()
92+ interned_modules = []
93+ 
94+ def intern_hook(package_exporter, module_name):
95+ interned_modules.append(module_name)
96+ 
97+ with PackageExporter(buffer) as exporter:
98+ handle = exporter.register_intern_hook(intern_hook)
99+ handle.remove()
100+ exporter.save_source_string("foo", "VALUE = 1", dependencies=False)
101+ 
102+ self.assertEqual(interned_modules, [])
103+ 
104+ def test_close(self):
105+ buffer = BytesIO()
106+ exporter = PackageExporter(buffer)
107+ exporter.save_source_string("foo", "VALUE = 3", dependencies=False)
108+ exporter.close()
109+ 
110+ buffer.seek(0)
111+ importer = PackageImporter(buffer)
112+ self.assertEqual(importer.import_module("foo").VALUE, 3)
113+ 
114+ def test_close_twice_raises(self):
115+ exporter = PackageExporter(BytesIO())
116+ 
117+ exporter.close()
118+ 
119+ with self.assertRaises(Exception):
120+ exporter.close()
121+ 
122+ 
123+if __name__ == "__main__":
124+ run_tests()