表格数据建模二十年:从MLP到TFM的架构演进与工程实践

发布时间:2026/9/18 8:02:07
表格数据建模二十年:从MLP到TFM的架构演进与工程实践 1. 为什么表格数据建模绕不开这二十年的反复折腾如果你做过几年机器学习相关的工作大概率对这样一个场景不陌生老板丢给你一份几百万行的表格里面有用户ID、注册时长、最近30天消费金额、点击次数、商品类目ID外加几个缺失值一堆的分类特征。你的第一反应不是上什么高大上的神经网络而是先跑一版XGBoost或者LightGBM做基线然后发现树模型效果居然还不错甚至比你后来精心调的MLP好上一截。这时候你心里可能会冒出一个疑问神经网络在图像和文本上都已经杀疯了为什么到了表格数据这里反而经常干不过一堆带正则化的决策树这个问题的答案恰恰就是“表格神经网络架构发展史从MLP到TFM模型”这个题目背后最核心的叙事线索。表格数据和我们熟悉的图像、文本数据有一个本质区别它没有天然的空间结构也没有天然的时序依赖。图像有像素之间的局部相关性文本有单词之间的顺序关系而表格数据里的每一列是什么含义、列和列之间是什么关系全靠人给的特征工程去定义。这就导致神经网络最擅长的“自动提取结构特征”这个能力在表格数据上很难直接发挥出来。MLP作为最朴素的神经网络理论上能逼近任何函数但在表格场景下它会遇到两个致命问题一是对特征之间的非线性交互拟合效率低二是对稀疏、高基数分类特征的处理非常笨拙。从2000年前后到2020年代初表格数据建模的主流一直是梯度提升决策树GBDT体系。GBDT能霸榜这么久靠的不是什么高深的数学理论而是它对表格数据“不均匀、有缺失、特征含义各异”这种天性的天然适配。但树模型也有天花板它学不到特征之间的乘法关系学不到某些全局性的结构模式而且在训练时需要反复扫描数据难以做到流式更新。于是一批研究者开始思考能不能设计出一种神经网络架构专门针对表格数据的特性做优化既保留神经网络的灵活性和可扩展性又能逼近甚至超越GBDT的效果这就是本文要聊的主线。从早期的EmbeddingMLP方案到后来引入注意力机制的Transformer式架构再到目前比较前沿的TFM模型我把这个缩写理解成针对表格数据的Transformer变体统称业内类似的工作还包括TabTransformer、FT-Transformer、SAINT等表格神经网络在这二十年里兜兜转转踩过很多坑也积累了不少真正有效的设计经验。本文不打算写成一篇干巴巴的论文综述而是想以一个实践者的视角把这条演进路线上的关键节点、背后的设计动机、以及我实际使用中的感受和教训一次讲清楚。不管你是刚接触表格深度学习的新手还是已经在用树模型做业务的工程师这篇文章都值得花二十分钟看完。2. 最初的二十年为什么MLP在表格上打不过GBDT2.1 MLP的底子并不差但它忽略了表格数据的“不均匀性”要说清楚表格神经网络的发展得先回到MLP本身。多层感知机由输入层、若干隐藏层和输出层组成每一层做线性变换加非线性激活按道理说只要隐藏层宽度和深度足够MLP可以逼近任意复杂的函数。这个结论在数学上没问题但在工程上有一道巨大的坎表格数据的特征分布极不均匀。图像数据进入神经网络之前像素值通常被归一化到[0,1]或者[-1,1]所有像素的语义尺度是统一的。文本数据经过Embedding之后每个词的向量维度也保持一致。但表格数据不是这样一列是年龄取值范围18到80另一列是消费金额取值范围从0到几百万还有一列是类别ID取值是几千个互不相关的整数。这种量纲和语义尺度差异巨大的多列输入如果直接喂给MLP优化过程会非常痛苦。损失函数对尺度的敏感度完全被大数值列主导小数值列的梯度信号会被淹没最终模型学出来就是个“偏科生”。这个问题可以通过特征归一化在一定程度上缓解但归一化只是第一步。真正麻烦的是稀疏高基数分类特征。比如说“用户所在城市”这一列有300多个取值如果做One-Hot编码输入维度被撑到几千维但每条样本里只有一维是1。这种极端的稀疏性对MLP来说几乎是灾难绝大部分Embedding参数在大多数样本上根本得不到更新训练效率极低还特别容易过拟合。这时候GBDT的优势就体现出来了。决策树在做节点分裂时本质上是在做“特征阈值划分”它天然不关心特征的绝对值大小只关心相对排序和划分点。而且树模型对缺失值有原生的处理逻辑不需要额外的填充步骤。这些特性让GBDT在处理表格数据的“不均匀性”时不需要像神经网络那样做大量前置的清洗和编码工作直接用原始特征就能拿到不错的效果。2.2 GBDT的隐式特征交互和神经网络显式建模的差距还有一个很关键的点就是特征交互的处理方式。表格数据里真正有价值的信息往往藏在特征组合里。比如“是否是新用户”和“最近7天消费次数”这两个特征单独看可能都一般但它们组合起来能非常强地预测“用户是否即将流失”。GBDT在处理这种交互时是隐式的每一棵树的决策路径本质上就是在做特征组合的判断。只要树深度够XGBoost和LightGBM可以自动捕捉到任意阶的特征交互虽然这种交互是沿着决策路径逐层组合的但实际效果已经足够好。MLP理论上也能建模高阶交互但它需要网络自己去学习哪些特征组合重要这在数据和算力有限的情况下效率远不如树模型的贪心分裂来得直接。更棘手的是MLP的神经元在整个特征空间上共享参数也就是说它对交互的建模是“全局性”的。而表格数据往往存在显著的局部模式某些特征组合只在特定的数据子集中有效。树的决策路径天然适合这种局部化建模MLP则显得有点“一视同仁”。我印象很深的一次实验某电商用户复购预测任务用LightGBM跑AUC能到0.78换成一个四层MLP做了精细的归一化和EmbeddingAUC只有0.75。这0.03的差距在业务上可能就是数百万营收的差别。那段时间很多团队都得出了类似的结论MLP在表格数据上不是不能用而是性价比太低。你需要花大量时间做特征工程、归一化、过拟合控制最后效果还可能不如默认参数的LightGBM。2.3 也是从这个时候开始“针对表格数据设计神经网络”成了研究方向正是这种“怎么调都打不过树模型”的挫败感催生了一批专门研究表格神经网络架构的工作。研究者们的思路大致分两类一类是用神经网络的组件去模拟树模型的行为比如Deep Forest、NODENeural Oblivious Decision Ensembles就是这种思路的代表另一类是给神经网络引入更适配表格数据的归纳偏置比如特征嵌入、注意力机制、特征维度上的残差连接等。前者的思路在当时看起来很吸引人但实际效果并没有取得压倒性优势而且训练成本比GBDT高得多后者则慢慢演变成了后来的Transformer系方案。现在回头看MLP在表格上表现不佳并不是神经网络这个方向本身错了而是当时的架构设计没有解决表格数据的核心痛点。这些痛点不解决再深的MLP也只是在错误的道路上越走越远。接下来要聊的Embedding化、残差连接、特征Token化等一系列改进本质上都是冲着这些痛点去的。3. 转折点Embedding化、残差连接与特征Token化3.1 把表格的每一列当成一个“词”来学如果说MLP时代最大的问题是“不知道怎么处理表格特征”那后续一系列工作的核心思路就是给每一列特征学习一个专属的向量表示让网络自己去理解这列特征的语义。这个思想借鉴了NLP里Embedding的概念——把离散的、无序的ID映射到一个连续的向量空间让语义相近的取值在向量空间里靠近。具体到表格数据做法是给每一列特征都建立一个Embedding表。类别特征直接查表得到向量连续特征先做分箱或者通过一个线性层映射成向量然后再查表。这样一来每个特征列都有自己的语义空间网络在后续层中可以在这个空间里做特征交互而不是在原始的、尺度混乱的数值空间里硬算。这里有一个我实际使用中觉得特别重要的细节连续特征的Embedding并没有一个绝对正确的做法。有些人喜欢先对连续特征做分箱再按类别特征处理有些人则是直接用线性变换把原始数值压到一个向量。分箱的好处是可以模拟树模型的分裂行为捕捉非线性关系但缺点是分箱边界的选择很敏感线性变换简单直接但表达力有限。我在多个数据集上对比过分箱普通Embedding在中小规模数据上一般比线性变换好但需要多调一个“箱数”的超参数。如果你的特征分布很偏比如长尾严重分箱前最好先做log变换或者rank归一化否则分箱后大部分样本会集中到同一个箱子里信息密度严重不均衡。3.2 残差连接在表格架构里为什么特别关键Embedding解决了特征表示的问题但它并没有解决“网络加深后训练不稳定”的问题。表格数据的特征维度虽然不算高通常几百维但这几百维特征之间的相关性差异很大有些特征高度相关有些几乎独立。直接堆叠多层全连接梯度在反传过程中很容易被这些复杂的相关性干扰导致训练震荡或收敛到次优解。残差连接在这里扮演的角色和它在ResNet里类似让梯度有一条高速公路可以直达输入层。但它在表格场景中还有一个额外的意义——等价于让模型保留原始的浅层特征表达。树模型天然具备“使用原始特征的能力”它可以在任意一层分裂的时候直接使用原始特征值而MLP如果网络很深前面的层会把原始特征反复变换到后面原始信息可能已经面目全非。残差连接可以让网络决定“我要不要在这一层做变换”如果这一层学到的东西不重要网络可以直接跳过它把上一层的输出原封不动地传给下一层。我在实测中观察到加入残差连接的MLP在表格任务上的收敛速度比普通MLP快很多尤其是embedding维度较大的情况下。这个现象背后的原因也不难理解embedding层初始化的参数接近随机如果全靠后面的深层网络去“纠错”优化压力很大有了残差连接前几层哪怕学得一般后面的层也可以基于相对稳定的输入特征表示继续优化。3.3 特征Token化从“一个特征一个数字”到“一个特征一个向量序列”MLP加Embedding加残差这套组合让表格神经网络的效果大幅提升但它仍然有一个结构性的短板特征之间的交互是通过隐式的全连接实现的网络不会明确地“知道”哪些特征值得交互哪些不值得。注意力机制的出现给了这个问题的另一种解法。特征Token化的思路非常直接把每个特征经过Embedding后看作一个Token多个Token组成一个序列然后用Transformer的注意力层让Token之间互相“看”。这样每一列特征都可以根据自己的语义去关注其他列中和自己相关的信息。举例来说在一个贷款风控模型里“年收入”这个Token可能会重点关注“工作年限”而“负债率”这个Token可能更关注“已有贷款笔数”。这种显式的交互建模方式比MLP的全连接更加高效——注意力权重告诉模型哪些交互重要模型把算力集中在这些重要的交互上。最早把Transformer直接搬到表格上的尝试之一是TabTransformer它主要针对类别特征做多层Transformer编码连续特征仍然走MLP分支。后来FT-Transformer进一步把连续特征也做了Token化让所有特征统一进入Transformer层效果进一步提升。再后来的SAINT等模型则引入了类似BERT的预训练思路在这个框架上做掩码重建让特征表示学到更多的数据内在结构。这套思路后来被统称为“表格Transformer”或直接叫TFM模型。不过这里必须提醒一句Transformer不是银弹。FT-Transformer在多个标准表格基准比如TabZilla这种大规模基准测试上的表现确实赶上甚至超过了调好参的GBDT但它的训练成本和推理成本比树模型高出一个数量级。而且Transformer对特征顺序和数量非常敏感特征一多注意力矩阵的复杂度和内存占用会快速上升。后面我会专门用一个章节来说这个“性能与代价的权衡”问题。4. TFM模型注意力机制给表格带来的真正的结构升级4.1 TFM模型的核心架构演进从MLP到特征Token化再到注意力循环我们把时间线拉回到“TFM模型”这个概念上。严格来说并不存在一个官方定义叫“TFM模型”的单一架构它更像是一个统称指向所有把Transformer结构用在表格数据上的模型家族。这个家族的内核是“表格数据特征Token化自注意力交互前馈网络输出”三段式结构。下面用一个简化的结构图来拆解它原始表格行数据一列一个字段 │ ▼ Feature Tokenizer特征Token化 │ ▼ [Token_1] [Token_2] ... [Token_n] ← 每列特征被映射成一个固定维度的向量 │ ▼ Transformer/注意力堆叠层 ← 在这里Token之间互相计算注意力权重 │ ▼ 池化 / [CLS] Token / 拼接 │ ▼ 预测头分类/回归输出这个结构其实和NLP里的BERT序列不太一样。BERT的Token有位置编码因为词序有含义表格数据里特征的顺序本身没有语义所以TFM模型一般不强制加入位置编码或者只加一个“可学习的列编号”向量来帮助模型区分不同列。这是一个很实际的工程细节——我见过有人直接把NLP里Transformer的位置编码搬到表格上结果效果不升反降就是因为位置编码让模型误以为特征之间有时序关系反而引入了噪音。特征Token化的具体实现可以稍微展开一点。假设你的表格有n个特征其中m个是类别特征p个是连续特征。类别特征通过查Embedding表得到一个维度为d的向量连续特征可以先用一个线性层或者分箱后的Embedding或者用数值编码器映射到同样的d维空间。最终一行样本被表示成一个形状为(n, d)的Tensor。这个Tensor可以理解成“一条由n个Token组成的序列”每个Token的宽度是d。之后标准的Transformer Encoder层在这个序列上做多头自注意力让特征之间充分交互。4.2 TFM模型在非线性拟合和特征交互上到底比MLP“多会”了什么要说清楚TFM比MLP强在哪里得回到注意力机制的本质。MLP对特征交互的建模是“隐式且全局共享”的每一层全连接的参数对所有样本都一样模型很难针对不同的样本选择不同的交互路径。而自注意力机制是“动态”的对不同的输入样本Token之间的注意力权重不同。这意味着模型可以根据每一条样本自身的特征取值动态决定哪些特征交互参与决策。举一个很直观的例子在“预测用户是否点击广告”这个任务中样本A是一位老用户样本B是一位新用户。老用户的点击行为可能和历史点击率、设备型号的交互更相关新用户的点击行为可能和落地页类型、广告创意的交互更相关。MLP对所有样本都使用同一套交互权重它只能用一套“平均最优”的交互策略而TFM可以对样本A和样本B分别算出不同的注意力矩阵真正做到“看人下菜碟”。这种动态特征交互能力是表格深度学习从MLP时代迈向注意力时代最核心的飞跃。实际数据上的表现也确实能说明问题。在多个公开数据集比如Forest Cover Type、Adult Census、Click Prediction上FT-Transformer的代表性实现相比传统MLP普遍有2到5个百分点的AUC提升。在特征交互密集的数据集上甚至有一定概率超过调参良好的XGBoost/LightGBM。当然这个结论也不能过度外推要知道在很多低维、特征独立的表格数据上TFM的优势并不明显树模型仍然是更划算的选择。4.3 TFM模型训练的几个独有细节特征顺序、Attention Head、数值稳定性关于TFM在实操中的注意事项我想分享三个踩过坑之后总结出来的要点。第一个是特征顺序不能随意摆放。虽然没有位置编码但Transformer对Token顺序依然不是完全无感的因为注意力矩阵的计算受Token顺序影响自注意力的QKV计算中特征的排列顺序会影响初始的注意力分布。在实验中我发现把相关性高的特征放在相邻位置训练收敛更快最终效果也略好。我常用的做法是先按照特征对目标变量的单变量相关性排序再按相关性从高到低摆放特征Token这个简单的小技巧对收敛速度有明显帮助。第二个是Attention Head的数量不宜盲目增加。表格数据的Token数量通常不会太多一般就是几十个到一两百个特征每个特征语义比较集中不像NLP里一个Token蕴含大量歧义需要多个Head去捕捉不同角度的语义。我尝试过头数从4增加到16效果并没有显著提升反而参数量涨了四倍训练时间明显变长。在表格场景里4到8个头一般就足够了再往上边际收益极低。第三个是数值稳定性。表格数据经过Embedding之后不同Token向量的尺度可能会有较大差异尤其是连续特征Embedding如果没做好归一化很容易导致注意力权重偏向某个大尺度的Token使其他Token的信息被压制。我的做法是在连续特征进入Embedding之前先做rank Gauss归一化也就是把连续值映射到标准正态分布的分位数上这样能有效抑制离群值的影响。另外注意力层输出之后接LayerNorm的时机也值得注意Pre-LN比Post-LN在表格任务上通常更稳定收敛更快这其实和NLP里大模型训练的经验一致。5. 性能与代价的权衡TFM模型比XGBoost慢凭什么选它5.1 从决策树的“分裂查找”到Transformer的“矩阵运算”计算模式完全不同聊到TFM的缺点最绕不开的就是计算效率。XGBoost和LightGBM的底层是直方图加速的分裂查找对高维稀疏特征支持好单机处理百万级数据毫无压力。TFM的底层是密集矩阵乘法和多头注意力计算尽管近年有一些核融合和优化但整体算力需求还是要高出一个量级。我在一个500万行、60个特征的数据集上做过对比LightGBM训练到全量迭代大约需要十几分钟FT-Transformer在小batch、GPU环境下训练到收敛需要数小时。这个差距在业务项目中不是小数目。但这并不意味着TFM完全没有用武之地。表格Transformer真正的优势场景有两个一是特征含义复杂、交互关系强的数据集比如推荐系统里的用户行为序列列、营销场景里的多触达特征、风险控制里的多种渠道交叉特征等二是需要特征表示复用的场景比如你打算做多任务学习、增量学习或者希望预训练一套特征表示然后在下游多个任务上微调——这种情况下TFM的Embedding表示比树模型的叶子节点分布更灵活泛化能力也更自然。5.2 我在真实场景里对TFM的选型判断标准基于上面的经验我给自己定了一条简化的选型判断标准在多个实际项目中验证过分享出来仅供参考判断维度更倾向选GBDT/XGBoost更倾向选TFM表格Transformer数据量级百万级以内、特征维度中等超大训练集千万级以上、算力充足特征类型大量稀疏类别特征、有大量缺失密集数值特征、有复杂交互模式推理延迟在线实时推理、毫秒级响应离线批量预测或延迟容忍度高任务复杂度单任务、特征工程到位多任务、需要共享特征表示团队资源希望快速上线、少调参有GPU环境和足够的时间做调参这个表不是绝对的但它能帮你在项目启动时快速判断方向。如果一开始就选错了路线后面花再多时间调参也补不回来。需要特别说明的是表格Transformer不是某些营销号说的“横空出世能替代一切的存在”。它是在表格深度学习这条路上把注意力机制引入数据建模之后自然进化出来的产物。MLP时代那些归一化、Embedding、残差连接的积累到TFM时代依然有效只是在架构层面换了一种更聪明的组合方式。理解了这条演进脉络你就知道为什么“从MLP到TFM”不只是一次模型升级更是一次关于“如何让神经网络尊重表格数据天性”的认知升级。5.3 训练TFM模型时绝对值得留意的工程细节如果你决定在项目里试一试TFM下面这几条工程层面的建议是我踩过坑之后的浓缩经验。第一重视Embedding层的参数初始化。TFM的输入质量几乎完全取决于特征Token化的质量如果Embedding初始化得不好后续的注意力层会在垃圾输入上做计算训多长时间都很难拉回来。推荐使用正态分布初始化标准差设成embedding维度平方根的倒数这个做法和NLP里常用的xavier初始化类似实验下来比较稳。第二用学习率调度器而不是固定学习率。TFM的优化地形比较复杂固定学习率很容易陷入震荡。我一般用WarmupCosineAnnealing的组合前10%的步数线性warmup到峰值学习率比如3e-4然后余弦退火到一个很小的学习率。这个设置几乎每次都能比固定学习率带来1到2个点的效果提升。第三如果不确定超参数怎么选先做小规模超参搜索。我的经验是embedding维度16到64、Transformer层数2到6、注意力头数4到8三个超参数对效果的影响最大。在小规模数据子集上先跑一轮随机搜索再在完整数据上复现最优参数这样能大幅节省GPU时间。6. 从MLP到TFM的演进中那些你没注意到的“隐藏主线”6.1 数据处理范式的变化从“手动特征工程”到“自动特征表示”如果你只是盯着网络结构的变化可能会漏掉一条更关键的暗线——数据处理方式的变化。MLP时代表格深度学习的前置工作极其繁重要对连续特征做归一化、分箱、WOE编码要对类别特征做One-Hot或目标编码要把所有特征拼接成一个固定维度的输入向量。这套流程非常依赖人工经验而且往往需要针对不同数据集做定制化调整。到了TFM时代数据处理的范式发生了本质变化每个特征列独立做Embedding连续特征和类别特征统一映射到同一个向量空间然后由网络自动学习特征之间的交互。这意味着你不再需要精心设计“特征之间的组合方式”网络自己会通过注意力机制找到有效的组合。这不代表特征工程没用了而是它的重心从“怎么把特征喂给模型”变成了“怎么定义特征的初始表示”——这一步看起来简单实际上对语义的理解要求更高。6.2 两种典型失败模式特征太少时的过拟合和特征太多时的注意力稀释TFM架构还有一个容易被忽视的边界情况——特征数量极端时的表现。我分别踩过两个方向的坑这里都说不仔细免得大家重蹈覆辙。特征太少时比如只有五六个特征TFM很容易过拟合。因为Transformer的核心优势在于建模高维特征的交互如果特征本身就很少模型参数相对输入信息量过多注意力机制反而成了“记忆训练样本”的工具。我试过一个只有8个特征的小数据集FT-Transformer训完训练AUC接近1.0测试AUC只有0.7过拟合极其严重。这种情况下MLP甚至GBDT都更容易控制方差。解决办法是给TFM加更强的正则化比如更大的Dropout0.3以上、更小的embedding维度、Early Stopping或者干脆不要用TFM。特征太多时比如几百个特征又会出现“注意力稀释”的问题。每个Token在计算注意力时都要和其他所有Token交互如果特征里有很多无关紧要的噪音列注意力权重会被分散真正重要的特征拿不到足够的注意力权重。我的经验是在特征进入TFM之前先用LightGBM做一个简单的特征重要性筛选去掉重要性极低的列或者用PCA把冗余的连续特征先降维。这个前处理步骤虽然“不神经”但在实际工程中极其有效。6.3 混合架构是新的趋势TFM负责交互决策树负责“兜底”最后一个趋势我想重点提一下因为它很可能是接下来表格深度学习的主流方向TFM模型和GBDT混合建模。思路很简单TFM负责提取高阶特征交互表示GBDT在TFM的输出通常是最后几层隐藏层向量基础上做最终预测反过来GBDT的叶子节点分布也可以作为TFM的额外特征输入让模型同时利用树模型的强单特征处理能力和神经网络的特征交互能力。这个方向的代表性思路可以参考TabNet、NODE以及一些后续的混合工作但真正让我觉得有价值的不是某个具体模型而是这种“优势互补”的建模哲学。我在一个客户流失预测任务中尝试过“LightGBM FT-Transformer集成”的方案FT-Transformer的输出作为GBDT的额外特征最终AUC比单独用任一模型高出接近1.5个百分点。训练成本虽然高了一些但业务收益非常明显。这也引出一个建议不要把MLP、GBDT、TFM看成竞争关系它们更像是不同工具各有各的适配场景。真正成熟的数据科学团队一定是在一个统一框架里动态组合这些工具而不是抱死一个模型走到底。理解了这个逻辑再回头看“从MLP到TFM的二十年”你会觉得这条演进路线其实特别清晰——每一个新架构的出现都是在解决前一个架构在特定问题上的短板而不是简单地推翻重来。7. 如果你想快速验证最基础的一版TFM模型实战拆解7.1 最小可运行的数据流从原始表格到预测结果讲了这么多理论是时候给出一段可以直接上手跑的代码思路了。为了不让代码过于琐碎我这里只给出核心的数据流骨架完整代码在GitHub上很多开源库都有比如pytorch-tabular这个库可以直接用。核心思想是让大家看到从原始表格到预测结果的完整链路长什么样。import torch import torch.nn as nn class FeatureTokenizer(nn.Module): 把每列特征映射成固定维度的Token向量 def __init__(self, feature_meta, d_model64): super().__init__() self.d_model d_model # 对每个特征根据类型建立独立的编码器 self.encoders nn.ModuleList() for meta in feature_meta: if meta[type] categorical: # 类别特征Embedding查表num_embeddings类目数embedding_dimd_model self.encoders.append(nn.Embedding(meta[num_classes], d_model)) else: # 连续特征线性层将标量映射到d_model维 self.encoders.append(nn.Linear(1, d_model)) def forward(self, x_cat, x_num): tokens [] # 类别和连续特征分别编码后追加到同一个Token序列里 for i, enc in enumerate(self.encoders): if isinstance(enc, nn.Embedding): tokens.append(enc(x_cat[:, i])) else: tokens.append(enc(x_num[:, i].unsqueeze(-1))) # tokens: list of [B, d_model]拼接得到 [B, num_features, d_model] return torch.stack(tokens, dim1) class TFMBlock(nn.Module): 单层Transformer Encoder多头注意力 前馈网络带残差和LayerNorm def __init__(self, d_model, nhead, dim_feedforward, dropout0.1): super().__init__() self.attn nn.MultiheadAttention(d_model, nhead, dropoutdropout, batch_firstTrue) self.ln1 nn.LayerNorm(d_model) self.ffn nn.Sequential( nn.Linear(d_model, dim_feedforward), nn.GELU(), nn.Dropout(dropout), nn.Linear(dim_feedforward, d_model), nn.Dropout(dropout) ) self.ln2 nn.LayerNorm(d_model) def forward(self, x): # x: [B, num_features, d_model] attn_out, _ self.attn(x, x, x) x self.ln1(x attn_out) ffn_out self.ffn(x) x self.ln2(x ffn_out) return x class SimpleTFM(nn.Module): 简化版TFM特征Token化 若干Transformer层 预测头 def __init__(self, feature_meta, d_model64, nhead8, num_layers4, num_classes1, dropout0.1): super().__init__() self.tokenizer FeatureTokenizer(feature_meta, d_model) self.blocks nn.ModuleList([ TFMBlock(d_model, nhead, d_model * 4, dropout) for _ in range(num_layers) ]) self.head nn.Linear(d_model, num_classes) def forward(self, x_cat, x_num): tokens self.tokenizer(x_cat, x_num) # [B, F, d_model] for block in self.blocks: tokens block(tokens) # 对所有Token取平均池化等价于一个简易的聚合表示 pooled tokens.mean(dim1) # [B, d_model] return self.head(pooled).squeeze(1)这段代码看起来不长但它把一个完整的TFM链路串起来了。FeatureTokenizer负责把每列特征变成Token序列TFMBlock实现标准的Transformer Encoder结构最后通过平均池化把多个Token的信息聚合成一个向量做预测。你拿到任何一份表格数据先做一下数据类型标注哪些列是类别、哪些是连续然后就能用这套骨架跑起来。7.2 训练时的关键超参数建议和几个坑超参数方面我从自己的几百次实验中总结了一个比较稳健的起点配置超参数建议值说明d_modelEmbedding维度32或64特征多时64特征少时32nhead注意力头数4或8头数太多容易过拟合且效率低num_layersTransformer层数2到6层数过深在小数据集上容易过拟合dropout0.1到0.3特征多时用更大的dropoutbatch_size256到1024根据显存和数据集规模调整学习率3e-4到1e-3配WarmupCosine退火效果更好优化器AdamW比Adam更稳配合weight_decay1e-5有一个坑特别值得提醒连续特征的缩放方式。很多人习惯用StandardScaler做z-score归一化但这个做法对TFM并不一定最好。我在多个数据集上对比后发现连续特征经过rank Gauss归一化排序后映射到标准正态分位数后TFM的最终效果普遍优于z-score。原因是rank Gauss归一化对离群值更鲁棒而且把连续特征的分位数关系也带了进来等于变相加入了一些非线性变换的信息。另外训练TFM模型时务必监控验证集损失而不是验证集AUC来Early Stop。AUC在训练早期上升很快但后期变化平缓容易让Early Stopping判断失误损失函数的变化更平滑对过拟合更敏感。这个小细节能让你的模型少训不少epoch效果还更好。7.3 从验证集到线上线下一致性表格Transformer部署的三个提醒最后聊几句部署层面的经验因为模型效果再好部署不上线也是白搭。第一个提醒是Embedding表的序列化。TFM推理时第一步就要查Embedding表这个映射关系的状态完全依赖训练时的特征编码顺序。线上服务必须保存一份训练时生成的“特征名→Embedding索引”映射文件否则上线后特征顺序一对不上结果全乱。很多初学者在本地跑通Demo后直接把模型文件传到线上结果特征列顺序不同导致预测结果和离线对不上排查半天才发现是Embedding索引错位。第二个提醒是推理延迟的优化空间。TFM的推理瓶颈主要在注意力层如果在线推理延迟要求很严格可以考虑做模型蒸馏用训练好的大TFM模型去指导一个小型MLP学习输出分布。小模型的效果虽然达不到大模型的水平但通常能保持90%以上的性能同时延迟降低一个数量级。第三个提醒是在线数据分布漂移的监控。GBDT对特征分布漂移相对不敏感因为树分裂只依赖特征的相对顺序TFM则对输入分布的变化敏感得多尤其是Embedding层的输入空间一旦被改变预测结果可能大变。因此表格Transformer上线后要对模型输入特征的分布做实时监控一旦发现某个特征的分位数分布明显偏移就要重新训练或微调。这一点在业务数据频繁变化的场景中尤其重要。8. 写在最后我实际用下来的个人判断断断续续用了快两年表格Transformer我的感受是它确实不是万能的但它把表格建模这个领域往前推了一大步。如果让我给团队推荐路线我不会一上来就上TFM而是会先跑一版LightGBM建立基线当业务对特征交互的要求明显提高、或者有多任务/增量学习的诉求时再引入TFM。你要是刚开始接触这个方向建议先从FT-Transformer开源实现跑通再自己动手写一套最小骨架。理解从Embedding到Token化再到自注意力交互这一层层递进的关系比记住任何具体模型的默认参数都重要。最后再分享一个小技巧无论用MLP还是TFM“把连续特征改成rank Gauss归一化的Embedding输入”这一个改动几乎在所有表格深度模型上都能带来稳定的效果提升。如果你现在的表格模型卡在瓶颈上别的先不用动先试试这一条大概率会有意外收获。