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

构建PyTorch DataLoader遇RuntimeError:张量尺寸不一致问题咨询

PyTorch DataLoader RuntimeError: Inconsistent Tensor Sizes

嘿,我来帮你拆解这个问题~首先解答你的核心疑惑:为什么遍历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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.28 07:31:31