PyTorch展示变换后图像时出现TypeError的解决求助
解决PyTorch图像变换中的类型不匹配错误
问题核心
你遇到的两类错误都是因为图像变换顺序违反了PyTorch变换的输入输出类型要求:
transforms.Normalize仅能处理Tensor类型图像,无法直接操作PIL Imagetransforms.RandomHorizontalFlip、transforms.Resize、transforms.CenterCrop等几何变换仅支持PIL Image输入,不能直接处理Tensor
原始代码中把Normalize放在PILToTensor之前,导致对PIL图像执行归一化触发第一个错误;调整顺序后若把Tensor转换放在几何变换之前,又会让几何变换接收Tensor输入,触发第二个错误。
修正后的变换代码
正确的顺序是:先对PIL图像完成所有几何/颜色变换,再转为Tensor,最后执行归一化。同时需要补充类型转换步骤,避免数值范围异常。
mean = [0.485, 0.456, 0.406] std = [0.229, 0.224, 0.225] train_transform = transforms.Compose([ transforms.RandomHorizontalFlip(p=0.5), transforms.Resize((256, 256), interpolation=torchvision.transforms.InterpolationMode.BILINEAR), transforms.CenterCrop(size=[224,224]), transforms.PILToTensor(), # 先将PIL图像转为uint8类型Tensor transforms.ConvertImageDtype(torch.float32), # 转换为float32类型,数值范围转为0-1 transforms.Normalize(mean, std) # 最后对Tensor执行归一化 ]) test_transform = transforms.Compose([ transforms.Resize((256, 256), interpolation=torchvision.transforms.InterpolationMode.BILINEAR), transforms.CenterCrop(size=[224,224]), transforms.PILToTensor(), transforms.ConvertImageDtype(torch.float32), transforms.Normalize(mean, std) ])
关键细节说明
transforms.PILToTensor()输出的是uint8类型Tensor(数值范围0-255),而Normalize要求输入float32类型Tensor(数值范围0-1),必须通过transforms.ConvertImageDtype(torch.float32)完成转换,否则会出现归一化后数值异常的问题。- 所有针对PIL图像的变换必须放在
PILToTensor之前,所有针对Tensor的变换必须放在PILToTensor之后。
数据集加载的额外优化
你当前的trainset和validset基于同一个原始ImageFolder创建,设置trainset.dataset.transform会同时影响validset的原始数据集(二者指向同一对象),建议修改为:
BATCH_SIZE = 32 # 分别创建独立的原始数据集,避免互相影响 trainset_full = torchvision.datasets.ImageFolder(root='CVPR2023_project_2_and_3_data/train/', loader=open_image) validset_full = torchvision.datasets.ImageFolder(root='CVPR2023_project_2_and_3_data/train/', loader=open_image) trainset_classes = trainset_full.classes.copy() subset_size = int(0.15*len(trainset_full)) indices = torch.randperm(len(trainset_full)) valid_indices = indices[:subset_size] train_indices = indices[subset_size:] trainset = Subset(trainset_full, train_indices) validset = Subset(validset_full, valid_indices) # 为独立的原始数据集设置变换 trainset_full.transform = train_transform validset_full.transform = test_transform trainloader = torch.utils.data.DataLoader(trainset, batch_size=BATCH_SIZE, shuffle=True, pin_memory=True) validloader = torch.utils.data.DataLoader(validset, batch_size=BATCH_SIZE, shuffle=False, pin_memory=True) testset = torchvision.datasets.ImageFolder(root='CVPR2023_project_2_and_3_data/test/', transform=test_transform, loader=Image.open) testloader = torch.utils.data.DataLoader(testset, batch_size=BATCH_SIZE, shuffle=False, pin_memory=True) print(f'entire train folder: {len(trainset_full)}, entire test folder: {len(testset)}, splitted trainset: {len(trainset)}, splitted validset: {len(validset)}')
内容的提问来源于stack exchange,提问作者Eminent Emperor Penguin
相关产品推荐
相关产品推荐

