已合并
inductor: use simt_template as kernel default compile_mode #37939
stonexxx创建于 6月9日
inductor: use simt_template as kernel default compile_mode #37939
已合并
共 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 | ||
| 58 | if __name__ == "__main__": | 61 | if __name__ == "__main__": |
| @@ -1,8 +1,13 @@ | |||
| 1 | from torch._inductor.codegen.common import DeviceOpOverrides, register_device_op_overrides | 1 | from 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 | ||
| 4 | class NewNPUDeviceOpOverrides(DeviceOpOverrides): | 6 | class 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 = None | 1274 | self.golden_var_list = None |
| 1275 | self.reduce_analysis = None | 1275 | self.reduce_analysis = None |
| 1276 | self.load_store_indexing = None | 1276 | self.load_store_indexing = None |
| 1277 | - self.npu_kernel_type = NPUKernelType.SIMD | 1277 | + 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_linter | 1281 | self.current_subblock_axis = set() # noqa: set_linter |
| 1279 | self.node_schedule = self.features.node_schedule | 1282 | self.node_schedule = self.features.node_schedule |
| 1280 | self.decide_codegen_dims_in_kernel() | 1283 | self.decide_codegen_dims_in_kernel() |