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

解决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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.05.09 19:27:34