Spaces:
Paused
Paused
| # 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} | |
| 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: | |
| 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 | |
| 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() | |