资讯详情

资讯详情

分类模型不玄学:从GBDT到神经网络的全面解析与实战

1. 内容整体设计与思路拆解1.1 为什么《模型不玄学》要从分类模型讲起如果你翻过几本机器学习的书大概率会注意到一个现象几乎所有教材在讲完线性回归之后紧接着就是逻辑回归和分类问题。这背后其实藏着一个很实际的原因——现实世界里的业务问题绝大多数都能被归约成“判断类别”。垃圾邮件识别是判断“是/否垃圾”医疗影像筛查是判断“有没有病灶”风控审核是判断“会不会逾期”电商推荐场景里的“预测点击率”也可以转化为二分类问题。分类模型是整个监督学习体系里应用面最广、落地价值最高的一类方法这一章把它放在开头去讲不是偶然。从学习路径来看分类模型的演化正好串起了一条清晰的技术脉络。早期工业界最常用的是逻辑回归简单、可解释、上线快但拟合能力有限随后树模型家族崛起尤其是GBDT这类梯度提升方法在结构化数据上打出了统治级表现再往后是深度学习时代神经网络把特征工程这件事彻底重构了图片、文本、语音这些非结构化数据终于有了统一的处理范式。把这三块放在同一章里对比着看你会发现它们并不是互相替代的关系而是不同数据形态、不同业务约束下的差异化选择。还有一个更关键的理由分类模型是你理解“模型评估”的最佳切入点。回归任务可以看MSE、MAE但分类任务引入的混淆矩阵、精确率、召回率、AUC这些概念才是机器学习里真正需要“决策思维”的地方。模型不是把准确率刷到99%就完事了在样本不均衡、误判代价不对称的真实场景里阈值怎么定、指标怎么取舍往往比模型本身更决定项目成败。1.2 这一章的阅读路径从树模型到神经网络的演进逻辑GBDT和神经网络放在同一章表面上看是因为它们都是当前工业界最常用的分类工具但更深层的原因是它们代表着两种完全不同的建模哲学。GBDT走的是“串行纠错”路线。每棵树都在拟合前面所有树的残差本质上是把一个复杂的非线性函数拆解成很多个简单的分段函数逐步逼近。它的优势在于对特征的理解极其充分能够自动处理特征之间的交互关系而且输出结果天然带有可解释性你可以很方便地算出每个特征的重要性。神经网络走的是“分层抽象”路线。原始输入经过一层一层的非线性变换低层学到的是局部模式高层学到的是语义概念最后通过一个softmax层输出类别概率。它在图像、语音、文本这些高度非结构化的数据上拥有压倒性优势因为不需要人手工设计特征端到端训练就能完成任务。我在实际项目里常用的做法是**结构化表格数据优先上GBDT非结构化数据直接上神经网络两者都有余力的时候用stacking把它们融合起来。**这也是工业界被验证过最稳的组合方式比如一些推荐系统比赛里的冠军方案基本上都是树模型和深度模型的融合。2. 核心细节解析与实操要点2.1 分类模型的底层逻辑不只是“分对”而已很多人学分类模型时容易陷入一个误区觉得分类就是“输出一个类别标签”。但实际工程里绝大多数分类模型输出的其实是一个概率分数然后再根据阈值决定判成哪个类别。以逻辑回归为例它的数学形式是p 1 / (1 exp(-(w^T x b)))这个p值代表“样本属于正类的概率”然后我们通常以0.5作为默认阈值p大于等于0.5判为正类否则判为负类。这里就引出一个很关键的工程问题**0.5这个阈值一定是最优的吗**在二分类样本分布均衡的时候0.5通常是个不错的起点但在现实业务里正负样本往往是极度不均衡的比如电商平台的点击率通常只有百分之几金融反欺诈场景里的欺诈样本可能只有千分之一甚至更低。我举一个实际例子来说明这个问题。假设我们做一个信贷风控模型正类是“违约用户”负类是“正常用户”模型在阈值取0.5时精确率和召回率可能都很低因为模型预测的概率普遍都很低。但如果把阈值降到0.1更多样本会被判为正类召回率升高了但误杀把正常用户判成违约也变多了。这时你需要根据业务成本去权衡到底哪个阈值更合适——每放过一个违约用户损失的是本金每误杀一个正常用户损失的是客户体验和潜在收益。这两者的代价完全不同阈值的选择就不该由模型本身决定而是由业务方和风控策略共同决定。实操中我建议你养成一个习惯**训练完模型之后不要急着看准确率先画出PR曲线或者ROC曲线然后沿着曲线找业务上最合适的阈值点。**这个过程往往比你花力气调模型的超参数更有效。再补充一个新手容易忽略的点分类模型的Loss到底是什么。很多人知道逻辑回归用交叉熵损失但不理解为什么不用MSE。核心原因在于MSE配合sigmoid激活函数时梯度更新速度在错误分类的区域会变得非常慢因为sigmoid的导数在两端趋近于0导致梯度消失。而交叉熵损失在错误分类时梯度更大模型能更快地修正错误。这个背后的数学直觉值得你记住。2.2 GBDT核心原理拆解梯度提升到底在提升什么GBDT的全称是Gradient Boosting Decision Tree直译过来就是“梯度提升决策树”。这个名字里三个词每一个都在描述一个关键设计选择。先看“Decision Tree”说明它的基学习器是决策树而且是CART回归树。注意即使是做分类任务GBDT底层的每棵树用的也是回归树因为每一轮拟合的目标是一个连续的梯度值而不是类别标签。这是初学者最容易搞混的地方。再看“Boosting”这是一种串行集成策略。每一棵新树都在纠正前面所有树的整体错误整个模型是累加结构F(x) F_{m-1}(x) η * h_m(x)其中h_m(x)是第m轮训练出来的树η是学习率。这个累加结构让模型可以一点一点逼近真实函数每一棵树只需要学会一小部分规律就行。最后是“Gradient”这是GBDT最精妙的点。传统的Boosting算法在每一轮算的是“残差”真实值减预测值但在GBDT里Friedman给出了一个更通用的框架——每一轮拟合的是损失函数的负梯度。**为什么是负梯度因为负梯度是损失函数在当前点下降最快的方向。**当损失函数取平方损失时负梯度正好等于残差但当我们换成其他损失函数比如对数损失、Huber损失负梯度的形式就变了但“沿着损失下降最快的方向走”这个核心思想不变。这种泛化使得GBDT能够适配各种不同的任务这也是它相较于AdaBoost的最大进步。2.3 GBDT与XGBoost、LightGBM的差异别再把它们混为一谈在网络上搜“gbdt”和“xgboost算法和gbdt算法的区别”之类的关键词能搜到大量回答但很多都讲得不够清楚。我用自己的实践来给你梳理一下。GBDT是一个算法框架而XGBoost、LightGBM、CatBoost是这个框架的高效工程实现。这三者在核心思想上是一致的都是在做梯度提升区别在于工程优化和算法细节的不同。XGBoost在GBDT的基础上做了三件重要的事。第一它对目标函数做了二阶泰勒展开同时使用一阶导数和二阶导数来拟合树的结构。二阶信息的引入让模型能更精确地逼近真实损失收敛也更快。第二它在目标函数里加上了正则化项包括叶子节点数量、叶子权重的L2范数这个约束让单棵树更加保守泛化能力更强。第三它在特征分裂时做了预排序训练前先对特征值排序这样分裂点的查找效率大幅提升。LightGBM跟XGBoost走的路线不一样。XGBoost用的是按层分裂level-wise每一层都要遍历所有特征的所有取值LightGBM用的是按叶子分裂leaf-wise每次只选增益最大的叶子节点去分裂配合直方图算法把特征值离散化成固定数量的bin特征分裂的候选点大幅减少训练速度快了很多很多。但需要注意leaf-wise的策略在小数据集上容易过拟合所以LightGBM提供了max_depth限制参数工程上建议不要设置成无穷大。我拿一个实际操作经验来说明选型差异。在我处理百万级样本的结构化数据时XGBoost的默认参数就已经能跑出不错的结果但训练时间可能要好几个小时换成LightGBM之后同样精度的模型只需要十几分钟。但在处理样本量比较小比如只有几千条的数据时LightGBM反而容易比XGBoost过拟合得更厉害这是因为leaf-wise贪婪分裂在小样本上会更快地学习到噪声。所以我的建议是**如果追求训练速度和内存开销优先LightGBM如果在意模型稳定性和默认参数的表现XGBoost依然是稳妥选择。**至于CatBoost它对类别特征有原生的处理方式适合大量类别特征的数据集比如广告点击预测里的用户ID、商品ID这类高基数类别特征。2.4 神经网络分类模型从全连接到卷积再到Transformer讨论神经网络做分类任务时如果只盯着“前面接几个全连接层”这种单一路径会漏掉今天这个领域最重要的分支结构。先看最基础的全连接网络。它的结构很简单输入层把特征向量喂进去中间通过若干全连接层进行非线性变换最后一层接softmax输出每个类别的概率。对于表格类数据——就是GBDT统治的那个领域——全连接网络能做的事相对有限它的优势主要体现在它能自动学习特征的高阶交互但缺点是参数量很大训练需要的数据量也多而且没有利用数据本身的局部结构信息。到了图像分类这个领域全连接网络的不足就很明显了。一张224x224的RGB图片展平之后是一个15万维的向量全连接第一层如果接512个神经元光这一层的参数量就是7680万个训练得不偿失。卷积神经网络CNN的出现恰好解决了这个问题。CNN的核心理念是局部感知和参数共享每个卷积核只需要在局部区域做感受野内的特征提取参数在不同位置共享模型参数爆炸的问题被解决了同时卷积操作天然对平移、缩放等形变具有稳定性。卷积网络的发展史其实就是图像分类比赛刷榜的历史。从AlexNet在2012年ImageNet上把错误率大幅压低开始到VGG用“更小的卷积核拿更深的网络”卷了起来再到ResNet提出残差连接机制使得上百层的网络可以稳定训练。**残差连接这个想法值得展开讲它不直接拟合目标函数F(x)而是拟合残差F(x)-x这个“至少恒等映射兜底”的机制让深层网络在反向传播时梯度可以沿捷径传递缓解了梯度消失问题。**后来Transformer架构里的每一层几乎都标配了残差连接就是在复用这个思想。对文本分类而言主流路线已经换成了预训练语言模型路线BERT、RoBERTa这些模型用大规模无监督语料预训练再在下游任务上做微调。而在推荐、搜索场景里神经网络往往长成“Embedding层多层神经网络”的样子先把稀疏离散特征映射成稠密向量再喂给神经网络学习特征交互最后输出CTR预估的点击概率。3. 实操过程与核心环节实现3.1 一个图片分类项目的完整落地流程我拿“基于卷积神经网络的手写数字识别”这个经典项目来拆解一下实操过程虽然它在今天已经是入门级任务但完整的工程链路很有参考价值。第一步是数据预处理。MNIST数据集每张图是28x28的灰度图值域在0到255。我习惯先把值归一化到[0,1]区间再考虑要不要做标准化。归一化能让梯度更新更平稳训练收敛速度更快。from tensorflow import keras (x_train, y_train), (x_test, y_test) keras.datasets.mnist.load_data() x_train x_train.astype(float32) / 255.0 x_test x_test.astype(float32) / 255.0第二步是搭建卷积网络。经典的LeNet结构在这个任务上效果很好卷积层提取局部纹理特征池化层做空间降采样全连接层做分类。核心参数我给出一个可以直接跑的配置model keras.Sequential([ keras.layers.Conv2D(32, kernel_size(3,3), activationrelu, input_shape(28,28,1)), keras.layers.MaxPooling2D(pool_size(2,2)), keras.layers.Conv2D(64, kernel_size(3,3), activationrelu), keras.layers.MaxPooling2D(pool_size(2,2)), keras.layers.Flatten(), keras.layers.Dropout(0.5), keras.layers.Dense(10, activationsoftmax) ])这里有几个细节值得解释。第一个卷积层用了32个3x3的卷积核输入通道是1灰度图输出32个特征图。第一个池化层把28x28变为14x14第二个池化层变为7x7。到了Flatten层64个7x7的特征图展平之后是3136维接Dropout(0.5)是为了防止全连接层过拟合这是经验值我试过0.3到0.7的范围0.5通常表现最稳。第三步是编译和训练。优化器我常用Adam初始学习率设0.001损失函数用交叉熵。训练轮数设置20到30轮之间配合早停策略防止在验证集上过拟合。model.compile(optimizeradam, losssparse_categorical_crossentropy, metrics[accuracy]) model.fit(x_train, y_train, batch_size128, epochs30, validation_split0.2)这套简单的CNN在MNIST测试集上可以轻松跑到99%以上的准确率。注意这里有个关键认知**MNIST在今天已经不是有挑战性的数据集了99%以上的准确率并不代表模型多厉害它只是一个入门练手项目。**如果你要挑战真实场景的图像分类CIFAR-10或ImageNet子集更有参考价值。3.2 GBDT实战代码解析从原始特征到概率输出我这里的GBDT示例用的是常见的开源实现它的接口很简单关键是理解每一步在做什么。假设我们要做一个二分类任务用户是否会再次购买。from sklearn.model_selection import train_test_split from sklearn.ensemble import GradientBoostingClassifier from sklearn.metrics import classification_report, roc_auc_score X_train, X_test, y_train, y_test train_test_split(X, y, test_size0.2, random_state42) gbdt GradientBoostingClassifier( n_estimators100, learning_rate0.1, max_depth3, subsample0.8, random_state42 ) gbdt.fit(X_train, y_train) y_pred_proba gbdt.predict_proba(X_test)[:, 1] y_pred (y_pred_proba 0.5).astype(int) print(classification_report(y_test, y_pred)) print(AUC:, roc_auc_score(y_test, y_pred_proba))这个例子看起来很基础但里面每个参数的选择都有讲究。max_depth设成3是控制树的复杂度防止单棵树太强。subsample等于0.8相当于每轮训练只随机采样80%的样本和随机森林里booststrap的思路类似这种随机性可以有效提升泛化能力。learning_rate设0.1是权衡了效果和训练速度之后的常用选择——学习率越小越精准但需要的树就越多。有一个经验我强烈建议你记住**训练GBDT时先固定学习率然后通过交叉验证找最优的树数量n_estimators最后再微调max_depth和min_samples_leaf。**别一来就所有参数一起搜维度太多反而分不清哪个参数起了作用。3.3 循环神经网络RNN的核心公式与文本分类场景词向量时代到Transformer时代之间RNN曾经是序列建模的主力。标题里提到的“标准循环神经网络vanilla RNN核心公式”这个东西学起来很容易被公式吓到但实际上拆开来看就三个步骤。假设在时间步t输入向量是x_t上一时刻的隐藏状态是h_{t-1}那么当前隐藏状态h_t的计算公式是h_t tanh(W_hh * h_{t-1} W_xh * x_t b_h)其中W_hh是隐藏状态到隐藏状态的权重矩阵W_xh是输入到隐藏状态的权重矩阵b_h是偏置项。当前时刻的输出可以表示为y_t W_hy * h_t b_y这就是所谓“循环”的含义——**同一个权重矩阵W_hh在每一个时间步上被反复使用每个时刻的隐藏状态既依赖当前输入又携带了前面所有时刻的信息。**我习惯用一个比喻来理解它RNN读序列的方式和人逐字阅读一句话很像读完“中”再读“国”读“国”的时候脑子里还带着“中”这个信息读到最后脑子里装的是整句话的语境。但vanilla RNN有个致命问题当序列很长的时候反向传播经过很多时间步梯度要么爆炸要么消失。梯度爆炸可以用梯度裁剪缓解梯度消失就麻烦很多它意味着早前时间步的信息很难传递到后面模型学不到长距离依赖。为了对付这个问题工程师们设计了LSTM和GRU。我自己的经验是**除非是做教学研究必须从vanilla RNN起步否则实际项目里直接上LSTM或GRU就对了。**GRU是LSTM的精简版参数更少、训练更快在数据量不是特别大的场景里效果几乎一样。但到了今天又在大量NLP任务中直接用Transformer比用RNN更稳妥毕竟并行计算的能力差距摆在那里。3.4 物理信息神经网络PINN把科学原理嵌进损失函数最近“物理信息神经网络”这个概念特别火我在热词里看到它出现得很频繁。简单说PINNPhysics-Informed Neural Network是一种把物理定律作为约束条件嵌入神经网络训练的模型设计思路很直接如果已知某个系统满足一个偏微分方程比如流体运动满足Navier-Stokes方程那就在训练网络的损失函数里加上一个“物理残差项”让网络的预测结果不仅要匹配观测数据还要尽量满足物理方程。它的损失函数长这样L_total L_data λ * L_physicsL_data是预测值和真实观测值之间的误差例如MSEL_physics是对偏微分方程残差的约束λ是两者的权重。举个例子如果你要预测一根杆子上的温度分布已知热量传导满足热传导方程你就可以在训练数据之外把热传导方程的内部点残差当作额外的监督信号喂给网络。这样即使训练数据很少网络也能因为物理约束的加入而学到合理的解。这个思路在做实验数据获取成本特别高的领域很有价值——比如流体力学、材料科学、医疗影像里的参数反演问题。不过注意PINN的训练难度比普通神经网络要高不少。物理残差项涉及网络输出对输入求偏导这意味着要计算二阶甚至更高阶的梯度内存开销大训练也容易不稳定。我的建议是上手PINN时先选择一维或二维的简单方程比如一维热传导或一维波动方程把流程走通之后再扩展到复杂场景。4. 常见问题与排查技巧实录4.1 明明训练集准确率99%测试集却崩了怎么办这个现象在分类模型里太常见了我在实际项目里踩过很多次。三个最可能的罪魁祸首分别是数据泄露、过拟合、训练测试分布不一致。数据泄露是最危险的。我遇到过一个经典案例在做特征工程时不小心把target的滞后变量也当成特征放进了模型结果训练AUC飙到0.99上线之后效果一塌糊涂。排查方法是检查特征列表有没有什么字段名字看起来就像“结果”的。过拟合的应对手段已经比较成熟。GBDT方面降低max_depth、提高min_samples_leaf、增大subsample的随机性、减少n_estimators或者调高learning_rate配合早停都是有效手段。神经网络方面Dropout、权重衰减、数据增强、BatchNorm这块也能有效抑制过拟合。我一般用一个原则来判断如果训练集和验证集的指标差距很大说明正则化不够优先增大Dropout比例或调大正则系数。至于训练测试分布不一致很多数据泄露实际上是这个原因。这里补充一个实操技巧**在划分数据集时尽量按时间切分而不是随机切分。**比如用户行为数据如果在时间上存在漂移随机切分会让模型“偷看”到未来信息测试集上表现很好但真正部署时效果会大打折扣。4.2 分类结果严重倾向多数类少数类几乎全错模型训练完一看混淆矩阵少数类的召回率几乎为0。这是典型的样本不均衡问题。理论上讲GBDT处理不均衡的能力其实不弱因为它是围绕损失函数最小化的算法不会像普通逻辑回归那样容易被多数类带偏但如果正负样本比例悬殊到1:1000这种量级任何模型都需要额外处理。我按效果从弱到强列几个常用方案调整class_weight给少数类更高的权重让模型更关注它们。在GBDT中调高subsample比例或者对少数类做过采样比如SMOTE。调低分类阈值这类方法在很多场景里直接能提高召回率。换用评估指标比如AUC或F1来指导训练。其中调阈值这条在工程上尤其好用。先训练模型得到概率分布然后在验证集上寻找满足业务要求的阈值不需要重新训练。我有个习惯是拿验证集画出PR曲线然后根据曲线找“肘部”位置作为阈值候选。4.3 网络结构的“正确性”为什么别人的网络到我的数据上不work很多初学者喜欢从GitHub拉一个效果很好的网络结构直接套到自己的数据上结果效果非常差于是开始怀疑代码有bug或者框架有问题。实际上绝大多数问题出在输入尺寸不匹配和数据规模不匹配上。举个例子用ImageNet预训练的ResNet做特征提取输入默认是224x224三通道。如果你自己的数据集图片分辨率只有96x96直接resize到224x224是可以的但这种强行放大可能丢细节不如直接改网络第一层的输入尺寸重新训练前面的卷积层只在分类层做微调。同理如果你数据量只有几千张小图拿ResNet152这种超深网络去训练哪怕冻结前面大部分层照样容易过拟合。这时换成ResNet18或者MobileNet效果反而更好。还有一类问题是激活函数和梯度相关。网络层数加深以后如果全部用sigmoid做激活函数浅层梯度很容易消失。现在的主流动法是在隐藏层用ReLU或GELU只在最后一层用softmax做输出。实用建议很直接复现别人的网络时先从论文或官方实现里把超参数抄一遍跑通一个baseline再去微调。别一上来就改结构、改学习率、改优化器每一次只改一个变量这样你才知道瓶颈在哪里。4.4 训练损失震荡不收敛的排查清单神经网络训练时最常见的问题之一就是loss反复横跳不下降或者是陷在平台期很久没有变化。我整理的排查顺序大致是下面这样的检查数据预处理。特征数值太大或太小都会影响收敛常用的做法是标准化到均值为0方差为1。图像数据一般归一化到[0,1]或者[-1,1]。检查学习率。学习率太大的话loss是震荡甚至发散的太小时loss下降特别慢。可以按0.1、0.01、0.001、0.0001几个数量级逐一试一遍找到最有下降趋势的那个。检查优化器设置。Adam对学习率相对不敏感但它内部维护的动量状态在超参数变化时也可能带来不稳定必要的时候重置优化器状态。检查梯度和数值稳定性。有用LogSumExp替代softmax交叉熵的细节规避负无穷带来的NaN问题。检查模型实现。如果是自己写的自定义层务必检查前向传播里的维度、转换以及各类操作的梯度路径有时候nan就来自某些除法操作分母为0。这些排查顺序背后有个价值观**先从最简单的因素查起把环境和数据弄干净了再怀疑模型本身。**很多情况下你以为的模型问题实际上是学习率问题。5. 模型选型与后续扩展建议5.1 GBDT还是神经网络到底怎么选这个问题是很多人在实际项目里反复纠结的痛点我给出一套“按场景直接套”的选择逻辑。如果你是处理表格型数据字段以数值型和类别型为主数据规模在几万到几千万之间那GBDT及其进化版本XGBoost、LightGBM、CatBoost是首选。原因很简单这类数据中特征之间的非线性交互模式复杂树模型天然擅长捕捉这些规律而且树模型对缺失值、量纲差异、异常值都有较强的鲁棒性不需要做大量的特征工程和归一化处理。如果你是处理图像、语音、文本这类非结构化数据直接上神经网络。卷积网络处理图像、Transformer处理文本、WaveNet或类似结构处理语音这些已经在工业界验证成熟。在这类数据上用树模型等于让模型从一小块一小块的像素里去理解形状效率极低。还有一种混合情况比如推荐系统里的CTR预估。这里的特征既有用户画像、物品属性等结构化特征又有行为序列等非结构化特征。实践中比较稳妥的方案是向量特征用Embedding神经网络部分手工特征用树模型部分最后一层融合输出。工业界不少广告系统就是这么做的效果比单一模型要好。所以别迷信“某个模型碾压所有模型”这种说法模型选择永远是基于数据形态和业务目标去做取舍。5.2 从分类模型到大模型这条技术路线其实有迹可循把眼光放远一点当前大模型狂飙的背景下分类模型并没有变成“旧技术”反而在以一种新的形态重生。以BERT为代表的预训练语言模型本质上解决的仍然是分类问题——只不过下游任务被做成了分类结构。模型先在大规模语料上做自监督预训练比如预测上下文中的某个词然后在下游任务里把[CLS]这个标志位的输出接一个全连接层加上softmax输出各个类别的概率。这个范式深刻揭示了分类模型的底层逻辑并不仅限于小模型只要表征足够好分类头甚至可以是一个简单的线性层。对于没有足够算力去做大模型预训练的个人开发者来说我建议的路线是先把GBDT和CNN/RNN这些经典模型吃透理解数据、特征、损失函数和评估指标之间的关系然后再去接触Transformer架构和预训练模型你会有一种“原来这些新东西不过都是在老配方上加了新调料”的通透感。这其实就是《模型不玄学》这一章想传递的核心信息——**分类模型的本质问题从未变过给定输入输出概率然后在不确定性里做决策。**变化的只是表征能力和计算范式。5.3 个人实操中的一点延伸多标签分类与多任务学习一线工作中除了最常规的二分类和多分类还会遇到多标签问题。比如一篇文章可能同时属于科技、互联网、职场三个标签这时候就不是softmax输出而是每个标签各接一个sigmoid用二分类交叉熵分别训练。多标签分类的实操细节往往会卡在评估上因为accuracy不好用了。我惯用的做法是算每个标签的AUC再取平均或者直接汇报F1的macro平均。这个逻辑和单标签分类有本质区别原因是各个标签的正负样本比例可能差异很大将其当作独立二分类问题分别评估会更合理一些。另外如果你的业务场景里多个分类任务之间存在共享特征可以考虑多任务学习。比如电商场景里点击率预估和转化率预估就可以共享底层的Embedding和隐藏层只在上层分开接预测头。这样做的好处是点击率样本量大、转化率样本量小两者共享底层可以让小样本任务借到大样本任务的建模能力这在实践里经常有奇效。包括GBDT和神经网络融合实际上也是一种多任务/多模型协作的思想延伸。这一章写到这里我个人最大的体会其实是**别被“新模型”迷了眼分类问题本身从来没有变简单过但它的工具箱一直在变厚。**理解清楚每个工具背后的假设和适用边界的区别比多学一个新的网络结构更有价值。
觉得有用,分享给同行:

为您的企业打造数字门面

稳重轻奢商务风格,端正雅致视觉,长效耐看不易过时。

立即咨询 →