如果你以为 mask modeling 就是把 BERT 的[MASK]token 移植到表格数据上,我可以先给你一个反面教材。去年我做电商订单的多任务建模,直接用标准 Transformer encoder 把每一行样本拼成一个"句子",随机把某些特征位替换成特殊 token,再做一个统一的重建任务。跑了三十个 epoch,下游 AUC 比直接用 LightGBM 还低了两个点。后来我认真拆了 LimiX-2 这个面向表格模型的 masked modeling 开源框架,才发现问题根本不在模型容量,而在于三个被忽略的细节:表格数据没有天然序列结构,位置编码反而是负资产;数值列和类别列的分布差异太大,单一输出头根本扛不住;更重要的是,掩码策略不能按 NLP 的随机 token 思路来,列和列之间的相关性结构完全不一样。这篇文章会围绕 LimiX-2 把这些点一一展开,重点是我在真实数据集上跑通它时用到的架构细节、掩码策略、数据工程处理和踩坑记录,适合正在做表格预训练、结构化数据竞赛,或者想给多个下游任务统一特征底座的人参考。
1. 为什么表格数据不能照搬 NLP 的掩码建模
1.1 表格不是句子,列之间没有"语序"
NLP 里掩码建模能成立,核心前提是语言具有强烈的顺序依赖和局部共现规律。你在 BERT 里把一句话中间的词挖掉,模型可以从左右两侧的上下文推断出被遮住的词,这是自然语言本身的冗余性带来的。表格数据完全相反。一张表的列顺序通常是业务字段的摆放顺序,比如用户 ID、注册时间、消费金额、商品类别,这个顺序不携带任何语义。你把"消费金额"这一列从字段顺序上挖掉,左右两边的列并不能给模型提供类似"语法成分"的信号。
我在实际尝试中的体会是:如果强行把一行样本压成一个序列,再加绝对位置编码,模型会把"第 3 列"和"第 7 列"这种位置身份当作规律去记忆,而不是去理解列名背后的业务含义。结果就是预训练阶段 loss 降得很快,但下游效果一塌糊涂。LimiX-2 的第一步改变就是不把一行当句子,而是把所有列看作一个无序的特征集合,每列是一个独立的 token,列与列之间的交互完全交给 attention 去学,不再假设左右顺序有信息。
1.2 LimiX-2 的设计起点:列级掩码而不是格子级掩码
LimiX-2 的掩码单位是"列",不是"格子"。这里需要区分一下:如果掩码的是一个单元格,模型很容易通过同一列里其他行的数值猜出结果,那就退化成了普通的缺失值插补,学不到列间依赖。LimiX-2 的做法是,每次迭代随机选择一定比例的列,把这一列在一整批样本上的 token 都替换成[MASK],让模型必须从其他没有被掩码的列去重建它。
这个设计对应到业务场景非常自然:你有一张用户表,掩掉"月消费金额",模型只能通过"用户等级""最近登录时间""历史订单数"等列去推断金额区间。这正是我们做特征工程时想要的跨列推理能力。我也是从这里才理解,masked modeling 在表格上的价值不是让模型学会"填空",而是让模型在重建的过程中被迫去刻画列与列之间的联合分布,这份分布就是后续所有任务的共享底座。
2. LimiX-2 的骨架拆解:列级 Tokenizer 与 Transformer Encoder
2.1 一行样本如何变成一组 token
LimiX-2 的输入构造可以理解成:每一行样本最终变成一组列向量,矩阵形状是[列数, hidden_size],送入 Transformer encoder。数值列、类别列都有各自的特征编码方式,编码后得到的向量不再区分你是数字还是字符串,后续全部走统一的 attention。
这里有个容易忽略的细节:LimiX-2 会给每个 token 叠加两种元信息。第一种是列类型 embedding,比如这个 token 来自数值列还是类别列;第二种是列索引 embedding,让模型知道当前 token 对应业务上的哪一列。但这跟 NLP 位置编码有本质区别,列索引 embedding 只是给模型一个"身份标识",并不参与序列顺序计算,所以不会引入虚假的先后关系。
2.2 数值列和类别列的编码方式差异
数值列的处理是 LimiX-2 里最值得琢磨的部分。作者没有直接做标准化然后过一个 Linear 层,而是先把每个数值特征单独做分位数变换,再通过分段线性映射到固定维度的向量。原因后面第四节我会详细说。类别列则走常规的 embedding 查表,但 LimiX-2 对低基数和高基数列分别做了处理:低基数列用完整词典,高基数列用 frequency-based 截断和哈希分桶,避免 embedding 矩阵过大。
在代码层面,一个典型的初始化大概是这样的:
from limix2 import LimixTabularModel, MaskConfig model = LimixTabularModel( hidden_size=256, num_layers=6, num_heads=8, numeric_encoding="quantile", categorical_encoding="hash_freq", mask_ratio=0.25, mask_mix=(0.7, 0.2, 0.1), ) model.fit_pretrain(X_train, batch_size=256, epochs=30)mask_mix是三路掩码的比例,我后面会解释三种掩码分别是什么。这段配置在我自己的数据集上基本能直接跑,hidden_size 256、6 层 encoder 对绝大多数中大规模表格也够用,不需要一上来就堆大模型,表格数据的 token 数量通常只有几十到几百,不像文本动辄上千 token。
2.3 [MASK] 的语义要跟缺失值分开
LimiX-2 在 token 层面额外维护了两个标志位:is_missing和is_masked。前者表示原始数据里这一列本来就是空的,后者表示是我们人为掩码的。这两个标志会被拼进 token 向量里,但模型在 loss 计算时会忽略真实缺失的格子,只对"被掩码且原本有值"的位置做重建。
这个设计非常关键。如果你把真实缺失也当作掩码目标,模型会逐渐学会一个错误对应关系:缺失的位置就等于要预测的位置,到了推理阶段看到真实缺失值时会启动一种很奇怪的"预测模式",而不是把它当作缺失处理。我在早期实验里就踩过这个坑,后面第四节会细说。
3. 三种掩码策略与训练目标:从直觉到参数
3.1 Random masking、Correlated masking 和 Mixture masking
LimiX-2 提供了三路掩码策略,而不是单一策略。
第一路是 random masking,随机选择一定比例的列做掩码。这是最基本的做法,好处是简单稳定,让模型均衡地重建所有列。
第二路是 correlated masking,先把列按相关性聚类,然后掩码其中一类里的所有列。举个例子,用户的"累计消费金额"和"最大单笔金额"相关性通常很高,如果只随机掩掉其中一个,模型可以从另一个轻松推断,任务就太简单了。把整组强相关列同时掩掉,模型只能借助更外围的弱相关列,这样学出来的表征会更鲁棒,不会过分依赖单点捷径。
第三路是 mixture masking 里的长列掩码,类似 NLP 里的 span masking。随机选某一列,掩掉它在一整个 batch 里的全部值,然后额外再带一两个其他列一起掩。这样可以模拟"整列数据在线上缺失"的场景。
我实际用下来比较顺手的配置是mask_mix=(0.7, 0.2, 0.1),即 70% 随机、20% 相关块、10% 长列掩码。这个配置不是 LimiX-2 官方唯一推荐的,但在我测试的三张表上都比纯随机掩码好 0.3 到 0.8 个 AUC 点。
3.2 重建目标怎么组合
LimiX-2 对数值列和类别列分别计算损失,再加权求和。数值列用的是归一化空间上的 MSE,类别列用的是交叉熵。这里有一个我一开始没注意到的点:数值列不要直接用原始尺度的 MSE。拿订单金额举例,几百块和几万块的量级差距会让模型把绝大部分容量花在拟合大额样本上,所以必须对数值列先做分布变换,把长尾拉平后再算 MSE。LimiX-2 默认走分位数到正态分布,本质上是让模型在标准化空间里做回归,评估时再映射回去。
损失计算可以用一个简化伪代码来理解:
def compute_loss(pred_num, target_num, pred_cat, target_cat): numeric_loss = F.mse_loss(pred_num, target_num) categorical_loss = F.cross_entropy( pred_cat.view(-1, pred_cat.size(-1)), target_cat.view(-1) ) return numeric_loss * 1.0 + categorical_loss * 1.0权重可以调,如果某张表类别列特别多,可以适当把类别权重降到 0.5,避免预训练被类别列主导。我在一张有 40 个类别列、5 个数值列的表上,把类别权重降到 0.6 之后,下游指标反而更稳。
3.3 掩码率的经验值
LimiX-2 里mask_ratio我一般设在 0.2 到 0.3 之间。太低会让任务很简单,模型直接走捷径,学不到深层交互;太高会让任务难到几乎重建不出来,训练非常不稳定,尤其数值列重建 loss 很难压下去。文本领域 15% 左右比较常见,表格领域 25% 左右我体感是性价比最高的区间。
我做过一组小规模对照,在 10 万行、30 个特征的数据上,掩码率从 0.15 提到 0.25,下游 AUC 提升了约 0.004;提到 0.35 时,预训练 loss 开始震荡,下游效果反而回落。所以如果你的数据列数比较少,比如只有 8 个特征,建议把掩码率降低到 0.15,不然一 mask 就是 20% 的列,信息太少。
4. 数据工程的四个关键处理:归一化、类别编码、缺失值、批采样
4.1 分位数归一化为什么比标准化稳
表格数据里数值列的分布千奇百怪,销售额、登录次数、用户年龄这些特征的分布形态完全不同。如果统一做 z-score 标准化,偏态严重的列会有很多极端值,模型在 MSE 重建时会被这些离群点牵着走。LimiX-2 推荐的做法是先做分位数变换到正态分布,再用标准化的向量表示。
这里有个实操细节:分位数变换的n_quantiles不要设得太大,1000 到 2000 就足够。设太大容易把训练集中的噪声也学进分位点,导致验证集和线上数据变换后分布不一致。我在一个百万行的数据集上看到,n_quantiles=5000时变换结果在线上数据集上出现了明显的分段跳变,改成 1000 后稳定很多。
4.2 高基数类别列的编码
类别列如果直接建 embedding,基数 100 万就会得到一个 100 万行的嵌入矩阵,显存直接爆炸。LimiX-2 的处理是双保险:先按频次截断,只保留出现次数超过某个阈值(比如 20 次)的类别,其余归入一个统一的<rare>token;如果截断后基数仍然很大,再走哈希分桶到例如 10000 个桶。这样能保证 embedding 矩阵大小可控。
这个方案会损失少量稀有类别信息,但对预训练来说关系不大。因为 LimiX-2 学的是列间联合分布,高频类别已经把大多数模式覆盖了,稀有类别本来样本量就少,硬学反而容易过拟合。实际业务里,如果某个高基数列本身业务意义很强,比如商品 ID,我会保留它为单独一路不做哈希,同时配合更小的 embedding 维度,比如 16 维。
4.3 缺失值和掩码千万不要混为一谈
这是 LimiX-2 数据工程里最容易被误解的一条。很多做表格预训练的人会把缺失值和[MASK]当成同一件事,认为"反正都是未知,让模型一起学"。但业务推理时,缺失是一个客观事实,掩码是人为制造的任务。LimiX-2 内部会生成两个 mask 矩阵:missing_mask标记原始空值,masked_mask标记被掩码位置。模型输入时两者同时拼接进 token,但计算 loss 时只评估masked_mask & ~missing_mask的区域。
我实际踩过这个坑:一开始我把缺失值直接填充为 0,再随机掩码重建。结果模型对真实缺失值的响应很奇怪,它把缺失当成了一个可以大胆猜测的信号,而不是一个需要谨慎对待的输入。改成双 mask 之后,模型在缺失率较高的列上的表现才真正稳定下来。如果你的数据本身缺失率很高,比如超过 40%,建议优先把缺失标志位做好,再去调别的超参数。
4.4 批采样要照顾类别平衡
表格预训练最常见的翻车点不是模型,而是采样。LimiX-2 默认按行随机采样,但如果标签存在严重类别不平衡,或者某些列的高频类别集中在特定批次,预训练出来的表征会偏向高频组。我的做法是做带权重的 batch sampler,或者干脆把下游 label 作为辅助监督信号只在微调阶段引入,预训练阶段完全自监督,避免标签泄漏。
5. 下游微调与效果验证:从预训练到落地
5.1 三种微调范式
LimiX-2 在拿到预训练 encoder 之后,提供了三种接地气的用法。
第一种是最常用的全量微调,在预训练模型上接一个简单的 MLP head,整个网络一起训练。适合下游数据量在几万行以上的场景。
第二种是 linear probing,冻结 encoder,只训练最后的分类或回归层。适合下游数据量很小、担心过拟合的场景。我测试下来,几千行样本时 linear probing 比全量微调稳定得多。
第三种是把预训练 encoder 当特征提取器,把每列 token 的输出拼成一条向量,喂给 LightGBM 或 XGBoost。这是我最喜欢的方式,因为表格数据上 GBM 依然是最强 baseline,用一个深度模型的特征去增强 GBM,往往比我单独优化深度模型效果好。
5.2 一组可以复现思路的对照实验
我在一个内部脱敏的贷款违约二分类数据集上跑过一组对比,数据规模 10 万行、35 个特征,其中类别列 12 个,数值列 23 个,缺失率约 18%。表格里的数字代表验证集 AUC。
| 方案 | AUC |
|---|---|
| LightGBM(默认参数,特征手工编码) | 0.838 |
| FT-Transformer(从头训练) | 0.842 |
| LimiX-2 随机初始化 + 微调 | 0.831 |
| LimiX-2 预训练 + 微调 | 0.851 |
| LimiX-2 预训练特征 + LightGBM | 0.855 |
最大的体会是两个点:第一,预训练相比随机初始化提升了整整两个点,说明这个数据集的列间依赖确实有大量可挖掘的结构;第二,纯深度模型还是没有超过"深度预训练特征 + LightGBM"的组合,表格下游任务里 GBM 依然是一个值得尊重的对手。
5.3 把预训练模型当成缺失值填补工具
LimiX-2 训练完成后,天然具备一个附加能力:指定任意一列被掩码,模型可以输出该列数值的预测分布或类别概率。这意味着你可以把它当成一个复杂的缺失值插补器来用。比如线上推理时用户没有填职业信息,你可以让模型基于其他列生成一个职业概率分布,而不是简单填众数。
不过要提醒一句:如果插补结果被拿去当特征再喂给同一个 LimiX-2 模型,会产生一定程度的循环依赖,可能导致局部特征被强化。我的习惯是用微调前的预训练模型来插补,插补完的特征再用于训练下游简单模型,这样信息泄漏的风险较低。
6. 运行 LimiX-2 时踩过的坑:从 loss 不收敛到显存爆炸
6.1 数值重建 loss 不收敛的排查链路
我第一次跑 LimiX-2 时,最头疼的是数值列 MSE loss 一直在震荡。当时我先怀疑学习率,从 1e-4 一路降到 1e-5,loss 还是不停在某个区间跳。后来逐项排查,发现是我在分位数变换后没有检查数值范围,某些列经过变换后带有极端尾部值,MSE 对离群点极其敏感。解决方法是把分位数变换后的值再进行一次截断,限制在 [-4, 4] 区间,或者改用 Huber loss 做数值重建。LimiX-2 官方设计里其实有 Huber 选项,但我一开始没注意。
排查顺序我建议这样走:先看输入是否有 NaN 和极端值,再看梯度是否出现异常大值,最后才考虑调学习率。表格预训练 task 不比 NLP,数据质量对收敛的影响远大于模型结构。
6.2 类别列 embedding 导致显存爆炸
有张表里有个"设备型号"列,基数接近 30 万。我一开始图省事直接建了完整 embedding,batch size 256 跑到一半显存就满了。LimiX-2 的哈希分桶参数我一开始设置成 50000,还是有点大,后来压到 8000 才流畅跑起来。这里我的心得是:分桶数不追求覆盖全部类别,8000 到 15000 已经能保留绝大多数频繁类别之间的区分度,再大只是浪费显存。
如果你要大批量训练,还有一个更省显存的做法:把高基数列单独走一个小的 MLP 编码器,而不是 embedding 查表,这样参数量从基数*维度变成固定参数量,代价是偶尔会有哈希碰撞导致的信息损失,但预训练阶段完全可接受。
6.3 分布式训练下掩码不一致
当我把 LimiX-2 从单卡升级到多卡 DDP 训练时,发现每个 rank 生成的掩码矩阵不一样,导致验证集上的 loss 波动很大。LimiX-2 默认按torch.manual_seed生成掩码,如果不额外同步,不同 GPU 上同一份数据会被掩掉不同列。我的处理方式是在每次 epoch 开始前统一用同一个随机种子生成掩码矩阵,然后 broadcast 到所有 rank;或者干脆在 DataLoader 层面实现一个全局掩码生成器,只由主进程生成。这个问题文档里写着,但很容易被忽略,遇到"多卡和单卡结果不一致"时优先查这里。
6.4 学习率调度与温度参数
最后一个是比较玄学但确实有效的细节。LimiX-2 预训练阶段用 warmup + linear decay 学习率,微调阶段用比预训练低一个量级的学习率。我在实验里发现,微调时如果直接沿用预训练的 1e-4,很容易在头几个 epoch 就把已经学到的表征破坏掉,特别是 linear probing 场景。另外,类别列输出的温度参数会显著影响重建难度,默认温度 1.0 偏保守;如果类别列特别多,可以适当降低温度到 0.8,让模型更有信心地输出,下游分类任务也会受益。
现在再让我跑一个表格项目,我的默认路径已经变得很固定:先花半天把数据清洗成 LimiX-2 能吃的格式,预训练一个 6 层的中等规模编码器,同时训练 LightGBM 和微调头,最后把两边的 logits 做一次简单融合。这套流程在我最近几个项目上没有输过纯 GBM 方案,在中小数据集上提升尤其明显。最后分享一个很实用的小技巧:做特征提取时不要只用最后一层的列 token,把每一层的平均池化结果拼起来,再作为下游特征,效果往往比单独取最后一层更好。