Google Research 推出零样本表格基础模型 TabFM
Introducing TabFM: A zero-shot foundation model for tabular data
Google Research 推出零样本表格基础模型 TabFM,一次前向即可对未见过的表做分类和回归,无需逐表训练、调参和特征工程。它把整张表当作统一提示,用行列交替注意力、行压缩和 Transformer 做上下文学习,全部基于数亿个结构因果模型合成的数据集预训练。
把表格预测改写成上下文学习之后,读者可以对照传统树模型流程判断单次前向能否省掉调参和特征工程。
自我们推出 TimesFM 以来,人们处理时间序列预测的方式已发生巨大转变。现在,我们正把同样的“零样本”逻辑带到表格数据上。
我们推出 TabFM,一款面向表格数据的新型基础模型,用以简化分类与回归工作流。
表格数据构成企业数据基础设施的骨干,并支撑着相当一部分关键的预测性机器学习应用。从预测客户流失到识别金融欺诈,表格回归与分类任务无处不在。多年来,有监督的基于树的算法,例如AdaBoost、XGBoost和随机森林等,历来主导这一领域,在结构化数据上提供稳健的性能。
然而,部署这些传统模型的生命周期构成了显著瓶颈。将 XGBoost 模型拟合到新数据集,并非只是一次 .fit() 步骤那么简单;它无一例外地需要繁琐的人工投入。数据科学家必须投入无数小时进行广泛的超参数优化和领域特定的特征工程,才能从原始数据中提取可靠信号。
另一方面,更广泛机器学习领域的最新进展——尤其是大语言模型(LLM)的演进——改变了我们与新任务交互的方式。LLM 已通过上下文学习(ICL)展现出零样本预测的非凡能力。该技术让预训练模型只需在输入上下文中提供示例和指令,即可学习新任务,而无需更新任何底层模型权重。
今天,我们推出 TabFM,一款专为表格数据分类与回归设计的基础模型。通过将表格预测构建为 ICL 问题,TabFM 消除了手动模型训练、超参数调优和复杂特征工程的需要。我们很高兴分享这一方法如何让用户在单次前向传播中,对先前未见的表格生成高质量预测。TabFM 现已在我们的 Hugging Face 和 GitHub 仓库中提供,并已直接在 Google Cloud BigQuery 中上线。
工作原理
传统机器学习范式依赖于更新针对给定数据集分布的模型参数。相比之下,ICL 范式完全绕过了这一点。TabFM 并不为每个新任务经历传统训练阶段,而是将整个数据集——既包括历史训练样本,也包括目标测试行——作为单一统一提示。模型在推理时直接从该上下文中学习解读列与行之间的关系。
然而,将 ICL 应用于表格数据并不像对自然语言进行分词那样直接。标准语言模型处理的是一维、有序的序列,但表格本质上是二维且无序的:交换两行或两列并不会改变数据的底层含义。为了在实现可扩展零样本预测的同时有效处理这些多样的表格结构,TabFM 将 TabPFN 与 TabICL 等架构的优势综合为一种新颖的混合设计。如下图所示,该架构依赖于三个关键机制:
- 交替的行与列注意力:首先,原始表格通过多层注意力模块进行处理。与 TabPFN 类似,这一步在列(特征)和行(样本)两个维度上交替施加注意力。通过持续在这两个维度上进行注意力计算,模型学习到丰富的表示,从而原生地捕捉复杂的特征交互与依赖关系。这种深度上下文化有效地完成了原本需要数据科学家进行繁琐手工特征构造的繁重工作。
- 行压缩:在这一上下文化之后,每一行丰富的交叉注意力信息被压缩为单个稠密向量表示。
- 上下文学习(ICL):最后,一个专用的 Transformer 在这一压缩嵌入序列上运行。采用 TabICL 的高效方法,对这些压缩后的行向量(而非原始未压缩网格)进行注意力计算,大幅降低了计算成本。这确保即使面对大得多的数据集,预测步骤仍保持很高的计算效率。
大规模合成数据训练
构建基础模型的典型做法是使用高容量神经网络,并在海量多样化数据上进行训练。然而,表格机器学习的一大障碍在于,高质量、多样化的表格数据集——尤其是能够反映真实工业数据分析所需的大规模表格——在开源领域极为稀缺。工业表格往往包含专有模式和敏感信息,因而无法用于广泛的预训练。
由于合成表格可以生成到任意规模,它们实际上是在这一规模上预训练基础模型的唯一可行选择。因此,TabFM 完全在数亿个合成数据集上训练。这些数据集使用纳入多种随机函数的结构因果模型(SCM)动态生成。这种大规模合成生成捕捉了真实世界表格数据中普遍存在的多样分布和复杂特征关系。因此,模型能够很好地泛化到未见过的真实世界表格,我们将在下文的基准测试中加以展示。
性能与基准测试
为了将 TabFM 与现有最先进方法进行严格对比,我们在 TabArena 上对其进行了评估。TabArena 是一个持续更新的基准系统,根据两两对战胜率计算 Elo 分数。这一全面评估涵盖 38 个分类数据集和 13 个回归数据集,样本量从 700 到 150,000 不等。
如下方性能图所示,我们对模型的两种不同配置进行了基准测试:
- TabFM:这代表模型的开箱即用能力。预测通过单次前向传播生成,无需调参或交叉验证。
- TabFM-Ensemble:该配置通过纳入交叉特征与 SVD(奇异值分解)特征进一步提升性能。我们使用非负最小二乘求解器计算 32 路集成的最优权重。对于分类任务,该变体还将 Platt 缩放 作为额外的校准步骤。
如需查看完整的 TabArena 基准结果——包括详细的逐折指标以及相对于特定基线模型的两两对战胜率——请访问我们的 GitHub 页面。
结论
通过将表格预测重构为上下文学习问题,TabFM 利用混合注意力架构与大规模合成训练数据,原生捕捉复杂的特征交互。该方法成功消除了人工特征工程、超参数优化和重复模型训练等传统瓶颈,并持续优于经过大量调优的行业标准监督学习算法。TabFM 将现代基础模型开箱即用的便利性直接带入表格机器学习工作流,使从业者能够在单次前向传播中生成高度准确的预测。
为使这些能力可直接用于企业分析,TabFM 现已原生集成到 Google Cloud BigQuery。从业者可以使用简单的 SQL 查询(AI.PREDICT)直接在其数据表上运行零样本回归和分类,从而无需自定义模型训练。阅读Google Cloud 公告了解更多信息,或查阅文档,即日起在 BigQuery 中开始使用。
致谢
本项目由 Erez Louidor Ilan、Taman Narayan、Shuxin Nie、Rajat Sen、Yichen Zhou、Joe Toth、Deqing Fu 与 Samet Oymak 共同完成。感谢 Kimberly Schwede 设计图形。
来源:Google Research · research.google