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
相关产品推荐
相关产品推荐

