CIFAR-100昆虫子类分类代码IndexError问题排查与修正求助
解决CIFAR-100昆虫分类筛选的IndexError及完整实现
错误原因分析
你遇到的IndexError: index 1 is out of bounds for axis 1 with size 1,本质是代码在尝试访问数组/张量第二个维度(axis=1)的第1个索引,但该维度实际只有1个元素(维度大小为1)。最常见的场景是:
- 错误地将CIFAR-100的一维标签数组当成了二维数组,比如用
labels[:,1]访问,但实际标签形状是(样本数,)或(样本数,1),axis=1没有索引1的位置。 - 混淆了粗分类标签和细分类标签的存储结构,比如误把粗分类标签的维度当成了包含细分类的二维结构。
完整可运行代码(含错误修复)
以下是针对CIFAR-100筛选Insects粗分类下5个细分类、并显示图像的代码,已修复索引错误:
import torchvision import torchvision.transforms as transforms import matplotlib.pyplot as plt import numpy as np # 1. 加载CIFAR-100数据集,同时获取粗分类和细分类标签 transform = transforms.Compose([transforms.ToTensor()]) trainset = torchvision.datasets.CIFAR100( root='./data', train=True, download=True, transform=transform ) # 2. 定义目标类别:粗分类"Insects"对应的细分类 insect_fine_classes = ["bee", "beetle", "butterfly", "caterpillar", "cockroach"] # 获取CIFAR-100的类别映射 fine_label_to_name = trainset.classes name_to_fine_label = {name: idx for idx, name in enumerate(fine_label_to_name)} # 获取目标细分类的标签索引 target_fine_labels = [name_to_fine_label[name] for name in insect_fine_classes] # 3. 筛选出符合条件的样本(细分类属于目标列表) filtered_images = [] filtered_labels = [] for img, fine_label in trainset: if fine_label in target_fine_labels: filtered_images.append(img) filtered_labels.append(fine_label) # 4. 显示筛选后的昆虫图像 def show_images(images, labels, class_names, num_rows=2, num_cols=3): fig, axes = plt.subplots(num_rows, num_cols, figsize=(12, 8)) axes = axes.flatten() for i, (img, label) in enumerate(zip(images[:num_rows*num_cols], labels[:num_rows*num_cols])): # 将张量转换为numpy并调整通道顺序(CIFAR是CHW,matplotlib是HWC) img_np = img.permute(1, 2, 0).numpy() # 因为ToTensor()会将像素值归一化到[0,1],这里还原到[0,255]方便显示 img_np = np.clip(img_np * 255, 0, 255).astype(np.uint8) axes[i].imshow(img_np) axes[i].set_title(class_names[label]) axes[i].axis('off') plt.tight_layout() plt.show() # 调用显示函数 show_images(filtered_images, filtered_labels, fine_label_to_name)
代码关键点说明(排查常见问题)
标签维度处理:
- CIFAR-100的
trainset.targets是细分类标签的一维列表(形状(50000,)),trainset.coarse_targets是粗分类标签的一维列表。如果你的代码之前错误地将标签转换为二维数组(比如用np.expand_dims),再尝试取[:,1]就会触发索引错误。本代码直接使用一维标签进行判断,避免了维度错误。
- CIFAR-100的
类别映射正确性:
- 确保你使用的类别名称与CIFAR-100官方定义一致,比如"Insects"是粗分类名称,对应的细分类确实是你列出的5种,避免因名称拼写错误导致筛选不到样本。
图像显示的通道调整:
- torchvision加载的图像是
(C, H, W)的张量格式,而matplotlib需要(H, W, C)的格式,所以必须用permute(1,2,0)转换通道顺序,否则会显示颜色错乱。
- torchvision加载的图像是
像素值还原:
ToTensor()会将图像像素值从[0,255]归一化到[0,1],显示前需要还原回[0,255]并转换为uint8类型,否则图像会偏暗或显示异常。
原错误的具体修复示例
如果你的第17行代码类似这样:
coarse_labels = trainset.coarse_targets[:, 1] # 错误写法
改成直接使用一维标签:
coarse_labels = trainset.coarse_targets # 正确写法,coarse_labels是一维列表
如果是需要通过粗分类筛选,先找到"Insects"对应的粗分类索引,再筛选:
# 获取粗分类名称映射 coarse_label_to_name = trainset.coarse_classes name_to_coarse_label = {name: idx for idx, name in enumerate(coarse_label_to_name)} insect_coarse_label = name_to_coarse_label["Insects"] # 通过粗分类筛选样本 filtered_images = [] filtered_labels = [] for img, fine_label, coarse_label in zip(trainset.data, trainset.targets, trainset.coarse_targets): if coarse_label == insect_coarse_label: filtered_images.append(img) filtered_labels.append(fine_label)
这种方式也能筛选出Insects粗分类下的所有细分类样本,再匹配你需要的5种即可。
内容的提问来源于stack exchange,提问作者Cyrus Noel Carano-o
相关产品推荐
相关产品推荐

