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