PyTorch单图像数据增强报错排查:GPU张量与形状问题
单张图像数据增强的两类错误排查与修正
错误1:GPU张量无法被CPU访问
错误原因
- 你将
test_image移至GPU(test_image.to(device)),但torchvision.transforms的多数操作(如ToPILImage)以及matplotlib.pyplot.imshow仅能处理CPU上的张量/图像数据。 - 直接把GPU张量传给
plt.imshow会触发类型错误,因为plt需要将张量转为numpy数组,而GPU张量无法直接完成转换。
修正代码
import matplotlib.pyplot as plt import torchvision.transforms as transforms def Show_Image(Image, Picture_Name): # 统一处理张量与PIL图像的显示逻辑 if isinstance(Image, torch.Tensor): # 先将张量移至CPU,再调整维度从(C,H,W)转为(H,W,C),最后转numpy Image = Image.cpu().permute(1, 2, 0).numpy() # 若张量经过归一化,需将数值还原到0-1范围 Image = Image.clip(0, 1) plt.imshow(Image) plt.title(Picture_Name) plt.show() # train_dl 为训练数据加载器 data_iter = iter(train_dl) images, label = next(data_iter) test_image = images[0] # 数据增强无需提前移至GPU,完成增强后可按需再转移 # test_image = test_image.to(device) Horizontal_Flipping_Transformation = transforms.Compose([ transforms.ToPILImage(), transforms.RandomHorizontalFlip(), transforms.ToTensor() # 最后转为张量即可,无需重复转PIL ]) Flipping_Img = Horizontal_Flipping_Transformation(test_image) Show_Image(test_image, 'Original Image') Show_Image(Flipping_Img, 'Flipped Image') # 若后续需要GPU训练,再将增强后的张量移至设备 # Flipping_Img = Flipping_Img.to(device)
错误2:直接读取单张图像时形状无效
错误原因
torchvision.io.read_image返回的张量形状为**(通道数C, 高度H, 宽度W),但plt.imshow要求的图像形状是(H, W, C)(RGB图像)或(H, W)**(灰度图像),维度不匹配导致报错。- 代码中
Show_Image(test_image, 'Original Image')的test_image未定义,应为img。
修正代码(两种方案)
方案1:调整张量维度后显示
import matplotlib.pyplot as plt import numpy as np import torchvision.transforms as transforms from torchvision.io import read_image def Show_Image(Image, Picture_Name): if isinstance(Image, torch.Tensor): # 调整维度顺序:(C,H,W) -> (H,W,C),再转numpy Image = Image.permute(1, 2, 0).numpy() # 处理不同数据类型的显示适配 if Image.dtype == np.float32: Image = Image.clip(0, 1) plt.imshow(Image) plt.title(Picture_Name) plt.show() img = read_image(r'baboon\n02486410_1.JPEG') Horizontal_Flipping_Transformation = transforms.Compose([ transforms.ToPILImage(), transforms.RandomHorizontalFlip(), transforms.ToTensor() ]) Flipping_Img = Horizontal_Flipping_Transformation(img) Show_Image(img, 'Original Image') # 修正变量名 Show_Image(Flipping_Img, 'Flipped Image')
方案2:用PIL直接读取图像(更简洁)
import matplotlib.pyplot as plt import torchvision.transforms as transforms from PIL import Image def Show_Image(Image, Picture_Name): plt.imshow(Image) plt.title(Picture_Name) plt.show() # 用PIL读取直接得到PIL图像,无需手动调整维度 img = Image.open(r'baboon\n02486410_1.JPEG') Horizontal_Flipping_Transformation = transforms.Compose([ transforms.RandomHorizontalFlip(), transforms.ToTensor() ]) Flipping_Img = Horizontal_Flipping_Transformation(img) Show_Image(img, 'Original Image') # 将增强后的张量转回PIL图像再显示 Show_Image(transforms.ToPILImage()(Flipping_Img), 'Flipped Image')
内容的提问来源于stack exchange,提问作者Rahul Anand
相关产品推荐
相关产品推荐

