ARTICLE DETAIL

资讯详情

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

如何用TabFM完成零样本表格分类?面向新手的10行代码分步实战指南

如何用TabFM完成零样本表格分类?面向新手的10行代码分步实战指南 如何用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/tabfmTabFM 零样本表格分类是谷歌研究院开源的表格基础模型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.pyimport numpy as np, pandas as pd from tabfm import TabFMClassifier, tabfm_v1_0_0_jax # 1. 加载预训练分类模型 model tabfm_v1_0_0.load(model_typeclassification) clf TabFMClassifier(modelmodel) # 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_typeclassification)自动下载并加载 v1.0.0 预训练权重2️⃣ 创建分类器TabFMClassifier(modelmodel)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(modelmodel) # 集成模式特征交叉 SVD 特征 NNLS 加权 概率校准 clf TabFMClassifier.ensemble(modelmodel)默认模式n_estimators32个随机数据视图做 logit 平均开箱即用集成模式额外启用特征交叉、SVD 特征、非负最小二乘权重和概率校准官方 TabArena 评测结果见 results/ 目录下的 parquet 文件新手建议先用默认模式跑通追求精度再切换集成模式。常见限制表格有多大能用TabFM 的上下文窗口是有界的默认500 个特征、100 行上下文超大表格会被采样处理而非整体输入关键参数max_num_features每个集成成员最多使用的特征数默认 500max_num_rows每个集成成员最多采样的行数n_estimators集成成员数默认 32batch_size控制推理内存详细 QA 见 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),仅供参考
返回列表