>_ DevTrendsja

言語

ホーム

言語

セクション

フロントエンド バックエンド モバイル DevOps AI / ML ゲーム開発 ブロックチェーン 組み込み セキュリティ
Python

PyTorchでのグラフニューラルネットワーク:余計な苦労もカスタムワークアラウンドもなく

ほとんどの開発者は、明確で規則的なデータ構造に慣れています。画像はピクセルグリッドであり、テキストはトークンのシーケンスに簡単に収まります。通常の畳み込みネットワークやTransformerはこれらを完璧に処理します。

データが本質的に非線形である場合に問題が発生します。銀行取引、ソーシャルネットワーク内のユーザー間の接続、または複雑な分子の化学式などを考えてみてください。グラフはどこにでも存在します。これを通常のPyTorchに投入しようとすると、スパース行列を手前で苦労して扱わなければならず、ノード集約のための複雑なループを書き、インデックスを追跡し続ける必要があります。

これがPyTorch Geometric(略称PyG)が作られた理由です。PyTorchの拡張機能であり、すべての低レベルな数学的処理を行い、グラフニューラルネットワークを操作するための直感的なAPIを提供します。

PyGの構造

すでにPyTorchでモデルの書き方を知っているなら、数時間でPyGにも慣れるでしょう。このフレームワークは同じ原則に従っています:同じモジュール torch.nn、おなじみのトレーニングループ、そして直接的なテンソル操作です。

グラフネットワークのすべての魔法は、Message Passingという概念を中心に構築されています。グラフ内の各頂点は、近隣からの情報を収集し、組み合わせ、それ自体の状態を更新します。

PyGは、このロジックを抽象化するベースクラス MessagePassingを提供します。内部では、3つのことだけを定義する必要があります:

  • 隣接ノードからメッセージがどのように形成されるか
  • これらのメッセージがどのように集約されるか(合計、平均、最大値)
  • ノード自体がどのように更新されるか

シンプルなグラフ畳み込み層の実装は次のようになります:

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)

すべてのグラフエッジをループで処理する必要はありません。PyGは自動的にこの操作を並列化し、コンパイルされたCUDAカーネル 덕분에高速に実行します。

基本例:ノード分類

科学論文とその間の引用関係で構成される標準データセット Coraを取り上げましょう。タスクは、テキストと他の出版物への接続に基づいて論文のカテゴリを予測することです。

通常の畳み込みモデル GCNとデータローダーは、文字通り数十行で組み立てられます:

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)

トレーニングループは、PyTorchに慣れている人にとっては完全に標準的に見えます:特徴行列 xとエッジテンソル edge_indexを渡 し、CrossEntropyを計算して loss.backward()を呼び出します。

グラフの主な問題とPyGの解決策

グラフ深層学習の主な課題はスケーリングです。バッチ内の画像は簡単に分割してGPUに部分的に送信できます。しかし、グラフではすべてのノードが相互に接続されています。グラフが数百万のノードと数十億のエッジに成長すると、最も強力なGPUのメモリにも収まらなくなります。

PyGの開発者はこの問題を解決するために多大な 노력을費やしました。このライブラリにはサブグラフサンプリングメカニズムが含まれています:

  • NeighborLoaderはバッチ内の各ノードに対してランダムな近傍のみを選択することで、メモリが爆発するのを防ぎます
  • ClusterGCNは大きなグラフを独立したピースにクラスタリングし、それらの上でネットワークをトレーニングします
  • GraphSAINTはトポロジーを保ちながらランダムなサブグラフをサンプリングします

これにより、PyGはメモリ不足になることなく、通常のGPUで巨大なグラフを処理できます。

標準で他に何が含まれているか

リポジトリには、研究論文からすでに実装されているアーキテクチャの印象的なコレクションが含まれています。古典的なGCN、GAT、GraphSAGEだけでなく、分子分析用のSchNetやDimeNet、3Dポイントクラウド操作用のPointNetなどの専門モデルも見つかります。

アルゴリズム以外にも、このライブラリにはソーシャルネットワークからバイオinformaticsまでの数百の標準データセット用のローダーが含まれています。これにより、パーサーを書く時間が大幅に節約されます。

インストールと落とし穴

バージョン2.3以降、ベースパッケージは非常に簡単にインストールできます:

pip install torch_geometric

これは始めるには十分です。ただし、巨大なグラフで最大速度が必要な場合は、C++/CUDA拡張子付きの追加ライブラリをインストールする必要があります: pyg-libtorch-scatter、および torch-sparse

ここで時々面倒ことになることがあります。これらのバイナリは特定のバージョンのPyTorchとCUDAに密接に結合しています。互換性のないビルドをインストールすると、PythonはC++ライブラリのインポートエラーで溢れかえります。プロジェクトのウェブサイトで互換性マトリックスを確認し、明示的なwheel URLで拡張機能をインストールしてください。

このフレームワークは誰に向いているか

データがテーブルやグリッドにうまく収まらない場合、このライブラリを試す価値があります:

  • レコメンデーションシステム:「ユーザー-アイテム」関係グラフは通常の行列因子分解より優れています。
  • 不正検知とフィンテック:疑わしい取引のチェーンや隠れた不正者のグループを見つけます。
  • バイオinformaticsと化学:分子特性の予測、創薬、タンパク質分析。
  • 3Dデータ処理:LiDARスキャンとポイントクラウドはグラフとして完全に表現できます。

PyGはグラフを操作するための事実上の業界標準になりました。こんな問題を解決する予定がある場合、ゼロからカスタムアーキテクチャを構築する価値は確かにありません—PyGは多くの時間とリソースを節約してくれるでしょう。

関連プロジェクト