>_ DevTrendspt

Idioma

Início

Linguagens

Seções

Frontend Backend Mobile DevOps AI / ML GameDev Blockchain Embarcados Segurança
Python

Redes Neurais de Grafos no PyTorch Sem Dor Extra e Soluções Alternativas Manuais

A maioria dos desenvolvedores está acostumada a trabalhar com estruturas de dados claras e regulares. Uma imagem é uma grade de pixels, e o texto se encaixa facilmente em uma sequência de tokens. Redes convolucionais regulares ou transformers lidam com isso perfeitamente.

Problemas surgem onde os dados são fundamentalmente não-lineares. Considere transações bancárias, conexões entre usuários em redes sociais ou fórmulas químicas de moléculas complexas. Grafos estão em ação em todo lugar. Se você tentar alimentá-los ao PyTorch regular, terá que lutar manualmente com matrizes esparsas, escrever loops complexos para agregação de nós e acompanhar os índices.

É exatamente para isso que o PyTorch Geometric (abreviado PyG) foi criado. É uma extensão do PyTorch que lida com toda a matemática de baixo nível para você e fornece uma API intuitiva para trabalhar com redes neurais de grafos.

Como o PyG é Estruturado

Se você já sabe como escrever modelos no PyTorch, ficará confortável com o PyG em poucas horas. O framework segue os mesmos princípios: os mesmos módulos torch.nn, loops de treinamento familiares e trabalho direto com tensores.

Toda a mágica das redes de grafos é construída em torno do conceito de Message Passing. Cada vértice em um grafo coleta informações de seus vizinhos, combina e atualiza seu próprio estado.

O PyG fornece uma classe base MessagePassing que abstrai essa lógica. Dentro dela, você só precisa definir três coisas:

  • Como uma mensagem é formada a partir de um nó vizinho
  • Como essas mensagens são agregadas (soma, média, máx)
  • Como o próprio nó é atualizado

Veja como fica a implementação de uma camada simples de convolução de grafos:

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)

Você não precisa escrever loops sobre todas as arestas do grafo. O PyG paraleliza essa operação por conta própria e a executa rapidamente graças a kernels CUDA compilados.

Exemplo Básico: Classificação de Nós

Vamos usar o dataset padrão Cora, que consiste em artigos científicos e citações entre eles. A tarefa é prever a categoria de um artigo com base em seu texto e conexões com outras publicações.

Um modelo convolucional regular GCN com carregamento de dados é montado em literalmente duas dezenas de linhas:

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)

O loop de treinamento parece completamente padrão para entusiastas do PyTorch: passamos a matriz de features x e o tensor de arestas edge_index, calculamos a CrossEntropy e chamamos loss.backward().

O Principal Problema dos Grafos e Como o PyG Resolve

O principal desafio do deep learning em grafos é a escalabilidade. Imagens em um batch podem ser facilmente divididas e enviadas para a GPU em partes. Mas em um grafo, todos os nós estão interconectados. Quando um grafo cresce para milhões de nós e bilhões de arestas, ele não cabe mais na memória mesmo das GPUs mais poderosas.

Os desenvolvedores do PyG gastaram muito esforço resolvendo esse problema. A biblioteca contém mecanismos de amostragem de subgrafos:

  • NeighborLoader seleciona apenas vizinhos aleatórios para cada nó no batch, impedindo que a memória exploda
  • ClusterGCN agrupa um grafo grande em peças independentes e treina a rede nelas
  • GraphSAINT amostra subgrafos aleatórios preservando sua topologia

Graças a isso, o PyG consegue processar grafos gigantes em GPUs normais sem ficar sem memória.

O Que Mais Vem Pronto para Uso

O repositório contém uma coleção impressionante de arquiteturas já implementadas de artigos de pesquisa. Você encontrará GCN clássico, GAT e GraphSAGE, além de modelos especializados como SchNet e DimeNet para análise molecular ou PointNet para trabalhar com nuvens de pontos 3D.

Além dos algoritmos, a biblioteca inclui carregadores para centenas de datasets padrão: de redes sociais a bioinformática. Isso economiza muito tempo na escrita de parsers.

Instalação e Armadilhas

Desde a versão 2.3, o pacote base instala extremamente facilmente:

pip install torch_geometric

Isso é suficiente para começar. No entanto, se você precisa de velocidade máxima em grafos enormes, precisará instalar bibliotecas adicionais com extensões C++/CUDA: pyg-lib, torch-scatter e torch-sparse.

É aqui que as coisas às vezes ficam complicadas. Esses binários estão fortemente acoplados a versões específicas do PyTorch e CUDA. Se você instalar uma versão incompatível, o Python vai inundá-lo com erros de importação de bibliotecas C++. Sempre verifique a matriz de compatibilidade no site do projeto e instale extensões com URLs wheel explícitas.

Para Quem Este Framework É

A biblioteca vale a pena tentar se seus dados não se encaixam bem em tabelas ou grades:

  • Sistemas de recomendação: o grafo de relacionamento "usuário-item" funciona melhor do que a fatoração de matriz regular.
  • Anti-fraude e fintech: encontrar cadeias de transações suspeitas e grupos ocultos de fraudadores.
  • Bioinformática e química: prever propriedades moleculares, descoberta de medicamentos e análise de proteínas.
  • Processamento de dados 3D: varreduras LiDAR e nuvens de pontos são perfeitamente representadas como grafos.

O PyG se tornou o padrão de fato na indústria para trabalhar com grafos. Se você está planejando resolver esses problemas, construir arquiteturas personalizadas do zero definitivamente não vale a pena — o PyG economizará muito tempo e recursos.

Projetos relacionados