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

自定义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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.18 12:02:29