news 2026/9/13 10:12:39

PyTorch Geometric 实验性功能实战:基于 contrib 包的 RBCD 图对抗攻击与 PGM 图神经网络可解释性

作者头像

张小明

前端开发工程师

1.2k 24
文章封面图
PyTorch Geometric 实验性功能实战:基于 contrib 包的 RBCD 图对抗攻击与 PGM 图神经网络可解释性

PyTorch Geometric 实验性功能实战:基于 contrib 包的 RBCD 图对抗攻击与 PGM 图神经网络可解释性

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

导读

本文围绕 examples/contrib/README.md 展开,系统讲解 PyTorch Geometric(PyG)中torch_geometric.contrib实验性功能包的四个官方示例:RBCD(Resource-based Critical Data)图对抗攻击的逃逸(Evasion)与投毒(Poisoning)两种场景,以及 PGM(Probabilistic Graphical Model,概率图模型)解释器在节点分类与图分类任务上的应用。读完本文,你将掌握GRBCDAttack/PRBCDAttack的调用方式、超参数调优要点,以及如何用PGMExplainer解释 GNN 的预测结果,并能直接运行仓库中的四个示例脚本进行验证。

认识 contrib 包:PyG 的实验性功能试验场

torch_geometric.contrib是 PyG 为早期、实验性代码提供的暂存区(staging area)。根据 examples/contrib/README.md 的说明,其中的模块未来可能被移入主库。这意味着:

  • API 可能变动:contrib 中的类和方法不保证向后兼容,升级 PyG 版本时需关注 CHANGELOG;
  • 示例即文档:由于处于实验阶段,官方文档对 contrib 功能的说明较少,examples/contrib 目录下的示例脚本就是最直接的用法参考;
  • 独立模块化:contrib 包按功能分为contrib.nn(网络与攻击模块)与contrib.explain(可解释性模块)等子包,详见 torch_geometric/contrib/nn/init.py 与 torch_geometric/contrib/explain/init.py。

本目录共包含四个示例,覆盖两大主题:

示例文件主题
rbcd_attack.pyRBCD(Resource-based Critical Data)攻击的逃逸示例
rbcd_attack_poisoning.pyRBCD 攻击结合数据投毒策略的示例
pgm_explainer_node_classification.pyPGM 解释器用于节点分类任务
pgm_explainer_graph_classification.pyPGM 解释器用于图分类任务

运行环境与前置依赖

四个示例均依赖 PyTorch 与 PyG 主库,其中投毒示例额外依赖higher库用于内层循环(inner-loop)优化:

pip install higher

rbcd_attack_poisoning.py在导入higher失败时会直接退出并提示安装命令(见 rbcd_attack_poisoning.py)。所有示例会自动选择cuda(若可用)否则cpu设备。数据文件默认下载到示例脚本同级的data/目录。

RBCD 图对抗攻击:原理与 API

攻击原理:松弛化的随机块坐标下降

RBCD 系列攻击源自论文Robustness of Graph Neural Networks at Scale,其核心思想是:只扰动邻接矩阵(增删边),不扰动节点特征,因此适用于任何能处理带权图、且对边权可微的 GNN 模型(如GCNConvGraphConv)。

其中两个攻击类定义在 torch_geometric/contrib/nn/models/rbcd_attack.py:

  • PRBCDAttack(投影随机块坐标下降):攻击期间将邻接矩阵的离散条目从{0, 1}松弛到[0, 1],通过梯度更新边权,再用投影操作保证松弛后的 L0 预算约束,最后采样得到离散的扰动图;
  • GRBCDAttack(贪心随机块坐标下降):共享 PRBCD 的梯度机制,但每一步贪心地基于梯度直接翻转边(torch.topk(gradient, step_size)取梯度最大的边置 1),实现见 rbcd_attack.py 中 GRBCDAttack._update。

两种攻击都通过随机块采样控制内存开销:每轮只在一批随机的候选边(至多block_size条)上计算梯度。由于块是"有放回采样后去重",实际块大小通常略小于设定值(见源码 rbcd_attack.py 的 docstring 说明)。二者可用于:

  • 局部攻击(local)与全局攻击(global):通过idx_attack指定攻击目标(单个节点或整个测试集);
  • 逃逸攻击(evasion,测试时)与投毒攻击(poisoning,训练时):分别对应两个示例脚本。

PRBCDAttack 核心参数

参数默认值含义与调优要点
model必填待评估的 GNN 模块
block_size必填每轮随机采样的候选边数量,是内存开销的主要来源;示例中取250_000
epochs125攻击轮数(贪心模式下预算耗尽可提前终止)
epochs_resampling100前多少轮进行块重采样,之后转为固定搜索空间微调
loss'prob_margin'衡量攻击强度的损失,可选'masked''margin''prob_margin''tanh_margin',或传入自定义可调用对象
metricloss用于监控/早停的第二个(可不可微)指标
lr1000边权更新学习率;PRBCD 最重要的超参数之一,最佳实践是让预算在几步内耗尽
is_undirectedTrue图是否为无向图
logTrue是否打印攻击进度(tqdm 进度条)

源码中的损失函数实现(见 rbcd_attack.py)给出了更精确的语义:

  • margin:真实类得分与最高非目标类得分之差m = -s_y + max_{y'≠y} s_y'
  • prob_margin:对 softmax 概率计算 margin,聚焦决策边界附近的节点;
  • tanh_margin:对 margin 取 tanh,同样关注边界节点;
  • masked:仅在预测正确的节点上计算交叉熵(argmax == labels的样本才计入损失)。

attack() 方法签名

pert_edge_index, perts = attack( x, # 节点特征矩阵 edge_index, # 边索引 labels, # 标签 budget, # 允许翻转(增删)的边数上限 idx_attack=None, # 攻击目标(节点索引/掩码),None 表示全部 **kwargs, # 透传给 GNN 模块的额外参数 )

返回值为(perturbed_edge_index, flipped_edges)二元组,即扰动后的边索引与具体被翻转的边列表。attack()内部会在 attack 方法 中维护attack_statistics字典记录每步的损失、投影前后的概率质量等,供后续分析或绘图。

实战一:逃逸攻击(Evasion)——测试时篡改图结构

examples/contrib/rbcd_attack.py 演示了在 Cora 数据集上分别对 GAT 做局部攻击、对 GCN 做全局攻击的完整流程。

构造可攻击的模型

示例中定义了两个模型类:

  1. GCN(rbcd_attack.py 第 20-44 行):两层GCNConv,关键点是normalize=False关闭卷积内部归一化,改为在forward中调用gcn_norm只归一化一次,并支持skip_norm标志——这样攻击过程中边权重变化时无需重复做昂贵的归一化;
  2. GAT(rbcd_attack.py 第 47-86 行):由于标准GATConv不接受边权,示例通过继承GATConv重写edge_update实现带边权的 GATWeightedGATConv)。其技巧是:将源/目标注意力系数相加模拟拼接,并用alpha + torch.log2(edge_attr)将边权以对数形式融入注意力(边权为 1 时 alpha 不变,为 0 时趋近 -Inf),从而规避后续 exp/softmax 的下溢问题,自环边权初始化为fill_value=1.

局部攻击:攻击单个节点

node_idx = 42 local_budget = 2 # 训练节点 42 的度为 2,即最多翻 2 条边 grbcd = GRBCDAttack(gat, block_size=250_000) prbcd = PRBCDAttack(gat, block_size=250_000, metric=metric, lr=2_000) # GRBCD:攻击单节点 pert_edge_index, perts = grbcd.attack( data.x, data.edge_index, data.y, budget=local_budget, idx_attack=[node_idx], )

metric定义为负的准确率(越小越好,与损失方向一致):

def metric(*args, **kwargs): return -accuracy(*args, **kwargs)

示例通过PRBCDAttack._probability_margin_loss(源码中为静态方法,rbcd_attack.py)计算攻击前后目标节点"真实类到最佳非目标类"的置信度边界(confidence margin),直观展示攻击效果:边界值从攻击前的正值跌向负值,意味着模型对目标节点的分类信心被显著破坏,随后打印被翻转的具体边(u, v)列表。

全局攻击:攻击整个测试集

# 扰动 5% 的边(无向图每条边存两份,故除以 2) global_budget = int(0.05 * data.edge_index.size(1) / 2) pert_edge_index, perts = grbcd.attack( data.x, data.edge_index, data.y, budget=global_budget, idx_attack=data.test_mask, # 目标为全部测试节点 )

全局攻击后,用copy.copy(data)复制数据并替换edge_index,重新评估 GCN 在测试集上的准确率,打印"Clean accuracy → Perturbed accuracy"的下降幅度。逃逸场景下模型参数在攻击前后保持不变,这正对应"测试时攻击"的定义。

学习率选择的启发式

源码 rbcd_attack.py 中 _update_edge_weights 显示,PRBCD 的实际学习率会按lr * budget / num_nodes / sqrt(max(0, epoch - epochs_resampling) + 1)做启发式修正,使其与预算、图规模无关,并在重采样阶段结束后(固定搜索空间)自然衰减。示例注释给出的经验法则是:选择一个能让预算在几步内耗尽的学习率(如lr=2_000),同时高学习率还能缓解边权松弛{0,1} → [0,1]带来的松弛间隙(relaxation gap)影响。GRBCD 由于是贪心翻转,在小预算下比 PRBCD 更快,但结果的一致性略逊(示例注释明确指出)。

实战二:投毒攻击(Poisoning)——训练时污染图结构

examples/contrib/rbcd_attack_poisoning.py 演示了训练时攻击:在 GCN 重新训练之前篡改邻接矩阵,观察最终模型性能下降。相比逃逸攻击,投毒需要模拟"攻击者修改图后,受害者在其上重新训练模型"的过程,因此必须引入双层优化(bi-level optimization)。

双层优化的实现:内层循环重训练

示例通过子类化PRBCDAttack并重写两个关键钩子来实现:

  • _forward(rbcd_attack_poisoning.py 第 44-55 行):每次前向先model.reset_parameters(),然后在torch.enable_grad()下用扰动后的图train(self.model, ped, n_epochs, lr, weight_decay)完整重训模型(50 轮、学习率 0.04、权重衰减 5e-4),模拟受害者对新图的适应过程;
  • _forward_and_gradient(rbcd_attack_poisoning.py 第 57-102 行):借助higher.innerloop_ctx构建可微的内层训练循环,让扰动边权的梯度能够穿过整个重训练过程反向传播;同时将梯度裁剪到范数0.5以保证数值稳定性。

示例开头有一句关键注释(rbcd_attack_poisoning.py 第 25-26 行):投毒场景下边权最终会被忽略,邻接矩阵的预处理(如归一化)应放在模型内部(参与反向传播),这正是示例 GCN 把gcn_norm放进forward的原因。

攻击流程与结果验证

prbcd = PoisoningPRBCDAttack(gcn, block_size=250_000, metric=metric, lr=100) global_budget = int(0.05 * data.edge_index.size(1) / 2) pert_edge_index, perts = prbcd.attack( data.x, data.edge_index, data.y, budget=global_budget, idx_attack=data.test_mask, ) # 用扰动后的图从零重训并评估 gcn.reset_parameters() pert_data = copy.copy(data) pert_data.edge_index = pert_edge_index train(gcn, pert_data) pert_acc = test(gcn, pert_data) print(f'PRBCD: Accuracy dropped from {clean_acc:.3f} to {pert_acc:.3f}')

注意验证阶段同样要reset_parameters()后重新训练,才符合"投毒影响后续训练"的真实语义。由于重训练引入随机性,示例注释提示投毒场景的数值波动比逃逸场景更大

利用 attack_statistics 绘制调试曲线

示例最后用 matplotlib 绘制攻击过程的诊断曲线(rbcd_attack_poisoning.py 第 126-144 行):

  • 左轴(红色实线):每步loss
  • 右轴(蓝色虚线/实线):prob_mass_after_update(投影前的边权概率质量)与prob_mass_after_projection(投影后实际使用的预算)。

这张图正是验证"学习率应让预算尽快耗尽"这一经验法则的工具:若蓝线迟迟达不到预算上限,说明学习率偏低,攻击在松弛空间中徘徊、效率低下。

深入 PGM 解释器:用概率图模型解释 GNN 预测

PGMExplainer实现了论文PGMExplainer: Probabilistic Graphical Model Explanations for Graph Neural Networks(arXiv:1903.03894),源码位于 torch_geometric/contrib/explain/pgm_explainer.py。它生成的Explanation对象提供node_maskpgm_stats两个核心输出,其中pgm_stats保存了每个节点由Chi-squared 检验计算出的 p 值,用于量化"该节点对预测的影响显著性"。

PGMExplainer 核心参数

参数默认值含义
feature_indexNone被扰动的特征索引;None表示扰动全部特征
perturbation_mode'randint'特征扰动方式:'randint''mean''zero''max''uniform'
perturbations_is_positive_onlyFalse是否限制扰动值为正
is_perturbation_scaledFalse是否归一化扰动特征的范围
num_samples100用于显著性检验的扰动采样次数
max_subgraph_sizeNone解释考虑的邻居节点数上限
significance_threshold0.05p 值阈值,低于该值判定节点对预测有显著影响
pred_threshold0.1判断扰动后输出与原始输出"不同"的缓冲阈值

各扰动模式的底层实现在 pgm_explainer.py 的 _perturb_features_on_nodes:randint将特征置为 0/1 随机值;mean/zero/max分别用列均值、0、列最大值替换;uniform则在0.05 * max(x)幅度内加均匀噪声。

与 Explainer 框架的集成

PGM 解释器不是独立运行的,而是作为算法插件接入 PyG 的统一可解释性框架torch_geometric.explain.Explainer。调用时需同时配置ModelConfig,声明任务类型(multiclass_classification)、任务层级(nodegraph)与返回类型(raw)。

实战三:节点分类解释(Cora)

examples/contrib/pgm_explainer_node_classification.py 在 Cora 上训练一个两层 GCN(关闭卷积内归一化,改由T.GCNNorm()变换预处理),然后解释节点 100 的预测:

explainer = Explainer( model=model, algorithm=PGMExplainer(), node_mask_type='attributes', explanation_type='phenomenon', model_config=ModelConfig(mode='multiclass_classification', task_level='node', return_type='raw')) node_idx = 100 explanation = explainer(x=data.x, edge_index=edge_index, index=node_idx, target=predicted_target, edge_weight=edge_weight) print(f'Significance of relevant neighbors: {explanation.pgm_stats}')

要点解析:

  • explanation_type='phenomenon'表示用模型预测结果(而非真实标签)作为解释目标,因此传入target=predicted_target
  • 训练采用F.nll_loss+log_softmax输出,与return_type='raw'兼容;
  • 解释器会自动提取目标节点的k_hop_subgraph邻域并执行扰动-显著性检验,最终pgm_stats给出每个相关邻居的 p 值,p 值低于significance_threshold(默认 0.05)的节点即为对预测起关键作用的邻居。

实战四:图分类解释(MNIST Superpixels)

examples/contrib/pgm_explainer_graph_classification.py 将 PGM 解释器用于图级分类:模型是带NNConv(边属性为 2 维坐标差)与graclus+max_pool分层池化的网络,数据集为 MNIST 超像素图(MNISTSuperpixels,经T.Cartesian(cat=False)变换生成边属性)。

explainer = Explainer( model=model, algorithm=PGMExplainer(perturb_feature_list=[0], perturbation_mode="mean"), explanation_type='phenomenon', node_mask_type="object", model_config=dict(mode="multiclass_classification", task_level="graph", return_type="raw"))

与节点分类示例的区别:

  • task_level='graph',解释目标是整张图;
  • node_mask_type='object'(节点分类示例为'attributes'),配合max_pool_x/global_mean_pool这类池化层;
  • perturb_feature_list=[0]+perturbation_mode='mean':只扰动第一个特征(超像素的 x 坐标),用列均值替换,扰动空间更小、解释更聚焦;
  • 解释器对测试集的每张图逐一解释,通过explanation.available_explanations遍历输出(如node_maskpgm_stats),示例仅处理前 3 张图。

测试覆盖:如何验证这些实验性功能

仓库在 test/contrib 下为这两类功能提供了配套测试,可作为理解行为边界的参考:

  • test/contrib/nn/models/test_rbcd_attack.py:覆盖 GRBCD/PRBCD 攻击的预算约束(扰动边数不超过 budget)、返回值形状等;
  • test/contrib/explain/test_pgm_explainer.py:覆盖 PGM 解释器在不同perturbation_mode、不同任务层级下的输出结构。

总结与实验路线

本目录是理解 PyG 实验性功能的快速入口,推荐按以下顺序动手:

  1. 运行 rbcd_attack.py,对比 GAT 局部攻击与 GCN 全局攻击的输出,观察置信度边界与准确率的下降;
  2. 调整lrbudgetblock_size,结合attack_statistics理解学习率-预算的关系;
  3. 运行 rbcd_attack_poisoning.py(需先pip install higher),观察双层优化下的投毒效果与诊断曲线;
  4. 依次运行两个 PGM 示例,对比node_mask_typetask_level配置差异对解释输出的影响,并尝试更换perturbation_mode观察显著性结果的变化。

需要注意,contrib 包内模块仍处于演进阶段,其 API 可能随版本调整;若在后续版本中遇到 API 变更,请以当前仓库源码与 CHANGELOG 为准。

【免费下载链接】pytorch_geometricGraph Neural Network Library for PyTorch项目地址: https://gitcode.com/GitHub_Trending/py/pytorch_geometric

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

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

RK3588S软实时化实战:从Ubuntu到工业控制的确定性优化

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/13 10:11:01

备份介质选型实战:10TB×30天背后的性能、成本与可靠性博弈

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/13 10:09:47

Chrome侧边栏投屏替代QtScrcpy的技术原理与实战

/* MD / 富文本中的 .toc(含博客园搬家等嵌套结构);.toc-box 在侧栏,不受影响 */#content_views .toc,/* 编辑器常在目录前后插入空 p(:empty 仍占 20px),一并去掉避免顶空隙 */#content_views.markdown_views > p:empty:has(+ .toc),#content_views.markdown_views …

作者头像 李华
网站建设 2026/9/13 10:07:36

BearPi-HM_Nano农业传感节点实战:E53_IA1驱动与华为IoT平台稳定接入

简介:本资源是一套基于BearPi-HM_Nano开发板的物联网农业监测系统完整工程,面向嵌入式初学者、物联网开发爱好者及高校课程实践者,解决农业环境数据采集、边缘设备联网与云平台对接等典型IoT落地问题。压缩包共305个文件,含57个C语…

作者头像 李华
网站建设 2026/9/13 10:05:43

AI改写工具核心技术解析与应用实践

1. 项目概述:AI改写工具如何重塑内容创作流程去年帮一位研究生朋友修改论文时,我第一次深度体验了某款AI改写工具。当时她正为查重率居高不下而焦虑,传统的手动改写耗时耗力。使用该工具的智能改写功能后,论文核心观点保留完整的前…

作者头像 李华