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

使用GroupKFold实现交叉验证时出现Key Error问题求助

问题说明

我有一个包含label、embeddings(特征列)、chr三列的数据框,想要按染色体分组做10折交叉验证,确保同一染色体(比如chr1)的所有行要么全在训练集,要么全在测试集,不拆分到两个集合里。我觉得代码写得没问题,但一直触发Key Error。以下是我的代码:

import numpy as np
from sklearn.model_selection import GroupKFold

X = np.array([np.array(x) for x in mini_df['embeddings']])
y = mini_df['label']
groups = mini_df['chromosome']
group_kfold = GroupKFold(n_splits=10)

# Initialize figure for plotting
plt.figure(figsize=(10, 6))

# Perform cross-validation and plot ROC curves for each fold
for i, (train_idx, val_idx) in enumerate(group_kfold.split(X, y, groups)):
    X_train_fold, X_val_fold = X[train_idx], X[val_idx]
    y_train_fold, y_val_fold = y[train_idx], y[val_idx]
    
    # Initialize classifier
    rf_classifier = RandomForestClassifier(n_estimators=n_trees, random_state=42, max_depth=max_depth, n_jobs=-1)
    
    # Train the classifier on this fold
    rf_classifier.fit(X_train_fold, y_train_fold)
    
    # Make predictions on the validation set
    y_pred_proba = rf_classifier.predict_proba(X_val_fold)[:, 1]
    
    # Calculate ROC curve
    fpr, tpr, _ = roc_curve(y_val_fold, y_pred_proba)
    
    # Calculate AUC
    roc_auc = auc(fpr, tpr)
    
    # Plot ROC curve for this fold
    plt.plot(fpr, tpr, lw=1, alpha=0.7, label=f'ROC Fold {i+1} (AUC = {roc_auc:.2f})')

# Plot ROC for random classifier
plt.plot([0, 1], [0, 1], linestyle='--', lw=2, color='r', label='Random', alpha=0.8)

# Add labels and legend
plt.xlabel('False Positive Rate')
plt.ylabel('True Positive Rate')
plt.title('ROC Curves for Random Forest Classifier')
plt.legend(loc='lower right')
plt.show()
错误原因及修正

触发Key Error的核心原因是列名不匹配:你数据框里的染色体列是chr,但代码里写的是mini_df['chromosome'],引用了不存在的列名,直接导致报错。

另外还要注意两个细节:

  1. 代码里用到的n_trees和max_depth必须提前定义,否则会触发NameError
  2. 缺少matplotlib.pyplot、RandomForestClassifier、roc_curve和auc的导入语句,运行时会报错
修正后的完整代码
import numpy as np
import matplotlib.pyplot as plt
from sklearn.model_selection import GroupKFold
from sklearn.ensemble import RandomForestClassifier
from sklearn.metrics import roc_curve, auc

# 提前定义模型参数,可根据需求调整
n_trees = 100
max_depth = 10

# 修正列名:使用数据框实际的'chr'列
X = np.array([np.array(x) for x in mini_df['embeddings']])
y = mini_df['label']
groups = mini_df['chr']  # 关键修正点
group_kfold = GroupKFold(n_splits=10)

plt.figure(figsize=(10, 6))

for i, (train_idx, val_idx) in enumerate(group_kfold.split(X, y, groups)):
    X_train_fold, X_val_fold = X[train_idx], X[val_idx]
    y_train_fold, y_val_fold = y[train_idx], y[val_idx]
    
    rf_classifier = RandomForestClassifier(n_estimators=n_trees, random_state=42, max_depth=max_depth, n_jobs=-1)
    rf_classifier.fit(X_train_fold, y_train_fold)
    
    y_pred_proba = rf_classifier.predict_proba(X_val_fold)[:, 1]
    fpr, tpr, _ = roc_curve(y_val_fold, y_pred_proba)
    roc_auc = auc(fpr, tpr)
    
    plt.plot(fpr, tpr, lw=1, alpha=0.7, label=f'ROC Fold {i+1} (AUC = {roc_auc:.2f})')

plt.plot([0, 1], [0, 1], linestyle='--', lw=2, color='r', label='Random', alpha=0.8)
plt.xlabel('False Positive Rate')
plt.ylabel('True Positive Rate')
plt.title('ROC Curves for Random Forest Classifier')
plt.legend(loc='lower right')
plt.show()
额外验证步骤
  • 运行print(mini_df.columns)确认数据框的列名,确保chr存在且拼写正确
  • 如果embeddings列的元素本身就是numpy数组,可以用X = np.stack(mini_df['embeddings'].values)替代列表推导,更高效

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 16:44:51