TabPFN 完整使用指南:小数据表格任务,fit 完直接出预测
【免费下载链接】TabPFN⚡ TabPFN: Foundation Model for Tabular Data ⚡项目地址: https://gitcode.com/GitHub_Trending/ta/TabPFN
如果你的表格只有几千行,还在为调参、交叉验证和特征工程反复折腾,可以试试 TabPFN。它是一个表格数据基础模型:预训练好的 Transformer 直接吃原始表格,小样本分类和回归一次前向传播就能出结果,不用训练、不用手调,还能直接处理缺失值,连缩放和 one-hot 都省了。下面带你看清安装、最小可跑示例、性能参数和一套避坑清单。
小数据场景,TabPFN 强在哪
小样本机器学习一直是传统方法的弱项:数据少,模型就靠调参和特征工程硬扛。TabPFN 把"训练"换成了"推理"——模型在海量合成数据上预训练完毕,你本地只需把训练集编码一次,预测时一次前向传播拿结果。
在小数据集(<10,000 样本)上对比传统方法,大致是:
- 准确率提升 15~25%
- 训练时间减少 90% 以上
- 不做特征工程也能拿到不错的基线
所以它常被用在这些地方:医疗诊断预测(样本贵)、金融风险评估(历史数据有限)、科学实验分析(采集成本高)、快速原型开发(要即时结果)。
安装与环境:一条 pip 命令的事
⚡ 先装:
pip install tabpfnTabPFN 安装之后,再对齐几条环境事实:
- 需要 Python 3.10+,3.10~3.14 均支持
- 建议用 GPU:约 8GB 显存就够,更大的数据集建议 16GB;macOS 的 Apple Silicon 自带 GPU 支持
- CPU 也能跑,但只适合小数据:默认 TabPFN-3 最多 5000 样本,老版本 1000
- 首次
fit会自动下载模型权重;TabPFN-2.5/2.6/3 首次使用会弹浏览器登录授权,token 本地缓存,只登一次 - 无浏览器的 CI 环境:设置环境变量
TABPFN_TOKEN即可跳过;Linux/Windows 的 AMD 显卡需要先装 ROCm 版 PyTorch 再装 TabPFN
想完全离线部署,克隆仓库后跑一次官方脚本拉全所有权重:
git clone https://gitcode.com/GitHub_Trending/ta/TabPFN python TabPFN/scripts/download_all_models.py最小可跑示例:分类 10 行、回归 3 行
分类(二分类、多分类同一个入口):
from sklearn.datasets import load_breast_cancer from sklearn.model_selection import train_test_split from tabpfn import TabPFNClassifier X, y = load_breast_cancer(return_X_y=True) X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.5) clf = TabPFNClassifier() clf.fit(X_train, y_train) # 首次运行自动下载权重 print(clf.predict(X_test)) # 类别标签 print(clf.predict_proba(X_test)) # 各类概率回归只换入口,协议完全一致:
reg = TabPFNRegressor(); reg.fit(X_train, y_train); reg.predict(X_test)更完整的脚本都在 examples/ 里,按需取用:
- 基础三件套:
tabpfn_for_binary_classification.py、tabpfn_for_multiclass_classification.py、tabpfn_for_regression.py - 微调:
finetune_classifier.py、finetune_regressor.py - KV 缓存加速预测:
kv_cache_fast_prediction.py - 超参调优:
tabpfn_classifier_with_tuning.py、tabpfn_regressor_with_tuning.py
默认用的是 TabPFN-3。想指定其它版本,一行切换:
from tabpfn.constants import ModelVersion clf = TabPFNClassifier.create_default_for_version(ModelVersion.V2_6)注意许可证:v2 权重是 Apache 2.0;2.5/2.6/3 权重为非商业许可,商用前先确认。
性能调优:批量预测、fit_mode 与数据规模上限
抓住三件事,速度就上去了:
- predict 一次喂完。每次
predict都会重算训练集表示,把 100 个样本拆成 100 次调用,耗时和开销约为一次调用的 100 倍。测试集很大时,按每块 1000 行分块。 - 别再额外预处理。不要缩放、不要 one-hot,原始列直接喂——TabPFN 内部有完整的预处理流水线,见 src/tabpfn/preprocessing/。
- 重复预测就开 KV 缓存。对同一训练集反复预测时,用
fit_mode="fit_with_cache"把训练集表示缓存住,examples/kv_cache_fast_prediction.py有完整演示。
规模上限(行数 × 特征数):
| 版本 | 推荐规模 |
|---|---|
| TabPFN-3(默认) | 1,000,000×200 / 100,000×2,000 / 1,000×20,000 |
| TabPFN-2.6 | 100,000 行、2,000 特征以内 |
| CPU 模式 | TabPFN-3 最多 5000 样本;老版本 1000 |
数据略超上限时,可以子采样,或者设ignore_pretraining_limits=True跳过体积护栏。
避坑清单:常见报错、环境变量与模型保存 🧯
按出现频率排,几个高频问题和对应处理:
- 加载模型报 pickle 错误→ 先
pip install tabpfn --upgrade,还报错就检查权重文件是否下载完整,必要时重新下载 - GPU 显存不足→ 换
device="cpu",或设置PYTORCH_CUDA_ALLOC_CONF="max_split_size_mb:512" - CPU 上慢→ 属正常现象,CPU 只适合小数据,大规模请直接上 GPU
- 数据里有缺失值?→ 支持,原始数据直接喂,不需要自己填
- fit_with_cache 模式 OOM→ 调小
TABPFN_MAX_BATCHED_TEST_ROWS(默认 32768)
常用的环境变量(export 或写进.env都行):
TABPFN_TOKEN:手动提供授权 token,用于无浏览器环境TABPFN_MODEL_CACHE_DIR:自定义权重缓存目录TABPFN_ALLOW_CPU_LARGE_DATASET=true:允许 CPU 跑超规模数据(会很慢)TABPFN_NO_BROWSER:禁止自动弹浏览器登录TABPFN_MAX_BATCHED_TEST_ROWS:缓存模式下单次前向最多喂多少测试行
训好的模型想跨进程复用,直接存盘:
from tabpfn.model_loading import save_fitted_tabpfn_model, load_fitted_tabpfn_model save_fitted_tabpfn_model(reg, "my_reg.tabpfn_fit") reg_loaded = load_fitted_tabpfn_model("my_reg.tabpfn_fit")完整演示在examples/save_and_load_model.py。
项目结构与生态:从示例目录到源码
核心代码在 src/tabpfn/,按功能分块:
classifier.py/regressor.py:分类器、回归器两大入口- architectures/:各代 Transformer 实现(TabPFN-2 / 2.5 / 2.6 / 3 / 3.5)
- preprocessing/:数据清洗与 Torch 加速的预处理流水线
- finetuning/:微调训练工具链
- scripts/:
download_all_models.py(离线拉权重)、convert_checkpoint_to_safetensors.py(权重格式转换)
核心之外还有一圈生态,可以按需取用:
- TabPFN Client:云端推理客户端,免硬件投入、自动扩缩、免维护
- TabPFN Extensions:
pip install tabpfn-extensions,提供 SHAP 解释与特征重要性、异常检测与合成数据、embeddings 提取、超多分类别等工具 - TabPFN UX:无代码图形界面,适合业务同学体验与原型验证
本地部署数据不出内网、可离线、可魔改;云端 API 零硬件投入、自动扩容——按你的合规和成本要求选即可。跑通之后,从 examples 挑一个脚本改改,基本就是你能交付的版本了。
【免费下载链接】TabPFN⚡ TabPFN: Foundation Model for Tabular Data ⚡项目地址: https://gitcode.com/GitHub_Trending/ta/TabPFN
创作声明:本文部分内容由AI辅助生成(AIGC),仅供参考