已合并
[Inductor] improve aot_load patch #45240
zhucehw创建于 8月24日
[Inductor] improve aot_load patch #45240
已合并
共 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_load | 7 | 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 library | 12 | + 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 | ||
| 12 | namespace torch::inductor { | 12 | namespace 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 | 19 | ||
| 19 | py::class_<AOTIModelContainerRunnerNpu>(m, "AOTIModelContainerRunnerNpu") | 20 | py::class_<AOTIModelContainerRunnerNpu>(m, "AOTIModelContainerRunnerNpu") |
| @@ -12,7 +12,7 @@ | |||
| 12 | 12 | ||
| 13 | namespace torch::inductor { | 13 | namespace torch::inductor { |
| 14 | 14 | ||
| 15 | -void initAOTIRunnerBindingsNpu(PyObject* module); | 15 | +void initAOTIRunnerBindingsNpu(); |
| 16 | 16 | ||
| 17 | } // namespace torch::inductor | 17 | } // namespace torch::inductor |
| 18 | -#endif | 18 | +#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 | ||
| 183 | std::string GetDeviceName() { | 183 | std::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 | +} |