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

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)

代码关键点说明(排查常见问题)

  1. 标签维度处理:

    • CIFAR-100的trainset.targets是细分类标签的一维列表(形状(50000,)),trainset.coarse_targets是粗分类标签的一维列表。如果你的代码之前错误地将标签转换为二维数组(比如用np.expand_dims),再尝试取[:,1]就会触发索引错误。本代码直接使用一维标签进行判断,避免了维度错误。
  2. 类别映射正确性:

    • 确保你使用的类别名称与CIFAR-100官方定义一致,比如"Insects"是粗分类名称,对应的细分类确实是你列出的5种,避免因名称拼写错误导致筛选不到样本。
  3. 图像显示的通道调整:

    • torchvision加载的图像是(C, H, W)的张量格式,而matplotlib需要(H, W, C)的格式,所以必须用permute(1,2,0)转换通道顺序,否则会显示颜色错乱。
  4. 像素值还原:

    • 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.31 09:56:19