2.14 模型评估与选择

1.训练集上的好成绩,到底值不值得信

说起来,前面这么多章我们把各种各样的模型都过了一遍,从最朴素的线性回归一直到boosting和聚类。模型训完了,手里总得有个数说它到底好不好。最直观的想法就是看它在训练集上的表现,训练集不就是模型反复见过的那批数据嘛,算出来的分数漂漂亮亮,看着也舒坦。可这件事恰恰最唬人。

我先讲一个我常给学生举的例子。小明备考期末考试,把过去十年的真题翻来覆去做,做到后来闭着眼睛都能答满分。结果真到了期末那张新卷子,分数一塌糊涂。问题出在哪呢,他其实把每道真题的答案都背下来了,并没有学到背后的知识点,题目换个数字或者换个问法就抓瞎。模型也是一模一样的脾气。决策树如果不加限制,能把训练集里每一个样本都单独分到一个叶子,训练准确率直奔 100%100\% 去(这里的 100%100\% 就是把所有训练样本都猜对),可这种本领搬到新数据上立刻现原形。这种把训练集学得太死、连噪声和细节都一并记进参数里的毛病,我们叫过拟合。

过拟合的根源,在于模型容量(模型能表达多复杂关系的那种能力)相对数据来说太富裕,把数据里那些偶然的、不具备普遍性的波动也当成规律学了进去。要把它揪出来,光看训练集没用,因为训练集的分数本来就被模型背得滚瓜烂熟。真正可靠的做法,是留出一部分模型从没见过的数据,用它来当试金石。这就有了两种必须留出的数据,一种叫验证集,一种叫测试集。

验证集用来在训练过程中挑模型、调超参数(超参数就是训练前由人指定的、不参与梯度更新的那些设置,比如决策树的最大深度)。测试集呢,训练彻底结束之后只跑一次,用来估计模型将来在真实新数据上能干多好。这两份数据都得严防死守,不能让模型提前偷看,更不能拿它们反复试错。我之前翻过一本讲机器学习实战的书,作者特别强调,测试集就像考场上那张密封的卷子,开考前谁都不许动,动一次就污染一次。

那么具体怎么切分这三份数据呢,最常见的做法是按比例切,比方说训练集占 60%60\%、验证集占 20%20\%、测试集占 20%20\%,比例可以按数据量大小灵活调整。数据量特别大的时候,验证集和测试集各留几千条就够估分了,剩下的全给训练。数据量小的时候呢,硬切走一部分就特别肉疼,下面我们就来谈这个。

2.K折交叉验证:把数据翻来覆去地用

数据量一大,切一点出去不痛不痒。可数据量小的时候,比如手头只有两三百条样本,再切个两成做验证集,验证集上就剩四五十条,估出来的分数抖动特别大。小张前阵子做一个小数据集的分类,同一套代码,今天验证集准确率 0.880.88,明天换一批划分就掉到 0.790.79,他自己都怀疑是不是代码出了bug。其实代码没错,错的只是评估方式不够稳。

这时我们用的法子叫K折交叉验证(K-fold cross validation)。思路说穿了,就是把数据翻来覆去地用,让每一条样本都轮上一遍验证集,再把多次结果取平均,分数就稳多了。具体怎么做呢,先把全部训练数据均匀切成K份(K就是份数这个数,常见的取 K=5K=5K=10K=10)。然后开始轮流,第一轮拿第 11 份当验证集,剩下 K1K-1 份合起来训练模型。第二轮拿第 22 份当验证集,剩下 K1K-1 份训练。如此类推,一共跑K轮。最后把K轮得到的K个分数取算术平均,作为这次评估的最终估计。

K个分数之间总会有差异,这份当验证集和那份当验证集结果不一定一样。这种差异本身也是有用的信息,可以用来判断模型有多依赖具体的数据划分。要是K个分数都挤在一起,说明模型比较稳定,换批数据也不会差太多。要是K个分数忽高忽低,那就要警惕了,要么模型容量不合适,要么数据本身有些古怪。

交叉验证里有一种特别严格的变体,叫留一法(leave-one-out)。它把每份只留一条样本当验证集,相当于 KK 等于样本总数的极端情况。它的估计准是准,可样本一多计算量就压死人了,所以平日里大家还是老老实实用 55 折、1010 折居多。

3.分类指标:光看准确率会被不均衡数据骗到

讲完了怎么切数据,我们再来谈指标本身。4.1 节我们聊过图像分类这种活,模型最后给每一个类别吐出一个分数。怎么把这个分数换算成一个具体的评估数字,里面门道不少。

最朴素的指标叫准确率(accuracy)。它就是全部样本里被猜对的那部分占比。设总样本数是 nnnn 表示样本总数),猜对的样本数是 n正确n_{\text{正确}},那么准确率就是:

accuracy=n正确n\text{accuracy}=\frac{n_{\text{正确}}}{n}

这个指标在类别均衡的时候很顶用。可一旦类别不均衡,准确率就开始骗人了。我有一位在医院做算法的朋友跟我吐槽过一件事。他们想用一个模型在体检人群里筛查某种罕见癌症,这种癌症的发病率大约是 1%1\%。他同事随手写了个特别简单的模型,不管三七二十一,对所有人都预测没病,跑出来准确率 0.990.99,欢天喜地觉得这模型神了。可这模型对真正的病人一个都没识别出来,放到临床上完全是场灾难。准确率 0.990.99 在这里成了一个体面的假象。

要拆穿这种假象,我们得把分类结果拆得更细,这就引出了混淆矩阵(confusion matrix)。以二分类为例,我们关心的是某一条样本的真实类别和模型预测的类别。真实类别有两种(正例和负例),预测也有两种(正例和负例),两两组合就是四个格子。这四个格子分别是:TP(true positive,真实是正例、预测也是正例的数量)、FP(false positive,真实是负例、却预测成了正例的数量)、FN(false negative,真实是正例、却预测成了负例的数量)、TN(true negative,真实是负例、预测也是负例的数量)。把这四个数凑成的那张二乘二的表格,就叫混淆矩阵。

有了这四个数,我们就能算出几个比准确率诚实得多的指标。先说精确率(precision),公式是:

precision=TPTP+FP\text{precision}=\frac{TP}{TP+FP}

分子是预测对的正例数,分母是所有被模型预测为正例的总数。精确率问的是,模型说是正例的那一批里,到底有多少真是正例。

再说召回率(recall),公式是:

recall=TPTP+FN\text{recall}=\frac{TP}{TP+FN}

分子还是预测对的正例数,分母换成所有真实的正例总数。召回率问的是,所有真正的正例里,模型抓回来多少。

F1是把这两者揉到一起的数,取精确率和召回率的调和平均(调和平均就是先求倒数、再取算术平均、最后再倒回来的一种平均方式):

F1=2×precision×recallprecision+recallF1=\frac{2\times\text{precision}\times\text{recall}}{\text{precision}+\text{recall}}

回到刚才那个罕见癌症的例子。一千人里十个人真有病,那个全判没病的模型 TP=0TP=0FP=0FP=0FN=10FN=10TN=990TN=990。代入进去,召回率是 00,F1也是 00,假象一下就被戳破了。这就是为什么医疗筛查这类场景,我们盯着召回率不放,宁可多花点钱多做几次复查,也别把真病人漏在家里。漏掉一个癌症病人,代价可能就是一条命。

反过来的例子也有。我认识一位做反垃圾邮件的朋友,他就更看重精确率。因为一旦把正常邮件误判成垃圾邮件,用户的重要通知可能就被丢进垃圾箱再也找不回来,这种误伤的体验非常糟糕。漏掉一两封垃圾邮件反倒无所谓。所以精确率和召回率之间到底怎么取舍,得看具体业务,并没有一套通用的答案。

4.ROC曲线与AUC:把阈值拉一遍看全貌

刚才那些指标都默认模型输出的是一个明确的类别。可很多模型吐出来的其实是一个介于 0011 之间的概率或者打分,到底判成正例还是负例,得我们自己拿一个阈值去切。阈值调低,更多样本会被判成正例。阈值调高,判成正例的就少了。光看某一个阈值下的指标,难免以偏概全。

ROC曲线(receiver operating characteristic curve)就是用来一眼看清不同阈值下表现的图。它的横轴是FPR(false positive rate,假正例率),纵轴是TPR(true positive rate,真正例率,其实就等于前面那个召回率)。它们的公式分别是:

FPR=FPFP+TNFPR=\frac{FP}{FP+TN} TPR=TPTP+FNTPR=\frac{TP}{TP+FN}

把阈值从最严调到最松,每调一次都能在图上点一个点,把这些点连起来就是ROC曲线。阈值最严的时候所有样本都判负例,TPTPFPFP 都落到 00,点落在左下角。阈值最松的时候所有样本都判正例,FNFNTNTN 都落到 00,点落在右上角。一条好模型的ROC曲线会尽量往左上角那个点靠,意思是它在保持低误报的同时还能抓住大量正例。

光凭一条曲线还是不好横向比较,于是我们再算曲线下方的面积,这就是AUC(area under curve)。AUC取值在 0011 之间,越接近 11 越好,等于 0.50.5 的时候模型和瞎猜差不多。说起来AUC还有一个特别直观的概率解释,它近似等于随机抽一个正例和一个负例时,模型给正例的打分高于负例的概率。AUC等于 0.90.9,意思就是十次这种两两比较里,平均有九次模型能把正例排在负例前面。

不过类别极不均衡的时候,比如前面那个发病率 1%1\% 的癌症,ROC曲线会显得格外好看,AUC虚高,给人一种模型很强的错觉。这种时候换PR曲线(precision recall curve)更靠谱。PR曲线横轴是召回率,纵轴是精确率,它对负例数量不那么敏感,反而能诚实地反映出在不均衡场景下模型的真实水平。所以做罕见病检测、欺诈识别这类活,看PR曲线心里更有底。

5.回归指标:MSE、MAE和R²

讲完了分类,我们再回到回归。还记得2.1 节我们用平方损失衡量预测值和真实值差多远吧,回归任务最常用的几个指标,跟平方损失其实是一脉相承的。

第一个是MSE(mean squared error,均方误差),公式是:

MSE=1ni=1n(yiy^i)2MSE=\frac{1}{n}\sum_{i=1}^{n}(y_i-\hat{y}_i)^2

这里 nn 是样本总数,ii 是样本编号,yiy_i 是第 ii 个样本的真实值,y^i\hat{y}_i 是模型给出的预测值(y^\hat{y} 顶上那个小帽子表示这是预测出来的),\sum 是求和符号,意思是从第 11 个样本到第 nn 个样本,把括号里的平方误差全加起来再除以 nn。MSE把每个误差平方后再平均,所以对那些离谱的大误差格外敏感。预测房价的时候,要是大多数房子误差几千块、唯独一栋豪宅预测偏了五十万,平方这一下就会把MSE整体顶高。

要是不想被这种少数大误差牵着鼻子走,可以换MAE(mean absolute error,平均绝对误差):

MAE=1ni=1nyiy^iMAE=\frac{1}{n}\sum_{i=1}^{n}|y_i-\hat{y}_i|

这里的 yiy^i|y_i-\hat{y}_i| 是误差的绝对值,绝对值符号 |\cdot| 把正负误差都拉成正的。MAE不放大极端误差,更稳健,也更贴近一个普通样本能犯多大的错。

还有一个特别讨喜的指标叫 R2R^2(决定系数)。它的直观含义是,模型比起最朴素的基线(就是用所有真实值的平均数 yˉ\bar{y} 去做预测)到底强了多少。yˉ\bar{y} 表示真实值的平均。公式是:

R2=1i=1n(yiy^i)2i=1n(yiyˉ)2R^2=1-\frac{\sum_{i=1}^{n}(y_i-\hat{y}_i)^2}{\sum_{i=1}^{n}(y_i-\bar{y})^2}

R2R^2 等于 11 的时候,模型完美预测。等于 00 的时候,模型只相当于拿平均值糊弄一下。要是连 00 都不到,那这个模型还不如直接输出平均值来得体面,挺丢人的。日常做回归,大家张口就报 R2R^2,因为它把误差和基线一对比,特别容易判断模型到底有没有学到位。

6.调参:网格搜索、随机搜索、贝叶斯优化

评估指标选定之后,下一步就是拿它去调超参数。深度学习模型超参数不算多,可传统机器学习里超参数一抓一大把,比如SVM的 CC(正则强度的倒数)、随机森林的树棵数、XGBoost的学习率。怎么在这些参数里挑出一组最合适的,本身就是一门学问。

最朴素的做法叫网格搜索(grid search)。它在每一个超参数上各预设几个候选值,再把所有组合穷举一遍,每一组都用交叉验证打分,挑分数最高的那组。比方说 CC{0.1,1,10}\{0.1,1,10\},最大深度取 {3,5,7}\{3,5,7\},一共 3×3=93\times3=9 组,全部跑一遍。这种方法直白好懂,配合Pipeline(把数据预处理和模型串成一条流水线的写法)和交叉验证用起来特别顺手。可一旦超参数一多,组合数量就指数膨胀,搜起来慢得让人望而却步。

随机搜索(random search)就是来缓解这个的。它不再老老实实把每个格子都走一遍,而是在预设的范围里随机抽固定数量的几组组合来试。听上去好像不太靠谱,可效果往往出乎意料地好,原因是大多数实际任务里,并不是每个超参数都同等重要,随机抽中好组合的概率比想象中高得多。同样的预算下,随机搜索覆盖到的参数空间通常比网格搜索宽。

再进一步就是贝叶斯优化(Bayesian optimization)。前两种要么纯随机要么死磕,贝叶斯优化则会学。它每跑完一组超参数,就更新一份对参数空间的概率估计,预测哪些区域更可能出好结果,下一组专门往那片区域钻。就好像一个有经验的采购员,跑了几次市场之后慢慢摸清哪几家店更实惠,后面的脚步就越来越准。它特别适合那种单次评估特别贵的场景,比如训练一个大模型要好几天的情况,跑一次就得心疼一次,自然希望每一组都尽量挑得准。

说到底,模型评估这件事,比训练本身更值得花心思。我之前看过一本讲运动员训练的书,里头有一句话大意是,平时训练得再漂亮,最终还是要拿到赛场上检验。机器学习也一样,所有的指标、所有的切分、所有的调参,都是为了让我们对模型上线之后的真实表现有一个不打折扣的判断。今天就先聊到这儿,下一篇我们看怎么把这些评估套路和sklearn的Pipeline串起来实操。

练习

Q1. 为什么类别不均衡时光看准确率会被骗?精确率和召回率分别问的是什么?

极端例子:发病率 1% 的癌症,模型对所有人都判没病,准确率仍有 0.99,但真正的病人一个没抓住,完全没用。精确率 TPTP+FP\frac{TP}{TP+FP} 问"模型说是正例的那批里有多少真是";召回率 TPTP+FN\frac{TP}{TP+FN} 问"所有真正的正例里模型抓回多少"。医疗筛查盯召回率(别漏病人),反垃圾邮件盯精确率(别误伤正常邮件)。

Q2. K 折交叉验证具体怎么操作?它比单切一次验证集好在哪?

把训练数据均匀切成 KK 份(常用 5 或 10),第 1 轮拿第 1 份当验证集、其余 K1K-1 份训练,第 2 轮换第 2 份当验证集,依此类推跑 KK 轮,最后把 KK 个分数取平均。好处是每条样本都轮上一遍验证集,估计更稳;尤其小数据集上单切一次分数抖动大,多切几刀取平均才靠得住。

Q3. MSE 和 MAE 在对待极端误差上有什么不同?R2=0.6R^2=0.6 说明什么?

MSE 把每个误差平方后再平均,对离谱的大误差格外敏感(一栋豪宅预测偏 50 万会把整体顶高);MAE 取绝对值,不放大极端误差,更稳健。R2=0.6R^2=0.6 表示模型比起"直接用平均值预测"这个基线,把误差压到了原来的 40%,说明模型比基线强但还远不完美(R2=1R^2=1 完美,=0=0 等于基线,小于 0 还不如基线)。

相关标签
机器学习模型评估交叉验证AUC