sklearn.model_selection.KFold

KFoldsklearn.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 结合优化超参数
Logo

开源鸿蒙跨平台开发社区汇聚开发者与厂商,共建“一次开发,多端部署”的开源生态,致力于降低跨端开发门槛,推动万物智联创新。

更多推荐