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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.10.06 15:57:02