Paper Reading: Graph Neural Network contextual embedding for Deep Learning on tabular data Graph Neural Network contextual embedding for Deep Learning on tabular data


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

论文概况

论文概况 详细
标题 《Graph Neural Network contextual embedding for Deep Learning on tabular data Graph Neural Network contextual embedding for Deep Learning on tabular data 》
作者 Mario Villaizán-Vallelado, Matteo Salvatori, Belén Carro, Antonio Javier Sanchez-Esguevillas
发表期刊 Neural Networks
发表年份 2024
期刊等级 新锐期刊分区表(2026 年 3 月)2 区 TOP,CCF-B
论文代码 https://github.com/MatteoSalvatori/INCE

作者单位:

  1. Artificial Intelligence Laboratory (AI-Lab), Telefonica I+D, Spain
  2. Universidad de Valladolid, Valladolid, 47011, Spain

研究动机

许多现实世界的应用(如医疗、金融、推荐系统)将大数据存储在表格形式中,每条记录(行)由一组异构的连续型和分类型特征(列)组成。因此,有效地从表格数据中学习具有极高的实际价值。深度学习在涉及人类技能的领域(如自然语言处理、计算机视觉)取得了突破性进展,激发了将其应用于表格数据的兴趣。然而,这种应用更具挑战性。因为与文本、图像等同质且具有固有结构的数据不同,表格数据具有两个阻碍深度学习性能的特性:

表格数据特性 说明
特征异构性 混合了各种统计分布(连续、分类)的特征,且特征间关系复杂
排列不变性 表格行的意义不依赖于列的顺序,这与句子中词序的重要性或图像中像素的空间相关性形成鲜明对比

基于树的集成模型(如 XGBoost, LightGBM, CatBoost)通常在表格数据上达到最优性能,它们预测准确、训练快速。尽管深度学习模型被广泛研究,但在与树模型的直接比较中往往表现不佳。然而,继续开发表格数据的深度学习模型是有充分动机的,因为树模型在持续学习、强化学习、或处理多模态数据(表格与图像、文本、音频相结合)等场景中存在局限性。基于上述背景,这篇论文旨在解决一个核心问题:如何设计一种更有效的深度学习模型,以提升其在监督学习任务中对表格数据的处理性能

文章贡献

本文提出了一个名为 INCE 的表格深度学习模型,其核心创新在于用 GNN 实现交互网络,用作表格特征的上下文嵌入器。模型将每个表格行转换为一个全连接图,其中每个节点对应一个特征的嵌入表示,并引入一个虚拟 CLS 节点。随后,一个 IN 层堆栈通过消息传递机制,动态地学习和更新图中所有节点(特征)之间连接的强度,从而建模复杂的特征交互。经过 IN 增强后的 CLS节点的最终表示,被用作整个表格行的上下文感知的全局嵌入,并送入一个 MLP 解码器进行预测。通过在公开数据集上与多个 baselines 的广泛对比实验,证明了 INCE 模型在深度学习模型中达到了最佳平均性能,并能与树模型形成有力竞争。

本文方法

本文提出的模型采用编码器-解码器的框架:

组件 功能
编码器 将原始表格特征映射到隐空间向量(嵌入),包含列嵌入上下文嵌入两个部分
解码器 一个针对具体任务(分类/回归)调整的 MLP,接收上下文嵌入并做出预测

screenshot-1777013790479

问题定义

模型聚焦于监督学习任务,给定一个表格数据集 \(D = \{x_i^{j_c}, x_i^{j_n}, y_i\}_{i=1}^N\),其中:

符号 含义
\(x_i^{j_n}\) 数值型特征集合 (\(j_n \in [1, M_{\text{num}}]\))
\(x_i^{j_c}\) 分类型特征集合 (\(j_c \in [1, M_{\text{cat}}]\))
\(y_i\) 标签
\(N\) 总行数
\(M = M_{\text{num}} + M_{\text{cat}}\) 总特征数

编码器

列嵌入

列嵌入将所有异构特征投影到同一个稠密的 \(l\) 维隐空间中:
screenshot-1777013877267
对于连续特征:

\[c_{i}^{j_{n}}=\operatorname{ReLU}\left(b^{j_{n}}+x_{i}^{j_{n}}\cdot W_{\text{num}}^{j_{n}}\right)\quad (W_{\text{num}}^{j_{n}}\in R^{l}) \]

对于分类特征:

\[c_{i}^{j_{c}}=b^{j_{c}}+h_{j_{c}}^{T} W_{cat}^{j_{c}}\quad (W_{cat}^{j_{c}}\in R^{\left|j_{c}\right|\times l}) \]

其中,涉及到的符号及其含义如下:

符号 含义
\(b\) 偏置
\(W\) 可学习权重/查找表
\(h\) one-hot 向量

上下文嵌入

为了解决列嵌入无法捕捉特征间关系的问题,引入了基于交互网络的上下文嵌入。
screenshot-1777014178454
首先进行全连接图构建,图架构的组件设置如下:

图组件 设置
节点 每个列嵌入 \(c_j\) 对应一个图节点 \(n_j\)
为每一对节点 \((n_{j1}, n_{j2})\) 创建两条独立的、有向的边(\(e_{j_1j_2}\)\(e_{j_2j_1}\)),形成一个全连接图。
虚拟节点 模仿 BERT,添加一个虚拟的 CLS 节点,与所有特征节点相连。其初始表示是一个可学习的参数向量。

注意,节点不使用位置编码,因为单独的列嵌入已足够区分它们。边的初始表示为空。接着将一个 IN 层堆栈作用于这个全连接图,以建模节点(特征)之间的交互并增强其表示。经过 IN 堆栈更新后,最终的 CLS 节点表示被用作整个表格行的全局上下文嵌入,传递给解码器。

交互网络

一个标准 IN 层的工作流程如下图所示,包含两步更新和残差连接:
screenshot-1777014248150
第一步进行更新边表示,每条边 \(e_{j_1 \rightarrow j_2}\) 的表示基于其连接的两个节点来更新:

\[e_{j_{1}\rightarrow j_{2}}^{\prime}=\operatorname{MLP}_{E}\left(\operatorname{Concat}\left(n_{j_{1}}, n_{j_{2}}, e_{j_{1}\rightarrow j_{2}}\right)\right) \]

其中,\(MLP_E\) 是一个在所有边上共享的神经网络。第二步更新节点表示,每个节点 \(n_j\) 聚合所有指向它的边的更新信息,并更新自身:

\[n_{j}^{\prime}=\operatorname{MLP}_{N}\left(\operatorname{Concat}\left(n_{j},\sum_{k\in\mathcal{N}} e_{k\rightarrow j}\right)\right) \]

其中,\(\mathcal{N}\) 是节点 \(n_j\) 的邻居集合,\(MLP_N\) 是一个在所有节点上共享的神经网络。第三步为残差连接,最终节点和边的表示是更新后的值与原始值的和:

\[\begin{align*} n_{j} &= n_{j}^{\prime} + n_{j} \\ e_{j_{1}\rightarrow j_{2}} &= e_{j_{1}\rightarrow j_{2}}^{\prime} + e_{j_{1}\rightarrow j_{2}} \end{align*} \]

解码器 \(MLP_{DEC}\) 接收最终的 CLS 上下文嵌入,其输出层的大小和激活函数根据具体任务(分类或回归)进行调整。

交互网络分析

可训练参数分析

一个由 \(n\) 层 IN 组成的堆栈,其可训练参数量 \(\mathcal{TP}(IN)\) 由各层的 \(MLP_E\)\(MLP_N\) 参数量之和决定。具体为:

\[\begin{aligned} \mathcal{TP}(IN) &= \sum_{i=1}^n [\mathcal{TP}(MLP_E^i) + \mathcal{TP}(MLP_N^i)] \\ \mathcal{TP}(MLP_N^i) &= (2 \cdot l^2 + l) + (d-1) \cdot (l^2 + l) \\ \mathcal{TP}(MLP_E^i) &= (K_i \cdot l^2 + l) + (d-1) \cdot (l^2 + l) \\ \text{其中,} K_i &= 2 \text{ 若 } i=1,\quad K_i=3 \text{ 若 } i>1 \end{aligned} \]

其中,涉及到的符号及其含义如下:

符号 含义
\(l\) 隐空间大小
\(d\) \(MLP_{E, N}\) 的深度
\(n\) 堆叠的 IN 层数

\(K_i\) 的不同是因为第一层 IN 的 \(MLP_E\) 输入是节点对 \((n_i, n_j)\),而后续层会额外接收前一层更新的边特征 \(e_{j_1 \rightarrow j_2}\) 作为输入。模型参数量随隐空间大小 \(l\) 二次方增长,随 IN 层数 \(n\) 和 MLP 深度 \(d\) 线性增长。增加 \(n\) 比增加 \(d\) 带来的参数量增长斜率更陡。
screenshot-1777015378316

对比 Transformer

Transformer 编码器也是基于上下文嵌入机制的模块,两种模块进行了深入的对比分析,都将列嵌入组织为带 CLS 节点的全连接图。两者都通过一种机制(注意力/IN 卷积)建模节点间交互。该机制能学习交互的强度,充当一种软链接剪枝。核心差异在于如何建模交互,Transformer 注意力机制(单头)为:

\[n_{i}^{\prime\alpha}=\sum_{j=1}^{M}\sum_{\beta=1}^{l}\omega_{i, j} V^{\alpha,\beta} n_{j}^{\beta} \]

其中,注意力权重 \(\omega_{i, j}\) 是拓扑空间(节点间)的非平凡算子,但在隐空间是对角的。而值矩阵 \(V^{\alpha,\beta}\) 在拓扑空间是对角的,在隐空间是非平凡的。IN 的更新机制则更为通用,\(MLP_N\) 的作用类似 \(V\),而 \(MLP_E\) 的作用类似 \(\omega\)。区别在于 \(MLP_E\) 是在拓扑空间和隐空间上均为非平凡的算子。这意味着对于一对节点 \((n_i, n_j)\)\(MLP_E\) 可以为隐空间的不同维度 \(\alpha, \beta\) 学习不同的交互强度。而 Transformer 的注意力权重 \(\omega_{i,j}\) 对整个隐空间是共享的标量。

两者在特征数增加时都面临挑战:标准多头自注意力和全连接图上的IN都具有相对于特征数的二次方复杂度。可能的解决方案包括使用高效注意力近似或构建更稀疏的图拓扑。

特征重要性

Transformer 模型可以通过 CLS 节点的注意力图来评估特征重要性。对于 IN,由于其学习的特征-特征交互是 \(l\) 维向量而非标量,论文提出了一种新的方法来评估特征的全局重要性。步骤如下:

  1. 训练与推断:在训练好的 INCE 模型上,对测试集样本进行推断,获取最后一层 IN 输出的所有特征-特征交互向量 \(\{ e_{j_1 \rightarrow j_2}^i \} \in R^l\),其中 \(i\) 是样本索引。
  2. 计算总体分布:计算整个测试集上所有交互向量的均值 \(\mu\) 和协方差矩阵 \(S\),以此代表交互的背景分布。
  3. 计算特征对重要性:对于每一对特定的特征 \((j_1, j_2)\),计算其所有样本交互向量 \(\{ e_{j_1 \rightarrow j_2}^i \}\) 的均值 \(\mu_{j_1, j_2}\)。然后,计算该均值向量与背景分布之间的马氏距离,作为这对特征交互强度的量化指标:

    \[I_{j_1, j_2} = \sqrt{(\mu_{j_1, j_2} - \mu)^T S^{-1} (\mu_{j_1, j_2} - \mu)} \]

    距离越大,表明这对特征之间的交互模式与随机背景差异越大,即它们的交互越强、越特殊。
  4. 聚合为特征重要性:将一个特征 \(j\) 的全局重要性 \(I_j\) 定义为其所有出边交互强度的总和,即它影响其他所有特征的总强度:

    \[I_j = \sum_{k \neq j} I_{j, k} \]

实验结果

数据集和实验设置

使用了 7 个公共表格数据集,并新增了 2 个具有大量特征的数据集,以探究模型在不同特征规模下的表现。数据集信息见下表。对数值特征进行零均值、单位方差标准化,对分类特征进行序数编码,缺失值用零填充。对每个数据集,使用 Optuna 库进行 50 次迭代的贝叶斯优化,每种配置进行 5 折交叉验证。
image

INCE 模型与以下模型进行了对比:

算法类型 对比模型
标准(树)模型 XGBoost, LightGBM, CatBoost
深度学习模型 MLP, DeepFM, TabTransformer, SAINT, FT-Transformer

为了强调上下文嵌入选择的关键性,论文构建了多个 INCE 变体,它们共享相同的列嵌入和解码器,但替换了上下文嵌入架构:、

INCE 变体 说明
INCE-GCN 使用图卷积网络
INCE-GAT 使用图注意力网络
INCE-Transformer 使用 Transformer 编码器
INCE 使用交互网络

对比实验

在 7 个数据集中的 6 个上,INCE 的表现超越了所有其他深度学习基线模型。在第七个数据集(HIGGS)上,与 SAINT 模型表现相当,但仍显著优于其他深度模型。这表明基于 IN 的上下文嵌入是当前表格数据深度学习中的强有力方法。INCE 在 2 个数据集上表现优于树模型,在其他 5 个数据集上取得了与树模型具有竞争力的结果。这表明 INCE 显著缩小了深度学习模型与树模型在表格数据上的性能差距。
image

在上下文嵌入机制方面,所有上下文嵌入都优于 MLP,证明了引入特征交互建模的有效性。GAT、Transformer 和 IN 优于 GCN。在全连接图中,GCN 平等对待所有邻居,而 GAT/Transformer/IN 能动态学习每条边的权重,区分不同邻居对目标节点更新的贡献重要性。IN 和 Transformer 取得了最佳结果,且 IN 在本基准测试中略胜一筹。这表明,在能学习动态边权的架构中,增加边权重学习机制的复杂性(IN 的 \(MLP_E\) 比 Transformer 的注意力机制更灵活)能带来进一步的性能提升。

超参数实验

实验表明,隐空间大小 \(l\) 需要针对每个数据集进行调优,但 \(MLP\) 深度 \(d\) 和 IN 层数 \(n\) 的影响在不同任务中呈现一致模式。配置 \(d=3\)\(n=2\) 是一个稳健的基线,在不同数据集和任务上都能取得良好性能。\(MLP\) 深度 \(d\) 对模型性能影响最大,但对参数量的影响相对较小。贝叶斯优化器会快速将 \(d\) 的搜索空间缩小到 \([3,4]\)。IN 层数 \(n\) 的影响次之,优化器通常在 \([2,3]\) 的范围内进行微调。由于图节点数(即特征数)较少(8 到 129 个),且图为全连接拓扑,两层 IN 足以让节点信息传播到全图。堆叠更多层对性能提升有限,且可能因数据集规模有限而导致过拟合。
image

计算时间

下图展示了不同因素对模型训练时间的影响,运行是以批次大小 256 的平均时间计算,并相对于无上下文嵌入的 MLP 进行归一化。计算时间随表格特征数量的增加而增加,在 INCE 的全连接图设定中,操作量随节点数(特征数)二次方增长。IN 层数 \(n\) 对计算成本影响最大,当特征数约 20 时,隐空间大小 \(l\) 的影响与 \(MLP\) 深度 \(d\) 的影响相当甚至更大。
image

对比 Transformer

从模型性能来看,INCE(IN)和 INCE-Transformer 表现相当,但 IN 在所选基准上略胜一筹,可能是 IN 更通用的交互建模机制。参数量方面,在相同 \(l, n\) 下 IN 的参数量少于 Transformer。这主要是由于 Transformer 中庞大的前馈网络层,当注意力头数 \(h > 2\) 时,即使不考虑前馈网络,Transformer 的参数量也可能超过 IN。
image

可解释性

为了直观展示上下文嵌入如何改进特征表示,论文以泰坦尼克号数据集为例进行了可视化分析。使用一个配置为 \(l=128, d=3, n=2\) 的 INCE 模型,对列嵌入和上下文嵌入的输出分别进行降维(例如 t-SNE 或 PCA)并可视化。列嵌入如左图所示,列嵌入是上下文无关的。例如,无论其他特征(如 pclass, age, family_size)的值如何,分类特征 title = Mrs 总是被投影到隐空间中的同一个固定点,这是一个静态的、无交互的表示。
上下文嵌入如右图所示,上下文嵌入是上下文感知的。这里展示的是最后一个 IN 层中,每个节点(特征)发送给 CLS 节点的消息向量的降维结果。可见同一分类特征(如 title)的嵌入不再集中于一个点,而是根据样本的不同上下文分散在不同的位置。这表明 IN 根据特征间的交互,为同一特征在不同样本中生成了动态的、个性化的表示。
image

论文将计算出的特征重要性排序,与通过 SHAP 计算出的特征重要性排序进行比较。在多个数据集上的实验表明,两者得出的重要性排序具有高度一致性(斯皮尔曼等级相关系数 > 0.8)。

优点和创新点

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

  1. 将表格数据行构建为全连接图,将每个特征视为节点,并引入一个虚拟 CLS 节点,从而利用图神经网络的排列不变性天然适应表格特征顺序无关的特性,提供了更优的数据结构表示。
  2. 引入在物理模拟中表现卓越的 GNN 交互网络作为核心上下文嵌入器,先更新边再聚合消息更新节点的两步消息传递机制,能够比标准 Transformer 注意力更精细地建模特征间复杂的、隐空间感知的交互强度。
  3. 通过上下文嵌入,将传统的、静态的列嵌入转变为动态的、上下文感知的特征表示,使得同一特征在不同样本中(因其他特征值不同)能够获得不同的嵌入向量,从而显著提升了模型的表征能力和判别精度。
posted @ 2026-04-30 15:21  乌漆WhiteMoon  阅读(24)  评论(0)    收藏  举报