已合并
[fix] cat refact of dsl change caused multi stream codegen fail #35144
zzll创建于 5月9日
[fix] cat refact of dsl change caused multi stream codegen fail #35144
已合并
共 4 个文件变更+169-75
| @@ -1,3 +1,5 @@ | |||
| 1 | +from unittest import skip | ||
| 2 | + | ||
| 1 | import torch | 3 | import torch |
| 2 | from torch.testing._internal.common_utils import run_tests, parametrize, instantiate_parametrized_tests | 4 | from torch.testing._internal.common_utils import run_tests, parametrize, instantiate_parametrized_tests |
| 3 | from testutils import TestUtils | 5 | from testutils import TestUtils |
| @@ -8,6 +10,8 @@ class TestVarMean(TestUtils): | |||
| 8 | def op_calc(self, input_element, dim): | 10 | def op_calc(self, input_element, dim): |
| 9 | return torch.var_mean(input_element, dim) | 11 | return torch.var_mean(input_element, dim) |
| 10 | 12 | ||
| 13 | + | ||
| 14 | + | ||
| 11 | 15 | ||
| 12 | 16 | ||
| 13 | 17 | ||
| @@ -19,7 +23,6 @@ class TestVarMean(TestUtils): | |||
| 19 | 23 | ||
| 20 | compiled_op_calc = torch.compile(self.op_calc, backend="inductor", dynamic=False) | 24 | compiled_op_calc = torch.compile(self.op_calc, backend="inductor", dynamic=False) |
| 21 | inductor_var, inductor_mean = compiled_op_calc(input_element, dim) | 25 | inductor_var, inductor_mean = compiled_op_calc(input_element, dim) |
| 22 | - | ||
| 23 | self.assertEqual(std_var, inductor_var, atol=1e-1, rtol=1e-1, equal_nan=True) | 26 | self.assertEqual(std_var, inductor_var, atol=1e-1, rtol=1e-1, equal_nan=True) |
| 24 | self.assertEqual(std_mean, inductor_mean, atol=1e-1, rtol=1e-1, equal_nan=True) | 27 | self.assertEqual(std_mean, inductor_mean, atol=1e-1, rtol=1e-1, equal_nan=True) |
| 25 | 28 | ||
| @@ -160,7 +160,7 @@ else: | |||
| 160 | 160 | ||
| 161 | pre_grad_custom_pass_fuc() | 161 | pre_grad_custom_pass_fuc() |
| 162 | post_grad_custom_pass_fuc() | 162 | post_grad_custom_pass_fuc() |
| 163 | - if os.environ.get("PARALLEL_SCHEDULER_OPTIMIZAR", "false").lower() == "true": | 163 | + if os.environ.get("ENABLE_PARALLEL_SCHEDULER", "false").lower() == "true": |
| 164 | from .fx_passes.parallel_scheduler_pass import parallel_scheduler | 164 | from .fx_passes.parallel_scheduler_pass import parallel_scheduler |
| 165 | 165 | ||
| 166 | parallel_scheduler() | 166 | parallel_scheduler() |
| @@ -1,29 +1,42 @@ | |||
| 1 | -import torch | ||
| 2 | -import torch_npu | ||
| 3 | -from torch._inductor import config | ||
| 4 | -from torch._inductor.utils import sympy_product | ||
| 5 | -from torch._inductor.ir import ExternKernelOut, InputBuffer, ReinterpretView | ||
| 6 | -import os | ||
| 7 | -from torch._inductor.virtualized import V | ||
| 8 | -from torch._inductor.utils import device_need_guard | ||
| 9 | -from torch.utils._ordered_set import OrderedSet | ||
| 10 | import traceback | 1 | import traceback |
| 11 | -from ..config import log | ||
| 12 | import typing | 2 | import typing |
| 13 | -from torch._inductor.scheduler import BaseSchedulerNode, FusedSchedulerNode, SchedulerNode, NopKernelSchedulerNode, ForeachKernelSchedulerNode, ExternKernelSchedulerNode | 3 | +from collections import defaultdict |
| 4 | + | ||
| 5 | +import torch | ||
| 6 | +import torch._inductor.ir as ir | ||
| 7 | +from torch._inductor import config | ||
| 14 | from torch._inductor.codegen.cuda_combined_scheduling import CUDACombinedScheduling | 8 | from torch._inductor.codegen.cuda_combined_scheduling import CUDACombinedScheduling |
| 15 | from torch._inductor.codegen.simd import SIMDScheduling | 9 | from torch._inductor.codegen.simd import SIMDScheduling |
| 16 | -from collections import defaultdict | 10 | +from torch._inductor.ir import ExternKernelOut, InputBuffer, ReinterpretView |
| 11 | +from torch._inductor.scheduler import ( | ||
| 12 | + BaseSchedulerNode, | ||
| 13 | + ExternKernelSchedulerNode, | ||
| 14 | + ForeachKernelSchedulerNode, | ||
| 15 | + FusedSchedulerNode, | ||
| 16 | + NopKernelSchedulerNode, | ||
| 17 | + SchedulerNode, | ||
| 18 | +) | ||
| 19 | +from torch._inductor.utils import device_need_guard, sympy_product | ||
| 20 | +from torch._inductor.virtualized import V | ||
| 21 | +from torch.utils._ordered_set import OrderedSet | ||
| 22 | + | ||
| 23 | +from ..codegen.catlass.catlass_kernel import CATLASSTemplateBuffer | ||
| 24 | +from ..config import log | ||
| 17 | from .parallelism_strategy_framework import ParallelGroupingStrategy | 25 | from .parallelism_strategy_framework import ParallelGroupingStrategy |
| 18 | from .utils.fx_pass_level import GroupType | 26 | from .utils.fx_pass_level import GroupType |
| 19 | -import torch._inductor.ir as ir | 27 | +from .utils.schedule_node_utils import is_multi_stream |
| 20 | -from ..codegen.catlass.catlass_kernel import CATLASSTemplateBuffer | ||
| 21 | 28 | ||
| 22 | 29 | ||
| 23 | def parallel_scheduler(): | 30 | def parallel_scheduler(): |
| 31 | + """ | ||
| 32 | + Patch the Scheduler._codegen method to support multi stream parallel scheduling. | ||
| 33 | + """ | ||
| 24 | original_codegen = torch._inductor.scheduler.Scheduler._codegen | 34 | original_codegen = torch._inductor.scheduler.Scheduler._codegen |
| 25 | - | 35 | + |
| 26 | def patched_codegen(self, nodes: list[BaseSchedulerNode]) -> None: | 36 | def patched_codegen(self, nodes: list[BaseSchedulerNode]) -> None: |
| 37 | + if not is_multi_stream(): | ||
| 38 | + original_codegen(self, nodes) | ||
| 39 | + return | ||
| 27 | parallel_strategy = ParallelGroupingStrategy() | 40 | parallel_strategy = ParallelGroupingStrategy() |
| 28 | groups = parallel_strategy.execute_strategy(nodes) | 41 | groups = parallel_strategy.execute_strategy(nodes) |
| 29 | is_parallel = True | 42 | is_parallel = True |
| @@ -36,6 +49,7 @@ def parallel_scheduler(): | |||
| 36 | return | 49 | return |
| 37 | if config.check_stack_no_cycles_TESTING_ONLY: | 50 | if config.check_stack_no_cycles_TESTING_ONLY: |
| 38 | import torch._dynamo.convert_frame | 51 | import torch._dynamo.convert_frame |
| 52 | + | ||
| 39 | stack = traceback.extract_stack() | 53 | stack = traceback.extract_stack() |
| 40 | seen = OrderedSet() | 54 | seen = OrderedSet() |
| 41 | for frame in reversed(stack): | 55 | for frame in reversed(stack): |
| @@ -47,9 +61,11 @@ def parallel_scheduler(): | |||
| 47 | key = (frame.filename, frame.lineno) | 61 | key = (frame.filename, frame.lineno) |
| 48 | 62 | ||
| 49 | if key in seen: | 63 | if key in seen: |
| 50 | - raise AssertionError(f"Duplicate stack frame {frame.filename}:{frame.lineno}; " | 64 | + raise AssertionError( |
| 51 | - "did you add a decorator to one of the functions in this stack " | 65 | + f"Duplicate stack frame {frame.filename}:{frame.lineno}; " |
| 52 | - "trace? If so, try using a context manager instead.") | 66 | + "did you add a decorator to one of the functions in this stack " |
| 67 | + "trace? If so, try using a context manager instead." | ||
| 68 | + ) | ||
| 53 | seen.add(key) | 69 | seen.add(key) |
| 54 | 70 | ||
| 55 | self.current_device = None | 71 | self.current_device = None |
| @@ -70,13 +86,9 @@ def parallel_scheduler(): | |||
| 70 | wrapper.pre_define_buffer = [] | 86 | wrapper.pre_define_buffer = [] |
| 71 | 87 | ||
| 72 | stream_vars = { | 88 | stream_vars = { |
| 73 | - key: f"stream_group_{key.lower()}" | 89 | + key: f"stream_group_{key.lower()}" for key in list(groups.keys()) |
| 74 | - for key in list(groups.keys()) | ||
| 75 | - } | ||
| 76 | - event_vars = { | ||
| 77 | - key: f"event_group_{key.lower()}" | ||
| 78 | - for key in list(groups.keys()) | ||
| 79 | } | 90 | } |
| 91 | + event_vars = {key: f"event_group_{key.lower()}" for key in list(groups.keys())} | ||
| 80 | stream_codegen(wrapper, groups, stream_vars, event_vars) | 92 | stream_codegen(wrapper, groups, stream_vars, event_vars) |
| 81 | 93 | ||
| 82 | node_to_group_id = {} | 94 | node_to_group_id = {} |
| @@ -92,7 +104,7 @@ def parallel_scheduler(): | |||
| 92 | 104 | ||
| 93 | group_to_buffer = build_buffer_producer_group(nodes, node_to_group_id) | 105 | group_to_buffer = build_buffer_producer_group(nodes, node_to_group_id) |
| 94 | 106 | ||
| 95 | - waited_events = set() | 107 | + waited_events = OrderedSet() |
| 96 | need_main_stream_event = [] | 108 | need_main_stream_event = [] |
| 97 | for node in nodes: | 109 | for node in nodes: |
| 98 | node_name = node.get_name() | 110 | node_name = node.get_name() |
| @@ -111,7 +123,9 @@ def parallel_scheduler(): | |||
| 111 | ): | 123 | ): |
| 112 | self.flush() | 124 | self.flush() |
| 113 | if device != self.current_device: | 125 | if device != self.current_device: |
| 114 | - if self.current_device and device_need_guard(self.current_device.type): | 126 | + if self.current_device and device_need_guard( |
| 127 | + self.current_device.type | ||
| 128 | + ): | ||
| 115 | V.graph.wrapper_code.codegen_device_guard_exit() | 129 | V.graph.wrapper_code.codegen_device_guard_exit() |
| 116 | self.current_device = device | 130 | self.current_device = device |
| 117 | if device_need_guard(device.type): | 131 | if device_need_guard(device.type): |
| @@ -120,10 +134,12 @@ def parallel_scheduler(): | |||
| 120 | V.graph.wrapper_code.codegen_device_guard_enter(device.index) | 134 | V.graph.wrapper_code.codegen_device_guard_enter(device.index) |
| 121 | self.buffer_names_to_free.update(node.last_usage) | 135 | self.buffer_names_to_free.update(node.last_usage) |
| 122 | tab_value = current_indent_level * wrapper.wrapper_call.tabwidth | 136 | tab_value = current_indent_level * wrapper.wrapper_call.tabwidth |
| 123 | - | 137 | + |
| 124 | if target_stream != current_stream: | 138 | if target_stream != current_stream: |
| 125 | if current_stream and current_stream != "main_stream" and current_event: | 139 | if current_stream and current_stream != "main_stream" and current_event: |
| 126 | - wrapper.writeline(" " * (tab_value) + f"{current_event}.record({current_stream})") | 140 | + wrapper.writeline( |
| 141 | + " " * (tab_value) + f"{current_event}.record({current_stream})" | ||
| 142 | + ) | ||
| 127 | need_main_stream_event.append(current_event) | 143 | need_main_stream_event.append(current_event) |
| 128 | 144 | ||
| 129 | while current_indent_level > 0: | 145 | while current_indent_level > 0: |
| @@ -133,15 +149,18 @@ def parallel_scheduler(): | |||
| 133 | add_need_pre_buf_define(group_to_buffer, group_id, wrapper) | 149 | add_need_pre_buf_define(group_to_buffer, group_id, wrapper) |
| 134 | wrapper.writeline(f"with torch_npu.npu.stream({target_stream}):") | 150 | wrapper.writeline(f"with torch_npu.npu.stream({target_stream}):") |
| 135 | current_indent_level += 1 | 151 | current_indent_level += 1 |
| 136 | - wrapper.writeline(" " * (current_indent_level * wrapper.wrapper_call.tabwidth) + f"{target_stream}.wait_event(main_event)") | 152 | + wrapper.writeline( |
| 153 | + " " * (current_indent_level * wrapper.wrapper_call.tabwidth) | ||
| 154 | + + f"{target_stream}.wait_event(main_event)" | ||
| 155 | + ) | ||
| 137 | 156 | ||
| 138 | if target_stream == "main_stream": | 157 | if target_stream == "main_stream": |
| 139 | for event_name in need_main_stream_event: | 158 | for event_name in need_main_stream_event: |
| 140 | if event_name not in waited_events: | 159 | if event_name not in waited_events: |
| 141 | intent_tab = 0 | 160 | intent_tab = 0 |
| 142 | wrapper.writeline( | 161 | wrapper.writeline( |
| 143 | - " " * (intent_tab * wrapper.wrapper_call.tabwidth) + | 162 | + " " * (intent_tab * wrapper.wrapper_call.tabwidth) |
| 144 | - f"{target_stream}.wait_event({event_name})" | 163 | + + f"{target_stream}.wait_event({event_name})" |
| 145 | ) | 164 | ) |
| 146 | waited_events.add(event_name) | 165 | waited_events.add(event_name) |
| 147 | 166 | ||
| @@ -152,8 +171,12 @@ def parallel_scheduler(): | |||
| 152 | build_multi_stream_buf_intent(node, tab_value, wrapper) | 171 | build_multi_stream_buf_intent(node, tab_value, wrapper) |
| 153 | 172 | ||
| 154 | if node.is_template(): | 173 | if node.is_template(): |
| 155 | - prologue, template_node, epilogue = node.get_prologue_template_epilogue(list(node.get_nodes())) | 174 | + prologue, template_node, epilogue = node.get_prologue_template_epilogue( |
| 156 | - self.get_backend(device).codegen_template(template_node, epilogue, prologue) | 175 | + list(node.get_nodes()) |
| 176 | + ) | ||
| 177 | + self.get_backend(device).codegen_template( | ||
| 178 | + template_node, epilogue, prologue | ||
| 179 | + ) | ||
| 157 | elif node.is_extern(): | 180 | elif node.is_extern(): |
| 158 | node = typing.cast(ExternKernelSchedulerNode, node) | 181 | node = typing.cast(ExternKernelSchedulerNode, node) |
| 159 | self.codegen_extern_call(node) | 182 | self.codegen_extern_call(node) |
| @@ -184,12 +207,12 @@ def parallel_scheduler(): | |||
| 184 | self.flush() | 207 | self.flush() |
| 185 | if current_stream != "main_stream" and current_event: | 208 | if current_stream != "main_stream" and current_event: |
| 186 | wrapper.writeline( | 209 | wrapper.writeline( |
| 187 | - " " * (current_indent_level * wrapper.wrapper_call.tabwidth) + | 210 | + " " * (current_indent_level * wrapper.wrapper_call.tabwidth) |
| 188 | - f"{current_event}.record({current_stream})" | 211 | + + f"{current_event}.record({current_stream})" |
| 189 | ) | 212 | ) |
| 190 | wrapper.writeline( | 213 | wrapper.writeline( |
| 191 | - " " * (current_indent_level * wrapper.wrapper_call.tabwidth) + | 214 | + " " * (current_indent_level * wrapper.wrapper_call.tabwidth) |
| 192 | - f"# end last stream {current_stream} context" | 215 | + + f"# end last stream {current_stream} context" |
| 193 | ) | 216 | ) |
| 194 | 217 | ||
| 195 | while current_indent_level > 0: | 218 | while current_indent_level > 0: |
| @@ -197,8 +220,11 @@ def parallel_scheduler(): | |||
| 197 | 220 | ||
| 198 | wrapper.writeline("\n# Wait for all parallel streams to complete") | 221 | wrapper.writeline("\n# Wait for all parallel streams to complete") |
| 199 | for key, event_name in event_vars.items(): | 222 | for key, event_name in event_vars.items(): |
| 200 | - if event_name not in waited_events and event_name != event_vars[GroupType.MAIN.name]: | 223 | + if ( |
| 201 | - wrapper.writeline(f"main_stream.wait_event({event_name})") | 224 | + event_name not in waited_events |
| 225 | + and event_name != event_vars[GroupType.MAIN.name] | ||
| 226 | + ): | ||
| 227 | + wrapper.writeline(f"main_stream.wait_event({event_name})") | ||
| 202 | 228 | ||
| 203 | if self.current_device and device_need_guard(self.current_device.type): | 229 | if self.current_device and device_need_guard(self.current_device.type): |
| 204 | V.graph.wrapper_code.codegen_device_guard_exit() | 230 | V.graph.wrapper_code.codegen_device_guard_exit() |
| @@ -207,22 +233,26 @@ def parallel_scheduler(): | |||
| 207 | torch._inductor.scheduler.Scheduler._codegen = patched_codegen | 233 | torch._inductor.scheduler.Scheduler._codegen = patched_codegen |
| 208 | 234 | ||
| 209 | original_codegen_assert = torch._inductor.ir.ExternKernel.codegen_size_asserts | 235 | original_codegen_assert = torch._inductor.ir.ExternKernel.codegen_size_asserts |
| 236 | + | ||
| 210 | def patched_codegen_size_asserts(self, wrapper) -> None: | 237 | def patched_codegen_size_asserts(self, wrapper) -> None: |
| 211 | if hasattr(wrapper, "buffer_define_multi_stream"): | 238 | if hasattr(wrapper, "buffer_define_multi_stream"): |
| 212 | - buffer_define_multi_stream = getattr(wrapper, "buffer_define_multi_stream") | 239 | + buffer_define_multi_stream = wrapper.buffer_define_multi_stream |
| 213 | if config.size_asserts and not V.graph.cpp_wrapper: | 240 | if config.size_asserts and not V.graph.cpp_wrapper: |
| 214 | if sympy_product(self.get_size()) == 0: | 241 | if sympy_product(self.get_size()) == 0: |
| 215 | return | 242 | return |
| 216 | size = V.graph.wrapper_code.codegen_shape_tuple(self.get_size()) | 243 | size = V.graph.wrapper_code.codegen_shape_tuple(self.get_size()) |
| 217 | stride = V.graph.wrapper_code.codegen_shape_tuple(self.get_stride()) | 244 | stride = V.graph.wrapper_code.codegen_shape_tuple(self.get_stride()) |
| 218 | multi_stream_intent = "" | 245 | multi_stream_intent = "" |
| 219 | - if self.get_name() in buffer_define_multi_stream.keys(): | 246 | + if self.get_name() in buffer_define_multi_stream: |
| 220 | - multi_stream_intent = " " * buffer_define_multi_stream.get(self.get_name()) | 247 | + multi_stream_intent = " " * buffer_define_multi_stream.get( |
| 248 | + self.get_name() | ||
| 249 | + ) | ||
| 221 | wrapper.writeline( | 250 | wrapper.writeline( |
| 222 | f"{multi_stream_intent}assert_size_stride({self.get_name()}, {size}, {stride})" | 251 | f"{multi_stream_intent}assert_size_stride({self.get_name()}, {size}, {stride})" |
| 223 | ) | 252 | ) |
| 224 | else: | 253 | else: |
| 225 | original_codegen_assert(self, wrapper) | 254 | original_codegen_assert(self, wrapper) |
| 255 | + | ||
| 226 | torch._inductor.ir.ExternKernel.codegen_size_asserts = patched_codegen_size_asserts | 256 | torch._inductor.ir.ExternKernel.codegen_size_asserts = patched_codegen_size_asserts |
| 227 | 257 | ||
| 228 | 258 | ||
| @@ -241,7 +271,19 @@ def stream_codegen(wrapper, groups, stream_vars, event_vars): | |||
| 241 | def add_need_pre_buf_define(group_to_buffer, group_id, wrapper): | 271 | def add_need_pre_buf_define(group_to_buffer, group_id, wrapper): |
| 242 | need_define_buffers = group_to_buffer[group_id] | 272 | need_define_buffers = group_to_buffer[group_id] |
| 243 | for buf in need_define_buffers: | 273 | for buf in need_define_buffers: |
| 244 | - if isinstance(buf, (ir.ComputedBuffer, ir.ExternKernelOut, CATLASSTemplateBuffer)) and hasattr(buf, "name") and buf.name not in V.graph.removed_buffers: | 274 | + if ( |
| 275 | + isinstance( | ||
| 276 | + buf, | ||
| 277 | + ( | ||
| 278 | + ir.ConcatKernel, | ||
| 279 | + ir.ComputedBuffer, | ||
| 280 | + ir.ExternKernelOut, | ||
| 281 | + CATLASSTemplateBuffer, | ||
| 282 | + ), | ||
| 283 | + ) | ||
| 284 | + and hasattr(buf, "name") | ||
| 285 | + and buf.name not in V.graph.removed_buffers | ||
| 286 | + ): | ||
| 245 | line = wrapper.make_buffer_allocation(buf) | 287 | line = wrapper.make_buffer_allocation(buf) |
| 246 | wrapper.writeline(line) | 288 | wrapper.writeline(line) |
| 247 | wrapper.pre_define_buffer.append(buf.name) | 289 | wrapper.pre_define_buffer.append(buf.name) |
| @@ -249,16 +291,20 @@ def add_need_pre_buf_define(group_to_buffer, group_id, wrapper): | |||
| 249 | 291 | ||
| 250 | def build_multi_stream_buf_intent(node, tab_value, wrapper): | 292 | def build_multi_stream_buf_intent(node, tab_value, wrapper): |
| 251 | if not hasattr(node, "multi_stream_intent"): | 293 | if not hasattr(node, "multi_stream_intent"): |
| 252 | - setattr(node, "multi_stream_intent", tab_value) | 294 | + node.multi_stream_intent = tab_value |
| 253 | - if hasattr(node, 'get_nodes'): | 295 | + if hasattr(node, "get_nodes"): |
| 254 | nodes = node.get_nodes() | 296 | nodes = node.get_nodes() |
| 255 | for n in nodes: | 297 | for n in nodes: |
| 256 | if not hasattr(n, "multi_stream_intent"): | 298 | if not hasattr(n, "multi_stream_intent"): |
| 257 | - setattr(n, "multi_stream_intent", tab_value) | 299 | + n.multi_stream_intent = tab_value |
| 258 | - | 300 | + |
| 259 | if getattr(node, "multi_stream_name", None) not in (None, "main_stream.npu_stream"): | 301 | if getattr(node, "multi_stream_name", None) not in (None, "main_stream.npu_stream"): |
| 260 | - if hasattr(node, "node") and isinstance(node.node, ExternKernelOut) and getattr(node.node, "python_kernel_name", None) is not None: | 302 | + if ( |
| 261 | - args = set() | 303 | + hasattr(node, "nodes") |
| 304 | + and isinstance(node.node, ExternKernelOut) | ||
| 305 | + and getattr(node.node, "python_kernel_name", None) is not None | ||
| 306 | + ): | ||
| 307 | + args = OrderedSet() | ||
| 262 | inputs = getattr(node.node, "inputs", None) | 308 | inputs = getattr(node.node, "inputs", None) |
| 263 | if inputs: | 309 | if inputs: |
| 264 | for input in inputs: | 310 | for input in inputs: |
| @@ -267,9 +313,9 @@ def build_multi_stream_buf_intent(node, tab_value, wrapper): | |||
| 267 | elif isinstance(input, ReinterpretView): | 313 | elif isinstance(input, ReinterpretView): |
| 268 | args.add(input.get_name()) | 314 | args.add(input.get_name()) |
| 269 | data = { | 315 | data = { |
| 270 | - "python_kernel_name": getattr(node.node, "python_kernel_name"), | 316 | + "python_kernel_name": node.node.python_kernel_name, |
| 271 | "args": args, | 317 | "args": args, |
| 272 | - "multi_stream_intent": tab_value | 318 | + "multi_stream_intent": tab_value, |
| 273 | } | 319 | } |
| 274 | wrapper.extern_node_intent_multi_stream.append(data) | 320 | wrapper.extern_node_intent_multi_stream.append(data) |
| 275 | if len(node.last_usage) > 0: | 321 | if len(node.last_usage) > 0: |
| @@ -281,7 +327,7 @@ def build_multi_stream_buf_intent(node, tab_value, wrapper): | |||
| 281 | for output in node_outputs: | 327 | for output in node_outputs: |
| 282 | wrapper.buffer_define_multi_stream[output.node.name] = tab_value | 328 | wrapper.buffer_define_multi_stream[output.node.name] = tab_value |
| 283 | elif getattr(node, "node", None) is None and hasattr(node, "snodes"): | 329 | elif getattr(node, "node", None) is None and hasattr(node, "snodes"): |
| 284 | - snodes = getattr(node, "snodes") | 330 | + snodes = node.snodes |
| 285 | for sn in snodes: | 331 | for sn in snodes: |
| 286 | sn_outputs = getattr(sn, "outputs", None) | 332 | sn_outputs = getattr(sn, "outputs", None) |
| 287 | if sn_outputs: | 333 | if sn_outputs: |
| @@ -299,19 +345,35 @@ def pre_node_multi_stream_flag(nodes, node_to_group_id, stream_vars): | |||
| 299 | else: | 345 | else: |
| 300 | target_stream = stream_vars[group_id] | 346 | target_stream = stream_vars[group_id] |
| 301 | if not hasattr(node, "multi_stream_name"): | 347 | if not hasattr(node, "multi_stream_name"): |
| 302 | - setattr(node, "multi_stream_name", f"{target_stream}.npu_stream") | 348 | + node.multi_stream_name = f"{target_stream}.npu_stream" |
| 303 | - if hasattr(node, 'get_nodes'): | 349 | + if hasattr(node, "get_nodes"): |
| 304 | nodes = node.get_nodes() | 350 | nodes = node.get_nodes() |
| 305 | for n in nodes: | 351 | for n in nodes: |
| 306 | if not hasattr(n, "multi_stream_name"): | 352 | if not hasattr(n, "multi_stream_name"): |
| 307 | - setattr(n, "multi_stream_name", f"{target_stream}.npu_stream") | 353 | + n.multi_stream_name = f"{target_stream}.npu_stream" |
| 308 | 354 | ||
| 309 | 355 | ||
| 310 | def build_buffer_producer_group(nodes, node_to_group_id): | 356 | def build_buffer_producer_group(nodes, node_to_group_id): |
| 311 | buffer_producer_group = {} | 357 | buffer_producer_group = {} |
| 358 | + group_to_buffer_names = defaultdict(list) | ||
| 312 | all_buffers = [] | 359 | all_buffers = [] |
| 360 | + concat_bufs = defaultdict(list) | ||
| 313 | for node in nodes: | 361 | for node in nodes: |
| 314 | gid = node_to_group_id.get(node.get_name()) | 362 | gid = node_to_group_id.get(node.get_name()) |
| 363 | + if ( | ||
| 364 | + hasattr(node, "node") | ||
| 365 | + and hasattr(node.node, "name") | ||
| 366 | + and isinstance(node.node, ir.ConcatKernel) | ||
| 367 | + and hasattr(node.node, "inputs") | ||
| 368 | + ): | ||
| 369 | + buf_name = node.node.name | ||
| 370 | + inputs = node.node.inputs | ||
| 371 | + for input in inputs: | ||
| 372 | + if isinstance( | ||
| 373 | + input, | ||
| 374 | + (ir.ComputedBuffer, ir.ExternKernelOut, CATLASSTemplateBuffer), | ||
| 375 | + ) and hasattr(input, "name"): | ||
| 376 | + concat_bufs[buf_name].append(input.name) | ||
| 315 | 377 | ||
| 316 | outputs = getattr(node, "outputs", None) | 378 | outputs = getattr(node, "outputs", None) |
| 317 | if outputs: | 379 | if outputs: |
| @@ -331,4 +393,22 @@ def build_buffer_producer_group(nodes, node_to_group_id): | |||
| 331 | for buf in all_buffers: | 393 | for buf in all_buffers: |
| 332 | if hasattr(buf, "name") and buf.name == buf_name: | 394 | if hasattr(buf, "name") and buf.name == buf_name: |
| 333 | group_to_buffer[gid].append(buf) | 395 | group_to_buffer[gid].append(buf) |
| 334 | - return group_to_buffer | 396 | + group_to_buffer_names[gid].append(buf_name) |
| 397 | + if len(concat_bufs) > 0: | ||
| 398 | + for concat_buf_name, input_buf_names in concat_bufs.items(): | ||
| 399 | + input_buf_names_set = OrderedSet(input_buf_names) | ||
| 400 | + for gid, buf_names in group_to_buffer_names.items(): | ||
| 401 | + buf_names_set = OrderedSet(buf_names) | ||
| 402 | + if gid != GroupType.MAIN.name and ( | ||
| 403 | + input_buf_names_set.issubset(buf_names_set) | ||
| 404 | + or buf_names_set.issubset(input_buf_names_set) | ||
| 405 | + ): | ||
| 406 | + for buf in all_buffers: | ||
| 407 | + if hasattr(buf, "name") and buf.name == concat_buf_name: | ||
| 408 | + group_to_buffer[gid].append(buf) | ||
| 409 | + group_to_buffer_names[gid].append(concat_buf_name) | ||
| 410 | + origin_gid = buffer_producer_group[concat_buf_name] | ||
| 411 | + group_to_buffer[origin_gid].remove(buf) | ||
| 412 | + group_to_buffer_names[origin_gid].remove(concat_buf_name) | ||
| 413 | + | ||
| 414 | + return group_to_buffer | ||
| @@ -1,21 +1,33 @@ | |||
| 1 | import os | 1 | import os |
| 2 | + | ||
| 3 | +from torch._inductor.codegen.cpp_wrapper_cpu import CppWrapperCpu | ||
| 4 | +from torch._inductor.codegen.cpp_wrapper_gpu import CppWrapperGpu | ||
| 2 | from torch._inductor.scheduler import BaseSchedulerNode | 5 | from torch._inductor.scheduler import BaseSchedulerNode |
| 3 | -from typing import Dict, List, Set | 6 | +from torch._inductor.virtualized import V |
| 7 | +from torch.utils._ordered_set import OrderedSet | ||
| 8 | + | ||
| 9 | +from ...codegen.cpp_wrapper import CppWrapperNpu | ||
| 4 | 10 | ||
| 5 | 11 | ||
| 6 | def is_multi_stream(): | 12 | def is_multi_stream(): |
| 7 | - return os.environ.get("PARALLEL_SCHEDULER_OPTIMIZAR", "false").lower() == "true" | 13 | + wrapper = V.graph.wrapper_code |
| 14 | + is_cpp_wrapper = isinstance(wrapper, (CppWrapperNpu, CppWrapperCpu, CppWrapperGpu)) | ||
| 15 | + env_parallel = ( | ||
| 16 | + os.environ.get("ENABLE_PARALLEL_SCHEDULER", "false").lower() == "true" | ||
| 17 | + ) | ||
| 18 | + return env_parallel and not is_cpp_wrapper | ||
| 19 | + | ||
| 8 | 20 | ||
| 9 | def find_first_overlap(pre_nodes, first_group_nodes): | 21 | def find_first_overlap(pre_nodes, first_group_nodes): |
| 10 | for i, item in enumerate(reversed(first_group_nodes)): | 22 | for i, item in enumerate(reversed(first_group_nodes)): |
| 11 | if item in pre_nodes: | 23 | if item in pre_nodes: |
| 12 | return len(first_group_nodes) - 1 - i | 24 | return len(first_group_nodes) - 1 - i |
| 13 | return None | 25 | return None |
| 14 | - | 26 | + |
| 15 | 27 | ||
| 16 | def make_disjoint(anc_sets): | 28 | def make_disjoint(anc_sets): |
| 17 | result = [] | 29 | result = [] |
| 18 | - seen = set() | 30 | + seen = OrderedSet() |
| 19 | for s in anc_sets: | 31 | for s in anc_sets: |
| 20 | cleaned = s - seen | 32 | cleaned = s - seen |
| 21 | result.append(cleaned) | 33 | result.append(cleaned) |
| @@ -24,12 +36,11 @@ def make_disjoint(anc_sets): | |||
| 24 | 36 | ||
| 25 | 37 | ||
| 26 | def get_predecessors( | 38 | def get_predecessors( |
| 27 | - node: BaseSchedulerNode, | 39 | + node: BaseSchedulerNode, name_to_node: dict[str, BaseSchedulerNode] |
| 28 | - name_to_node: Dict[str, BaseSchedulerNode] | 40 | +) -> OrderedSet[BaseSchedulerNode]: |
| 29 | -) -> Set[BaseSchedulerNode]: | 41 | + preds = OrderedSet() |
| 30 | - preds = set() | ||
| 31 | 42 | ||
| 32 | - if hasattr(node, 'mpi_node') and node.mpi_node is not None: | 43 | + if hasattr(node, "mpi_node") and node.mpi_node is not None: |
| 33 | for pred_node in node.mpi_node.pred_nodes: | 44 | for pred_node in node.mpi_node.pred_nodes: |
| 34 | if isinstance(pred_node, BaseSchedulerNode): | 45 | if isinstance(pred_node, BaseSchedulerNode): |
| 35 | preds.add(pred_node) | 46 | preds.add(pred_node) |
| @@ -37,26 +48,26 @@ def get_predecessors( | |||
| 37 | if preds: | 48 | if preds: |
| 38 | return preds | 49 | return preds |
| 39 | 50 | ||
| 40 | - if hasattr(node, 'ancestors') and node.ancestors: | 51 | + if hasattr(node, "ancestors") and node.ancestors: |
| 41 | for name in node.ancestors: | 52 | for name in node.ancestors: |
| 42 | if name in name_to_node: | 53 | if name in name_to_node: |
| 43 | preds.add(name_to_node[name]) | 54 | preds.add(name_to_node[name]) |
| 44 | 55 | ||
| 45 | - if hasattr(node, 'read_writes') and node.read_writes.reads: | 56 | + if hasattr(node, "read_writes") and node.read_writes.reads: |
| 46 | for dep in node.read_writes.reads: | 57 | for dep in node.read_writes.reads: |
| 47 | - if hasattr(dep, 'name') and dep.name in name_to_node: | 58 | + if hasattr(dep, "name") and dep.name in name_to_node: |
| 48 | preds.add(name_to_node[dep.name]) | 59 | preds.add(name_to_node[dep.name]) |
| 49 | 60 | ||
| 50 | - if hasattr(node, 'unmet_dependencies'): | 61 | + if hasattr(node, "unmet_dependencies"): |
| 51 | for dep in node.unmet_dependencies: | 62 | for dep in node.unmet_dependencies: |
| 52 | - if hasattr(dep, 'name') and dep.name in name_to_node: | 63 | + if hasattr(dep, "name") and dep.name in name_to_node: |
| 53 | preds.add(name_to_node[dep.name]) | 64 | preds.add(name_to_node[dep.name]) |
| 54 | return preds | 65 | return preds |
| 55 | 66 | ||
| 56 | 67 | ||
| 57 | -def get_successors_names(node: BaseSchedulerNode) -> List[str]: | 68 | +def get_successors_names(node: BaseSchedulerNode) -> list[str]: |
| 58 | - succ_nodes = set() | 69 | + succ_nodes = OrderedSet() |
| 59 | - if hasattr(node, 'mpi_node') and node.mpi_node is not None: | 70 | + if hasattr(node, "mpi_node") and node.mpi_node is not None: |
| 60 | for succ in node.mpi_node.succ_nodes: | 71 | for succ in node.mpi_node.succ_nodes: |
| 61 | if isinstance(succ, BaseSchedulerNode): | 72 | if isinstance(succ, BaseSchedulerNode): |
| 62 | succ_nodes.add(succ) | 73 | succ_nodes.add(succ) |