PyTorch DataLoader遍历速度远慢于直接访问数据集的优化求助
PyTorch DataLoader遍历速度远慢于直接访问数据集的问题
问题描述
我在PyTorch训练机器学习模型时遇到严重性能瓶颈:遍历DataLoader的速度明显慢于直接访问数据集,导致训练过程中等待数据的时间过长,大幅降低训练效率。
对比示例:
- 遍历DataLoader耗时超15秒:
for inputs,labels in tqdm(dataloader): pass
- 直接遍历数据集耗时不到1秒:
for inputs,labels in tqdm(zip(dataloader.dataset.data, dataloader.dataset.targets)): pass
已尝试关闭shuffle功能、调整num_workers参数,但未能显著缩小耗时差距。当前CPU和内存使用率远未达上限,I/O性能也不是限制因素,但DataLoader的数据加载耗时仍远超预期。
基础复现示例
import torch from tqdm import tqdm from torchvision import datasets, transforms transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)), ]) trainset = datasets.MNIST('MNINST', download=True, train=True, transform=transform) trainloader = torch.utils.data.DataLoader(trainset, batch_size=1, shuffle=False)
运行遍历代码:
for data,targets in tqdm(trainloader): pass for data,targets in tqdm(zip(trainloader.dataset.data,trainloader.dataset.targets)): pass
测试结果显示两者耗时差距极为明显。
补充测试1(增大batch_size+开启shuffle)
随着batch_size增大,问题表现更突出。测试代码如下:
import torch from tqdm import tqdm from torchvision import datasets, transforms transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)), ]) batch_size=64 trainset = datasets.MNIST('MNINST', download=True, train=True, transform=transform) trainloader = torch.utils.data.DataLoader(trainset, batch_size=batch_size, shuffle=True) for data,targets in tqdm(trainloader): pass indices = torch.randperm(len(trainset)) for i in tqdm(range(0,len(indices),batch_size)): data = [] targets = [] for j in range(i,i+batch_size): if j < len(indices): data.append(trainset.data[indices[j]]) targets.append(trainset.targets[indices[j]]) data = torch.utils.data.default_collate(data) targets = torch.utils.data.default_collate(targets) tensor = (data.to(torch.float) / 256).unsqueeze(0) normalized = transforms.functional.normalize(tensor, (0.5,), (0.5,))
测试结果显示,DataLoader的耗时仍远高于手动实现的带shuffle和batch处理的加载逻辑。
补充测试2(自定义无transform数据集)
使用无数据预处理的自定义Dataset测试:
import torch from torch.utils.data import Dataset, DataLoader from tqdm import tqdm class CustomDataset(Dataset): def __init__(self, data, labels): self.data = data self.labels = labels def __len__(self): return len(self.data) def __getitem__(self, idx): sample = self.data[idx],self.labels[idx] return sample n=100000 data = torch.randn(n, 3, 28, 28) labels = torch.randint(0, 10, (n,)) custom_dataset = CustomDataset(data, labels) batch_size = 1 dataloader = DataLoader(custom_dataset, batch_size=batch_size, shuffle=False) for inputs, labels in tqdm(dataloader): pass for inputs, labels in tqdm(zip(dataloader.dataset.data,dataloader.dataset.labels)): pass
测试结果依然显示,DataLoader的遍历速度显著慢于直接访问数据集的方式。
核心需求
寻求有效解决方案,加快DataLoader的数据加载速度,缩小其与直接访问数据集的耗时差距,提升整体训练效率。
内容的提问来源于stack exchange,提问作者triple_double
相关产品推荐
相关产品推荐

