NVIDIA Kumo Tabular 定义表格预测精度与效率新边界
NVIDIA Kumo Tabular 属于 NVIDIA Kumo Structured 模型集合,是一个面向表格数据的开源基础模型,现已在 Hugging Face 上发布。给定一个带有标签行的表格,它能在单次前向传播中预测新行的标签,支持分类和回归任务,且无需训练、调参或特征工程。该模型仅使用合成数据进行预训练,提供三种规模(参数量从 28M 到 215M),可通过我们的 开源库 运行,并以 OpenMDW-1.1 许可证发布,允许商业使用。它在 TabArena、BeyondArena、TALENT 和 ScoringBench 这四项基准测试中均位列第一。
- 模型代码: https://github.com/NVIDIA/structured-data-models
- 模型权重: https://huggingface.co/nvidia/Kumo-Tabular
迈向表格基础模型的转变
表格数据是企业机器学习的基石。客户记录、交易、传感器日志、理赔和订单都存储表中,基于这些数据预测流失、违约、需求或价格,是工业界最常见的机器学习任务之一。过去二十年来,这项工作一直依赖梯度提升树,效果也一直不错。但围绕这些模型的生命周期几乎没有变化:每一个新问题都意味着收集标签、工程化特征、搜索超参数、验证并部署一个对表格通用结构一无所知、需从零开始学习每个任务的模型。
大语言模型展示了处理新任务的另一种方式。只需在提示中提供少量示例,预训练模型即可在不更新任何权重参数的情况下解决任务。这就是 上下文学习,它同样适用于表格:一个在数百万个表格上预训练的模型,可以将带标签的表格作为上下文读取,直接预测新行的标签。
我们今日正式发布了 NVIDIA Kumo Tabular(GitHub、HuggingFace),这是一个用于表格分类和回归任务的开放基础模型。给定包含标记行的表格以及需要预测的行,Kumo Tabular 可以在单次前向传播中返回类别概率或数值预测。
Kumo Tabular 工作原理
Kumo Tabular 是一个基于表格结构构建的 Transformer,采用了 TabICL 和 TabPFN 中提出的列注意力、行注意力及上下文注意力机制。为了预测标签,它需要完成三项任务:(1)理解每个值在其所属列中的含义;(2)理解同一行中各列之间的交互作用;以及 (3)建立带有已知标签的上下文行与带有未知标签的查询行之间的关联。Kumo Tabular 通过以下方式实现这一目标:
单元格嵌入:一组单元格构成一个 Token。数值型和分类型值分别经过傅里叶特征(即学习到的频率的正弦和余弦)处理,且针对不同类型使用独立的权重。缺失值无需插补,而是被特殊处理。最后,上下文中的每个 Token 都会接受一个标签嵌入。
行嵌入:接下来,我们交替多次使用两种注意力机制,将每一行转化为嵌入表示。列注意力沿单列向下查看,通过归纳自注意力学习该值在其所属列分布中的意义,例如,判断 42 是典型值还是极端值。因此,其计算成本随行数呈线性增长。行注意力则横向查看单行内的各个 Token,借助旋转位置编码来区分不同列,从而学习特征之间的交互作用。每行还包含四个可学习的 [CLS] Token,它们共同作为该行的最终读出结果。经过这种行压缩后,最终阶段的计算成本不再依赖于列数。
In-context Learning(上下文学习):最后由一个 Transformer 处理行嵌入。上下文行之间相互注意,而查询行只注意上下文行。因此每次预测只取决于上下文和该行自身,与其他同时被预测的行无关。由于上下文从不关注查询,其 key 和 value 只需计算一次,即可在后续预测中复用。查询行采用 Test-GQA,缩小了每次预测所需读取的缓存。一个输出头将每个查询行转换为分类任务的类别概率,或回归任务的 999 个分位数,进而得到点预测和不确定性估计。
Length-aware Attention Temperature(长度感知注意力温度):随着 key 数量增加,Softmax 注意力会逐渐发散。在几百行数据上保持锐利的注意力,到了几万行时可能就失效了——而这恰恰是推理时表格远大于典型训练表格的场景。因此 Kumo Tabular 对每个 query 乘上一个随 key 数量的对数增长的温度系数,且每个注意力头单独学习各自的系数。这样无论表格变长还是变宽,注意力都能保持锐利。
Kumo Tabular 是如何构建的
Kumo Tabular 完全在人工生成的表格上预训练。每张训练表格都采样自一个 结构因果模型(Structural Causal Model, SCM),过程分为下面六个步骤:
我们首先为整张表生成一个配置,涵盖其规模、任务类型、机制及缺失情况。随后,一个随机因果图将隐变量连接起来,并在从根节点到叶节点的每个节点上通过随机抽取的函数进行评估(例如线性映射、小型神经网络、树模型或高斯过程)。部分节点被设为数值型或类别型列,其中一个作为目标列,其余则保持隐藏状态,类似于真实数据背后未被观测到的成因。后处理阶段会对列组进行相关性构建、截断异常值并注入缺失值,随后通过快速树集成检测剔除那些无可学习信号的表格。由于生成器本质上是程序化采样器而非训练模型,它能无限地产生表格,每个表格都拥有全新的因果图与生成机制。
现实世界中的表格往往杂乱无章,因此我们在生成器中融入了更多这种不完美性。数据值会以多种模式缺失,部分特征会被粗化,导致重复行在标签上可能不一致;某些类别型列包含众多层级,而回归目标值则可能具有厚尾特性。一个经过数百万此类表格训练过的模型,能够在无需任何数据清洗的情况下处理这些不完美之处。
在每个人工生成的表格上,模型将大多数行及其标签作为上下文输入,并学习预测剩余行的标签。分类任务使用交叉熵损失,回归任务使用分位数损失。分类和回归由两个独立的模型分别训练。与 TabICLv2 类似,训练过程分为三个阶段。第一阶段耗时最长,使用包含 1,024 行、最多 100 列的表格,教导模型识别表格的结构特征。第二阶段将上下文长度从 400 行变化至 10,240 行,第三阶段则将其扩展至 60,000 行,列数仍保持最多 100 列。总计,Kumo Tabular-Small/Medium/Large 模型分别观察了约 3,500 万/7,100 万/1.37 亿张人工表格。
我们的训练配方和人工数据生成器将很快发布。
性能表现
我们以默认配置运行了 Kumo Tabular 的三个不同规格,并在完整的 TabArena 排行榜上进行了测试。该排行榜涵盖了经过调优的梯度提升树、AutoGluon 以及最新的表格基础模型。在统一的单块 RTX 6000 Pro 评估环境下,Kumo Tabular 以 1950 的 ELO 得分位居榜首,且推理速度比 LimiX-2 快 17 倍。这三个规格均刷新了表格预测模型在精度与效率 Pareto 前沿上的最先进水平:
我们还在 BeyondArena、TALENT 和 ScoringBench 上评估了 Kumo Tabular。在 BeyondArena 上,Kumo Tabular 获得 1418 的 ELO 得分和 7.78% 的可改进性分数,位居排行榜第一。在 TALENT 上,它在分类准确率、分类对数损失和回归 RMSE 三项指标的综合排名中均列首位,平均排名分别为 6.67、3.98 和 4.22。ScoringBench 是一个针对预测分布的基准测试,Kumo Tabular 的大模型和中模型分别以平均排名第一和第二位列前二。
局限性
Kumo Tabular 仅支持数值型和类别型列,而文本、图像或时间戳等类型可通过内置预处理配方转换为特征。单次前向传播最多覆盖 10 个类别,库通过纠错输出码机制将其扩展至任意数量的类别。当表格数据远超训练范围,或查询行的分布与上下文行显著不同时,模型精度可能会下降。因此,与所有预测模型一样,建议在部署前使用自己的保留数据验证精度和校准效果。
演示
Kumo Tabular 通过 NVIDIA 新发布的 GPU 原生库 structured-data-models 运行。该库会在首次使用时从 Hub 下载权重,并提供我们评测中所用的预处理、集成学习以及多分类处理功能。只需以下几行代码,就能从一个 pandas.DataFrame 得到预测结果:
import sdm # structured-data-models
# Tensorize tabular data:
table = sdm.TableTensor.from_pandas(pd.load_csv(...), device="cuda")
na_mask = table["target"].isnan()
model = sdm.models.KumoTabular(device="cuda")
pred = model(
# In-context examples (features/targets):
x_context=table[~na_mask].drop_columns("target"),
y_context=table[~na_mask, "target"],
# Prediction examples (features):
x_query=table[na_mask].drop_column("target"),
)
开始使用 Kumo Tabular
Kumo Tabular 以 OpenMDW License Agreement 1.1 版发布。NVIDIA 相信可信 AI 是一项共同责任,我们已建立相关政策和实践,支持广泛的 AI 应用开发。在按照我们的服务条款下载或使用模型时,开发者应与相关模型支持团队合作,确保模型满足相关行业和用例的要求,并防范意外的产品误用。如发现模型质量问题、风险、安全漏洞或 NVIDIA AI 相关问题,请在此反馈。


