自定义ResNet训练时PyTorch DataLoader报stack尺寸不匹配RuntimeError
解决ResNet训练时DataLoader的张量尺寸不匹配问题
你遇到的RuntimeError: stack expects each tensor to be equal size,核心原因是数据集里混了单通道灰度图(对应[1,224,224])或四通道RGBA图(对应[4,224,224])——即使你认为都是彩色图,实际数据集里存在这类异常图片。Resize只修改图像尺寸,不会统一通道数,导致DataLoader批量堆叠张量时失败。
具体修复方案
1. 强制统一图像为三通道RGB
在__getitem__中打开图像后,直接调用convert('RGB'),把灰度图自动扩展为三通道,RGBA图去掉透明通道转为RGB:
image_data = Image.open(self.imgn_list[idx]).convert('RGB')
2. 过滤损坏/异常图片
数据集可能存在无法正常打开的图片,添加异常捕获逻辑跳过这类数据:
def __getitem__(self, idx): img_path = self.imgn_list[idx] try: image_data = Image.open(img_path).convert('RGB') if self.transforms: sample = self.transforms(image_data) return sample, self.img_label[idx] except Exception as e: print(f"跳过异常图片: {img_path}, 错误: {str(e)}") # 递归获取下一张有效图片,避免中断训练 return self.__getitem__((idx + 1) % len(self))
3. 避免Transforms变量名冲突
你的代码中transforms=transforms.Compose(...)存在变量名冲突(transforms既是torchvision的模块名,又是自定义变量),建议重命名:
from torchvision import transforms as T train_transforms = T.Compose([ T.Resize(size=(224, 224)), T.ToTensor() ])
修改后的完整Dataset类
class cnd_data(torch.utils.data.Dataset): def __init__(self, file_path, train=True, transforms=None): self.train = train self.transforms = transforms self.cat_img_path = os.path.join(file_path, 'data/kagglecatsanddogs/PetImages/Cat') self.dog_img_path = os.path.join(file_path, 'data/kagglecatsanddogs/PetImages/Dog') self.cat_list = natsort.natsorted(glob.glob(self.cat_img_path + '/*.jpg')) self.dog_list = natsort.natsorted(glob.glob(self.dog_img_path + '/*.jpg')) if self.train: self.imgn_list = self.cat_list[:12000] + self.dog_list[:12000] self.img_label = [0]*12000 + [1]*12000 else: self.imgn_list = self.cat_list[12000:] + self.dog_list[12000:] self.img_label = [0]*500 + [1]*500 def __len__(self): return len(self.img_label) def __getitem__(self, idx): img_path = self.imgn_list[idx] try: # 强制转换为三通道RGB image_data = Image.open(img_path).convert('RGB') if self.transforms: sample = self.transforms(image_data) return sample, self.img_label[idx] except Exception as e: print(f"跳过异常图片: {img_path}, 错误信息: {str(e)}") # 递归获取下一张有效图片 return self.__getitem__((idx + 1) % len(self))
内容的提问来源于stack exchange,提问作者GD G
相关产品推荐
相关产品推荐

