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

未标记图像数据集SESEMI变换标签生成代码IndexError报错排查求助

问题排查与修复方案

咱们一步步拆解你遇到的两个核心问题:标签与变换不匹配,以及展示时触发的IndexError。

1. 直接引发IndexError的数据源错误

你通过random_split得到的unlabeled_trainset是一个Dataset对象,但你的UnlabeledDataset类要求传入np.ndarray类型的data参数。当你把Dataset对象传给它后,在__getitem__中执行data = self.data[idx]时,得到的是原ImageFolder返回的**(图像, 原始标签)**元组,而非单纯的图像数据!

这会导致一系列连锁问题:

  • 你的transform被错误地作用在元组上,而非图像
  • 后续SesemiTransform处理的也是元组,生成的label完全不符合预期,最终labels的结构或取值超出了unsup_classes的索引范围,触发报错。

修复方法

重构UnlabeledDataset,让它正确处理传入的Dataset对象,提取图像数据:

class UnlabeledDataset(Dataset):
    """ Unlabeled training Dataset class. """
    def __init__(self, dataset: Dataset, transform: Optional[transforms.Compose] = None):
        # 将参数类型从np.ndarray改为Dataset
        self.dataset = dataset
        self.transform = transform
        self.sesemi_transform = SesemiTransform()
    def __len__(self):
        return len(self.dataset)
    def __getitem__(self, idx: int) -> tuple:
        # 从原Dataset中取出图像,忽略原始标签
        image, _ = self.dataset[idx]
        if self.transform is not None:
            image = self.transform(image)
        # 对图像应用SESEMI变换,返回变换后的图像与对应标签
        transformed_image, label = self.sesemi_transform(image)
        return transformed_image, label

同时修改数据集加载代码:

trainset = torchvision.datasets.ImageFolder(root=train_img, transform=transform)
labeled_trainset, unlabeled_trainset = torch.utils.data.random_split(trainset, [1000, 2300])
# 传入Dataset对象而非numpy数组
unsupervised_trainset = UnlabeledDataset(unlabeled_trainset)

2. SesemiTransform中vflip的实现错误(标签与变换不匹配)

你定义的classes中第5项是'vflip',但对应的代码逻辑是旋转180度+水平翻转,和垂直翻转完全不是同一个操作!

修复方法

修正SesemiTransform的tf_type=5分支逻辑:

class SesemiTransform:    
    """ Torchvision-style transform to apply SESEMI augmentation to image. """
    classes = ('0', '90', '180', '270', 'hflip', 'vflip')
    def __call__(self, x):
        tf_type = random.randint(0, len(self.classes) - 1)
        if tf_type == 0:
            # 无变换,直接pass更简洁
            pass
        elif tf_type == 1:
            x = transforms.functional.rotate(x, 90)
        elif tf_type == 2:
            x = transforms.functional.rotate(x, 180)
        elif tf_type == 3:
            x = transforms.functional.rotate(x, 270)
        elif tf_type == 4:
            x = transforms.functional.hflip(x)
        elif tf_type == 5:
            # 修正为垂直翻转
            x = transforms.functional.vflip(x)
        return x, tf_type

3. 验证展示代码的正确性

修复前两个问题后,你可以先打印labels的内容,确认其取值范围是否在0-5之间:

dataiter = iter(unlabeled_trainloader)
images, labels = dataiter.next()
print("Labels张量内容:", labels)
print("Label取值:", labels.numpy())

确认无误后再运行展示代码,就能正常显示图像与对应的变换标签了。

内容的提问来源于stack exchange,提问作者beginner

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.04.30 13:12:34