已合并
inductor: use simt_template as kernel default compile_mode #37939
stonexxx创建于 6月9日
inductor: use simt_template as kernel default compile_mode #37939
已合并
stonexxx创建于 6月9日
3 个文件变更+16-5
@@ -48,11 +48,14 @@ class TestNpuDevice(TestCase):
48 excepted = "torch.npu.set_device(0)"48 excepted = "torch.npu.set_device(0)"
49 self.assertEqual(result, excepted)49 self.assertEqual(result, excepted)
50 50 
51- def test_import_get_raw_stream_as(self):
52 overrides = NewNPUDeviceOpOverrides()51 overrides = NewNPUDeviceOpOverrides()
53- result = overrides.import_get_raw_stream_as("test_name")52+ test_name = "test_name_npu"
54- excepted = "from torch_npu._C import _npu_getCurrentRawStream as test_name"53+ result = overrides.import_get_raw_stream_as(test_name)
55- self.assertEqual(result, excepted)54+ expected = f"from torch._C import _npu_getCurrentRawStream as {test_name}"
55+ import torch_npu
56+ if hasattr(torch_npu._C, "_npu_getCurrentRawStreamNoWait"):
57+ expected = f"from torch_npu._C import _npu_getCurrentRawStreamNoWait as {test_name}"
58+ self.assertEqual(result, expected)
56 59 
57 60 
58if __name__ == "__main__":61if __name__ == "__main__":
@@ -1,8 +1,13 @@
1from torch._inductor.codegen.common import DeviceOpOverrides, register_device_op_overrides1from torch._inductor.codegen.common import DeviceOpOverrides, register_device_op_overrides
2+import torch_npu
3+from torch_npu._inductor.codegen.catlass.catlass_utils import try_import_catlass
2 4 
3 5 
4class NewNPUDeviceOpOverrides(DeviceOpOverrides):6class NewNPUDeviceOpOverrides(DeviceOpOverrides):
5 def import_get_raw_stream_as(self, name):7 def import_get_raw_stream_as(self, name):
8+ enabled_catlass = try_import_catlass()
9+ if not enabled_catlass and hasattr(torch_npu._C, "_npu_getCurrentRawStreamNoWait"):
10+ return f"from torch_npu._C import _npu_getCurrentRawStreamNoWait as {name}"
6 return f"from torch_npu._C import _npu_getCurrentRawStream as {name}"11 return f"from torch_npu._C import _npu_getCurrentRawStream as {name}"
7 12 
8 def set_device(self, device_idx):13 def set_device(self, device_idx):
@@ -1274,7 +1274,10 @@ class NPUIndexTritonKernel(TritonKernel):
1274 self.golden_var_list = None1274 self.golden_var_list = None
1275 self.reduce_analysis = None1275 self.reduce_analysis = None
1276 self.load_store_indexing = None1276 self.load_store_indexing = None
1277- self.npu_kernel_type = NPUKernelType.SIMD1277+ if npu_config.is_ascend950:
1278+ self.npu_kernel_type = NPUKernelType.SIMT_TEMPLATE
1279+ else:
1280+ self.npu_kernel_type = NPUKernelType.SIMD
1278 self.current_subblock_axis = set() # noqa: set_linter1281 self.current_subblock_axis = set() # noqa: set_linter
1279 self.node_schedule = self.features.node_schedule1282 self.node_schedule = self.features.node_schedule
1280 self.decide_codegen_dims_in_kernel()1283 self.decide_codegen_dims_in_kernel()