6个调参旋钮提升TabFM预测精度:n_estimators集成、特征交叉、SVD与概率校准完全指南
【免费下载链接】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 开发的预训练表格数据模型,支持零样本分类与回归。它不训练、不微调,而是把训练集当作"上下文"读入模型,通过**上下文学习(in-context learning)**对新样本即时预测。想提升 TabFM 预测精度?无需训练,只需调好 6 个"旋钮"——n_estimators集成数量、特征交叉、SVD 结构特征、概率校准、NNLS 加权融合与上下文窗口控制,即可在不改变模型权重的前提下显著改善效果。
TabFM 工作原理:为什么"调参"就够
TabFM 提供 scikit-learn 风格的两个估计器,直接fit+predict就能用:
TabFMClassifier:分类器,见 TabFMClassifierTabFMRegressor:回归器,见 TabFMRegressor
它的精度来自集成推理:内部类 EnsembleGenerator 会生成多份"数据视图"(不同的特征顺序、归一化方式、行采样、类别扰动),每个视图独立预测,再汇总成最终结果。你调的参数,本质上是控制这些视图怎么生成、怎么汇总。
6个旋钮总览
| 旋钮 | 默认值 | 作用 | 推荐调整场景 |
|---|---|---|---|
n_estimators | 32 | 集成成员数(数据视图数) | 基线,一般不动 |
n_feature_crosses | 0(关闭) | 随机生成特征交叉列 | 特征间存在交互效应 |
n_svd_features | 0(关闭) | 生成 SVD 降维结构特征 | 特征数多、数据冗余高 |
binary/multiclass_calibration_method | None | 概率校准(Platt / vector) | 关心 AUC、log loss 等概率指标 |
enable_nnls/nnls_beta | False/0.75 | NNLS 加权融合各成员 | 默认均匀平均效果不佳 |
max_num_rows/max_num_features | None/500 | 控制上下文窗口大小 | 大表推理、内存不足 |
旋钮1:n_estimators——集成成员数量
每个"成员"是一份独立的训练数据视图(不同特征顺序 + 归一化 + 行采样),预测结果取平均。成员越多,视图越多样,预测越稳健,但推理耗时线性增长。
- 默认值 32是精度与速度的平衡点,官方预设保持不变
- 数据量小(几百行)时降到 8~16 可显著提速
- 追求极限精度且有算力时,可尝试 64
旋钮2:n_feature_crosses——随机特征交叉
设为"sqrt"后,每个集成成员会额外获得 √(特征数) 个随机特征交叉列(两列相乘),让模型能"看到"类似面积 × 层数这样的交互信号,而不必自己学。
- 偶数索引成员不加交叉列、奇数索引成员加满,形成增强/非增强混合视图,这正是集成多样性的来源(见 _get_member_n_features_list)
- 何时开启:业务上存在明显的特征交互(如面积×单价),或表格较宽
- 何时关闭:特征很少(<5 个)或推理要快
旋钮3:n_svd_features——SVD 结构特征
同样支持"sqrt"档位:内部对预处理后的训练表做 TruncatedSVD,把前若干主成分作为新列拼回特征,相当于免费获得一组"数据压缩视角"。
- 适合:特征数多、相关性高、列冗余的宽表
- 配合
total_svd_pool可控制 SVD 特征池总量,避免生成过多
旋钮4:概率校准——让概率"说真话"
分类器提供两种校准方法(见 TabFMClassifier 参数):
binary_calibration_method="platt":二分类问题的 Platt 缩放,改善校准曲线multiclass_calibration_method="vector":多分类的逐类向量校准
校准在验证集/OOF 交叉预测上学习,用calibration_lambda(默认1e-2)做 L2 正则防过拟合。如果你的业务指标是 log loss、AUC 或需要可解释的风险概率,强烈建议开启;只关心分类准确率的可以不动。
旋钮5:enable_nnls / nnls_beta——NNLS 加权融合
默认情况下 32 个成员均匀平均。enable_nnls=True会在验证集上用非负最小二乘(NNLS)学习每个成员的融合权重,让"更靠谱的视图"贡献更大;nnls_beta(默认0.75)控制学习权重与均匀权重的混合比例。
⚠️ 注意:enable_nnls与average_logits=True(默认)互斥,开启 NNLS 时自动切换为概率平均(见 参数校验)。学习融合权重的数据量较大时(默认阈值 2000 行)才划算。
旋钮6:上下文窗口——max_num_rows 与 max_num_features
TabFM 靠上下文学习,上下文越长信息越多,但内存和耗时越高:
max_num_features(默认 500):每个成员采样的特征数上限max_num_rows(默认不限):每个成员读入的训练行数上限,大表必须设置,否则整张表都会进入上下文
此外norm_methods控制每个成员的归一化方式(默认["none", "power"],可选quantile、robust等,见 PreprocessingPipeline),feat_shuffle_method与class_shift则控制特征重排和类别标签扰动带来的多样性。
一步到位:ensemble() 预设
不想逐个调?官方提供了预设工厂方法 TabFMClassifier.ensemble 和 TabFMRegressor.ensemble,一次开启:n_feature_crosses="sqrt"+n_svd_features="sqrt"+enable_nnls=True+ 概率校准(分类器)。回归任务对比示例见 examples/tabarena_regression_example.py,分类任务对比见 examples/tabarena_classification_example.py。
实践建议:按场景选配置
| 场景 | 推荐配置 |
|---|---|
| 快速基线评估 | 默认TabFMClassifier(model=...),不动任何旋钮 |
| 追求精度上限 | 直接用ensemble()预设 |
| 大表(>10万行) | 默认预设 +max_num_rows=10000控制内存 |
| 特征很少(<10列) | 保持默认,关闭特征交叉/SVD 以免过增强 |
| 风控/信贷等概率敏感场景 | 显式开启platt/vector校准 |
快速上手
git clone https://gitcode.com/gh_mirrors/ta/tabfm cd tabfm pip install -e .[jax] # 或 pip install -e .[pytorch] python examples/classification_example.py入门示例可参考 examples/classification_example.py 与 examples/regression_example.py;完整参数说明见 README.md 和 CHANGELOG.md。
💡 提示:默认
load()下载的预训练权重受tabfm-non-commercial-v1.0许可约束,仅限非商业用途;商用场景请留意许可条款。
总结
TabFM 是"零训练"的表格基础模型,精度优化不靠训练,而靠集成推理的 6 个旋钮:
n_estimators:控制集成规模,默认 32 已够用n_feature_crosses="sqrt":为宽表补充交互信号n_svd_features="sqrt":为高维冗余表补充结构视角- 概率校准(
platt/vector):让概率输出可信赖 enable_nnls:用数据说话,加权融合各成员max_num_rows/max_num_features:给大表装上限,稳住内存
多数任务从默认配置出发,再按需启用ensemble()预设,就是最划算的调参路径 🚀
【免费下载链接】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),仅供参考