如何用TabFM完成零样本表格分类?面向新手的10行代码分步实战指南
【免费下载链接】tabfmTabFM (Tabular Foundation Model) is a pretrained tabular foundation model developed by Google Research for tabular data regression and classification.项目地址: https://gitcode.com/gh_mirrors/ta/tabfm
TabFM 零样本表格分类是谷歌研究院开源的表格基础模型(Tabular Foundation Model),它不需要在你的数据上训练,只需"读"一遍训练集,就能直接对新样本做分类预测。本文面向新手,用 10 行代码带你完成第一次零样本表格分类 🚀
什么是 TabFM:不用训练的表格分类器
传统流程:准备数据 → 选模型 → 训练调参 → 评估,动辄几小时。
TabFM 的流程:加载预训练权重 → fit(训练集) → predict(测试集),完事。
它的原理是上下文学习(In-Context Learning):把训练数据当作"上下文"喂给模型,模型据此直接预测新样本,全程不更新任何参数。因此:
- 支持数值列 + 类别列混合的表格,无需手工特征工程
- 适合小数据场景——传统模型数据太少学不动,TabFM 靠预训练知识兜底
- 兼容 scikit-learn 接口,
fit / predict / predict_proba用法与熟悉的ClassifierMixin一致
一键安装步骤
git clone https://gitcode.com/gh_mirrors/ta/tabfm cd tabfm pip install -e .[jax] # JAX 后端(CPU) # pip install -e .[pytorch] # 或 PyTorch 后端环境要求 Python ≥ 3.11。首次load()会自动从 Hugging Face 下载预训练权重,耐心等待即可。
10行代码完成零样本表格分类
以下是核心代码,对应完整可运行脚本见 classification_example.py:
import numpy as np, pandas as pd from tabfm import TabFMClassifier, tabfm_v1_0_0_jax # 1. 加载预训练分类模型 model = tabfm_v1_0_0.load(model_type="classification") clf = TabFMClassifier(model=model) # 2. 准备混合类型表格(数值列 + 类别列) X_train = pd.DataFrame({"age": [25.0, 45.0, 35.0, 50.0], "job": ["eng", "mgr", "eng", "mgr"], "income": [8e4, 12e4, 9e4, 13e4]}) y_train = np.array(["low", "high", "low", "high"]) X_test = pd.DataFrame({"age": [30.0, 48.0], "job": ["eng", "mgr"], "income": [85e3, 125e3]}) # 3. fit 只是准备编码器,predict 立即出结果 clf.fit(X_train, y_train) print(clf.predict(X_test)) print(clf.predict_proba(X_test)) # 各类概率逐行解读
| 步骤 | 代码 | 说明 |
|---|---|---|
| 1️⃣ 加载模型 | tabfm_v1_0_0.load(model_type="classification") | 自动下载并加载 v1.0.0 预训练权重 |
| 2️⃣ 创建分类器 | TabFMClassifier(model=model) | sklearn 风格封装,内部自动做类别编码和数值缩放 |
| 3️⃣ 喂入数据 | clf.fit(X_train, y_train) | 不训练,只准备 Ordinal 编码器、标量归一化等数据变换 |
| 4️⃣ 即时预测 | clf.predict(X_test) | 零样本推理,毫秒级返回 |
PyTorch 用户只需把加载行换成from tabfm import tabfm_v1_0_0_pytorch as tabfm_v1_0_0,其余完全相同(封装类定义在 classifier_and_regressor.py)。
默认模式 vs 集成模式:精度更高一档
TabFMClassifier有两种用法,在 tabarena_classification_example.py 中可看到两者在同一任务上的对比:
# 默认模式:简单平均 logit,速度最快 clf = TabFMClassifier(model=model) # 集成模式:特征交叉 + SVD 特征 + NNLS 加权 + 概率校准 clf = TabFMClassifier.ensemble(model=model)- 默认模式:
n_estimators=32个随机数据视图做 logit 平均,开箱即用 - 集成模式:额外启用特征交叉、SVD 特征、非负最小二乘权重和概率校准,官方 TabArena 评测结果见 results/ 目录下的 parquet 文件
新手建议先用默认模式跑通,追求精度再切换集成模式。
常见限制:表格有多大能用?
TabFM 的上下文窗口是有界的(默认500 个特征、100 行上下文),超大表格会被采样处理而非整体输入,关键参数:
max_num_features:每个集成成员最多使用的特征数(默认 500)max_num_rows:每个集成成员最多采样的行数n_estimators:集成成员数(默认 32)batch_size:控制推理内存
详细 Q&A 见 README.md 的 FAQ 章节;更多参数说明见 classifier_and_regressor.py。
⚠️ 一个重要提醒:许可证
源码为 Apache-2.0,但默认预训练权重采用tabfm-non-commercial-v1.0许可,仅限非商业、非生产用途。商业落地请自行评估合规性。
动手清单 ✅
- 按上文安装 TabFM
- 运行
python examples/classification_example.py验证环境 - 换成自己的 DataFrame,体验
fit + predict两步出结果 - 精度不够时切换到
TabFMClassifier.ensemble
更多细节可阅读 CHANGELOG.md 了解 v1.0.1 的修复内容。祝你玩转零样本表格分类!
【免费下载链接】tabfmTabFM (Tabular Foundation Model) is a pretrained tabular foundation model developed by Google Research for tabular data regression and classification.项目地址: https://gitcode.com/gh_mirrors/ta/tabfm
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考