如何修复plt.imshow单张图像显示的负片异常?
CIFAR10图像单张显示负片,批量显示正常的问题解决
问题描述
在本地Python 3.11的conda环境中可视化CIFAR10图像,已安装numpy、matplotlib、PyTorch(通过conda install pytorch torchvision torchaudio pytorch-cuda=11.8 -c pytorch -c nvidia安装)及scikit-image。使用指定代码时,批量显示图像(子图循环方式)的结果与Google Colab一致,但单独执行plt.imshow(X[0].permute(1, 2, 0))显示单张图像时,本地呈现负片效果(如青蛙腹部白色变为黑色),保存图像也存在该问题。仅对图像做了Resize和CenterCrop变换,未启用归一化操作。
复现代码:
import torch import matplotlib.pyplot as plt import torchvision.transforms as T import torchvision.datasets as datasets transform = T.Compose([ T.Resize(256), T.CenterCrop(224), T.ToTensor(), # T.Normalize( # mean=[0.485, 0.456, 0.406], # std=[0.229, 0.224, 0.225]) ]) train_dataset = datasets.CIFAR10( root='./data', download=True, transform=transform ) train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=16, shuffle=False, num_workers=4) X,y = next(iter(train_loader)) # 批量显示正常 fig, axes = plt.subplots(2, 8, figsize=(12, 4)) for i, ax in enumerate(axes.flat): image = X[i].permute(1, 2, 0) ax.imshow(image) ax.axis('off') plt.show() # 单张显示负片 plt.imshow(X[0].permute(1, 2, 0)) plt.show()
原因分析
问题源于本地matplotlib对PyTorch Tensor的单张显示逻辑差异:
- 批量显示时,
ax.imshow()对Tensor的内部转换逻辑与plt.imshow()不同,本地matplotlib的后端(如TkAgg/QtAgg)对单张Tensor的渲染规则,和Colab使用的Agg后端不一致。 - 直接传入PyTorch Tensor到
plt.imshow()时,本地matplotlib可能错误地自动调整像素值范围(比如将0-1的float值误判为其他区间,导致颜色反转),而子图循环中的ax.imshow()则能正确识别Tensor的数值范围。
解决方法
将PyTorch Tensor转换为numpy数组后再传入plt.imshow(),确保matplotlib正确识别像素值范围,以下是两种可行方案:
方案1:直接转换为numpy数组(GPU张量需先移至CPU)
# 张量在CPU上时 plt.imshow(X[0].permute(1, 2, 0).numpy()) # 张量在GPU上时 plt.imshow(X[0].permute(1, 2, 0).cpu().numpy()) plt.axis('off') plt.show()
方案2:转换为uint8格式(模拟原始图像0-255数值范围)
# 转换为0-255的uint8数组 img_np = (X[0].permute(1, 2, 0) * 255).byte().cpu().numpy() plt.imshow(img_np) plt.axis('off') plt.show()
可选验证:强制钳制像素值范围
如果图像存在异常数值(本例中不会出现),可先将数值限制在0-1区间再转换:
img_clamped = torch.clamp(X[0].permute(1, 2, 0), 0, 1) plt.imshow(img_clamped.cpu().numpy()) plt.axis('off') plt.show()
内容的提问来源于stack exchange,提问作者yaodao vang
相关产品推荐
相关产品推荐

