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

多标签数据集划分训练/验证/测试集时报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个样本。
修复方案
  1. 先打印len(all_data)确认过滤后的实际有效样本总数,禁止硬编码采样数量,采样时直接使用有效样本长度作为np.random.choice的入参。
  2. 核对CSV文件的实际表头,确保代码中读取的字段名和表头完全一致,注意下划线、连字符、斜杠的写法差异。
  3. 修正数据集划分的索引边界,避免样本丢失。

修复后的核心代码段参考:

# 先统计有效样本数
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.26 11:48:16