原文:Custom refit strategy of a grid search with cross-validation,The scikit-learn developers。本文依据 2026-10-05 读取的 scikit-learn 1.9.1 文档翻译整理;原例代码为 BSD-3-Clause,完整许可保留在本文下方。
网格搜索通常按某一项得分挑选参数,但实际需求可能同时关心误报、漏报和计算成本。GridSearchCV 的 refit 参数可以接收一个函数:让这个函数查看所有候选的交叉验证结果,返回最终候选的索引,再由网格搜索对象用选定参数在整个开发集上重新训练。
原例用一半带标签的数据完成模型选择,把另一半留作独立评估。它演示“先保证精确率,再保留召回率接近的模型,最后参考评分时间”的自定义流程。这些阈值和取舍是示范选择,不是普遍适用的最优策略。

把数字识别转成“是不是 8”
数据来自 scikit-learn 内置的 digits。为使例子容易理解,原文不做十分类,而是把目标变成布尔值:当前数字是否为 8。每张 8×8 像素图片展平成 64 维向量:
from sklearn import datasets
from sklearn.model_selection import train_test_split
digits = datasets.load_digits()
n_samples = len(digits.images)
X = digits.images.reshape((n_samples, -1))
y = digits.target == 8
print(
f"The number of images is {X.shape[0]} "
f"and each image contains {X.shape[1]} pixels"
)
X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=0.5, random_state=0
)
源页记录数据共有 1797 张图片,每张 64 个像素;该切分得到 898 个开发样本和 899 个评估样本。这里保留原例 test_size=0.5 与 random_state=0,原例没有传入 stratify。如果换成类别很少或严重失衡的数据,应重新设计切分,而不是照搬比例和随机种子。
精确率(precision)关心“被预测为 8 的样本里,有多少真的是 8”;召回率(recall)关心“实际为 8 的样本里,有多少被找出来”。它们评估的是正类,即这里的 True,不能混同为整个数据集的准确率。
理解自定义 refit 的输入与返回值
启用 scoring=["precision", "recall"] 后,cv_results_ 中包含两种指标的均值、折间标准差和排名,同时还有参数字典及拟合、评分耗时。refit 函数接收这个结果字典,返回候选在原数组中的索引。
原文用下面的辅助函数打印候选的指标。输出中的 std_test_precision 和 std_test_recall 是各候选跨交叉验证折的波动:
import pandas as pd
def print_dataframe(filtered_cv_results):
for mean_precision, std_precision, mean_recall, std_recall, params in zip(
filtered_cv_results["mean_test_precision"],
filtered_cv_results["std_test_precision"],
filtered_cv_results["mean_test_recall"],
filtered_cv_results["std_test_recall"],
filtered_cv_results["params"],
):
print(
f"precision: {mean_precision:0.3f} (±{std_precision:0.03f}),"
f" recall: {mean_recall:0.3f} (±{std_recall:0.03f}),"
f" for {params}"
)
print()
接下来的代码保留原文的实际筛选运算,仅缩短打印标签和翻译注释。它有几个值得先说明的边界:精确率使用严格大于 0.98;召回率筛选中的标准差是“候选模型平均召回率之间”的样本标准差;最终比较的是评分耗时,而非独立测量的纯推理延迟。
def refit_strategy(cv_results):
precision_threshold = 0.98
cv_results_ = pd.DataFrame(cv_results)
print_dataframe(cv_results_)
high_precision_cv_results = cv_results_[
cv_results_["mean_test_precision"] > precision_threshold
]
print_dataframe(high_precision_cv_results)
high_precision_cv_results = high_precision_cv_results[
[
"mean_score_time",
"mean_test_recall",
"std_test_recall",
"mean_test_precision",
"std_test_precision",
"rank_test_recall",
"rank_test_precision",
"params",
]
]
# 原例在候选均值之间计算标准差,并非取最佳模型的折间标准差。
best_recall_std = high_precision_cv_results["mean_test_recall"].std()
best_recall = high_precision_cv_results["mean_test_recall"].max()
best_recall_threshold = best_recall - best_recall_std
high_recall_cv_results = high_precision_cv_results[
high_precision_cv_results["mean_test_recall"] > best_recall_threshold
]
print_dataframe(high_recall_cv_results)
selected = high_recall_cv_results["mean_score_time"].idxmin()
print(high_recall_cv_results.loc[selected])
return selected
不能把这段实现描述成“在最佳模型折间一标准差内挑选”,更不能进一步称作标准误差规则。原例虽然保留了 std_test_recall 一列,却没有用它计算筛选阈值;真正参与计算的是 mean_test_recall.std()。这一区别会影响哪些候选被保留。
mean_score_time 是交叉验证中评分阶段的平均秒数,包含取得预测结果及计算评分等开销。它可以用于这个示例的相对取舍,但不等同于固定硬件、固定批量下独立测量的服务推理延迟,计时噪声还可能改变并列候选的排序。
定义 SVC 参数网格并拟合
原例搜索 RBF 核和线性核。RBF 核对两个 gamma 值和四个 C 值组合,共八项;线性核只搜索四个 C 值,共十二个候选:
from sklearn.model_selection import GridSearchCV
from sklearn.svm import SVC
scores = ["precision", "recall"]
tuned_parameters = [
{"kernel": ["rbf"], "gamma": [1e-3, 1e-4], "C": [1, 10, 100, 1000]},
{"kernel": ["linear"], "C": [1, 10, 100, 1000]},
]
grid_search = GridSearchCV(
SVC(), tuned_parameters, scoring=scores, refit=refit_strategy
)
grid_search.fit(X_train, y_train)
当前源页的默认交叉验证为 5 折,分类任务采用不打乱的分层折分。没有显式传 cv 的旧版本可能有不同默认值,复现时应记录软件版本与分割规则。先前的随机留出切分与内部交叉验证是两个不同层次。
原文记录的候选结果如下,保留三位小数及其折间标准差。它们是源站构建文档时的结果,不是本次实测:
| 核 / C / gamma | 精确率 | 召回率 |
|---|---|---|
| RBF / 1 / 0.001 | 1.000 ± 0.000 | 0.854 ± 0.063 |
| RBF / 1 / 0.0001 | 1.000 ± 0.000 | 0.257 ± 0.061 |
| RBF / 10 / 0.001 | 1.000 ± 0.000 | 0.877 ± 0.069 |
| RBF / 10 / 0.0001 | 0.968 ± 0.039 | 0.780 ± 0.083 |
| RBF / 100 / 0.001 | 1.000 ± 0.000 | 0.877 ± 0.069 |
| RBF / 100 / 0.0001 | 0.905 ± 0.058 | 0.889 ± 0.074 |
| RBF / 1000 / 0.001 | 1.000 ± 0.000 | 0.877 ± 0.069 |
| RBF / 1000 / 0.0001 | 0.904 ± 0.058 | 0.890 ± 0.073 |
| linear / 1 / 不适用 | 0.695 ± 0.073 | 0.743 ± 0.065 |
| linear / 10 / 不适用 | 0.643 ± 0.066 | 0.757 ± 0.066 |
| linear / 100 / 不适用 | 0.611 ± 0.028 | 0.744 ± 0.044 |
| linear / 1000 / 不适用 | 0.618 ± 0.039 | 0.744 ± 0.044 |
第一轮有五项精确率达标;第二轮排除了召回率约为 0.257 的候选,余下四项再比较评分时间。源页选中索引 6,即 {"C": 1000, "gamma": 0.001, "kernel": "rbf"},记录的 mean_score_time 为 0.005937 秒,平均召回率为 0.877206。它不意味着这组参数在其他运行环境中一定更快。
只在选择结束后使用独立评估集
GridSearchCV 已用所选参数在整个开发集上重新拟合,因此可以直接调用 grid_search.predict。这里的“整个”指传给 fit 的开发集,不包括留出的 X_test:
from sklearn.metrics import classification_report
print(grid_search.best_params_)
y_pred = grid_search.predict(X_test)
print(classification_report(y_test, y_pred))
源页的独立评估报告为:
precision recall f1-score support
False 0.99 1.00 0.99 807
True 1.00 0.87 0.93 92
accuracy 0.99 899
macro avg 0.99 0.93 0.96 899
weighted avg 0.99 0.99 0.99 899
原文也提醒,这个问题过于容易,超参数的性能平台较平坦,若干候选在指标上并列。因此,这个例子的价值在于展示选择机制,而不是提供通用的数字识别参数。不要根据这份留出报告反复调整阈值,再把同一留出集称为未参与模型选择的评估集。
编者修订:把空候选和 NaN 明确处理掉
原策略换数据后可能失效:没有候选精确率大于 0.98 时,筛选结果为空;仅剩一个候选时,pandas 默认样本标准差会是 NaN;所有召回率相同时,标准差为零,原来的严格大于又会把最佳候选排除。拟合或评分失败产生的非有限值也需要处理。
下面是另列的修订版。它保留“候选均值之间的样本标准差”这个原例语义,不暗中换成折间标准差;单候选时显式取零;召回率边界改为大于等于;无精确率达标者时抛出错误,而不是偷偷放宽 0.98 要求。返回值仍为原候选索引:
import numpy as np
import pandas as pd
def refit_strategy_checked(cv_results):
frame = pd.DataFrame(cv_results)
required = ["mean_test_precision", "mean_test_recall", "mean_score_time"]
finite = np.isfinite(frame[required].to_numpy(dtype=float)).all(axis=1)
valid = frame.loc[finite & (frame["mean_score_time"] >= 0)]
eligible = valid.loc[valid["mean_test_precision"] > 0.98]
if eligible.empty:
raise ValueError("没有有限且精确率严格大于0.98的候选;请审查搜索结果。")
recalls = eligible["mean_test_recall"]
spread = float(recalls.std(ddof=1)) if len(eligible) > 1 else 0.0
threshold = float(recalls.max()) - spread
shortlist = eligible.loc[recalls >= threshold]
if shortlist.empty:
raise ValueError("召回率筛选得到空集;请审查输入结果。")
return int(shortlist["mean_score_time"].idxmin())
要使用修订版,在前面的 GridSearchCV 中把 refit=refit_strategy 改成 refit=refit_strategy_checked。可另行选择 error_score="raise" 让拟合错误直接暴露,但这也改变了原例默认继续收集结果的行为。阈值不应由修订函数替业务决定。
使用可调用 refit 时,可通过 best_index_、best_params_ 和 best_estimator_ 取得结果;不要假定存在唯一的 best_score_,当前 API 对这种情况不提供该属性。保留的多指标结果应从 cv_results_ 明确读取。
代码审核:本轮仅做静态审查,未安装依赖、训练模型或测量性能。例中未发现硬编码秘密、命令注入或不可信模型反序列化;资源耗时和数据切分仍需在实际环境验证。没有发现其他问题不等于无漏洞。
作者与版权:The scikit-learn developers,Copyright © 2007–2026。代码来源 SPDX-License-Identifier: BSD-3-Clause,完整条件与免责声明见本文下方 及项目许可证。本文明确区分原例、原例输出和编者修订,不表示原作者为修订版背书。
scikit-learn 示例代码 BSD 3-Clause 许可全文
BSD 3-Clause License Copyright (c) 2007-2026 The scikit-learn developers. All rights reserved. Redistribution and use in source and binary forms, with or without modification, are permitted provided that the following conditions are met: * Redistributions of source code must retain the above copyright notice, this list of conditions and the following disclaimer. * Redistributions in binary form must reproduce the above copyright notice, this list of conditions and the following disclaimer in the documentation and/or other materials provided with the distribution. * Neither the name of the copyright holder nor the names of its contributors may be used to endorse or promote products derived from this software without specific prior written permission. THIS SOFTWARE IS PROVIDED BY THE COPYRIGHT HOLDERS AND CONTRIBUTORS "AS IS" AND ANY EXPRESS OR IMPLIED WARRANTIES, INCLUDING, BUT NOT LIMITED TO, THE IMPLIED WARRANTIES OF MERCHANTABILITY AND FITNESS FOR A PARTICULAR PURPOSE ARE DISCLAIMED. IN NO EVENT SHALL THE COPYRIGHT HOLDER OR CONTRIBUTORS BE LIABLE FOR ANY DIRECT, INDIRECT, INCIDENTAL, SPECIAL, EXEMPLARY, OR CONSEQUENTIAL DAMAGES (INCLUDING, BUT NOT LIMITED TO, PROCUREMENT OF SUBSTITUTE GOODS OR SERVICES; LOSS OF USE, DATA, OR PROFITS; OR BUSINESS INTERRUPTION) HOWEVER CAUSED AND ON ANY THEORY OF LIABILITY, WHETHER IN CONTRACT, STRICT LIABILITY, OR TORT (INCLUDING NEGLIGENCE OR OTHERWISE) ARISING IN ANY WAY OUT OF THE USE OF THIS SOFTWARE, EVEN IF ADVISED OF THE POSSIBILITY OF SUCH DAMAGE. Source: https://raw.githubusercontent.com/scikit-learn/scikit-learn/main/COPYING












暂无评论内容