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

已完成train_test_split划分的数据集如何正确实现k-fold交叉验证

K折交叉验证接入位置与实现方法

你第一步拆分出的20%测试集必须全程独立封存,不参与任何训练、验证、超参数调整环节,仅在所有K折流程结束、确定最终模型与超参数后,用来做一次最终泛化性能评估,否则会造成数据泄露,结果完全不可信。

K折交叉验证的作用范围是你拆分出的80% x_trainval/y_trainval 集合,直接替换你原来第二次固定切分10%验证集的逻辑即可。


具体执行流程

  • 保留原有数据归一化、测试集分层拆分的代码,这部分不需要改动
  • 删除原有固定切分x_train/x_val、创建固定train/val Dataset与DataLoader的代码
  • 引入分层K折拆分器(分类任务优先用分层拆分,保证每折类别分布和整体一致,避免验证结果波动),对x_trainval/y_trainval做K折切分,常用折数为5或10
  • 逐折迭代:
    • 用当前折的训练索引、验证索引,分别从x_trainval/y_trainval中取出对应折的训练数据、验证数据
    • 基于当前折的数据创建训练集、验证集,以及对应的DataLoader
    • 必须从头初始化全新的模型、优化器、学习率调度器,禁止复用前一折训练好的权重,否则验证结果失真
    • 按设定的epoch数完成当前折的训练与验证,记录当前折的验证集指标(准确率、F1等)
  • 所有折训练完成后,计算K折验证指标的均值与标准差,评估当前超参数(学习率、batch size、模型结构、正则强度等)的真实效果,以此为依据调整超参数
  • 超参数确定后,可以用全部x_trainval数据重新训练一个最终模型,或者选择K折中验证表现最好的模型,放到封存的测试集上跑一次最终评估,得到模型真实泛化能力

改造后代码示例

首先补充需要导入的依赖:

from sklearn.model_selection import StratifiedKFold
import torch
import numpy as np
# 你原有其他导入(preprocessing、Dataset、DataLoader、模型定义等)保持不变

原有预处理与测试集拆分部分完全保留:

x_norm = preprocessing.normalize(x, axis=0)
x = x_norm

# 这部分测试集拆分保留,测试集全程封存不参与K折
x_trainval, x_test, y_trainval, y_test = train_test_split(
    x, y, test_size=0.2, random_state=0, stratify=df["label"]
)

# 测试集Dataset和DataLoader可以提前创建,最后测试才用
test_dataset = classifierdataset(
    torch.from_numpy(x_test).float(), 
    torch.from_numpy(y_test).long()
)
test_loader = DataLoader(dataset=test_dataset, batch_size=1)

替换原有固定train/val拆分的部分为K折逻辑:

# 超参数定义保持不变
EPOCHS = 10
BATCH_SIZE = 16
LEARNING_RATE = 0.0007
N_SPLITS = 5  # K折的折数,常用5或10

# 初始化分层K折拆分器
skf = StratifiedKFold(n_splits=N_SPLITS, shuffle=True, random_state=0)
all_fold_val_acc = []  # 记录每折的验证准确率

for fold, (train_idx, val_idx) in enumerate(skf.split(x_trainval, y_trainval)):
    print(f"===== 开始训练第 {fold+1}/{N_SPLITS} 折 =====")
    # 取出当前折的训练、验证数据
    x_train_fold = x_trainval[train_idx]
    y_train_fold = y_trainval[train_idx]
    x_val_fold = x_trainval[val_idx]
    y_val_fold = y_trainval[val_idx]

    # 创建当前折的Dataset与DataLoader
    train_fold_dataset = classifierdataset(
        torch.from_numpy(x_train_fold).float(),
        torch.from_numpy(y_train_fold).long()
    )
    val_fold_dataset = classifierdataset(
        torch.from_numpy(x_val_fold).float(),
        torch.from_numpy(y_val_fold).long()
    )
    train_loader = DataLoader(dataset=train_fold_dataset, batch_size=BATCH_SIZE, shuffle=True)
    val_loader = DataLoader(dataset=val_fold_dataset, batch_size=1)

    # !!!每折必须重新初始化模型、优化器,替换成你自己的模型类即可
    model = YourClassifierModel()  # 这里换成你自己定义的PyTorch模型类
    criterion = torch.nn.CrossEntropyLoss()
    optimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE)

    # 当前折的训练循环
    for epoch in range(EPOCHS):
        model.train()
        train_loss = 0
        for batch_x, batch_y in train_loader:
            optimizer.zero_grad()
            outputs = model(batch_x)
            loss = criterion(outputs, batch_y)
            loss.backward()
            optimizer.step()
            train_loss += loss.item()
        
        # 当前epoch验证
        model.eval()
        val_correct = 0
        val_total = 0
        with torch.no_grad():
            for batch_x, batch_y in val_loader:
                outputs = model(batch_x)
                _, pred = torch.max(outputs.data, 1)
                val_total += batch_y.size(0)
                val_correct += (pred == batch_y).sum().item()
        val_acc = val_correct / val_total
        print(f"Fold {fold+1} Epoch {epoch+1}/{EPOCHS}, Train Loss: {train_loss/len(train_loader):.4f}, Val Acc: {val_acc:.4f}")
    
    # 记录当前折最终验证准确率
    all_fold_val_acc.append(val_acc)
    print(f"第 {fold+1} 折训练完成,验证准确率: {val_acc:.4f}\n")

# 输出K折整体验证结果
print(f"===== 所有折训练完成 =====")
print(f"平均验证准确率: {np.mean(all_fold_val_acc):.4f} ± {np.std(all_fold_val_acc):.4f}")

注意事项

  • K折本身不会直接解决过拟合,它的作用是提供比单次固定切分更稳定、更可靠的泛化性能评估结果,帮你更准确地判断正则化、早停、模型结构调整等防过拟合手段的实际效果
  • 训练时训练集DataLoader建议开启shuffle=True,验证、测试集不需要打乱
  • 如果要使用早停,早停的判断依据是当前折的验证集指标,不要跨折复用早停状态
  • 最终测试集只能在所有调参、K折验证完成后使用一次,不要根据测试集结果反过来调整模型,否则测试集失去意义

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.28 03:20:03