【GNN】Message Passing Neural Network(MPNN)を解説する

ソーシャルネットワークで「知人の知人はどんな人物か」を推測する場面を想像してください。あなたが直接知らない人でも、共通の友人が多ければ、その人の傾向をある程度推測できます。グラフ上の機械学習も同じ発想です。あるノードの性質を予測するために、周囲のノードから情報を「メッセージ」として受け取り、それを集めて自分の表現を更新する。この繰り返しがメッセージパッシングの本質です。

MPNN(Message Passing Neural Network) は、2017年にGilmerらが提案したGNNの統一的フレームワークです。GCN・GAT・GraphSAGEなど一見バラバラに見える多様なGNNアルゴリズムを、「メッセージ関数 $M$・集約 $\bigoplus$・更新関数 $U$」という3つの部品で統一的に説明できます。

MPNNを理解すると、次の2つの恩恵が得られます。まず、PyTorch Geometric(PyG)の MessagePassing 基底クラスを正しく使いこなせるようになります。PyGのほぼすべてのGNNレイヤーはMPNNの形式で実装されており、message()update() をオーバーライドするだけで独自のGNNを組めます。次に、既存手法の設計意図を読み解けるようになります。GINがSUM集約にこだわる理由、GATが注意重みを使う理由——これらがMPNNの枠組みの中で「どの部品を変えたか」として明確になります。

本記事の内容

  • メッセージパッシングの直感と数学的定義
  • MPNNの一般フレームワーク(Gilmer 2017):メッセージ・集約・更新・Readoutの4段
  • GCN・GAT・GraphSAGE・GINとMPNNの対応
  • 集約関数(SUM / MEAN / MAX)の識別能力の違い
  • kホップ受容野と過平滑化の問題
  • NumPyによる数値実装と検証

メッセージパッシングとは何か

グラフは「ノード(頂点)」と「エッジ(辺)」の2つの要素からなります。化学分子であれば原子がノード、結合がエッジです。ソーシャルネットワークなら人物がノード、友人関係がエッジです。このようなグラフ構造データに対して予測や分類を行うのがGNNの目的です。

しかし畳み込みニューラルネットワーク(CNN)は画像のような「格子構造」を前提としており、そのままではグラフには適用できません。グラフのノードは不規則な数の隣接ノードを持ち、どの方向が「右隣」かという概念がありません。そこで必要になるのが「近傍から情報を集めて自分を更新する」という操作の抽象化です。

MPNNメッセージパッシング1ステップの概念図

上の図はメッセージパッシングの1ステップを示しています。ターゲットノード $i$(赤)に対して、直接エッジで繋がっている近傍ノード $j_1, j_2$(青)が「メッセージ」を送ります。非近傍のノード(灰色)は1ステップでは関与しません。受け取ったメッセージを集約し、ノード $i$ の表現を更新します。この操作を層ごとに繰り返すことで、2層目では「近傍の近傍」まで情報が届きます。


MPNNの一般フレームワーク

メッセージ関数 $M_t$

Gilmerら(2017)の論文 Neural Message Passing for Quantum Chemistry では、ノード $i$ と近傍ノード $j$ の間で計算されるメッセージを次のように定義しています。

$$ \begin{equation} m_{ij}^{(t)} = M_t\!\left(h_i^{(t)},\, h_j^{(t)},\, e_{ij}\right) \end{equation} $$

ここで各変数の意味は次のとおりです。

  • $h_i^{(t)}$:$t$ 層目($t$ ステップ目)におけるノード $i$ の特徴ベクトル(隠れ状態)
  • $h_j^{(t)}$:同じく $t$ 層目の近傍ノード $j$ の特徴ベクトル
  • $e_{ij}$:ノード $i$ と $j$ を結ぶエッジの特徴量(分子なら結合の種類、グラフによってはなくてもよい)
  • $M_t$:メッセージ関数。MLPなどの非線形関数が使われる

「ノード $j$ がノード $i$ に送るメッセージ」が $m_{ij}^{(t)}$ です。このメッセージはノード $j$ の情報だけでなく、受け取る側のノード $i$ の情報と、エッジの属性も組み込んで計算されます。エッジ特徴量が不要な場合(多くのグラフタスクでは省略される)は $M_t(h_j^{(t)})$ のように簡略化されます。

集約(Aggregation)

ノード $i$ の近傍 $\mathcal{N}_i$ 全体からのメッセージを1つにまとめる操作が集約です。

$$ \begin{equation} \tilde{m}_i^{(t)} = \bigoplus_{j \in \mathcal{N}_i} m_{ij}^{(t)} \end{equation} $$

記号 $\bigoplus$ は集約演算子の総称です。具体的には総和(SUM)、平均(MEAN)、最大値(MAX)などが使われます。集約関数は置換不変(permutation invariant)であることが必要です。つまり、近傍ノードの順番が変わっても結果が変わらなければなりません。グラフのノードには「1番目の近傍」「2番目の近傍」という自然な順番がないので、この性質は必須です。

更新関数 $U_t$

集約されたメッセージ $\tilde{m}_i^{(t)}$ と、現在のノード表現 $h_i^{(t)}$ を使って、次の層の特徴量を計算します。

$$ \begin{equation} h_i^{(t+1)} = U_t\!\left(h_i^{(t)},\, \tilde{m}_i^{(t)}\right) \end{equation} $$

更新関数 $U_t$ にも MLPやGRU(Gated Recurrent Unit)などが使われます。GRUを使う場合、$h_i^{(t)}$ を隠れ状態、$\tilde{m}_i^{(t)}$ を入力として見なすことで、時系列のような状態更新が可能になります。

MPNNの3段フレームワーク(メッセージ計算・集約・更新)

上の図がMPNNの全体像です。①メッセージ計算→②集約→③更新という3段の処理を $T$ 回繰り返します。各ステップでノードの受容野が1ホップずつ広がり、最終的には全体的な構造情報を取り込んだ表現が得られます。


kホップ受容野の拡大

MPNNの重要な性質として、層数(ステップ数)を増やすほど、各ノードが参照できる近傍の範囲が広がる点があります。

  • $t=0$:自分自身のみ(0ホップ)
  • $t=1$:直接繋がっている1ホップ近傍まで
  • $t=k$:$k$ ホップ以内の全ノードを間接的に参照

kホップ受容野の拡大(0ホップ・1ホップ・2ホップ)

この図では0ホップ・1ホップ・2ホップでの受容野の広がりを示しています。赤がターゲットノード(ノード0)、オレンジが1ホップ以内、緑が2ホップ以内です。層を1つ重ねるだけで参照できるノードが大幅に増えることがわかります。

CNNの畳み込みカーネルが局所的なピクセルパッチを見るのと同様、MPNNの1層は局所的な近傍を見ます。層を重ねることで受容野が広がる点も対応しています。ただしグラフの場合、各ノードの「近傍の数」(次数)が不均一なため、受容野の広がり方は一様ではありません。

では、受容野が広ければ広いほどよいのでしょうか?実はそうではなく、深すぎると「過平滑化(over-smoothing)」という問題が起きます。この問題については後ほど詳しく説明します。


Readout関数:グラフ全体の表現を作る

ノード分類(各ノードにラベルを付けるタスク)であれば、最終層のノード表現 $h_i^{(T)}$ をそのまま分類器に入力すれば十分です。しかしグラフ分類(グラフ全体に1つのラベルを付けるタスク)では、グラフ全体を1つのベクトルに圧縮する必要があります。これを行うのが Readout関数 $R$ です。

$$ \begin{equation} \hat{y}_G = R\!\left(\left\{h_i^{(T)}\right\}_{i \in G}\right) \end{equation} $$

Readout関数もノードの順番に依存しない(置換不変な)操作である必要があります。よく使われる選択肢は次のとおりです。

  • 平均プーリング:$\hat{y}_G = \frac{1}{|G|} \sum_{i \in G} h_i^{(T)}$(シンプルで安定)
  • 総和プーリング:$\hat{y}_G = \sum_{i \in G} h_i^{(T)}$(グラフの大きさ情報を保持)
  • 最大プーリング:$\hat{y}_G = \max_{i \in G} h_i^{(T)}$(最も顕著な特徴を抽出)
  • 階層的プーリング(DiffPool など):グラフ構造ごと段階的に粗視化する高度な方法

MPNNのReadout関数:ノード表現からグラフ全体の表現へ

図のように、T層のメッセージパッシング後に得られた全ノードの表現をReadout関数に入力し、グラフ全体の表現 $\hat{y}_G$ を出力します。この $\hat{y}_G$ を下流の分類器(線形層など)に渡すことで、分子の毒性予測や化合物の生理活性推定といったグラフ分類タスクを解けます。


GCN・GAT・GraphSAGE・GINとMPNNの対応

MPNNは非常に一般的な枠組みです。有名なGNNアルゴリズムのほとんどは、メッセージ関数 $M$・集約 $\bigoplus$・更新関数 $U$ の具体的な選択として解釈できます。

主要GNNのMPNNフレームワーク対応表(GCN・GAT・GraphSAGE・GIN)

対応表が示すように、各手法の違いは主に「集約をどう重み付けするか」と「更新で何を組み合わせるか」に集約されます。以下で各手法を詳しく見ていきましょう。

GCN(Graph Convolutional Network)

Kipf & Welling(2017)が提案したGCNは、次の式で表されます。

$$ \begin{equation} H^{(l+1)} = \sigma\!\left(\hat{D}^{-1/2} \hat{A} \hat{D}^{-1/2} H^{(l)} W^{(l)}\right) \end{equation} $$

ここで $\hat{A} = A + I$(自己ループを加えた隣接行列)、$\hat{D}$ は $\hat{A}$ の次数行列です。$\hat{D}^{-1/2} \hat{A} \hat{D}^{-1/2}$ による正規化が、次数の異なるノード間での集約を安定させます。

MPNNの言葉で言うと、メッセージ関数は正規化された隣接行列によるノード特徴量の線形変換、集約は平均に相当します。シンプルな実装と安定した性能から最も広く使われるGNNの一つです。ただし、すべての近傍を均等に扱うため、重要な近傍とそうでない近傍を区別できないという弱点があります。

GAT(Graph Attention Network)

Velickovicら(2018)のGATは、近傍ノードの重要度を注意機構(Attention)で学習します。

$$ \begin{equation} \alpha_{ij} = \frac{\exp\!\left(\mathrm{LeakyReLU}\!\left(\bm{a}^\top [W h_i^{(l)} \| W h_j^{(l)}]\right)\right)}{\sum_{k \in \mathcal{N}_i} \exp\!\left(\mathrm{LeakyReLU}\!\left(\bm{a}^\top [W h_i^{(l)} \| W h_k^{(l)}]\right)\right)} \end{equation} $$

注意重み $\alpha_{ij}$ を使った集約は次のようになります。

$$ \begin{equation} h_i^{(l+1)} = \sigma\!\left(\sum_{j \in \mathcal{N}_i} \alpha_{ij} W h_j^{(l)}\right) \end{equation} $$

ここで $\bm{a}$ は学習されるベクトル、$\|$ は結合(concatenation)を表します。注意重みを使うことで、「隣接ノードの中でどれが重要か」をタスクに応じて適応的に学習できます。これはGCNの均等集約に対する大きな改善です。

GraphSAGE(Graph Sample and Aggregation)

Hamiltonら(2017)のGraphSAGEは、スケーラビリティを重視した設計です。

$$ \begin{equation} h_i^{(l+1)} = \sigma\!\left(W \cdot \left[h_i^{(l)} \| \mathrm{Agg}\!\left(\{h_j^{(l)}, j \in \mathcal{N}_i\}\right)\right]\right) \end{equation} $$

集約には平均・最大値・LSTM(近傍をランダム順でLSTMに通す)を選べます。大きな特徴は、自分自身の表現と集約結果をconcatenateしてから線形変換する点です。自分自身の情報が集約に埋もれにくくなっています。また、大規模グラフでは全近傍でなく一部をサンプリングして集約する点がスケーラビリティの源です。

GIN(Graph Isomorphism Network)

Xuら(2019)のGINは、グラフの識別能力という理論的観点から設計されています。

$$ \begin{equation} h_i^{(l+1)} = \mathrm{MLP}^{(l)}\!\left((1 + \epsilon^{(l)}) \cdot h_i^{(l)} + \sum_{j \in \mathcal{N}_i} h_j^{(l)}\right) \end{equation} $$

ここで $\epsilon$ は学習可能なスカラーです。GINの鍵はSUM集約の使用にあります。

式の意味を分解すると、$(1 + \epsilon) \cdot h_i^{(l)}$ は自分自身の特徴量に重みをかけた項、$\sum_{j \in \mathcal{N}_i} h_j^{(l)}$ は近傍の特徴量のSUM集約です。この和全体をMLPに通すことで、「集約した後に非線形変換する」のではなく「集約した後の表現をMLPが適切に分離できるように変換する」という設計になっています。$\epsilon$ は学習可能なパラメータとして自己ループの強さを調整しますが、固定値($\epsilon=0$)にしても多くのタスクで十分な性能が出ることが報告されています。

これらの4手法の比較をまとめると、GCNは実装のシンプルさと安定性で定番、GATは注意機構で解釈可能性と精度を両立、GraphSAGEは大規模グラフへのスケーラビリティ、GINは理論的に最強の識別能力を提供するという役割分担があります。MPNNのフレームワークを知っていれば、「次の手法はどの部品を変えているか」という視点で新しいGNNを素早く理解できます。


集約関数の識別能力:なぜGINはSUMにこだわるのか

集約関数の選択は予測精度に直結する重要な設計判断です。SUM・MEAN・MAX の3つを比較してみましょう。

集約関数の違い(SUM・MEAN・MAX)の比較

4つの近傍ノードが特徴量 $[1.0, 2.0], [3.0, 1.0], [2.0, 3.0], [0.5, 1.5]$ を持つ場合、SUM は $[6.5, 7.5]$、MEAN は $[1.625, 1.875]$、MAX は $[3.0, 3.0]$ となります。

MEANは「近傍の数」を区別できない問題があります。 以下の例を見てください。

GINの動機:SUM集約とMEAN集約の識別能力の違い

全ノードが同じ特徴量 0.5 を持つとき、4ノードのグラフAと2ノードのグラフBでは、MEAN集約の結果がどちらも 0.5 になり区別できません。一方SUM集約では A=2.0、B=1.0 となり正しく区別できます。

Xuら(2019)の理論解析では、「ノードの多重集合(multiset)を区別する能力という観点では、SUM集約を用いたMPNNがWeisfeiler-Lehman(WL)グラフ同型テストと同等の識別能力を持つ最強の集約」であることが証明されています。MAXは最大値しか見ないため集合の規模情報を完全に失い、MEANは平均しか見ないため要素数情報を失います。GINがSUMにこだわる理由はここにあります。

ただし実用上は、タスクによって最適な集約が異なります。分子特性予測のようにグラフの「組成」(各原子が何個あるか)が重要なタスクにはSUM、相対的な比率が重要なタスクにはMEAN、最大の活性化を捉えたいタスクにはMAXが有効です。


NumPy実装:メッセージパッシングの数値例

理論を数式で理解したら、実際に手を動かして確認しましょう。3ノードのグラフを例に、メッセージパッシング1ステップを NumPy で実装します。

グラフの構造は次のとおりです。

  • ノード数:3
  • エッジ:0-1, 1-2, 0-2(完全グラフ)
  • 初期特徴量:$h_0 = [1, 0]$、$h_1 = [0, 1]$、$h_2 = [1, 1]$
import numpy as np

np.random.seed(42)

# 隣接行列(ノード間の接続関係)
A = np.array([[0, 1, 1],
              [1, 0, 1],
              [1, 1, 0]], dtype=float)

# ノード特徴量 (ノード数 × 特徴次元) = (3, 2)
H = np.array([[1.0, 0.0],   # ノード0
              [0.0, 1.0],   # ノード1
              [1.0, 1.0]])  # ノード2

# 重み行列(特徴次元 × 特徴次元)
W = np.array([[0.5, 0.3],
              [0.2, 0.6]])

# ステップ①:メッセージ計算 M = H @ W
M = H @ W
print("メッセージ M = H @ W:")
print(M)
# [[0.5  0.3 ]   <- ノード0のメッセージ
#  [0.2  0.6 ]   <- ノード1のメッセージ
#  [0.7  0.9 ]]  <- ノード2のメッセージ

# ステップ②:集約(SUM): H_agg = A @ M
H_agg = A @ M
print("\n集約後 H_agg = A @ M:")
print(H_agg)
# [[0.9  1.5]  <- ノード0が受け取る(ノード1+ノード2のメッセージ和)
#  [1.2  1.2]  <- ノード1が受け取る(ノード0+ノード2のメッセージ和)
#  [0.7  0.9]] <- ノード2が受け取る(ノード0+ノード1のメッセージ和)

# ステップ③:更新 ReLU(H + H_agg)(簡略化した残差更新)
H_new = np.maximum(0, H + H_agg)
print("\n更新後 H_new = ReLU(H + H_agg):")
print(H_new)
# [[1.9  1.5]
#  [1.2  2.2]
#  [1.7  1.9]]

計算結果を確認しましょう。ノード0は元の特徴量 $[1, 0]$ から更新後 $[1.9, 1.5]$ に変化しています。ノード1と2からの情報(メッセージ $[0.9, 1.5]$)を受け取り、それを自分の表現に加算した結果です。

特に注目してほしいのはノード0の2次元目です。元は0でしたが、近傍ノード(特にノード1は $h_{1,1}=1$、ノード2は $h_{2,1}=1$)からの情報を受け取ることで 1.5 に増加しました。近傍の情報が「伝播」している様子が数値として確認できます。

NumPy実装によるメッセージパッシング1ステップの数値例

左のグラフは更新前と更新後のノード特徴量の比較です。全ノードの特徴量が更新後に増加(または維持)していることが確認できます。右のグラフは各ノードへの集約メッセージ($A \cdot M$ の各行)を示しています。ノード0とノード1はどちらも2つの近傍からメッセージを受け取り、それぞれ異なる集約値になっています。完全グラフ(全ノード間にエッジあり)では1ステップで全ノードの情報が共有されることが見て取れます。


過平滑化(Over-smoothing):深さの落とし穴

kホップ受容野の拡大を見ると「層を深くすれば深くするほど良い」と思えます。しかし実際には、GNNは層を深くしすぎると性能が急激に低下します。その原因が「過平滑化(over-smoothing)」です。

過平滑化とは、メッセージパッシングを多数回繰り返した結果、すべてのノードの表現が似通ってしまい、区別がつかなくなる現象です。直感的に言えば「情報が全体に混ざりすぎて、各ノードの個性が失われる」状態です。

数学的には、正規化された隣接行列 $\hat{D}^{-1/2} \hat{A} \hat{D}^{-1/2}$ の固有値が1以下であることから、繰り返し適用すると特徴量が低周波成分(グラフ全体の平均的な値)に収束することが示されます。

過平滑化:層数と特徴量分散の低下

このグラフは、30ノードのランダムグラフに対して正規化された隣接行列を繰り返し掛けたときの、ノード特徴量の分散の変化を示しています。層数が増えるにつれて分散が急速に低下し、層5〜6以降ではほぼ0に収束しています。これは全ノードの表現がほぼ同一になったことを意味し、ノードごとの予測に必要な「個性」が失われた状態です。

実用的には、多くのタスクで 2〜4層のGNNが最適とされています。過平滑化への対策として、残差接続(ResNet的なスキップ接続)、DropEdge(ランダムにエッジを除去して情報の過度な混合を防ぐ)、PairNorm(特徴量の分散を一定に保つ正規化)などの手法が提案されています。過平滑化の詳しい分析と対策については、

画像なし
GNNの過平滑化とその対策
過平滑化の理論的背景と残差接続・DropEdge・PairNormによる対策を解説します

MPNNの実用例:ノード分類タスク

最後に、MPNNがノード分類タスクをどのように解くかを視覚的に確認しましょう。

MPNNによるノード分類:メッセージ伝播の前後

7ノードのグラフで、ノード0と2がクラスAのラベルを持ち、ノード3と6がクラスBのラベルを持つとします(半教師あり設定)。伝播前は大部分のノードが「未分類(灰色)」です。2層のMPNNを適用すると、ラベル付きノードからの情報が2ホップ以内の近傍に伝播し、ほとんどのノードが正しく分類されます。

このような半教師あり設定でのノード分類がGNNの最も成功した応用の一つです。引用グラフ(論文と引用関係のグラフ)でのトピック分類、ソーシャルネットワークでのコミュニティ発見、タンパク質相互作用グラフでの機能予測などがあります。

グラフ分類タスクでは、分子グラフの毒性予測・溶解性予測・生理活性予測がよく知られた応用例です。化合物の原子(ノード)・化学結合(エッジ)をグラフとして表現し、MPNNでノード表現を学習したのちReadoutでグラフ全体の表現を作り、最終的に分類器に渡します。Gilmerら(2017)の元論文も量子化学の分子特性予測を対象としており、原子特徴量(原子番号・電荷・結合数など)とエッジ特徴量(結合の種類・結合距離)を使ったMPNNが量子力学的な計算(DFT計算)を高精度に近似できることを示しました。


MPNNと隣接行列の行列表現

メッセージパッシングを行列の言葉に翻訳すると、GCNの更新式が次のように表せます。

$$ \begin{equation} H^{(l+1)} = \sigma\!\left(\hat{A}_\mathrm{norm}\, H^{(l)}\, W^{(l)}\right) \end{equation} $$

ここで $\hat{A}_\mathrm{norm} = \hat{D}^{-1/2} \hat{A} \hat{D}^{-1/2}$ が正規化された隣接行列です。

$\hat{A}_\mathrm{norm} H^{(l)}$ という行列積が「集約」に対応します。$i$ 行目に注目すると、

$$ \begin{equation} \left(\hat{A}_\mathrm{norm} H^{(l)}\right)_i = \sum_{j} (\hat{A}_\mathrm{norm})_{ij} \cdot h_j^{(l)} \end{equation} $$

となり、各近傍 $j$ の特徴ベクトル $h_j^{(l)}$ を正規化係数 $(\hat{A}_\mathrm{norm})_{ij} = \frac{1}{\sqrt{d_i d_j}}$($d_i, d_j$ はノードの次数)で重み付けして足し合わせる操作になっています。

ここで $\frac{1}{\sqrt{d_i d_j}}$ という係数の意味を考えましょう。次数が高いノード(多くのエッジを持つノード)は多くの近傍からメッセージを受け取るため、集約値が大きくなりやすいです。この係数はその不均衡を補正するためのものです。具体的には、送信側ノード $j$ の次数 $d_j$ が大きいほど $j$ から届く各メッセージの寄与を小さくし、受信側ノード $i$ の次数 $d_i$ が大きいほど全体の集約値をスケールダウンします。

この行列形式を使うと、メッセージパッシングの1ステップ全体(全ノード同時)が単一の行列積で表せることがわかります。NumPy実装では H_agg = A @ M という1行で全ノードの集約が完了しますが、これは行列積がまさにこの操作を実行しているためです。


PyTorch Geometricとの対応

理解を深めるために、PyTorch Geometric(PyG)の MessagePassing クラスとMPNNの対応を確認しておきましょう。

from torch_geometric.nn import MessagePassing
import torch
import torch.nn.functional as F

class SimpleGNN(MessagePassing):
    def __init__(self, in_channels, out_channels):
        # aggr="add" がSUM集約、"mean" が平均、"max" が最大値集約に対応
        super().__init__(aggr="add")
        self.linear = torch.nn.Linear(in_channels, out_channels)

    def forward(self, x, edge_index):
        # x: ノード特徴量 (ノード数, in_channels)
        # edge_index: エッジリスト (2, エッジ数)
        x = self.linear(x)
        # propagate がメッセージ計算→集約→更新を一括実行
        return self.propagate(edge_index, x=x)

    def message(self, x_j):
        # x_j: 送信元ノード(ノードj)の特徴量 → メッセージ関数 M の実装
        # ここでは恒等写像(x_j をそのままメッセージとして送る)
        return x_j

    def update(self, aggr_out):
        # aggr_out: 集約後の値 → 更新関数 U の実装
        return F.relu(aggr_out)

message() メソッドがメッセージ関数 $M$、propagate()aggr= 引数が集約 $\bigoplus$、update() メソッドが更新関数 $U$ に正確に対応しています。GATであれば message() 内で注意重みを計算し、aggr="add" で重み付き和を行います。GraphSAGEであれば update() 内で自分自身の特徴量との concatenation を行います。

このクラスを継承するだけで、式(1)〜(3)のフレームワーク内の任意のGNNを実装できます。


まとめ

本記事では、MPNNの一般フレームワークとその主要な要素について解説しました。

  • メッセージ関数 $M_t$:近傍ノードとエッジの情報からメッセージを計算する。MLPなどを使用。
  • 集約 $\bigoplus$:全近傍からのメッセージをSUM/MEAN/MAXなどで集約する。置換不変性が必須。
  • 更新関数 $U_t$:集約結果と自分の表現から次層の表現を作る。
  • Readout 関数 $R$:グラフ分類ではノード表現をグラフ全体の表現に圧縮する。
  • 主要GNNはMPNNの特殊形:GCN(均等集約)、GAT(注意重み付き集約)、GraphSAGE(concat+集約)、GIN(SUM+MLP)はすべてこの枠組みに収まる。
  • 識別能力の理論的保証:SUM集約はWLテストと同等の最強の識別能力を持つ(GINの動機)。
  • kホップ受容野と過平滑化:層を増やすと受容野が広がるが、深くなりすぎると全ノードが均一化(過平滑化)する。実用的には2〜4層が最適。

GNNを実装したい方は、まず GCN(最もシンプル)を実装し、次に GAT(注意機構)、GraphSAGE(スケーラビリティ)、GIN(理論的最強)と進んでいくとMPNNの全体像が掴みやすくなります。

以下の記事も参考にしてください。

画像なし
GCN(Graph Convolutional Network)をPyTorchで実装する
MPNNの最もシンプルな実装であるGCNの理論と実装を解説します
GAT(Graph Attention Network)の仕組みとPyTorch実装
注意機構を使ってMPNNを拡張したGATの理論と実装を解説します