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

PyTorch展示变换后图像时出现TypeError的解决求助

解决PyTorch图像变换中的类型不匹配错误

问题核心

你遇到的两类错误都是因为图像变换顺序违反了PyTorch变换的输入输出类型要求:

  • transforms.Normalize仅能处理Tensor类型图像,无法直接操作PIL Image
  • transforms.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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.02 14:39:55