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

PyTorch中ConcatDataset如何实现不同数据集的非均匀采样?

实现方案

PyTorch 中数据集的采样逻辑由Sampler组件控制,你不需要修改ConcatDataset的定义,只需要自定义符合采样规则的采样器传入 DataLoader 即可,具体实现如下:


1. 严格固定循环规律的实现

如果你需要严格按照「CustomDataset1、CustomDataset1、CustomDataset2」的顺序循环采样,用自定义采样器实现:

import torch
from torch.utils.data import Sampler, ConcatDataset
from typing import List

class FixedRatioCycleSampler(Sampler[int]):
    def __init__(self, concat_dataset: ConcatDataset, ratios: List[int], shuffle_subset: bool = True):
        """
        按指定比例循环从ConcatDataset的子数据集采样
        :param concat_dataset: 拼接后的ConcatDataset实例
        :param ratios: 各子数据集的采样比例,你的需求对应传[2, 1]
        :param shuffle_subset: 子数据集内部是否打乱采样顺序
        """
        self.sub_datasets = concat_dataset.datasets
        self.sub_cum_sizes = concat_dataset.cumulative_sizes
        self.ratios = ratios
        self.shuffle_subset = shuffle_subset
        # 各子数据集在ConcatDataset中的全局索引偏移量
        self.offset = [0] + self.sub_cum_sizes[:-1].tolist()
        # 计算单epoch总采样数,这里取最长子数据集对齐,你也可以自定义为固定数值
        max_sub_len = max([len(ds) for ds in self.sub_datasets])
        self.total_samples = sum(ratios) * (max_sub_len // min(ratios) + 1)

    def __iter__(self):
        # 生成每个子数据集的采样索引池,长度不够时自动循环采样
        sub_index_pools = []
        for ds_idx in range(len(self.sub_datasets)):
            # 生成子数据集内部索引
            inner_indices = torch.arange(len(self.sub_datasets[ds_idx]))
            if self.shuffle_subset:
                inner_indices = inner_indices[torch.randperm(len(inner_indices))]
            # 循环填充到足够长度
            repeat = (self.total_samples // self.ratios[ds_idx]) + 1
            inner_indices = inner_indices.repeat(repeat)[:self.total_samples // self.ratios[ds_idx]]
            # 加上偏移量得到全局索引
            inner_indices += self.offset[ds_idx]
            sub_index_pools.append(inner_indices.tolist())
        
        # 按固定比例拼接索引,实现[ds1, ds1, ds2]循环
        final_indices = []
        ptrs = [0] * len(self.sub_datasets)
        while len(final_indices) < self.total_samples:
            for ds_idx in range(len(self.sub_datasets)):
                for _ in range(self.ratios[ds_idx]):
                    if len(final_indices) >= self.total_samples:
                        break
                    final_indices.append(sub_index_pools[ds_idx][ptrs[ds_idx]])
                    ptrs[ds_idx] += 1
                if len(final_indices) >= self.total_samples:
                    break
        return iter(final_indices)

    def __len__(self):
        return self.total_samples

2. 自定义采样器的使用方法

创建 DataLoader 时传入自定义采样器即可:

# 原有数据集拼接逻辑
dataset = ConcatDataset([CustomDataset1, CustomDataset2])
# 初始化采样器,设置2:1的采样比例
sampler = FixedRatioCycleSampler(concat_dataset=dataset, ratios=[2, 1], shuffle_subset=True)
# 传入DataLoader,注意设置sampler后不要重复设置shuffle参数
# batch_size设为3的话,每个batch刚好对应[ds1, ds1, ds2]的采样规律
dataloader = torch.utils.data.DataLoader(dataset, batch_size=3, sampler=sampler)

可选简化方案(仅保证整体比例,不要求固定循环顺序)

如果你不需要严格的固定循环顺序,只需要整体采样比例趋近于2:1,可以直接用PyTorch内置的WeightedRandomSampler,实现更简单:

dataset = ConcatDataset([CustomDataset1, CustomDataset2])
len1 = len(CustomDataset1)
len2 = len(CustomDataset2)
# 给所有样本设置权重:CustomDataset1每个样本权重为2,CustomDataset2每个为1
weights = [2.0] * len1 + [1.0] * len2
# num_samples可以自定义为你需要的单epoch采样总数
sampler = torch.utils.data.WeightedRandomSampler(weights, num_samples=len(weights), replacement=True)
dataloader = torch.utils.data.DataLoader(dataset, batch_size=3, sampler=sampler)

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.26 23:57:04