Paper Reading: From GNNs to Trees: Multi-Granular Interpretability for Graph Neural Networks


Paper Reading 是从个人角度进行的一些总结分享,受到个人关注点的侧重和实力所限,可能有理解不到位的地方。具体的细节还需要以原文的内容为准,博客中的图表若未另外说明则均来自原文。

论文概况 详细
标题 《From GNNs to Trees: Multi-Granular Interpretability for Graph Neural Networks》
作者 Jie Yang, Yuwen Wang, Kaixuan Chen, Tongya Zheng, Yihe Zhou, Zhenbang Xiao, Ji Cao, Mingli Song, Shunyu Liu
发表会议/期刊 ICLR (International Conference on Learning Representations)
发表年份 2025
会议/期刊等级 CCF-A
论文代码 https://github.com/dutyj2020/TIF

作者单位:

  1. Zhejiang University(浙江大学)
  2. State Key Laboratory of Blockchain and Data Security, Zhejiang University(浙江大学区块链与数据安全全国重点实验室)
  3. Hangzhou High-Tech Zone (Binjiang) Institute of Blockchain and Data Security(杭州高新区(滨江)区块链与数据安全研究院)
  4. Big Graph Center, Hangzhou City University(杭州城市大学大图中心)
  5. Nanyang Technological University(南洋理工大学)

研究动机

图神经网络(GNN)在图分类任务中取得了广泛的应用,但模型的可解释性一直是一个难题。可解释的图神经网络旨在揭示模型预测背后的推理过程,将决策归因于图中特定的信息子结构。现有的可解释方法主要分为两类:基于子图的事后解释方法和内在解释方法,它们通过识别关键子图来解释模型决策。然而这些方法存在一个共同的局限性:过度关注局部结构,容易忽略整个图中的长程依赖关系。
image

为了解决局部结构的局限性,近期工作 GIP 引入了可学习的图粗化机制来实现全局可解释性。GIP 通过将粗化后的图实例与多个自解释图原型对齐,从全局视角揭示模型的推理过程。然而 GIP 将图粗化到固定粒度,只能捕获输出层的图连通性,无法反映中间层次的多粒度信息。现实世界中的图任务往往涉及不同粒度层次的关系,例如蛋白质中的相关交互可以跨越功能团、氨基酸和蛋白质结构域等多个层次。固定粒度的粗化方式会丢失图内在结构中必要的细节信息,也限制了模型对不同大小和结构图类型的适应能力。

综上所述,现有的可解释 GNN 方法要么关注局部子图而忽略全局依赖,要么采用固定粒度的全局粗化而丢失中间层次的结构信息。本文针对这一问题,提出了一种将 GNN 转化为层次化树结构的方法,通过多粒度粗化来同时捕获局部和全局的结构信息,从而提供更全面的可解释性。

文章贡献

本文针对可解释图神经网络在多粒度结构信息捕获方面的不足,提出了一种树状可解释框架(Tree-like Interpretable Framework, TIF)。TIF 将普通的 GNN 转化为层次化的树结构,树的每一层包含不同粒度的粗化图作为树节点。TIF 由三个主要模块组成:层次图粗化模块通过迭代式图粗化将原始图(树的根节点)压缩为越来越粗的图(子节点),捕获多粒度结构信息;可学习图扰动模块为每个父节点的粗化过程引入多个可学习扰动,增强树的表示多样性和鲁棒性;自适应路由模块为每个非叶节点分配路由器,动态选择从根到叶的最有信息量的路径,用于模型预测和可解释性。在五个真实世界数据集和三个合成数据集上的实验表明,TIF 在预测性能上与主流 GNN 和可解释 GNN 方法相当或更优,在解释准确率、一致性、路径一致性和路径重要性等可解释性指标上显著优于现有方法。

本文方法

TIF 的整体框架如图所示,层次图粗化模块负责构建树的深度(纵向),可学习图扰动模块负责构建树的广度(横向),自适应路由模块负责在树结构中动态选择最有信息量的根到叶路径。通过以上模块,TIF 将输入图转化为多粒度的树结构,并基于该结构进行可解释的图分类。
image

层次图粗化模块

该模块的目标是通过迭代式图粗化来构建树的层次结构,每一层对应一个粒度级别的粗化图。首先使用图卷积网络(GCN)提取每层的节点嵌入 \(Z^{(l)}\),更新规则为:

\[Z^{(l)} = \sigma(\hat{D}^{-\frac{1}{2}}\hat{A}\hat{D}^{-\frac{1}{2}}Z^{(l-1)}W^{(l)}) \]

其中 \(Z^{(0)} = X\) 表示输入特征矩阵,\(\hat{A} = A + I\) 是添加自环的邻接矩阵,\(\hat{D}\)\(\hat{A}\) 的度矩阵,\(W^{(l)}\) 是权重矩阵,\(\sigma(\cdot)\) 是激活函数。然后通过一个带 softmax 输出的多层感知机(MLP)生成聚类分配矩阵 \(S^{(l)}\)

\[S^{(l)} = \text{softmax}(\text{MLP}^{(l)}(Z^{(l)}; \Theta_{\text{MLP}})) \]

其中 \(\Theta_{\text{MLP}}\) 表示可训练参数,\(S^{(l)}_{ij}\) 表示节点 \(v_i\) 属于簇 \(j\) 的概率。基于聚类分配矩阵,生成粗化图的新邻接矩阵 \(A^{(l+1)}\) 和新嵌入矩阵 \(X^{(l+1)}\)

\[X^{(l+1)} = \sum_{i=1}^{N} S^{(l)\top}_{ji} Z^{(l)}_i, \quad \forall j = 1, \ldots, K^{(l)} \]

\[A^{(l+1)} = \sum_{i=1}^{N}\sum_{k=1}^{N} S^{(l)\top}_{ji} A^{(l)}_{ik} S^{(l)}_{kj}, \quad \forall j = 1, \ldots, K^{(l)} \]

其中 \(N\) 是当前层的节点数,\(K^{(l)}\) 是第 \(l\) 层的簇数量。为了在粗化过程中保持图的连通性,本文引入边预测损失来约束粗化过程:

\[L_{\text{link}} = -\sum_{i,j}[A_{ij}\log\hat{A}_{ij} + (1-A_{ij})\log(1-\hat{A}_{ij})] \]

其中 \(A_{ij}\) 是原始邻接矩阵,\(\hat{A}_{ij}\) 是粗化后的邻接矩阵。

可学习图扰动模块

该模块为树中每个父节点的粗化过程引入多个可学习扰动,增强表示的多样性和鲁棒性。以树中第 \(l\) 层第 \(k\) 个节点的扩展过程为例,首先定义 \(M\) 个可学习扰动矩阵 \(P_{l,k} = \{P^{(l),k(1)}, P^{(l),k(2)}, \ldots, P^{(l),k(M)}\}\),用这些扰动矩阵对聚类分配矩阵 \(S^{(l),k}\) 进行扰动:

\[S^{(l),k(i)} = S^{(l),k} + P^{(l),k(i)}, \quad i = 1, 2, \ldots, M \]

其中 \(S^{(l),k}\) 表示第 \(l\) 层第 \(k\) 个节点的原始聚类分配矩阵。扰动后的节点嵌入 \(X^{(l),k(i)}\) 基于扰动后的分配矩阵计算:

\[X^{(l),k(i)} = S^{(l),k(i)\top}Z^{(l),k} = S^{(l),k\top}Z^{(l),k} + P^{(l),k(i)\top}Z^{(l),k} \]

其中 \(Z^{(l),k}\) 是第 \(l\) 层第 \(k\) 个父节点扩展时的节点嵌入矩阵。为了保证扰动后的嵌入既有效又多样,本文引入了两个正则化项。相似性正则化确保每个扰动嵌入 \(X^{(l),i}\) 与原始嵌入 \(X^{(l)}\) 保持接近,在施加扰动的同时保留重要的图结构:

\[L_{\text{similarity}} = \sum_{l=1}^{L}\sum_{k=1}^{K^{(l)}}\sum_{i=1}^{M}\lambda_i\|X^{(l),k(i)} - X^{(l),k}\|_2 \]

其中 \(\lambda_i\) 控制相似性项的强度,\(M\)\(K^{(l)}\)\(L\) 分别表示父节点的分支数、每层父节点数和树的层数。多样性正则化促进扰动嵌入之间的差异,确保每个分支代表原始图的不同变体:

\[L_{\text{diversity}} = \sum_{l=1}^{L}\sum_{k=1}^{K^{(l)}}\mu\sum_{i\neq j}\|X^{(l),k(i)} - X^{(l),k(j)}\|_2 \]

其中 \(\mu\) 控制多样性程度。两个正则化项构成总的扰动正则化损失 \(L_{\text{perturb}} = L_{\text{similarity}} + L_{\text{diversity}}\),在保留核心图结构和分支多样性之间取得平衡。

自适应路由模块

该模块为树模型的每一层每个非叶节点分配路由器,用于动态选择层次结构中信息量最大的根到叶路径。首先为第 \(l\) 层的非叶节点 \(k\) 分配一个路由器,将可学习图扰动模块生成的扰动嵌入拼接作为输入:

\[\hat{Z}^{(l),k} = \text{MLP}([Z^{(l),k(1)}; Z^{(l),k(2)}; \ldots; Z^{(l),k(M)}]) \]

然后路由器基于最终节点嵌入 \(\hat{Z}^{(l),k}\) 生成一组路由 logits \(r^{(l),k}\),表示选择每条路径的可能性:

\[r^{(l),k} = W^{(2),r,k} \cdot \sigma(W^{(1),r,k} \cdot \hat{Z}^{(l),k} + b^{(1),r,k}) + b^{(2),r,k} \]

其中 \(W^{(1),r,k}\)\(W^{(2),r,k}\) 是父节点 \(k\) 的权重矩阵,\(b^{(1),r,k}\)\(b^{(2),r,k}\) 是偏置项,\(\sigma\) 是非线性激活函数(如 ReLU)。路由 logits 经 softmax 函数转换为概率分布 \(p^{(l),k,i} = \text{softmax}(r^{(l),k})_i\),选择概率最大的路径 \(\hat{i}_{l,k} = \arg\max_i p^{(l),k,i}\),并据此更新下一层的节点嵌入和邻接矩阵:

\[X^{(l+1),\hat{i}_{l,k}}_{\text{pooled}} = S^{(l),\hat{i}_{l,k}\top}Z^{(l)}, \quad A^{(l+1),\hat{i}_{l,k}}_{\text{pooled}} = S^{(l),\hat{i}_{l,k}\top}A^{(l)}S^{(l),\hat{i}_{l,k}} \]

为了鼓励对多条路径的探索,引入基于熵的正则化来促进路径选择过程的多样性:

\[L_{\text{entropy}} = -\sum_{l=1}^{L}\sum_{k=1}^{K^{(l)}}\sum_{i=1}^{M}p^{(l),k,i}\log(p^{(l),k,i}) \]

其中 \(p^{(l),k,i}\) 是第 \(l\) 层父节点 \(k\) 分配给每条路径的概率。

基于神经树结构的可解释分类

该模块基于构建的层次化树模型进行图分类,通过追踪所选的根到叶路径来整合树不同层的多粒度信息。对于每个测试图 \(G_t\),计算树中每层的路径选择概率:

\[p(\text{Path}^{(l),k} | G_t) = \frac{\exp(f(\text{Path}^{(l),k}, G_t))}{\sum_j \exp(f(\text{Path}^{(l),j}, G_t))} \]

其中 \(f(\text{Path}^{(l),k}, G_t)\) 是衡量第 \(l\) 层路径 \(k\) 对图 \(G_t\) 分类相关性的评分函数。选择概率最高的路径 \(\hat{k}^{(l)} = \arg\max_k p(\text{Path}^{(l),k} | G_t)\),该过程在每层迭代执行直到到达叶节点,得到路径序列 \(\{\hat{k}_1, \hat{k}_2, \ldots, \hat{k}_L\}\) 作为多粒度解释结果,其中 \(L\) 是树的总层数。

最后所选路径的嵌入 \(Z_{\hat{k}_L}\) 直接用作最终嵌入 \(\hat{Z} = Z_{\hat{k}_L}\),通过评分函数 \(f(\cdot)\) 和 softmax 获得分类概率分布 \(h_i = \text{softmax}(f(\hat{Z}))\)。使用交叉熵损失作为优化目标:

\[L_{\text{CE}} = \frac{1}{M}\sum_{i=1}^{M}\text{CrsEnt}(h_i, y_i) \]

其中 \(M\) 是批大小,\(y_i\) 是真实概率分布。最终的总损失函数结合了分类损失、边预测损失、熵正则化和扰动正则化:

\[L_{\text{total}} = L_{\text{CE}} + \alpha_1 L_{\text{link}} + \alpha_2 L_{\text{perturb}} + \alpha_3 L_{\text{entropy}} \]

其中 \(\alpha_1\)\(\alpha_2\)\(\alpha_3\) 分别控制边预测正则化、扰动正则化和熵正则化的强度。

实验结果

数据集和实验设置

本文在五个真实世界数据集和三个合成数据集上进行了实验,合成数据集由特定粒度层次的结构组合构成,具有多层次的粒度信息结构,用于更好地展示框架的可解释性。

数据集类别 数据集名称 任务/分类类型
真实世界数据集 ENZYMES、PROTEINS、D&D 生物信息学(蛋白质数据集)
MUTAG 生物信息学(分子数据集)
COLLAB 社交网络(科学合作数据集)
合成数据集 GraphCycle Cycle 与 Non-cycle 二分类
GraphFive Wheel、Grid、Tree、Ladder、Star 五分类
MultipleCycle Pure Cycle、Pure Chain、Hybrid Cycle、Hybrid Chain 四分类

对比方法分为四类:

类别 方法/模型
广泛使用的 GNN 模型 GCN、DGCNN、DiffPool、RWNN、GraphSAGE
基于子图的可解释 GNN GNNExplainer、SubgraphX、XGNN、ProtGNN、KerGNN、π-GNN、GIB、GSAT、CAL
基于全局的可解释 GNN GIP
神经树变体 Bi-Tree(TIF 简化版本)

评价指标方面,预测性能使用分类准确率和 F1 分数,解释性能包括四个指标:

评估指标 计算/衡量方式
解释准确率 用训练好的 GNN 预测不同方法生成的解释,并使用预测置信度作为衡量
一致性 用随机游走图核计算生成解释与真实标签之间的相似度
路径一致性 重复输入测试样本,记录路径选择的一致率
路径重要性 分析路径利用频率,并使用归一化熵衡量

预测性能对比

在预测性能方面,本文在八个数据集上与广泛使用的 GNN 模型和可解释 GNN 模型进行了比较,结果如表所示。实验结果表明 TIF 在预测性能上与主流 GNN 模型相当或更优,在 MUTAG 数据集上准确率提升了 0.09% 到 35.77%,在八个数据集中的六个上取得了最高或次高的 F1 分数。与现有的可解释 GNN 模型相比,TIF 在八个数据集中的六个上取得了更高的准确率,在四个数据集上取得了更高的 F1 分数。在 D&D 数据集上 TIF 的准确率达到 84.19%,F1 分数达到 81.01%,优于所有对比方法。在合成数据集 GraphCycle 和 MultipleCycle 上,TIF 分别达到 84.77% 和 69.04% 的准确率,均位列所有方法之首。可见 TIF 在保持可解释性的同时没有牺牲预测性能。
image

解释性能对比

在解释性能方面,本文与基于子图和基于全局的可解释方法进行了比较,结果如论文表所示。在解释准确率上,与基于子图的可解释方法相比,TIF 在八个数据集中的五个上取得了最高的解释准确率,在其余数据集上位列第二。与唯一的基于全局的可解释基线 GIP 相比,TIF 在大多数数据集上也取得了可比的结果。在 D&D 数据集上 TIF 的解释准确率达到 89.11%,在 PROTEINS 上达到 87.62%,均显著优于 GIP 的 83.47% 和 86.04%。
image

在一致性方面,本文在两个合成数据集上计算了不同方法生成的解释与真实标签之间的相似度,如图(a)所示。TIF 的一致性显著高于大多数基于子图和基于全局的可解释方法,说明其生成的解释与真实结构更接近。在路径一致性方面,本文将 TIF 与其简化版本 Bi-Tree 进行了比较,如图(b)所示。TIF 在所有数据集上的路径选择一致性均高于 Bi-Tree,表明自适应路由模块在保持决策路径稳定性方面的有效性。在路径重要性方面,如图(c)所示,TIF 的路径重要性分布在所有数据集上较为均衡,没有单一路径过度占据主导地位。说明模型没有过度依赖少数决策路径,有助于提升模型的鲁棒性。
image

消融实验

本文对压缩比 \(q\)、路径数 \(N\) 和路由复杂度进行了消融实验。在压缩比实验中,本文在 D&D 数据集上将 \(q\) 设置为 \(\{0.1, 0.2, 0.3, 0.5\}\) 进行测试,结果如图(a)所示。实验结果表明分类和解释准确率在 \(q\) 过高或过低时都会下降,较低的 \(q\) 保留了噪声结构,较高的 \(q\) 导致关键信息丢失,说明存在一个平衡点来兼顾信息保留和噪声控制。在路径数实验中,本文将 \(N\) 设置为 \(\{2, 4, 8\}\) 进行测试,结果如图(b)所示。实验结果表明 4 条路径在分类和解释性能上均取得最佳结果。较少的路径限制了信息融合,较多的路径引入噪声,4 条路径在信息利用和噪声控制之间提供了有效的平衡。
image

在路由复杂度和扰动效果实验中,本文将基于 MLP 的路由模块替换为更简单的线性结构(去除层间自适应路由机制,w/o IAR),并替换扰动模块(w/o PM),结果如图所示。TIF 在分类和解释任务上均优于两个变体,表明路由结构能更好地捕获路径之间的复杂关系,扰动结构能有效捕获和学习有利于分类和可解释性的信息。
image

定性分析

本文可视化展示了 TIF 在 MultipleCycle 数据集上的推理过程,如图所示。可以观察到 TIF 在较细的解释中有效捕获了局部子结构,在较粗的解释中捕获了全局图模式,确保不同粒度的关键特征得到保留。自适应路由模块根据多粒度复杂度动态选择树中信息量最大的路径。
image

本文还将 TIF 与 GIP 在 MultipleCycle 数据集上生成的解释进行了对比,如图所示。TIF 从细粒度局部交互到粗粒度全局结构的能力提供了更透明和可解释的决策过程,说明了不同层次的图信息如何贡献于最终的模型预测。
image

优点和讨论

个人认为,本文有如下一些优点和创新点可供参考学习:

  1. 将 GNN 转化为层次化树结构来实现多粒度可解释性,思路新颖且符合直觉。本文通过树的深度和广度两个维度分别对应粒度层次和分支多样性,将多粒度解释问题映射为树结构上的路径选择问题。
  2. 三个模块的设计协同有效,层次图粗化模块负责纵向粒度层次,可学习图扰动模块负责横向分支多样性,自适应路由模块负责路径选择。
  3. 扰动模块同时引入相似性正则化和多样性正则化来平衡结构保留和分支差异,相似性正则化保证扰动后的嵌入不偏离原始图结构太远,多样性正则化保证不同分支学到不同的变体,两个约束共同作用使树结构既有信息量又不失稳定性。
  4. 评价指标体系较为完整,不仅关注预测性能,还包括了四个解释性能指标。路径一致性和路径重要性两个指标从稳定性和均衡性两个角度评估决策路径的质量。

这篇论文提出的模型和其他神经树的结构很类似,本质上是全局图输入根节点后通过路由机制逐层选择分支,最终到达一条路径末端的最粗粒度图再送入分类头决策。所以关键在于如何去定义树结构中的非叶节点的表示方式,这篇文章巧妙的地方在于使用了逐层粗化的图(邻接矩阵+嵌入)作为节点内容,对应传统决策树中那个被选择出来的特征和阈值。从这个角度来看,我觉得这篇论文是不错的。

posted @ 2026-09-02 14:45  乌漆WhiteMoon  阅读(15)  评论(0)    收藏  举报