从零理解图卷积网络(GCN):理论推导与 PyTorch 实现
传统神经网络在表格数据,图像和序列数据上取得了很好的效果,但现实也充满了大量图结构数据例如:人际关系网,蛋白质结构,分子键等。图神经网络(GNN)是处理图结构数据的强大工具,而图卷积网络(GCN)作为其经典代表,通过“聚合邻居信息”的思想实现了节点级别的表示学习。本文将结合手写笔记与代码,带你从数学原理到 Python 实现,彻底搞懂 GCN。
一、GCN 理论笔记完整版
1. 基本符号与数据示例
节点特征矩阵
:
其中第一列是年龄,第二列是分数。
邻接矩阵
:
度矩阵
:
2. 图卷积网络更新公式
图卷积网络(GCN)的核心公式为:
其中:
-
表示第
层的节点特征
-
是可学习的权重矩阵
-
是非线性激活函数
-
是归一化的邻接矩阵
3. 归一化邻接矩阵的计算
3.1 添加自环
首先给邻接矩阵添加自环:
3.2 计算带自环的度矩阵
3.3 对称归一化
GCN 使用对称归一化:
其中:
4. 为什么需要归一化?(手工计算示例)
4.1 未归一化的情况
当时
计算第一层聚合特征
:
节点1的聚合结果:
节点2的聚合结果:
问题:聚合后节点1变为,节点2变为
。聚合结果主要取决于邻居数量,邻居数量越多则聚合特征更大。不同节点输出结果尺度不统一会造成:
(1)训练不稳定:梯度消失或爆炸。
(2)模型梯度被度数大的节点主导。
(3)无法公平比较节点:不是因为特征重要而是因为邻居多。
4.2 归一化后的情况
使用归一化邻接矩阵 :
归一化后 近似为(python计算后保留两位小数):
节点1的聚合结果:
节点3的聚合结果:
优点:归一化后,聚合结果既考虑邻居特征,又考虑了节点的重要性,输出更稳定。
5. 归一化邻间矩阵
的广播机制实现推导
5.1 数学公式
归一化公式:
由于
则
由于D是对角矩阵即则
设 为度向量,其中
5.2 广播机制推导
令 为度向量的平方根倒数。
第一步:构造列向量
在代码中:row_sum_inv_sqrt.view(-1, 1)
第二步:左乘
在代码中:normalized_adj = row_sum_inv_sqrt * adj
第三步:构造行向量
在代码中:row_sum_inv_sqrt.view(1, -1)
第四步:右乘
在代码中:normalized_adj = normalized_adj * row_sum_inv_sqrt其中表示逐元素相乘的哈达玛乘积。
5.3 为什么这样高效?
传统方法需要显式构造对角矩阵:
D = torch.diag(row_sum_inv_sqrt)
normalized_adj = D @ adj @ D
复杂度:,需要两个矩阵乘法。
广播方法:
# 只需要逐元素乘法
row_sum_inv_sqrt_col = row_sum_inv_sqrt.view(-1, 1) # (n,1)
row_sum_inv_sqrt_row = row_sum_inv_sqrt.view(1, -1) # (1,n)
normalized_adj = row_sum_inv_sqrt_col * adj * row_sum_inv_sqrt_row
复杂度:,利用广播机制,无需显式构造对角矩阵,对如果节点数较多内存友好。
6. 其他注意事项
归一化后,聚合权重与节点度数相关:
- 高度数节点:对邻居的影响较小(权重小)
- 低度数节点:对邻居的影响较大(权重大)
这更符合现实世界的图结构特性。
7. 完整计算流程总结
-
输入:节点特征
,邻接矩阵
-
添加自环:
-
计算度向量:
-
计算归一化向量:
-
广播归一化:
-
特征聚合:
-
线性变换:
-
激活函数:
二、关键总结
| 要点 | 说明 |
|---|---|
| 归一化目的 | 防止梯度爆炸/消失,平衡节点重要性 |
| 对称归一化 | |
| 广播机制 | 避免构造对角矩阵,降低计算复杂度 |
| 自环添加 | 确保节点自身特征参与聚合 |
| 时间复杂度 | 从 |
# 广播机制核心代码总结
def normalize_adj(adj):
d = adj.sum(dim=1) # 度向量
d_inv_sqrt = d.pow(-0.5) # d^{-1/2}
d_inv_sqrt[torch.isinf(d_inv_sqrt)] = 0 # 处理孤立节点
# 广播归一化
d_inv_sqrt = d_inv_sqrt.view(-1, 1) # 列向量
norm_adj = d_inv_sqrt * adj # 左乘
d_inv_sqrt = d_inv_sqrt.view(1, -1) # 行向量
norm_adj = norm_adj * d_inv_sqrt # 右乘
return norm_adj
最终效果:通过广播机制高效实现GCN归一化,既保证了数值稳定性,又提升了计算效率。
三、GraphConvLayer 完整代码实现
# -*- coding: utf-8 -*-
"""
Created on Thu Jan 22 22:10:26 2026
@author: Bear
"""
import torch
import torch.nn as nn
class GraphConvLayer(nn.Module):
def __init__(self, in_features, out_features,
activation=None, dropout=0.1):
super().__init__()
self.in_features = in_features
self.out_features = out_features
self.linear = nn.Linear(in_features, out_features)
self.activation = activation
self.dropout = nn.Dropout(dropout) if dropout>0 else None
def compute_normalized_adjacent(self, adj, add_self_loops=True):
"""计算归一化邻接矩阵:D^{-1/2} A D^{-1/2}"""
n_nodes = adj.shape[0]
# 是否添加自环
if add_self_loops:
adj = adj + torch.eye(n_nodes)
# 计算度矩阵(对角线上为每个节点的度数)
row_sum = torch.sum(adj, dim=1)
row_sum_inv_sqrt = row_sum ** (-0.5)
row_sum_inv_sqrt[torch.isinf(row_sum_inv_sqrt)] = 0
# 计算归一化的邻接矩阵(使用广播机制)
row_sum_inv_sqrt = row_sum_inv_sqrt.view(-1, 1) # 列向量 (n, 1)
normalized_adj = row_sum_inv_sqrt * adj # 左乘 D^{-1/2}
row_sum_inv_sqrt = row_sum_inv_sqrt.view(1, -1) # 行向量 (1, n)
normalized_adj = normalized_adj * row_sum_inv_sqrt # 右乘 D^{-1/2}
return normalized_adj # (n, n)
def forward(self, x, adj, add_self_loops=True):
# 归一化邻接矩阵
normalized_adj = self.compute_normalized_adjacent(adj)
# 前向传播:聚合邻居特征
x = normalized_adj @ x # (n, f)
# dropout正则化
if self.dropout is not None:
x = self.dropout(x)
# 线性变换
x = self.linear(x)
# 激活函数
if self.activation is not None:
x = self.activation(x)
return x
# 测试案例
adj = torch.tensor([[0, 1, 0, 1],
[1, 0, 1, 0],
[0, 1, 0, 0],
[1, 0, 0, 0]]).to(torch.float32)
# 节点特征:[年龄, 分数]
x = torch.tensor([[18, 95],
[19, 87],
[20, 90],
[18, 89]]).to(torch.float32)
gcn = GraphConvLayer(2, 4)
output = gcn(x, adj)
print("输出形状:", output.shape)