如何解决PyTorch自定义数据集经DataLoader加载后形状异常的问题
问题修复方案
问题原因
- 你的数据集总样本量仅为4,你设置的
batch_size=64远大于数据集总样本数,PyTorch的DataLoader默认会将所有不足batch_size的剩余样本合并为一个batch返回,因此你拿到的batch第一维度为4是正常的batch维度,对应当前batch包含的样本总数。 - 你代码中的
one列表存储的是每次迭代拿到的batch数据,而非单个样本数据,如果你预期获取单个样本,直接从dataset索引取值即可,不需要遍历DataLoader。 - 额外注意:
F.one_hot默认返回torch.int64类型的张量,GAN训练通常需要浮点型输入,建议新增类型转换避免后续训练报错。
修复方法
方案1:调整batch_size或数据集大小
如果需要batch第一维度符合设置值,要么扩充数据集到至少64个样本,要么将batch_size调整为小于等于4的值:
# 调整batch_size为符合数据集大小的值 dataloader = DataLoader(dataset, batch_size=2, shuffle=True) # 此时每个batch的shape为 [2, 1274, 22]
方案2:丢弃不足batch_size的批次
如果你不需要保留不满的批次,可添加drop_last=True参数,注意你当前仅4个样本的情况下,设置batch_size=64加该参数会导致没有数据返回:
dataloader = DataLoader(dataset, batch_size=64, shuffle=True, drop_last=True)
方案3:修改数据集类返回浮点型张量(推荐)
适配GAN训练的输入类型要求:
def __getitem__(self, index): 'Generates one sample of data' mylist = self.list_IDs[index] # 新增.float()转浮点类型 X = F.one_hot(mylist, num_classes=len(alphabet)).float() y = self.labels[index] return X, y
方案4:自定义collate_fn固定batch大小(可选)
如果需要强制每个batch为64的固定大小,不足的部分用填充值补全,可以自定义collate_fn实现:
def collate_fn(batch): # batch是 [(X1,y1), (X2,y2), ...] 的列表 Xs, ys = zip(*batch) Xs = list(Xs) # 不足64个的话用零张量填充 while len(Xs) < 64: Xs.append(torch.zeros_like(Xs[0])) return torch.stack(Xs), torch.tensor(ys) dataloader = DataLoader(dataset, batch_size=64, shuffle=True, collate_fn=collate_fn) # 此时返回的batch shape为 [64, 1274, 22]
内容的提问来源于stack exchange,提问作者khashayar ehteshami
相关产品推荐
相关产品推荐

