已合并
fix: preserve deferred Triton backend loading (v2.12.0) #44737
fix: preserve deferred Triton backend loading (v2.12.0) #44737
已合并
黄桂军创建于 8月17日
共 2 个文件变更+20-24
@@ -705,6 +705,23 @@ class TorchCompileTriggerTests(unittest.TestCase):
705 """705 """
706 )706 )
707 707 
708+ # Verify creating an Inductor wrapper does not load the backend prematurely.
709+ def test_inductor_backend_load_is_deferred_until_first_call(self):
710+ self.run_in_subprocess(
711+ """
712+ import sys
713+ import torch
714+ import torch_npu
715+ from torch_npu.utils import _dynamo
716+ 
717+ torch.compile(lambda x: x + 1, backend="inductor")
718+ 
719+ assert _dynamo._lazy_dynamo_setup.has_run
720+ assert not _dynamo._lazy_inductor_setup.has_run
721+ assert "torch_npu._inductor" not in sys.modules
722+ """
723+ )
724+ 
708 # Verify lazy setup completes before compile backend lookup.725 # Verify lazy setup completes before compile backend lookup.
709 def test_compile_triggers_setup_before_backend_lookup(self):726 def test_compile_triggers_setup_before_backend_lookup(self):
710 self.run_in_subprocess(727 self.run_in_subprocess(
@@ -166,6 +166,7 @@ class _NpuBackendScope:
166 try:166 try:
167 os.environ["TORCHINDUCTOR_NPU_BACKEND"] = self.backend167 os.environ["TORCHINDUCTOR_NPU_BACKEND"] = self.backend
168 register_inductor_npu()168 register_inductor_npu()
169+ _lazy_inductor_setup()
169 if self.backend == "ascendc":170 if self.backend == "ascendc":
170 from torch_npu._inductor.deterministic_cache import (171 from torch_npu._inductor.deterministic_cache import (
171 patch_npu_deterministic_level_cache_keys,172 patch_npu_deterministic_level_cache_keys,
@@ -221,6 +222,7 @@ def patch_inductor_wrapper():
221 self._config["npu_backend"] = _ConfigEntry(cfg, "npu_backend")222 self._config["npu_backend"] = _ConfigEntry(cfg, "npu_backend")
222 else:223 else:
223 self._config["npu_backend"] = _ConfigEntry(cfg)224 self._config["npu_backend"] = _ConfigEntry(cfg)
225+ 
224 return ori_dict226 return ori_dict
225 227 
226 def new_init(self, mode, options, dynamic, name=None):228 def new_init(self, mode, options, dynamic, name=None):
@@ -231,13 +233,10 @@ def patch_inductor_wrapper():
231 src_init(self, mode, options, dynamic, name)233 src_init(self, mode, options, dynamic, name)
232 else:234 else:
233 src_init(self, mode, options, dynamic)235 src_init(self, mode, options, dynamic)
234- shape_handling_requested = self._npu_shape_handling_requested
235 finally:236 finally:
236 del self._npu_defer_shape_handling237 del self._npu_defer_shape_handling
237 del self._npu_shape_handling_requested238 del self._npu_shape_handling_requested
238- _setup_inductor_for_compile(self.config)239+ _lazy_dynamo_setup()
239- if shape_handling_requested:
240- torch_npu._inductor.patch_shape_handling()
241 backend = _resolve_npu_backend_from_wrapper(self)240 backend = _resolve_npu_backend_from_wrapper(self)
242 if backend == "mlir":241 if backend == "mlir":
243 with _NpuBackendScope(backend):242 with _NpuBackendScope(backend):
@@ -652,26 +651,6 @@ def _lazy_inductor_setup():
652 _inject_inductor_npu_backend_config()651 _inject_inductor_npu_backend_config()
653 652 
654 653 
655-def _setup_inductor_for_compile(options=None):
656- """Initialize the NPU Inductor backend selected for this compile call."""
657- _lazy_dynamo_setup()
658- 
659- option_backend = options.get("npu_backend") if isinstance(options, dict) else None
660- selected_backend = _resolve_npu_backend(option_backend)
661- 
662- old_backend = os.environ.get("TORCHINDUCTOR_NPU_BACKEND")
663- if selected_backend not in (None, "", "default"):
664- os.environ["TORCHINDUCTOR_NPU_BACKEND"] = selected_backend
665- try:
666- _lazy_inductor_setup()
667- finally:
668- if old_backend is None:
669- os.environ.pop("TORCHINDUCTOR_NPU_BACKEND", None)
670- else:
671- os.environ["TORCHINDUCTOR_NPU_BACKEND"] = old_backend
672- return selected_backend
673- 
674- 
675@run_once654@run_once
676def install_npugraph_mark_step_trigger():655def install_npugraph_mark_step_trigger():
677 """Expose the public NPUGraph step API without importing compiler internals."""656 """Expose the public NPUGraph step API without importing compiler internals."""