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

PyTorch加载无标签图像验证集遇多类错误求解决方案

PyTorch加载无标签图像数据集的问题排查与解决

问题背景

我有一个用于CNN分类任务的无标签图像验证数据集,需要用PyTorch加载。由于没有标注信息,无法使用torchvision.datasets.ImageFolder(),因此需要自定义Dataset类并传入torch.utils.data.DataLoader(),但编写过程中遇到了多个错误。


初始尝试与第一个错误

自定义Dataset类代码:

import torch
from torch.utils.data import Dataset
from torchvision.io import read_image
import os


class Dset(Dataset):
    def __init__(self, dir: str, transform=None) -> None:
        self.transform = transform
        self.images = os.listdir(dir)
        self.dir = dir
    
    def __getitem__(self, index: int) -> torch.Tensor:
        image = read_image(f'{self.dir}/{self.images[index]}')
        if self.transform is not None:
            image = self.transform(image)
        return image

    def __len__(self) -> int:
        return len(self.images)

加载数据集代码:

from torchvision import transforms

batch_size = 64
transform = transforms.Compose([transforms.Grayscale(), transforms.ToTensor()])
data = Dset('data', transform=transform)
trainloader = torch.utils.data.DataLoader(data, batch_size=batch_size, shuffle=True)

images, labels = iter(trainloader)

错误信息:

TypeError: Input image tensor permitted channel values are [1, 3], but found 4

第一次更新:修复Alpha通道问题

修改Dataset类,指定读取模式为RGB以忽略Alpha通道:

import torch
from torch.utils.data import Dataset
from torchvision.io import read_image, ImageReadMode
import os


class Dset(Dataset):
    def __init__(self, dir: str, transform=None) -> None:
        self.transform = transform
        self.images = os.listdir(dir)
        self.dir = dir
    
    def __getitem__(self, index: int) -> torch.Tensor:
        image = read_image(f'{self.dir}/{self.images[index]}', mode=ImageReadMode.RGB)
        if self.transform is not None:
            image = self.transform(image)
        return image

    def __len__(self) -> int:
        return len(self.images)

新错误信息:

TypeError: pic should be PIL Image or ndarray. Got <class 'torch.Tensor'>

第二次更新:调整Transform后的新错误

修改Transform,添加ToPILImage()转换:

from torchvision import transforms

batch_size = 64
transform = transforms.Compose(
[transforms.ToPILImage(), transforms.Resize((512, 512)),
transforms.Grayscale(), transforms.ToTensor()]
)
data = Dset('data', transform=transform)
trainloader = torch.utils.data.DataLoader(data, batch_size=batch_size, shuffle=True)

images = iter(trainloader)[0]

错误信息:

TypeError: '_SingleProcessDataLoaderIter' object is not subscriptable

完整解决方案

1. 自定义Dataset类的正确实现

read_image返回的是张量,而大部分torchvision变换默认只支持PIL Image或numpy数组,因此直接用PIL读取图像更简洁,避免额外类型转换:

import torch
from torch.utils.data import Dataset
from PIL import Image
import os


class UnlabeledImageDataset(Dataset):
    def __init__(self, img_dir: str, transform=None) -> None:
        self.img_dir = img_dir
        self.transform = transform
        # 过滤非图像文件,避免读取错误
        self.img_paths = [
            os.path.join(img_dir, fname) 
            for fname in os.listdir(img_dir)
            if fname.lower().endswith(('.png', '.jpg', '.jpeg', '.bmp'))
        ]
    
    def __getitem__(self, index: int) -> torch.Tensor:
        # 用PIL读取图像并转为RGB格式
        img = Image.open(self.img_paths[index]).convert('RGB')
        if self.transform:
            img = self.transform(img)
        return img

    def __len__(self) -> int:
        return len(self.img_paths)

2. Transform的正确配置

无需额外添加ToPILImage(),直接对PIL图像应用变换即可:

from torchvision import transforms

batch_size = 64
transform = transforms.Compose([
    transforms.Resize((512, 512)),
    transforms.Grayscale(),
    transforms.ToTensor()
])

3. DataLoader的正确迭代方式

iter(trainloader)返回的是迭代器,需用next()获取批次数据,而非下标访问:

data = UnlabeledImageDataset('data', transform=transform)
trainloader = torch.utils.data.DataLoader(data, batch_size=batch_size, shuffle=True)

# 获取第一个批次的图像数据
images = next(iter(trainloader))
# 遍历所有批次的写法
for images in trainloader:
    # 执行推理等操作
    pass

内容的提问来源于stack exchange,提问作者dimicorn

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.29 22:30:31