分层K折交叉验证训练/验证循环中索引越界错误排查
在PyTorch实现分层K折交叉验证时,首次折叠可正常运行,但进入第二次折叠时,在代码行y_train_fold, y_val_fold = labels[train_index], labels[val_index]处触发索引越界错误。输入数据形状为(864,19,500),对应二分类标签长度为864。尝试过用SMOTE平衡类别分布、检查标签的长度与形状,但问题仍未解决;当设置num_folds=4时,第二次折叠时标签形状变为24,导致出现“索引24超出维度0的大小24”的越界错误。
相关代码
def __init__(self, model, num_folds=5, batch_size=32, epochs=10, lr=0.001, betas=(0.9, 0.999), eps=1e-8): """ Initialize the ModelTrainer with specified parameters. Args: model (torch.nn.Module): The PyTorch model to be trained and validated. num_folds (int): The number of folds for stratified k-fold cross-validation. Default is 5. batch_size (int): Batch size for training and validation. Default is 32. epochs (int): Number of epochs for training. Default is 10. lr (float): Learning rate for the optimizer. Default is 0.001. betas (tuple): Coefficients used for computing running averages of gradient and its square. Default is (0.9, 0.999). eps (float): Term added to the denominator to improve numerical stability in the optimizer. Default is 1e-8. """ self.model = model self.num_folds = num_folds self.batch_size = batch_size self.epochs = epochs self.lr = lr self.betas = betas self.eps = eps def train_and_validate(self, data, labels): """ Train and validate the model using stratified k-fold cross-validation. Args: data (torch.Tensor): The input data for training and validation. labels (torch.Tensor): The labels corresponding to the input data. Returns: auc_scores (list): A list of AUC scores for each fold. sensitivities (list): Sensitivities at 100% specificity for each fold. mean_auc (float): Mean AUC score across all folds. mean_sensitivity (float): Mean sensitivity at 100% specificity across all folds. """ # Initialize StratifiedKFold skf = StratifiedKFold(n_splits=self.num_folds, shuffle=True, random_state=42) # Lists to store AUC scores and sensitivities for each fold auc_scores = [] sensitivities = [] # Lists to store ROC curve data for each fold tprs = [] mean_fpr = np.linspace(0, 1, 100) # Iterate over folds for fold, (train_index, val_index) in enumerate(skf.split(data, labels)): # Get the data for this fold X_train_fold, X_val_fold = data[train_index], data[val_index] y_train_fold, y_val_fold = labels[train_index], labels[val_index] # Create PyTorch datasets and data loaders for this fold train_dataset = TensorDataset(X_train_fold, y_train_fold) val_dataset = TensorDataset(X_val_fold, y_val_fold) train_loader = DataLoader(train_dataset, batch_size=self.batch_size, shuffle=True) val_loader = DataLoader(val_dataset, batch_size=self.batch_size) # Define loss function and optimizer criterion = nn.BCELoss() optimizer = torch.optim.Adam(self.model.parameters(), lr=self.lr, betas=self.betas, eps=self.eps) # Training loop for this fold for epoch in range(self.epochs): self.model.train() for inputs, labels in train_loader: optimizer.zero_grad() outputs = self.model(inputs) loss = criterion(outputs, labels) loss.backward() optimizer.step() # Validation loop for this fold self.model.eval() val_outputs_list = [] val_labels_list = [] for inputs, labels in val_loader: with torch.no_grad(): val_outputs = self.model(inputs) val_outputs_list.append(val_outputs.numpy()) val_labels_list.append(labels.numpy()) val_outputs_np = np.concatenate(val_outputs_list) val_labels_np = np.concatenate(val_labels_list) # Calculate ROC curve for this fold fpr, tpr, thresholds = roc_curve(val_labels_np, val_outputs_np) roc_auc = auc(fpr, tpr) auc_scores.append(roc_auc) # Calculate sensitivity at 100% specificity sensitivity = np.interp(1e-3, fpr, tpr) sensitivities.append(sensitivity) # Store ROC curve data for this fold tprs.append(np.interp(mean_fpr, fpr, tpr)) tprs[-1][0] = 0.0 # Plot ROC curve for this fold plt.plot(fpr, tpr, lw=1, alpha=0.3) # Plot settings plt.xlim([-0.05, 1.05]) plt.ylim([-0.05, 1.05]) plt.xlabel('False Positive Rate') plt.ylabel('True Positive Rate') plt.title('Receiver Operating Characteristic (ROC) Curve') plt.show() # Calculate mean ROC curve mean_tpr = np.mean(tprs, axis=0) mean_auc = auc(mean_fpr, mean_tpr) # Plot mean ROC curve with thicker line plt.plot(mean_fpr, mean_tpr, color='black', lw=2, linestyle='--', label=f'Mean ROC (AUC = {mean_auc:.2f})') plt.legend(loc='lower right') plt.show() # Calculate mean sensitivity at 100% specificity mean_sensitivity = np.mean(sensitivities) return auc_scores, sensitivities, mean_auc, mean_sensitivity
问题根源与解决办法
核心问题:变量名冲突覆盖全局标签
训练和验证循环中,使用了labels作为局部变量(for inputs, labels in train_loader),这会覆盖函数参数传入的全局labels张量。第一次折叠结束后,全局的labels已经被替换为最后一个batch的标签张量(形状为(batch_size,),比如用户遇到的24),导致第二次折叠时用这个小张量去索引train_index/val_index(长度远大于24),直接触发索引越界。
修复步骤
重命名局部变量:将训练/验证循环中的
labels改为其他名称,比如batch_labels:# 训练循环修改 for inputs, batch_labels in train_loader: optimizer.zero_grad() outputs = self.model(inputs) loss = criterion(outputs, batch_labels) loss.backward() optimizer.step() # 验证循环修改 for inputs, batch_labels in val_loader: with torch.no_grad(): val_outputs = self.model(inputs) val_outputs_list.append(val_outputs.numpy()) val_labels_list.append(batch_labels.numpy())验证标签维度:确保传入的
labels是一维张量(形状(864,)),而非二维的(864,1)。如果是二维,可通过labels = labels.squeeze(dim=1)转换,因为StratifiedKFold要求标签为一维。可选:每次折叠重新初始化模型:如果希望各折叠完全独立,避免前一次折叠的模型参数影响后续折叠,可在每次折叠开始前重新初始化模型(比如在for循环内创建模型实例,而非传入已初始化的模型)。
内容的提问来源于stack exchange,提问作者Manu Jack Pel

