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

如何在PyTorch中使用DataLoader实现有放回采样并自定义迭代次数

问题解答

可以通过DataLoader实现该需求

PyTorch官方提供的DataLoader完全支持该需求,你可以通过搭配内置的WeightedRandomSampler实现有放回采样+自定义总迭代次数的效果,不需要自己手动写采样逻辑,还能复用DataLoader的多线程加载、自动批处理等特性。
实现步骤如下:

  1. 核心逻辑是用WeightedRandomSampler的num_samples参数控制总采样样本数,数值设为总迭代次数 × 批量大小即可,同时开启replacement=True启用有放回采样。
  2. 代码示例:
import torch
from torch.utils.data import DataLoader, WeightedRandomSampler, Dataset

# 你的自定义数据集
class MyDataset(Dataset):
    def __init__(self, N):
        self.N = N
        self.x = torch.rand(self.N, 10)
        self.y = torch.randint(0, 3, (self.N,))

    def __len__(self):
        return self.N

    def __getitem__(self, idx):
        return self.x[idx], self.y[idx]

# 实例化数据集
N = 100
dataset = MyDataset(N)

# 自定义参数
batch_size = 3 # 批量大小m
total_iters = 20 # 可传入的总迭代次数,不需要和N/m绑定

# 定义有放回采样器
sampler = WeightedRandomSampler(
    weights=torch.ones(len(dataset)), # 所有样本采样权重相等,即均匀采样
    num_samples=total_iters * batch_size, # 总采样数=迭代次数*批量大小
    replacement=True
)

# 构造DataLoader
dataloader = DataLoader(
    dataset,
    batch_size=batch_size,
    sampler=sampler,
    # 可按需加num_workers、pin_memory等参数
)

# 训练迭代,刚好跑total_iters次,每次都是有放回采样的批量
for batch_x, batch_y in dataloader:
    # 你的训练逻辑
    pass

注意:指定sampler参数后,不需要再设置DataLoader的shuffle参数,采样逻辑完全由sampler控制。

其他替代实现方案

如果不使用DataLoader,除了你提到的numpy手动采样方案外,还有两种更适配PyTorch生态的实现:

  • 用PyTorch原生的torch.randint生成索引,避免numpy和torch张量的格式转换开销:
for _ in range(total_iters):
    # 直接生成有放回的批量索引
    indices = torch.randint(low=0, high=len(dataset), size=(batch_size,))
    batch_x, batch_y = dataset[indices]
    # 训练逻辑
  • 自定义继承自Sampler的采样类,适合需要更复杂采样规则(比如动态调整样本权重、按类别约束采样比例)的场景,灵活度更高。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.29 11:18:01