如何将自定义数据集转换为MNIST格式并适配模型解决IndexError报错
错误根因
这个错误的触发原因和你写的数据变换代码没有直接关系,问题出在标签处理逻辑:torch.eye(10)会生成一个10行的单位矩阵,仅支持传入09范围内的索引值做行选择,你自有数据集中的标签存在超出09的取值,调用index_select时就会触发索引越界。
修复方案
- 首先确认自有数据集的标签分布,运行以下代码遍历数据集,确认标签的取值范围和分类数量:
# 替换your_dataset为你自己定义的数据集实例 label_list = [] for img, label in your_dataset: label_list.append(label) print(f"标签最小值: {min(label_list)}, 标签最大值: {max(label_list)}") print(f"总分类数: {len(set(label_list))}")
- 根据上面的输出结果做对应修改:
- 如果你的数据集分类数刚好是10,但标签是从1开始计数(取值110):只需要在加载标签时统一减1,把标签映射到09区间即可。
- 如果你的数据集分类数不等于10:若坚持不修改模型结构,需要对数据集做筛选/合并,把分类数调整为10,保证标签落在0~9区间;如果可以接受少量代码调整,只需要把独热编码生成代码里的
torch.eye(10)改成对应你实际分类数的torch.eye(N)(N为你的分类数),同时修改模型最后输出层的维度为N即可,不需要调整整体架构。
- 可选优化:你当前的数据变换顺序存在冗余,
ToTensor()操作会把PIL图像转为张量,应该放在所有图像维度变换操作之后,调整后的变换逻辑运行效率更高:
data_transform_test = transforms.Compose([ transforms.Grayscale(num_output_channels=1), transforms.Resize((28, 28)), transforms.ToTensor(), transforms.Normalize((0.5,), (0.5,)) ])
内容的提问来源于stack exchange,提问作者autopilot38
相关产品推荐
相关产品推荐

