# Copyright (c) 2026 SandAI. All Rights Reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. # You may obtain a copy of the License at # # http://www.apache.org/licenses/LICENSE-2.0 # # Unless required by applicable law or agreed to in writing, software # distributed under the License is distributed on an "AS IS" BASIS, # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. from enum import Enum from typing import Any, Dict, List import torch from magi_compiler.tokenflow.green_ctx import GreenCtxManager from torch import fx class FX_NODE_OP(Enum): PLACEHOLDER = "placeholder" GET_ATTR = "get_attr" CALL_FUNCTION = "call_function" CALL_METHOD = "call_method" CALL_MODULE = "call_module" OUTPUT = "output" class LaneType(Enum): COMPUTE = "compute" OTHERS = "others" class NodeOnlyExecutor: def __init__(self, graph_module: fx.GraphModule, device: torch.device): self.graph_module = graph_module self.device = device self.name_to_node = {node.name: node for node in graph_module.graph.nodes} @classmethod def replace_nodes_in_args(cls, args, value_map): if isinstance(args, torch.fx.Node): return value_map[args.name] elif isinstance(args, (list, tuple)): return type(args)(cls.replace_nodes_in_args(a, value_map) for a in args) elif isinstance(args, dict): return {k: cls.replace_nodes_in_args(v, value_map) for k, v in args.items()} else: return args def execute(self, node: fx.Node, value_map: Dict[str, Any], stream: torch.cuda.Stream = None) -> None: node_name = node.name if node.op == FX_NODE_OP.PLACEHOLDER.value: if node_name in value_map: return value_map[node_name] raise RuntimeError("PLACEHOLDER节点不应被执行。") args = self.replace_nodes_in_args(node.args, value_map) if node.args and node.args != () else () kwargs = self.replace_nodes_in_args(node.kwargs, value_map) if node.kwargs and node.kwargs != {} else {} with torch.cuda.stream(stream): # with nullcontext(): if node.op == FX_NODE_OP.GET_ATTR.value: attr_val = self.graph_module for attr in node.target.split("."): attr_val = getattr(attr_val, attr) result = attr_val elif node.op == FX_NODE_OP.CALL_FUNCTION.value: result = node.target(*args, **kwargs) elif node.op == FX_NODE_OP.CALL_METHOD.value: obj = args[0] method = getattr(obj, node.target) result = method(*args[1:], **kwargs) elif node.op == FX_NODE_OP.CALL_MODULE.value: submod = self.graph_module for mod_name in node.target.split("."): submod = getattr(submod, mod_name) result = submod(*args, **kwargs) elif node.op == FX_NODE_OP.OUTPUT.value: result = args[0] if len(args) == 1 else args else: raise NotImplementedError(f"不支持的op类型: {node.op}") assert result is not None, f"节点 {node_name} 执行未返回结果。" value_map[node_name] = result class GraphRawExecutor: def __init__(self, graph_module: fx.GraphModule, device: torch.device = None): # 基础属性初始化 self.graph = graph_module.graph self.module = graph_module self.device = device self.topological_nodes = [node for node in self.graph.nodes] self.name_to_node = {node.name: node for node in self.topological_nodes} self.node_executor = NodeOnlyExecutor(self.module, self.device) self.value_map: Dict[str, Any] = {} self.stream_map: Dict[str, torch.cuda.Stream] = {} for node in self.topological_nodes: if node.op == FX_NODE_OP.PLACEHOLDER.value: continue self.stream_map[node.name] = torch.cuda.default_stream(device=self.device) # cuda_graph_mgr().run(func, *args, layer_number=layer_number, **kwargs) # @cuda_graph_enable_if(condition=lambda: True) def execute(self, *inputs) -> Any: for idx, node in enumerate(self.topological_nodes): if node.op == FX_NODE_OP.PLACEHOLDER.value: self.value_map[node.name] = inputs[idx] continue self.node_executor.execute(node=node, value_map=self.value_map, stream=self.stream_map[node.name]) if node.op == FX_NODE_OP.OUTPUT.value: output_result = self.value_map[node.name] break return output_result def synchronize(self): for stream in self.stream_map.values(): if stream is not None: stream.synchronize() def cleanup(self): pass class GraphNormalExecutor: def __init__(self, graph_module: torch.fx.GraphModule, device: torch.device = None): self.graph = graph_module.graph self.module = graph_module self.device = device or torch.device("cuda" if torch.cuda.is_available() else "cpu") self.name_to_node = {node.name: node for node in self.graph.nodes} self.topological_names = [node.name for node in self.graph.nodes] self.value_map = {node.name: None for node in self.graph.nodes} self.dependencies = {node.name: [] for node in self.graph.nodes} self.rev_dependencies = {node.name: [] for node in self.graph.nodes} self._resolve_node_dependencies() self.stream_map: dict[str, torch.cuda.Stream] = {} self.event_map: dict[str, torch.cuda.Event] = {} for node_name in self.topological_names: if self.name_to_node[node_name].op == FX_NODE_OP.PLACEHOLDER.value: continue self.stream_map[node_name] = torch.cuda.Stream(device=self.device) self.event_map[node_name] = torch.cuda.Event(enable_timing=False, blocking=False) self.node_executor = NodeOnlyExecutor(self.module, self.device) def _resolve_node_dependencies(self): for node in self.graph.nodes: dep_node_names = [] def extract_deps(arg): if isinstance(arg, torch.fx.Node): dep_node_names.append(arg.name) elif isinstance(arg, (tuple, list)): for a in arg: extract_deps(a) elif isinstance(arg, dict): for v in arg.values(): extract_deps(v) extract_deps(node.args) extract_deps(node.kwargs) dep_node_names = list(set(dep_node_names)) self.dependencies[node.name] = dep_node_names for dep_name in dep_node_names: self.rev_dependencies[dep_name].append(node.name) def wait_for_dependencies(self, node_name: str): for dep_name in self.dependencies[node_name]: if self.name_to_node[dep_name].op == FX_NODE_OP.PLACEHOLDER.value: continue self.stream_map[node_name].wait_event(self.event_map[dep_name]) def _replace_nodes_in_args(self, args): if isinstance(args, torch.fx.Node): return self.value_map[args.name] elif isinstance(args, (tuple, list)): return type(args)(self._replace_nodes_in_args(a) for a in args) elif isinstance(args, dict): return {k: self._replace_nodes_in_args(v) for k, v in args.items()} else: return args def execute(self, *inputs) -> Any: self.value_map = {} placeholder_names = [node.name for node in self.graph.nodes if node.op == FX_NODE_OP.PLACEHOLDER.value] assert len(placeholder_names) == len(inputs), f"输入数量不匹配:图需要 {len(placeholder_names)} 个输入,但提供了 {len(inputs)} 个。" for i, node_name in enumerate(placeholder_names): self.value_map[node_name] = inputs[i] for idx, node_name in enumerate(self.topological_names): if self.name_to_node[node_name].op == FX_NODE_OP.PLACEHOLDER.value: continue node = self.name_to_node[node_name] stream = self.stream_map[node_name] self.wait_for_dependencies(node_name) self.node_executor.execute(node=node, value_map=self.value_map, stream=stream) self.event_map[node_name].record(stream) if node.op == FX_NODE_OP.OUTPUT.value: return self.value_map[node_name] raise RuntimeError("图中未找到OUTPUT节点,执行未完成。") def synchronize(self): for stream in self.stream_map.values(): stream.synchronize() def cleanup(self): pass class GraphStageConfig: def __init__(self, name: str, sm_dict: Dict[str, int], lane_node_dict: Dict[str, List[str]]): self.name = name self.lane_sm_dict = sm_dict self.lane_node_dict = lane_node_dict class GraphOptimizer: @staticmethod def generate_stages_per_op(graph: torch.fx.Graph) -> List[GraphStageConfig]: stages = [] for idx, node in enumerate(graph.nodes): if node.op == FX_NODE_OP.PLACEHOLDER.value: continue stage_name = f"stage_{idx}_{node.name}" stages.append( GraphStageConfig( name=stage_name, sm_dict={LaneType.COMPUTE.value: GreenCtxManager(0).max_sm}, lane_node_dict={LaneType.COMPUTE.value: [node.name]}, ) ) return stages @staticmethod def generate_stages_all_in_one(graph: torch.fx.Graph) -> List[GraphStageConfig]: all_node_names = [node.name for node in graph.nodes if node.op != FX_NODE_OP.PLACEHOLDER.value] return [ GraphStageConfig( name="stage_all_in_one", sm_dict={LaneType.COMPUTE.value: 132}, lane_node_dict={LaneType.COMPUTE.value: all_node_names}, ) ] class GraphStageExecutor: def __init__(self, graph_module: torch.fx.GraphModule, stage_configs: List[GraphStageConfig], device: torch.device = None): self.graph_module = graph_module self.graph = graph_module.graph self.stage_configs = stage_configs self.device = device or torch.device("cuda" if torch.cuda.is_available() else "cpu") self.name_to_node = {node.name: node for node in self.graph.nodes} self.value_map = {node.name: None for node in self.graph.nodes} self.stage_green_manager: Dict[str, GreenCtxManager] = {} self.stage_lane_stream: Dict[str, Dict[str, torch.cuda.Stream]] = {} self.stage_lane_event: Dict[str, Dict[str, torch.cuda.Event]] = {} for idx, stage_config in enumerate(stage_configs): stage_name = stage_config.name self.stage_green_manager[stage_name] = GreenCtxManager(device_index=self.device.index) self.stage_lane_stream[stage_name] = {} self.stage_lane_event[stage_name] = {} for lane_type in stage_config.lane_node_dict.keys(): cur_green_manager = self.stage_green_manager[stage_name] cur_sm_count = stage_config.lane_sm_dict.get(lane_type, 0) assert cur_sm_count >= 0, f"阶段 {stage_name} 的泳道 {lane_type} 的 SM 数量不能为负数" self.stage_lane_stream[stage_name][lane_type] = ( cur_green_manager.create_stream(sm_count=cur_sm_count) if cur_sm_count else None ) self.stage_lane_event[stage_name][lane_type] = torch.cuda.Event(blocking=False) self.input_nodes = [node for node in self.graph.nodes if node.op == FX_NODE_OP.PLACEHOLDER.value] self.output_node = next(node for node in self.graph.nodes if node.op == FX_NODE_OP.OUTPUT.value) self.node_executor = NodeOnlyExecutor(self.graph_module, self.device) def _replace_nodes_in_args(self, args): if isinstance(args, torch.fx.Node): return self.value_map[args.name] elif isinstance(args, (list, tuple)): return type(args)(self._replace_nodes_in_args(a) for a in args) elif isinstance(args, dict): return {k: self._replace_nodes_in_args(v) for k, v in args.items()} else: return args def wait_for_stage_dependencies(self, stage_name: str): stage_idx = next(i for i, sc in enumerate(self.stage_configs) if sc.name == stage_name) if stage_idx == 0: return curr_stage = self.stage_configs[stage_idx] prev_stage = self.stage_configs[stage_idx - 1] for curr_lane in curr_stage.lane_node_dict.keys(): curr_stream = self.stage_lane_stream[curr_stage.name][curr_lane] for prev_lane in prev_stage.lane_node_dict.keys(): prev_event = self.stage_lane_event[prev_stage.name][prev_lane] curr_stream.wait_event(prev_event) def execute(self, *inputs) -> Any: assert len(inputs) == len(self.input_nodes), "输入数量与PLACEHOLDER节点数量不匹配。" for idx, input_node in enumerate(self.input_nodes): self.value_map[input_node.name] = inputs[idx] for idx, stage in enumerate(self.stage_configs): stage_name = stage.name self.wait_for_stage_dependencies(stage_name) for lane_type, node_names in stage.lane_node_dict.items(): stream = self.stage_lane_stream[stage_name][lane_type] for node_name in node_names: node = self.name_to_node[node_name] self.node_executor.execute(node=node, value_map=self.value_map, stream=stream) event = self.stage_lane_event[stage_name][lane_type] event.record(stream) self.synchronize() return self.value_map[self.output_node.name] def synchronize(self): for stage_config in self.stage_configs: stage_name = stage_config.name for lane_type in stage_config.lane_node_dict.keys(): stream = self.stage_lane_stream[stage_name][lane_type] if stream is not None: stream.synchronize() def cleanup(self): for stage_manager in self.stage_green_manager.values(): stage_manager.cleanup()