已合并
[Inductor] Inductor support multi process, add Device guard for AOTInductor #38634
zhucehw创建于 6月16日
[Inductor] Inductor support multi process, add Device guard for AOTInductor #38634
已合并
zhucehw创建于 6月16日
7 个文件变更+80-120
@@ -603,88 +603,55 @@ class TestMultiStreamPass(TestUtils):
603 2,603 2,
604 )604 )
605 605 
606- 
607 def multi_stream_test(606 def multi_stream_test(
608 self,607 self,
609- arg0_1,608+ arg0_1, arg2_1, arg3_1, arg4_1, arg5_1, arg6_1, arg7_1, arg8_1, arg9_1,
610- arg1_1, arg2_1, arg3_1, arg4_1, arg5_1, arg6_1, arg7_1, arg8_1, arg9_1,609+ arg10_1, arg23_1, arg24_1, arg25_1, arg26_1, arg27_1, arg28_1,
611- arg10_1, arg11_1, arg12_1, arg13_1, arg14_1, arg15_1, arg16_1, arg17_1,
612- arg18_1, arg19_1, arg20_1, arg21_1, arg22_1,
613- arg23_1, arg24_1, arg25_1, arg26_1, arg27_1, arg28_1,
614 arg29_1, arg30_1, arg31_1, arg32_1, arg33_1, arg34_1,610 arg29_1, arg30_1, arg31_1, arg32_1, arg33_1, arg34_1,
615 arg35_1, arg36_1, arg37_1, arg38_1, arg39_1611 arg35_1, arg36_1, arg37_1, arg38_1, arg39_1
616 ):612 ):
617- slice_2 = torch.ops.aten.slice.Tensor(arg0_1, 1, 0, 3)613+ slice_2 = torch.ops.aten.slice.Tensor(arg0_1, 1, 0, 1)
618 sum_1 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg2_1, slice_2), [1])614 sum_1 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg2_1, slice_2), [1])
619- slice_4 = torch.ops.aten.slice.Tensor(arg0_1, 1, 3, 5)615+ slice_4 = torch.ops.aten.slice.Tensor(arg0_1, 1, 1, 2)
620 sum_2 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg3_1, slice_4), [1])616 sum_2 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg3_1, slice_4), [1])
621- slice_6 = torch.ops.aten.slice.Tensor(arg0_1, 1, 5, 6)617+ slice_6 = torch.ops.aten.slice.Tensor(arg0_1, 1, 2, 3)
622 sum_3 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg4_1, slice_6), [1])618 sum_3 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg4_1, slice_6), [1])
623- slice_8 = torch.ops.aten.slice.Tensor(arg0_1, 1, 6, 8)619+ slice_8 = torch.ops.aten.slice.Tensor(arg0_1, 1, 3, 4)
624 sum_4 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg5_1, slice_8), [1])620 sum_4 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg5_1, slice_8), [1])
625- slice_10 = torch.ops.aten.slice.Tensor(arg0_1, 1, 8, 14)621+ slice_10 = torch.ops.aten.slice.Tensor(arg0_1, 1, 4, 6)
626 sum_5 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg6_1, slice_10), [1])622 sum_5 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg6_1, slice_10), [1])
627- slice_12 = torch.ops.aten.slice.Tensor(arg0_1, 1, 14, 15)623+ slice_12 = torch.ops.aten.slice.Tensor(arg0_1, 1, 6, 7)
628 sum_6 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg7_1, slice_12), [1])624 sum_6 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg7_1, slice_12), [1])
629- slice_14 = torch.ops.aten.slice.Tensor(arg0_1, 1, 15, 16)625+ slice_14 = torch.ops.aten.slice.Tensor(arg0_1, 1, 7, 8)
630 sum_7 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg8_1, slice_14), [1])626 sum_7 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg8_1, slice_14), [1])
631- slice_16 = torch.ops.aten.slice.Tensor(arg0_1, 1, 16, 17)627+ slice_16 = torch.ops.aten.slice.Tensor(arg0_1, 1, 8, 9)
632 sum_8 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg9_1, slice_16), [1])628 sum_8 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg9_1, slice_16), [1])
633- slice_18 = torch.ops.aten.slice.Tensor(arg0_1, 1, 17, 18)629+ slice_18 = torch.ops.aten.slice.Tensor(arg0_1, 1, 9, 10)
634 sum_9 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg10_1, slice_18), [1])630 sum_9 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg10_1, slice_18), [1])
635- slice_20 = torch.ops.aten.slice.Tensor(arg0_1, 1, 18, 25)631+ 
636- sum_10 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg11_1, slice_20), [1])
637- slice_22 = torch.ops.aten.slice.Tensor(arg0_1, 1, 25, 28)
638- sum_11 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg12_1, slice_22), [1])
639- slice_24 = torch.ops.aten.slice.Tensor(arg0_1, 1, 28, 36)
640- sum_12 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg13_1, slice_24), [1])
641- slice_26 = torch.ops.aten.slice.Tensor(arg0_1, 1, 36, 37)
642- sum_13 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg14_1, slice_26), [1])
643- slice_28 = torch.ops.aten.slice.Tensor(arg0_1, 1, 37, 43)
644- sum_14 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg15_1, slice_28), [1])
645- slice_30 = torch.ops.aten.slice.Tensor(arg0_1, 1, 43, 52)
646- sum_15 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg16_1, slice_30), [1])
647- slice_32 = torch.ops.aten.slice.Tensor(arg0_1, 1, 52, 57)
648- sum_16 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg17_1, slice_32), [1])
649- slice_34 = torch.ops.aten.slice.Tensor(arg0_1, 1, 57, 58)
650- sum_17 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg18_1, slice_34), [1])
651- slice_36 = torch.ops.aten.slice.Tensor(arg0_1, 1, 58, 59)
652- sum_18 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg19_1, slice_36), [1])
653- slice_38 = torch.ops.aten.slice.Tensor(arg0_1, 1, 59, 60)
654- sum_19 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg20_1, slice_38), [1])
655- slice_40 = torch.ops.aten.slice.Tensor(arg0_1, 1, 60, 72)
656- sum_20 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg21_1, slice_40), [1])
657- slice_42 = torch.ops.aten.slice.Tensor(arg0_1, 1, 72, 172)
658- sum_21 = torch.ops.aten.sum.dim_IntList(torch.ops.aten.embedding.default(arg22_1, slice_42), [1])
659 cat = torch.ops.aten.cat.default(632 cat = torch.ops.aten.cat.default(
660 [sum_1, sum_2, sum_3, sum_4, sum_5,633 [sum_1, sum_2, sum_3, sum_4, sum_5,
661- sum_6, sum_7, sum_8, sum_9, sum_10,634+ sum_6, sum_7, sum_8, sum_9],
662- sum_11, sum_12, sum_13, sum_14, sum_15,
663- sum_16, sum_17, sum_18, sum_19, sum_20, sum_21],
664 1635 1
665 )636 )
666 add_relu = torch.ops.aten.add.Tensor(cat, cat)637 add_relu = torch.ops.aten.add.Tensor(cat, cat)
667 638 
668- # branch A
669 a = torch.ops.aten.relu.default(torch.ops.aten.mm.default(arg23_1, arg24_1))639 a = torch.ops.aten.relu.default(torch.ops.aten.mm.default(arg23_1, arg24_1))
670 a = torch.ops.aten.relu.default(torch.ops.aten.mm.default(a, arg25_1))640 a = torch.ops.aten.relu.default(torch.ops.aten.mm.default(a, arg25_1))
671 a = torch.ops.aten.relu.default(torch.ops.aten.mm.default(a, arg26_1))641 a = torch.ops.aten.relu.default(torch.ops.aten.mm.default(a, arg26_1))
672 a = torch.ops.aten.relu.default(torch.ops.aten.mm.default(a, arg27_1))642 a = torch.ops.aten.relu.default(torch.ops.aten.mm.default(a, arg27_1))
673 mm_4 = torch.ops.aten.relu.default(torch.ops.aten.mm.default(a, arg28_1))643 mm_4 = torch.ops.aten.relu.default(torch.ops.aten.mm.default(a, arg28_1))
674 644 
675- # branch B
676 b = torch.ops.aten.relu.default(torch.ops.aten.mm.default(arg29_1, arg30_1))645 b = torch.ops.aten.relu.default(torch.ops.aten.mm.default(arg29_1, arg30_1))
677 b = torch.ops.aten.relu.default(torch.ops.aten.mm.default(b, arg31_1))646 b = torch.ops.aten.relu.default(torch.ops.aten.mm.default(b, arg31_1))
678 b = torch.ops.aten.relu.default(torch.ops.aten.mm.default(b, arg32_1))647 b = torch.ops.aten.relu.default(torch.ops.aten.mm.default(b, arg32_1))
679 b = torch.ops.aten.relu.default(torch.ops.aten.mm.default(b, arg33_1))648 b = torch.ops.aten.relu.default(torch.ops.aten.mm.default(b, arg33_1))
680 mm_8 = torch.ops.aten.relu.default(torch.ops.aten.mm.default(b, arg34_1))649 mm_8 = torch.ops.aten.relu.default(torch.ops.aten.mm.default(b, arg34_1))
681 650 
682- # merge
683 mm_8_t = torch.ops.aten.permute.default(mm_8, [1, 0])651 mm_8_t = torch.ops.aten.permute.default(mm_8, [1, 0])
684 merge = torch.ops.aten.mm.default(mm_4, mm_8_t)652 merge = torch.ops.aten.mm.default(mm_4, mm_8_t)
685 merge = torch.ops.aten.mm.default(merge, arg35_1)653 merge = torch.ops.aten.mm.default(merge, arg35_1)
686 654 
687- # MAIN
688 add = torch.ops.aten.add.Tensor(merge, add_relu)655 add = torch.ops.aten.add.Tensor(merge, add_relu)
689 mm_15 = torch.ops.aten.mm.default(arg36_1, add)656 mm_15 = torch.ops.aten.mm.default(arg36_1, add)
690 relu_15 = torch.ops.aten.relu.default(mm_15)657 relu_15 = torch.ops.aten.relu.default(mm_15)
@@ -698,51 +665,46 @@ class TestMultiStreamPass(TestUtils):
698 665 
699 @patch("torch_npu._inductor.fx_passes.parallel_scheduler_pass.is_multi_stream", return_value=True)666 @patch("torch_npu._inductor.fx_passes.parallel_scheduler_pass.is_multi_stream", return_value=True)
700 def test_multi_stream_compile_case(self, mock_multi_stream):667 def test_multi_stream_compile_case(self, mock_multi_stream):
701- arg0_1 = torch.randint(0, 99, (64, 199), dtype=torch.int64, device="npu")668+ arg0_1 = torch.randint(0, 16, (8, 40), dtype=torch.int64, device="npu")
702- arg1_1 = torch.randn(100, 64, device="npu")669+ arg1_1 = torch.randn(16, 8, device="npu")
703- arg2_1 = torch.randn(100, 64, device="npu")670+ arg2_1 = torch.randn(16, 8, device="npu")
704- arg3_1 = torch.randn(100, 64, device="npu")671+ arg3_1 = torch.randn(16, 8, device="npu")
705- arg4_1 = torch.randn(100, 64, device="npu")672+ arg4_1 = torch.randn(16, 8, device="npu")
706- arg5_1 = torch.randn(100, 64, device="npu")673+ arg5_1 = torch.randn(16, 8, device="npu")
707- arg6_1 = torch.randn(100, 64, device="npu")674+ arg6_1 = torch.randn(16, 8, device="npu")
708- arg7_1 = torch.randn(100, 64, device="npu")675+ arg7_1 = torch.randn(16, 8, device="npu")
709- arg8_1 = torch.randn(100, 64, device="npu")676+ arg8_1 = torch.randn(16, 8, device="npu")
710- arg9_1 = torch.randn(100, 64, device="npu")677+ arg9_1 = torch.randn(16, 8, device="npu")
711- arg10_1 = torch.randn(100, 64, device="npu")678+ arg10_1 = torch.randn(16, 8, device="npu")
712- arg11_1 = torch.randn(100, 64, device="npu")
713- arg12_1 = torch.randn(100, 64, device="npu")
714- arg13_1 = torch.randn(100, 64, device="npu")
715- arg14_1 = torch.randn(100, 64, device="npu")
716- arg15_1 = torch.randn(100, 64, device="npu")
717- arg16_1 = torch.randn(100, 64, device="npu")
718- arg17_1 = torch.randn(100, 64, device="npu")
719- arg18_1 = torch.randn(100, 64, device="npu")
720- arg19_1 = torch.randn(100, 64, device="npu")
721- arg20_1 = torch.randn(100, 64, device="npu")
722- arg21_1 = torch.randn(100, 64, device="npu")
723- arg22_1 = torch.randn(100, 64, device="npu")
724- arg23_1 = torch.randn(64, 64, device="npu")
725- arg24_1 = torch.randn(64, 64, device="npu")
726- arg25_1 = torch.randn(64, 64, device="npu")
727- arg26_1 = torch.randn(64, 64, device="npu")
728- arg27_1 = torch.randn(64, 64, device="npu")
729- arg28_1 = torch.randn(64, 64, device="npu")
730- arg29_1 = torch.randn(64, 64, device="npu")
731- arg30_1 = torch.randn(64, 64, device="npu")
732- arg31_1 = torch.randn(64, 64, device="npu")
733- arg32_1 = torch.randn(64, 64, device="npu")
734- arg33_1 = torch.randn(64, 64, device="npu")
735- arg34_1 = torch.randn(64, 64, device="npu")
736- arg35_1 = torch.randn(64, 1344, device="npu")
737- arg36_1 = torch.randn(64, 64, device="npu")
738- arg37_1 = torch.randn(64, 1344, device="npu")
739- arg38_1 = torch.randn(64, 64, device="npu")
740- arg39_1 = torch.randn(64, 1344, device="npu")
741 679 
742- std_result = self.multi_stream_test(arg0_1, arg1_1, arg2_1, arg3_1, arg4_1, arg5_1, arg6_1, arg7_1, arg8_1, arg9_1, arg10_1, arg11_1, arg12_1, arg13_1, arg14_1, arg15_1, arg16_1, arg17_1, arg18_1, arg19_1, arg20_1, arg21_1, arg22_1, arg23_1, arg24_1, arg25_1, arg26_1, arg27_1, arg28_1, arg29_1, arg30_1, arg31_1, arg32_1, arg33_1, arg34_1, arg35_1, arg36_1, arg37_1, arg38_1, arg39_1)680+ arg23_1 = torch.randn(8, 8, device="npu")
681+ arg24_1 = torch.randn(8, 8, device="npu")
682+ arg25_1 = torch.randn(8, 8, device="npu")
683+ arg26_1 = torch.randn(8, 8, device="npu")
684+ arg27_1 = torch.randn(8, 8, device="npu")
685+ arg28_1 = torch.randn(8, 8, device="npu")
686+ arg29_1 = torch.randn(8, 8, device="npu")
687+ arg30_1 = torch.randn(8, 8, device="npu")
688+ arg31_1 = torch.randn(8, 8, device="npu")
689+ arg32_1 = torch.randn(8, 8, device="npu")
690+ arg33_1 = torch.randn(8, 8, device="npu")
691+ arg34_1 = torch.randn(8, 8, device="npu")
692+ arg35_1 = torch.randn(8, 72, device="npu")
693+ arg36_1 = torch.randn(8, 8, device="npu")
694+ arg37_1 = torch.randn(8, 72, device="npu")
695+ arg38_1 = torch.randn(8, 8, device="npu")
696+ arg39_1 = torch.randn(8, 72, device="npu")
697+ 
698+ std_result = self.multi_stream_test(arg0_1, arg2_1, arg3_1, arg4_1, arg5_1, arg6_1, arg7_1, arg8_1, arg9_1,
699+ arg10_1, arg23_1, arg24_1, arg25_1, arg26_1, arg27_1, arg28_1, arg29_1,
700+ arg30_1, arg31_1, arg32_1, arg33_1, arg34_1, arg35_1, arg36_1, arg37_1,
701+ arg38_1, arg39_1)
743 with torch.no_grad():702 with torch.no_grad():
744 compiled_op_calc = torch.compile(self.multi_stream_test, backend="inductor")703 compiled_op_calc = torch.compile(self.multi_stream_test, backend="inductor")
745- inductor_result = compiled_op_calc(arg0_1, arg1_1, arg2_1, arg3_1, arg4_1, arg5_1, arg6_1, arg7_1, arg8_1, arg9_1, arg10_1, arg11_1, arg12_1, arg13_1, arg14_1, arg15_1, arg16_1, arg17_1, arg18_1, arg19_1, arg20_1, arg21_1, arg22_1, arg23_1, arg24_1, arg25_1, arg26_1, arg27_1, arg28_1, arg29_1, arg30_1, arg31_1, arg32_1, arg33_1, arg34_1, arg35_1, arg36_1, arg37_1, arg38_1, arg39_1)704+ inductor_result = compiled_op_calc(arg0_1, arg2_1, arg3_1, arg4_1, arg5_1, arg6_1, arg7_1, arg8_1, arg9_1,
705+ arg10_1, arg23_1, arg24_1, arg25_1, arg26_1, arg27_1, arg28_1, arg29_1,
706+ arg30_1, arg31_1, arg32_1, arg33_1, arg34_1, arg35_1, arg36_1, arg37_1,
707+ arg38_1, arg39_1)
746 self.assertEqual(std_result, inductor_result, atol=1e-3, rtol=1e-3)708 self.assertEqual(std_result, inductor_result, atol=1e-3, rtol=1e-3)
747 709 
748 def op_calc(self, arg0_1, arg1_1, arg2_1, arg3_1, arg4_1, arg5_1):710 def op_calc(self, arg0_1, arg1_1, arg2_1, arg3_1, arg4_1, arg5_1):
@@ -789,4 +751,4 @@ instantiate_parametrized_tests(TestMultiStreamPass)
789 751 
790 752 
791if __name__ == "__main__":753if __name__ == "__main__":
792- run_tests()754+ run_tests()
@@ -307,11 +307,6 @@ def _load_triton_backend():
307 add_additional_op()307 add_additional_op()
308 os.environ["TORCHINDUCTOR_COMPREHENSIVE_PADDING"] = "0"308 os.environ["TORCHINDUCTOR_COMPREHENSIVE_PADDING"] = "0"
309 torch._inductor.config.comprehensive_padding = False309 torch._inductor.config.comprehensive_padding = False
310- compile_threads = int(
311- os.environ.get("TORCHINDUCTOR_COMPILE_THREADS") or "1"
312- )
313- os.environ["TORCHINDUCTOR_COMPILE_THREADS"] = str(compile_threads)
314- torch._inductor.config.compile_threads = compile_threads
315 310 
316 _fasta_autotune = os.environ.get("FASTAUTOTUNE", "0") == "1"311 _fasta_autotune = os.environ.get("FASTAUTOTUNE", "0") == "1"
317 _fasta_autotune_method = os.getenv("AUTOTUNE_METHOD", "Expert")312 _fasta_autotune_method = os.getenv("AUTOTUNE_METHOD", "Expert")
@@ -114,6 +114,7 @@ def gen_common_triton_imports():
114 """114 """
115 import torch115 import torch
116 import torch_npu116 import torch_npu
117+ torch_npu.npu._initialized = True
117 from torch_npu._inductor.runtime import triton_heuristics as triton_heuristics118 from torch_npu._inductor.runtime import triton_heuristics as triton_heuristics
118 from torch_npu._inductor.runtime import triton_helpers119 from torch_npu._inductor.runtime import triton_helpers
119 from torch_npu._inductor.runtime.triton_helpers import libdevice, extension, math as tl_math120 from torch_npu._inductor.runtime.triton_helpers import libdevice, extension, math as tl_math
@@ -83,8 +83,10 @@ class NPUWrapperCodeGen(_NPUKernelCodegenMixin, PythonWrapperCodegen):
83 super().write_triton_header_once()83 super().write_triton_header_once()
84 import_str = f"""84 import_str = f"""
85 import torch_npu85 import torch_npu
86+ torch_npu.npu._initialized = torch_npu.npu.is_initialized()
86 has_initialized = False87 has_initialized = False
87 """88 """
89+ 
88 if config.triton.autotune_at_compile_time:90 if config.triton.autotune_at_compile_time:
89 self.kernel_autotune_calls.splice(import_str)91 self.kernel_autotune_calls.splice(import_str)
90 self.kernel_autotune_calls.splice(92 self.kernel_autotune_calls.splice(
@@ -13,14 +13,12 @@ namespace torch::aot_inductor {
13 13 
14inline void delete_npu_guard(void* ptr)14inline void delete_npu_guard(void* ptr)
15{15{
16- AOTI_TORCH_ERROR_CODE_CHECK(16+ AOTI_TORCH_ERROR_CODE_CHECK(aoti_torch_delete_npu_guard(reinterpret_cast<NPUGuardHandle>(ptr)));
17- aoti_torch_delete_npu_guard(reinterpret_cast<NPUGuardHandle>(ptr)));
18}17}
19 18 
20inline void delete_npu_stream_guard(void* ptr)19inline void delete_npu_stream_guard(void* ptr)
21{20{
22- AOTI_TORCH_ERROR_CODE_CHECK(21+ AOTI_TORCH_ERROR_CODE_CHECK(aoti_torch_delete_npu_stream_guard(reinterpret_cast<NPUStreamGuardHandle>(ptr)));
23- aoti_torch_delete_npu_stream_guard(reinterpret_cast<NPUStreamGuardHandle>(ptr)));
24}22}
25 23 
26class AOTINpuGuard {24class AOTINpuGuard {
@@ -34,8 +32,7 @@ public:
34 32 
35 void set_index(int32_t device_index)33 void set_index(int32_t device_index)
36 {34 {
37- AOTI_TORCH_ERROR_CODE_CHECK(35+ AOTI_TORCH_ERROR_CODE_CHECK(aoti_torch_npu_guard_set_index(guard_.get(), device_index));
38- aoti_torch_npu_guard_set_index(guard_.get(), device_index));
39 }36 }
40 37 
41private:38private:
@@ -47,8 +44,7 @@ public:
47 AOTINpuStreamGuard(aclrtStream stream, int32_t device_index): guard_(nullptr, delete_npu_stream_guard)44 AOTINpuStreamGuard(aclrtStream stream, int32_t device_index): guard_(nullptr, delete_npu_stream_guard)
48 {45 {
49 NPUStreamGuardHandle ptr = nullptr;46 NPUStreamGuardHandle ptr = nullptr;
50- AOTI_TORCH_ERROR_CODE_CHECK(47+ AOTI_TORCH_ERROR_CODE_CHECK(aoti_torch_create_npu_stream_guard(stream, device_index, &ptr));
51- aoti_torch_create_npu_stream_guard(stream, device_index, &ptr));
52 guard_.reset(ptr);48 guard_.reset(ptr);
53 }49 }
54 50 
@@ -16,26 +16,21 @@ AOTI_TORCH_EXPORT AOTITorchError aoti_torch_create_npu_guard(
16 NPUGuardHandle* ret_guard // returns new reference16 NPUGuardHandle* ret_guard // returns new reference
17);17);
18 18 
19-AOTI_TORCH_EXPORT AOTITorchError19+AOTI_TORCH_EXPORT AOTITorchError aoti_torch_delete_npu_guard(NPUGuardHandle guard);
20-aoti_torch_delete_npu_guard(NPUGuardHandle guard);
21 20 
22-AOTI_TORCH_EXPORT AOTITorchError21+AOTI_TORCH_EXPORT AOTITorchError aoti_torch_npu_guard_set_index(NPUGuardHandle guard, int32_t device_index);
23-aoti_torch_npu_guard_set_index(NPUGuardHandle guard, int32_t device_index);
24 22 
25struct NPUStreamGuardOpaque;23struct NPUStreamGuardOpaque;
26using NPUStreamGuardHandle = NPUStreamGuardOpaque*;24using NPUStreamGuardHandle = NPUStreamGuardOpaque*;
27 25 
28AOTI_TORCH_EXPORT AOTITorchError aoti_torch_create_npu_stream_guard(26AOTI_TORCH_EXPORT AOTITorchError aoti_torch_create_npu_stream_guard(
29- void* stream,27+ void* stream, int32_t device_index,
30- int32_t device_index,
31 NPUStreamGuardHandle* ret_guard // returns new reference28 NPUStreamGuardHandle* ret_guard // returns new reference
32);29);
33 30 
34-AOTI_TORCH_EXPORT AOTITorchError31+AOTI_TORCH_EXPORT AOTITorchError aoti_torch_delete_npu_stream_guard(NPUStreamGuardHandle guard);
35-aoti_torch_delete_npu_stream_guard(NPUStreamGuardHandle guard);
36 32 
37-AOTI_TORCH_EXPORT AOTITorchError33+AOTI_TORCH_EXPORT AOTITorchError aoti_torch_get_current_npu_stream(int32_t device_index, void** ret_stream);
38-aoti_torch_get_current_npu_stream(int32_t device_index, void** ret_stream);
39 34 
40AOTI_TORCH_EXPORT AOTITorchError35AOTI_TORCH_EXPORT AOTITorchError
41aoti_torch_get_current_npu_device(int32_t* device_index);36aoti_torch_get_current_npu_device(int32_t* device_index);
@@ -50,4 +45,4 @@ AOTI_TORCH_EXPORT AOTITorchError aoti_torch_get_current_sycl_queue(void** ret);
50#endif45#endif
51 46 
52#endif // USE_NPU47#endif // USE_NPU
53-#endif // AOTI_TORCH_SHIM_NPU48+#endif // AOTI_TORCH_SHIM_NPU
@@ -7,6 +7,7 @@
7#include <torch/csrc/inductor/aoti_torch/utils.h>7#include <torch/csrc/inductor/aoti_torch/utils.h>
8#include <torch/csrc/inductor/inductor_ops.h>8#include <torch/csrc/inductor/inductor_ops.h>
9#include <torch_npu/csrc/aten/common/from_blob.h>9#include <torch_npu/csrc/aten/common/from_blob.h>
10+#include <torch_npu/csrc/core/npu/NPUGuard.h>
10#include <torch_npu/csrc/core/npu/NPUStream.h>11#include <torch_npu/csrc/core/npu/NPUStream.h>
11#include <torch_npu/csrc/core/npu/NPUCachingAllocator.h>12#include <torch_npu/csrc/core/npu/NPUCachingAllocator.h>
12#include <torch_npu/csrc/inductor/aoti_torch/c/shim_npu.h>13#include <torch_npu/csrc/inductor/aoti_torch/c/shim_npu.h>
@@ -33,20 +34,24 @@ static c10::Device c10_device(int32_t device_type, int32_t device_index)
33 34 
34AOTITorchError aoti_torch_create_npu_guard(int32_t device_index, NPUGuardHandle* ret_guard)35AOTITorchError aoti_torch_create_npu_guard(int32_t device_index, NPUGuardHandle* ret_guard)
35{36{
36- // todo: implement create npu guard logic37+ AOTI_TORCH_CONVERT_EXCEPTION_TO_ERROR_CODE({
37- return AOTI_TORCH_SUCCESS;38+ TORCH_CHECK(ret_guard != nullptr, "ret_guard is nullptr");
39+ *ret_guard = nullptr;
40+ c10_npu::NPUGuard* guard = new c10_npu::NPUGuard(static_cast<c10::DeviceIndex>(device_index));
41+ *ret_guard = reinterpret_cast<NPUGuardHandle>(guard);
42+ });
38}43}
39 44 
40AOTITorchError aoti_torch_delete_npu_guard(NPUGuardHandle guard)45AOTITorchError aoti_torch_delete_npu_guard(NPUGuardHandle guard)
41{46{
42- // todo: implement delete npu guard logic47+ AOTI_TORCH_CONVERT_EXCEPTION_TO_ERROR_CODE({ delete reinterpret_cast<c10_npu::NPUGuard*>(guard); });
43- return AOTI_TORCH_SUCCESS;
44}48}
45 49 
46AOTITorchError aoti_torch_npu_guard_set_index(NPUGuardHandle guard, int32_t device_index)50AOTITorchError aoti_torch_npu_guard_set_index(NPUGuardHandle guard, int32_t device_index)
47{51{
48- // todo: implement npu guard set index logic52+ AOTI_TORCH_CONVERT_EXCEPTION_TO_ERROR_CODE({
49- return AOTI_TORCH_SUCCESS;53+ reinterpret_cast<c10_npu::NPUGuard*>(guard)->set_index(static_cast<c10::DeviceIndex>(device_index));
54+ });
50}55}
51 56 
52AOTITorchError aoti_torch_create_npu_stream_guard(57AOTITorchError aoti_torch_create_npu_stream_guard(
@@ -78,6 +83,8 @@ AOTITorchError aoti_torch_create_tensor_from_blob_npu(void* data, int64_t ndim,
78 AtenTensorHandle* ret_new_tensor)83 AtenTensorHandle* ret_new_tensor)
79{84{
80 AOTI_TORCH_CONVERT_EXCEPTION_TO_ERROR_CODE({85 AOTI_TORCH_CONVERT_EXCEPTION_TO_ERROR_CODE({
86+ TORCH_CHECK(ret_new_tensor != nullptr, "ret_new_tensor is nullptr");
87+ *ret_new_tensor = nullptr;
81 c10::IntArrayRef sizes(sizes_ptr, ndim);88 c10::IntArrayRef sizes(sizes_ptr, ndim);
82 c10::IntArrayRef strides(strides_ptr, ndim);89 c10::IntArrayRef strides(strides_ptr, ndim);
83 c10::Device device = c10_device(device_type, device_index);90 c10::Device device = c10_device(device_type, device_index);
@@ -96,8 +103,10 @@ AOTITorchError aoti_torch_create_tensor_from_blob_npu_v2(void* data, int64_t ndi
96 const uint8_t* opaque_metadata, int64_t opaque_metadata_size)103 const uint8_t* opaque_metadata, int64_t opaque_metadata_size)
97{104{
98 AOTI_TORCH_CONVERT_EXCEPTION_TO_ERROR_CODE({105 AOTI_TORCH_CONVERT_EXCEPTION_TO_ERROR_CODE({
106+ TORCH_CHECK(ret_new_tensor != nullptr, "ret_new_tensor is nullptr");
107+ *ret_new_tensor = nullptr;
99 if (layout == static_cast<int32_t>(at::kMkldnn)) {108 if (layout == static_cast<int32_t>(at::kMkldnn)) {
100- throw std::runtime_error("do not support mkldnn on npu.");109+ TORCH_CHECK(false, "do not support mkldnn on npu.");
101 } else {110 } else {
102 aoti_torch_create_tensor_from_blob_npu(data, ndim, sizes_ptr, strides_ptr, storage_offset, dtype,111 aoti_torch_create_tensor_from_blob_npu(data, ndim, sizes_ptr, strides_ptr, storage_offset, dtype,
103 device_type, device_index, ret_new_tensor);112 device_type, device_index, ret_new_tensor);