为何使用同一张PyTorch张量时,PIL与plt.imshow显示的图像存在差异?
为何使用同一张PyTorch张量时,PIL与plt.imshow显示的图像存在差异?
嗨,我来帮你捋捋为啥同一张张量转成图像后,PIL和plt.imshow显示出来不一样~
其实核心问题出在数值范围和数据类型的兼容性上,咱们一点点说:
1. 两者对输入数据的要求不一样
plt.imshow特别“贴心”,不管你给的是0-1的浮点数数组,还是0-255的整数数组,它都会自动调整显示范围,把数据映射到适合人眼观看的灰度区间,所以你能看到正常的图像。- 但PIL的"L"灰度模式就很“较真”:它要求输入必须是0-255之间的uint8类型整数。如果你的PyTorch张量是经过归一化的(比如常见的0-1浮点数范围),直接转成numpy数组丢给PIL的话,大部分小于1的浮点数会被直接截断成0,显示出来自然就偏暗甚至全黑,和plt的效果天差地别。
2. 你的代码里缺了关键的转换步骤
看你写的代码,直接把张量转成的numpy数组传给了PIL,没做数值缩放和类型转换。咱们来改改这个函数,就能让两者显示一致了:
import matplotlib.pyplot as plt from PIL import Image import numpy as np def tensor_to_pil(image_tensor): # 提取单通道的numpy数组(先移到CPU,压缩维度) img_np = image_tensor[0].cpu().squeeze().numpy() # plt显示部分(保持你的原逻辑就行) plt.figure() plt.imshow(img_np, cmap='gray') plt.title("plt.imshow 显示效果") # 处理PIL图像的关键步骤: # 第一步:把数值范围转成0-255 # 如果你的张量是0-1范围的浮点数,直接乘255;如果是-1到1的范围,先转成0-1再乘 # 这里先假设是0-1的情况 img_np_scaled = (img_np * 255).astype(np.uint8) # 第二步:用处理好的uint8数组创建PIL图像 pil_image = Image.fromarray(img_np_scaled, "L") # 显示PIL图像 pil_image.show(title="PIL 显示效果") return pil_image
额外说明
如果你的张量是用torchvision.transforms.Normalize做过归一化(比如范围是-1到1),那在缩放前还要先把数值转成0-1:
# 先把-1到1的范围转成0-1 img_np = (img_np + 1) / 2 # 再缩放转成uint8 img_np_scaled = (img_np * 255).astype(np.uint8)
这样修改后,PIL和plt显示的图像就会一致啦~
备注:内容来源于stack exchange,提问作者Hjin
相关产品推荐
相关产品推荐

