OneScience-Group/MeshGraphNet
022
1import torch2import torch.nn as nn3from torch import Tensor4 5try:6 import dgl # noqa: F401 for docs7 from dgl import DGLGraph8except ImportError:9 raise ImportError(10 "Mesh Graph Net requires the DGL library. Install the "11 )12from dataclasses import dataclass13from itertools import chain14from typing import Callable, List, Tuple, Union15 16import onescience # noqa: F401 for docs17from onescience.modules.edge.mesh_edge_block import MeshEdgeBlock18from onescience.modules.mlp.mesh_graph_mlp import MeshGraphMLP19from onescience.modules.node.mesh_node_block import MeshNodeBlock20 21from onescience.modules.utils.gnnlayer_utils import CuGraphCSC, set_checkpoint_fn22from onescience.modules.layer.activations import get_activation23from onescience.modules.meta import ModelMetaData24from onescience.modules.module import Module25 26 27@dataclass28class MetaData(ModelMetaData):29 name: str = "MeshGraphNet"30 # Optimization, no JIT as DGLGraph causes trouble31 jit: bool = False32 cuda_graphs: bool = False33 amp_cpu: bool = False34 amp_gpu: bool = True35 torch_fx: bool = False36 # Inference37 onnx: bool = False38 # Physics informed39 func_torch: bool = True40 auto_grad: bool = True41 42 43class MeshGraphNet(Module):44 """45 MeshGraphNet 网络架构。46 47 该模型基于 "Learning mesh-based simulation with graph networks" (Pfaff et al., 2020) 实现。48 它采用 Encode-Process-Decode 架构:49 1. **Encoder**: 将节点和边的物理特征映射到高维隐空间。50 2. **Processor**: 通过多层消息传递(Message Passing)在图中传播信息,更新节点和边的隐状态。51 3. **Decoder**: 将处理后的节点特征解码回物理空间(例如加速度或速度增量)。52 53 本实现使用 MeshGraphMLP、MeshEdgeBlock 和 MeshNodeBlock 构建。54 55 Args:56 input_dim_nodes (int): 输入节点特征的维度。57 input_dim_edges (int): 输入边特征的维度。58 output_dim (int): 输出特征的维度(通常是节点状态的更新量)。59 processor_size (int, optional): 消息传递块(Processor Block)的数量。默认值: 15。60 mlp_activation_fn (Union[str, List[str]], optional): MLP 中使用的激活函数。默认值: 'relu'。61 num_layers_node_processor (int, optional): 处理器中节点更新 MLP 的层数。默认值: 2。62 num_layers_edge_processor (int, optional): 处理器中边更新 MLP 的层数。默认值: 2。63 hidden_dim_processor (int, optional): 处理器中隐层的特征维度。默认值: 128。64 hidden_dim_node_encoder (int, optional): 节点编码器的隐层维度。默认值: 128。65 num_layers_node_encoder (Union[int, None], optional): 节点编码器的层数。如果为 None,则不使用编码器。默认值: 2。66 hidden_dim_edge_encoder (int, optional): 边编码器的隐层维度。默认值: 128。67 num_layers_edge_encoder (Union[int, None], optional): 边编码器的层数。如果为 None,则不使用编码器。默认值: 2。68 hidden_dim_node_decoder (int, optional): 节点解码器的隐层维度。默认值: 128。69 num_layers_node_decoder (Union[int, None], optional): 节点解码器的层数。如果为 None,则不使用解码器。默认值: 2。70 aggregation (str, optional): 消息聚合方式,可选 "sum", "mean" 等。默认值: "sum"。71 do_concat_trick (bool, optional): 是否使用拼接优化技巧 (MLP+idx+sum) 以节省显存。默认值: False。72 num_processor_checkpoint_segments (int, optional): 梯度检查点 (Gradient Checkpointing) 的分段数。0 表示禁用。默认值: 0。73 recompute_activation (bool, optional): 是否重计算激活函数以节省显存。默认值: False。74 75 形状:76 输入 node_features: (N, input_dim_nodes),其中 N 为节点总数。77 输入 edge_features: (M, input_dim_edges),其中 M 为边总数。78 输入 graph: DGLGraph 或 CuGraphCSC,定义图拓扑结构。79 输出: (N, output_dim),解码后的节点物理量。80 81 """82 83 def __init__(84 self,85 input_dim_nodes: int,86 input_dim_edges: int,87 output_dim: int,88 processor_size: int = 15,89 mlp_activation_fn: Union[str, List[str]] = "relu",90 num_layers_node_processor: int = 2,91 num_layers_edge_processor: int = 2,92 hidden_dim_processor: int = 128,93 hidden_dim_node_encoder: int = 128,94 num_layers_node_encoder: Union[int, None] = 2,95 hidden_dim_edge_encoder: int = 128,96 num_layers_edge_encoder: Union[int, None] = 2,97 hidden_dim_node_decoder: int = 128,98 num_layers_node_decoder: Union[int, None] = 2,99 aggregation: str = "sum",100 do_concat_trick: bool = False,101 num_processor_checkpoint_segments: int = 0,102 recompute_activation: bool = False,103 ):104 super().__init__(meta=MetaData())105 106 activation_fn = get_activation(mlp_activation_fn)107 108 # 1. Edge Encoder109 self.edge_encoder = MeshGraphMLP(110 input_dim=input_dim_edges,111 output_dim=hidden_dim_processor,112 hidden_dim=hidden_dim_edge_encoder,113 hidden_layers=num_layers_edge_encoder,114 activation_fn=activation_fn,115 norm_type="LayerNorm",116 recompute_activation=recompute_activation,117 )118 119 # 2. Node Encoder120 self.node_encoder = MeshGraphMLP(121 input_dim=input_dim_nodes,122 output_dim=hidden_dim_processor,123 hidden_dim=hidden_dim_node_encoder,124 hidden_layers=num_layers_node_encoder,125 activation_fn=activation_fn,126 norm_type="LayerNorm",127 recompute_activation=recompute_activation,128 )129 130 # 3. Node Decoder131 self.node_decoder = MeshGraphMLP(132 input_dim=hidden_dim_processor,133 output_dim=output_dim,134 hidden_dim=hidden_dim_node_decoder,135 hidden_layers=num_layers_node_decoder,136 activation_fn=activation_fn,137 norm_type=None,138 recompute_activation=recompute_activation,139 )140 141 # 4. Processor (Core GNN)142 self.processor = MeshGraphNetProcessor(143 processor_size=processor_size,144 input_dim_node=hidden_dim_processor,145 input_dim_edge=hidden_dim_processor,146 num_layers_node=num_layers_node_processor,147 num_layers_edge=num_layers_edge_processor,148 aggregation=aggregation,149 norm_type="LayerNorm",150 activation_fn=activation_fn,151 do_concat_trick=do_concat_trick,152 num_processor_checkpoint_segments=num_processor_checkpoint_segments,153 )154 155 def forward(156 self,157 node_features: Tensor,158 edge_features: Tensor,159 graph: Union[DGLGraph, List[DGLGraph], CuGraphCSC],160 ) -> Tensor:161 edge_features = self.edge_encoder(edge_features)162 node_features = self.node_encoder(node_features)163 x = self.processor(node_features, edge_features, graph)164 x = self.node_decoder(x)165 return x166 167 168class MeshGraphNetProcessor(nn.Module):169 """170 MeshGraphNet 核心处理器 (Processor)。171 172 该模块由一系列堆叠的消息传递块 (Message Passing Blocks) 组成。173 每个块包含两个步骤:174 1. **Edge Block**: 使用 MeshEdgeBlock 更新边特征。175 2. **Node Block**: 使用 MeshNodeBlock 聚合边信息并更新节点特征。176 177 支持梯度检查点 (Gradient Checkpointing) 以减少大规模图训练时的显存占用。178 179 Args:180 processor_size (int, optional): 处理器包含的消息传递层数。默认值: 15。181 input_dim_node (int, optional): 输入节点特征维度。默认值: 128。182 input_dim_edge (int, optional): 输入边特征维度。默认值: 128。183 num_layers_node (int, optional): 节点更新 MLP 的层数。默认值: 2。184 num_layers_edge (int, optional): 边更新 MLP 的层数。默认值: 2。185 aggregation (str, optional): 消息聚合方式 ("sum", "mean" 等)。默认值: "sum"。186 norm_type (str, optional): 归一化类型。默认值: "LayerNorm"。187 activation_fn (nn.Module, optional): 激活函数。默认值: nn.ReLU()。188 do_concat_trick (bool, optional): 是否启用显存优化技巧。默认值: False。189 num_processor_checkpoint_segments (int, optional): 梯度检查点分段数。默认值: 0 (禁用)。190 191 形状:192 输入 node_features: (N, input_dim_node)193 输入 edge_features: (M, input_dim_edge)194 输入 graph: DGLGraph195 输出: (N, input_dim_node) - 仅返回更新后的节点特征。196 197 """198 199 def __init__(200 self,201 processor_size: int = 15,202 input_dim_node: int = 128,203 input_dim_edge: int = 128,204 num_layers_node: int = 2,205 num_layers_edge: int = 2,206 aggregation: str = "sum",207 norm_type: str = "LayerNorm",208 activation_fn: nn.Module = nn.ReLU(),209 do_concat_trick: bool = False,210 num_processor_checkpoint_segments: int = 0,211 ):212 super().__init__()213 self.processor_size = processor_size214 self.num_processor_checkpoint_segments = num_processor_checkpoint_segments215 216 edge_blocks = []217 node_blocks = []218 219 for _ in range(self.processor_size):220 edge_blocks.append(221 MeshEdgeBlock(222 input_dim_nodes=input_dim_node,223 input_dim_edges=input_dim_edge,224 output_dim=input_dim_edge,225 hidden_dim=input_dim_edge,226 hidden_layers=num_layers_edge,227 activation_fn=activation_fn,228 norm_type=norm_type,229 do_concat_trick=do_concat_trick,230 recompute_activation=False231 )232 )233 node_blocks.append(234 MeshNodeBlock(235 aggregation=aggregation,236 input_dim_nodes=input_dim_node,237 input_dim_edges=input_dim_edge,238 output_dim=input_dim_node,239 hidden_dim=input_dim_node,240 hidden_layers=num_layers_node,241 activation_fn=activation_fn,242 norm_type=norm_type,243 recompute_activation=False244 )245 )246 247 # 按照 Edge -> Node 的顺序交替排列248 layers = list(chain(*zip(edge_blocks, node_blocks)))249 250 self.processor_layers = nn.ModuleList(layers)251 self.num_processor_layers = len(self.processor_layers)252 self.set_checkpoint_segments(self.num_processor_checkpoint_segments)253 254 def set_checkpoint_segments(self, checkpoint_segments: int):255 if checkpoint_segments > 0:256 if self.num_processor_layers % checkpoint_segments != 0:257 raise ValueError(258 "Processor layers must be a multiple of checkpoint_segments"259 )260 segment_size = self.num_processor_layers // checkpoint_segments261 self.checkpoint_segments = []262 for i in range(0, self.num_processor_layers, segment_size):263 self.checkpoint_segments.append((i, i + segment_size))264 self.checkpoint_fn = set_checkpoint_fn(True)265 else:266 self.checkpoint_fn = set_checkpoint_fn(False)267 self.checkpoint_segments = [(0, self.num_processor_layers)]268 269 def run_function(270 self, segment_start: int, segment_end: int271 ) -> Callable[272 [Tensor, Tensor, Union[DGLGraph, List[DGLGraph]]], Tuple[Tensor, Tensor]273 ]:274 segment = self.processor_layers[segment_start:segment_end]275 276 def custom_forward(277 node_features: Tensor,278 edge_features: Tensor,279 graph: Union[DGLGraph, List[DGLGraph]],280 ) -> Tuple[Tensor, Tensor]:281 for module in segment:282 edge_features, node_features = module(283 edge_features, node_features, graph284 )285 return edge_features, node_features286 287 return custom_forward288 289 @torch.jit.unused290 def forward(291 self,292 node_features: Tensor,293 edge_features: Tensor,294 graph: Union[DGLGraph, List[DGLGraph], CuGraphCSC],295 ) -> Tensor:296 for segment_start, segment_end in self.checkpoint_segments:297 edge_features, node_features = self.checkpoint_fn(298 self.run_function(segment_start, segment_end),299 node_features,300 edge_features,301 graph,302 use_reentrant=False,303 preserve_rng_state=False,304 )305 306 return node_features307 