构建PyTorch DataLoader遇RuntimeError:张量尺寸不一致问题咨询
嘿,我来帮你拆解这个问题~首先解答你的核心疑惑:为什么遍历DataLoader就报错,而不是把数据输入网络时才出问题?
PyTorch的DataLoader在构建batch的阶段,会调用默认的
collate_fn把一个batch里的所有样本打包成tensor。这个打包过程发生在你遍历DataLoader的瞬间(也就是for i, batch in enumerate(train_loader)这一行),而不是等你把数据喂进模型。当样本的尺寸不一致时,torch.stack就会直接抛出尺寸不匹配的错误——因为它根本没办法把不同形状的tensor堆叠到一起。
接下来看你的代码里的具体问题和修复方案:
1. 图片尺寸不一致的核心原因
你目前只使用了transforms.Resize(224),这个方法默认是保持图片原宽高比的:它会把图片的短边缩放到224,长边按比例自动缩放。比如一张448x224的图,缩放后会变成224x112;另一张224x224的图缩放后还是224x224。这样不同图片的尺寸依然不统一,转成numpy数组后,collate_fn尝试堆叠它们就会直接报错。
修复方案:
改成强制缩放到固定尺寸,或者结合裁剪(更适合训练场景):
- 强制固定尺寸:
transforms.Resize((224, 224))(直接把所有图片缩放到224x224,忽略原比例) - 比例保留+中心裁剪:
transforms.Compose([transforms.Resize(256), transforms.CenterCrop(224)])(先把短边缩到256,再中心裁剪成224x224,保留图片主体) - 训练时还可以用
transforms.RandomResizedCrop(224),增加数据增强效果
2. __len__方法的逻辑错误
你现在的__len__返回的是len(os.listdir(self.root_dir)),但这个数值和csv文件里的图片数量不一定一致。比如如果csv里的条目数比文件夹里的图片少,当index超过csv的长度时,self.csv_file['Image'][index]就会抛出索引越界的错误。
修复方案:
把__len__改成返回csv文件的实际条目数:
def __len__(self): return len(self.csv_file)
3. 可能的label取值错误
你当前的label是self.csv_file['Image'][index]——也就是图片文件名,这应该不是你真正想要的类别标签吧?通常这类任务的csv里会有类似Id的列来表示图片对应的类别,比如:
label = self.csv_file['Id'][index]
如果确实是用文件名当标签,可以忽略这条,但大概率是写错了。
4. Transform初始化逻辑优化
你在__init__里直接把self.transform赋值为transforms.Resize(224),忽略了参数传入的transform。应该改成:
def __init__(self, data_file, root_dir, transform=None): self.csv_file = pd.read_csv(data_file) self.root_dir = root_dir # 如果用户没传transform,就用默认的固定尺寸缩放+转tensor if transform is None: self.transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor() # 建议加上ToTensor,直接转成PyTorch tensor并归一化像素值 ]) else: self.transform = transform
加上ToTensor()的好处是:直接把PIL图片转成PyTorch tensor,不需要手动转numpy,还会把像素值从0-255归一化到0-1,符合PyTorch的输入要求。
5. 遍历循环的冗余代码
你循环里的(i, batch)这行代码没有实际作用,可以去掉,或者改成打印batch信息来验证:
for i, batch in enumerate(train_loader): print(f"Batch {i}: Image shape {batch['image'].shape}, Labels {batch['label']}") # 测试时可以只看前几个batch if i == 3: break
修复后的完整代码示例
import os import pandas as pd from PIL import Image import torch from torch.utils.data import Dataset, DataLoader import torchvision.transforms as transforms class WhaleData(Dataset): def __init__(self, data_file, root_dir, transform=None): self.csv_file = pd.read_csv(data_file) self.root_dir = root_dir # 设置默认transform if transform is None: self.transform = transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor() ]) else: self.transform = transform def __len__(self): return len(self.csv_file) def __getitem__(self, index): img_name = os.path.join(self.root_dir, self.csv_file['Image'][index]) image = Image.open(img_name).convert('RGB') # 确保是RGB格式,避免灰度图的通道问题 if self.transform: image = self.transform(image) # 假设csv里有Id列作为类别标签 label = self.csv_file['Id'][index] sample = {'image': image, 'label': label} return sample trainset = WhaleData( data_file='/mnt/55-91e8-b2383e89165f/Ryan/1234/train.csv', root_dir='/mnt/4d55-91e8-b2383e89165f/Ryan/1234/train' ) train_loader = DataLoader(trainset, batch_size=4, shuffle=True, num_workers=2) # 测试遍历 for i, batch in enumerate(train_loader): print(f"Batch {i}: Image shape {batch['image'].shape}, Labels {batch['label']}") if i == 3: break
内容的提问来源于stack exchange,提问作者Ryan

