Sieci neuronowe na grafach w PyTorch bez zbędnego bólu i niestandardowych obejść
Większość programistów jest przyzwyczajona do pracy z przejrzystymi, regularnymi strukturami danych. Obraz to siatka pikseli, a tekst łatwo mieści się w sekwencji tokenów. Zwykłe sieci splotowe lub transformery radzą sobie z tymi danymi doskonale.
Problemy pojawiają się tam, gdzie dane mają fundamentalnie nieliniowy charakter. Weźmy transakcje bankowe, połączenia między użytkownikami w sieciach społecznościowych czy wzory chemiczne złożonych cząsteczek. Grafy są wszędzie. Jeśli spróbujesz przekazać je do zwykłego PyTorch, będziesz musiał ręcznie zmagać się z rzadkimi macierzami, pisać złożone pętle agregacji węzłów i śledzić indeksy.
Właśnie do tego został stworzony PyTorch Geometric (w skrócie PyG). To rozszerzenie do PyTorch, które obsługuje za Ciebie całą niskopoziomową matematykę i udostępnia intuicyjny API do pracy z sieciami neuronowymi na grafach.
Jak zbudowany jest PyG
Jeśli wiesz już, jak pisać modele w PyTorch, w ciągu kilku godzin poczujesz się komfortowo z PyG. Framework podąża za tymi samymi zasadami: te same moduły torch.nn, znajome pętle trenowania i bezpośrednia praca z tensorami.
Cała magia sieci grafowych opiera się na koncepcji Message Passing. Każdy wierzchołek w grafie zbiera informacje od swoich sąsiadów, łączy je i aktualizuje swój własny stan.
PyG udostępnia klasę bazową MessagePassing, która abstrahuje tę logikę. Wewnątrz musisz zdefiniować tylko trzy rzeczy:
- Jak wiadomość jest tworzona z węzła sąsiada
- Jak te wiadomości są agregowane (suma, średnia, maksimum)
- Jak sam węzeł jest aktualizowany
Oto jak wygląda implementacja prostej warstwy splotowej na grafie:
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)
Nie musisz pisać pętli po wszystkich krawędziach grafu. PyG sam parallelizuje tę operację i wykonuje ją szybko dzięki skompilowanym jądrom CUDA.
Podstawowy przykład: klasyfikacja węzłów
Weźmy standardowy zbiór danych Cora, który składa się z artykułów naukowych i cytowań między nimi. Zadanie polega na przewidzeniu kategorii artykułu na podstawie jego tekstu i połączeń z innymi publikacjami.
Zwykły model splotowy GCN z ładowaniem danych można złożyć dosłownie w dwóch tuzinach linii:
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)
Pętla trenowania wygląda całkowicie standardowo dla entuzjastów PyTorch: przekazujemy macierz cech x i tensor krawędzi edge_index, obliczamy CrossEntropy i wywołujemy loss.backward().
Główny problem grafów i jak rozwiązuje go PyG
Głównym wyzwaniem głębokiego uczenia na grafach jest skalowalność. Obrazy w partii można łatwo podzielić i wysłać na GPU częściami. Ale w grafie wszystkie węzły są ze sobą połączone. Gdy graf rośnie do milionów węzłów i miliardów krawędzi, nie mieści się już w pamięci nawet najpotężniejszych GPU.
Deweloperzy PyG włożyli wiele wysiłku w rozwiązanie tego problemu. Biblioteka zawiera mechanizmy próbkowania podgrafów:
NeighborLoaderwybiera tylko losowych sąsiadów dla każdego węzła w partii, zapobiegając eksplozji pamięciClusterGCNklastruje duży graf na niezależne części i trenuje sieć na nichGraphSAINTpróbkuje losowe podgrafy, zachowując ich topologię
Dzięki temu PyG może przetwarzać ogromne grafy na zwykłych GPU bez wyczerpania pamięci.
Co jeszcze znajdziesz out of the box
Repozytorium zawiera imponującą kolekcję już zaimplementowanych architektur z artykułów badawczych. Znajdziesz klasyczne GCN, GAT i GraphSAGE, a także wyspecjalizowane modele jak SchNet i DimeNet do analizy molekularnej czy PointNet do pracy z chmurami punktów 3D.
Oprócz algorytmów biblioteka zawiera ładowarki setek standardowych zbiorów danych: od sieci społecznościowych po bioinformatykę. To oszczędza mnóstwo czasu na pisaniu parserów.
Instalacja i podwodne kamienie
Od wersji 2.3 podstawowy pakiet instaluje się niezwykle łatwo:
pip install torch_geometric
To wystarczy, żeby zacząć. Jeśli jednak potrzebujesz maksymalnej prędkości na ogromnych grafach, musisz zainstalować dodatkowe biblioteki z rozszerzeniami C++/CUDA: pyg-lib, torch-scatter i torch-sparse.
Tutaj czasem robi się trudno. Te pliki binarne są ściśle powiązane z konkretnymi wersjami PyTorch i CUDA. Jeśli zainstalujesz niekompatybilną wersję, Python zaleje Cię błędami importu bibliotek C++. Zawsze sprawdzaj macierz kompatybilności na stronie projektu i instaluj rozszerzenia z jawnymi URL-ami wheel.
Dla kogo jest ten framework
Biblioteka warto wypróbować, jeśli Twoje dane nie pasują dobrze do tabel ani siatek:
- Systemy rekomendacyjne: graf relacji „użytkownik-przedmiot" działa lepiej niż zwykła faktoryzacja macierzy.
- Antyfraud i fintech: znajdowanie łańcuchów podejrzanych transakcji i ukrytych grup oszustów.
- Bioinformatyka i chemia: przewidywanie właściwości molekularnych, odkrywanie leków i analiza białek.
- Przetwarzanie danych 3D: skany LiDAR i chmury punktów są doskonale reprezentowane jako grafy.
PyG stał się de facto standardem w branży do pracy z grafami. Jeśli planujesz rozwiązywać takie problemy, budowanie niestandardowych architektur od zera zdecydowanie nie jest tego warte — PyG zaoszczędzi Ci mnóstwo czasu i zasobów.
Powiązane projekty