交叉验证和网格搜索之代码实现
1 课程概览
本课通过代码实现交叉验证和网格搜索,对KNN算法的超参数n_neighbors进行调优。详细讲解GridSearchCV的使用流程、参数含义、返回属性,以及如何获取最优超参数组合。同时强调该方法仅为超参数调优的一种思路,并非绝对最优。
2 核心概念与定义
- GridSearchCV:网格搜索+交叉验证的API,自动寻找最优超参数。
- param_grid:参数字典,定义超参数可能出现的值。
- cv:交叉验证折数,如cv=4表示四折交叉验证。
- best_score_:交叉验证中最好的得分。
- best_estimator_:最优超参数对应的模型对象。
- cv_results_:每次交叉验证的详细结果集。
3 算法与模型详解
3.1 交叉验证原理回顾
将训练集分成N份,每次取一份做验证集,其余N-1份做训练集,进行N次评估取平均值。
四折交叉验证流程:
| 轮次 | 验证集 | 训练集 | 准确率 |
|---|---|---|---|
| 第1次 | 第1份 | 第2、3、4份 | 准确率1 |
| 第2次 | 第2份 | 第1、3、4份 | 准确率2 |
| 第3次 | 第3份 | 第1、2、4份 | 准确率3 |
| 第4次 | 第4份 | 第1、2、3份 | 准确率4 |
最终得分 = 四次准确率的平均值
3.2 网格搜索原理
接收超参数可能出现的值,对每个值进行交叉验证,获取最优超参数组合。
示例:n_neighbors=[1,2,3,5,7,10],结合四折验证:
- 总执行次数 = 6个参数 × 4折 = 24次
3.3 GridSearchCV参数说明
| 参数 | 说明 | 示例 |
|---|---|---|
| estimator | 要计算最优超参的模型对象 | KNeighborsClassifier() |
| param_grid | 超参数可能出现的值(字典) | {'n_neighbors': [1,3,5,7,10]} |
| cv | 交叉验证折数 | 4 |
3.4 GridSearchCV返回属性
| 属性 | 说明 |
|---|---|
| best_score_ | 交叉验证中最好的结果 |
| best_estimator_ | 最优超参数对应的模型对象 |
| cv_results_ | 每次交叉验证的详细结果 |
3.5 完整调优流程
- 加载数据集
- 切分训练集和测试集(8:2)
- 特征工程:标准化(训练集fit_transform,测试集transform)
- 模型训练:
- 创建模型对象(不指定超参数)
- 定义参数字典
- 创建GridSearchCV对象
- fit训练
- 打印最优超参组合
- 模型评估
4 数学原理与推导
4.1 交叉验证平均得分
$$\text{Score} = \frac{1}{N}\sum_{i=1}^{N} \text{Score}_i$$
4.2 网格搜索执行次数
$$\text{Total} = P \times N$$
- P:参数组合数
- N:折数
5 代码示例
from sklearn.datasets import load_iris
from sklearn.model_selection import train_test_split, GridSearchCV
from sklearn.preprocessing import StandardScaler
from sklearn.neighbors import KNeighborsClassifier
from sklearn.metrics import accuracy_score
def dm01_grid_search_cv():
# 1. 加载数据集
iris_data = load_iris()
# 2. 切分训练集和测试集(8:2)
x_train, x_test, y_train, y_test = train_test_split(
iris_data.data, iris_data.target,
test_size=0.2, random_state=22
)
# 3. 特征工程:标准化
transfer = StandardScaler()
x_train = transfer.fit_transform(x_train)
x_test = transfer.transform(x_test)
# 4. 模型训练
# 4.1 创建模型对象(不指定超参数)
estimator = KNeighborsClassifier()
# 4.2 定义参数字典
param_dict = {'n_neighbors': [i for i in range(1, 11)]}
# 4.3 创建GridSearchCV对象(4折交叉验证)
estimator = GridSearchCV(estimator, param_grid=param_dict, cv=4)
# 4.4 模型训练(必须有这一步)
estimator.fit(x_train, y_train)
# 5. 打印最优超参组合
print("最优评分:", estimator.best_score_)
print("最优估计器:", estimator.best_estimator_)
print("交叉验证结果:", estimator.cv_results_)
# 6. 模型评估
y_predict = estimator.predict(x_test)
print("准确率:", accuracy_score(y_test, y_predict))
if __name__ == '__main__':
dm01_grid_search_cv()
输出示例:
最优评分: 0.9666666666666667
最优估计器: KNeighborsClassifier(n_neighbors=3)
准确率: 0.9666666666666667
6 重难点与易错提醒
- ❗重点:创建模型对象时不指定超参数,由GridSearchCV寻找最优值。
- ❗重点:必须调用fit方法训练,否则无法获取best_score_等属性。
- ❗重点:cv=4表示四折交叉验证,每个超参数值都进行4次验证。
- ⚠️易错:忘记调用fit方法,导致best_score_属性不存在。
- ⚠️易错:param_grid是字典格式,键为参数名,值为参数可能出现的列表。
- 💡深入理解:网格搜索结果受test_size、random_state、cv等多个参数影响,不是绝对的。
7 课堂问答精选
Q: 网格搜索找到的最优超参数一定是最准确的吗?
A: 不一定。网格搜索结果受多个因素影响:
- test_size(测试集比例)
- random_state(随机种子)
- cv(折数) 不同的参数组合可能得到不同的最优超参数。网格搜索仅是寻找最优超参数的一种思路,不要过度依赖。
Q: 为什么网格搜索结果与预期不同?
A: 因为网格搜索的底层是分折数验证的,不同的折数、不同的随机种子、不同的测试集比例都会影响最终结果。例如:
- random_state=22,cv=4,最优n_neighbors=3
- random_state=25,cv=5,最优n_neighbors可能变为5或10
8 本课小结
- GridSearchCV API:estimator(模型)、param_grid(参数字典)、cv(折数)。
- 流程:创建模型→定义参数字典→创建GridSearchCV→fit训练→查看best_score_。
- 关键属性:best_score_(最优评分)、best_estimator_(最优模型)、cv_results_(验证结果)。
- 网格搜索仅为超参数调优的一种思路,结果受多个参数影响,不是绝对的。
9 延伸思考与实践
- 实践:修改cv值为5或10,观察最优超参数变化。
- 思考:为什么网格搜索结果不是绝对的?
- 预习:手写数字识别案例。