图神经网络学习知识总结
一、核心概念
1.1 异构图 (Heterogeneous Graph)
普通的同构图只包含一种节点类型和一种边类型。异构图则可以包含多种类型的节点和多条类型的边,能更真实地建模现实世界的复杂关系。
在 PyTorch Geometric 中使用 HeteroData 构建异构图:
from torch_geometric.data import HeteroData
g = HeteroData()
g['node_type_A'].x = features_A # 节点的特征矩阵
g['node_type_B'].x = features_B
g['A', 'edge_type', 'B'].edge_index = edge_AB # 从A到B的边核心要素:
node_types:所有节点类型的列表edge_types:所有边类型的列表(三元组格式:(src, relation, dst))metadata:由[node_types, edge_types]组成,是异构 GNN 模型的必需参数
1.2 元路径 (Meta-path)
在异构图中,不同类型的节点通过不同的关系连接。元路径定义了一种跨节点类型的路径模式,比如:
- 投保人 → 保单 → 车辆(表示一个人拥有的车辆信息)
- 机构 → 保单 → 报价单(表示一个机构经手的报价)
HAN 模型的核心就是自动学习不同元路径的重要性权重,而不是手工指定哪些路径有用。
二、HAN 模型 (Heterogeneous Graph Attention Network)
2.1 原理
HAN 是处理异构图的一个经典模型,它包含两层注意力机制:
- 节点级注意力 (Node-level Attention):对于每条元路径,学习同一个元路径下邻居节点对目标节点的重要性
- 语义级注意力 (Semantic-level Attention):学习不同元路径之间的重要性权重
2.2 PyTorch Geometric 中的使用
from torch_geometric.nn import HANConv, Linear
class HANModel(torch.nn.Module):
def __init__(self, in_channels, hidden_channels, out_channels, heads, dropout, metadata):
super().__init__()
# in_channels: dict, 如 {'A': 10, 'B': 20}
self.conv = HANConv(in_channels, hidden_channels, heads=heads,
dropout=dropout, metadata=metadata)
self.classifier = Linear(hidden_channels, out_channels)
def forward(self, x_dict, edge_index_dict):
x = self.conv(x_dict, edge_index_dict) # 输出仍是 dict
return self.classifier(x['target_node']) # 取目标节点的表示做分类关键参数说明:
in_channels:一个字典{node_type: feature_dim},分别指定每种节点的输入特征维度hidden_channels:隐藏层维度,所有节点类型共用heads:注意力头数(多头注意力机制),头的数量越大,模型捕捉不同关系的能力越强metadata:[list(node_types), list(edge_types)],描述图的全部结构dropout:在注意力权重上应用 dropout,防止过拟合
2.3 多层 HAN + 残差连接 + BatchNorm
可以将多个 HANConv 堆叠,并加入 BatchNorm 和残差连接来提升训练稳定性:
class EnhancedHAN(torch.nn.Module):
def __init__(self, in_channels, hidden_channels, ...):
self.conv1 = HANConv(in_channels, hidden_channels, heads=heads, metadata=metadata)
self.conv2 = HANConv(hidden_channels, hidden_channels, heads=heads, metadata=metadata)
# 投影层:让输入和输出的维度对齐,用于残差连接
self.proj = ModuleDict({k: Linear(v, hidden_channels) for k, v in in_channels.items()})
# 为每种节点的 hidden 表示单独一个 BatchNorm
self.bn1 = ModuleDict({k: BatchNorm1d(hidden_channels) for k in in_channels.keys()})
def forward(self, x_dict, edge_index_dict):
x_proj = {k: self.proj[k](v) for k, v in x_dict.items()}
x1 = self.conv1(x_dict, edge_index_dict)
x1 = {k: F.relu(self.bn1[k](v)) for k, v in x1.items()}
x1 = {k: self.dropout(v) for k, v in x1.items()}
x2 = self.conv2(x1, edge_index_dict)
x2 = {k: self.bn2[k](v) + x_proj[k] for k, v in x2.items()} # 残差连接
x2 = {k: F.relu(v) for k, v in x2.items()}
return self.lin(x2['target_node'])设计要点:
- BatchNorm 按节点类型分开:因为不同类型的节点特征分布不同,不能用同一个 BN
- 残差连接:x_proj 先把原始特征投影到 hidden_channels 维度,再与第二层输出相加
- Dropout:放在激活函数之后
2.4 GNN Embedding 提取
GNN 的中间层输出(通常是 x['target_node'])可以作为该节点的低维稠密表示(embedding),用于下游任务:
def get_embeddings(model, g):
"""提取目标节点的 embedding 向量"""
model.eval()
with torch.no_grad():
x = model.conv1(g.x_dict, g.edge_index_dict)
x = {k: F.relu(v) for k, v in x.items()}
return x['target_node'].cpu().numpy()这个 embedding 融合了图结构信息和节点自身特征,可以作为特征输入到其他模型中(如 LightGBM)。
三、图构建流程
3.1 从表格数据到异构图
将结构化表格(CSV/DataFrame)转换为异构图的关键步骤:
Step 1 — 确定节点类型 从业务表中识别出不同实体,每行数据关联多个实体。
Step 2 — 生成节点 ID 用实体的关键属性字段组合生成唯一 ID。节点 ID 的粒度很重要:
- 太细(如每行一个节点)→ 图失去共享信息的优势
- 太粗(如所有行合并成一个节点)→ 节点数太少,图退化
Step 3 — 定义节点特征 为每种节点类型选择相关的特征列。特征按类型分类:
- 数值型特征:连续值,用
RobustScaler或StandardScaler标准化 - 类别型特征:离散值,用
LabelEncoder编码后再处理
Step 4 — 构建边 (Edge) 边表示节点之间的关系。从表格的关联字段中抽取出边,一般需要去掉重复的边(drop_duplicates)。
Step 5 — 添加反向边和自环边
- 反向边:对于每条
(A, rel, B)边,添加(B, rel_by, A)反向边,让信息可以双向流通 - 自环边:每个节点添加一条指向自己的边
(node, self, node),保留节点自身信息
# 自环边
g['p', 'self', 'p'].edge_index = torch.stack([
torch.arange(num_p), torch.arange(num_p)
])3.2 特征工程
- 缺失值处理:数值型用中位数填充并添加
_missing标记列;类别型填充"MISSING" - 标准化:
StandardScaler(标准正态分布)或RobustScaler(用中位数和四分位数,抗异常值) - 编码:
LabelEncoder将类别型文本转为整数 - 训练/推理一致性:训练时 fit scaler/encoder 并保存,推理时用保存的 transformer 做 transform
3.3 数据集划分
使用分层采样(stratify)保证训练/验证/测试集中各类别的比例与原始数据一致:
from sklearn.model_selection import train_test_split
train_idx, temp_idx = train_test_split(
indices, test_size=0.3, random_state=42, stratify=labels
)
val_idx, test_idx = train_test_split(
temp_idx, test_size=0.5, random_state=42, stratify=labels[temp_idx]
)在 PyTorch Geometric 中通过布尔掩码标记哪些节点属于哪个集合:
g['p'].train_mask = torch.zeros(num_p, dtype=torch.bool)
g['p'].train_mask[train_idx] = True四、损失函数
4.1 类别不平衡问题
当正负样本比例严重失衡(如 1:10 甚至更低),标准交叉熵会让模型偏向预测多数类。
4.2 类别权重 (Class Weight)
最简单的处理方式:给少数类更大的权重。
class_weight = torch.tensor([
1.0,
count_neg / count_pos # 少数类权重 = 多数类数量 / 少数类数量
])
criterion = nn.CrossEntropyLoss(weight=class_weight)4.3 Focal Loss
比加权交叉熵更进一步:降低已分类正确样本的 loss,让模型更关注难分类样本。
class FocalLoss(nn.Module):
def __init__(self, alpha=0.6, gamma=1.5, weight=None):
super().__init__()
self.alpha = alpha # 平衡正负样本
self.gamma = gamma # 聚焦参数,越大越关注难样本
def forward(self, inputs, targets):
# 数值稳定版:用 log_softmax
log_probs = F.log_softmax(inputs, dim=1)
ce_loss = F.nll_loss(log_probs, targets, weight=self.weight, reduction='none')
pt = torch.exp(-ce_loss) # pt = 预测正确的概率
focal_loss = self.alpha * (1 - pt) ** self.gamma * ce_loss
return focal_loss.mean()参数调节策略:
alpha越大,少数类的损失权重越大 → 提升召回率gamma越大,简单样本的 loss 被压制得越厉害 → 模型更关注难样本- 降低
alpha+ 降低gamma→ 降低假阳性(让模型不要过度倾向于预测少数类)
五、训练技巧
5.1 早停 (Early Stopping)
在验证集指标不再提升时停止训练,防止过拟合:
if val_metric > best_val_metric:
best_val_metric = val_metric
best_state = model.state_dict().copy() # 保存最佳模型参数
patience = 0
else:
patience += 1
if patience >= early_stop_patience:
break关键:保存 state_dict().copy(),因为 state_dict 是浅拷贝,不 copy 的话后续训练会覆盖。
5.2 学习率调度 (LR Scheduler)
- ReduceLROnPlateau:当验证指标不再提升时降低学习率(factor=0.5,每次减半)
- CosineAnnealingLR:余弦退火,适合自监督预训练(平滑衰减)
5.3 梯度裁剪 (Gradient Clipping)
防止梯度爆炸,尤其在异构图消息传递中很重要:
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()5.4 固定随机种子
保证实验可复现:
def set_seed(seed=42):
np.random.seed(seed)
torch.manual_seed(seed)
if torch.cuda.is_available():
torch.cuda.manual_seed_all(seed)六、阈值调优
分类模型默认用 0.5 作为阈值,但在不平衡分类中需要根据业务需求调整。
6.1 基于 Precision-Recall 曲线
from sklearn.metrics import precision_recall_curve
precision, recall, thresholds = precision_recall_curve(y_true, y_proba)
# 在满足最低召回率的前提下,选使 F1 或综合评分最高的阈值
for threshold in np.arange(0.05, 0.95, step):
pred = (proba >= threshold).astype(int)
# 计算 recall, precision, f1, fp, fn...
if recall >= min_recall:
score = f1 - fp * 0.001 # 在召回率达标的前提下优化 F16.2 阈值选择的原则
- 降低阈值 → 召回率上升、假阳性增多(宁可错杀)
- 提高阈值 → 精确率上升、假阴性增多(宁可漏过)
- 根据对召回率和精确率的相对重视程度选择
七、评估指标
| 指标 | 含义 | 使用场景 |
|---|---|---|
| Accuracy | 整体正确率 | 在不平衡数据中意义有限 |
| Precision | 预测为正的样本中真正为正的比例 | 关注误报成本 |
| Recall | 真正的正样本中被预测出的比例 | 关注漏报成本 |
| F1 Score | Precision 和 Recall 的调和平均 | 综合衡量 |
| ROC-AUC | 排序能力(对正负样本的区分度) | 通用,但极端不平衡下可能虚高 |
| PR-AUC | 精确率-召回率曲线下面积 | 不平衡数据中比 ROC-AUC 更真实 |
| 混淆矩阵 | TP/TN/FP/FN 四格表 | 直观理解模型行为 |
八、GNN + GBDT 融合
8.1 动机
GNN 学习图的拓扑结构信息(节点间关系),GBDT(如 LightGBM)擅长处理表格特征和非线性关系。两者结合可以互补。
8.2 融合方法
# 1. 用 GNN 提取目标节点的 embedding
emb = get_gnn_embeddings(model, g_train)
# 2. 拼接 embedding + 原始特征
X = np.concatenate([emb, raw_features], axis=1)
# 3. 用 LightGBM 在拼接特征上训练
lgb_model = LGBMClassifier(
n_estimators=500, learning_rate=0.05, max_depth=8,
class_weight='balanced', subsample=0.8, colsample_bytree=0.8,
reg_alpha=0.1, reg_lambda=0.1
)
lgb_model.fit(X_train, y_train)关键优势:GNN embedding 捕捉了图中邻居传递过来的信息,LightGBM 在此基础上做精细化决策,通常比单独使用 GNN 的分类头效果更好。
九、社区检测 (Community Detection)
9.1 相似度图构建
基于 GNN embedding 的余弦相似度构建一个全连接加权图(只保留相似度高于阈值的边):
from sklearn.metrics.pairwise import cosine_similarity
sim_matrix = cosine_similarity(embeddings)
# 只保留 sim > threshold 的边
G = nx.Graph()
idx = np.argwhere(sim_matrix > threshold)
edges = [(int(i), int(j)) for i, j in idx if i < j]
G.add_edges_from(edges)9.2 Louvain 社区发现
import community as community_louvain
partition = community_louvain.best_partition(G)
# partition: {node_id: community_id}9.3 社区特征
对每个社区计算统计量,赋予给社区内的每个节点:
comm_size:社区大小fraud_ratio:社区内正样本比例fraud_count:社区内正样本数量
这些特征反映了"团伙行为"——如果一个节点所在的社区里正样本比例高,该节点也更可能是正样本。
十、自监督预训练 (Self-Supervised Pre-training)
10.1 特征 Mask + 重建
借鉴 NLP 中的 MLM (Masked Language Modeling),对图节点特征做 mask,训练模型重建被 mask 的特征:
编码器 (Encoder):HAN,将 masked 特征 + 图结构 → embedding
解码器 (Decoder):轻量 MLP,将 embedding → 原始特征维度
# Mask 策略
mask = torch.rand(n_nodes, n_feats) < mask_rate # 35% 的特征被 mask
# 被 mask 的位置:
# - 大部分置 0
# - 小部分 (replace_rate=10%) 用随机值替换
# 未被 mask 的特征加轻微高斯噪声 (noise_std=0.05)
# 只计算被 mask 位置的 MSE Loss
loss = F.mse_loss(reconstructed[mask], original[mask])10.2 预训练 → 微调 (Fine-tuning)
# 1. 预训练 encoder
encoder = pretrain(encoder, decoder, g, epochs=150)
# 2. 将预训练权重加载到下游模型
downstream_model.load_pretrained_encoder(encoder)
# 3. 用少量标注数据微调
downstream_model = train(downstream_model, g)优势:预训练让编码器学到了图的结构先验和特征分布,在标注数据有限时能显著提升性能。
十一、半监督自训练 (Self-Training / Pseudo-Labeling)
11.1 原理
- 用已标注数据训练模型
- 对未标注数据做预测
- 将高置信度的预测作为"伪标签"
- 把伪标签样本加入训练集
- 重新训练模型
- 重复 2-5 若干轮
11.2 实现要点
# 筛选高置信度样本
high_conf_mask = (proba >= 0.85) | (proba <= 0.15)
pseudo_labels = (proba >= 0.5).astype(int)
# 合并数据集
df_augmented = pd.concat([df_train, df_pseudo], ignore_index=True)
# 重新训练
model = train(model, g_augmented)注意:每一轮训练需要重新构建图(因为训练数据变了),且分类阈值和置信度阈值需要根据验证集谨慎选择,否则伪标签中的噪声会累积。
十二、模型可解释性
12.1 扰动法 (Perturbation-based Attribution)
通过逐一"移除"节点/边/特征并观察预测变化来归因风险来源:
节点级归因:
- 将某个节点的特征全部重置为 baseline(如 0)
- 观察目标节点预测概率的下降幅度
- 下降越大 → 该节点对风险的贡献越大
特征级归因:
- 将某个特征的取值重置为 baseline
- 观察预测变化
- 选出下降幅度最大的 Top-K 个特征
边级归因:
- 删除某条边(切断信息传导)
- 观察预测变化
- 如果删除关联边后风险骤降 → 风险主要来自图拓扑关系(关联网络),而非自身特征
12.2 核心思想
这种方法的本质是 Ablation Study:用干预手段"关掉"模型的某部分输入,看输出变化,从而推断各部分的重要性。
十三、可视化
13.1 训练过程可视化
- Loss 曲线:训练和验证 loss 随 epoch 的变化,判断是否过拟合
- Accuracy 曲线:训练和验证准确率的变化
13.2 模型评估可视化
- 混淆矩阵热力图:直观展示 TP/TN/FP/FN
- ROC 曲线:展示不同阈值下 TPR vs FPR 的 trade-off
- PR 曲线:在不平衡数据中比 ROC 更有参考价值
- 预测概率分布直方图:观察模型对正负样本的区分程度
- 阈值-指标曲线:不同阈值下的 Recall/Precision/F1/FP 变化
13.3 FLOPs 和参数量统计
from thop import profile
flops, params = profile(model, inputs=(g.x_dict, g.edge_index_dict), verbose=False)
print(f"GFLOPs: {flops / 1e9:.4f}, Params: {params:,}")十四、完整技术栈
| 组件 | 工具 |
|---|---|
| 图数据结构 | torch_geometric.data.HeteroData |
| GNN 模型 | torch_geometric.nn.HANConv |
| 深度学习框架 | torch, torch.nn |
| 传统 ML | sklearn (scalers, metrics, split) |
| 梯度提升 | lightgbm.LGBMClassifier |
| 图分析 | networkx (图构建), community (Louvain) |
| 模型分析 | thop (FLOPs), 手动实现的扰动归因 |
| 可视化 | matplotlib, seaborn |
十五、GNN 实践中的常见坑与经验
- 节点 ID 粒度决定图的信息传播效率——太细会孤立,太粗会过度平滑
- 反向边是必须的——GNN 消息传递默认是单向的,不加反向边信息不对称
- 自环边能防止节点自身信息在多层卷积中丢失
- 类别不平衡不要只用 Accuracy,要关注 Precision/Recall 和 PR-AUC
- 阈值是超参数,不是固定的 0.5,需要在验证集上调优
- 训练和推理的特征处理必须一致——Scaler 和 Encoder 要保存和复用
- 早停保存
state_dict().copy()——这是最容易被忽视的 bug - 元路径定义了模型能学到什么关系——好的图结构设计比模型调参更重要
- GNN embedding 是很好的特征——不一定要用 GNN 做端到端分类,提取 embedding + LightGBM 往往效果更好
- 自监督预训练在小标注数据集上提升明显,但在标注充足时提升有限