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())