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

PyTorch中多数据集差异化子采样(按Epoch重采样)拼接实现问询

针对动态子采样需求的两种实现方案

方案1:基于IterableDataset的自动Epoch重采样

直接实现一个包装类继承IterableDataset,利用其__iter__方法在每个Epoch被调用时自动重新采样,完美契合你的需求:

import torch
from torch.utils.data import IterableDataset, ConcatDataset, Dataset

class SubsampledIterableDataset(IterableDataset):
    def __init__(self, base_dataset, target_sample_count):
        self.base_dataset = base_dataset
        self.target_size = target_sample_count  # 与小数据集规模对齐

    def __iter__(self):
        # 每个Epoch随机生成采样索引
        sample_indices = torch.randperm(len(self.base_dataset))[:self.target_size].tolist()
        for idx in sample_indices:
            yield self.base_dataset[idx]

# 实际使用
small_ds = YourSmallDataset()
large_ds = YourLargeDataset()
# 将大数据集包装为动态采样版本
subsampled_large = SubsampledIterableDataset(large_ds, len(small_ds))
# 拼接两个数据集
combined_ds = ConcatDataset([subsampled_large, small_ds])
# 传入DataLoader
train_loader = torch.utils.data.DataLoader(combined_ds, batch_size=32)

每次DataLoader启动一个新Epoch时,都会触发__iter__方法重新采样,无需额外手动操作。

方案2:Map-style Dataset的Epoch级采样更新

Map-style Dataset没有内置的Epoch钩子,但可以通过自定义包装类+手动触发更新的方式实现:

import torch
from torch.utils.data import Dataset, ConcatDataset

class EpochSubsampledMapDataset(Dataset):
    def __init__(self, base_dataset, target_sample_count):
        self.base_dataset = base_dataset
        self.target_size = target_sample_count
        self.current_sample_indices = []

    def refresh_samples(self):
        # 每个Epoch开始前调用,更新采样索引
        self.current_sample_indices = torch.randperm(len(self.base_dataset))[:self.target_size].tolist()

    def __len__(self):
        return self.target_size

    def __getitem__(self, idx):
        return self.base_dataset[self.current_sample_indices[idx]]

# 实际使用
small_ds = YourSmallDataset()
large_ds = YourLargeDataset()
subsampled_large = EpochSubsampledMapDataset(large_ds, len(small_ds))
combined_ds = ConcatDataset([subsampled_large, small_ds])
train_loader = torch.utils.data.DataLoader(combined_ds, batch_size=32)

# 训练循环中手动触发更新
for epoch in range(10):
    subsampled_large.refresh_samples()
    for batch in train_loader:
        # 训练逻辑
        pass

这种方案保留了Map-style Dataset的随机访问特性,适合需要依赖该特性的场景,代价是每个Epoch开始前需要手动调用刷新方法。

补充说明

  • 如果需要采样的可复现性,可以在采样时固定随机种子(比如torch.manual_seed(epoch + 固定值))
  • 两种方案都能满足你"拼接后传入DataLoader"的要求,无需修改DataLoader的sampler参数

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.18 19:50:44