Paper Reading: Cross-Feature Interactive Tabular Data Modeling With Multiplex Graph Neural Networks
Paper Reading 是从个人角度进行的一些总结分享,受到个人关注点的侧重和实力所限,可能有理解不到位的地方。具体的细节还需要以原文的内容为准,博客中的图表若未另外说明则均来自原文。
| 论文概况 | 详细 |
|---|---|
| 标题 | 《Cross-Feature Interactive Tabular Data Modeling With Multiplex Graph Neural Networks》 |
| 作者 | Mang Ye, Yi Yu, Ziqin Shen, Wei Yu, Qingyan Zeng |
| 发表期刊 | IEEE Transactions on Knowledge and Data Engineering(TKDE) |
| 发表年份 | 2024 |
| 期刊等级 | 新锐期刊分区表(2026 年 3 月)2 区 TOP,CCF-A |
| 论文代码 | 文中未公开 |
作者单位:
- The National Engineering Research Center for Multimedia Software, School of Computer Science, Wuhan University, Wuhan, Hubei 430072, China.
- The Taikang Center for Life and Medical Sciences, Wuhan University, Wuhan, Hubei 430072, China.
- Aier Eye Hospital, Wuhan University, Wuhan, Hubei 430072, China.
研究动机
表格数据是现实世界应用中最常见的数据形式,蕴含着大量对预测分析至关重要的信息,如推荐系统、在线广告、医疗诊断、欺诈检测。DNNs(如CNN、RNN、Transformer)在计算机视觉和自然语言处理等结构化数据领域取得了巨大成功,这激发了研究者将其应用于表格数据建模的兴趣。与图像、文本等数据不同,表格数据具有两个的特性:
| 表格数据特性 | 说明 |
|---|---|
| 排列不变性 | 表格中特征列的顺序排列不影响样本的标签,与图像像素或文本单词的顺序重要性有所不同 |
| 局部依赖性 | 预测标签往往只依赖于一小部分相关的特征,而不是所有特征,即特征间的关系是局部且异构的 |
当前用于表格数据的深度学习方法(如混合模型、基于 Transformer 的模型)存在明显局限:它们要么强调全局特征交互而忽视了局部依赖性,要么只关注局部交互而忽略了排列不变性。许多方法难以有效建模异构特征之间的复杂交互关系,或者会引入冗余计算和高复杂度。

基于上述背景,论文旨在解决将深度神经网络,特别是图神经网络,应用于表格数据建模时所面临的两个挑战:
- 如何有效建模异构特征间的交互:传统方法(如简单拼接或直接特征交互)会导致维度灾难或高计算成本,且不符合人类直觉。例如,判断肥胖程度需同时考虑身高和体重,而非单一特征。
- 如何为特征交互构建合适的图结构:图结构是 GNNs 有效工作的核心,现有构建方法各有缺陷:手工构建的连通性(如全连接图)可能导致冗余连接或信息缺失,结构固定不灵活。自动学习的连通性缺乏人类先验知识,可能导致学习效率低或可解释性差。
文章贡献
针对表格数据的排列不变性和局部依赖性两大特性,以及现有深度方法在建模特征交互和构建图结构方面的不足,本文提出了一种名为多路交叉特征交互网络的深度学习模型。MPCFIN 将表格数据的特征建模为图结构,其中每个特征列对应一个图节点。模型的自适应交叉特征嵌入模块通过学习特征间的相关性矩阵,动态地为每个特征生成融合了其最相关特征的交叉特征表示,从而有效捕获异构特征交互。多路图神经网络模块则构建了一个融合手工规则(CLS 连接与最小连通子图)与数据驱动学习(基于度量的图结构学习)的双路图,并利用交互网络进行消息传递,以同时建模特征间先验的与潜在的复杂关系。最终模型通过聚合双路图中的全局信息进行预测,在提升性能的同时保证了良好的可解释性。通过在 7 个数据集上进行的大量实验,本文从性能对比、消融分析、复杂度评估和可解释性可视化等多个方面,全面验证了 MPCFIN 在预测准确性、模型可靠性和逻辑合理性上均优于现有的多种深度神经网络基线模型。
预备知识
图神经网络的消息传递机制是本文的理论基础,该框架抽象了不同 GNN 模型的共性,包含三个函数:
| 消息传递机制组件 | 符号 | 功能 |
|---|---|---|
| 消息函数 | \(\phi\) | 定义如何根据源节点和目标节点的特征生成消息 |
| 聚合函数 | \(\oplus\) | 定义如何聚合来自邻居节点的消息,常用操作有求和、均值、最大值、注意力 |
| 更新函数 | \(\gamma\) | 定义如何结合节点自身特征和聚合后的消息来更新节点表示 |
在第 \(t\) 次迭代时,消息传递形式化表示为:
其中,涉及到的符号及其含义为:
| 符号 | 含义 |
|---|---|
| \(h_i^{j(t)}\) | 表示节点 \(h_i^j\) 在第 \(t\) 次迭代后的表示 |
| \(h_i^{j(0)}\) | 通过交叉特征嵌入得到的初始特征 |
经过 \(T\) 次迭代后,提取 CLS 节点的表示向量 \(h_i^{0(T)}\) 用于最终的预测任务。
本文方法
问题公式化
一个表格数据集被定义为 \(\mathcal{D} = \{x_i^j, y_i\}_{i \in [1,N], j \in [1,M]}\),其中:
| 符号 | 含义 |
|---|---|
| \(x_i^j\) | 第 \(i\) 行(样本)的第 \(j\) 列(特征)的值 |
| \(y_i\) | 第 \(i\) 行样本的标签,回归任务为连续值,分类任务为离散值 |
| \(N\) | 样本总数 |
| \(M\) | 特征总数 |
| \(\mathcal{X} = \{x_i | i \in [1, N]\}\) | 特征集合 |
| \(\mathcal{Y} = \{y_i | i \in [1, N]\}\) | 标签集合 |
| \(\mathcal{S}_{num}\) | 数值型特征的索引集合 |
| \(\mathcal{S}_{cat}\) | 分类型特征的索引集合 |
目标是学习一个函数 \(\hat{y_i} = f(x_i)\),以最小化预测误差:$ \min \sum_{i=1}^{N} \mathcal{L}(f(x_i), y_i) $。其中 \(\mathcal{L}\) 是损失函数,如回归用均方误差 MSE,分类用交叉熵损失。
MPCFIN 框架
论文提出多路交叉特征交互网络来对表格数据的每个样本行进行建模。图构建是其核心思想,对于一个样本 \((x_i, y_i)\):
- 节点生成:每个特征 \(x_i^j\) 通过交叉特征嵌入模块映射为潜在向量 \(h_i^j \in \mathbb{R}^d\),作为一个节点。
- 添加 CLS 节点:引入一个额外的 CLS 节点 \(h_i^0 \in \mathbb{R}^d\),用于聚合全局信息。
- 图定义:由此构成一个图结构 \(\mathcal{G}_i = \{\mathcal{V}, \mathcal{A}\}\)。
其中图结构的相关定义如下:
| 符号 | 含义 |
|---|---|
| 节点集 | \(\mathcal{V} = \{v_0, v_1, ..., v_M\} = \{h_i^0, h_i^1, ..., h_i^M\}\),其中 \(v_0\) 是 CLS 节点 |
| 邻接矩阵 | \(\mathcal{A} \in \{0,1\}^{(M+1)\times(M+1)}\),被所有样本共享,定义了节点间的连接关系 |
| 边集 | \(\mathcal{E} = \{e(v_i, v_j) | v_i, v_j \in \mathcal{V} \wedge \mathcal{A}_{ij}=1\}\) |
| 邻居节点集 | \(\mathcal{N}(i) = \{j | v_j \in \mathcal{V} \wedge \mathcal{A}_{ij}=1\}\) |
MPCFIN 模型框架如下,具体架构包含交叉特征嵌入模块和多路 GNN 模块。

交叉特征嵌入模块
交叉特征嵌入模块的目标将原始异构特征 \(x_i^j\) 融合交叉特征信息,并投影到统一、稠密的潜在空间,得到节点嵌入 \(h_i^j \in \mathbb{R}^d\),核心组件是自适应特征组合器。该组件第一步是进行特征统一化,对数值型和分类型特征分别进行标准化和离散化处理:
其中,涉及到的符号及其含义为:
| 符号 | 含义 |
|---|---|
| \(f_n\) | 标准缩放器 |
| \(f_d\) | 标签编码器 |
接着第二步是学习特征相关性矩阵,初始化可学习的列嵌入 \(E_{col} = (c_1, c_2, ..., c_M) \in \mathbb{R}^{d \times M}\) 来表示每个特征列。接着计算可能性得分矩阵 \(P \in \mathbb{R}^{M \times M}\),其元素 \(P[i,j]\) 表示特征 \(i\) 和 \(j\) 的关联概率,通过列嵌入的余弦相似度计算:
通过激活函数 \(\sigma\) 和阈值 \(T\) 得到相关性邻接矩阵 \(A_c\):
其中,涉及到的符号及其含义为:
| 符号 | 含义 |
|---|---|
| \(b\) | 可学习的偏置 |
| \(A_c[i,j]=1\) | 表示第 \(j\) 个特征与第 \(i\) 个特征最相关 |
第三步是生成交叉特征,根据 \(A_c\),为每个特征 \(x_i^j\) 找到其最相关的特征集合,并将其值拼接,形成交叉特征 \(x_i^{\prime j}\):
其中 \(x_c^j\) 是相关特征的集合。最后投影为节点嵌入,将每个交叉特征 \(x_i^{\prime j}\) 通过一个独立的 MLP(投影器)映射到 \(d\) 维空间,得到最终的节点嵌入向量:
这里为每个特征使用非共享的 MLP 参数 (\(b^j\), \(W^j\)),以获得更具表达力的表示。
多路图神经网络模块
多路图神经网络模块用于构建双图结构,分别进行消息传递,并融合结果进行预测。首先进行构建多路图结构,为节点集合 \(\mathcal{V} = \{h_i^0, h_i^1, ..., h_i^M\}\)(已加入CLS节点 \(h_i^0\))构建两个并行的图。第一个图是手工构建图 \(\mathcal{G}_i^{hcs}\),它主要通过结合最小连通子图与先验知识实现。构建规则如下:
- CLS 节点与所有特征节点双向连接。
- 特征节点之间按最小连通子图(如顺序连接)方式连接,避免全连接带来的冗余。邻接矩阵:\(\mathcal{A}^{hcs} = f_{\text{mcsg}}(\{v_i | i \in \mathcal{S}_{nor}\}) \cup A^{cls}\),其中 \(A^{cls}\) 表示 CLS 节点的连接。
第二个图是自动学习图 \(\mathcal{G}_i^{gsl}\),它通过图结构学习器数据驱动地学习节点间关系。构建流程为:
- 基于节点嵌入计算概率矩阵:\(P^{gsl}[j,k]=\exp^{-\|h_{i}^{j}-h_{i}^{k}\|^2}\)。
- 应用 Gumbel-Softmax 和 Top-k 采样,将概率矩阵转化为稀疏的、无向的硬邻接矩阵 \(\mathcal{A}^{gsl}\)。

接着进行图内消息传递,使用的 GNN 为交互网络,它为每条边 \(e(v_m, v_n)\) 初始化一个嵌入,例如为其两端节点嵌入的均值。接着在消息传递中,同时更新边嵌入和节点嵌入:
| 消息传递操作 | 说明 |
|---|---|
| 边嵌入更新 | 聚合其两端节点信息及自身信息 |
| 节点嵌入更新 | 聚合与其相连的所有边的信息及自身信息 |
通过 MLP 处理聚合后的信息,并使用残差连接进行更新。两个图 \(\mathcal{G}_i^{hcs}\) 和 \(\mathcal{G}_i^{gsl}\) 分别通过独立的交互网络,得到更新后的节点嵌入集合:

最后进行特征融合与预测,从两个图的输出中分别提取全局表示:
对自动学习图的CLS表示 \(z_i^{gsl}\) 进行 Add & Norm 操作(残差连接 + 层归一化),以稳定训练:
最后将两个 CLS 表示拼接,输入一个 MLP 进行最终预测:
复杂度分析
模型整体复杂度 \(O_{\text{MPCFIN}}\) 主要由三部分主导:
| 模块 | 复杂度 |
|---|---|
| 交叉特征交互 | 计算特征间相关性矩阵,复杂度为 \(O(M^2 \cdot d)\) |
| 图结构学习与消息传递 | 构建图和学习结构为 \(O(M^2 \cdot d)\),GNN 进行 \(T\) 轮消息传递的复杂度为 \(O(T \cdot M \cdot d)\) |
| 输出模块 | MLP 预测器的复杂度为 \(O(L \cdot d^2)\) |
由于特征数量 \(M\) 通常远大于迭代次数 \(T\) 和层数 \(L\),模型的计算复杂度可简化为:
这表明模型复杂度主要与特征数量的平方和嵌入维度成正比。
实验结果
数据集和实验设置
共使用 7 个数据集,包括 1 个私有临床数据集和 6 个公共数据集。私有数据集来自武汉爱尔眼科医院集团,包含 4 种眼病的诊断数据(1 个正常标签,3 个疾病标签),共 44 个特征。公共数据集涵盖分类与回归任务,规模、特征数和类别数各异,具体信息见下表:

用于对比的基线模型包括:
| 算法类型 | 对比模型 |
|---|---|
| 传统机器学习 | SVM, MLP |
| 混合模型 | NODE。 |
| 基于 Transformer 的模型 | TabTransformer, SAINT |
| 基于 GNN 的模型 | INCE, T2G-Former |
实验设置如下:
| 实验设置 | 说明 |
|---|---|
| 预处理 | 标准化数值特征,用标签编码器处理类别特征,缺失值补零。 |
| 数据划分 | 80% 训练,20% 测试,训练集中再划分 10% 作为验证集用于超参数调优。 |
| 超参数 | 使用 Optuna 库自动调优,超参数包括:隐藏向量维度(16, 32, 64, 128)、交互网络层数(1-4)、MLP 层数(1-4)。 |
| 训练 | 使用Adam优化器(学习率0.001),批量大小256,训练200轮并采用早停策略。结果报告为3次独立运行的平均值。 |
| 评估指标 | 分类任务用准确率(ACC)和AUC;回归任务用均方误差(MSE)。在医疗数据集上额外评估了假阴性率和假阳性率。 |
对比实验
MPCFIN 在 5 个公开数据集上表现最佳,在剩余的 1 个数据集上获得次优结果,综合性能超越所有对比的 DNN 模型。基于 GNN 的方法(如 INCE, T2G-Former, MPCFIN)整体上优于其他类型的 DNN,间接证明了用图结构建模表格数据的有效性。MPCFIN 的稳定性略逊于基于注意力机制的模型,如 SAINT,可能是因为 MPCFIN 的图结构学习模块(GSL)在每一步训练时都基于输入特征动态学习边连接,在样本有限时,边的生成可能不稳定。

在需要高可靠性的医疗诊断场景中,MPCFIN 展现出独特优势。在整体准确率(ACC)和 AUC 上,MPCFIN 排名第二,略低于最优的 T2G-Former。在更关键的假阴性率和假阳性率指标上,MPCFIN 全面优于 T2G-Former 及其他所有基线模型。这表明 MPCFIN 在降低误诊和漏诊风险方面表现更佳,这对于临床实践具有重大价值。

消融实验
为验证各模块有效性,在三个数据集上进行了消融研究。
| 研究主题 | 对比方法 | 结论 |
|---|---|---|
| 不同交叉特征策略 | 单特征、可重复随机采样、不可重复随机采样、自适应采样 (Ours) | 自适应采样方法效果最好、非重复采样优于可重复采样,因后者会造成信息冗余。所有交叉特征方法均优于单特征方法。 |
| 不同组合方式 | 单特征+全连接图、单特征+多路图、交叉特征+全连接图、交叉特征+多路图(Ours) | 交叉特征+多路图的组合取得了最佳性能,证明了两个核心设计的协同作用。 |
| 不同图结构 | 仅自动学习图、仅手工构建图、多路图 (Ours) | 多路图显著优于单独使用任何一种子图,证明了融合先验知识与数据驱动学习的必要性。 |
| Add&Norm 层位置 | 尝试将其加在自动学习图后、手工构建图后、两者都加或都不加。 | 实验表明,仅将 Add&Norm 层加在自动学习图路径之后效果最佳,有助于稳定该路径的训练。 |

超参数敏感性分析
对于 AFC 中的阈值 T,实验发现不同数据集对阈值 T 的敏感度不同。为公平起见,论文统一将其设为 0.5,作为一个在探索特征交互和避免冗余之间取得平衡的稳健基线。对于 GSL 中的 Top-k 比例,实验探究了不同 k 值,即每个节点连接的邻居比例的影响。结果表明,需要避免 k 值过高(图过密,类似全连接图)或过低(图过稀疏,信息传递不畅),需取得平衡。论文还对学得的图结构进行了可视化,证明了其合理性。

计算复杂度对比
在私有数据集上对比了 GNN 类模型的复杂度。参数量方面 MPCFIN (138.5K) 远小于 T2G-Former (1664.4K),大于 INCE (44.844K)。计算量 (MACs) 方面 MPCFIN (3724M) 显著低于基于注意力的 T2G-Former (30380.6M),但高于简单的 INCE (630.383M)。MPCFIN 在模型性能和计算效率之间取得了较好的平衡,其主要计算开销在于多路 GNN 模块。

自适应特征组合器的可解释性
通过可视化 AFC 学习到的相关性矩阵(未经过阈值化的原始矩阵 P),证明了模型能学习到符合人类直觉的特征关联。加州住房数据集显示“收入中位数”与“房屋年龄”、“平均房间数”强相关;“平均入住人数”与“平均房间数”、“平均卧室数”强相关,这些都与现实认知一致。成人收入数据集显示“工作阶级”与“教育水平”强相关,符合社会经验。这表明 MPCFIN 的交叉特征模块不仅有效,而且具有良好的可解释性。

优点和创新点
个人认为,本文有如下一些优点和创新点可供参考学习:
- 通过可学习的特征相关性矩阵,为每个特征自动构建融合其最相关特征的交叉特征表示,有效建模了表格数据中异构特征的交互关系。
- 将手工构建的连通性(基于先验知识)与自动学习的连通性(数据驱动)结合在一个框架内,为特征交互构建了更优且更鲁棒的图拓扑。
- 在医疗诊断等应用中展现出可靠性与可解释性,模型不仅在多个基准数据集上取得先进性能,更在私有医疗数据集上实现了最低的假阴性与假阳性率,且其学习的特征相关性矩阵直观合理。

浙公网安备 33010602011771号