PyTorch DataLoader中自定义CT数据集索引越界问题求助
问题解决:CT数据集训练时索引越界错误
核心错误修复
首先看代码中MyDataset类的__len__方法,存在明显笔误:
def __len__(self): if self.dataset_name == "cifar10": return len(self.cifar10) if self.dataset_name == "ct": return len(self.ct) if self.dataset_name == "satellite": return len(self.ct) # 此处错误,应返回len(self.satellite)
将satellite分支的返回值改为len(self.satellite),该错误虽未在satellite数据集上触发,但会影响代码逻辑一致性。
CT数据集索引越界的针对性排查
检查数据集长度与配置匹配
打印len(data_train)和len(data_unlabeled),确认两者实际长度是否一致,同时核对NUM_TRAIN(来自config.py)的值是否超过CT训练集的实际样本数。若训练过程中用NUM_TRAIN限定样本索引范围,当该值大于数据集实际长度时,会在多轮训练后触发索引越界。修正MyDataset的train_flag参数逻辑
当前CT和satellite分支的__init__方法完全忽略train_flag参数,始终加载train目录:if self.dataset_name == "ct": self.ct = ImageFolder(root='/Dataset/radiology_ai/CT/Split-CT-abd/train', transform=transf)若
data_unlabeled需要加载验证集而非训练集,应修改为:if self.dataset_name == "ct": root = '/Dataset/radiology_ai/CT/Split-CT-abd/train' if train_flag else '/Dataset/radiology_ai/CT/Split-CT-abd/val' self.ct = ImageFolder(root=root, transform=transf)这能避免unlabeled数据集与train数据集重复,也可排除因逻辑错误导致的索引问题。
排除多进程加载干扰
训练若干轮后才出现错误,可能与DataLoader多进程加载有关。尝试设置DataLoader(num_workers=0),若错误消失,说明是多进程环境下数据集状态被意外修改,此时需检查是否有动态修改数据集的代码,或确保数据集在多进程中被正确初始化。验证数据集完整性
遍历CT训练目录,统计实际样本数量,与len(data_train)对比,确认是否存在损坏文件导致ImageFolder未正确统计样本数。
内容的提问来源于stack exchange,提问作者Kanza
相关产品推荐
相关产品推荐

