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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.30 20:21:07