从矩阵公式到 PyTorch 代码:手写一个两层 GCN

Chen Xi
Chen Xi

上一篇文章从消息传递的角度介绍了 GCN、GraphSAGE 和 GAT。公式看懂之后,我一直觉得还差一步:矩阵里的每一项,放进代码到底是什么?

这篇文章只做一件事——不依赖 PyTorch Geometric 的卷积层,用 PyTorch 写出一个最小的两层 GCN。代码以 Cora 这类节点分类数据为背景,但不展开数据下载和准确率比较,重点是公式与实现之间的对应关系。

从一层 GCN 开始

设图的邻接矩阵为 AA,节点特征为 XX。GCN 首先给每个节点加一条指向自己的边:

A^=A+I\hat{A}=A+I

然后根据加入自环后的节点度,构造对称归一化邻接矩阵:

S=D^12A^D^12S= \hat{D}^{-\frac{1}{2}} \hat{A} \hat{D}^{-\frac{1}{2}}

一层 GCN 可以写为:

H(l+1)=σ(SH(l)W(l))H^{(l+1)} = \sigma \left( S H^{(l)} W^{(l)} \right)

这里的计算可以拆成两步:

  1. H(l)W(l)H^{(l)}W^{(l)}:对每个节点做相同的线性变换;
  2. S()S(\cdot):按照图结构聚合自身和邻居的表示。

矩阵乘法满足结合律,所以先聚合再线性变换也可以。实际工程中应根据特征维度、隐藏维度和稀疏矩阵的计算代价选择顺序。

image

先把边变成归一化稀疏矩阵

图数据通常不会直接保存一个 N×NN\times N 的稠密邻接矩阵,而是保存两行边索引:

1
edge_index.shape = [2, E]

第一行是源节点,第二行是目标节点。下面的函数加入反向边和自环,再计算每条边的归一化权重:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
import torch


def normalized_adjacency(edge_index, num_nodes, device):
src, dst = edge_index.to(device)

# 无向图:补上反向边,并加入自环
nodes = torch.arange(num_nodes, device=device)
edges = torch.cat([
torch.stack([src, dst]),
torch.stack([dst, src]),
torch.stack([nodes, nodes]),
], dim=1)

# 输入可能已包含反向边,先去重再计算度
edges = torch.unique(edges, dim=1)
src, dst = edges

# 计算加入自环后的度
degree = torch.zeros(num_nodes, device=device)
degree.scatter_add_(0, dst, torch.ones_like(dst, dtype=torch.float))

# 每条边 u -> v 的权重为 1 / sqrt(deg(u) * deg(v))
weight = degree[src].pow(-0.5) * degree[dst].pow(-0.5)

indices = torch.stack([dst, src])
adjacency = torch.sparse_coo_tensor(
indices,
weight,
size=(num_nodes, num_nodes),
device=device,
)
return adjacency.coalesce()

代码中的 indices = [dst, src] 容易让人疑惑。稀疏矩阵第 vv 行、第 uu 列的值,表示节点 uu 向节点 vv 发送消息。因此边方向是 src -> dst,矩阵坐标却写成 [dst, src]

输入数据可能本来就同时保存正向边和反向边,因此代码在补边之后调用 torch.unique 去重。最后的 coalesce() 会整理稀疏坐标。真实任务中的重复边可能代表交互次数或边权,是否应该直接去重,需要结合数据语义决定;这里把它当作简单无权无向图处理。

GCN 层其实很短

有了归一化矩阵 SS,一层 GCN 只剩线性变换和稀疏矩阵乘法:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
from torch import nn


class GCNLayer(nn.Module):
def __init__(self, in_features, out_features):
super().__init__()
self.weight = nn.Parameter(
torch.empty(in_features, out_features)
)
self.bias = nn.Parameter(torch.zeros(out_features))
nn.init.xavier_uniform_(self.weight)

def forward(self, x, adjacency):
support = x @ self.weight
aggregated = torch.sparse.mm(adjacency, support)
return aggregated + self.bias

假设输入特征为:

XRN×FX\in\mathbb{R}^{N\times F}

权重矩阵为:

WRF×HW\in\mathbb{R}^{F\times H}

那么 support 的形状是 N×HN\times H。左乘 N×NN\times N 的稀疏邻接矩阵后,输出仍然是 N×HN\times H,只是每一行已经融合了邻居信息。

这里没有在 GCNLayer 内部写 ReLU。卷积层只负责传播,激活、Dropout 和层间组合交给模型本身,结构会更清楚。

组成两层节点分类模型

第一层把原始特征映射到隐藏空间,第二层把隐藏表示映射到类别数:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
import torch.nn.functional as F


class GCN(nn.Module):
def __init__(self, in_features, hidden, num_classes, dropout=0.5):
super().__init__()
self.gcn1 = GCNLayer(in_features, hidden)
self.gcn2 = GCNLayer(hidden, num_classes)
self.dropout = dropout

def forward(self, x, adjacency):
x = self.gcn1(x, adjacency)
x = F.relu(x)
x = F.dropout(x, p=self.dropout, training=self.training)
return self.gcn2(x, adjacency)

两层 GCN 连续聚合两次。一层后,每个节点包含一跳邻居的信息;第二层继续传播后,理论上可以利用两跳邻域。

这不代表层数越多越好。不断左乘归一化邻接矩阵,会让相邻节点的表示逐渐接近。层数太深时,不同类别的节点可能变得难以区分,这就是常说的过平滑。

训练时为什么只对一部分节点算损失

Cora 一类节点分类任务通常提供 train_maskval_masktest_mask。模型前向传播会使用整张图和所有节点特征,但损失只在训练节点上计算:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
device = x.device
adjacency = normalized_adjacency(
edge_index=edge_index,
num_nodes=x.size(0),
device=device,
)

model = GCN(
in_features=x.size(1),
hidden=16,
num_classes=int(y.max()) + 1,
).to(device)

optimizer = torch.optim.Adam(
model.parameters(),
lr=0.01,
weight_decay=5e-4,
)

for epoch in range(200):
model.train()
optimizer.zero_grad()

logits = model(x, adjacency)
loss = F.cross_entropy(
logits[train_mask],
y[train_mask],
)

loss.backward()
optimizer.step()

model.eval()
with torch.no_grad():
prediction = model(x, adjacency).argmax(dim=1)
val_acc = (
prediction[val_mask] == y[val_mask]
).float().mean()

这种设置属于转导学习:验证节点和测试节点可以出现在训练使用的图结构中,它们的标签不能进入损失。如果任务要求预测训练时从未出现的新节点,就需要重新考虑数据划分和归纳式模型。

三个很容易写错的地方

1. 忘记加入自环

不加自环时,节点更新只依赖邻居。原始特征不能直接保留下来,孤立节点甚至无法获得有效消息。自环不是装饰,而是传播规则的一部分。

2. 归一化方向写反

如果 edge_index 表示 src -> dst,稀疏矩阵坐标通常应为 [dst, src]。方向写反后,程序仍能运行,但传播语义已经改变。这类错误比语法错误更难发现。

3. 在所有节点上计算训练损失

1
2
# 错误示例
loss = F.cross_entropy(logits, y)

这会直接使用验证集和测试集标签,得到的结果没有意义。正确做法是始终用 train_mask 限定监督信号。

最后看回公式

手写完以后,一层 GCN 就不再神秘了:

1
2
support = x @ weight
output = sparse_adjacency @ support

第一行学习“特征应该怎样变换”,第二行决定“节点应该从谁那里接收信息”。GCN 的核心正是把可学习的特征变换和固定的图结构传播放到同一个计算图里。

真正应用时,还要处理有向边、边权、重复边、大图采样和动态图切分。但只要理解了这个最小实现,再阅读 PyTorch Geometric 等框架的 GCNConv,就能分清哪些是模型本身,哪些只是框架替我们完成的工程细节。