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

如何将StratifiedKFold交叉验证每折的索引转换为对应DataFrame

解决方法

核心原理是StratifiedKFold.split()返回的是整数位置索引,直接用Pandas的iloc属性按位置取原始DataFrame的对应行即可,无需提前将数据集转为Numpy数组。

推荐方案(直接复用原始DataFrame)

修改你的代码如下,跳过转Numpy数组的步骤,直接对原始DataFrame做索引:

import numpy as np
import pandas as pd
from sklearn.model_selection import StratifiedKFold

# 直接使用原始DataFrame/Series,不要转为Numpy数组
X = X_trainval  # X_trainval为原始特征DataFrame
y = y_trainval  # y_trainval为原始标签Series/ DataFrame
skf = StratifiedKFold(n_splits=4, random_state=None, shuffle=False)

for train_index, test_index in skf.split(X, y):
    print("TRAIN:", train_index, "TEST:", test_index)
    # 用iloc按位置索引,直接得到DataFrame格式结果
    X_traincv, X_testcv = X.iloc[train_index], X.iloc[test_index]
    y_traincv, y_testcv = y.iloc[train_index], y.iloc[test_index]

兼容方案(已转为Numpy数组时使用)

如果你已经提前将数据集转为了Numpy数组,可以手动把拆分后的数组转回DataFrame,提前保留原始的列名和索引即可:

import numpy as np
import pandas as pd
from sklearn.model_selection import StratifiedKFold

X = np.array(X_trainval)
y = np.array(y_trainval)
# 提前保留原始特征列名、行索引、标签名
feature_cols = X_trainval.columns.tolist()
original_index = X_trainval.index
label_name = y_trainval.name

skf = StratifiedKFold(n_splits=4, random_state=None, shuffle=False)

for train_index, test_index in skf.split(X, y):
    print("TRAIN:", train_index, "TEST:", test_index)
    # 将Numpy数组转回DataFrame/Series
    X_traincv = pd.DataFrame(X[train_index], columns=feature_cols, index=original_index[train_index])
    X_testcv = pd.DataFrame(X[test_index], columns=feature_cols, index=original_index[test_index])
    y_traincv = pd.Series(y[train_index], name=label_name, index=original_index[train_index])
    y_testcv = pd.Series(y[test_index], name=label_name, index=original_index[test_index])

注意事项

  • 如果你不需要保留原始的行索引,可以省略index参数,拆分后的DataFrame会自动生成从0开始的新索引
  • 若设置了shuffle=True,iloc仍会按照返回的位置索引正确取数,和是否打乱无关

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.04 16:15:00