MP_PDE / models /pde.py
yushuang88's picture
Upload folder using huggingface_hub
c92f17c verified
Raw
History Blame Contribute Delete
10.1 kB
"""Paper-priority MP-PDE model for experiment E3."""
from __future__ import annotations
from typing import Iterable, Sequence, Tuple
import torch
from torch import Tensor, nn
class Swish(nn.Module):
def forward(self, values: Tensor) -> Tensor:
return values * torch.sigmoid(values)
class TwoLayerMLP(nn.Module):
def __init__(self, input_dim: int, hidden_dim: int, output_dim: int):
super().__init__()
self.network = nn.Sequential(
nn.Linear(input_dim, hidden_dim),
Swish(),
nn.Linear(hidden_dim, output_dim),
Swish(),
)
def forward(self, values: Tensor) -> Tensor:
return self.network(values)
def periodic_neighbor_indices(num_nodes: int, offsets: Iterable[int], device: torch.device | None = None) -> Tensor:
"""Return source indices [num_nodes, num_neighbors] for each target node."""
offsets_tensor = torch.as_tensor(tuple(offsets), dtype=torch.long, device=device)
if num_nodes < 7:
raise ValueError(f"The six-neighbor periodic graph requires num_nodes>=7, found {num_nodes}")
if offsets_tensor.numel() != 6 or offsets_tensor.unique().numel() != 6 or torch.any(offsets_tensor == 0):
raise ValueError("MP-PDE E3 requires exactly six distinct non-zero neighbor offsets")
target = torch.arange(num_nodes, dtype=torch.long, device=device)[:, None]
return torch.remainder(target + offsets_tensor[None, :], num_nodes)
class MessagePassingLayer(nn.Module):
"""Equation (8)--(9) processor layer with sum aggregation."""
def __init__(self, hidden_dim: int, history_dim: int, parameter_dim: int, affine_norm: bool):
super().__init__()
edge_dim = 2 * hidden_dim + history_dim + 1 + parameter_dim
node_dim = 2 * hidden_dim + parameter_dim
self.edge_mlp = TwoLayerMLP(edge_dim, hidden_dim, hidden_dim)
self.node_mlp = TwoLayerMLP(node_dim, hidden_dim, hidden_dim)
self.norm = nn.InstanceNorm1d(hidden_dim, affine=affine_norm, track_running_stats=False)
def forward(
self,
hidden: Tensor,
history: Tensor,
x: Tensor,
parameters: Tensor,
neighbors: Tensor,
domain_length: float,
) -> Tensor:
batch, num_nodes, hidden_dim = hidden.shape
num_neighbors = neighbors.shape[1]
source_hidden = hidden[:, neighbors, :]
target_hidden = hidden[:, :, None, :].expand(-1, -1, num_neighbors, -1)
source_history = history[:, neighbors, :]
history_difference = history[:, :, None, :] - source_history
source_x = x[:, neighbors]
displacement = x[:, :, None] - source_x
displacement = torch.remainder(displacement + 0.5 * domain_length, domain_length) - 0.5 * domain_length
theta_edges = parameters[:, None, None, :].expand(-1, num_nodes, num_neighbors, -1)
edge_input = torch.cat(
(target_hidden, source_hidden, history_difference, displacement[..., None], theta_edges), dim=-1
)
messages = self.edge_mlp(edge_input)
aggregated = messages.sum(dim=2)
theta_nodes = parameters[:, None, :].expand(-1, num_nodes, -1)
node_update = self.node_mlp(torch.cat((hidden, aggregated, theta_nodes), dim=-1))
return self.norm((hidden + node_update).transpose(1, 2)).transpose(1, 2)
class TemporalDecoder(nn.Module):
def __init__(
self,
hidden_dim: int,
time_window: int,
middle_channels: int = 8,
kernels: Sequence[int] = (16, 26),
strides: Sequence[int] = (3, 1),
):
super().__init__()
if len(kernels) != 2 or len(strides) != 2:
raise ValueError("The paper-priority decoder requires exactly two convolutions")
length_after_first = (hidden_dim - int(kernels[0])) // int(strides[0]) + 1
output_length = (length_after_first - int(kernels[1])) // int(strides[1]) + 1
if output_length != time_window:
raise ValueError(
f"Decoder does not close hidden={hidden_dim} to K={time_window}: output length={output_length}"
)
self.network = nn.Sequential(
nn.Conv1d(1, middle_channels, kernel_size=int(kernels[0]), stride=int(strides[0])),
Swish(),
nn.Conv1d(middle_channels, 1, kernel_size=int(kernels[1]), stride=int(strides[1])),
)
def forward(self, hidden: Tensor) -> Tensor:
batch, num_nodes, hidden_dim = hidden.shape
decoded = self.network(hidden.reshape(batch * num_nodes, 1, hidden_dim))
return decoded.reshape(batch, num_nodes, decoded.shape[-1])
class MPPDESolver(nn.Module):
"""Message-passing neural solver mapping K E3 states to the next K states."""
def __init__(
self,
time_window: int = 25,
hidden_dim: int = 164,
message_passing_layers: int = 6,
neighbor_offsets: Sequence[int] = (-3, -2, -1, 1, 2, 3),
domain_length: float = 16.0,
final_time: float = 4.0,
parameter_maxima: Sequence[float] = (3.0, 0.4, 1.0),
scale_coordinates: bool = True,
scale_parameters: bool = True,
instance_norm_affine: bool = False,
decoder_middle_channels: int = 8,
decoder_kernels: Sequence[int] = (16, 26),
decoder_strides: Sequence[int] = (3, 1),
):
super().__init__()
if time_window <= 0 or hidden_dim <= 0 or message_passing_layers <= 0:
raise ValueError("time_window, hidden_dim, and message_passing_layers must be positive")
self.time_window = int(time_window)
self.hidden_dim = int(hidden_dim)
self.neighbor_offsets: Tuple[int, ...] = tuple(int(value) for value in neighbor_offsets)
self.domain_length = float(domain_length)
self.final_time = float(final_time)
self.scale_coordinates = bool(scale_coordinates)
self.scale_parameters = bool(scale_parameters)
maxima = torch.tensor(tuple(float(value) for value in parameter_maxima), dtype=torch.float32)
if maxima.shape != (3,) or torch.any(maxima <= 0.0):
raise ValueError("parameter_maxima must contain three positive values")
self.register_buffer("parameter_maxima", maxima, persistent=False)
self.encoder = TwoLayerMLP(self.time_window + 1 + 1 + 3, self.hidden_dim, self.hidden_dim)
self.processor = nn.ModuleList(
MessagePassingLayer(self.hidden_dim, self.time_window, 3, instance_norm_affine)
for _ in range(int(message_passing_layers))
)
self.decoder = TemporalDecoder(
self.hidden_dim, self.time_window, decoder_middle_channels, decoder_kernels, decoder_strides
)
def _canonicalize_inputs(
self, history: Tensor, x: Tensor, current_time: Tensor | float, parameters: Tensor
) -> Tuple[Tensor, Tensor, Tensor, Tensor]:
if history.ndim != 3:
raise ValueError(f"history must have shape [B,N,K], found {tuple(history.shape)}")
batch, num_nodes, window = history.shape
if window != self.time_window:
raise ValueError(f"Expected history K={self.time_window}, found {window}")
if num_nodes < 7:
raise ValueError(f"MP-PDE requires N>=7, found {num_nodes}")
if parameters.shape != (batch, 3):
raise ValueError(f"parameters must have shape [{batch},3], found {tuple(parameters.shape)}")
if x.ndim == 1:
if x.shape[0] != num_nodes:
raise ValueError(f"x length {x.shape[0]} does not match N={num_nodes}")
x = x[None, :].expand(batch, -1)
elif x.shape != (batch, num_nodes):
raise ValueError(f"x must have shape [N] or [B,N], found {tuple(x.shape)}")
time = torch.as_tensor(current_time, dtype=history.dtype, device=history.device)
if time.ndim == 0:
time = time.expand(batch)
elif time.shape == (batch, 1):
time = time[:, 0]
elif time.shape != (batch,):
raise ValueError(f"current_time must be scalar or shape [B], found {tuple(time.shape)}")
return history, x.to(history), time, parameters.to(history)
def forward(
self,
history: Tensor,
x: Tensor,
current_time: Tensor | float,
parameters: Tensor,
dt: Tensor | float,
*,
return_derivative: bool = False,
) -> Tensor | Tuple[Tensor, Tensor]:
history, x, current_time, parameters = self._canonicalize_inputs(history, x, current_time, parameters)
batch, num_nodes, _ = history.shape
scaled_x = x / self.domain_length if self.scale_coordinates else x
scaled_time = current_time / self.final_time if self.scale_coordinates else current_time
scaled_parameters = parameters / self.parameter_maxima.to(parameters) if self.scale_parameters else parameters
time_feature = scaled_time[:, None, None].expand(-1, num_nodes, 1)
parameter_features = scaled_parameters[:, None, :].expand(-1, num_nodes, -1)
encoded = torch.cat((history, scaled_x[..., None], time_feature, parameter_features), dim=-1)
hidden = self.encoder(encoded)
neighbors = periodic_neighbor_indices(num_nodes, self.neighbor_offsets, history.device)
for layer in self.processor:
hidden = layer(hidden, history, x, scaled_parameters, neighbors, self.domain_length)
derivative = self.decoder(hidden)
step = torch.as_tensor(dt, dtype=history.dtype, device=history.device)
if step.ndim == 0:
step = step.expand(batch)
elif step.shape != (batch,):
raise ValueError(f"dt must be scalar or shape [B], found {tuple(step.shape)}")
offsets = torch.arange(1, self.time_window + 1, dtype=history.dtype, device=history.device)
delta_times = step[:, None, None] * offsets[None, None, :]
prediction = history[:, :, -1:] + delta_times * derivative
return (prediction, derivative) if return_derivative else prediction