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

PyTorch中嵌套交叉验证(Nested Cross Validation)的实现及代码验证

在PyTorch中实现嵌套交叉验证的正确方式

可以在PyTorch中实现嵌套交叉验证,它能更严谨地评估模型的泛化能力,尤其适合你这种4000张图像的多类别分类任务。你的现有代码思路方向是对的,但存在几个关键错误,会导致评估结果不可靠,下面逐一说明并给出修正方案。

你的代码存在的核心问题

  • 模型/优化器未重置:所有fold共用同一个模型和优化器实例,参数会在不同fold间累积更新,完全破坏了交叉验证的独立性。
  • 内层循环逻辑混乱:你把epoch循环放在了内层交叉验证外面,导致每个epoch都重新划分一次训练/验证集并从头训练,这不符合嵌套交叉验证的流程——内层交叉验证应该是为了找到当前外层fold下的最佳模型状态,而非每个epoch都重复交叉验证。
  • 标签混淆:验证集的loss被打印成"test loss",容易和外层的测试集loss混淆。
  • 数据集初始化错误:Dataset类的__init__中super().__init__(self)参数错误,VisionDataset的第一个参数应该是数据集根目录,这里模拟数据可以传空字符串或None。

正确的嵌套交叉验证实现流程

  1. 外层K折:将整个数据集划分为K份,每次取1份作为测试集,剩下的K-1份作为"训练+验证"数据集。
  2. 内层K折:在"训练+验证"数据集中再做K折交叉验证,选出当前外层fold下的最佳超参数(或最佳模型权重)。
  3. 重新训练:用"训练+验证"的全部数据,以内层交叉验证得到的最佳配置重新训练模型。
  4. 测试评估:用训练好的模型在外层测试集上评估性能,记录结果。
  5. 重复外层fold:完成所有外层fold后,取所有测试结果的平均值作为最终的模型泛化能力评估。

修正后的代码实现

import torch 
from torch.utils.data import DataLoader, SubsetRandomSampler
from sklearn.model_selection import KFold
from torchvision import datasets

# 模拟数据集参数
input_size = (256, 3, 224, 224)  # 修正原代码中尺寸不一致的问题
target_size = (256,)

class Dataset(datasets.VisionDataset):
    def __init__(self):   
        # 修正VisionDataset初始化参数,传入空字符串作为根目录
        super().__init__("")
        self.images = torch.rand(input_size).float()
        self.targets = torch.randint(0, 3, target_size)

    def __getitem__(self, index: int):
        return self.images[index], self.targets[index]

    def __len__(self) -> int:
        return len(self.images)

class BasicModel(torch.nn.Module):
    def __init__(self):
        super().__init__()
        self.conv = torch.nn.Conv2d(3, 16, kernel_size=(5,5))
        self.adp = torch.nn.AdaptiveAvgPool2d(1)
        self.linear = torch.nn.Linear(16, 3)
    
    def forward(self, x):
        x = self.conv(x)
        x = self.adp(x)
        x = x.view(x.size(0), -1)
        x = self.linear(x)
        return x

def train_model(model, train_loader, criterion, optimizer, num_epochs):
    """封装训练逻辑,方便复用"""
    model.train()
    for epoch in range(num_epochs):
        total_loss = 0.0
        for images, targets in train_loader:
            optimizer.zero_grad()
            outputs = model(images)
            loss = criterion(outputs, targets)
            loss.backward()
            optimizer.step()
            total_loss += loss.item() * images.size(0)
        avg_loss = total_loss / len(train_loader.sampler)
        print(f"Epoch {epoch+1}/{num_epochs}, Train Loss: {avg_loss:.4f}")
    return model

def evaluate_model(model, loader, criterion):
    """封装评估逻辑,返回平均loss和准确率"""
    model.eval()
    total_loss = 0.0
    correct = 0
    total = 0
    with torch.no_grad():
        for images, targets in loader:
            outputs = model(images)
            loss = criterion(outputs, targets)
            total_loss += loss.item() * images.size(0)
            _, predicted = torch.max(outputs.data, 1)
            total += targets.size(0)
            correct += (predicted == targets).sum().item()
    avg_loss = total_loss / total
    accuracy = correct / total
    return avg_loss, accuracy

# 初始化数据集
data = Dataset()
data_ids = list(range(len(data)))
k_fold = 5
num_epochs = 2

# 外层K折交叉验证(测试集划分)
outer_kfold = KFold(n_splits=k_fold, shuffle=True, random_state=42)
outer_test_results = []

for outer_fold, (remain_ids, test_ids) in enumerate(outer_kfold.split(data_ids), 1):
    print(f"\n===== Outer Fold {outer_fold}/{k_fold} =====")
    
    # 内层K折交叉验证(验证集划分,用于选择最佳模型)
    inner_kfold = KFold(n_splits=k_fold-1, shuffle=True, random_state=42)
    inner_best_acc = 0.0
    inner_best_model = None
    
    for inner_fold, (train_ids, val_ids) in enumerate(inner_kfold.split(remain_ids), 1):
        print(f"\n--- Inner Fold {inner_fold}/{k_fold-1} ---")
        
        # 每个内层fold都重新初始化模型、优化器、损失函数
        model = BasicModel()
        criterion = torch.nn.CrossEntropyLoss()
        optimizer = torch.optim.Adam(model.parameters(), lr=0.01)
        
        # 创建数据加载器
        train_sampler = SubsetRandomSampler(train_ids)
        train_loader = DataLoader(data, sampler=train_sampler, batch_size=2)
        
        val_sampler = SubsetRandomSampler(val_ids)
        val_loader = DataLoader(data, sampler=val_sampler, batch_size=2)
        
        # 训练模型
        model = train_model(model, train_loader, criterion, optimizer, num_epochs)
        
        # 验证模型
        val_loss, val_acc = evaluate_model(model, val_loader, criterion)
        print(f"Inner Fold {inner_fold} Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}")
        
        # 保存当前内层fold的最佳模型
        if val_acc > inner_best_acc:
            inner_best_acc = val_acc
            inner_best_model = model.state_dict()
    
    # 用剩余数据的全部样本重新训练最佳模型
    print(f"\n--- Retraining on full remain data with best inner model ---")
    final_model = BasicModel()
    final_model.load_state_dict(inner_best_model)
    final_optimizer = torch.optim.Adam(final_model.parameters(), lr=0.01)
    
    remain_sampler = SubsetRandomSampler(remain_ids)
    remain_loader = DataLoader(data, sampler=remain_sampler, batch_size=2)
    final_model = train_model(final_model, remain_loader, criterion, final_optimizer, num_epochs)
    
    # 在外层测试集上评估
    test_sampler = SubsetRandomSampler(test_ids)
    test_loader = DataLoader(data, sampler=test_sampler, batch_size=2)
    
    test_loss, test_acc = evaluate_model(final_model, test_loader, criterion)
    print(f"Outer Fold {outer_fold} Test Loss: {test_loss:.4f}, Test Acc: {test_acc:.4f}")
    outer_test_results.append(test_acc)

# 计算最终的平均测试准确率
avg_test_acc = sum(outer_test_results) / len(outer_test_results)
print(f"\n===== Final Average Test Accuracy: {avg_test_acc:.4f} =====")

关键改进点说明

  • 模型/优化器重置:每个内层fold都重新初始化模型和优化器,确保每个fold的训练独立。
  • 分离训练/评估逻辑:封装train_model和evaluate_model函数,让代码更清晰,避免重复冗余。
  • 正确的嵌套流程:内层交叉验证选最佳模型,再用全部训练数据重新训练,最后在测试集评估,符合嵌套交叉验证的严谨性要求。
  • 添加准确率评估:除了loss,增加了分类任务更关注的准确率指标,评估更全面。
  • 固定随机种子:设置random_state=42,让实验结果可复现。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.02 22:20:41