别再一上来就上神经网络了:我在 3 万行表格数据上,让 MLP 输给了 LightGBM

🔑 关键词:LightGBM, 表格数据, 树模型 vs 神经网络, XGBoost, 小数据集建模

📖 摘要:用一次真实的用户流失预测项目,对比 LightGBM 和 MLP 在小规模表格数据上的表现差异,讲清楚为什么树模型在 10 万行以下的数据上往往更难被打败,以及什么时候才该掏出神经网络,附具体参数和操作步骤。

我花了两个星期,就为了让神经网络输给 LightGBM

图片

2023 年春天接的一个活,用户流失预测。样本 38241 行——这个数我到现在还记得,因为它在我脚本的注释里躺了很久,我一直没舍得删。正样本率 4.7%,62 个特征,其中 19 个是类别型,最大基数的那个是“注册渠道”,43 个取值。

我当时的反应很典型:都什么年代了还上 GBDT,太土了吧。于是我花了大概 10 天时间搭了一个 MLP。结构是 256-128-64,中间加了 BatchNorm 和 Dropout(0.3),Adam,lr 1e-3,batch size 256,早停 patience 设 20,跑 5 折交叉验证。类别特征做了 one-hot 加 embedding,数值特征做了 StandardScaler。

结果 5 折 AUC 卡在 0.786,我调到第 12 版的时候到了 0.794,然后就再也上不去了。那天晚上我有点烦躁,就把 LightGBM 装上了,参数基本是随手写的默认值改了两个,5 折跑完 AUC 0.812,总耗时 47 秒。

对,47 秒。我前面那 10 天,一大半时间是在等 GPU 和调那些我其实没完全搞明白的正则项。

不是玄学,这事有论文专门解释过

图片

后来我去翻了一下,发现踩这个坑的人多得是。Grinsztajn、Oyallon 和 Varoquaux 在 NeurIPS 2022 的 Datasets and Benchmarks track 上发过一篇,叫《Why do tree-based models still outperform deep learning on typical tabular data?》,他们在四十多个中规模数据集上做了系统对比,结论是树模型在中等规模数据上平均胜出,而且差距不是噪声级别的。

早一点还有 Shwartz-Ziv 和 Armon 发在 Information Fusion 上的那篇《Tabular Data: Deep Learning is Not All You Need》,结论差不多。Kaggle 上大部分表格类比赛,冠军方案里 XGBoost、LightGBM、CatBoost 基本都是主力,神经网络更多是拿来和 GBDT 做 ensemble 加一点点分。

为什么?我自己的理解是三件事:

第一,表格数据的决策边界往往是对齐坐标轴的。年龄超过 35 岁是一个条件,不是“年龄和收入的某个线性组合超过某个值”。树的分裂天生就是轴对齐的,而神经网络要先学会把数据旋转过去,这得多花很多参数和数据。

图片

第二,表格里总有一堆没用甚至有害的特征。树模型对不相关特征相当鲁棒,它顶多不分裂那根;MLP 不一样,权重会被那些噪声特征牵走,你得靠正则化和特征筛选硬扛。

第三,样本量。3 万多行对 MLP 来说太少了,它有几万甚至几十万个参数,过拟合几乎是必然的。树模型靠分裂点的贪心选择,容量天然受数据规模限制。

那什么时候该老老实实上神经网络

我不是说神经网络不行,我自己做图像和文本的时候从来没想过用树。判断标准我大致是这么几条:

  • 数据维度高且稀疏(比如文本 TF-IDF 出来几十万维),或者模态是图像、音频、视频、序列,那没得选。
  • 样本量上到百万级,且特征之间有明显的高阶交互,神经网络开始能压过 GBDT。
  • 需要迁移学习或者预训练,比如你有 BERT 或者 CLIP 可以直接微调,这个树模型完全做不了。
  • 需要端到端可微的 pipeline,比如推荐系统里 embedding 要和其他塔一起训。
  • 数据结构本身有局部相关性,比如时序里的卷积、图结构里的 GNN。

图片

反过来说,如果你的场景是“一份几十列、几万到几十万行的宽表,要预测一个标签”,那我劝你先把 LightGBM 跑起来,别急着写 nn.Module。

如果你手上现在就有这么一份表,我会这么走

第一步,5 分钟建 baseline。LightGBM 装好用 pip install lightgbm 就完事,参数我一般起手是:

num_leaves=31
learning_rate=0.05
n_estimators=2000
min_data_in_leaf=20
feature_fraction=0.8
bagging_fraction=0.8
bagging_freq=1
early_stopping_rounds=100
objective='binary'

图片

类别特征不要 one-hot,直接把它们列进 categorical_feature 参数,LightGBM 内部会用类似 Fischer 最优分裂的方式处理,43 个取值的渠道变量它扛得住。

第二步,看特征重要度和 SHAP,砍掉明显没用的。这一步经常能再涨 0.003~0.008 的 AUC。

第三步,如果 AUC 还差得远,先想清楚是特征不够还是模型不够。我见过太多人模型换了一圈,最后发现是标签定义有问题,或者有个强特征泄漏没进特征池。

第四步,实在需要,再上 MLP 或者 FT-Transformer,然后和 LightGBM 做加权平均。到这一步通常也就再涨 0.005 左右,值不值看你自己的时间成本。

有个数字我印象挺深:在我那个项目里,从 LightGBM 0.812 折腾到最后的 ensemble 0.818,多花了我大概 6 天。6 天换 0.006 的 AUC,业务上大概就是挽回几个用户的事。领导当时还挺满意,但我心里清楚,那 6 天的边际收益低得可怕。

图片

一句可能有点讨打的话

现在很多教程一讲机器学习就上 MNIST 或者 CIFAR-10,把人训练出一种条件反射,觉得“模型越深越现代”。但工业界真正每天在跑的模型,很大一部分还是几棵 boosting 树。它们不酷,论文不好发,面试的时候讲出来也不够唬人。

可它是真的能干活。

所以我的默认流程已经变成:先跑 LightGBM 拿到 baseline,再决定要不要动神经网络。如果神经网络打不过这个 baseline,那说明数据本身还没到需要它的程度,不是你代码写得不好。承认这一点,能省下很多个凌晨两点。

🏷️ 标签: