多标签数据集划分训练/验证/测试集时报IndexError索引越界错误
问题描述
我正在开展多标签图像分类任务,需要将自有数据集划分为训练集、测试集与验证集。目前持有一份CSV标注文件,存储了所有图像ID及其对应的类别标签,单张图像可同时归属多个类别。执行数据集划分逻辑时,程序持续抛出如下错误:
Index Error: index 1034 is out of bounds for axis 0 with size 0
对应的实现代码如下:
import argparse import csv import os import numpy as np from PIL import Image from tqdm import tqdm def save_csv(data, path, fieldnames=['image_path', 'airplane', 'bare_soil', 'buildings', 'cars', 'chaparral', 'court', 'dock', 'field', 'grass', 'mobile_home', 'pavement', 'sand', 'sea', 'ship', 'tanks', 'trees', 'water']): with open(path, 'w', newline='') as csv_file: writer = csv.DictWriter(csv_file, fieldnames=fieldnames) writer.writeheader() for row in data: writer.writerow(dict(zip(fieldnames, row))) annotation = 'LandUse_Multilabeled.csv' all_data = [] with open(annotation) as csv_file: # 按CSV格式解析 reader = csv.DictReader(csv_file) # tqdm用于显示进度条 # CSV中每一行对应一张图像 for row in tqdm(reader, total=reader.line_num): # 提取图像ID用于拼接图像文件路径 img_id = row['IMAGE\LABEL'] airplane = row['airplane'] bare_soil = row['bare-soil'] buildings = row['buildings'] cars = row['cars'] chaparral = row['chaparral'] court = row['court'] dock = row['dock'] field = row['field'] grass = row['grass'] mobile_home = row['mobile-home'] pavement = row['pavement'] sand = row['sand'] sea = row['sea'] ship = row['ship'] tanks = row['tanks'] trees = row['trees'] water = row['water'] img_name = os.path.join('/notebooks', 'All_Images', str(img_id) + '.tif') # 检查文件是否存在 if os.path.exists(img_name): # 检查图像尺寸是否为60*80、是否为3通道RGB格式 img = Image.open(img_name) if img.size == (60, 80) and img.mode == "RGB": all_data.append([img_name, airplane, bare_soil, buildings, cars, chaparral, court, dock, field, grass, mobile_home, pavement, sand, sea, ship, tanks, trees, water]) print(all_data) else: print("文件不存在: ", img_name) # 设置随机数种子保证结果可复现 np.random.seed(42) # 将列表转为Numpy数组 all_data = np.asarray(all_data) # 随机采样2100个样本 inds = np.random.choice(2100, replace=False) # 划分训练、验证、测试集并保存为CSV save_csv(all_data[inds][:1470], os.path.join('/notebooks', 'train2.csv')) save_csv(all_data[inds][1471:1680], os.path.join('/notebooks','val2.csv')) save_csv(all_data[inds][1681:2100], os.path.join('/notebooks', 'test2.csv'))
报错原因
报错的核心是all_data最终是空数组(axis 0维度长度为0),此时访问任何非0索引都会触发越界,具体由两个问题导致:
- 硬编码采样总数为2100,但前置逻辑做了两层过滤:检查图像文件是否存在、检查图像尺寸和通道格式,不符合要求的样本都会被剔除,最终
all_data的有效样本长度可能远小于2100,甚至为0。 - CSV字段读取存在命名不匹配风险:代码中写的
row['IMAGE\LABEL']、row['bare-soil']、row['mobile-home']如果和CSV实际表头不一致,会直接触发KeyError,导致没有任何样本被加入all_data。 - 额外问题:划分数据集时索引边界错位,会丢失2个样本。
修复方案
- 先打印
len(all_data)确认过滤后的实际有效样本总数,禁止硬编码采样数量,采样时直接使用有效样本长度作为np.random.choice的入参。 - 核对CSV文件的实际表头,确保代码中读取的字段名和表头完全一致,注意下划线、连字符、斜杠的写法差异。
- 修正数据集划分的索引边界,避免样本丢失。
修复后的核心代码段参考:
# 先统计有效样本数 valid_sample_num = len(all_data) print(f"有效样本总数:{valid_sample_num}") # 按7:1:2比例划分 train_num = int(valid_sample_num * 0.7) val_num = int(valid_sample_num * 0.1) np.random.seed(42) all_data = np.asarray(all_data) # 用实际有效样本数生成随机索引 inds = np.random.choice(valid_sample_num, valid_sample_num, replace=False) # 修正划分边界,避免丢样本 save_csv(all_data[inds][:train_num], os.path.join('/notebooks', 'train2.csv')) save_csv(all_data[inds][train_num:train_num+val_num], os.path.join('/notebooks','val2.csv')) save_csv(all_data[inds][train_num+val_num:], os.path.join('/notebooks', 'test2.csv'))
排查提示:如果打印出的valid_sample_num为0,优先核对CSV表头字段名、图像存储路径、图像尺寸校验规则是否和实际数据匹配。
内容的提问来源于stack exchange,提问作者sss
相关产品推荐
相关产品推荐

