图嵌入与图神经网络:从 node2vec 到 GCN 的完整图谱

系统覆盖图嵌入与图神经网络(GNN):图嵌入的动机与三类方法(随机游走/node2vec、矩阵分解、GraphSAGE)、图神经网络核心(消息传递、聚合函数、GCN/GAT)、节点分类/链路预测/图分类任务、训练与评估(划分/度量)、PyTorch Geometric 实战、以及与图数据库的联合应用。

引言

图数据库解决了「关系怎么存、怎么查」,但「关系怎么学」——从图结构里自动提取预测信号——需要图表示学习:把节点、边、整张图编码成低维向量(图嵌入),或用**图神经网络(GNN)**端到端学习。这两者已成为风控、推荐、社交分析、知识图谱推理的标准组件。

本文系统讲图嵌入与 GNN:先从「为什么要嵌入」讲起,覆盖三类图嵌入方法(随机游走/node2vec、矩阵分解、GraphSAGE);接着进入 GNN 核心——消息传递范式、聚合函数、GCN 与 GAT;再讲三大任务(节点分类、链路预测、图分类)与训练评估;最后用 PyTorch Geometric 落地,并与 Neo4j 联合应用。

前置:/graphdb-algorithms-practice/(图算法)、/graphdb-recommendation-system/(图推荐)、/graphdb-fraud-detection/(图特征)、[[ai-ml]](ML 基础)。


目录


1. 为什么要图嵌入

1.1 图数据无法直接喂给 ML

传统 ML 输入是「固定维度向量」(表格行、图像像素)
图数据的难点:节点数量不固定、结构信息(邻居关系)在向量外

图嵌入 = 把节点/边/图映射到低维向量空间
        保留「结构相似 → 向量相近」的性质

1.2 嵌入的三个层次

节点嵌入  :每个节点 → 向量(用于节点分类/相似度)
边嵌入    :每条边 → 向量(节点向量拼接/点积)
整图嵌入  :每张图 → 向量(用于图分类/相似图检索)

1.3 应用场景

场景嵌入用途
风控把账户嵌入喂给风险模型
推荐用户/物品向量做相似度(见推荐篇)
社交用户向量聚类找群体
知识图谱实体/关系嵌入用于推理补全
生物/分子分子图嵌入预测性质

一句话总结:图嵌入把「关系结构」压缩成向量,让传统 ML 能消费图数据——结构相似的实体在向量空间里也相近。


2. 随机游走类:DeepWalk 与 node2vec

2.1 DeepWalk:把图当「句子」

思想:在图上做随机游走,得到节点序列(像句子中的词)
      → 用 Word2Vec 训练,让「同一游走里近邻」的节点向量相近

游走:从 v 出发随机选邻居走 L 步,得到序列 [v, v1, v2, ...]
      → 序列集合 = 语料库 → Skip-gram 训练嵌入

2.2 node2vec:偏置游走,捕捉同质性与结构性

node2vec 的改进是让游走偏向「深度优先(探索社区)」或「广度优先(探索枢纽结构)」:

偏置参数:
  p(返回参数):控制回到上一个节点的概率(小 p → 局部性)
  q(进出参数):控制走向更远邻居的概率(小 q → 广度/结构)
同质性(小 q):邻居结构相似的节点向量相近(同社区)
结构性(大 q):枢纽/桥节点向量相近(同角色)
from node2vec import Node2Vec

edges = [('U1','M1'), ('U1','M2'), ('U2','M2'), ('U2','M3')]  # 边列表
model = Node2Vec(edges, dimensions=64,
                 walk_length=20, num_walks=50,
                 p=1.0, q=0.5,    # q<1 → 偏重结构
                 workers=4).fit()
vec = model.wv['U1']   # 节点 U1 的 64 维向量

2.3 局限性

1. 转导式(transductive):新节点加入要重新训练
2. 只利用结构、不利用节点属性
3. 大规模图的游走内存开销大

一句话总结:随机游走类把图转成「语料」,用 Word2Vec 训练节点向量——node2vec 的 p/q 偏置让嵌入能倾向「社区」或「结构角色」,缺点是转导式、不看属性。


3. 矩阵分解与图拉普拉斯嵌入

3.1 邻接矩阵分解

设邻接矩阵 A(N×N),希望找到 U(N×d)使 A ≈ U·Uᵀ
→ 用 SVD / 谱方法求解,得到节点嵌入 U

局限:N 大时 A 是 N² 稀疏矩阵,SVD 昂贵 → 只适合中等规模

3.2 图拉普拉斯谱嵌入(Laplacian Eigenmaps)

图拉普拉斯 L = D - A(D 为度矩阵)
谱嵌入:取 L 的「最小非零特征值」对应的特征向量
       → 保留「局部相似性」(相连的节点向量相近)

意义:这是 GCN 的数学基础——GCN 的卷积本质是谱域的局部近似

3.3 矩阵分解 vs 随机游走

方法复杂度规模用途
SVD 分解O(N²)~O(N³)小理论清晰
谱嵌入O(N²)中理解拉普拉斯
DeepWalk/node2vec近似线性大工业落地主流

一句话总结:矩阵分解与谱嵌入是图嵌入的「数学原生」方法,复杂度随节点数平方增长——理解它是理解 GCN 谱卷积的钥匙,工业上大规模仍靠随机游走。


4. 消息传递:GNN 的核心范式

4.1 消息传递(Message Passing)

GNN 的核心思想:每个节点迭代地「收集邻居消息、聚合、更新自己」:

第 k 层:
  ① 消息(Message)   :每个邻居向节点发送自己的表示
     m_{u→v} = Message(h_u^(k-1))
  ② 聚合(Aggregate) :节点聚合所有邻居消息
     a_v = AGG({m_{u→v} | u ∈ N(v)})
  ③ 更新(Update)    :结合自身表示生成新表示
     h_v^(k) = Update(h_v^(k-1), a_v)

k 层后,h_v 聚合了 k 跳邻居的信息(感受野 = k 跳)

4.2 聚合函数的三种选择

均值聚合(Mean) :a_v = mean({h_u})        → 平滑、稳定(GCN 用)
求和聚合(Sum)   :a_v = sum({h_u})         → 保度数敏感(GIN 用)
注意力聚合(Attn):a_v = Σ α_uv · h_u       → 邻居重要性加权(GAT 用)

4.3 层数与过平滑(Over-smoothing)

层数太深 → 所有节点表示趋同(过平滑)→ 分类能力退化
经验:2-3 层 GNN 最常用(感受野 2-3 跳)
深图学习需残差、跳跃连接、GPR 等专门设计

一句话总结:GNN 的引擎是「消息传递」——每层让节点聚合邻居消息并更新自身,层数决定感受野;聚合函数的选择(均值/求和/注意力)决定了模型的表达能力。


5. GCN 与 GAT:两种经典聚合

5.1 GCN:谱卷积的局部近似

GCN 的聚合公式(均值 + 自连接 + 对称归一化):
  h_v^(k) = ReLU( W · Σ_{u∈N(v)∪{v}} (1/√(d_u·d_v)) · h_u^(k-1) )

要点:
  1. 归一化用度 d_u·d_v 的平方根 → 抑制高连接节点的消息放大
  2. 加自连接(v 也聚合自己)→ 保留自身信息
  3. 本质上是对图拉普拉斯的 1 阶局部近似
# 示意:GCN 单层聚合(伪代码)
def gcn_layer(h, adj, W, D):
    A_hat = adj + I                          # 加自连接
    D_hat = degree_diag(A_hat)
    A_norm = D_hat**(-0.5) @ A_hat @ D_hat**(-0.5)   # 对称归一化
    return relu(A_norm @ h @ W)

5.2 GAT:注意力聚合

GAT 给每个邻居学一个重要性权重 α_uv:
  α_uv = softmax( LeakyReLU( a·[W·h_u ; W·h_v] ) )
聚合 = Σ α_uv · W · h_u

优势:无需预设度归一化,自适应地放大重要邻居、忽略噪声邻居
多头注意力:并行多组 α → 拼接/平均,更稳定

5.3 选型

模型聚合优点局限
GCN均值+度归一简单、稳健均质化、忽略度差异
GAT注意力自适应、强表达训练略贵
GraphSAGE采样+均值/LSTM可采样大图采样引入偏差
GIN求和最强表达力(WL 等价)度敏感易过拟合

一句话总结:GCN 用「度归一化的均值聚合」把谱卷积拉成局部计算,GAT 用注意力学邻居权重更自适应——GraphSAGE 以采样支持大规模图,是工业落地的常客。


6. 三大任务:节点分类、链路预测、图分类

6.1 节点分类(Node Classification)

目标:预测节点标签(黑产账户/非黑产、欺诈订单/正常订单)
做法:节点嵌入 → 接 MLP/softmax 分类器
监督:一部分节点有标签,其余预测(半监督)
# 示意
node_emb = gnn_forward(graph)        # (N, d)
logits = linear(node_emb)            # (N, num_classes)
loss = cross_entropy(logits[masked], y[masked])   # 只在有标签节点上算
目标:预测「两个节点之间是否会出现/应该存在关系」
做法:正样本(已有边)+ 负采样(随机不存在的边)
     打分函数 score(u,v) = φ(h_u, h_v)   (点积/拼接+MLP)
损失:二分类(有无边)或排序损失
应用:推荐(用户-物品边)、知识图谱补全

6.3 图分类(Graph Classification)

目标:整张图一个标签(分子是否有药性、网络是否异常)
做法:节点嵌入 → 读出(Readout:全局池化 mean/sum)→ 分类器

一句话总结:GNN 三大任务——节点分类给每个节点打分、链路预测判「边会不会出现」、图分类给整图定性,训练都建立在消息传递得到的嵌入之上。


7. 训练与评估

7.1 数据划分(转导式 vs 归纳式)

转导式(transductive):整图训练,随机遮住部分标签测试
  → 测试节点在训练时「看到」了图的连接
归纳式(inductive)   :训练图/测试图分开(GraphSAGE、GNN 通用)
  → 更接近真实「新节点上线」场景

7.2 指标

任务指标
节点分类Accuracy、F1(类别不平衡)、AUC
链路预测AUC、AP、Hit@K
图分类Accuracy、F1、AUC
# 链路预测指标示例
from sklearn.metrics import roc_auc_score
auc = roc_auc_score(y_true, pos_prob)      # y_true: 是否存在边

7.3 常见坑

1. 标签泄漏:测试节点「通过邻居传播」偷看标签 → 用 inductive 划分验证
2. 类别不平衡:欺诈场景正样本极少 → F1/AUC 而非准确率
3. 过平滑:层数别太深
4. 嵌入维度:64-256 常见,过大易过拟合

一句话总结:训练要注意「归纳 vs 转导」的划分是否匹配真实场景,评估用 F1/AUC 应对不平衡——标签泄漏与过平滑是 GNN 最容易翻车的两个坑。


8. PyTorch Geometric 实战

8.1 从边列表构造 Data

import torch
from torch_geometric.data import Data

edge_index = torch.tensor([[0, 1, 1, 2, 2, 3],      # 源节点
                           [1, 0, 2, 1, 3, 2]],     # 目标节点
                          dtype=torch.long)
x = torch.randn(4, 16)      # 4 个节点的 16 维特征
y = torch.tensor([0, 1, 0, 1])   # 节点标签
data = Data(x=x, edge_index=edge_index, y=y)

8.2 定义一个 GCN 模型

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

class GCN(torch.nn.Module):
    def __init__(self, in_dim, hid_dim, out_dim):
        super().__init__()
        self.conv1 = GCNConv(in_dim, hid_dim)
        self.conv2 = GCNConv(hid_dim, out_dim)

    def forward(self, x, edge_index):
        x = F.relu(self.conv1(x, edge_index))
        x = F.dropout(x, training=self.training)
        return self.conv2(x, edge_index)   # 节点分类 logits

8.3 训练循环

model = GCN(in_dim=16, hid_dim=32, out_dim=2)
opt = torch.optim.Adam(model.parameters(), lr=0.01)
loss_fn = torch.nn.CrossEntropyLoss()

for epoch in range(200):
    model.train()
    logits = model(data.x, data.edge_index)
    loss = loss_fn(logits[data.train_mask], data.y[data.train_mask])
    opt.zero_grad(); loss.backward(); opt.step()
    if epoch % 20 == 0:
        print(epoch, loss.item())

8.4 链路预测版本要点

# 负采样:随机取「不存在的边」作为负样本
neg_edge_index = random_non_existing_edges(data)   # (2, num_neg)
# 打分:node_emb[u] · node_emb[v]
score = (emb[src] * emb[dst]).sum(dim=1)
loss = bce(score, labels)     # labels: 正边=1, 负边=0

一句话总结:PyTorch Geometric 让 GNN 落地很直接——构造 edge_index Data、定义 GCN/GAT 层、用掩码损失训练;链路预测加一步负采样与边打分即可。


9. 与图数据库的联合应用

9.1 联合工作流

Neo4j(存与查)→ 导出边列表/子图 → Python(GNN 训练)→ 嵌入
→ 写回 Neo4j 节点属性 → 在线查询用嵌入相似度/向量检索

9.2 Neo4j GDS + 嵌入

// GDS 原生支持 node2vec 等(老版本),新版本可算相似度
CALL gds.nodeSimilarity.stream('graph')
YIELD node1, node2, similarity
// 也可把 Python 算好的嵌入存为属性,配合向量索引做 TopK
# Python 侧:读取子图 → node2vec → 写回
from neo4j import GraphDatabase
driver = GraphDatabase.driver("bolt://localhost:7687")
with driver.session() as s:
    edges = list(s.run("MATCH (u)-[r]->(m) RETURN id(u), id(m)"))
    # ... 训练嵌入 ...
    s.run("UNWIND $rows AS row SET n.emb = row.emb", rows=vec_rows)

9.3 适合图数据库的在线嵌入应用

1. 向量相似度 TopK:把嵌入存向量索引,新查询秒级返回相似节点
2. 图特征:把 GNN 学的嵌入当「图特征」喂给风控/推荐模型
3. 结构检索:找到「结构与某节点最像」的节点(反欺诈同构)

一句话总结:图数据库负责「存与查」,GNN 负责「学与嵌入」,两者通过「导出边→训练→写回」联合——嵌入存回 Neo4j 后,在线查询直接做向量相似度。


10. 速查表

需求方法
节点/图 → 向量node2vec / 谱嵌入 / GraphSAGE
邻居聚合GCN(度归一均值)/ GAT(注意力)
大图GraphSAGE 采样
节点分类嵌入 + MLP
链路预测边打分 + 负采样
图分类全局池化 + 分类器
训练注意归纳划分、F1/AUC、防泄漏
落地库PyTorch Geometric
与 Neo4j 联合导出边 → 训练 → 写回嵌入

一句话记忆:图嵌入把关系压成向量——node2vec 走随机游走、谱嵌入走拉普拉斯;GNN 靠消息传递聚合邻居,GCN 度归一均值、GAT 注意力加权;三大任务(节点分类/链路预测/图分类)都建立在嵌入之上;训练防标签泄漏与过平滑,落地用 PyTorch Geometric,嵌入写回 Neo4j 供在线向量检索——图数据库管「存查」、GNN 管「学习」。

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「graphdb」更多文章

  1. 图驱动推荐系统:从协同过滤到图嵌入的实战路径
  2. 图数据建模模式与反模式:从关系思维到图谱思维
  3. 反欺诈与风控图谱实战:关联分析识别黑产团伙