Paper Reading:GraphChef: Decision-Tree Recipes to Explain Graph Neural Networks

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

论文概况 详细
标题 《GraphChef: Decision-Tree Recipes to Explain Graph Neural Networks》
作者 Peter Müller, Lukas Faber, Karolis Martinkus, Roger Wattenhofer
发表会议/期刊 International Conference on Learning Representations(ICLR)
发表年份 2024
会议/期刊等级 CCF-A
论文代码 https://interpretable-gnn.netlify.app/

作者单位:

  1. ETH Zurich, Switzerland(苏黎世联邦理工学院)

研究动机

图(Graph)在化学、工程、社会科学和交通等诸多领域中被用来表示复杂的关联数据,图神经网络(GNN)作为处理图数据的流行模型。它在图分类任务中取得了广泛应用,例如判断一个以图表示的蛋白质是否为酶。然而 GNN 本质上是一个黑盒模型,其决策过程不透明,用户无法理解模型为何做出某个预测。现有 GNN 解释方法大致可分为五类:梯度方法、互信息方法、反事实方法、子图方法和示例方法。前三种方法计算节点或边级别的重要性分数,生成热力图来突出输入中哪些部分重要。子图方法通过展示重要子图或原型来解释。示例方法则给出相似的参考图。这些方法存在两个核心局限:

  1. 只能回答"哪些输入重要",无法回答"如何使用这些输入"。例如在 PROTEINS 数据集中,现有方法可以指出两个 Sheet 节点是重要的,但无法回答"Sheet 之间是否需要连接""是否必须恰好有两个 Sheet"等问题。
  2. 只能针对单个图进行解释,无法给出数据集层面的整体规律。用户需要逐一查看数十个图的局部解释,难以从中归纳出数据集层面的分类依据。


本文希望将解释推进一步,不仅理解哪些输入重要,还要理解它们如何被使用;不仅解释单个图,还要理解整个数据集,回答"什么使一个蛋白质成为酶"这样的问题。

文章贡献

针对 GNN 缺乏可解释性的问题,本文提出了一个自解释图神经网络模型 GraphChef 。GraphChef 的思路是将决策树集成到 GNN 消息传递框架中,使得训练完成后能生成一组人类可理解的规则(recipe),同时解释数据集中所有图的分类依据。在方法设计上,本文首先提出了一种新的 GNN 层 dish,灵感来自分布式计算中的 stone-age 模型,通过 Gumbel-Softmax 将节点内部状态从连续向量变为离散的 one-hot 类别值,使得消息聚合简化为对邻居状态的计数。随后将训练好的 dish GNN 中的所有神经网络模块蒸馏为决策树,包括编码器、各层更新函数和解码器,每个决策树的分支可以基于节点自身状态、邻居计数或状态间的比较来判断,形成层次化的可解释 recipe。本文还提出了一种剪枝方法,同时利用训练集和验证集的精度作为约束条件,在不损失或极少损失精度的前提下大幅缩减决策树规模。实验在 6 个合成数据集和 7 个真实数据集上进行,结果表明 GraphChef 的精度与 GIN 基线相当,解释质量与现有解释方法具有竞争力。

本文方法

GraphChef 的起点是 GIN 模型。在标准的消息传递 GNN 框架中,每个节点在每个层中计算消息发送给邻居,然后聚合收到的消息并更新自身状态。GIN 的消息函数为恒等映射,聚合方式为逐元素求和,更新函数为可学习的神经网络。GIN 的更新公式如下:

\[h_v^{l+1} = f_\theta^l(h_v^l, \sum_{w \in Nb(v)} h_w^l) \]

其中 \(h_v^l\) 是节点 \(v\) 在第 \(l\) 层的内部状态,\(f_\theta^l\) 是可学习的神经网络,\(Nb(v)\)\(v\) 的邻居集合。GIN 的内部状态是 \(d\) 维实数向量 \(h_v^l \in \mathbb{R}^d\),这些连续向量允许特征之间形成复杂关系,使得可解释性非常困难。本文的改进是对 GNN 层和编码器层施加 Gumbel-Softmax,将内部状态变为 one-hot 类别值:

\[h_v^{l+1} = \text{Gumbel}(f_\theta^l(h_v^l, \sum_{w \in Nb(v)} h_w^l)) \]

经过这一变换,内部状态 \(h_v^l\) 变成 one-hot 类别值,聚合步骤中的求和就变成了统计每个状态类别的邻居数量。该设计的理论动机来自分布式计算领域:Loukas (2020) 证明了消息传递 GNN(如 GIN)等价于 LOCAL 分布式计算模型,其中节点可以执行任意局部计算。而更简单的 stone-age 模型(Emek and Wattenhofer, 2013)限制节点只能传输和处理有限数量的状态类别,类似于“一、二、三、许多”的计数方式,超过某个阈值的邻居数之间不可区分。这种简化模型仍然能解决许多分布式计算问题。dish 层就是 stone-age 模型在 GNN 中的对应实现。

对于使用 \(\log(d)\) 比特连续嵌入空间的 GIN,理论上可以构建一个具有 \(d\) 个类别状态的 dish GNN 达到相同的理论表达力,实践中本文追求低数量的类别状态以确保人类可解释性。编码器对初始节点特征 \(x_v\) 编码得到第一层 dish 的初始状态。解码器层对于节点分类使用跳跃连接,将所有中间状态用于最终预测;对于图分类则对每层的节点状态做求和池化,将各层求和结果用于最终预测。如图所示,一个 GraphChef 层中节点的新状态由三部分信息决定:上一层的类别状态(绿色)、每个状态类别中的邻居数量(蓝色)、以及状态之间的二值比较(黄色)。
image

从 dish 到 GraphChef

由于 dish GNN 中所有状态都是离散的 one-hot 类别值,可以将训练好的神经网络模块蒸馏为决策树。需要蒸馏的部分包括 dish 层的更新函数、编码器和解码器,由于状态是类别的,蒸馏就变成了一个分类问题,即决策树学习预测类别状态。蒸馏后的 GraphChef 层如下所示:

\[h_v^{l+1} = \text{TREE}_l(h_v^l, \sum_{w \in Nb(v)} h_w^l) \]

GraphChef 层仍然遵循消息传递框架,只是将神经网络加 Gumbel-Softmax 替换为决策树。为了让决策树更小,本文还引入了成对 delta 特征 \(\Delta\)。决策树通常不擅长比较两个特征的大小关系,delta 特征专门解决这个问题。令 \(c_i^l\) 为状态 \(i\) 的邻居计数,delta 特征对所有有序状态对计算二值比较:

\[\Delta(c_v^l) = \mathbf{1}[c_i^l > c_j^l] \quad (i \in S, j \neq i \in S) \]

\[h_v^{l+1} = \text{TREE}_l(h_v^l, c_v^l, \Delta(c_v^l)) \]

决策树可以访问三组特征:节点自身的上一状态、每个状态类别的邻居计数、以及两个状态类别计数之间的比较。如图所示,决策树的分支节点可以基于这三类特征进行判断:(a) 判断节点当前处于哪个状态;(b) 判断节点在某个状态中是否有特定数量的邻居;(c) 判断节点在一个状态中的邻居是否多于另一个状态。每个决策节点的分支类型可以直接解释为 GNN 推理的一个步骤,将所有节点的解释组合起来就构成了 GraphChef 的 recipe。
image

剪枝 GraphChef

足够深的决策树可以作为通用函数近似器,但本文需要小而浅的决策树来保证人类可理解性。为此本文提出了一种剪枝方法,采用的质量标准为:当用一个叶节点替换一个内部决策节点时,如果 (i) 验证集精度不下降且 (ii) 训练精度不低于验证精度,则接受替换。不允许验证精度下降是为了防止过度剪枝,而允许训练精度下降可以移除因过拟合产生的决策节点,训练精度不低于验证精度则是另一个防止过度剪枝的保障。

具体流程上,首先按覆盖的数据点数量对所有内部决策节点排序,依次尝试用叶节点替换,直到找不到可以无损替换的节点。然后继续迭代,每次移除验证精度下降最小的节点,允许轻微的精度下降。最终记录 10 个剪枝级别:无损剪枝后的状态以及在 10% 步长上进一步剪枝的结果。这些级别在用户界面中可供选择,用户可以在精度和树大小之间权衡。

计算解释分数

GraphChef recipe 除了提供数据集层面的解释,还可以为单个图计算热力图式的重要性分数,类似于现有图解释方法。计算过程逐层进行:在输入层,每个节点是自身唯一的解释。在每个 GraphChef 层中,对决策树的每个特征计算 Tree-Shap 值(Lundberg et al., 2018),用于衡量该特征对预测的重要性。然后根据特征类型分别传播重要性:

特征类型 重要性分配规则
状态特征 重要性保持在节点自身,通过 Tree-SHAP 值加权后更新。
消息特征 重要性均匀分配到该状态的所有邻居节点。
Delta 特征 正重要性分配给多数类邻居,负重要性分配给少数类邻居。

最后对所有分数归一化使总和为 1。在解码器层,通过跳跃连接也考虑中间状态(节点分类)或中间池化计数(图分类)。

实验结果

数据集和实验设置

本文使用了 6 个合成数据集和 7 个真实数据集。合成数据集带有解释的 ground truth,用于验证解释的正确性。真实数据集用于评估精度和 recipe 的可读性。各数据集的统计信息如下表所示。

数据集 任务类型 图数量 类别数 平均节点数 平均边数 特征维度 来源
Infection 节点分类 1 7 1000 3973 2 Faber et al. (2021)
Negative Evidence 节点分类 1 2 2000 102394 3 Faber et al. (2021)
BA-Shapes 节点分类 1 4 700 4110 0 Ying et al. (2019)
Tree-Cycle 节点分类 1 2 871 1942 0 Ying et al. (2019)
Tree-Grid 节点分类 1 2 1231 3130 0 Ying et al. (2019)
BA-2Motifs 图分类 1000 2 25 50.96 0 Luo et al. (2020)
MUTAG 图分类 188 2 17.93 39.59 7 Debnath et al. (1991)
Mutagenicity 图分类 4337 2 30.32 61.54 14 Kazius et al. (2005)
BBBP 图分类 2039 2 24.06 51.91 9 Wu et al. (2017b)
PROTEINS 图分类 1113 2 39.06 145.63 3 Borgwardt et al. (2005)
REDDIT-BINARY 图分类 2000 2 429.63 995.51 0 Borgwardt et al. (2005)
IMDB-BINARY 图分类 1000 2 19.77 193.06 0 Borgwardt et al. (2005)
COLLAB 图分类 5000 3 74.49 4914.43 0 Borgwardt et al. (2005)

使用的对比方法包括:

方法 类型 出处
GIN 图同构网络(GNN 基线) Xu et al. (2019)
dish GNN 类别状态 GNN(本文中间模型) 本文
DT 决策树(非图基线) 标准方法
DT+degrees 决策树 + 度特征 标准方法
Gradient 梯度解释方法 Baldassarre and Azizpour (2019)
GNNExplainer 互信息解释方法 Ying et al. (2019)
PGMExplainer 概率图模型解释方法 Vu and Thai (2020)

实验采用 10 折交叉验证,不同数据集划分分别训练 GraphChef 和 GIN 基线。GNN 训练 1500 个 epoch,允许基于验证损失早停(patience 为 100)。每个划分使用验证集分数进行早停。GIN 和 dish GNN 的更新函数使用 2 层 MLP,中间加入 Batch Normalization 和 ReLU 激活函数。使用 5 层图卷积。GIN 的隐藏维度为 16,GraphChef 的状态空间大小为 10。对于 GraphChef,训练集进一步划分出 holdout 集用于决策树剪枝。蒸馏后的每棵决策树最多 100 个节点。GraphChef 还支持自调节超参数:训练后可以检查 recipe 实际使用了多少层和状态,若发现冗余则重新训练。除 COLLAB 外所有数据集均可在普通 CPU 上训练。

对比实验

精度对比实验在两组数据集上比较 GIN、dish GNN 和 GraphChef(含无损剪枝版本)的分类精度。如表所示,GraphChef recipe 的精度与 GIN 非常接近。在需要复杂图推理的数据集上,GraphChef 也明显优于仅使用度特征的决策树基线。模型简化,即从连续向量到类别状态再到决策树没有带来精度下降。值得注意的是,剪枝在某些数据集上甚至提高了测试精度,这可能是由于剪枝过程引入的正则化效果。
image

在更高维特征的数据集上(如 Cora、CiteSeer、PubMed 和 OBGN-Arxiv),结果则较为复杂。如表所示,在 PubMed 上 GraphChef 与 GIN 表现相当,在 Cora 上 dish GNN 有小幅下降而蒸馏为决策树后下降明显,在 CiteSeer 上 dish GNN 和决策树蒸馏都导致明显下降。主要原因是高维特征空间难以压缩到少量类别状态,例如 Cora 数据集需要从 1433 维压缩到 10 个类别。此外,剪枝方法需要叶子节点数量平方级别的运行次数,在 OBGN-Arxiv 等大规模数据集上面临可扩展性瓶颈。
image

解释效果实验

解释效果实验利用合成数据集的 ground truth 来评估 GraphChef 的重要性分数。遵循已有工作的做法,对图中每个节点计算重要性分数,取分数最高的 \(n\) 个节点作为解释(\(n\) 为 ground truth 中的节点数),解释精度为解释中正确节点的比例。如表所示,GraphChef 的解释分数与现有解释方法具有竞争力。在 Infection 数据集上达到 0.95,在 Saturation 上达到 1.00,在 BA-Shapes 上达到 0.94。
image

此外,本文还展示了 GraphChef 和 GNNExplainer 在多个数据集上选取的重要节点对比。如图所示,GraphChef 的重要性分数在多个数据集上与 GNNExplainer 的表现相当。在合成数据集上,GraphChef 的 recipe 发现了数据集的一些已知缺陷:某些数据集(如 BA-2Motifs 和 Tree-Cycle)并不需要找到完整的 motif 就能正确分类,GraphChef 基于更简单的结构特征(如节点度数)即可达到正确预测,这反映了现有数据集设计中的已知缺陷。
image

剪枝效果实验

剪枝效果实验比较了剪枝前后的决策树规模和精度。树规模以所有决策树中决策节点的总数衡量。如表所示,剪枝大幅减少了决策树规模。在合成数据集上平均可以剪去约 62% 的节点而不损失精度,在真实数据集上甚至可以剪去约 84% 的节点。如果接受小幅精度下降,合成数据集和真实数据集分别可以剪去 68% 和 87% 的节点。在多种剪枝设置中,本文提出的同时使用训练精度和验证精度的方法表现最优。
image

recipe 分析实验

recipe 分析实验展示了如何阅读 GraphChef recipe 来获取数据集层面的分析结果。以 Reddit-Binary 数据集为例,如图和表所示,recipe 阅读采用类似动态规划的方式逐层理解:先理解第一层的类别状态含义,再基于这些理解来解释下一层。Reddit-Binary 的 recipe 揭示了一个 Q/A 图需要满足的条件:至少 15 个用户处于特定状态,这些用户要么是与至少 2 个核心用户互动的不活跃用户,要么是主要与不活跃用户互动的活跃或核心用户。这一 recipe 与已有研究中对 Q/A 图更"星状"结构的观察一致,并且进一步给出了量化标准:核心用户的度数阈值为 46 或更高,非核心通信最多容忍 2 个活跃用户。
image
image

在 MUTAG 数据集上,recipe 揭示了一个分子具有致突变性需要至少 12 个非氧原子和 8 个满足特定条件的原子,这些条件涉及与 NO2 基团的关联。在 BA-2Motifs 数据集上,GraphChef 仅学会了识别 house 节点,cycle 图通过偏置项被归为"非 house"类别,揭示了该数据集分类任务中偏置项可被利用的已知缺陷。在 Tree-Cycle 数据集上,GraphChef 仅通过节点度数检查即可定位 cycle 节点,无需识别完整的环结构。这些发现与文献中关于这些数据集已知缺陷的分析一致。

优点和创新点

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

  1. 与现有 GNN 解释方法聚焦单个图不同,GraphChef 的 recipe 同时解释数据集中所有图的分类逻辑,能得到数据集层面的可解释内容。
  2. 从 LOCAL 模型到 stone-age 模型的简化过程中得到启发设计了 dish 层,将连续向量变为类别状态,再将神经网络模块蒸馏为决策树。渐进式的模型设计使得能够在中间阶段(dish GNN)检查精度损失,定位可解释性代价的来源。
  3. 剪枝标准同时使用训练精度和验证精度作为约束,训练精度允许下降以去除过拟合节点,验证精度不允许下降以防止过度剪枝,实现了树规模控制。
posted @ 2026-09-10 01:40  乌漆WhiteMoon  阅读(10)  评论(0)    收藏  举报