已合并
test(package): add testcase for PackageExporter additional APIs #37842
PAGEMRW创建于 6月8日
test(package): add testcase for PackageExporter additional APIs #37842
已合并
从已删除 :test-package-exporter-additional-api-v2.10.0合入到Ascend/pytorchv2.10.0
共 1 个文件变更+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() | ||