from __future__ import annotations import torch from torch import nn from torch.nn import functional as F class MeshGraphGCN(nn.Module): def __init__(self, features: int = 8) -> None: super().__init__() self.input = nn.Linear(features, 32, bias=False) self.hidden = nn.Linear(32, 16, bias=False) self.output = nn.Linear(16, 2) def forward( self, features: torch.Tensor, normalized_adjacency: torch.Tensor, ) -> torch.Tensor: hidden = normalized_adjacency @ features hidden = F.gelu(self.input(hidden)) hidden = F.dropout(hidden, p=0.15, training=self.training) hidden = normalized_adjacency @ hidden hidden = F.gelu(self.hidden(hidden)) return self.output(hidden) def normalize_adjacency(adjacency: torch.Tensor) -> torch.Tensor: with_self_loops = adjacency + torch.eye(len(adjacency)) degree = with_self_loops.sum(dim=1).clamp(min=1) inverse_sqrt = degree.pow(-0.5) return inverse_sqrt[:, None] * with_self_loops * inverse_sqrt[None, :] def parameter_count(model: nn.Module) -> int: return sum(parameter.numel() for parameter in model.parameters())