如何让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
相关产品推荐
相关产品推荐

