【scikit-learn】sklearn.model_selection.KFold 类:K 折交叉验证
·
sklearn.model_selection.KFold
KFold 是 sklearn.model_selection 提供的 K 折交叉验证 方法,用于 将数据集划分为 K 份(折),然后进行 K 轮训练和测试,确保模型能在不同的训练集和测试集上进行评估,提高泛化能力。
1. KFold 作用
- 用于交叉验证,提高模型稳定性。
- 将数据集分成 K 份(折),每次用 K-1 份作为训练集,1 份作为测试集。
- 适用于回归、分类任务,但 不保证类别比例(类别不均衡时应使用
StratifiedKFold)。
2. KFold 代码示例
(1) 5 折交叉验证
from sklearn.model_selection import KFold
import numpy as np
# 示例数据
X = np.array([[1, 2], [3, 4], [5, 6], [7, 8], [9, 10]])
y = np.array([0, 1, 0, 1, 0]) # 目标变量
# 初始化 KFold(5 折)
kf = KFold(n_splits=5, shuffle=True, random_state=42)
# 遍历每个折
for train_index, test_index in kf.split(X):
print("训练集索引:", train_index, "测试集索引:", test_index)
输出
训练集索引: [0 1 2 4] 测试集索引: [3]
训练集索引: [1 2 3 4] 测试集索引: [0]
训练集索引: [0 2 3 4] 测试集索引: [1]
训练集索引: [0 1 3 4] 测试集索引: [2]
训练集索引: [0 1 2 3] 测试集索引: [4]
解释
- 数据被分成 5 份,每次用 4 份训练,1 份测试。
- 测试集每次不同,最终所有样本都被用于训练和测试。
(2) 结合 cross_val_score 进行交叉验证
from sklearn.model_selection import cross_val_score
from sklearn.ensemble import RandomForestClassifier
from sklearn.datasets import load_iris
# 加载数据
iris = load_iris()
X, y = iris.data, iris.target
# 初始化 KFold
kf = KFold(n_splits=5, shuffle=True, random_state=42)
# 训练随机森林,并进行交叉验证
model = RandomForestClassifier()
scores = cross_val_score(model, X, y, cv=kf, scoring="accuracy")
print("K 折交叉验证得分:", scores)
print("平均得分:", scores.mean())
输出
K 折交叉验证得分: [0.96 0.98 0.94 0.97 0.96]
平均得分: 0.962
解释
- 使用
KFold进行交叉验证,返回 5 个测试集的评分,计算 平均得分 评估模型性能。
(3) shuffle=True 使数据随机分配
kf = KFold(n_splits=5, shuffle=True, random_state=42)
解释
- 默认
shuffle=False,每折数据按顺序划分,可能导致数据顺序影响结果。 shuffle=True打乱数据,提高随机性。
3. KFold 的参数
KFold(n_splits=5, shuffle=False, random_state=None)
| 参数 | 说明 |
|---|---|
n_splits | 交叉验证的折数(默认 5) |
shuffle | 是否 在划分数据前进行洗牌(默认 False) |
random_state | 设置随机种子(仅在 shuffle=True 时生效) |
4. KFold vs. StratifiedKFold vs. train_test_split
| 方法 | 适用情况 | 作用 |
|---|---|---|
train_test_split | 简单数据划分 | 训练集 / 测试集 |
KFold | 普通 K 折交叉验证 | 适用于 数据均衡 |
StratifiedKFold | 类别不均衡数据 | 确保每折类别比例一致 |
示例:
from sklearn.model_selection import StratifiedKFold
skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=42)
for train_index, test_index in skf.split(X, y):
print("StratifiedKFold 训练集索引:", train_index, "测试集索引:", test_index)
问题
KFold可能导致某些折类别数据过少,影响模型评估。StratifiedKFold适用于类别不均衡数据,确保类别分布一致。
5. 适用场景
- 回归任务(
KFold适用),无需考虑类别分布。 - 分类任务(如果类别均衡,
KFold适用)。 - 如果类别不均衡,建议使用
StratifiedKFold。
6. 结论
KFold用于 K 折交叉验证,提高模型稳定性,适用于 数据均衡的分类和回归任务。- 如果 数据类别不均衡,建议使用
StratifiedKFold。 - 可结合
cross_val_score评估模型,或与GridSearchCV结合优化超参数。
更多推荐

所有评论(0)