>_ DevTrendsfr

Langue

Accueil

Langages

Sections

Frontend Backend Mobile DevOps AI / ML GameDev Blockchain Embarqué Sécurité
Python

Les réseaux de neurones sur graphes en PyTorch sans douleur ni bidouillage

La plupart des développeurs sont habitués à travailler avec des structures de données claires et régulières. Une image est une grille de pixels, et le texte s'intègre facilement dans une séquence de jetons. Les réseaux convolutifs classiques ou les transformers gèrent parfaitement ces données.

Les problèmes surviennent lorsque les données sont fondamentalement non linéaires. Prenez les transactions bancaires, les connexions entre utilisateurs dans les réseaux sociaux ou les formules chimiques de molécules complexes. Les graphes sont partout. Si vous essayez de les transmettre à PyTorch classique, vous devrez lutter manuellement avec les matrices creuses, écrire des boucles complexes pour l'agrégation des nœuds et suivre les indices.

C'est exactement pour cela que PyTorch Geometric (abréégé PyG) a été créé. C'est une extension pour PyTorch qui gère toutes les mathématiques de bas niveau pour vous et fournit une API intuitive pour travailler avec les réseaux de neurones sur graphes.

Comment PyG est structuré

Si vous savez déjà comment écrire des modèles en PyTorch, vous serez à l'aise avec PyG en quelques heures. Le framework suit les mêmes principes : les mêmes modules torch.nn, des boucles d'entraînement familières et un travail direct avec les tenseurs.

Toute la magie des réseaux de graphes repose sur le concept de Message Passing. Chaque sommet d'un graphe collecte les informations de ses voisins, les combine et met à jour son propre état.

PyG fournit une classe de base MessagePassing qui abstrait cette logique. À l'intérieur, vous n'avez besoin de définir que trois choses :

  • Comment un message est formé à partir d'un nœud voisin
  • Comment ces messages sont agrégés (somme, moyenne, max)
  • Comment le nœud lui-même est mis à jour

Voici à quoi ressemble l'implémentation d'une simple couche de convolution sur graphe :

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)

Vous n'avez pas besoin d'écrire des boucles sur tous les bords du graphe. PyG parallélise lui-même cette opération et l'exécute rapidement grâce aux noyaux CUDA compilés.

Exemple de base : Classification de nœuds

Prenons le jeu de données standard Cora, qui se compose d'articles scientifiques et de citations entre eux. La tâche consiste à prédire la catégorie d'un article en fonction de son texte et de ses connexions avec d'autres publications.

Un modèle convolutif classique GCN avec chargement de données est assemblé en à peine deux dizaines de lignes :

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)

La boucle d'entraînement semble tout à fait standard pour les passionnés de PyTorch : nous transmettons la matrice de caractéristiques x et le tenseur d'arêtes edge_index, calculons la CrossEntropy et appelons loss.backward().

Le problème principal des graphes et comment PyG le résout

Le principal défi de l'apprentissage profond sur graphes est la mise à l'échelle. Les images d'un lot peuvent être facilement divisées et envoyées au GPU par parties. Mais dans un graphe, tous les nœuds sont interconnectés. Lorsqu'un graphe atteint des millions de nœuds et des milliards d'arêtes, il ne tient plus dans la mémoire même des GPU les plus puissants.

Les développeurs de PyG ont consacré beaucoup d'efforts à résoudre ce problème. La bibliothèque contient des mécanismes d'échantillonnage de sous-graphes :

  • NeighborLoader sélectionne uniquement des voisins aléatoires pour chaque nœud du lot, empêchant la mémoire d'exploser
  • ClusterGCN regroupe un grand graphe en morceaux indépendants et entraîne le réseau dessus
  • GraphSAINT échantillonne des sous-graphes aléatoires tout en préservant leur topologie

Grâce à cela, PyG peut venir à bout de graphes géants sur des GPU ordinaires sans manquer de mémoire.

Ce qu'il y a encore dans la boîte

Le dépôt contient une impressionnante collection d'architectures déjà implémentées à partir d'articles de recherche. Vous y trouverez les classiques GCN, GAT et GraphSAGE, ainsi que des modèles spécialisés comme SchNet et DimeNet pour l'analyse moléculaire ou PointNet pour travailler avec des nuages de points 3D.

Au-delà des algorithmes, la bibliothèque comprend des chargeurs pour des centaines de jeux de données standard : des réseaux sociaux à la bioinformatique. Cela fait gagner beaucoup de temps sur l'écriture d'analyseurs.

Installation et pièges

Depuis la version 2.3, le package de base s'installe extrêmement facilement :

pip install torch_geometric

C'est suffisant pour commencer. Cependant, si vous avez besoin d'une vitesse maximale sur de grands graphes, vous devrez installer des bibliothèques supplémentaires avec des extensions C++/CUDA : pyg-lib, torch-scatter et torch-sparse.

C'est là que les choses se compliquent parfois. Ces binaires sont étroitement liés à des versions spécifiques de PyTorch et CUDA. Si vous installez une version incompatible, Python vous submergera d'erreurs d'importation de bibliothèques C++. Vérifiez toujours la matrice de compatibilité sur le site du projet et installez les extensions avec des URLs wheel explicites.

Pour qui est ce framework

La bibliothèque mérite d'être essayée si vos données ne s'adaptent pas bien aux tableaux ou aux grilles :

  • Systèmes de recommandation : le graphe de relation « utilisateur-article » fonctionne mieux que la factorisation matricielle classique.
  • Anti-fraude et fintech : trouver des chaînes de transactions suspectes et des groupes cachés de fraudeurs.
  • Bioinformatique et chimie : prédiction des propriétés moléculaires, découverte de médicaments et analyse des protéines.
  • Traitement de données 3D : les scans LiDAR et les nuages de points sont parfaitement représentés comme des graphes.

PyG est devenu la norme de facto dans l'industrie pour travailler avec les graphes. Si vous prévoyez de résoudre de tels problèmes, construire des architectures personnalisées à partir de zéro ne vaut certainement pas la peine — PyG vous fera gagner beaucoup de temps et de ressources.

Projets similaires