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

如何修复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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.16 23:55:00