引言
训练一个模型很容易,但「这个模型到底靠不靠谱、该不该上线」是更难的问题。评估不是跑完 fit() 后打印一个分数就结束——它决定着你的一切后续决策:特征要不要删、超参往哪调、能不能换更强的模型、是不是已经过拟合了。
本文把模型评估拆成四层递进:数据划分(为什么必须有独立的测试集)、交叉验证(怎么把有限的验证做得更稳)、偏差方差与学习曲线(怎么判断是欠拟合还是过拟合、对症下药)、超参调优(网格/随机搜索怎么避免调出「虚假的最好」)。最后给出一张「评估 → 诊断 → 决策」的完整工作流图。
前置:[[ml]] 专题的分类/回归基础(https://plumephp.com/ml-supervised-regression/、https://plumephp.com/ml-supervised-classification/)。本专题聚焦「评估方法论」的动手实践。
目录
- 1. 为什么评估比训练更难
- 2. 数据划分:训练/验证/测试三件套
- 3. 交叉验证:把有限数据用出信心
- 4. 偏差方差权衡与欠拟合/过拟合
- 5. 学习曲线:诊断问题的利器
- 6. 缓解过拟合的常用手段
- 7. 超参调优:网格搜索与随机搜索
- 8. 模型对比与最终选择
- 9. 总结:评估驱动决策的工作流
- 延伸阅读
1. 为什么评估比训练更难
1.1 目标不是「记住训练数据」
模型的目标是泛化——在没见过的数据上表现好。训练集上的高分可能是「背答案」。评估的本质是回答:换一批新数据,它还行不行?
1.2 评估的三重陷阱
| 陷阱 | 表现 | 解法 |
|---|---|---|
| 用训练集评估 | 分数虚高 | 独立测试集 |
| 反复用测试集 | 测试集被「过拟合」 | 再留一份真测试集 |
| 单次切分运气 | 分数不稳定 | 交叉验证 |
2. 数据划分:训练/验证/测试三件套
2.1 三层划分
全部数据
├── 训练集 (60-80%) → 训练模型,调参
├── 验证集 (10-20%) → 选模型/选超参(反复使用)
└── 测试集 (10-20%) → 最终评估(只用一次!)
2.2 为什么需要「验证集」而不是直接在测试集上调参
调参本质是在测试集上试错。试得越多,测试集越失真。所以选超参用验证集,测真实水平才用测试集。
from sklearn.model_selection import train_test_split
X_temp, X_test, y_temp, y_test = train_test_split(
X, y, test_size=0.2, random_state=42)
X_train, X_val, y_train, y_val = train_test_split(
X_temp, y_temp, test_size=0.25, random_state=42) # 0.25*0.8=20%
print(X_train.shape, X_val.shape, X_test.shape)
2.3 数据量小怎么办:交叉验证
数据太少时三份切分每份都太薄,用交叉验证把验证环节复用起来。
3. 交叉验证:把有限数据用出信心
3.1 K 折交叉验证原理
数据切成 K 份(常 K=5 或 10)
第1轮:1折做验证,其余 K-1 折训练
第2轮:2折做验证,其余训练
...
第K轮:K折做验证
最终:K 个验证分数的平均 + 标准差
3.2 动手做
from sklearn.model_selection import cross_val_score
from sklearn.ensemble import RandomForestClassifier
rf = RandomForestClassifier(n_estimators=100, random_state=42)
scores = cross_val_score(rf, X, y, cv=5, scoring='roc_auc')
print("每折 AUC:", scores.round(3))
print("平均 AUC:", scores.mean().round(3))
print("标准差:", scores.std().round(3))
3.3 读懂结果
- 平均分高 + 标准差小 → 模型稳定可靠。
- 平均分高 + 标准差大 → 对数据划分敏感,可能要加正则或减复杂度。
3.4 分层交叉验证:类别不平衡时的必要选项
from sklearn.model_selection import StratifiedKFold
cv = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)
scores = cross_val_score(rf, X, y, cv=cv, scoring='roc_auc')
分类任务里,分层交叉验证保持每折类别比例,比普通 K 折更可靠。
4. 偏差方差权衡与欠拟合/过拟合
4.1 两个错误来源
| 来源 | 直觉 | 表现 |
|---|---|---|
| 偏差 | 模型太简单,学不动 | 训练分也低(欠拟合) |
| 方差 | 模型太灵活,学过头 | 训练高、测试低(过拟合) |
训练分高 + 测试分高 → 好模型
训练分低 + 测试分低 → 欠拟合(高偏差)
训练分高 + 测试分低 → 过拟合(高方差)
4.2 诊断表
| 现象 | 诊断 | 处方 |
|---|---|---|
| 训练低、测试低 | 欠拟合 | 加特征/换强模型/减正则 |
| 训练高、测试低 | 过拟合 | 加正则/减特征/加数据/减复杂度 |
| 训练中、测试中 | 正常 | 微调提升 |
4.3 一个直观例子:多项式次数
import numpy as np
from sklearn.preprocessing import PolynomialFeatures
from sklearn.linear_model import LinearRegression
from sklearn.pipeline import make_pipeline
np.random.seed(42)
x = np.linspace(-3, 3, 100)
y_true = 0.5*x**2 + x
y = y_true + np.random.normal(0, 2, 100)
for degree in [1, 2, 15]:
pipe = make_pipeline(PolynomialFeatures(degree), LinearRegression())
pipe.fit(x.reshape(-1,1), y)
# degree=1 欠拟合,degree=2 合适,degree=15 过拟合
5. 学习曲线:诊断问题的利器
5.1 学习曲线是什么
画「训练样本量 → 训练分/验证分」的变化:欠拟合与过拟合的曲线形状截然不同。
from sklearn.model_selection import learning_curve
train_sizes, train_scores, val_scores = learning_curve(
rf, X, y, cv=5, train_sizes=[0.2, 0.4, 0.6, 0.8, 1.0],
scoring='roc_auc')
train_mean = train_scores.mean(axis=1)
val_mean = val_scores.mean(axis=1)
import matplotlib.pyplot as plt
plt.plot(train_sizes, train_mean, 'o-', label='训练分')
plt.plot(train_sizes, val_mean, 'o-', label='验证分')
plt.xlabel('训练样本量')
plt.ylabel('AUC')
plt.legend()
plt.show()
5.2 两种典型形状
| 曲线形态 | 诊断 | 对策 |
|---|---|---|
| 训练分高、验证分低,两线间隔大 | 过拟合(方差大) | 加正则/减复杂度/加数据 |
| 两线都低且收敛到低处 | 欠拟合(偏差大) | 换更强模型/加特征 |
| 两线都高且接近 | 健康 | 继续微调 |
学习曲线的核心价值:判断「加数据」能不能救回来。若两线间隔大但仍在收敛,加数据有效;若已平躺,加数据没用,得换模型。
6. 缓解过拟合的常用手段
6.1 从「复杂度」与「数据」两端下手
| 手段 | 方向 | 例子 |
|---|---|---|
| 增加正则 | 压复杂度 | L2(岭)、L1(Lasso)、Dropout |
| 减少特征 | 压复杂度 | 特征选择 |
| 简化模型 | 压复杂度 | 降多项式次数、减树深 |
| 增加数据 | 提泛化 | 收集数据/数据增强 |
| 集成 | 降方差 | 随机森林、Bagging |
6.2 代码示例
# 决策树限制深度 = 减复杂度
from sklearn.tree import DecisionTreeClassifier
dt_deep = DecisionTreeClassifier(max_depth=None) # 易过拟合
dt_lim = DecisionTreeClassifier(max_depth=5) # 控制复杂度
# 集成降低方差
from sklearn.ensemble import RandomForestClassifier
rf = RandomForestClassifier(n_estimators=200, max_depth=8, random_state=42)
7. 超参调优:网格搜索与随机搜索
7.1 为什么不能直接「试最好的」
在验证集上试 100 组参数,选最好的一组,这组参数本身就是在「过拟合验证集」。用交叉验证 + 独立测试集能降低这种风险,但试太多仍会虚高。
7.2 网格搜索(GridSearchCV)
from sklearn.model_selection import GridSearchCV
param_grid = {
'n_estimators': [50, 100, 200],
'max_depth': [5, 10, None],
'min_samples_split': [2, 5, 10],
}
search = GridSearchCV(
RandomForestClassifier(random_state=42),
param_grid, cv=5, scoring='roc_auc', n_jobs=-1)
search.fit(X, y)
print("最优参数:", search.best_params_)
print("最优交叉验证分:", search.best_score_.round(3))
7.3 随机搜索(RandomizedSearchCV)
参数空间大时,网格搜索穷举太慢;随机搜索随机采样,通常更快找到近似最优。
from sklearn.model_selection import RandomizedSearchCV
from scipy.stats import randint, uniform
param_dist = {
'n_estimators': randint(50, 300),
'max_depth': randint(3, 15),
'min_samples_split': randint(2, 10),
}
search = RandomizedSearchCV(
RandomForestClassifier(random_state=42),
param_dist, n_iter=50, cv=5, scoring='roc_auc', random_state=42, n_jobs=-1)
search.fit(X, y)
print("最优:", search.best_params_)
7.4 调参铁律
- 验证集/交叉验证调参,测试集只做最终评估。
- 参数范围基于领域直觉,别盲目放大。
- 记录每次实验(参数、分数、数据划分),可复现。
8. 模型对比与最终选择
8.1 公平对比:同一数据划分、同一指标
models = {
'LogisticRegression': LogisticRegression(max_iter=1000),
'DecisionTree': DecisionTreeClassifier(max_depth=8, random_state=42),
'RandomForest': RandomForestClassifier(n_estimators=100, random_state=42),
'GradientBoosting': GradientBoostingClassifier(random_state=42),
}
for name, model in models.items():
scores = cross_val_score(model, X, y, cv=5, scoring='roc_auc')
print(f"{name:20s} AUC={scores.mean():.3f} ±{scores.std():.3f}")
8.2 选型的三个维度
| 维度 | 考量 |
|---|---|
| 性能 | 交叉验证平均分与稳定性 |
| 解释性 | 业务要不要「讲得清为什么」 |
| 成本 | 训练/推理时间、内存、部署难度 |
8.3 最终决策
# 选定的模型只在最终测试集上评估一次
best = RandomForestClassifier(n_estimators=200, max_depth=10, random_state=42)
best.fit(X_train, y_train)
final_auc = roc_auc_score(y_test, best.predict_proba(X_test)[:, 1])
print("最终测试集 AUC:", round(final_auc, 3))
9. 总结:评估驱动决策的工作流
9.1 完整流程图
数据划分(留测试集) → 交叉验证训练/调参 → 诊断(学习曲线/偏差方差)
→ 对症优化(正则/加数据/换模型) → 选模型 → 测试集最终评估(仅一次!)
9.2 决策速查
| 情况 | 行动 |
|---|---|
| 训练低、验证低 | 换强模型/加特征 |
| 训练高、验证低 | 加正则/减复杂度 |
| 验证分波动大 | 分层CV/加数据 |
| 调参后验证分涨、测试分跌 | 调参过度,回归保守参数 |
9.3 一句话心法
评估的本质是对抗「自我欺骗」:独立数据 + 稳定验证 + 一次测试,三层防线守住的,是模型真实泛化能力的真相。
延伸阅读
- https://plumephp.com/ml-supervised-classification/ — 评估指标(精确率/召回率/F1/AUC)详解
- https://plumephp.com/ml-supervised-regression/ — 回归评估指标(MSE/MAE/R²)
- https://plumephp.com/ml-feature-engineering/ — 特征侧优化与评估协作
- [[ai-ml]] 专题的模型评估深度文章
- scikit-learn 模型选择文档
继续阅读
探索更多技术文章
浏览归档,发现更多关于系统设计、工具链和工程实践的内容。