解决AttributeError:numpy.ndarray无imshow属性,实现CIFAR100子图绘制
解决
AttributeError: 'numpy.ndarray' object has no attribute 'imshow'在绘制CIFAR100子图时的问题 这个错误的核心原因是个常见的小疏漏:plt.subplots(4,5)返回的axes是2维NumPy数组(形状为(4,5)),但你用一维索引axes[idx]去访问它。比如idx=0会拿到整行的5个Axes对象(是一个数组),idx=5会尝试取不存在的第6行,自然调用imshow会报错——数组没有这个方法。
下面是两种可行的修复方案,任选其一即可:
方案1:将子图轴数组展平为一维
修改plot函数,把2维的axes展平成一维数组,这样就能直接用idx遍历单个子图:
# 确保cifar100_mean和std是正确的标准化参数 cifar100_mean = [0.5071, 0.4867, 0.4408] cifar100_std = [0.2675, 0.2565, 0.2761] dm = torch.as_tensor(cifar100_mean, **setup)[:, None, None] ds = torch.as_tensor(cifar100_std, **setup)[:, None, None] def plot(tensor, labels=None): tensor = tensor.clone().detach() # 反标准化并限制在0-1区间,适配imshow的要求 tensor.mul_(ds).add_(dm).clamp_(0, 1) if tensor.shape[0] == 1: return plt.imshow(tensor[0].permute(1, 2, 0).cpu()) else: fig, axes = plt.subplots(4, 5, figsize=(12, 12)) # 关键操作:将2维axes数组展平为1维 axes = axes.flatten() for idx, img in enumerate(tensor): # 转换为CPU上的numpy数组供imshow使用 img_np = img.permute(1, 2, 0).cpu().numpy() axes[idx].imshow(img_np) # 可选:添加类别标签作为子图标题 if labels: axes[idx].set_title(labels[idx], fontsize=10) # 隐藏坐标轴让图像更整洁 axes[idx].axis('off') plt.tight_layout() # 自动调整子图间距 plt.show()
方案2:使用二维索引定位子图
如果你不想展平数组,可以通过计算行和列的索引来精准访问每个子图:
cifar100_mean = [0.5071, 0.4867, 0.4408] cifar100_std = [0.2675, 0.2565, 0.2761] dm = torch.as_tensor(cifar100_mean, **setup)[:, None, None] ds = torch.as_tensor(cifar100_std, **setup)[:, None, None] def plot(tensor, labels=None): tensor = tensor.clone().detach() tensor.mul_(ds).add_(dm).clamp_(0, 1) if tensor.shape[0] == 1: return plt.imshow(tensor[0].permute(1, 2, 0).cpu()) else: rows, cols = 4, 5 fig, axes = plt.subplots(rows, cols, figsize=(12, 12)) for idx, img in enumerate(tensor): # 计算当前子图的行、列索引 row = idx // cols col = idx % cols img_np = img.permute(1, 2, 0).cpu().numpy() axes[row, col].imshow(img_np) if labels: axes[row, col].set_title(labels[idx], fontsize=10) axes[row, col].axis('off') plt.tight_layout() plt.show()
额外优化:完善ground_truth构建代码
你的数据加载部分可以简化,避免混淆,同时确保逻辑正确:
from torchvision import datasets, transforms import torch device = torch.device('cuda') if torch.cuda.is_available() else torch.device('cpu') setup = {'device': device} num_images = 20 data_path = './data' # 替换为你的数据存储路径 # 初始化CIFAR100验证集 valid_transform = transforms.ToTensor() validset = datasets.CIFAR100(root=data_path, train=False, download=True, transform=valid_transform) # 构建包含不同类别的ground_truth和标签 ground_truth, labels = [], [] idx = 125 while len(labels) < num_images: img, label = validset[idx] # 直接从dataset取数据,无需通过loader idx += 1 if label not in labels: labels.append(label) ground_truth.append(img.to(device)) ground_truth = torch.stack(ground_truth) labels_print = [validset.classes[l] for l in labels] # 调用plot函数,传入标签显示子图标题 plot(ground_truth, labels_print) print(labels_print)
这样修改后,你就能成功绘制4行5列的CIFAR100图像子图,不会再触发那个属性错误了。
内容的提问来源于stack exchange,提问作者Unknown Shin
相关产品推荐
相关产品推荐

