Paper Reading: Table2Graph: Transforming Tabular Data to Unifed Weighted Graph


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

论文概况 详细
标题 《Table2Graph: Transforming Tabular Data to Unifed Weighted Graph》
作者 Kaixiong Zhou, Zirui Liu, Rui Chen, Li Li, Soo-Hyun Choi, Xia Hu
发表会议 The Thirty-First International Joint Conference on Artificial Intelligence (IJCAI-22)
发表年份 2022
会议等级 CCF-B
论文代码 文中未公开

作者单位:

  1. Department of Computer Science, Rice University
  2. Samsung Research America
  3. Samsung Electronics

研究动机

表格数据是许多现实世界应用(如推荐系统、在线广告)中最常见的数据形式,其中建模特征(列)之间复杂的交互关系是提升模型预测性能的关键。例如,在在线广告中,用户所在地区与广告内容语言之间的交互,能有效预测用户的点击行为。现有特征交互建模方法存在局限性:

方法类型 局限
传统方法 如逻辑回归、因子分解机这类低阶模型只能捕获低阶交互,能力有限。
深度神经网络 能够隐式地学习高阶交互,但交互信息在网络隐藏层中相互纠缠,导致优化过程复杂(易陷入局部最优)。
启发式固定图 基于领域知识预先定义特征间连接可能构造出次优的图,因为它遗漏重要的真实连接或引入噪声连接,并且难以推广到其他表格数据上。
基于样本的注意力图 为每个样本独立计算注意力权重来构建加权图(如 Fi-GNN)。虽然能灵活建模样本特定的交互,但为每个样本都计算一个图是极其耗时的。

基于上述背景,论文提出了一个问题:是否存在一种端到端的框架,能够用一个统一的图来建模特征交互?统一的图旨在捕获整个表格数据集中所有样本共享的、常见且重要的特征交互模式。与样本级的图构建相比,它无需为每个样本进行耗时的注意力计算,从而在推理时实现高效预测。具体而言,论文旨在解决构建这样一个统一图时所面临的两个挑战:

  1. 如何稳定地提取关键交互模式:表格数据中的样本可能呈现多样化的特征交互模式,会导致统一图的学习过程不稳定。
  2. 如何正则化统一图的结构:一个过于平滑(所有节点均匀连接)的图无法突出关键模式,而一个过于稀疏的图则容易对某些特定的交互模式过拟合,缺乏对多样样本的适应性。

文章贡献

本文提出了一个名为 Table2Graph​ 的框架,旨在将表格数据的特征交互建模问题转化为一个统一的加权图学习问题。其核心模型通过学习一个全局共享的概率邻接矩阵来显式地编码所有特征列之间的交互强度。为了解决图学习过程中的不稳定性和稀疏性控制难题,模型引入了强化学习策略以稳定地强化对提升预测性能至关重要的关键交互连接,并设计了一个可微的稀疏性约束损失来正则化图的连接结构,避免其过于平滑或稀疏。最终,模型在一个包含预测任务损失、强化损失和稀疏性约束损失的联合损失函数指导下,端到端地训练出这个统一图及相应的图神经网络模型,以实现高效且准确的预测。通过在多个真实世界和合成数据集上进行的广泛实验,本文严谨地验证了该框架的有效性、效率和可解释性。

本文方法

Table2Graph的核心思想是将表格数据的特征交互建模问题转化为一个统一的图学习问题,框架包含三个关键部分:统一图的构建、利用强化学习稳定训练、以及加入稀疏性约束的联合优化。
image

问题定义

给定表格数据,目标是学习一个映射函数 \(f\),以便预测未见样本的标签 \(\hat{y}^{(i)}=f(x^{(i)})\)。其中,\(x^{(i)}=[x_1^{(i)},\cdots,x_m^{(i)}]\) 表示第 \(i\) 个样本的 \(m\) 个特征取值,\(y^{(i)}\) 是其标签。论文旨在学习一个统一的图,以建模所有样本共享的特征交互,并应用 GNN 来实现映射函数 \(f\)
定义特征交互图为 \(G=(\mathcal{V},\mathcal{E})\),表格中的每一列(特征)被视为一个节点 \(\mathcal{V}\),例如第 \(j\) 列对应节点 \(j\)。节点之间的连接 \(\mathcal{E}\) 表示特征之间的交互关系。用矩阵 \(A \in R^{m \times m}\) 来表示特征交互图,称为概率邻接矩阵。矩阵的每一行都被归一化,总和为 1。元素 \(A_{jk} \ge 0\) 表示节点 \(j\)\(k\) 之间的边权重,即特征 \(x_j\)\(x_k\) 之间的交互强度。

统一图构建

统一图构建的目标是为整个数据集学习一个全局共享的概率邻接矩阵 \(A\),然后在此基础上应用 GNN 进行特征交互学习和预测。为表格的 \(m\) 个特征列学习一个列嵌入 \(E \in R^{m \times d}\),每行代表一个特征,通过自注意力机制计算概率邻接矩阵 \(A\)

\[A = \text{Softmax}(\sigma(E W_l) \sigma(E W_r)^\top) \in R^{m \times m} \]

其中 \(W_l, W_r \in R^{d \times d’}\) 是可训练矩阵,\(\sigma\) 是激活函数。矩阵的每一行应用 Softmax 进行归一化,使 \(A\) 具有概率解释,\(A_{jk}\) 表示特征 \(j\)\(k\) 的交互强度。

特征交互学习

对于样本 \(x^{(i)} = [x_1^{(i)}, \cdots, x_m^{(i)}]\),将其特征值转换为初始特征嵌入 \(X_0^{(i)} \in R^{m \times d}\),作为图中 \(m\) 个节点的初始表示。在构建的图 \(A\) 上应用 K 层 GNN 来学习高阶特征交互。第 \(k\) 层的图卷积操作为:

\[X_k^{(i)} = X_0^{(i)} + \sigma(A X_{k-1}^{(i)} W_k) \]

其中 \(W_k \in R^{d \times d}\) 是可训练矩阵。这里加入了初始连接 \(X_0^{(i)}\) 以促进梯度流动,有助于更好地训练初始特征嵌入。经过 \(K\) 层聚合后,将最终的特征嵌入 \(X_K^{(i)}\) 拼接并输入预测层,得到预测结果 \(\hat{y}^{(i)}\)。模型通过任务损失(如分类任务的交叉熵)进行优化:

\[L_{task} = y^{(i)}\log(\hat{y}^{(i)}) + (1-y^{(i)})\log(1-\hat{y}^{(i)}) \]

统一图训练

如果仅用任务损失 \(L_{task}\) 端到端地优化邻接矩阵 \(A\),会面临两个问题:

  1. \(A\) 倾向于学习一个稠密矩阵,无法突出关键交互;
  2. 不同样本的交互模式各异,导致 \(A\) 的训练不稳定。

为解决此问题,论文引入强化学习来稳定地强化关键的共享交互连接,其灵感来源于神经架构搜索。首先重要链接采样,对于矩阵 \(A\) 的每一行(对应一个特征),基于其定义的概率分布(多项式分布)采样 \(s\) 个交互对:

\[\mathcal{I}_i = \text{RowSample}(A[i,:], s) = \{(i, j_1), \cdots, (i, j_s)\} \]

其中 \(\mathcal{I} = \bigcup_i \mathcal{I}_i\) 是本次采样的所有边集合,这迫使模型关注高权重的关键交互。将任务损失 \(L_{task}\) 的倒数作为奖励信号 \(R = 1 / L_{task}\),采用 REINFORCE 规则构建强化损失 \(L_{rl}(A)\)

\[$L_{rl}(A) = -\lambda_1 * \mathbb{E}_{\mathcal{I} \sim A} [\sum_{(i,j)\in\mathcal{I}} (R - R_{avg}) \log A_{ij}] \]

其中 \(\lambda_1\) 是超参数,\(R_{avg}\) 是奖励的滑动平均值。如果当前奖励 \(R\) 高于历史平均 \(R_{avg}\),则增加所采样链接 \((i,j)\) 的权重 \(A_{ij}\);反之则降低。这驱使 \(A\) 去学习那些能持续提升模型性能(降低任务损失)的、为多数样本所共享的关键特征交互。在实践中,通过采样一个 \(\mathcal{I}\) 来近似期望计算,以简化训练。

联合训练与稀疏性约束

为了控制图的稀疏度,避免其变得过于平滑(无法突出关键交互)或过于稀疏(对特定样本过拟合),论文设计了可微的稀疏性损失 \(L_{sp}(A)\)

\[L_{sp}(A) = -\frac{\lambda_2}{m} \mathbf{1}^T \log(A \odot A) \mathbf{1} + \frac{\lambda_3}{m} ||A||_F^2 \]

其中 \(\mathbf{1}\) 是全1向量,\(\odot\) 是逐元素相乘,\(||\cdot||_F\) 是Frobenius范数,\(\lambda_2, \lambda_3\) 是超参数。损失函数的第一项最大化每行分布的锐度 \(\sum_j A_{ij}^2\)。由于每行已是概率向量(和为 1),最大化其平方和会促使其分布稀疏化,即某些值更大,其余更小。第二项是惩罚 \(A\) 的Frobenius 范数,防止其变得过于稀疏。框架通过最小化一个联合损失来端到端地学习统一的图 \(A\) 和 GNN 模型:

\[L = L_{task} + L_{rl}(A) + L_{sp}(A) \]

这个总损失综合了预测准确性、关键交互的稳定学习和图结构的正则化三个目标,是 Table2Graph 框架的核心优化目标。

实验结果

数据集

本文使用了 4 个数据集:

数据集 说明
Creditcard 金融欺诈检测数据集,284,807个交易样本,28个数值型匿名特征。
Criteo 在线广告数据集,4500万点击记录,39个特征域。
Synthetic 一个合成数据集,用于评估特征交互检测的准确性。
MovieLens 协同过滤常用数据集,包含 3706 个项目(即特征列),任务是为用户预测个性化偏好分数并排序。

其中合成数据集为一个回归任务,被明确定义为包含四组真实特征交互的函数,其真实交互为:$ {[x_0, x_1, x_2], [x_3, x_4], [x_5, x_6], [x_7, x_8, x_9]}$。

\[y = \frac{1}{1+x_0^2+x_1^2+x_2^2} + \sqrt{e^{x_3+x_4}} + |x_5 + x_6| + x_7x_8x_9 \]

在真实数据集上使用 AUC 和 LogLoss 评估预测性能,在合成数据集上使用交互检测的 AUC 评估。在 MovieLens 上,将 Table2Graph 集成到 FISM 和 NAIS 两个推荐框架中,用 HR@10 和 NDCG@10 评估。使用的对比模型如下:

算法类型 对比模型
一阶模型 线性回归
因子分解机类 FM、DeepFM、AFM
基于树的模型 RF、DT
深度神经网络 MLP、DeepCrossing、NFM、CIN
图学习方法 Fi-GNN、Fixed-GNN

实现方面使用一个三层 GNN 模型,超参数 \(\lambda_1, \lambda_2, \lambda_3\) 通过网格搜索确定。

对比实验

在 Creditcard 和 Criteo 数据集上,Table2Graph 的 AUC 和 LogLoss 均优于所有基线,表明统一的图建模能有效提升下游任务性能。在合成数据集上,Table2Graph 的交互检测 AUC 达到 0.9714,远超其他方法,证明了其能准确捕获真实的特征交互模式。高阶模型(如随机森林、NFM)通常优于一阶和 FM 类模型,说明捕捉高阶交互的重要性。图学习方法总体上优于传统基线,证明了显式建模交互图的优势。Table2Graph > Fi-GNN > Fixed-GNN 说明:统一的全局图优于为每个样本单独建图(Fi-GNN),因为后者易受样本噪声干扰;学习到的自适应图优于固定的全连接图(Fixed-GNN),因为后者引入了噪声连接。
image

在 MovieLens 上,无论是基于 FISM 还是 NAIS 框架,集成 Table2Graph 的模型在 HR 和 NDCG 指标上均达到最高。这表明 Table2Graph 能有效学习大规模项目-项目交互图。其采用的强化学习和稀疏性约束有助于加权对多数用户有益的通用项目交互,同时避免对特定交互模式的过拟合。
image

模型超参数

在 MovieLens 上研究了超参数 \(\lambda_1\)(强化损失)、\(\lambda_2\)\(\lambda_3\)(稀疏性损失)的影响。当 \(\lambda_1=0\)(即不使用强化学习)时,性能显著下降。这验证了 RL 对于稳定学习关键交互模式、降低图学习方差的必要性。适当的 \(\lambda_2\)\(\lambda_3\) 值(如 \(\lambda_2 \rightarrow 0, \lambda_3 \rightarrow 10^{-4}\))能获得最佳性能。这证实了正则化在避免图过于稀疏(过拟合)和过于平滑(无法突出重点)之间的关键平衡作用。
image

运行效率

在训练和测试时间上,Table2Graph 与 Fixed-GNN(固定图)效率相当,且远快于 Fi-GNN。Table2Graph 在训练阶段学习一个统一的图,在测试阶段直接复用,无需额外计算。而 Fi-GNN 需要为每个测试样本重新计算注意力图,导致巨大的计算开销。这证明了 Table2Graph 在实际部署中具有显著的时间优势。
image

可视化分析

论文可视化了在合成数据集上学到的邻接矩阵 \(A\),结果显示学到的图中最强的交互边恰好对应着四组真实的特征交互 \(\{[0,1,2], [3,4], [5,6], [7,8,9]\}\),为模型的可解释性提供了直观证据。
image

优点和创新点

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

  1. 提出的 Table2Graph 框架将表格数据的特征交互建模问题转化为一个可学习的、统一的概率邻接矩阵,实现了从隐式或启发式交互建模到端到端显式图学习的范式转变。
  2. 为解决图学习不稳定和难以聚焦关键交互的难题,引入强化学习策略来优化邻接矩阵,稳定地强化对提升整体任务性能至关重要的通用特征交互连接。
  3. 设计了一种新颖的可微稀疏性约束损失,能够自动地调节所学图的稀疏程度,有效避免了图结构因过于平滑而无法突出重点,或因过于稀疏而对特定样本过拟合。
posted @ 2026-04-27 00:39  乌漆WhiteMoon  阅读(18)  评论(0)    收藏  举报