AI图神经网络折腾手记

先说结论:GNN 适合处理那些数据之间存在复杂关系、且关系本身也包含信息的场景。

社交网络、推荐系统、化学分子结构、交通路网、金融风控、知识图谱——这些场景的共同点就是:不能只看单个节点的特征,还要看它跟谁连着、连的方式是什么、周围的节点都在干什么。

图神经网络能解决什么

先说结论:GNN 适合处理那些数据之间存在复杂关系、且关系本身也包含信息的场景。

社交网络、推荐系统、化学分子结构、交通路网、金融风控、知识图谱——这些场景的共同点就是:不能只看单个节点的特征,还要看它跟谁连着、连的方式是什么、周围的节点都在干什么。

以社交网络异常检测为例:

  • 单看一个用户的特征:发帖数量、登录时间、好友数
  • 用图神经网络看:这个用户的好友圈里有没有其他异常账号?他的互动网络跟正常用户有什么区别?他的连接模式是自然的还是刻意构建的?

前者只看个体,后者看关系网络。

从图数据到模型输入

刚开始搞的时候,我以为图数据就是邻接矩阵加上节点特征矩阵,这事儿不复杂。但实际跑了一圈才发现,现实世界的图数据比教科书上的例子麻烦多了。

图数据的常见表示方式

# 方式1:邻接矩阵 + 特征矩阵
import torch
import numpy as np

# 假设有 100 个节点,每个节点 50 维特征
num_nodes = 100
feature_dim = 50

# 邻接矩阵:100x100,0表示无边,1表示有边
adj_matrix = torch.zeros((num_nodes, num_nodes))
adj_matrix[0, 1] = 1  # 节点0指向节点1

# 节点特征:100x50
node_features = torch.randn(num_nodes, feature_dim)

这种方式看着规整,但真实图一进来就暴露问题:

  1. 稀疏矩阵存储效率低:社交网络百万级节点,邻接矩阵存储成稠密就爆内存了
  2. 不支持边特征:边本身也带信息(比如关系类型、权重),用 0/1 表示不完
  3. 不方便动态扩展:新节点加进来就要整个矩阵扩维
# 方式2:边列表 + 特征矩阵
edge_index = torch.tensor([
    [0, 1, 2, 3],  # 源节点
    [4, 5, 6, 7]   # 目标节点
], dtype=torch.long)

# 边特征(可选):每条边 10 维
edge_attr = torch.randn(4, 10)

这种方式在 PyTorch Geometric 和 DGL 里用得比较多。边列表省空间,扩展也方便。

# 方式3:NetworkX 图对象(适合数据预处理)
import networkx as nx

G = nx.Graph()
G.add_node(0, feature=[0.1, 0.2, 0.3])
G.add_node(1, feature=[0.4, 0.5, 0.6])
G.add_edge(0, 1, weight=0.8)

# 转成 PyTorch Geometric 格式
from torch_geometric.utils import from_networkx
data = from_networkx(G)

NetworkX 在预处理阶段很好用,它有各种图算法和可视化工具,但真正训练时还是会转成张量形式。

实际数据准备踩坑

在做社交网络项目时,我踩过几个典型的坑:

坑1:节点ID不连续

真实数据里经常是这样:节点ID是用户ID,比如 10001、10002、10003,不是 0、1、2。直接拿来建邻接矩阵会浪费大量空间。

# 解决方案:做ID映射
original_ids = [10001, 10002, 10003, 20001]
id_map = {old_id: new_id for new_id, old_id in enumerate(original_ids)}

# 映射后的邻接矩阵:4x4,而不是 20002x20002

坑2:图太大,内存不够

百万级节点的图,邻接矩阵存储就够呛,更别说训练时还要存各种中间结果。

# 解决方案:分批采样 + 子图训练
from torch_geometric.loader import NeighborSampler

train_loader = NeighborSampler(data.edge_index, size=[10, 10],
                               batch_size=1024, shuffle=True)

PyTorch Geometric 的 NeighborSampler 会采样每个节点的邻居,构成小的子图进行训练,避免一次性加载全图。

坑3:有向图 vs 无向图搞错

社交网络里的"关注"关系是有向的,但论文里很多例子用的都是无向图。直接照搬会搞反信息流向。

# 检查图是否是有向的
print(f"是否为有向图: {G.is_directed()}")

# 如果需要把有向图转为无向图(注意是否合理)
G_undirected = G.to_undirected()

图卷积网络实现

搞完数据准备,接下来是模型。图神经网络里最基础也最常用的是图卷积网络(GCN),它本质上是卷积在图结构上的推广。

GCN 核心公式

GCN 的核心思想是:节点的表示聚合它邻居的信息。

$$H^{(l+1)} = \sigma(\tilde{D}^{-1/2} \tilde{A} \tilde{D}^{-1/2} H^{(l)} W^{(l)})$$

看起来很复杂,实际代码实现就几行:

import torch
import torch.nn as nn
import torch.nn.functional as F

class GCNLayer(nn.Module):
    def __init__(self, in_features, out_features):
        super(GCNLayer, self).__init__()
        self.linear = nn.Linear(in_features, out_features)

    def forward(self, x, adj):
        # x: 节点特征 [num_nodes, in_features]
        # adj: 邻接矩阵 [num_nodes, num_nodes]

        # 添加自环(节点自己也算邻居)
        adj = adj + torch.eye(adj.size(0), device=adj.device)

        # 计算度矩阵的对角线
        degree = torch.sum(adj, dim=1)
        d_inv_sqrt = torch.pow(degree, -0.5)
        d_inv_sqrt[torch.isinf(d_inv_sqrt)] = 0.0

        # 归一化邻接矩阵
        adj_norm = adj * d_inv_sqrt.view(-1, 1) * d_inv_sqrt.view(1, -1)

        # 特征变换 + 邻居聚合
        x = self.linear(x)
        x = torch.matmul(adj_norm, x)

        return F.relu(x)

这个实现虽然简单,但跟实际框架比还差不少:没有批处理、不支持边特征、计算效率也不高。

用 PyTorch Geometric 实现

实际项目里我会直接用 PyTorch Geometric,它已经把各种 GNN 模块封装好了:

from torch_geometric.nn import GCNConv
from torch_geometric.data import Data

class GCN(nn.Module):
    def __init__(self, num_features, hidden_dim, num_classes, num_layers=2):
        super(GCN, self).__init__()
        self.convs = nn.ModuleList()

        # 输入层
        self.convs.append(GCNConv(num_features, hidden_dim))

        # 隐藏层
        for _ in range(num_layers - 2):
            self.convs.append(GCNConv(hidden_dim, hidden_dim))

        # 输出层
        self.convs.append(GCNConv(hidden_dim, num_classes))

    def forward(self, data):
        x, edge_index = data.x, data.edge_index

        for i, conv in enumerate(self.convs[:-1]):
            x = conv(x, edge_index)
            x = F.relu(x)
            x = F.dropout(x, p=0.5, training=self.training)

        x = self.convs[-1](x, edge_index)

        return F.log_softmax(x, dim=1)

PyTorch Geometric 的 GCNConv 已经内置了自环添加和归一化,用起来简洁很多。

GAT 实践:注意力机制加持

GCN 的缺点是:所有邻居一视同仁,不区分哪个邻居更重要。GAT(Graph Attention Network)引入注意力机制,让模型自己决定哪些邻居该重点参考。

from torch_geometric.nn import GATConv

class GAT(nn.Module):
    def __init__(self, num_features, hidden_dim, num_classes, heads=4):
        super(GAT, self).__init__()
        self.conv1 = GATConv(num_features, hidden_dim, heads=heads, dropout=0.6)
        self.conv2 = GATConv(hidden_dim * heads, num_classes, heads=1, dropout=0.6)

    def forward(self, data):
        x, edge_index = data.x, data.edge_index

        x = F.dropout(x, p=0.6, training=self.training)
        x = self.conv1(x, edge_index)
        x = F.elu(x)
        x = F.dropout(x, p=0.6, training=self.training)
        x = self.conv2(x, edge_index)

        return F.log_softmax(x, dim=1)

在社交网络异常检测这个场景里,GAT 表现确实比 GCN 好一点:它可以自动学习哪些互动关系更重要,比如某些好友的异常行为更值得关注。

节点分类完整训练流程

搞完模型,接下来就是训练了。这里以 Cora 数据集(学术论文分类)为例,展示完整的节点分类流程。

from torch_geometric.datasets import Planetoid
from torch_geometric.utils import train_test_split_edges

# 加载数据集
dataset = Planetoid(root='/tmp/Cora', name='Cora')
data = dataset[0]

# 数据集信息
print(f"节点数: {data.num_nodes}")
print(f"边数: {data.num_edges}")
print(f"特征维度: {data.num_node_features}")
print(f"类别数: {dataset.num_classes}")

# 划分训练集/验证集/测试集
data.train_mask = torch.zeros(data.num_nodes, dtype=torch.bool)
data.val_mask = torch.zeros(data.num_nodes, dtype=torch.bool)
data.test_mask = torch.zeros(data.num_nodes, dtype=torch.bool)

data.train_mask[:140] = True
data.val_mask[140:640] = True
data.test_mask[640:] = True

# 初始化模型
model = GCN(num_features=data.num_node_features,
            hidden_dim=16,
            num_classes=dataset.num_classes,
            num_layers=2).to(device)

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

# 训练循环
def train():
    model.train()
    optimizer.zero_grad()
    out = model(data)
    loss = F.nll_loss(out[data.train_mask], data.y[data.train_mask])
    loss.backward()
    optimizer.step()
    return loss.item()

def test(mask):
    model.eval()
    with torch.no_grad():
        out = model(data)
        pred = out[mask].argmax(dim=1)
        acc = (pred == data.y[mask]).sum().item() / mask.sum().item()
    return acc

# 实际训练
best_val_acc = 0
for epoch in range(200):
    loss = train()
    val_acc = test(data.val_mask)
    test_acc = test(data.test_mask)

    if val_acc > best_val_acc:
        best_val_acc = val_acc
        torch.save(model.state_dict(), 'best_gcn.pth')

    if epoch % 20 == 0:
        print(f'Epoch: {epoch:03d}, Loss: {loss:.4f}, '
              f'Val Acc: {val_acc:.4f}, Test Acc: {test_acc:.4f}')

print(f'Best Val Acc: {best_val_acc:.4f}')

这个流程看起来跟普通深度学习差不多,但有几个关键区别:

  1. 半监督学习:Cora 数据集只标注了 140 个节点(总共 2708 个),模型通过图结构推断未标注节点的标签
  2. 图结构固定:训练过程中 edge_index 不变,只更新节点特征表示
  3. 梯度传播路径:梯度通过边反向传播,邻居节点互相影响

训练调优踩坑

在实际项目里,我遇到过几个典型的训练问题:

坑1:过度平滑(Over-smoothing)

堆叠太多 GCN 层,所有节点的表示趋于一致,分类效果反而变差。

# 解决方案:不要堆太多层,通常 2-3 层就够
# 或者用残差连接
class ResidualGCN(nn.Module):
    def __init__(self, num_features, hidden_dim, num_classes):
        super(ResidualGCN, self).__init__()
        self.conv1 = GCNConv(num_features, hidden_dim)
        self.conv2 = GCNConv(hidden_dim, hidden_dim)
        self.conv3 = GCNConv(hidden_dim, num_classes)

    def forward(self, data):
        x, edge_index = data.x, data.edge_index

        x1 = F.relu(self.conv1(x, edge_index))
        x2 = F.relu(self.conv2(x1, edge_index)) + x1  # 残差连接
        x3 = self.conv3(x2, edge_index)

        return F.log_softmax(x3, dim=1)

坑2:标签泄露(Label Leakage)

在数据划分时,不小心把强相关的节点放到了训练集和测试集,导致测试集性能虚高。

# 解决方案:确保训练/测试节点在图结构上尽量隔离
from torch_geometric.utils import index_to_mask

# 按比例随机划分,但考虑图的连通性
def split_by_graph_component(data, train_ratio=0.1):
    # 先划分图的连通分量,再在每个分量内划分
    # 避免同一连通分量的节点同时出现在训练和测试集
    pass

坑3:图不平衡

某些类别节点特别少,模型容易偏向多数类。

# 解决方案:加权损失函数
from sklearn.utils.class_weight import compute_class_weight

class_weights = compute_class_weight('balanced',
                                     classes=np.unique(data.y.numpy()),
                                     y=data.y.numpy())
class_weights = torch.tensor(class_weights, dtype=torch.float).to(device)

loss = F.nll_loss(out[data.train_mask], data.y[data.train_mask],
                  weight=class_weights)

图级任务实践:图分类

节点分类只是图神经网络的一种应用,实际项目里还会遇到图分类任务(比如化学分子性质预测)。

from torch_geometric.nn import global_mean_pool
from torch_geometric.datasets import TUDataset
from torch_geometric.loader import DataLoader

# 加载 MUTAG 数据集(分子分类)
dataset = TUDataset(root='/tmp/MUTAG', name='MUTAG')
loader = DataLoader(dataset, batch_size=32, shuffle=True)

class GraphGCN(nn.Module):
    def __init__(self, num_features, hidden_dim, num_classes):
        super(GraphGCN, self).__init__()
        self.conv1 = GCNConv(num_features, hidden_dim)
        self.conv2 = GCNConv(hidden_dim, hidden_dim)
        self.lin = nn.Linear(hidden_dim, num_classes)

    def forward(self, data):
        x, edge_index, batch = data.x, data.edge_index, data.batch

        x = self.conv1(x, edge_index)
        x = F.relu(x)
        x = self.conv2(x, edge_index)
        x = F.relu(x)

        # 图级别池化:把节点特征聚合成图特征
        x = global_mean_pool(x, batch)

        x = self.lin(x)
        return F.log_softmax(x, dim=1)

这里的关键是 global_mean_pool,它把同一个 batch 里每个图的节点特征聚合起来,形成图级别的表示。

实际项目建议

从理论走到实践,这几点建议能帮你少走弯路:

1. 不要一开始就上复杂模型

GCN 虽然简单,但很多场景已经够用了。先用 GCN 跑通流程,有需要再试 GAT、GraphSAGE、GIN。

# 简单优先
model = GCN(num_features=data.num_node_features,
            hidden_dim=16,
            num_classes=dataset.num_classes)

2. 数据预处理比模型调参重要

图数据清洗、特征工程、ID 映射、子图采样,这些步骤做好比换几个高级模型提升明显。

3. 小图上本地验证,大图上分布式训练

先用小数据集(Cora、CiteSeer)验证流程和超参数,再上真实的大图。

# 小图快速验证
small_dataset = Planetoid(root='/tmp/Cora', name='Cora')

# 大图需要考虑内存和计算
# 可以用 PyTorch Geometric 的分布式训练

4. 监控图结构的变化

训练过程中图结构通常不变,但要注意预处理是否改变了图的性质(比如去除孤立节点、限制度数)。

结尾

图神经网络不是万能药,但在处理关系型数据时确实比常规深度学习顺手。从社交网络分析到化学分子预测,从推荐系统到知识图谱,它帮我们把那些"节点缠着节点"的问题理清了不少。

这趟从节点到图的实践走下来,最大的感受是:理论看再多不如跑一遍代码,跑一遍代码不如踩几个坑。那些看起来复杂的公式,拆成矩阵乘法、邻接矩阵、节点特征,其实也没那么神秘。

图神经网络还在持续演进,新的模型和框架不断出现,但核心思想没变:让节点看到自己的邻居,让关系结构参与到特征学习里来。这大概是它最朴素也最有效的地方。

版权声明: 本文首发于 指尖魔法屋-AI图神经网络折腾手记https://blog.thinkmoon.cn/post/223-deep-dive-gnn-graph-neural-networks/) 转载或引用必须申明原指尖魔法屋来源及源地址!