已合并
[Inductor] improve aot_load patch #45240
zhucehw创建于 8月24日
[Inductor] improve aot_load patch #45240
已合并
zhucehw创建于 8月24日
共 5 个文件变更+11-23
@@ -183,6 +183,7 @@ set_property(TARGET aoti_example PROPERTY CXX_STANDARD 17)
183## 使用约束183## 使用约束
184 184 
185暂不支持叠加Catlass,仅做功能兼容支持。185暂不支持叠加Catlass,仅做功能兼容支持。
186+不推荐使用社区已标记为deprecated的torch._export.aot_compile()/torch._export.aot_load()接口, 如需使用,请避免在同一进程内部多次调用
186 187 
187## 设备支持说明188## 设备支持说明
188 189 
@@ -7,25 +7,11 @@ def patch_aot_load():
7 origin_aot_load = torch._export.aot_load7 origin_aot_load = torch._export.aot_load
8 8 
9 def aot_load_npu(so_path: str, device: str) -> Callable:9 def aot_load_npu(so_path: str, device: str) -> Callable:
10- """
11- Loads a shared library generated by aot_compile and returns a callable
12 10 
13- Args:11+ if device == "npu" or device.startswith("npu:"):
14- so_path: Path to the shared library12+ runner = torch._C._aoti.AOTIModelContainerRunnerNpu(so_path, 1, device)
15- 13+ else:
16- Returns:
17- A callable
18- """
19- try:
20 return origin_aot_load(so_path, device)14 return origin_aot_load(so_path, device)
21- except RuntimeError as e:
22- if device == "npu" or device.startswith("npu:"):
23- import torch_npu
24- runner = torch_npu._C._aoti.AOTIModelContainerRunnerNpu(so_path, 1, device)
25- else:
26- raise RuntimeError(
27- f"Failed to load model with community logic: {e}"
28- ) from e
29 15 
30 def optimized(*args, **kwargs):16 def optimized(*args, **kwargs):
31 call_spec = runner.get_call_spec()17 call_spec = runner.get_call_spec()
@@ -11,9 +11,10 @@
11 11 
12namespace torch::inductor {12namespace torch::inductor {
13 13 
14-void initAOTIRunnerBindingsNpu(PyObject* module) {14+void initAOTIRunnerBindingsNpu() {
15+ py::module module = py::module::import("torch._C");
15 auto rootModule = py::handle(module).cast<py::module>();16 auto rootModule = py::handle(module).cast<py::module>();
16- auto m = rootModule.def_submodule("_aoti");17+ auto m = py::cast<py::module>(rootModule.attr("_aoti"));
17 18 
18#ifdef USE_NPU19#ifdef USE_NPU
19 py::class_<AOTIModelContainerRunnerNpu>(m, "AOTIModelContainerRunnerNpu")20 py::class_<AOTIModelContainerRunnerNpu>(m, "AOTIModelContainerRunnerNpu")
@@ -12,7 +12,7 @@
12 12 
13namespace torch::inductor {13namespace torch::inductor {
14 14 
15-void initAOTIRunnerBindingsNpu(PyObject* module);15+void initAOTIRunnerBindingsNpu();
16 16 
17} // namespace torch::inductor17} // namespace torch::inductor
18-#endif18+#endif
@@ -177,7 +177,7 @@ void RegisterNPUDeviceProperties(PyObject* module) {
177 177 
178 m.def("_npu_isHistoryEnabled", []() { return c10_npu::NPUCachingAllocator::isHistoryEnabled(); });178 m.def("_npu_isHistoryEnabled", []() { return c10_npu::NPUCachingAllocator::isHistoryEnabled(); });
179 179 
180- torch::inductor::initAOTIRunnerBindingsNpu(module);180+ torch::inductor::initAOTIRunnerBindingsNpu();
181}181}
182 182 
183std::string GetDeviceName() {183std::string GetDeviceName() {
@@ -2418,4 +2418,4 @@ void initCommMethods() {
2418 py::arg("out"),2418 py::arg("out"),
2419 py::arg("dim"),2419 py::arg("dim"),
2420 py::call_guard<py::gil_scoped_release>());2420 py::call_guard<py::gil_scoped_release>());
2421-}2421+}