重塑PIL图像转换的PyTorch张量后显示多幅灰度图问题
torchvision加载花卉数据集图像显示异常排查
问题现象
使用torchvision.datasets.ImageFolder加载花卉数据集时,经ToTensor()转换后的张量通过reshape调整维度后显示错乱,原本的彩色图像呈现为多幅灰度图拼接效果,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
相关产品推荐
相关产品推荐

