Skip to content
0

图神经网络学习知识总结 ​

一、核心概念 ​

1.1 异构图 (Heterogeneous Graph) ​

普通的同构图只包含一种节点类型和一种边类型。异构图则可以包含多种类型的节点和多条类型的边,能更真实地建模现实世界的复杂关系。

在 PyTorch Geometric 中使用 HeteroData 构建异构图:

python
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 是处理异构图的一个经典模型,它包含两层注意力机制:

  1. 节点级注意力 (Node-level Attention):对于每条元路径,学习同一个元路径下邻居节点对目标节点的重要性
  2. 语义级注意力 (Semantic-level Attention):学习不同元路径之间的重要性权重

2.2 PyTorch Geometric 中的使用 ​

python
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 和残差连接来提升训练稳定性:

python
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),用于下游任务:

python
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),保留节点自身信息
python
# 自环边
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)保证训练/验证/测试集中各类别的比例与原始数据一致:

python
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 中通过布尔掩码标记哪些节点属于哪个集合:

python
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) ​

最简单的处理方式:给少数类更大的权重。

python
class_weight = torch.tensor([
    1.0,
    count_neg / count_pos   # 少数类权重 = 多数类数量 / 少数类数量
])
criterion = nn.CrossEntropyLoss(weight=class_weight)

4.3 Focal Loss ​

比加权交叉熵更进一步:降低已分类正确样本的 loss,让模型更关注难分类样本。

python
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) ​

在验证集指标不再提升时停止训练,防止过拟合:

python
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) ​

防止梯度爆炸,尤其在异构图消息传递中很重要:

python
loss.backward()
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
optimizer.step()

5.4 固定随机种子 ​

保证实验可复现:

python
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 曲线 ​

python
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  # 在召回率达标的前提下优化 F1

6.2 阈值选择的原则 ​

  • 降低阈值 → 召回率上升、假阳性增多(宁可错杀)
  • 提高阈值 → 精确率上升、假阴性增多(宁可漏过)
  • 根据对召回率和精确率的相对重视程度选择

七、评估指标 ​

指标含义使用场景
Accuracy整体正确率在不平衡数据中意义有限
Precision预测为正的样本中真正为正的比例关注误报成本
Recall真正的正样本中被预测出的比例关注漏报成本
F1 ScorePrecision 和 Recall 的调和平均综合衡量
ROC-AUC排序能力(对正负样本的区分度)通用,但极端不平衡下可能虚高
PR-AUC精确率-召回率曲线下面积不平衡数据中比 ROC-AUC 更真实
混淆矩阵TP/TN/FP/FN 四格表直观理解模型行为

八、GNN + GBDT 融合 ​

8.1 动机 ​

GNN 学习图的拓扑结构信息(节点间关系),GBDT(如 LightGBM)擅长处理表格特征和非线性关系。两者结合可以互补。

8.2 融合方法 ​

python
# 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 的余弦相似度构建一个全连接加权图(只保留相似度高于阈值的边):

python
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 社区发现 ​

python
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 → 原始特征维度

python
# 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) ​

python
# 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 原理 ​

  1. 用已标注数据训练模型
  2. 对未标注数据做预测
  3. 将高置信度的预测作为"伪标签"
  4. 把伪标签样本加入训练集
  5. 重新训练模型
  6. 重复 2-5 若干轮

11.2 实现要点 ​

python
# 筛选高置信度样本
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 和参数量统计 ​

python
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
传统 MLsklearn (scalers, metrics, split)
梯度提升lightgbm.LGBMClassifier
图分析networkx (图构建), community (Louvain)
模型分析thop (FLOPs), 手动实现的扰动归因
可视化matplotlib, seaborn

十五、GNN 实践中的常见坑与经验 ​

  1. 节点 ID 粒度决定图的信息传播效率——太细会孤立,太粗会过度平滑
  2. 反向边是必须的——GNN 消息传递默认是单向的,不加反向边信息不对称
  3. 自环边能防止节点自身信息在多层卷积中丢失
  4. 类别不平衡不要只用 Accuracy,要关注 Precision/Recall 和 PR-AUC
  5. 阈值是超参数,不是固定的 0.5,需要在验证集上调优
  6. 训练和推理的特征处理必须一致——Scaler 和 Encoder 要保存和复用
  7. 早停保存 state_dict().copy()——这是最容易被忽视的 bug
  8. 元路径定义了模型能学到什么关系——好的图结构设计比模型调参更重要
  9. GNN embedding 是很好的特征——不一定要用 GNN 做端到端分类,提取 embedding + LightGBM 往往效果更好
  10. 自监督预训练在小标注数据集上提升明显,但在标注充足时提升有限
最近更新