如何使用plt展示PyTorch DataLoader加载的batch图像张量
解答
完全可以使用matplotlib.pyplot实现张量的可视化需求。你当前拿到的x张量维度为torch.Size([16, 3, 448, 448]),符合PyTorch默认的NCHW(批次大小、通道数、高度、宽度)存储格式,而matplotlib绘图要求单张图像输入为HWC(高度、宽度、通道数)格式,同时需要匹配0-1或0-255的像素取值范围,具体实现步骤如下:
前置依赖导入
import matplotlib.pyplot as plt import torch
单张图像可视化
以展示批次中第一张图像为例:
# 1. 取出单张图像张量,此时维度为 [3, 448, 448] single_img = x[0] # 2. 调整通道顺序,将通道维度放到最后,调整后维度为 [448, 448, 3] single_img = single_img.permute(1, 2, 0) # 3. 若预处理时做了归一化操作,需要先反归一化恢复正常像素范围(无归一化可跳过该步) # 示例为ImageNet常用归一化参数的反归一化,可替换为你自己的预处理参数 mean = torch.tensor([0.485, 0.456, 0.406]) std = torch.tensor([0.229, 0.224, 0.225]) single_img = single_img * std + mean # 4. 张量转CPU并转为numpy数组,同时截断像素值到0-1范围避免绘图报错 single_img = single_img.cpu().clip(0, 1).numpy() # 5. 绘图展示 plt.imshow(single_img) plt.axis('off') # 隐藏坐标轴 plt.show()
批次全量图像网格展示
如果需要一次性展示16张批次内的所有图像,可以按4行4列网格布局实现:
plt.figure(figsize=(16, 16)) # 反归一化参数,无归一化可删除相关逻辑 mean = torch.tensor([0.485, 0.456, 0.406]) std = torch.tensor([0.229, 0.224, 0.225]) for i in range(16): # 单张图像预处理 img = x[i].permute(1, 2, 0).cpu() img = img * std + mean img = img.clip(0, 1).numpy() # 子图绘制 plt.subplot(4, 4, i+1) plt.imshow(img) plt.title(f"对应标签: {y_true[i].item()}") # 可按需显示对应标签 plt.axis('off') plt.tight_layout() # 自动调整子图间距避免重叠 plt.show()
注意事项
- 若你的图像张量本身存储的是0-255范围的整数值,无需做反归一化,直接将数组转为
uint8类型即可:single_img = single_img.numpy().astype('uint8') - 如果是单通道灰度图,需要先squeeze掉通道维度,或者在
imshow方法中指定cmap='gray'参数
内容的提问来源于stack exchange,提问作者dextrin
相关产品推荐
相关产品推荐

