PyTorch使用DataLoader加载自定义Dataset时多出维度如何解决
错误原因
你的__getitem__方法没有使用传入的索引参数取对应位置的单条样本,每次直接返回全量的数据集张量。DataLoader会收集batch_size次__getitem__的返回结果,在最外层新增维度拼接为批次张量,因此你会得到多出来的311维度。
修正后的代码
from torch.utils.data import Dataset, DataLoader class KD_Train(Dataset): def __init__(self, a, b): self.imgs = a self.labels = b def __len__(self): return len(self.imgs) def __getitem__(self, index): # 按传入的索引取对应单样本 return self.imgs[index], self.labels[index] kdt = KD_Train(x[train], y[train]) train_data_loader = DataLoader( kdt, batch_size=64, shuffle=True, num_workers=0 ) for step, (a, b) in enumerate(train_data_loader): print(a.shape) # 输出为 torch.Size([64, 3, 224, 224]),符合预期 break
简化方案
如果没有自定义Dataset的特殊逻辑,可以直接用PyTorch内置的TensorDataset省略类定义步骤:
from torch.utils.data import TensorDataset, DataLoader kdt = TensorDataset(x[train], y[train]) train_data_loader = DataLoader(kdt, batch_size=64, shuffle=True, num_workers=0)
内容的提问来源于stack exchange,提问作者Tim
相关产品推荐
相关产品推荐

