ARTICLE DETAIL

资讯详情

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

零训练秒出结果:TabFM表格大模型安装与快速入门完整教程(JAX/PyTorch双后端10分钟上手)

零训练秒出结果:TabFM表格大模型安装与快速入门完整教程(JAX/PyTorch双后端10分钟上手) 零训练秒出结果TabFM表格大模型安装与快速入门完整教程JAX/PyTorch双后端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 是什么为什么值得你花 10 分钟上手TabFMTabular Foundation Model表格数据基础大模型是由 Google Research 开发的预训练表格大模型专为分类与回归任务设计。与传统机器学习不同TabFM 采用上下文学习in-context learning机制推理时无需任何训练参数它直接把你的训练集当作上下文读入对新样本零训练、秒出预测结果开箱即用。它的核心亮点零训练推理不需要梯度下降不需要调参训练fit()只是准备编码器兼容 scikit-learnTabFMClassifier/TabFMRegressor与 sklearn 接口一致平滑迁移混合列类型数值列 类别列可以直接混在一个 DataFrame 里喂进去⚙️JAX / PyTorch 双后端按你的技术栈自由选择JAX 还支持 GPU/TPU 加速⚠️许可证提醒源码为 Apache-2.0但预训练权重tabfm_v1_0_0.load()自动下载受tabfm-non-commercial-v1.0限制仅限非商业、非生产环境使用。二、TabFM 安装指南一键安装步骤JAX/PyTorch 双后端环境要求依赖项要求Python≥ 3.11Hugging Face Hub用于自动下载预训练权重JAX 后端jax0.10.1 flax0.12.7flax.nnx APIPyTorch 后端torch2.12.1CPU 或对应 CUDA 的 GPU 版本完整依赖版本清单见 requirements.txt。最快安装方法克隆仓库后只需一条 pip 命令按后端选择即可JAX 后端CPUgit clone https://gitcode.com/gh_mirrors/ta/tabfm cd tabfm pip install -e .[jax]JAX 后端GPU 加速pip install -e .[jax,cuda]PyTorch 后端CPU/GPUpip install -e .[pytorch] 小提示如果你要用 PyTorch GPU请确认已先安装与你的 CUDA 版本匹配的 PyTorch再安装 TabFM。三、TabFM 快速入门3 步完成零训练分类预测以 TabFM v1.0.0 为例分类任务只需 3 步加载模型 → fit → predict。import numpy as np import pandas as pd from tabfm import TabFMClassifier # 第 1 步加载预训练模型自动从 Hugging Face 下载权重 # JAX 后端 from tabfm import tabfm_v1_0_0_jax as tabfm_v1_0_0 model tabfm_v1_0_0.load() # PyTorch 后端则改为 # from tabfm import tabfm_v1_0_0_pytorch as tabfm_v1_0_0 # 第 2 步初始化 sklearn 兼容分类器 clf TabFMClassifier(modelmodel) # 第 3 步准备数据数值 类别混合列 X_train pd.DataFrame({ age: [25.0, 45.0, 35.0, 50.0], job: [engineer, manager, engineer, manager], income: [80000, 120000, 90000, 130000], }) y_train np.array([low_risk, high_risk, low_risk, high_risk]) X_test pd.DataFrame({ age: [30.0, 48.0], job: [engineer, manager], income: [85000, 125000], }) clf.fit(X_train, y_train) # 仅准备编码器/标度器不做训练 predictions clf.predict(X_test) # 秒出类别 probabilities clf.predict_proba(X_test) # 类别概率注意fit()在这里不是训练而是准备序数编码器和数值标度器真正的学习发生在模型把训练集读入上下文的推理过程中。回归任务预测房价示例回归用法与分类几乎一致只需把模型类型指定为regressionfrom tabfm import TabFMRegressor model tabfm_v1_0_0.load(model_typeregression) reg TabFMRegressor(modelmodel) X_train pd.DataFrame({ sqft: [1200, 2500, 1500, 3000], neighborhood: [A, B, A, C], }) y_train np.array([250000, 550000, 310000, 620000]) reg.fit(X_train, y_train) print(reg.predict(X_test)) # 输出预测房价 可运行的完整脚本在 examples/ 目录下直接python examples/classification_example.py或python examples/regression_example.py即可执行修改文件内注释可切换 JAX / PyTorch 后端分类示例examples/classification_example.py回归示例examples/regression_example.pyTabArena 基准示例examples/tabarena_classification_example.py、examples/tabarena_regression_example.py四、进阶技巧控制上下文窗口与推理性能TabFM 通过有界的上下文窗口做上下文学习因此超大表格建议先采样或拆分。sklearn 估计器通过 4 个关键参数控制实际行为参数默认值作用max_num_features500最大特征数上限max_num_rows100上下文行数上限n_estimators—对多组采样上下文做集成提升稳定性inference_batch_size—控制推理显存/内存占用如果数据集超过上限TabFM 会用采样后的上下文行工作而非一次吞下整张表。模型评估基准结果TabArena 分类/回归、单模型与集成版保存在 results/ 目录下的 parquet 文件中可用于效果对比。五、常见问题 FAQQ1第一次运行会很慢JAX 首次运行会编译内核可能耗时几分钟属正常现象后续调用会快很多。Q2有技术报告或论文吗目前仓库尚未附带技术报告发布后 README.md 会更新链接。Q3如何运行单元测试# 全量测试需同时安装 JAX 和 PyTorch PYTHONPATH. python3 -m unittest discover -s tabfm/src/ -p *_test.py # 单文件测试 PYTHONPATH. python3 -m unittest tabfm/src/pytorch/model_test.py六、项目源码导航模块说明tabfm/src/classifier_and_regressor.pysklearn 兼容的TabFMClassifier/TabFMRegressor核心实现tabfm/src/jax/JAX 后端模型定义model.py、v1.0.0 加载tabfm_v1_0_0.py、内存高效注意力memory_efficient_attention.pytabfm/src/pytorch/PyTorch 后端模型定义与 v1.0.0 加载tabfm/src/hugging_face/Hugging Face 权重转换与上传工具requirements.txt各后端锁定的依赖版本从克隆仓库到跑通第一次预测全程不超过 10 分钟——TabFM 把预训练 上下文学习带进了表格数据领域让你彻底告别调参熬夜零训练秒出结果。【免费下载链接】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),仅供参考
返回列表