jiadisu
Switch back to Docker SDK with local pkgs
e6066e8
Raw
History Blame Contribute Delete
14.8 kB
# 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()