未标记图像数据集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
相关产品推荐
相关产品推荐

