Paper Reading:TREE-G: Decision Trees Contesting Graph Neural Networks
Paper Reading 是从个人角度进行的一些总结分享,受到个人关注点的侧重和实力所限,可能有理解不到位的地方。具体的细节还需要以原文的内容为准,博客中的图表若未另外说明则均来自原文。
| 论文概况 | 详细 |
|---|---|
| 标题 | 《TREE-G: Decision Trees Contesting Graph Neural Networks》 |
| 作者 | Maya Bechler-Speicher, Ran Globerson, Gal Gilad-Bachrach |
| 发表会议/期刊 | Thirty-Eighth AAAI Conference on Artificial Intelligence(AAAI) |
| 发表年份 | 2024 |
| 会议/期刊等级 | CCF-A |
| 论文代码 | github.com/mayabechlerspeicher/TREE-G |
作者单位:
- Blavatnik School of Computer Science, Tel-Aviv University
- Department of Bio-Medical Engineering and Edmond J. Safra Center for Bioinformatics, Tel-Aviv University
研究动机
决策树(Decision Trees, DTs)是表格数据领域的经典方法,以其高精度、易用性和可解释性著称。然而,当数据从表格形式转变为图结构时,决策树的优势能否延续,是一个值得深入探讨的问题。图结构数据在社交网络、分子生物学、药物发现等领域广泛存在。图数据同时具有表格数据的特征(如社交网络中节点的年龄、学历等属性)和拓扑结构信息。现有的图学习方法主要分为两类:
| 方法类别 | 核心机制 | 主要局限 |
|---|---|---|
| 图核方法 | 依赖预定义的嵌入和相似度计算 | 特征工程依赖领域知识,且难以捕捉节点特征与图结构的交互 |
| 图神经网络(GNN) | 通过迭代式邻居聚合学习节点表示 | 作为深度学习方法,在表格数据上并不总是优于树模型,且可解释性有限 |
将决策树适配到图数据上的现有尝试主要通过特征工程实现,即在训练或推理前用预定义函数计算图论特征(如度中心性、聚类系数等)并拼接到原始特征中。这种方式存在工作量大、需要领域知识,且无法捕捉节点特征与图结构之间的交互关系等局限。
文章贡献
针对决策树难以有效处理图结构数据的问题,本文提出了一种专为图数据设计的新型决策树 TREE-G。TREE-G 的核心设计是引入了一个新的分裂函数,该函数将节点特征沿图中的 walks 进行传播,同时通过一种新颖的指针机制动态生成候选顶点子集,使得分裂节点能够利用先前分裂中计算的信息来聚焦于图的关键子结构。TREE-G 保留标准决策树的贪心训练框架,可以直接嵌入梯度提升树(GBT)等集成方法中作为弱学习器使用。在理论层面,本文证明了 TREE-G 满足图标注任务中的排列不变性和顶点标注任务中的排列等变性,且其表达能力严格强于标准决策树,甚至可以表达 GNN 无法区分的分类规则。在实验方面,TREE-G 在 17 个图预测和顶点预测任务上超越其他树模型,在 10 个图分类任务中有 7 个优于 GNN,在 7 个顶点分类任务中有 4 个优于 GNN。此外,本文还提出了基于子集频率的可解释性机制,可以可视化模型决策过程中关注的顶点和边。
本文方法
TREE-G 处理的图由顶点集 \(V\)(大小为 \(n\))和邻接矩阵 \(A\)(\(n \times n\),可以有向或无向)定义。每个顶点关联一个实值特征向量,所有顶点的特征向量堆叠成矩阵 \(X\),其中 \(f_k\) 表示 \(X\) 的第 \(k\) 列,即第 \(k\) 个特征在所有顶点上的取值。\(A^d\) 的 \((i,j)\) 元素表示从顶点 \(i\) 到顶点 \(j\) 长度为 \(d\) 的 walk 数量。TREE-G 支持图标注(graph labeling,为整个图分配标签)和顶点标注(vertex labeling,为单个顶点分配标签)两类任务,每类均包含分类和回归。

分裂函数
标准决策树的分裂规则为 \(x_k > \theta\),即比较第 \(k\) 个特征值与阈值。TREE-G 将其替换为更复杂的分裂函数 \(\phi_{k,d,*,\rho,r}\),该函数整合了顶点特征、邻接矩阵和顶点子集信息。对于顶点标注任务,分裂函数定义为:
其中 \(k\) 是特征索引,\(d\) 是 walk 传播长度,\(*\)是指向祖先分裂节点的指针,\(\rho \in \{+,-\}\)指示使用祖先生成的哪个子集,\(r\) 是 walk 限制类型。\(M_r(S_{*,\rho})\) 是根据子集 \(S_{*,\rho}\) 和 walk 类型 \(r\) 构建的掩码矩阵,\(\circ\) 表示逐元素乘法。当 \(d=0\) 时 \(A^0 = I\),分裂函数退化为标准决策树的分裂规则。当 \(d=1\) 时,\(A^1 f_k\) 的第 \(i\) 个元素即为顶点 \(i\) 的邻居在第 \(k\) 个特征上的取值之和。对于有向图,可以使用 \(A^T\) 的幂来考虑反方向的 walk。
Walk限制类型
TREE-G 支持四种 walk 限制方式,通过掩码矩阵 \(M_r\) 实现:
| 类型 | 符号 | 描述 | 实现方式 |
|---|---|---|---|
| Source Walks | \(r=1\) | 仅考虑从子集 \(S\) 中顶点出发的 walk | 将 \(A^d\) 中不属于 \(S\) 的顶点对应的列置零 |
| Cycle Walks | \(r=2\) | 仅考虑从 \(S\) 中顶点出发且回到同一顶点的 walk | 仅保留 \(A^d\) 中 \(S\) 内顶点对应的主对角线元素 |
| Target Walks | \(r=3\) | 仅考虑终止于 \(S\) 中顶点的 walk | 将不属于 \(S\) 的顶点对应的行置零 |
| Target-Source Walks | \(r=4\) | 仅考虑从 \(S\) 中顶点出发且终止于 \(S\) 中顶点的 walk | 将不属于 \(S\) 的顶点对应的行和列都置零 |
在顶点标注任务中,由于只使用向量的第 \(i\)个 元素,Target 和 Target-Source 类型的掩码会导致冗余计算,因此不使用。在图标注任务中,由于需要对整个向量进行聚合,行掩码会影响聚合结果(如 min 聚合),因此聚合仅在选定子集的元素上进行。
子集生成机制
每个分裂节点 \(u\) 在执行分裂函数的同时,还会将当前使用的顶点子集 \(S_{*,\rho_u}\) 分割为两个新子集 \(S_{u,+}\) 和 \(S_{u,-}\),供树中的后代节点使用。具体定义如下:
子集的生成是针对每个图动态计算的,但使用相同的规则。子集定义对图同构不变,且不假设图的大小,因此 TREE-G 可以应用于训练中未见过的图大小。根节点指向自身,使用图中所有顶点 \(V\) 作为子集,指针参数 \(*\) 和方向参数 \(\rho\) 唯一确定了每个节点使用的子集。

如图所示,每个分裂节点中的虚线箭头指向其子集来源的祖先节点,\(\rho\) 值标注在箭头上。

图标注任务的聚合函数
在图标注任务中,分裂函数不再取向量的某个元素与阈值比较,而是对整个向量进行聚合以产生标量。聚合函数必须是排列不变的,以保证在图同构下结果一致。可选的聚合函数包括\(sum\)、\(mean\)、\(min\)、\(max\),由额外参数 AGG 指定。因此图标注任务的分裂函数为:
当聚合函数为 \(sum\) 时,子集的生成使用缩放后的阈值 \(\theta/|S_{*,\rho}|\),因为任何大于该缩放阈值的元素都会对总和产生贡献。
训练过程
TREE-G 保留了标准决策树的贪心训练过程。首先选择优化准则,如 Gini 系数或 \(L_2\) 损失,然后对每个叶节点进行网格搜索以找到能最大程度降低损失的最优分裂参数。将带来最大损失降低的叶节点转换为分裂节点,使用找到的最优参数。该过程重复直到满足停止条件,如树大小限制、叶节点最小样本数、最小增益等。与标准决策树的唯一区别在于网格搜索需要调优额外的参数,例如 \(k, d, *, \rho, r\) 及图标注任务中的 AGG。通过限制 walk 长度 \(d \leq 2\) 和祖先距离 \(a \leq 2\) 可以有效控制搜索空间,实验表明这已经足以实现高性能。

理论性质
本文从理论上对 TREE-G 进行了分析:
- 排列不变性/等变性(Lemma 4.1):TREE-G 在图标注任务中对顶点排列不变,在顶点标注任务中对顶点排列等变。证明思路是展示分裂函数的每个组件都是排列等变的(图标注中的聚合函数是排列不变的),阈值比较在顶点标注中保持等变性,在图标注中保持不变性。
- 计算复杂度(Lemma 4.2):搜索最优分裂参数的运行时间与特征数、最大walk长度和最大祖先距离成线性关系。每个分裂节点计算的动态特征数不超过 \(4 \times 4 \times (2^a + 1) \times d_{max} \times l\),其中 4 来自 walk 类型数,4 来自聚合函数数(图标注任务),\(2^a + 1\)来自子集数,\(d_{max}\)来自 walk 长度,\(l\) 来自特征数。
- 表达能力(Lemma 4.3):存在 TREE-G 可分离但禁用子集后无法分离的图。具体构造了位于 \(\{\pm1\} \times \{\pm1\}\) 平面上的两个 4 顶点图 \(G_1\) 和 \(G_2\),每个顶点有 \(x\) 和 \(y\) 坐标两个特征。由于两图拓扑同构,任何不变的图论特征取值相同,禁用子集的 TREE-G 无法区分。但使用子集后,根节点以 \(f_1\) 分裂生成子集 \(\{(1,1),(1,-1)\}\),随后在长度 2 的 walk 上传播 \(f_2\) 并施加该子集掩码,\(G_1\) 和 \(G_2\) 产生不同的结果向量,可以被 TREE-G 区分。
- 超越 1-WL(Lemma 4.4):存在 GNN 无法分离但 TREE-G 可以分离的图。GNN 的判别能力受 1-WL 测试限制,而 TREE-G 通过计算长度 3 的 cycle walk 数量(\(A^3\) 的对角线元素)可以区分两个 4-正则非同构图,它们的 3-cycle 数量不同。
实验结果
数据集和实验设置
本文使用到的数据集如下:
| 数据集 | 任务类型 | 规模 | 来源 |
|---|---|---|---|
| Mutag | 图分类(二分类) | 188 个图,平均 17.93 个顶点 | TUDatasets |
| NCI1 | 图分类(二分类) | 4110 个图,平均 29.87 个顶点 | TUDatasets |
| Proteins | 图分类(二分类) | 1113 个图,平均 39.06 个顶点 | TUDatasets |
| D&D | 图分类(二分类) | 1178 个图,平均 284.32 个顶点 | TUDatasets |
| Enzymes | 图分类(6 分类) | 600 个图,平均 32.63 个顶点 | TUDatasets |
| PTC | 图分类(二分类) | 344 个图,平均 14 个顶点 | TUDatasets |
| Mutagenicity | 图分类(二分类) | 4337 个图,平均 30.32 个顶点 | TUDatasets |
| IMDb-B | 图分类(二分类) | 1000 个图,平均 19 个顶点 | TUDatasets |
| IMDb-M | 图分类(3 分类) | 1500 个图,平均 13 个顶点 | TUDatasets |
| molHIV | 图分类(二分类) | 41127 个图,平均 25.5 个顶点 | OGB |
| Cora | 顶点分类(7 分类) | 2708 个顶点,10556 条边 | Planetoid |
| Citeseer | 顶点分类(6 分类) | 3327 个顶点,9104 条边 | Planetoid |
| Pubmed | 顶点分类(3 分类) | 19717 个顶点,88648 条边 | Planetoid |
| Arxiv | 顶点分类(40 分类) | 169343 个顶点,1166243 条边 | OGB |
| Cornell | 顶点分类(5 分类) | 183 个顶点,295 条边 | 异质图 |
| Actor | 顶点分类(5 分类) | 7600 个顶点,33544 条边 | 异质图 |
| Country | 顶点回归 | 3217 个顶点,12684 条边 | 选举网络 |
本文选择的对比方法有:
| 方法 | 类型 | 出处 |
|---|---|---|
| GCN | GNN | Kipf and Welling 2017 |
| GIN | GNN | Xu et al. 2019 |
| GAT | GNN | Veličković et al. 2018 |
| GATv2 | GNN | Brody et al. 2022 |
| GraphSAGE | GNN | Hamilton et al. 2018 |
| WL | 图核 | Shervashidze et al. 2011 |
| RW | 图核 | Gärtner et al. 2003 |
| BGNN | GNN+DT混合 | Ivanov and Prokhorenkova 2021 |
| XGraphBoost | GNN+DT混合 | Deng et al. 2021 |
| TREE-G (d=0, a=0) | 消融(禁用拓扑和子集) | 本文 |
| TREE-G (a=0) | 消融(禁用子集) | 本文 |
图预测任务采用 10 折交叉验证报告平均准确率和标准差,molHIV 数据集使用 OGB 官方预定义划分并报告 AUC。顶点预测任务使用预定义划分报告平均准确率和标准差,Country 任务为回归任务报告 RMSE。TREE-G 使用 GBT 框架,估计器数量为 {20, 50},学习率 0.1,最大深度 10,最大 walk 长度 \(d \in \{0,1,2\}\),最大祖先距离 \(a \in \{0,1,2\}\)。GNN 使用 ReLU 激活函数,{2,3,5} 层,{32,64} 隐藏维度,Adam 优化器训练 1000 轮。TREE-G 在 CPU 上运行,使用稀疏矩阵乘法提升大规模图的效率。
图分类对比实验
实验结果如表所示。TREE-G 在所有图分类任务上全面优于其他树模型方法,包括 XGraphBoost、BGNN 和 TREE-G 的消融变体。在 10 个图分类任务中有 9 个优于图核方法,有 7 个优于 GNN。在 IMDb-M 数据集上,TREE-G 相比 GNN 方法的领先幅度超过 6.4 个百分点。在 Mutag 数据集上领先超过 4.9 个百分点,在 Proteins 数据集上领先超过 2.4 个百分点。在 molHIV 大规模数据集上,TREE-G 达到 83.5% 的 AUC,优于所有 GNN 方法,其中 GATv2 为 81.9%,GIN 为 77.8%。值得注意的是,TREE-G 始终优于 TREE-G (a=0),这表明子集机制对性能提升很重要。实验结果说明,专门为图数据设计的纯决策树方法可以匹敌甚至超越主流 GNN 和图核方法。


顶点分类对比实验
实验结果如表所示。TREE-G 在 7 个顶点分类任务中有 4 个优于所有 GNN 方法,在其余任务上与领先的 GNN 方法持平。在 Cora 数据集上 TREE-G 达到 83.5% 的准确率,优于 GATv2 的 83.1% 和 GAT 的 83.0%。在 Actor 异质图数据集上 TREE-G 达到 37.0%,大幅领先 GATv2 的 34.2% 和 GAT 的 34.0%。在 Arxiv 大规模数据集上 TREE-G 达到74.7%,优于 GATv2 的 74.0% 和 GIN 的 73.8%。TREE-G 同样始终优于其消融变体 TREE-G (a=0) 和 TREE-G (d=0, a=0),进一步验证了拓扑传播和子集机制的有效性。

Walk 类型消融实验
如图所示,本文设计了四个合成任务来验证四种 walk 类型各自的必要性。每个任务对应一种特定的 walk 计数模式,从红色顶点出发、在红色顶点间循环、终止于红色顶点、红色顶点间的 walk。实验中对每个任务分别使用单一 walk 类型和排除最优类型后的三种 walk 类型进行对比。结果表明每种 walk 类型在其对应的任务上都优于其他类型,说明四种 walk 类型并非冗余,各自有独特的表达能力。

可解释性分析
TREE-G 的可解释性机制基于顶点在选定子集中的出现频率。对于图 \(G\) 在树 \(T\) 中的预测路径,统计每个顶点 \(i\) 在该路径上的选定子集中出现的次数 \(n_T(i)\),然后转换为排名 \(r_T(i)\),排名越高表示越重要。集成模型中顶点 \(i\) 的重要性得分为所有树中排名的加权平均,权重为该树预测值 \(|y_T|\):
重要性得分非负且归一化为 1。如图所示,在 Red Isolated Vertex 合成任务中,TREE-G 的注意力集中在孤立顶点上,尤其是红色孤立顶点。在 Mutagenicity 任务中,TREE-G 关注 \(NO_2\) 基团和碳环结构,这些正是已知的致突变子结构。边级别的重要性通过统计预测路径上使用该边的分裂节点数量来计算。

优点和讨论
个人认为,本文有如下一些优点和创新点可供参考学习:
- 通过指针 \(*\) 和方向 \(\rho\) 参数,让分裂节点可以引用祖先节点生成的子集,实现了在决策树框架内动态学习和复用图子结构的能力。子集由树的结构和数据共同决定,无需预先指定。
- TREE-G 仅修改了分裂函数,保留了贪心训练、剪枝、集成等全部标准决策树机制,可以直接作为 GBT 的弱学习器使用,使得决策树在表格数据上的成功能够迁移到图数据。
- 本文不仅给出了排列不变性/等变性和计算复杂度的证明,还通过 Lemma 4.3 和 Lemma 4.4 分别证明了 TREE-G 强于标准决策树且不受 1-WL 测试限制。模型还具有可解释性,在没有注入领域知识的前提下发现了与领域知识一致的子结构。
目前图学习主要还是利用以消息传递范式为基础的 GNN 和图 transformer 作为主干网络,决策树类的方法很少被考虑。这是因为图数据结构有实例间的关联信息,难以抽象成单个实例以适配决策树的输入和分裂函数。这篇论文把整张图和这个顶点的信息一起输入,规避单个样本如何定义的问题,并设计了针对图的分裂函数,这是一个很有创意的做法。不过目前这方面的工作由于需要输入完整的图,还难以回答应用到动态图时的适配问题。但是如果在图学习问题上需要使用到树的一些性质,可以参考这篇论文的建模思路。

浙公网安备 33010602011771号