PyTorch中如何结合使用K折交叉验证与Weighted Random Sampler
PyTorch中结合加权随机采样与K折交叉验证的实现方案
核心原则:WeightedRandomSampler必须针对每折的训练子集单独初始化,绝对不能作用于验证集,也不能提前基于全数据集计算全局采样权重。
实现逻辑
K折交叉验证会把全数据集拆成K份互斥的子集,每一轮用其中K-1份做训练、1份做验证。你之前的加权采样逻辑是针对全数据集写的,直接套K折会有两个问题:一是采样范围覆盖了验证集,打乱验证集的真实类别分布,导致验证指标失效;二是采样权重和每折实际参与训练的样本不匹配,达不到类别平衡的效果。
推荐搭配StratifiedKFold(分层K折)使用,它会在拆分时保证每折的类别占比和全数据集一致,比普通K折更适配类别不平衡场景,能大幅降低折间指标波动。
具体实现步骤
- 提前提取全数据集的标签列表,供分层K折拆分使用
- 遍历每一轮折拆分,分别拿到当前折的训练集索引、验证集索引,生成对应的训练/验证子集
- 仅针对当前折的训练子集计算样本权重,初始化专属的WeightedRandomSampler
- 构造DataLoader时,训练集传入上述sampler,验证集不设置sampler、保持shuffle=False,保证验证集分布和真实场景一致
- 每折训练前必须重新初始化模型、优化器,不能复用之前折的训练参数
参考代码
import torch import numpy as np from torch.utils.data import DataLoader, WeightedRandomSampler, Subset from sklearn.model_selection import StratifiedKFold # 基础超参数,和你原有配置保持一致即可 K = 5 CLASS_WEIGHTS = [5, 1, 1] BATCH_SIZE = 32 NUM_WORKERS = 2 EPOCHS = 20 # 提前提取全量标签,供分层K折拆分使用 # ds为你之前加载的完整三分类数据集 all_labels = np.array([ds[i]["targets"] for i in range(len(ds))]) kfold = StratifiedKFold(n_splits=K, shuffle=True, random_state=42) fold_metrics = [] for fold_id, (train_idx, val_idx) in enumerate(kfold.split(np.arange(len(ds)), all_labels)): print(f"=== 训练第 {fold_id+1}/{K} 折 ===") # 拆分当前折的训练、验证子集 train_subset = Subset(ds, train_idx) val_subset = Subset(ds, val_idx) # 仅为当前折的训练样本计算权重 train_weights = [CLASS_WEIGHTS[ds[i]["targets"]] for i in train_idx] # 初始化当前训练集专属的加权采样器 train_sampler = WeightedRandomSampler( weights=train_weights, num_samples=len(train_weights), replacement=True ) # 构造DataLoader train_loader = DataLoader( train_subset, batch_size=BATCH_SIZE, sampler=train_sampler, num_workers=NUM_WORKERS ) # 验证集绝对不能加sampler,保持默认顺序采样 val_loader = DataLoader( val_subset, batch_size=BATCH_SIZE, shuffle=False, num_workers=NUM_WORKERS ) # 每折必须重新初始化模型、优化器、损失函数 model = YourTriClassModel() # 替换为你自己的三分类模型结构 optimizer = torch.optim.Adam(model.parameters(), lr=1e-3, weight_decay=1e-4) # 加weight_decay可进一步缓解过拟合 criterion = torch.nn.CrossEntropyLoss(weight=torch.tensor(CLASS_WEIGHTS, dtype=torch.float32)) # 正常执行当前折的训练、验证流程即可 best_f1 = 0 for epoch in range(EPOCHS): model.train() for batch in train_loader: # 写入你原有的训练逻辑:前向传播、损失计算、反向传播、参数更新 x, y = batch["feature"], batch["targets"] pred = model(x) loss = criterion(pred, y) optimizer.zero_grad() loss.backward() optimizer.step() # 验证阶段 model.eval() # 写入你原有的验证逻辑,计算准确率、F1等指标 # 可以搭配早停机制,验证集指标长期不提升就提前终止当前折训练,进一步缓解过拟合 fold_metrics.append(best_f1) print(f"{K}折交叉验证平均F1分数:{np.mean(fold_metrics):.4f}")
常见避坑点
- 不要在验证集使用加权采样:验证集的作用是模拟真实无偏的数据分布,加权采样会人为改变类别占比,得到的验证指标完全不具备参考性,无法判断过拟合程度
- 不要基于全数据集提前计算采样权重:每折的训练子集是动态变化的,全量权重和当前训练子集不匹配,会导致采样逻辑失效
- 如果过拟合问题仍然明显,可以在交叉验证基础上搭配Dropout、权重衰减、早停、数据增强等正则化手段,泛化能力提升会更明显
内容的提问来源于stack exchange,提问作者Asma Bouzidi
相关产品推荐
相关产品推荐

