2026-01-23


从零理解图卷积网络(GCN):理论推导与 PyTorch 实现

传统神经网络在表格数据,图像和序列数据上取得了很好的效果,但现实也充满了大量图结构数据例如:人际关系网,蛋白质结构,分子键等。图神经网络(GNN)是处理图结构数据的强大工具,而图卷积网络(GCN)作为其经典代表,通过“聚合邻居信息”的思想实现了节点级别的表示学习。本文将结合手写笔记与代码,带你从数学原理到 Python 实现,彻底搞懂 GCN。

一、GCN 理论笔记完整版

1. 基本符号与数据示例

节点特征矩阵 X:

X = \begin{pmatrix} 18 & 95 \\ 19 & 87 \\ 20 & 90 \\ 18 & 89 \end{pmatrix}
其中第一列是年龄,第二列是分数。

邻接矩阵 A:

A = \begin{pmatrix} 0 & 1 & 0 & 1 \\ 1 & 0 & 1 & 0 \\ 0 & 1 & 0 & 0 \\ 1 & 0 & 0 & 0 \end{pmatrix}
A_{ij} = \begin{cases} 1 & \text{节点i与节点j相连} \\ 0 & \text{节点i与节点j不相连} \end{cases}

度矩阵 D:

D = \begin{pmatrix} 2 & 0 & 0 & 0 \\ 0 & 2 & 0 & 0 \\ 0 & 0 & 1 & 0 \\ 0 & 0 & 0 & 1 \end{pmatrix}
D_{ii} = \sum_j A_{ij} = \text{节点i的度数}


2. 图卷积网络更新公式

图卷积网络(GCN)的核心公式为:

H^{(l+1)} = \sigma\left( \hat{A} H^{(l)} W^{(l)} \right)

其中:

  • H^{(l)} 表示第 l 层的节点特征
  • W^{(l)} 是可学习的权重矩阵
  • \sigma 是非线性激活函数
  • \hat{A} 是归一化的邻接矩阵

3. 归一化邻接矩阵的计算

3.1 添加自环

首先给邻接矩阵添加自环:

\tilde{A} =A+I= \begin{pmatrix} 1 & 1 & 0 & 0 \\ 1 & 1 & 1 & 0 \\ 0 & 1 & 1 & 0 \\ 1 & 0 & 0 & 1 \end{pmatrix}

3.2 计算带自环的度矩阵

\tilde{D} = \begin{pmatrix} 3 & 0 & 0 & 0 \\ 0 & 3 & 0 & 0 \\ 0 & 0 & 2 & 0 \\ 0 & 0 & 0 & 2 \end{pmatrix}

3.3 对称归一化

GCN 使用对称归一化:
\hat{A} = \tilde{D}^{-\frac{1}{2}} \tilde{A} \tilde{D}^{-\frac{1}{2}}

其中:
\tilde{D}^{-\frac{1}{2}} = \begin{pmatrix} \frac{1}{\sqrt{3}} & 0 & 0 & 0 \\ 0 & \frac{1}{\sqrt{3}} & 0 & 0 \\ 0 & 0 & \frac{1}{\sqrt{2}} & 0 \\ 0 & 0 & 0 & \frac{1}{\sqrt{2}} \end{pmatrix}


4. 为什么需要归一化?(手工计算示例)

4.1 未归一化的情况

当l=1时H^{1}=X 计算第一层聚合特征AX:

节点1的聚合结果:
(AX)_{11} = 1 \times 18 + 1 \times 19 + 0 \times 20 + 1 \times 18 = 55
(AX)_{12} = 1 \times 95 + 1 \times 87 + 0 \times 90 + 1 \times 89 = 271
节点2的聚合结果:
(AX)_{31} = 0 \times 18 + 1 \times 19 + 1 \times 20 + 0 \times 18 = 39
(AX)_{32} = 0 \times 95 + 1 \times 87 + 1 \times 90 + 0 \times 89 = 177

问题:聚合后节点1变为(55,271),节点2变为(39,177)。聚合结果主要取决于邻居数量,邻居数量越多则聚合特征更大。不同节点输出结果尺度不统一会造成:
(1)训练不稳定:梯度消失或爆炸。
(2)模型梯度被度数大的节点主导。
(3)无法公平比较节点:不是因为特征重要而是因为邻居多。


4.2 归一化后的情况

使用归一化邻接矩阵 \hat{A}X:

归一化后 \hat{A} 近似为(python计算后保留两位小数):
\hat{A} \approx \begin{pmatrix} 0.33 & 0.33 & 0 & 0.33 \\ 0.33 & 0.33 & 0.33 & 0 \\ 0 & 0.5 & 0.5 & 0 \\ 0.5 & 0 & 0 & 0.5 \end{pmatrix}

节点1的聚合结果:
(\hat{A}X)_{11} = 18 \times 0.33 + 19 \times 0.33 + 0 \times 20 + 18 \times 0.33 \approx 18.15
(\hat{A}X)_{12} = 95 \times 0.33 + 87 \times 0.33 + 0 \times 90 + 89 \times 0.33 \approx 89.43

节点3的聚合结果:
(\hat{A}X)_{31} = 18 \times 0 + 19 \times 0.5 + 20 \times 0.5 + 18 \times 0 = 19.5
(\hat{A}X)_{32}= 95 \times 0 + 87 \times 0.5 + 90 \times 0.5 + 89 \times 0 = 88.5

优点:归一化后,聚合结果既考虑邻居特征,又考虑了节点的重要性,输出更稳定。


5. 归一化邻间矩阵\hat{A} 的广播机制实现推导

5.1 数学公式

归一化公式:
\hat{A} = D^{-\frac{1}{2}} A D^{-\frac{1}{2}}
由于
(D^{-\frac{1}{2}} A)_{ik} = \sum_{m}D^{-\frac{1}{2}}_{im}A_{mk}
则
\hat{A}=(D^{-\frac{1}{2}} AD^{-\frac{1}{2}})_{ij}=\sum_{k}(D^{-\frac{1}{2}}A_{ik})D^{-\frac{1}{2}}\sum_{k}\sum_{m}(D^{-\frac{1}{2}}_{im}A_{mk})D^{-\frac{1}{2}}_{kj}
由于D是对角矩阵即D^{-\frac{1}{2}}_{ij} = \begin{cases} 1 & \text{i=j} \\ 0 & \text{i≠j} \end{cases}则
\hat{A}=(D^{-\frac{1}{2}} AD^{-\frac{1}{2}})_{ii}=D^{-\frac{1}{2}} _{ii}A_{ij}D^{-\frac{1}{2}} _{jj}

设 \mathbf{d} 为度向量,其中 \mathbf{d}_i = D_{ii}

5.2 广播机制推导

令 \mathbf{v} = \mathbf{d}^{-\frac{1}{2}} 为度向量的平方根倒数。

第一步:构造列向量
\mathbf{v}_{\text{col}} = \begin{pmatrix} v_1 \\ v_2 \\ \vdots \\ v_n \end{pmatrix}
在代码中:row_sum_inv_sqrt.view(-1, 1)

第二步:左乘 D^{-\frac{1}{2}}
B_{ij} = v_i \odot A_{ij}
在代码中:normalized_adj = row_sum_inv_sqrt * adj

第三步:构造行向量
\mathbf{v}_{\text{row}} = \begin{pmatrix} v_1 & v_2 & \cdots & v_n \end{pmatrix}
在代码中:row_sum_inv_sqrt.view(1, -1)

第四步:右乘 D^{-\frac{1}{2}}
\hat{A}_{ij} = v_i \odot A_{ij} \odot v_j
在代码中:normalized_adj = normalized_adj * row_sum_inv_sqrt其中\odot表示逐元素相乘的哈达玛乘积。


5.3 为什么这样高效?

传统方法需要显式构造对角矩阵:

D = torch.diag(row_sum_inv_sqrt)
normalized_adj = D @ adj @ D

复杂度:O(n^3),需要两个矩阵乘法。

广播方法:

# 只需要逐元素乘法
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

复杂度:O(n^2),利用广播机制,无需显式构造对角矩阵,对如果节点数较多内存友好。


6. 其他注意事项

归一化后,聚合权重与节点度数相关:

  • 高度数节点:对邻居的影响较小(权重小)
  • 低度数节点:对邻居的影响较大(权重大)
    这更符合现实世界的图结构特性。

7. 完整计算流程总结

  1. 输入:节点特征 X,邻接矩阵 A
  2. 添加自环:\tilde{A} = A + I
  3. 计算度向量:\mathbf{d} = \text{sum}(\tilde{A}, \text{dim}=1)
  4. 计算归一化向量:\mathbf{v} = \mathbf{d}^{-\frac{1}{2}}
  5. 广播归一化:\hat{A} = \mathbf{v}_{\text{col}} \odot \tilde{A} \odot \mathbf{v}_{\text{row}}
  6. 特征聚合:Z = \hat{A} X
  7. 线性变换:H = Z W
  8. 激活函数:H' = \sigma(H)

二、关键总结

要点 说明
归一化目的 防止梯度爆炸/消失,平衡节点重要性
对称归一化 \hat{A} = D^{-1/2} A D^{-1/2}
广播机制 避免构造对角矩阵,降低计算复杂度
自环添加 确保节点自身特征参与聚合
时间复杂度 从 O(n^3) 优化到 O(n^2)
# 广播机制核心代码总结
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)

最后编辑于 :
©著作权归作者所有,转载或内容合作请联系作者
【社区内容提示】社区部分内容疑似由AI辅助生成,浏览时请结合常识与多方信息审慎甄别。
平台声明:文章内容(如有图片或视频亦包括在内)由作者上传并发布,文章内容仅代表作者本人观点,简书系信息发布平台,仅提供信息存储服务。

相关阅读更多精彩内容

友情链接更多精彩内容