- 人工智能
- 机器学习
- 深度学习
- 图计算
【免费下载链接】dgl
Python package built to ease deep learning on graph, on top of existing DL frameworks.
本文基于 DGL 官方仓库中的 caregnn 示例,完整讲解 CAmouflage-REsistant GNN(CARE-GNN)这一面向欺诈检测的图神经网络模型的原理、数据集、训练脚本与性能复现方法。读完本文,你将掌握如何在 DGL 中加载 FraudDataset 内置欺诈数据集(Amazon / YelpChi),运行全图版与采样版两套训练管线,并理解其基于强化学习的相似度门控机制在源码中的具体实现。
一、模型背景:伪装欺诈检测为什么需要 CARE-GNN
CARE-GNN(CAmouflage-REsistant GNN)由论文 Enhancing Graph Neural Network-based Fraud Detectors against Camouflaged Fraudsters(CIKM'20)提出。在电商评论与金融反欺诈场景中,欺诈者(fraudster)会通过伪装行为隐藏自身与真实受害者之间的紧密关系,例如伪造大量同评分、同时段、同商品的关联,导致普通 GNN 的邻居聚合被"噪声邻居"污染。CARE-GNN 的核心思想是:
- 对每个关系类型(relation/edge type)上的邻居进行相似度门控采样,只聚合与中心节点最相似的 Top-p% 邻居;
- 引入一个强化学习(RL)模块,在训练过程中动态调整每个关系类型的采样比例 p,使得模型能够自适应地抑制被伪装的关系,同时保留信息量高的关系。
DGL 示例由 Kay Liu 在 AWS 上海 AI Lab 的 SDE 实习期间实现,示例代码位于 examples/pytorch/caregnn,其中包含两套可独立运行的训练入口:main.py(全图训练)与 main_sampling.py(邻居采样训练)。
二、依赖环境
README 中给出的示例开发环境如下(以当前仓库版本运行时代码为准,示例本身面向 DGL 0.7.x 设计):
- Python 3.7.10
- PyTorch 1.8.1
- dgl 0.7.1
- scikit-learn 0.23.2
其中 scikit-learn 用于计算recall_score与roc_auc_score两项评价指标(见 main.py),PyTorch 承担模型构建与优化,DGL 提供异构图数据结构与消息传递原语。
三、数据集:DGL 内置 FraudDataset
两个数据集均为 DGL 内置的 FraudDataset(位于dgl.data.FraudDataset),是从真实工业数据构建的多关系图(multi-relational graph),每个图只有一个节点类型,包含三种关系类型,具有类别不均衡与特征不一致等真实噪声特性。图的构建逻辑在 fraud.py:读取.mat原始文件中的三个邻接矩阵,分别以(node_type, relation, node_type)的三元组构造dgl.heterograph,并注入feature、label与train_mask/val_mask/test_mask。
Amazon(虚假用户检测)
- 节点:11,944(user)
- 边:
- U-P-U:351,216(共同购买过同一商品)
- U-S-U:7,132,958(一周内给出相同星级评分)
- U-V-U:2,073,474(评论文本 TF-IDF 相似度位居前 5%)
- 类别:Positive(欺诈)821;Negative(良性)7,818;Unlabeled 3,305
- 正负比:1 : 10.5
- 节点特征维度:25
YelpChi(虚假评论检测)
- 节点:45,954(review)
- 边:
- R-U-R:98,630(同一用户发布的评论)
- R-T-R:1,147,232(同一商品同月发布的评论)
- R-S-R:6,805,486(同一商品相同星级评分的评论)
- 类别:Positive(垃圾评论)6,677;Negative(正常评论)39,277
- 正负比:1 : 5.9
- 节点特征维度:32
值得注意的是FraudDataset的默认划分参数为 train_size=0.7、val_size=0.1、random_seed=717,而本示例在 main.py 中通过dgl.data.FraudDataset(args.dataset, train_size=0.4)将训练集比例调整为 0.4(测试集约为 0.5)。数据集的划分实现见 fraud.py,其中 Amazon 数据集中索引 0~3304 的未标注节点会被排除在划分之外。数据集会按(random_seed, train_size, val_size)生成哈希键并缓存到本地,见 fraud.py。
四、全图训练版:main.py
4.1 训练流程概览
main.py 的main()分为四步:
- 数据准备:加载
FraudDataset,取出图、特征、标签与三类 mask;从 mask 中解析出 train/val/test 节点索引;此外还专门提取正类训练节点索引rl_idx,用于驱动强化学习模块(main.py)。 - 模型构建:以特征维度、类别数、隐藏维度、层数、激活函数(
tanh)、RL 步长与图的全规范边类型(graph.canonical_etypes)构造CAREGNN。 - 训练组件:由于类别不均衡,损失函数使用按类别计数加权的 CrossEntropyLoss——
th.nn.CrossEntropyLoss(weight=1 / cnt);优化器为 Adam(lr=0.01,weight_decay=0.001);开启--early-stop时创建 patience=100 的 EarlyStopping。 - 训练循环:每轮前向得到
logits_gnn, logits_sim两个输出,总损失为CE(logits_gnn) + sim_weight * CE(logits_sim);随后调用model.RLModule(graph, epoch, rl_idx)更新各关系的采样比例 p;若启用 early stopping,则在验证 AUC 连续 100 轮不提升时终止训练,并将最优权重保存到es_checkpoint.pt。
训练、验证与测试均基于 sklearn 的recall_score(正类召回)与roc_auc_score(AUC,使用 softmax 后的正类概率),输出格式如:
Epoch 0, Train: Recall: 0.xxxx AUC: 0.xxxx Loss: x.xxxx | Val: Recall: 0.xxxx AUC: 0.xxxx Loss: x.xxxx4.2 全部命令行参数
| 参数 | 类型 | 默认值 | 说明 |
|---|---|---|---|
--dataset | str | amazon | 数据集名称,可选yelp或amazon |
--gpu | int | -1 | GPU 索引,-1 表示使用 CPU |
--hid_dim | int | 64 | 隐藏层维度 |
--num_layers | int | 1 | CARE-GNN 层数 |
--max_epoch | int | 30 | 最大训练轮数 |
--lr | float | 0.01 | 学习率 |
--weight_decay | float | 0.001 | 权重衰减 |
--step_size | float | 0.02 | RL 动作步长(论文公式中的 λ2) |
--sim_weight | float | 2 | 相似度损失权重(论文公式中的 λ1) |
--early-stop | flag | False | 是否启用早停 |
4.3 运行方式
在examples/pytorch/caregnn目录下执行:
# 全图训练 + 早停(默认 Amazon) python main.py --early-stop # 使用 GPU python main.py --gpu 0 # 切换为 Yelp 数据集 python main.py --dataset yelp # 组合使用 python main.py --dataset yelp --gpu 0 --early-stop --num_layers 2 --hid_dim 128程序入口在启动训练前会调用th.manual_seed(717)(main.py),保证结果可复现。
五、模型结构源码解析:model.py
5.1 CAREConv:单层卷积
CAREConv 是 CARE-GNN 的核心单层实现,在__init__中为每个边类型维护四个状态:
p[etype] = 0.5:当前采样比例(初始为 0.5,即保留 Top 50% 相似邻居);last_avg_dist[etype] = 0:上一轮该关系的平均相似度距离;f[etype] = []:RL 动作历史(+1/-1 序列);cvg[etype] = False:该关系的 RL 是否已收敛。
单层前向传播对应论文公式 8、9(model.py):
- 对每个规范边类型
etype,通过g.apply_edges(self._calc_distance, etype=etype)计算邻居距离。_calc_distance对应论文公式 2:d = || tanh(MLP(h_src)) - tanh(MLP(h_dst)) ||_1,即用 MLP 嵌入后的 L1 距离度量中心节点与邻居的相似度(距离越小越相似)。 - 调用
_top_p_sampling按比例 p 保留距离最小的前ceil(in_degree * p)条入边。当前实现基于np.argpartition完成部分排序,代码注释说明其效率较低,可优化方向是 DGL 的dgl.sampling.select_top_p。 - 对采样后的边执行
g.send_and_recv,用fn.copy_u("h","m")+fn.mean("m","h_etype")完成按边类型聚合。 - 关系间聚合:将各关系的聚合结果按 p 加权求和(
h_homo = Σ hr * p),再加上中心节点自身特征feat,经激活后通过self.linear线性投影输出。这段实现了论文公式 9 的均值型 inter-relation aggregator(此处权重即 RL 输出的 p)。
5.2 CAREGNN:多层堆叠与 RL 模块
CAREGNN 负责按num_layers堆叠CAREConv:单层时直接输出类别数维度;多层时按"输入层(hid_dim)→中间层(hid_dim)→输出层(num_classes)"的结构组织。前向传播返回两个输出:
feat:GNN 主分支的 logits;sim:tanh(self.layers[0].MLP(feat)),即公式 4 定义的相似度学习分支输出,用于辅助损失。
RLModule(model.py)实现了论文公式 5~7 的强化学习更新:对每个尚未收敛的 (layer, etype) 组合,取正类训练节点rl_idx的入边距离均值avg_dist(公式 5):
- 若
last_avg_dist < avg_dist(距离变大,说明该关系噪声增多),则 p 减小step_size(下限为 0,公式 6),并记录动作 -1; - 否则 p 增大
step_size(上限为 1),并记录动作 +1; - 当
epoch >= 9且最近 10 个动作之和的绝对值<= 2时,判定该关系 RL 已收敛,此后不再调整(公式 7)。
六、采样训练版:main_sampling.py
6.1 与全图版的差异
README 中特别注明:采样版本根据 DGL NodeDataLoader 的特性做了修改——论文公式 2 中原本使用"最后一层嵌入"计算相似度,本采样版改用"当前层在上一轮 epoch 得到的嵌入"来度量中心节点与邻居的相似度。这一改动是为了适应 mini-batch 采样场景下无法直接获得全图多层嵌入的现实约束。
6.2 CARESampler 采样器
model_sampling.py 定义了CARESampler,继承自dgl.dataloading.BlockSampler:
sample_frontier中,对每个边类型基于上轮缓存的距离矩阵dists与当前 p 值,为每个种子节点挑选ceil(in_degree * p)条距离最小的入边,生成边掩码后通过dgl.edge_subgraph构造 frontier;sample_blocks自底向上(reversed(range(num_layers)))逐层采样,并用dgl.to_block将 frontier 转为 message-passing block,同时保留原始边 ID(dgl.EID),供 RL 模块查询。
6.3 训练流程
main_sampling.py 每轮 epoch 开始时(main_sampling.py):
- 对每个层 i、每个边类型计算距离
dist[etype] = L1(tanh(MLP(feat_i))),缓存在dists中; - 用当前 p 构造
CARESampler; - 通过
dgl.dataloading.DataLoader以batch_size=256、shuffle=True构建训练/验证/测试 mini-batch,逐 batch 前向、计算加权 CE 损失并反向更新; - 每轮结束后调用
model.RLModule(graph, epoch, rl_idx, dists)更新 p 值。
采样版额外提供--batch_size(默认 256)与--num_workers(默认 4)两个参数;在 GPU 模式下会将num_workers强制设为 0(main_sampling.py)。
6.4 运行方式
# 默认以 Amazon + 全图版同参数运行采样训练 python main_sampling.py # 常用组合 python main_sampling.py --dataset yelp --gpu 0 --early-stop --batch_size 512七、早停机制:utils.py
utils.py 中的EarlyStopping类以验证 AUC 为监控指标:当指标不升反降时计数 +1,连续patience轮(示例中为 100)不提升即触发早停;每当验证指标刷新最优值时,将当前模型权重保存为es_checkpoint.pt,测试阶段直接加载该最优权重(main.py、main_sampling.py)。
八、性能复现结果
README 中报告的结果遵循论文设定:在30 个 epoch 内取最佳验证结果,随机种子统一为seed=717(论文原报告未给出测试集结果,用-表示)。以下为论文、DGL 全图版与 DGL 采样版在 Amazon 与 Yelp 上的 AUC / Recall 对比:
| Dataset | Amazon | Yelp |
|---|---|---|
| Metric (val / test) | Max Epoch 30 | Max Epoch 30 |
| AUC (val/test) — paper reported | 0.8973 / - | 0.7570 / - |
| AUC (val/test) — DGL full graph | 0.8849 / 0.8922 | 0.6856 / 0.6867 |
| AUC (val/test) — DGL sampling | 0.9350 / 0.9331 | 0.7857 / 0.7890 |
| Recall (val/test) — paper reported | 0.8848 / - | 0.7192 / - |
| Recall (val/test) — DGL full graph | 0.8615 / 0.8544 | 0.6667 / 0.6619 |
| Recall (val/test) — DGL sampling | 0.9130 / 0.9045 | 0.7537 / 0.7540 |
从表中可以看到,DGL 采样版在两个数据集上的 val/test AUC 与 Recall 均高于全图版与论文原始报告值。需要说明的是,这些数字是该示例在特定环境与seed=717下的复现结果,实际运行受硬件、DGL/PyTorch 版本与随机种子影响,可能存在合理波动。
九、在仓库中继续深入
- caregnn 示例目录:
main.py/model.py/main_sampling.py/model_sampling.py/utils.py全部源码; - FraudDataset 数据集实现:数据集下载、异构图构造、划分与缓存逻辑;
- DGL 文档指南(中文版)与 异构图消息传递 API:理解
apply_edges、send_and_recv、canonical_etypes等本示例高频使用的接口。
使用提示:运行示例前请确认 DGL 与 PyTorch 环境已正确安装,数据集首次加载时会自动从 DGL 官方数据源下载(约几百 MB 量级),FraudDataset会按划分参数缓存处理后的二进制图文件(见 fraud.py),重复运行无需重新下载;若需调整数据划分,可修改train_size/val_size/random_seed参数以控制缓存键。
- 人工智能
- 机器学习
- 深度学习
- 图计算
【免费下载链接】dgl
Python package built to ease deep learning on graph, on top of existing DL frameworks.
相关推荐
如何在macOS Finder中实现视频文件的完美预览:QLVideo终极指南
如何在macOS Finder中实现视频文件的完美预览:QLVideo终极指南 QLVideo 是一款专为macOS设计的开源视频扩展工具,它让Finder能够
音视频桌面应用终极指南:如何用curl-impersonate突破网站指纹检测,完美伪装Chrome和Firefox浏览器
终极指南:如何用curl impersonate突破网站指纹检测,完美伪装Chrome和Firefox浏览器 curl impersonate是一个特殊的cur
网络安全网络开发工具DGL 中的 APPNP(Personalized PageRank 图神经网络)MXNet 实现与实战指南
DGL 中的 APPNP(Personalized PageRank 图神经网络)MXNet 实现与实战指南 导读 APPNP(Approximate Pers
人工智能机器学习深度学习图计算
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考