File size: 1,194 Bytes
3de20f8 | 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 | 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())
|