Team Ai
Modelpublic

OneScience-Group/MeshGraphNet

sourceHugging Faceapache-2.0updated 2mo agoView on Hugging Face
0likes22downloads
meshgraphnet.py307 linesDownload Raw Back to model
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