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

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.29 15:43:32