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/ |
作者单位:
- ETH Zurich, Switzerland(苏黎世联邦理工学院)
研究动机
图(Graph)在化学、工程、社会科学和交通等诸多领域中被用来表示复杂的关联数据,图神经网络(GNN)作为处理图数据的流行模型。它在图分类任务中取得了广泛应用,例如判断一个以图表示的蛋白质是否为酶。然而 GNN 本质上是一个黑盒模型,其决策过程不透明,用户无法理解模型为何做出某个预测。现有 GNN 解释方法大致可分为五类:梯度方法、互信息方法、反事实方法、子图方法和示例方法。前三种方法计算节点或边级别的重要性分数,生成热力图来突出输入中哪些部分重要。子图方法通过展示重要子图或原型来解释。示例方法则给出相似的参考图。这些方法存在两个核心局限:
- 只能回答"哪些输入重要",无法回答"如何使用这些输入"。例如在 PROTEINS 数据集中,现有方法可以指出两个 Sheet 节点是重要的,但无法回答"Sheet 之间是否需要连接""是否必须恰好有两个 Sheet"等问题。
- 只能针对单个图进行解释,无法给出数据集层面的整体规律。用户需要逐一查看数十个图的局部解释,难以从中归纳出数据集层面的分类依据。

本文希望将解释推进一步,不仅理解哪些输入重要,还要理解它们如何被使用;不仅解释单个图,还要理解整个数据集,回答"什么使一个蛋白质成为酶"这样的问题。
文章贡献
针对 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\) 是节点 \(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\) 变成 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 层中节点的新状态由三部分信息决定:上一层的类别状态(绿色)、每个状态类别中的邻居数量(蓝色)、以及状态之间的二值比较(黄色)。

从 dish 到 GraphChef
由于 dish GNN 中所有状态都是离散的 one-hot 类别值,可以将训练好的神经网络模块蒸馏为决策树。需要蒸馏的部分包括 dish 层的更新函数、编码器和解码器,由于状态是类别的,蒸馏就变成了一个分类问题,即决策树学习预测类别状态。蒸馏后的 GraphChef 层如下所示:
GraphChef 层仍然遵循消息传递框架,只是将神经网络加 Gumbel-Softmax 替换为决策树。为了让决策树更小,本文还引入了成对 delta 特征 \(\Delta\)。决策树通常不擅长比较两个特征的大小关系,delta 特征专门解决这个问题。令 \(c_i^l\) 为状态 \(i\) 的邻居计数,delta 特征对所有有序状态对计算二值比较:
决策树可以访问三组特征:节点自身的上一状态、每个状态类别的邻居计数、以及两个状态类别计数之间的比较。如图所示,决策树的分支节点可以基于这三类特征进行判断:(a) 判断节点当前处于哪个状态;(b) 判断节点在某个状态中是否有特定数量的邻居;(c) 判断节点在一个状态中的邻居是否多于另一个状态。每个决策节点的分支类型可以直接解释为 GNN 推理的一个步骤,将所有节点的解释组合起来就构成了 GraphChef 的 recipe。

剪枝 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 也明显优于仅使用度特征的决策树基线。模型简化,即从连续向量到类别状态再到决策树没有带来精度下降。值得注意的是,剪枝在某些数据集上甚至提高了测试精度,这可能是由于剪枝过程引入的正则化效果。

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

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

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

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

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


在 MUTAG 数据集上,recipe 揭示了一个分子具有致突变性需要至少 12 个非氧原子和 8 个满足特定条件的原子,这些条件涉及与 NO2 基团的关联。在 BA-2Motifs 数据集上,GraphChef 仅学会了识别 house 节点,cycle 图通过偏置项被归为"非 house"类别,揭示了该数据集分类任务中偏置项可被利用的已知缺陷。在 Tree-Cycle 数据集上,GraphChef 仅通过节点度数检查即可定位 cycle 节点,无需识别完整的环结构。这些发现与文献中关于这些数据集已知缺陷的分析一致。
优点和创新点
个人认为,本文有如下一些优点和创新点可供参考学习:
- 与现有 GNN 解释方法聚焦单个图不同,GraphChef 的 recipe 同时解释数据集中所有图的分类逻辑,能得到数据集层面的可解释内容。
- 从 LOCAL 模型到 stone-age 模型的简化过程中得到启发设计了 dish 层,将连续向量变为类别状态,再将神经网络模块蒸馏为决策树。渐进式的模型设计使得能够在中间阶段(dish GNN)检查精度损失,定位可解释性代价的来源。
- 剪枝标准同时使用训练精度和验证精度作为约束,训练精度允许下降以去除过拟合节点,验证精度不允许下降以防止过度剪枝,实现了树规模控制。

浙公网安备 33010602011771号