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

重塑PIL图像转换的PyTorch张量后显示多幅灰度图问题

torchvision加载花卉数据集图像显示异常排查

问题现象

使用torchvision.datasets.ImageFolder加载花卉数据集时,经ToTensor()转换后的张量通过reshape调整维度后显示错乱,原本的彩色图像呈现为多幅灰度图拼接效果,PIL直接读取的原图像显示正常:
PIL正常显示的原花卉图像

原问题代码如下:

# creating a flower dataset
f_ds = torchvision.datasets.ImageFolder(data_path)
# a transform to convert images to tensors
to_tensor = torchvision.transforms.ToTensor()


for idx, (img, label) in enumerate(f_ds):
    if idx == 2100:
        # random PIL image
        display(img)
        print(img.size, img.mode) # W * H

        # convert the same arrray to_tensor (a torchvision transform to convert images to pytorch tensor)
        img_tensor = to_tensor(img)
        print("After converting to torch tensor: ", img_tensor.shape)
        C, H, W = img_tensor.shape  
        
        # the same image reshaped to match matplotlib
        reshaped_img_tensor = img_tensor.reshape(H, W, C) # i think the problem is here...
        print('After reshaping img_tensor: ', reshaped_img_tensor.shape)
        new_arr = (reshaped_img_tensor.numpy()*255.0).astype(np.uint8)
        print("dtype:", new_arr.dtype, "min:", new_arr.min(), "max:",new_arr.max(), "shape:", new_arr.shape)
        display(new_arr, 'RGB')) 
        break

问题根因

判断完全正确,问题出在维度调整操作上:

  • ToTensor()输出的PyTorch图像张量维度顺序为[通道数(C), 图像高度(H), 图像宽度(W)],而常规图像显示库(PIL、matplotlib等)要求的输入维度顺序是[图像高度(H), 图像宽度(W), 通道数(C)]。
  • reshape()操作只会按照内存中元素的连续存储顺序重新分配张量形状,不会调换维度的实际排布逻辑。直接reshape会把原本连续存储的R通道整幅、G通道整幅、B通道整幅数据硬拆分配到新形状下,最终呈现出三幅灰度图拼接的错乱效果。
  • 额外语法问题:原代码最后一行display(new_arr, 'RGB'))多了一个右括号,运行时会触发语法错误,需要删除多余括号。

修复方案

维度顺序调换不能用reshape(),需要使用专门的维度置换接口,两种常用实现方式:

  • 对PyTorch张量,使用permute()方法指定新的维度顺序:原张量维度索引对应关系为0=C、1=H、2=W,因此目标维度顺序对应索引为(1,2,0),代码为:
    reshaped_img_tensor = img_tensor.permute(1, 2, 0)
    
  • 如果已经将张量转为numpy数组,可使用numpy的transpose()方法实现相同的维度调换:
    new_arr = (img_tensor.numpy().transpose(1, 2, 0)*255.0).astype(np.uint8)
    

修正后的关键代码段:

for idx, (img, label) in enumerate(f_ds):
    if idx == 2100:
        display(img)
        print(img.size, img.mode)
        img_tensor = to_tensor(img)
        print("After converting to torch tensor: ", img_tensor.shape)
        C, H, W = img_tensor.shape  
        # 替换原reshape逻辑,使用permute调换维度顺序
        reshaped_img_tensor = img_tensor.permute(1, 2, 0)
        print('After reshaping img_tensor: ', reshaped_img_tensor.shape)
        new_arr = (reshaped_img_tensor.numpy()*255.0).astype(np.uint8)
        print("dtype:", new_arr.dtype, "min:", new_arr.min(), "max:",new_arr.max(), "shape:", new_arr.shape)
        # 删除多余右括号
        display(new_arr, 'RGB') 
        break

内容的提问来源于stack exchange,提问作者Loki

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.27 19:51:19