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

如何让PyTorch DataLoader使用动态批量大小列表训练?

实现PyTorch DataLoader动态批量大小训练

要实现按自定义的动态批量大小列表加载数据,核心是自定义BatchSampler——PyTorch中负责生成批量索引的组件,替代默认的固定批量采样逻辑。

步骤1:自定义动态批量采样器

继承Sampler类,根据给定的批量大小列表切分样本索引:

import torch
from torch.utils.data import TensorDataset, DataLoader, Sampler

class DynamicBatchSampler(Sampler):
    def __init__(self, dataset_size, batch_sizes):
        self.dataset_size = dataset_size
        self.batch_sizes = batch_sizes
        # 强制校验批量总和与样本数一致
        assert sum(batch_sizes) == dataset_size, "批量大小列表总和必须等于样本总数"

    def __iter__(self):
        # 生成样本索引,如需打乱数据,把arange换成randperm即可
        indices = torch.arange(self.dataset_size).tolist()
        start_idx = 0
        for batch_size in self.batch_sizes:
            end_idx = start_idx + batch_size
            yield indices[start_idx:end_idx]
            start_idx = end_idx

    def __len__(self):
        # 返回批量的数量,即列表长度
        return len(self.batch_sizes)

步骤2:配置DataLoader

把自定义采样器传给DataLoader的batch_sampler参数,同时必须将batch_size设为None:

# 假设x_train.shape=(8400,4),y_train是对应标签
train_dataset = TensorDataset(x_train, y_train)

# 你的动态批量大小列表,总和为8400
list_batch_size = [30, 60, 110, ..., 231]  # 替换为实际列表

# 初始化采样器
dynamic_sampler = DynamicBatchSampler(len(train_dataset), list_batch_size)

# 创建DataLoader
dataloader_train = DataLoader(train_dataset, batch_sampler=dynamic_sampler)

步骤3:训练时使用

直接迭代DataLoader即可,每次拿到的批量大小会严格遵循你定义的列表:

# 假设model、criterion、optimizer已定义
for batch_x, batch_y in dataloader_train:
    outputs = model(batch_x)
    loss = criterion(outputs, batch_y)
    
    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

额外说明

  • 如果需要每个epoch打乱数据,只需把采样器里的torch.arange改成torch.randperm,这样每次迭代都会生成随机顺序的索引块。
  • 务必保证list_batch_size的元素总和等于样本总数,否则采样器会触发断言报错,避免出现数据遗漏或重复加载的问题。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.05 08:29:54