news 2026/9/23 4:33:08

DGL 实现 CARE-GNN:面向伪装欺诈检测器的抗伪装图神经网络实战指南

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
DGL 实现 CARE-GNN:面向伪装欺诈检测器的抗伪装图神经网络实战指南
  • 人工智能
  • 机器学习
  • 深度学习
  • 图计算

【免费下载链接】dgl

Python package built to ease deep learning on graph, on top of existing DL frameworks.

项目地址:https://gitcode.com/gh_mirrors/dg/dgl
点击查看免费下载

本文基于 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_scoreroc_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,并注入featurelabeltrain_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()分为四步:

  1. 数据准备:加载FraudDataset,取出图、特征、标签与三类 mask;从 mask 中解析出 train/val/test 节点索引;此外还专门提取正类训练节点索引rl_idx,用于驱动强化学习模块(main.py)。
  2. 模型构建:以特征维度、类别数、隐藏维度、层数、激活函数(tanh)、RL 步长与图的全规范边类型(graph.canonical_etypes)构造CAREGNN
  3. 训练组件:由于类别不均衡,损失函数使用按类别计数加权的 CrossEntropyLoss——th.nn.CrossEntropyLoss(weight=1 / cnt);优化器为 Adam(lr=0.01weight_decay=0.001);开启--early-stop时创建 patience=100 的 EarlyStopping。
  4. 训练循环:每轮前向得到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.xxxx

4.2 全部命令行参数

参数类型默认值说明
--datasetstramazon数据集名称,可选yelpamazon
--gpuint-1GPU 索引,-1 表示使用 CPU
--hid_dimint64隐藏层维度
--num_layersint1CARE-GNN 层数
--max_epochint30最大训练轮数
--lrfloat0.01学习率
--weight_decayfloat0.001权重衰减
--step_sizefloat0.02RL 动作步长(论文公式中的 λ2)
--sim_weightfloat2相似度损失权重(论文公式中的 λ1)
--early-stopflagFalse是否启用早停

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):

  1. 对每个规范边类型etype,通过g.apply_edges(self._calc_distance, etype=etype)计算邻居距离。_calc_distance对应论文公式 2:d = || tanh(MLP(h_src)) - tanh(MLP(h_dst)) ||_1,即用 MLP 嵌入后的 L1 距离度量中心节点与邻居的相似度(距离越小越相似)。
  2. 调用_top_p_sampling按比例 p 保留距离最小的前ceil(in_degree * p)条入边。当前实现基于np.argpartition完成部分排序,代码注释说明其效率较低,可优化方向是 DGL 的dgl.sampling.select_top_p
  3. 对采样后的边执行g.send_and_recv,用fn.copy_u("h","m")+fn.mean("m","h_etype")完成按边类型聚合。
  4. 关系间聚合:将各关系的聚合结果按 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;
  • simtanh(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):

  1. 对每个层 i、每个边类型计算距离dist[etype] = L1(tanh(MLP(feat_i))),缓存在dists中;
  2. 用当前 p 构造CARESampler
  3. 通过dgl.dataloading.DataLoaderbatch_size=256shuffle=True构建训练/验证/测试 mini-batch,逐 batch 前向、计算加权 CE 损失并反向更新;
  4. 每轮结束后调用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 对比:

DatasetAmazonYelp
Metric (val / test)Max Epoch 30Max Epoch 30
AUC (val/test) — paper reported0.8973 / -0.7570 / -
AUC (val/test) — DGL full graph0.8849 / 0.89220.6856 / 0.6867
AUC (val/test) — DGL sampling0.9350 / 0.93310.7857 / 0.7890
Recall (val/test) — paper reported0.8848 / -0.7192 / -
Recall (val/test) — DGL full graph0.8615 / 0.85440.6667 / 0.6619
Recall (val/test) — DGL sampling0.9130 / 0.90450.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_edgessend_and_recvcanonical_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.

项目地址:https://gitcode.com/gh_mirrors/dg/dgl
点击查看免费下载

创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考

版权声明: 本文来自互联网用户投稿,该文观点仅代表作者本人,不代表本站立场。本站仅提供信息存储空间服务,不拥有所有权,不承担相关法律责任。如若内容造成侵权/违法违规/事实不符,请联系邮箱:809451989@qq.com进行投诉反馈,一经查实,立即删除!
网站建设 2026/9/23 4:32:51

从0到1搭建AI Agent平台:架构设计与工程实践

最近一年&#xff0c;"AI Agent"这个词几乎被聊烂了。我身边不少开发者分成了两拨&#xff1a;一拨觉得Agent无非就是"大模型加一个循环调用"&#xff0c;另一拨正在认真琢磨怎么把Agent变成公司里真正能上岗、能交付成果的"数字同事"。我属于后…

作者头像 李华
网站建设 2026/9/23 4:32:17

轻量应用服务器:云服务器部署的极简方案与选型实战

1. 轻量应用服务器到底是什么先说个我自己的经历。前几年给一个小创业团队做官网&#xff0c;老板开口就是“上云”&#xff0c;我第一反应是去ECS控制台选配置。选完系统盘、数据盘、带宽、安全组规则&#xff0c;再配一堆乱七八糟的选项&#xff0c;折腾了一下午。后来换了轻…

作者头像 李华
网站建设 2026/9/23 4:31:33

QNX虚拟化部署实战:VirtualBox中构建实时微内核环境

1. QNX不是Linux&#xff0c;也不是Windows——它是一台“工业级精密钟表”很多人第一次听说QNX&#xff0c;是在车载芯片的新闻里&#xff1a;高通8155平台用QNX做仪表系统&#xff0c;黑莓当年靠它撑起企业安全终端&#xff0c;特斯拉早期座舱原型机跑的也是QNX。但当你打开V…

作者头像 李华
网站建设 2026/9/23 4:31:32

SpringBoot+Vue儿童性教育网站管理系统架构与权限设计解析

之前带团队做未成年人教育类产品时&#xff0c;我们内部反复讨论过“儿童性教育内容该怎么管理”这个问题。这个品类很特殊&#xff0c;不像数学语文&#xff0c;老师可以用一套标准课件讲遍所有年级&#xff0c;它必须分龄、分场景、内容要经过严格的科学审核&#xff0c;后台…

作者头像 李华
网站建设 2026/9/23 4:29:15

STM32CubeMX + VS Code 从零搭建第一个STM32工程完整指南

1. 为什么第一个STM32工程值得认真对待很多人学STM32的方式是&#xff1a;装好Keil&#xff0c;找个现成工程&#xff0c;编译下载&#xff0c;灯亮了&#xff0c;就算入门了。但真到了要自己从零搭一个工程、换一颗不同封装的芯片、或者把代码交给同事接手的时候&#xff0c;问…

作者头像 李华
网站建设 2026/9/23 4:24:21

三相离网逆变器VSG控制:惯量阻尼整定与电压波形优化调试

做离网逆变器的人应该都有体会&#xff0c;负载一突加&#xff0c;母线电压抖一下&#xff0c;频率跟着掉一截。尤其是带电机、整流设备这类负载的时候&#xff0c;传统PQ控制根本没法独立支撑&#xff0c;下垂控制虽然能分功率&#xff0c;但频率变化太硬&#xff0c;没有惯量…

作者头像 李华