>_ DevTrendsnl

Taal

Home

Talen

Secties

Frontend Backend Mobiel DevOps AI / ML GameDev Blockchain Embedded Beveiliging
Python

Graph Neural Networks in PyTorch Zonder Extra Pijn en Aangepaste Workarounds

De meeste ontwikkelaars zijn gewend om met duidelijke, reguliere datastructuren te werken. Een afbeelding is een raster van pixels en tekst past gemakkelijk in een reeks tokens. Reguliere convolutionele netwerken of transformers verwerken dit perfect.

Problemen ontstaan wanneer data fundamenteel niet-lineair is. Neem banktransacties, verbindingen tussen gebruikers in sociale netwerken, of chemische formules van complexe moleculen. Grafen zijn overal aan het werk. Als je probeert ze aan reguliere PyTorch te voeren, moet je handmatig worstelen met sparse matrices, complexe lussen schrijven voor node aggregatie, en indices bijhouden.

Dit is precies waar PyTorch Geometric (afgekort PyG) voor is gemaakt. Het is een extensie voor PyTorch die alle low-level wiskunde voor je afhandelt en een intuïtieve API biedt voor het werken met graph neural networks.

Hoe PyG is Gestructureerd

Als je al weet hoe je modellen in PyTorch schrijft, dan voel je je binnen een paar uur thuis in PyG. Het framework volgt dezelfde principes: dezelfde modules torch.nn, vertrouwde trainingslussen, en directe werk met tensors.

Alle magie van graph networks draait om het Message Passing concept. Elke vertex in een graaf verzamelt informatie van zijn buren, combineert het, en werkt zijn eigen staat bij.

PyG biedt een base class MessagePassing die deze logica abstraheert. Binnenin hoef je maar drie dingen te definiëren:

  • Hoe een bericht wordt gevormd vanuit een buur-node
  • Hoe deze berichten worden geaggregeerd (sum, mean, max)
  • Hoe de node zelf wordt bijgewerkt

Zo ziet een eenvoudige graph convolution layer implementatie eruit:

import torch
from torch.nn import Sequential, Linear, ReLU
from torch_geometric.nn import MessagePassing

class EdgeConv(MessagePassing):
    def __init__(self, in_channels, out_channels):
        super().__init__(aggr="max")
        self.mlp = Sequential(
            Linear(2 * in_channels, out_channels),
            ReLU(),
            Linear(out_channels, out_channels),
        )

    def forward(self, x, edge_index):
        # x задает фичи узлов, edge_index отвечает за связи между ними
        return self.propagate(edge_index, x=x)

    def message(self, x_j, x_i):
        # x_i — текущий узел, x_j — его сосед
        edge_features = torch.cat([x_i, x_j - x_i], dim=-1)
        return self.mlp(edge_features)

Je hoeft geen lussen te schrijven over alle graaf-edges. PyG paralleliseert deze operatie zelf en voert het snel uit dankzij gecompileerde CUDA kernels.

Basisvoorbeeld: Node Classificatie

Laten we de standaard dataset Cora nemen, die bestaat uit wetenschappelijke artikelen en citaties ertussen. De taak is om de categorie van een artikel te voorspellen op basis van de tekst en verbindingen met andere publicaties.

Een regulier convolutioneel model GCN met data loading wordt letterlijk in een paar dozijn regels in elkaar gezet:

import torch
from torch_geometric.datasets import Planetoid
from torch_geometric.nn import GCNConv

# Загружаем датасет
dataset = Planetoid(root='.', name='Cora')

class GCN(torch.nn.Module):
    def __init__(self, in_channels, hidden_channels, out_channels):
        super().__init__()
        self.conv1 = GCNConv(in_channels, hidden_channels)
        self.conv2 = GCNConv(hidden_channels, out_channels)

    def forward(self, x, edge_index):
        x = self.conv1(x, edge_index).relu()
        x = self.conv2(x, edge_index)
        return x

model = GCN(dataset.num_features, 16, dataset.num_classes)

De trainingslus ziet er volledig standaard uit voor PyTorch-enthousiastelingen: we geven de feature matrix x en de edge tensor edge_index door, berekenen CrossEntropy, en roepen loss.backward() aan.

Het Hoofdprobleem van Grafen en Hoe PyG Het Oplost

De belangrijkste uitdaging van graph deep learning is schaalbaarheid. Afbeeldingen in een batch kunnen gemakkelijk worden gesplitst en in delen naar de GPU worden gestuurd. Maar in een graaf zijn alle nodes met elkaar verbonden. Wanneer een graaf groeit tot miljoenen nodes en miljarden edges, past het niet meer in het geheugen van zelfs de krachtigste GPU's.

PyG-ontwikkelaars hebben veel moeite gestoken in het oplossen van dit probleem. De library bevat subgraph sampling mechanismen:

  • NeighborLoader selecteert alleen willekeurige buren voor elke node in de batch, waardoor geheugen niet explodeert
  • ClusterGCN groepeert een grote graaf in onafhankelijke stukken en traint het netwerk erop
  • GraphSAINT samplet willekeurige subgraven terwijl de topologie behouden blijft

Dankzij dit kan PyG enorme graven verwerken op reguliere GPU's zonder out-of-memory errors.

Wat Er Nog Meer Inzit

De repository bevat een indrukwekkende collectie van al geïmplementeerde architecturen uit onderzoekspapers. Je vindt klassieke GCN, GAT en GraphSAGE, maar ook gespecialiseerde modellen zoals SchNet en DimeNet voor moleculaire analyse of PointNet voor het werken met 3D point clouds.

Naast algoritmen bevat de library loaders voor honderden standaard datasets: van sociale netwerken tot bio-informatica. Dit bespaart veel tijd op het schrijven van parsers.

Installatie en Valstrikken

Sinds versie 2.3 installeert het basispakket extreem gemakkelijk:

pip install torch_geometric

Dit is genoeg om te beginnen. Als je echter maximale snelheid nodig hebt op enorme graven, moet je extra libraries installeren met C++/CUDA extensies: pyg-lib, torch-scatter en torch-sparse.

Dit is waar het soms lastig wordt. Deze binaries zijn strikt gekoppeld aan specifieke versies van PyTorch en CUDA. Als je een incompatibele build installeert, zal Python je overspoelen met C++ library import errors. Controleer altijd de compatibiliteitsmatrix op de projectwebsite en installeer extensies met expliciete wheel URLs.

Voor Wie Dit Framework Is

De library is het proberen waard als je data niet goed past in tabellen of rasters:

  • Recommendersystemen: de "user-item" relatiegraaf werkt beter dan reguliere matrixfactorisatie.
  • Anti-fraud en fintech: het vinden van ketens van verdachte transacties en verborgen groepen van oplichters.
  • Bio-informatica en chemie: het voorspellen van moleculaire eigenschappen, drug discovery en eiwitanalyse.
  • 3D dataverwerking: LiDAR scans en point clouds worden perfect gerepresenteerd als graven.

PyG is de de facto standaard in de industrie geworden voor het werken met graven. Als je van plan bent om dergelijke problemen op te lossen, is het bouwen van aangepaste architecturen helemaal niet de moeite waard — PyG zal je veel tijd en resources besparen.

Gerelateerde projecten