模型评估与选择:交叉验证、偏差方差、过拟合与调参实战

系统讲解机器学习模型评估与选择:训练/验证/测试集划分、交叉验证原理、偏差方差权衡与学习曲线、过拟合识别与缓解、网格搜索与随机搜索超参调优、模型对比与最终选型决策。

引言

训练一个模型很容易,但「这个模型到底靠不靠谱、该不该上线」是更难的问题。评估不是跑完 fit() 后打印一个分数就结束——它决定着你的一切后续决策:特征要不要删、超参往哪调、能不能换更强的模型、是不是已经过拟合了。

本文把模型评估拆成四层递进:数据划分(为什么必须有独立的测试集)、交叉验证(怎么把有限的验证做得更稳)、偏差方差与学习曲线(怎么判断是欠拟合还是过拟合、对症下药)、超参调优(网格/随机搜索怎么避免调出「虚假的最好」)。最后给出一张「评估 → 诊断 → 决策」的完整工作流图。

前置:[[ml]] 专题的分类/回归基础(https://plumephp.com/ml-supervised-regression/、https://plumephp.com/ml-supervised-classification/)。本专题聚焦「评估方法论」的动手实践。


目录


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 模型选择文档

继续阅读

探索更多技术文章

浏览归档,发现更多关于系统设计、工具链和工程实践的内容。

全部文章 返回首页

「ml」更多文章

  1. 集成学习实战:Bagging、随机森林、梯度提升与 Stacking
  2. 迁移学习实战:预训练模型、特征提取与微调全流程
  3. 计算机视觉入门实战:图像处理与 CNN 图像分类