ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

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

TabFM如何无缝集成到现有sklearn项目?TabFMClassifier与TabFMRegressor API使用完全指南 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/tabfmTabFMTabular 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 或 PyTorchTabFM 支持 JAXCPU/GPU与 PyTorchCPU/GPU双后端要求Python ≥ 3.11详见 pyproject.tomlgit 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(modelmodel) # sklearn 风格估计器 clf.fit(X_train, y_train) # 只做特征编码/集成准备不训练 probs clf.predict_proba(X_test) # 类概率 preds clf.predict(X_test) # 类别预测关键点步骤说明load()加载分类权重回归需加model_typeregressionfit(X, y)自动完成类别列序数编码、数值标准化、集成视图构建predict / predict_proba基于多集成成员 概率校准输出结果fit()内部会自动识别日期型文本列、按出现顺序或频率编码类别列appearance/frequency源码见 TransformToNumerical。四、TabFMRegressor 三步走免训练回归预测 回归流程与分类完全对称官方示例在 examples/regression_example.pyfrom tabfm import TabFMRegressor from tabfm import tabfm_v1_0_0_jax as tabfm_v1_0_0 model tabfm_v1_0_0.load(model_typeregression) reg TabFMRegressor(modelmodel) 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默认或frequencyverboseTrue打印每列被判定为数值/类别/日期的分类结果调试利器。追求精度时可一行启用官方ensemble 增强预设特征交叉 SVD 特征 NNLS 加权融合 概率校准clf TabFMClassifier.ensemble(modelmodel) # 见 [ensemble 预设](https://link.gitcode.com/i/f1e7e5aec7e709b4a87d5cd9ab702f7d) reg TabFMRegressor.ensemble(modelmodel) # 回归版预设六、如何替换现有 sklearn 模型Pipeline 中的正确姿势 TabFMClassifier/TabFMRegressor继承自BaseEstimatorClassifierMixin/RegressorMixin因此可以像普通 sklearn 估计器一样放进Pipeline末位from sklearn.pipeline import Pipeline pipe Pipeline([ (drop_cols, ColumnTransformer([(drop, drop, [id])])), (tabfm, TabFMClassifier(modelmodel)), ]) pipe.fit(X_train, y_train)实践建议把 TabFM 放在 Pipeline 最后——它自带完整预处理编码、标准化、异常值裁剪前面只需做列筛选、去重等轻量操作不要给DataFrame留重名列会直接报错提示重命名大表请提前采样或分片官方 FAQ 明确说明上下文窗口有限超过max_num_rows的数据会用采样行推理见 README FAQ换后端只改一行 importtabfm_v1_0_0_jax↔tabfm_v1_0_0_pytorch估计器代码零改动。七、常见问题速查 ✅首次运行很慢JAX 后端首次编译模型执行可能需几分钟属正常现象类别数超限训练类别数超过model.max_classes时fit()会抛出ValueError想要更快推理PyTorch 后端支持cache_contextTrue预缓存上下文 K/V含 int8 量化显著降低重复预测延迟但 JAX 后端暂不支持版本信息当前发布版本见 tabfm/init.py1.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),仅供参考
返回列表