Paper Reading: TabR: Tabular Deep Learning Meets Nearest Neighbors


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

论文概况 详细
标题 《TabR: Tabular Deep Learning Meets Nearest Neighbors》
作者 Yury Gorishniy, Ivan Rubachev, Nikolay Kartashev, Daniil Shlenskii, Akim Kotelnikov, Artem Babenko
发表会议 The Twelfth International Conference on Learning Representations(ICLR 2024)
发表年份 2024
会议等级 CCF-A
论文代码 https://github.com/yandex-research/tabular-dl-tabr

作者单位:

  1. Yandex
  2. HSE

研究动机

在表格数据的机器学习任务中,基于 GBDT 的非深度学习模型长期以来一直是首选且强大的解决方案。近年来,针对表格数据的深度学习模型受到越来越多关注,但其性能通常仍难以超越精心调优的 GBDT 模型,尤其是在小到中等规模的数据集上。一个有望提升表格 DL 性能的研究方向是设计检索增强模型,这类模型的基本思路是:对于一个目标样本,从训练数据中检索出相关样本(如最近邻),并利用这些检索到的上下文样本的特征和标签信息来辅助预测。尽管已经存在一些检索增强的表格 DL 模型,但它们仅比经过适当调优的 MLP 带来非常有限的性能提升,同时它们在架构上通常非常复杂且计算成本高昂。基于以上背景,本文旨在解决的问题是:能否设计一个既高效又强大的检索增强型表格深度学习模型?

文章贡献

本文提出了一种基于检索的表格深度学习模型 TabR,其核心在于设计了一个类 KNN 的检索模块,并将其集成到一个简单的前馈神经网络中。在预测时,TabR 会为目标样本从其训练数据中检索出称为上下文对象的最相似的样本。TabR 的关键组件是检索模块中的相似性计算与值聚合机制。TabR 的相似性计算摒弃了标准注意力机制中的查询-键点积,转而采用基于键的 L2 距离来衡量样本间的相似性。在聚合上下文信息时,值聚合机制不仅利用了上下文对象的标签,还通过一个小型网络引入了校正项动态调整标签的贡献。在多个公开基准测试中,TabR 取得了最佳的平均性能,并成为多个数据集上的最优方法。

本文方法

符号和术语定义

对于一个给定的表格数据监督学习问题,本文的符号定义如下,上下文允许时下标 \(i\) 可以省略。为了简化表述,论文在大多数地方假设 \(x_i\) 只包含连续(或数值型)特征。数据集被划分为三个互斥的部分 \(\overline{1, n} = I_{train} \cup I_{val} \cup I_{test}\),模型需要做出预测的输入样本被称为输入对象或目标对象。

符号 含义
\(\{(x_i, y_i)\}_{i=1}^{n}\) 数据集
\(x_i \in X\) 表示第 \(i\) 个样本的特征
\(y_i \in Y\) 表示第 \(i\) 个样本的标签
\(Y = \{0, 1\}\) 二分类
\(Y = \{1, \dots, C\}\) 多分类
\(Y = \mathbb{R}\) 回归
\((I_{train})\) 训练集
\((I_{val})\) 验证集
\((I_{test})\) 测试集

在使用检索技术时,本文规定了以下术语,本文中的所有输入对象使用相同的候选集 \(I_{cand} = I_{train}\)

术语 含义
候选集 检索操作从一个上下文候选集 \(I_{cand} \subseteq I_{train}\) 中进行
上下文对象 被检索到的对象被称为上下文对象,或简称为上下文
可选操作 目标对象可以(作为一个特殊情况)被包含在自己的上下文集合中

整体设计思路

TabR 的设计采用了一种增量式方法,从一个简单的、无检索的基础架构开始,逐步添加并优化检索组件。考虑一个通用的、无检索的前馈网络 \(f(x) = P(E(x))\),可非正式地划分为两部分:

网络组件 符号 功能
编码器 \(E: X \rightarrow \mathbb{R}^d\) 将原始特征映射到中间表示
预测器 \(P: \mathbb{R}^d \rightarrow \hat{Y}\) 基于中间表示进行最终预测

为了将其转变为基于检索的模型,本文在编码器 \(E\) 之后,以一个残差连接的形式添加了检索模块 \(R\)。相关符号的定义如下:

符号 含义
\(\tilde{x} = E(x) \in \mathbb{R}^d\) 目标对象的中间表示
\(\{\tilde{x}_i\}_{i \in I_{cand}} \subset \mathbb{R}^d\) 所有候选对象的中间表示
\(\{y_i\}_{i \in I_{cand}} \subset Y\) 候选对象的标签

搜狗高速浏览器截图20260318004437

编码器和预测器并非本文重点,因此保持了简单设计。

网络组件设计 符号 功能
编码器 \(E\) 可包含 0 个或多个块(Block),每个块由线性层、ReLU、Dropout 和 Layer Normalization 组成
预测器 \(P\) 可包含 1 个或多个块

搜狗高速浏览器截图20260318005611

检索模块的设计

检索模块 \(R\) 的灵感来源于 KNN,其流程为:

  1. 相似性计算:通过相似性模块 \(\mathcal{S}\) 计算目标对象与所有候选对象的相似性分数。
  2. Top-m 选择:根据相似性分数选择 Top-m 个候选对象作为上下文(m=96)。
  3. 值提取:通过值模块 \(\mathcal{V}\) 计算每个被选中上下文对象的值表示。
  4. 加权求和:对上下文对象的值表示进行加权求和,权重由 Softmax 计算得出。求和后与目标对象的表示 \(\tilde{x}\) 相加,形成增强后的表示,传递给预测器 \(P\)

搜狗高速浏览器截图20260318005750

本文聚焦于设计相似性模块 \(\mathcal{S}\)值模块 \(\mathcal{V}\),此时将编码器 \(E\) 暂时被设定为线性层\(N_E=0\)),且不启用数值特征嵌入。
image

Step-0: 类标准注意力的基线

使用标准自注意力(忽略 Top-m 操作)的思路作为最基础的版本:

\[\mathcal{S}(\tilde{x}, \tilde{x_i}) = W_Q(\tilde{x})^T W_K(\tilde{x_i}) \cdot d^{-1/2} \quad \mathcal{V}(\tilde{x}, \tilde{x_i}, y_i) = W_V(\tilde{x_i}) \]

其中 \(W_Q, W_K, W_V\) 是线性层,且将目标对象无条件地作为其上下文的第 (m+1) 个对象。此配置的性能与 MLP 相似,说明标准自注意力用于检索是次优策略。

Step-1: 添加上下文标签

一个改进是尝试利用上下文对象的标签,例如将其加入到值模块:

\[\mathcal{S}(\tilde{x}, \tilde{x_i}) = W_Q(\tilde{x})^T W_K(\tilde{x_i}) \cdot d^{-1/2} \quad \mathcal{V}(\tilde{x}, \tilde{x_i}, y_i) = \underline{W_Y(y_i)} + W_V(\tilde{x_i}) \]

下划线部分是新增的 \(W_Y: Y \rightarrow \mathbb{R}^d\),对于分类任务它是一个嵌入表,对于回归任务它是一个线性层。但实验结果显示,仅添加标签并未带来性能提升。这个实验结果很反直觉,可能是因为从标准注意力借鉴的相似性模块 \(\mathcal{S}\) 无法有效利用标签这一有价值的信号。

Step-2: 改进相似性模块

通过实验发现,移除查询(即移除 \(W_Q\),等同于 \(W_Q = W_K\))并使用 \(L_2\) 距离代替点积,能在多个数据集上显著提升性能:

\[\mathcal{S}(\tilde{x}, \tilde{x_i}) = \underline{-\|W_K(\tilde{x}) - W_K(\tilde{x_i})\|^2} \cdot d^{-1/2} \quad \mathcal{V}(\tilde{x}, \tilde{x_i}, y_i) = W_Y(y_i) + W_V(\tilde{x_i}) \]

后续的消融实验证明,上下文标签、仅用键表示、\(L_2\)距离这三个要素缺一不可,移除任何一个都会导致性能回落到 MLP 水平。

Step-3: 改进值模块

受到近期回归算法 DNNR 的启发,为了让值模块 \(\mathcal{V}\) 更具表达力,将目标对象的表示 \(\tilde{x}\) 也考虑进来:

\[\begin{align*} \mathcal{S}(\tilde{x}, \tilde{x_i}) = -\|W_K(\tilde{x}) - W_K(\tilde{x_i})\|^2 \cdot d^{-1/2} \\ \mathcal{V}(\tilde{x}, \tilde{x_i}, y_i) = W_Y(y_i) + \underline{T(W_K(\tilde{x}) - W_K(\tilde{x_i}))} \\ T(\cdot) = \text{LinearWithoutBias}(\text{Dropout}(\text{ReLU}(\text{Linear}(\cdot)))) \end{align*} \]

其中 \(T\) 是一个小型网络。直观上,\(W_Y(y_i)\) 是第 \(i\) 个上下文对象的原始标签贡献,而 \(T(W_K(\tilde{x}) - W_K(\tilde{x_i}))\) 可被视为校正项,它将键空间的差异转换为标签嵌入空间的差异。如实验结果所示,这个新的值模块进一步提升了多个数据集的性能。

Step-4: TabR

最终通过经验观察,在相似性模块中省略缩放项 \(d^{-1/2}\) 以及不将目标对象包含在其自身上下文中,能带来平均更好的结果。本文将得到的模型命名为 TabR,其检索模块 \(R\) 的正式、完整描述如下:

\[\begin{align*} k &= W_K(\tilde{x}), \quad k_i = W_K(\tilde{x_i}) \\ \mathcal{S}(\tilde{x}, \tilde{x_i}) &= -\|k - k_i\|^2 \\ \mathcal{V}(\tilde{x}, \tilde{x_i}, y_i) &= W_Y(y_i) + T(k - k_i) \\ \end{align*} \]

其中 \(W_K\) 是一个线性层,\(W_Y\) 对分类任务是嵌入表,对回归任务是线性层。默认情况下,目标对象不包含在其自身上下文中,相似性分数不进行缩放。

\[T(\cdot) = \text{LinearWithoutBias}(\text{Dropout}(\text{ReLU}(\text{Linear}(\cdot)))) \]

实验数据

数据集和实验设置

论文主要使用先前文献中的数据集,如下表所示。对任一算法,在每个数据集上使用验证集进行超参数调优和早停。在选定最佳超参数后,使用 15 个随机种子在测试集上评估,并报告性能指标的平均值。比较任意两个算法时,会将标准偏差考虑在内。为了获得同一类型模型的集成结果,将 15 个随机种子分成 3 组(每组 5 个模型),在每组内平均预测结果,然后报告这 3 个集成模型的平均性能。
image

本文使用两个版本的 TabR:

TabR 版本 说明
TabR 完整配置,在编码器和预测器中拥有所有可调自由度
TabR-S 简单配置,不使用数值特征嵌入,编码器是线性的(\(N_E=0\)),预测器只有一个块(\(N_P=1\)

评估检索增强 DL 模型

本节将 TabR 与现有的检索增强方案,以及完全参数化的深度学习模型进行比较。从实验数据可见,TabR 是唯一能在许多数据集上相比 MLP 带来性能提升的检索增强模型。完整的 TabR 在多个数据集(CA, OT, BL, WE, CO)上优于 MLP-PLR(前置研究报告的平均排名最高的参数化 DL 模型),而在其余数据集上(除 MI 外)表现与 MLP-PLR 相当。相比先前的检索增强方案,TabR 不仅性能更高,还克服了它们的多项限制,如对分类任务的不兼容、扩展性问题。结果表明,检索技术数值特征嵌入是两种能改善表格 DL 模型优化性能的组件。
image

与 GBDT 的比较

TabR 与基于 GBDT 的模型进行比较,包括 XGBoost、LightGBM 和 CatBoost。经过调优的 TabR 在多个数据集(CH, CA, HO, HI, WE, CO)上相比调优的 GBDT 模型提供了改进,而在其余数据集上(除 MI 外)也具有竞争力。TabR 的默认配置同样表现出色,与经过精心调优的 GBDT(如 CatBoost)相比具有竞争力。
image
接着使用前置工作中,用来证明 GBDT 在样本数 ≤5 万的中小规模任务上优于参数化 DL 模型的。结果可见 MLP-PLR 在此基准上确实略微落后于 GBDT,TabR 在平均性能上超越了 GBDT。
image

扩展分析

冻结上下文

在 TabR 的标准设定中,每个训练批次都需要为当前模型编码所有候选对象并计算相似性,这在大型数据集上会变得非常缓慢。例如,在包含 300 多万个样本的数据集上,训练一个 TabR 模型需超过 18 小时。观察到在训练过程中,对于平均的训练样本,其上下文(即 Top-m 候选对象及其对应的相似性分布)会逐渐稳定下来。因此可以在训练进行到固定的 epoch 后,执行一次上下文冻结。这意味着最后一次为所有训练样本计算其最新的上下文,并在剩余的整个训练过程中重用这些被冻结的上下文,不再动态更新。
搜狗高速浏览器截图20260318014246

在一些数据集上,这个技巧可以在不显著损失性能的前提下加速 TabR 的训练,在更大的数据集上提速效果更明显。例如在 300 多万个样本的数据集上,该方法实现了近 7 倍的训练加速,从 18 小时 9 分钟缩短到 3 小时 15 分钟,同时仍保持 RMSE 性能。
image

取消重新训练

在模型训练后获取新的、未见过的训练数据是常见的实际场景,TabR 允许一种无需重新训练即可利用新数据的方法。该方法直接将新的训练数据添加到用于检索的候选集(\(I_{cand}\))中,当已训练好的 TabR 模型在做预测时,可以同时从原始的候选集和新数据中进行检索。实验数据表明这种策略是可行的,此外这种方法还可用于扩展TabR到大规模数据集。
搜狗高速浏览器截图20260318014717

优点和创新点

  1. 本文将标准注意力中的查询-键点积相似性,替换为基于键的 L2 距离计算,实现了模型性能的提升。
  2. 在聚合上下文信息时,不仅使用上下文对象的标签嵌入,还新增了一个基于特征空间差异的校正项,使模型能动态调整邻居信息的贡献。
  3. 用单个、定制化的类注意力模块取代先前检索模型中复杂、低效的多层多头 Transformer 交互,实现了性能与效率的平衡。
posted @ 2026-03-18 01:58  乌漆WhiteMoon  阅读(87)  评论(0)    收藏  举报