☰
TabFM如何无缝集成到现有sklearn项目?TabFMClassifier与TabFMRegressor API使用完全指南
2026/10/3 1:28:01 网站建设 项目流程

TabFM如何无缝集成到现有sklearn项目?TabFMClassifier与TabFMRegressor API使用完全指南

【免费下载链接】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)是 Google Research 推出的表格数据预训练基础模型,天然兼容scikit-learn接口。通过TabFMClassifier与TabFMRegressor两个核心类,你可以零训练、零调参地对混合类型表格数据做零样本分类与回归——调用方式和RandomForest几乎一样,直接嵌入现有 sklearn 项目。本文给出从安装到 API 调用的完整上手指南。

一、TabFM 是什么:免训练的数据表"大模型" 🧠

传统 sklearn 模型需要你训练;TabFM 走的是另一条路线——上下文学习(In-Context Learning):

  • 推理时不训练任何参数,而是把你的训练数据当作"上下文"喂给模型;
  • 模型读完上下文后,直接对新样本即时预测;
  • 自动处理数值列、类别列、日期列混合的数据表,内置编解码与集成推理管线。

核心 API 定义在 tabfm/init.py 中,两类估计器源码位于 tabfm/src/classifier_and_regressor.py。

二、最快安装步骤:一条命令接入 JAX 或 PyTorch

TabFM 支持 JAX(CPU/GPU)与 PyTorch(CPU/GPU)双后端,要求Python ≥ 3.11,详见 pyproject.toml:

git clone https://gitcode.com/gh_mirrors/ta/tabfm.git cd tabfm # JAX 后端(CPU) pip install -e .[jax] # JAX 后端(GPU) pip install -e .[jax,cuda] # PyTorch 后端(CPU/GPU) pip install -e .[pytorch]

首次load()会自动从 Hugging Face 下载TabFM v1.0.0预训练权重并缓存。

⚠️ 重要提醒:源码是 Apache-2.0 协议,但默认预训练权重受tabfm-non-commercial-v1.0协议约束,仅限非商业、非生产用途。

三、TabFMClassifier 快速上手:5 行代码完成零样本分类 📊

以 examples/classification_example.py 为蓝本,最小可用流程如下:

import numpy as np import pandas as pd from tabfm import TabFMClassifier from tabfm import tabfm_v1_0_0_jax as tabfm_v1_0_0 # 或 tabfm_v1_0_0_pytorch model = tabfm_v1_0_0.load() # 加载预训练权重 clf = TabFMClassifier(model=model) # sklearn 风格估计器 clf.fit(X_train, y_train) # 只做特征编码/集成准备,不训练 probs = clf.predict_proba(X_test) # 类概率 preds = clf.predict(X_test) # 类别预测

关键点:

步骤说明
load()加载分类权重(回归需加model_type="regression")
fit(X, y)自动完成类别列序数编码、数值标准化、集成视图构建
predict / predict_proba基于多集成成员 + 概率校准输出结果

fit()内部会自动识别日期型文本列、按出现顺序或频率编码类别列(appearance/frequency),源码见 TransformToNumerical。

四、TabFMRegressor 三步走:免训练回归预测 📈

回归流程与分类完全对称,官方示例在 examples/regression_example.py:

from tabfm import TabFMRegressor from tabfm import tabfm_v1_0_0_jax as tabfm_v1_0_0 model = tabfm_v1_0_0.load(model_type="regression") reg = TabFMRegressor(model=model) reg.fit(X_train, y_train) predictions = reg.predict(X_test)

内部细节:目标值先做StandardScaler标准化送入模型,输出再逆变换回原始量纲(TabFMRegressor.fit),所以你无需手动缩放 y。

五、核心参数清单:n_estimators 与 ensemble 预设 🔧

两个估计器共享一套集成推理参数(完整参数表):

  • n_estimators(默认 32):集成成员数量,越多越稳但越慢;
  • max_num_features(默认 500)/max_num_rows:单成员特征与上下文行数上限;
  • random_state(默认 42):控制所有随机组件,保证结果可复现;
  • cat_encoder_mode:类别编码顺序,"appearance"(默认)或"frequency";
  • verbose=True:打印每列被判定为数值/类别/日期的分类结果,调试利器。

追求精度时,可一行启用官方"ensemble 增强预设"(特征交叉 + SVD 特征 + NNLS 加权融合 + 概率校准):

clf = TabFMClassifier.ensemble(model=model) # 见 [ensemble 预设](https://link.gitcode.com/i/f1e7e5aec7e709b4a87d5cd9ab702f7d) reg = TabFMRegressor.ensemble(model=model) # 回归版预设

六、如何替换现有 sklearn 模型:Pipeline 中的正确姿势 🔌

TabFMClassifier/TabFMRegressor继承自BaseEstimator+ClassifierMixin/RegressorMixin,因此可以像普通 sklearn 估计器一样放进Pipeline末位:

from sklearn.pipeline import Pipeline pipe = Pipeline([ ("drop_cols", ColumnTransformer([("drop", "drop", ["id"])])), ("tabfm", TabFMClassifier(model=model)), ]) pipe.fit(X_train, y_train)

实践建议:

  1. 把 TabFM 放在 Pipeline 最后——它自带完整预处理(编码、标准化、异常值裁剪),前面只需做列筛选、去重等轻量操作;
  2. 不要给DataFrame留重名列,会直接报错提示重命名;
  3. 大表请提前采样或分片,官方 FAQ 明确说明上下文窗口有限,超过max_num_rows的数据会用采样行推理(见 README FAQ);
  4. 换后端只改一行 import:tabfm_v1_0_0_jax↔tabfm_v1_0_0_pytorch,估计器代码零改动。

七、常见问题速查 ✅

  • 首次运行很慢?JAX 后端首次编译+模型执行可能需几分钟,属正常现象;
  • 类别数超限?训练类别数超过model.max_classes时fit()会抛出ValueError;
  • 想要更快推理?PyTorch 后端支持cache_context=True预缓存上下文 K/V(含 int8 量化),显著降低重复预测延迟,但 JAX 后端暂不支持;
  • 版本信息:当前发布版本见 tabfm/init.py(1.0.1),变更记录见 CHANGELOG.md。

八、项目文件导航

文件作用
tabfm/src/classifier_and_regressor.pyTabFMClassifier/TabFMRegressor及全部预处理组件
tabfm/src/jax/tabfm_v1_0_0.pyJAX 后端权重加载入口
tabfm/src/pytorch/tabfm_v1_0_0.pyPyTorch 后端权重加载入口
examples/classification_example.py分类可运行示例
examples/regression_example.py回归可运行示例
results/官方评测结果(parquet 格式)

总结:TabFM 让"基础模型"第一次以标准 sklearn 估计器的形态进入表格数据工作流——load → TabFMClassifier/TabFMRegressor → fit → predict,四步即可把零样本预测能力嫁接进你现有的 Pipeline。

【免费下载链接】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),仅供参考

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询