数值、类别、日期混合列也不怕: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
TabFM 是 Google Research 开源的表格数据基础模型(Tabular Foundation Model),支持零样本分类与回归。它最"省心"的一点是:你的 DataFrame 里数值、类别、日期列随便混着来——调用fit()时,TabFM 会自动完成列类型检测与预处理流水线,无需手动写任何编码或缩放代码。本文带你看懂这条隐藏流水线是怎么工作的。
三步上手:混合类型数据直接喂给它 🚀
以分类任务为例,只需四行核心代码(完整版见 examples/classification_example.py):
from tabfm import TabFMClassifier import tabfm model = tabfm.tabfm_v1_0_0_jax.load(model_type="classification") clf = TabFMClassifier(model=model) X_train = pd.DataFrame({ "age": [25.0, 45.0, 35.0], # 数值列 "job": ["engineer", "manager", "engineer"], # 类别列 "join_date": ["2023-01-05", "2022-07-19", "2021-11-30"], # 日期列(文本) }) clf.fit(X_train, y_train) # 类型检测 + 编码 + 缩放,全部自动完成fit()这一行背后,就藏着本文要拆解的完整流水线。
第一关:列类型自动检测是怎么做到的?
核心逻辑在 tabfm/src/classifier_and_regressor.py 的TransformToNumerical.fit()方法中。对每一列,它按以下优先级判断:
| 判断顺序 | 检测条件 | 归类结果 |
|---|---|---|
| 1 | dtype 本身是datetime64系列 | 日期列 |
| 2 | 文本列中"长得像日期"的值超过 80% | 日期列 |
| 3 | dtype 是数值型 | 数值列 |
| 4 | 其余情况 | 类别列(兜底) |
其中最巧妙的是第 2 条——文本日期的启发式检测(_looks_like_datetime函数):
- 先用
pd.to_numeric试转,能转成数字的列直接排除(避免把"编号"误判为日期); - 对超过 500 行的列,固定种子随机抽样 500 行再做解析,保证检测速度且不依赖模型随机种子(列类型必须稳定,否则换个种子特征结构就变了);
- 用
pd.to_datetime(..., errors="coerce", format="mixed")宽容解析,只要 20% 以上能解析为日期就判定为日期列——即使日期里混了几个脏数据也没关系。
一个贴心细节:开启verbose=True后,fit()会打印每一列被分到了哪个类别,帮你核对检测结果。
第二关:三类列各走各的专属编码器
检测完成后,TabFM 用 sklearn 的ColumnTransformer让三类列并行进入各自的编码分支:
🔢 数值列 → 缺失值填补
数值列走SimpleImputer,自动把缺失值用训练集的均值/众数补齐,你不必提前处理 NaN。
🏷️ 类别列 → 智能序数编码
CategoricalOrdinalEncoder支持三种编码顺序(mode参数):
appearance(默认):按首次出现顺序编号;frequency:按出现频率降序编号,最常见的类别拿到 0 号——对预测任务往往更有效;alphabetical:按字母升序,对齐 sklearnLabelEncoder习惯。
同时内置两个"抗脏数据"设计:
- 稀有类别过滤:默认
min_cat_frequency=2,训练中出现次数不足 2 次的类别直接归为"未知"; - 统一未知码:测试集中出现的新类别和缺失值,统一编码为
-1,模型见过这个标记,不会崩。
📅 日期列 → 一列变五列
DatetimeTransformer会把每个日期列展开为 5 个数值特征:Unix 纳秒时间戳 + 年 + 月 + 日 + 星期几,并统一转 UTC。这样"季节性""周期规律"对模型来说都是显式可见的信号;缺失日期用训练集日期均值填补。
第三关:数值特征精修流水线
编码完成后,数据进入PreprocessingPipeline,依次经过三道工序:
- 标准缩放:
CustomStandardScaler做标准化,并把结果裁剪到 [-100, 100],防止极端值干扰注意力机制; - 可选归一化:可选
power(Yeo-Johnson 幂变换)、quantile(分位数变换)、quantile_rtdl(加噪分位数变换,借鉴表格深度学习研究)、robust(鲁棒缩放)等方法,默认集成none和power两种做集成投票; - 异常值钳制:
OutlierRemover用两阶段 Z-score 法(阈值 4 倍标准差)识别异常值,重算统计量后做对数式裁剪,让极端异常点不再"绑架"整个分布。
此外,UniqueFeatureFilter会顺手删掉"只有一种取值的列"——常数列对模型毫无信息量,白白占位。
为什么要做这么多"花样"?🤔
流水线末尾还有一个EnsembleGenerator:它对同一份编码后的数据做特征重排、类别值置换、类别标签偏移等操作,生成多个不同视角的数据副本,交给多个模型副本分别推理再融合。而前面那些"稳定的类型检测 + 统一编码"正是集成可靠的前提——只有每个副本的特征 schema 完全一致,投票结果才可解释、可复现。
关键文件清单
- 预处理流水线全部源码:tabfm/src/classifier_and_regressor.py
- 分类示例(含混合类型数据):examples/classification_example.py
- 回归示例:examples/regression_example.py
- 安装与快速上手说明:README.md
- 依赖版本清单:requirements.txt
小结
TabFM 的fit()看似只是"喂数据",实际默默完成了类型检测 → 分类编码 → 日期展开 → 缺失填补 → 缩放归一化 → 异常值处理六道工序。对新手来说,这意味着你可以把业务里最常见的"数值 + 类别 + 日期"混合表格直接交给它,把精力留给特征设计本身,而不是数据清洗脚本。
【免费下载链接】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),仅供参考