You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何验证基于外部分组特征的StratifiedKFold交叉验证行为?

问题解答

代码正确性判断

你的代码思路是对的——通过将groups传入StratifiedKFold.split()的第二个参数(替代原本的分类目标变量),确实能生成按groups分层的训练/测试拆分,保证每个fold中groups的类别分布与整体数据一致。但有两点需要注意:

  1. StratifiedKFold本身是为分类任务设计的,它会把传入的groups当作离散分类标签处理,因此你的groups必须是离散型变量(如果是连续型需要先离散化)。
  2. 代码中skf.split(training, groups)里的training应该和cross_val_score中的X保持一致,否则会出现索引不匹配的问题,建议统一改为X。

验证方法

可以通过以下两种方式验证交叉验证是否符合预期:

1. 手动检查每个fold的group分布

遍历拆分后的索引,统计每个训练集/测试集的groups类别占比,与整体数据的分布对比:

import pandas as pd

# 计算整体数据的group分布
overall_group_dist = groups.value_counts(normalize=True).sort_index()
print("整体group分布:\n", overall_group_dist)
print("-" * 60)

# 遍历每个fold,检查分布
skf = StratifiedKFold(n_splits=5, shuffle=True, random_state=57)
for fold_num, (train_idx, test_idx) in enumerate(skf.split(X, groups), 1):
    # 获取当前fold的训练/测试集groups
    train_groups = groups.iloc[train_idx]
    test_groups = groups.iloc[test_idx]
    
    # 计算分布
    train_dist = train_groups.value_counts(normalize=True).sort_index()
    test_dist = test_groups.value_counts(normalize=True).sort_index()
    
    print(f"Fold {fold_num} 训练集group分布:\n", train_dist)
    print(f"Fold {fold_num} 测试集group分布:\n", test_dist)
    print("-" * 60)

如果每个fold的训练/测试集groups占比与整体分布偏差很小(比如小于5%),说明分层效果符合预期。

2. 对比模型得分的稳定性

按groups分层的交叉验证,各fold的模型得分(比如回归任务的R²、MAE)方差应该远小于随机拆分的得分方差。你可以同时用KFold(无分层)做一次交叉验证,对比两者的得分标准差:

from sklearn.model_selection import KFold

# 分层拆分的得分
stratified_scores = cross_val_score(regr, X, y, cv=skf.split(X, groups))
# 随机拆分的得分
random_scores = cross_val_score(regr, X, y, cv=KFold(n_splits=5, shuffle=True, random_state=57))

print("分层交叉验证得分标准差:", stratified_scores.std())
print("随机交叉验证得分标准差:", random_scores.std())

如果分层后的得分标准差更小,说明分层有效降低了数据分布差异带来的得分波动。

更规范的替代方案

如果你不想“误用”StratifiedKFold,可以使用StratifiedGroupKFold(需sklearn版本≥0.24),它专门支持按外部groups分层,同时可适配回归任务:

from sklearn.model_selection import StratifiedGroupKFold

sgkf = StratifiedGroupKFold(n_splits=5, shuffle=True, random_state=57)
scores = cross_val_score(regr, X, y, groups=groups, cv=sgkf)

内容的提问来源于stack exchange,提问作者hexolitemax

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.10 06:48:36