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

如何解决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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.09.30 17:45:03