Grafneuronnät (GNN)
Kunna förklara meddelandepassning och träna ett GNN på en liten graf.
Förkunskaper
Intuition
Ett grafneuronnät arbetar på data där relationerna är lika viktiga som objekten: molekyler, sociala nätverk, vägnät, citeringsgrafer — och kunskapsgrafer som den här plattformens.
Meddelandepassning är hela idén, och den upprepas i varje lager:
1. Varje nod skickar ett meddelande till sina grannar
2. Varje nod aggregerar de meddelanden den fått (summa, medel, max)
3. Varje nod uppdaterar sin representation utifrån aggregatet och sitt eget tillstånd
Efter lager har varje nod information från alla noder inom steg. Det är därför djupet i ett GNN betyder något helt annat än i ett CNN: det styr hur långt information reser i grafen, inte hur abstrakta features blir.
Aggregeringen måste vara permutationsinvariant — grannarna har ingen ordning. Därför summa, medelvärde eller max, aldrig konkatenering.
Formellt
Tre uppgiftstyper, och de kräver olika uppdelning:
| Typ | Exempel | Uppdelning |
|---|---|---|
| Nodnivå | klassificera användare i ett nätverk | maskera noder i samma graf |
| Kantnivå | förutsäg länkar | maskera kanter |
| Grafnivå | molekylens egenskap | dela grafer mellan mängder |
Nodnivå-uppdelningen är speciell: hela grafen syns under träning, men bara vissa noders etiketter används. Det kallas transduktiv inlärning och gör att vanliga läckagekontroller inte räcker — grannarnas features syns även för testnoder.
Två problem som gör djupa GNN svåra:
| Problem | Vad som händer |
|---|---|
| Översläpning (over-smoothing) | efter många lager konvergerar alla nodrepresentationer mot samma vektor |
| Översnävning (over-squashing) | information från exponentiellt många noder pressas genom en smal flaskhals |
Därför är de flesta GNN grunda — 2 till 4 lager. Att gå djupare kräver residualkopplingar, normalisering och ibland grafomskrivning.
Arkitekturfamiljer:
| Modell | Aggregering |
|---|---|
| GCN | normaliserat medelvärde över grannar |
| GraphSAGE | sampla ett fast antal grannar — skalbart till stora grafer |
| GAT | attention över grannar, inlärda vikter |
| GIN | summa plus MLP; bevisat lika uttrycksfullt som Weisfeiler–Lehman-testet |
GIN-resultatet är teoretiskt viktigt: det sätter en övre gräns för vad meddelandepassande GNN kan skilja på. Två grafer som WL-testet inte kan särskilja kan inte heller ett standard-GNN särskilja — till exempel vissa regelbundna grafer.
Praktiska råd:
- Börja med en baslinje utan graf. En MLP på nodernas egna features, eller logistisk regression. Slår GNN:en inte den tillför grafen ingenting.
- Använd 2–3 lager.
- Lägg till nodgrad och andra strukturella features — ofta större effekt än en finare arkitektur.
- För stora grafer: grannsampling (GraphSAGE) eller klusterbaserad batchning.
Råd 1 är det som oftast hoppas över, och det avgör om projektet är meningsfullt.
Kod
import torch, torch.nn as nn, torch.nn.functional as F
class GCNLager(nn.Module):
"""h' = σ(D^-1/2 Â D^-1/2 h W), Â = A + I"""
def __init__(self, cin, cut):
super().__init__()
self.lin = nn.Linear(cin, cut)
def forward(self, h, A_norm):
return A_norm @ self.lin(h)
def normalisera(A):
A_hat = A + torch.eye(len(A)) # lägg till självloopar
grad = A_hat.sum(1)
d = torch.diag(grad.pow(-0.5))
return d @ A_hat @ d
class GCN(nn.Module):
def __init__(self, in_dim, dolt, n_klasser, lager=2, dropout=0.5):
super().__init__()
dims = [in_dim] + [dolt] * (lager - 1) + [n_klasser]
self.lager = nn.ModuleList([GCNLager(dims[i], dims[i + 1]) for i in range(lager)])
self.dropout = dropout
def forward(self, h, A_norm):
for i, lg in enumerate(self.lager):
h = lg(h, A_norm)
if i < len(self.lager) - 1:
h = F.dropout(F.relu(h), self.dropout, self.training)
return h
# Liten graf: fyra noder, en fyrkant
A = torch.tensor([[0., 1., 0., 1.],
[1., 0., 1., 0.],
[0., 1., 0., 1.],
[1., 0., 1., 0.]])
An = normalisera(A)
h = torch.randn(4, 8)
modell = GCN(8, 16, 2)
print(modell(h, An).shape) # torch.Size([4, 2])
# Översläpning: mät hur lika representationerna blir med djupet
def likhet(h):
hn = F.normalize(h, dim=-1)
s = hn @ hn.T
n = len(h)
return float((s.sum() - s.diag().sum()) / (n * (n - 1)))
h = torch.randn(4, 8)
for djup in (1, 2, 4, 8, 16):
m = GCN(8, 8, 8, lager=djup)
m.eval()
with torch.no_grad():
print(f" {djup:>2} lager: genomsnittlig cosinuslikhet {likhet(m(h, An)):.4f}")
# 1 lager: 0.1...
# 16 lager: nära 1.0 ← alla noder har blivit likadana: översläpning
# Transduktiv nodklassificering: maskera ETIKETTER, inte noder
def dela_noder(n, andel_train=0.6, andel_val=0.2, fro=0):
g = torch.Generator().manual_seed(fro)
perm = torch.randperm(n, generator=g)
a, b = int(n * andel_train), int(n * (andel_train + andel_val))
mask = lambda idx: torch.zeros(n, dtype=torch.bool).index_fill_(0, idx, True)
return mask(perm[:a]), mask(perm[a:b]), mask(perm[b:])
# Baslinje UTAN graf — kör alltid denna först
class MLPBaslinje(nn.Module):
def __init__(self, in_dim, dolt, n_klasser):
super().__init__()
self.f = nn.Sequential(nn.Linear(in_dim, dolt), nn.ReLU(),
nn.Dropout(0.5), nn.Linear(dolt, n_klasser))
def forward(self, h, A_norm=None):
return self.f(h)
# Slår GCN inte denna tillför grafstrukturen ingenting till uppgiften.
Behärskning innebär
- Förklarar meddelandepassning
- Vet varför djupa GNN är svåra
- Väljer uppgiftstyp och uppdelning korrekt
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — Semi-Supervised Classification with Graph Convolutional Networks — arXiv (öppen åtkomst; licens per artikel)
- arXiv — How Powerful are Graph Neural Networks? (GIN) — arXiv (öppen åtkomst; licens per artikel)
- PyTorch Geometric — dokumentation (MIT) — MIT