google/tabfm-1.0.0-pytorch

Hugging Face Models Trending 模型

摘要

Google Research发布TabFM,这是一个使用PyTorch的零样本表格基础模型,用于分类和回归,无需微调。

任务:表格分类 标签:tabfm, 表格, 表格回归, 零样本, 上下文学习, pytorch, 基础模型, 表格分类, 许可证:其他, 地区:美国
查看原文
查看缓存全文

缓存时间: 2026/07/02 17:36

google/tabfm-1.0.0-pytorch · Hugging Face

来源:https://huggingface.co/google/tabfm-1.0.0-pytorch TabFM 是谷歌研究推出的零样本表格基础模型。它支持对包含混合数值列和类别列的结构化/表格数据进行分类和回归,无需微调或超参数搜索——训练样本作为上下文传入,预测在一次前向传播中完成。

此仓库包含 PyTorch 权重。JAX/Flax 权重请参见 google/tabfm-1.0.0-jax (https://huggingface.co/google/tabfm-1.0.0-jax)。

https://huggingface.co/google/tabfm-1.0.0-pytorch#getting-started 入门

pip install tabfm[pytorch]

分类:

from tabfm import TabFMClassifier, tabfm_v1_0_0_pytorch as tabfm_v1_0_0

model = tabfm_v1_0_0.load(model_type="classification")
clf = TabFMClassifier(model=model)
clf.fit(X_train, y_train)
probs = clf.predict_proba(X_test)

回归:

from tabfm import TabFMRegressor, tabfm_v1_0_0_pytorch as tabfm_v1_0_0

model = tabfm_v1_0_0.load(model_type="regression")
reg = TabFMRegressor(model=model)
reg.fit(X_train, y_train)
preds = reg.predict(X_test)

你也可以直接通过 HuggingFace Hub API 加载:

from tabfm.src.pytorch.model import TabFM

clf_model = TabFM.from_pretrained("google/tabfm-1.0.0-pytorch", subfolder="classification")
reg_model = TabFM.from_pretrained("google/tabfm-1.0.0-pytorch", subfolder="regression")

https://huggingface.co/google/tabfm-1.0.0-pytorch#available-checkpoints 可用检查点

子文件夹任务is_classifier
classification/分类(最多10个类别)True
regression/回归False

https://huggingface.co/google/tabfm-1.0.0-pytorch#developers-and-affiliations 开发者与所属机构

由谷歌研究 (https://research.google/) 团队开发。

https://huggingface.co/google/tabfm-1.0.0-pytorch#intended-use 预期用途

  • 包含数值列和/或类别列的表格数据
  • 二分类和多分类(最多10个类别)
  • 连续目标变量的回归
  • 零样本推理:无需特定数据集训练或超参数调优
  • 支持 DataFrame (pandas) 或 numpy 数组

https://huggingface.co/google/tabfm-1.0.0-pytorch#not-intended-for 非预期用途

  • 图像、音频、视频或原始文本
  • 超过10个输出类别(模型硬限制)
  • 需要任务特定微调的任务
  • 非表格结构化数据(图、序列)
  • 商业用途(请参见下方许可)

https://huggingface.co/google/tabfm-1.0.0-pytorch#model-architecture 模型架构

TabFM 使用交替的行注意力和列注意力来捕捉特征交互和行级模式:

  1. 列注意力(Set Transformer):利用傅里叶特征和每个分组的线性投影嵌入每个单元格,然后通过诱导自注意力跨行聚合。
  2. 行压缩:CLS 令牌通过行级注意力(使用旋转位置编码 RoPE)将每一行汇总为稠密向量。
  3. ICL Transformer:一个 24 块因果 Transformer 对压缩后的行向量进行操作,将训练行视为上下文,并为测试行输出预测。

关键超参数:

参数
嵌入维度256
列注意力块3(4头,256个诱导点)
行注意力块3(8头,8个CLS令牌)
ICL Transformer 块24(8头)
前馈因子4
最大类别数10
激活函数SwiGLU
傅里叶特征32个频率

https://huggingface.co/google/tabfm-1.0.0-pytorch#training-data-and-priors 训练数据与先验

TabFM 是在数亿个使用结构因果模型 (SCMs) 动态生成的合成数据集上训练的。选择合成数据是因为缺乏多样、高质量的开源表格数据集,并且为了避免真实工业数据带来的隐私/许可问题。SCM 先验编码了关于表格任务中常见因果结构和特征关系的归纳偏差。

https://huggingface.co/google/tabfm-1.0.0-pytorch#performance 性能

TabFM 在 TabArena (https://tabarena.ai/) 上进行了评估,涵盖 51 个数据集(38 个分类,13 个回归)。在零样本模式下——单次前向传播,无超参数搜索——TabFM 优于经过大量调优的有监督基线方法,包括梯度提升树。TabFMClassifier.ensemble() 预设(特征交叉、SVD 特征、NNLS 混合)可实现更进一步的提升。

完整的基准测试详情请参见谷歌研究博客文章 (https://research.google/blog/introducing-tabfm-a-zero-shot-foundation-model-for-tabular-data/)。

https://huggingface.co/google/tabfm-1.0.0-pytorch#ethical-considerations 伦理考量

TabFM 完全在合成数据上训练。在特定真实世界领域、少数群体或边缘分布上的性能尚未完全表征。用户应在代表其使用场景的留出数据上评估模型,然后再在高风险环境中部署。

https://huggingface.co/google/tabfm-1.0.0-pytorch#limitations 局限性

  • 分类最多10个类别(架构硬限制)
  • 内存使用量随训练行数扩展(所有行都作为上下文传入)
  • 针对最多500个特征的表进行了优化;在非常宽的表格上表现可能下降
  • 不保证在所有数据集上都能匹配特定任务微调后的模型
  • 非谷歌官方支持产品

https://huggingface.co/google/tabfm-1.0.0-pytorch#license 许可

此仓库中的模型权重根据 TabFM 非商业许可 v1.0 发布——请参见 LICENSE (https://huggingface.co/google/tabfm-1.0.0-pytorch/tree/main/LICENSE)。源代码通过 google-research/tabfm (https://github.com/google-research/tabfm) 采用 Apache 2.0 许可。

https://huggingface.co/google/tabfm-1.0.0-pytorch#version 版本

1.0.0

https://huggingface.co/google/tabfm-1.0.0-pytorch#citation 引用

@article{tabfm2026,
  title  = {TabFM: A Zero-Shot Foundation Model for Tabular Data},
  author = {Google Research},
  year   = {2026},
  url    = {https://research.google/blog/introducing-tabfm-a-zero-shot-foundation-model-for-tabular-data/}
}

相似文章