引言
图数据库解决了「关系怎么存、怎么查」,但「关系怎么学」——从图结构里自动提取预测信号——需要图表示学习:把节点、边、整张图编码成低维向量(图嵌入),或用**图神经网络(GNN)**端到端学习。这两者已成为风控、推荐、社交分析、知识图谱推理的标准组件。
本文系统讲图嵌入与 GNN:先从「为什么要嵌入」讲起,覆盖三类图嵌入方法(随机游走/node2vec、矩阵分解、GraphSAGE);接着进入 GNN 核心——消息传递范式、聚合函数、GCN 与 GAT;再讲三大任务(节点分类、链路预测、图分类)与训练评估;最后用 PyTorch Geometric 落地,并与 Neo4j 联合应用。
前置:/graphdb-algorithms-practice/(图算法)、/graphdb-recommendation-system/(图推荐)、/graphdb-fraud-detection/(图特征)、[[ai-ml]](ML 基础)。
目录
- 1. 为什么要图嵌入
- 2. 随机游走类:DeepWalk 与 node2vec
- 3. 矩阵分解与图拉普拉斯嵌入
- 4. 消息传递:GNN 的核心范式
- 5. GCN 与 GAT:两种经典聚合
- 6. 三大任务:节点分类、链路预测、图分类
- 7. 训练与评估
- 8. PyTorch Geometric 实战
- 9. 与图数据库的联合应用
- 10. 速查表
- 延伸阅读
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]) # 只在有标签节点上算
6.2 链路预测(Link Prediction)
目标:预测「两个节点之间是否会出现/应该存在关系」
做法:正样本(已有边)+ 负采样(随机不存在的边)
打分函数 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 管「学习」。
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。