リンク予測とグラフ分類
エッジ予測タスク、負例サンプリング、グラフレベルのプーリング、グラフ分類用GINConvについて学習します。
「リンク予測とグラフ分類」はCoddyKit上の無料Learn AI with Pythonレッスンです。 これはレッスン4/4です。 下記で完全なレッスンを無料で読むことができます。その後、ブラウザ内の組み込みコードエディタと24時間対応のAIチューターでハンズオン演習できます。 これはLearn AI with Python学習パスの一部であり、ウェブとCoddyKitアプリ全体で進捗が同期されます。 Learn AI with Pythonコースには全4レッスンが含まれています。
2つの新しいグラフタスク
ノードの分類以外にも、GNNは次のタスクに対応します。
- リンク予測:2つのノード間に辺が存在するかを予測します(友人の推薦、薬物相互作用など)
- グラフ分類:グラフ全体にラベルを付けます(この分子は毒性を持つかなど)
リンク予測の設定
リンク予測では、まずGNNでノード埋め込みを計算し、次に候補となるノードのペアをスコアリングします。スコアが高いほど、そのペアを結ぶ辺が存在する可能性が高いとモデルが判断したことを意味します。
辺のスコアリング
辺の一般的なスコアとして、2つのノード埋め込みの内積を使います。score = dot(h_u, h_v) のように計算します。埋め込みが似ているほど内積は大きくなり、リンクが存在する可能性が高いと予測されます。
h = gnn(data.x, data.edge_index) # node embeddings
score = (h[u] * h[v]).sum(dim=-1) # dot product per pair負例サンプリング
グラフには、存在する辺(正例)だけが記録されています。分類器を学習するには、辺が存在しないケースも必要です。負例サンプリングでは、接続されていないノードのペアをランダムに選び、負例として扱うことで、訓練データのバランスを整えます。
from torch_geometric.utils import negative_sampling
neg_edge_index = negative_sampling(
edge_index=data.edge_index,
num_nodes=data.num_nodes,
num_neg_samples=data.edge_index.size(1),
)BCEWithLogitsLoss
リンク予測は、辺があるかないかを判定する二値問題です。正例と負例のペアをスコアリングし、それぞれに1と0のラベルを付け、BCEWithLogitsLossで学習します。これはシグモイド関数と二値交差エントロピーを組み合わせ、数値的に安定して計算します。
import torch
pos = (h[pos_u] * h[pos_v]).sum(-1)
neg = (h[neg_u] * h[neg_v]).sum(-1)
scores = torch.cat([pos, neg])
labels = torch.cat([torch.ones_like(pos), torch.zeros_like(neg)])
loss = torch.nn.functional.binary_cross_entropy_with_logits(scores, labels)グラフ分類への切り替え
グラフ分類では、ノードごとではなくグラフごとに1つのベクトルが必要です。GNN層でノード埋め込みを生成した後、それらをプーリングしてグラフレベルの1つの表現にまとめます。
global_mean_pool
global_mean_poolは、グラフ内のすべてのノード埋め込みを平均し、グラフの大きさにかかわらず固定サイズのベクトルを1つ生成します。グラフをまとめて処理する際は、batchインデックスによって各ノードがどのグラフに属するかを指定します。
from torch_geometric.nn import global_mean_pool
h = gnn(x, edge_index) # [num_nodes, dim]
hg = global_mean_pool(h, batch) # [num_graphs, dim]
logits = classifier(hg)プーリングが重要な理由
プーリングによって、モデルはノードの順序やグラフの大きさに依存しなくなります。そのため、同型な2つのグラフからは同じプーリング済みベクトルが得られます。平均プーリングは単純な方法で、合計プーリングや最大プーリングは感度の異なる代替手法です。
GINConv
GINConv(Graph Isomorphism Network)は、より表現力の高い畳み込みです。MLPと加算による集約を使用し、グラフレベルのタスクにおけるメッセージパッシングの識別能力を最大限に高めるよう設計されています。
from torch_geometric.nn import GINConv
import torch
mlp = torch.nn.Sequential(
torch.nn.Linear(in_dim, hid),
torch.nn.ReLU(),
torch.nn.Linear(hid, hid),
)
conv = GINConv(mlp)Weisfeiler-Lemanとの関係
GINは、非同型グラフを区別する古典的なアルゴリズムであるWeisfeiler-Leman(WL)テストと同等の能力を持つように設計されています。より単純な多くのGNNでは区別できないグラフも、GINならWLテストの限界内で区別できます。そのため、グラフ分類に適しています。
適切なツールの選択
タスクに応じてアーキテクチャを選択します。
- リンク予測:GNN埋め込み + 内積によるスコアリング + ネガティブサンプリング + BCE損失
- グラフ分類:GINConvのような表現力の高い畳み込み + グローバルプーリング + 分類器
確認テスト
理解度を確認しましょう。
まとめ
リンク予測とグラフ分類について学びました。
- エッジスコア =
dot(h_u, h_v)。ネガティブサンプリングとBCEWithLogitsLossを使って学習します - global_mean_poolはノード埋め込みをグラフレベルのベクトルに変換します
- GINConvは非常に表現力が高く、Weisfeiler-Lemanテストに匹敵します
よくある質問
「リンク予測とグラフ分類」レッスンは無料ですか?
はい。「リンク予測とグラフ分類」の完全なテキストはこのウェブで無料で読めます。インタラクティブに演習し(組み込みコードエディタと24時間対応のAIチューター)、Learn AI with Pythonコースの残りをアンロックするには、CoddyKit PROにアップグレードしてください。 Learn AI with Pythonコースには全4レッスンが含まれています。
「リンク予測とグラフ分類」で何を学びますか?
エッジ予測タスク、負例サンプリング、グラフレベルのプーリング、グラフ分類用GINConvについて学習します。 ブラウザで直接実行するハンズオンコードでLearn AI with Pythonを演習し、24時間対応のAIチューターがレッスンを進める中での質問に答えます。
Learn AI with Pythonを始めるのに経験は必要ですか?
事前経験は必要ありません。CoddyKitのLearn AI with Pythonは初級者から上級者向けに構成されているため、ここから始めるか最初から始めて、自分のペースで進むことができます。 これはレッスン4/4です。
「リンク予測とグラフ分類」レッスンにはどのくらい時間がかかりますか?
ほとんどのCoddyKitレッスンは約5~10分かかります。各レッスンはコンパクトでインタラクティブなので、着実に進歩し、ウェブとアプリ全体で正確に前回の場所から再開できます。
このLearn AI with Pythonレッスンでコードを書いて実行できますか?
はい。すべてのLearn AI with Pythonレッスンに組み込みコードエディタが含まれているため、ブラウザでリアルコードを書いて実行し、即座のAIフィードバックを取得できます。ローカル設定は不要です。
このコースのすべてのレッスン
- 機械学習のためのグラフ理論
- グラフ畳み込みネットワーク(GCN)
- GNNによるノード分類
- リンク予測とグラフ分類