如何均匀划分COCO数据集?解决按1/3比例分划的代码问题
问题
我下载了完整的COCO数据集,需要修改instances_train2017.json文件实现以下需求:
- 将完整训练集按所有类别均匀分成1/3
- 例如,若
class_1有100条标注,修改后的标注文件需保留该类约33条(100/3)的标注及对应图像
我写了如下代码,但运行耗时且生成结果错误:
import json from collections import defaultdict import random # Path to your local COCO-format JSON annotation file original_annotation_file = 'coco/annotations/instances_train2017.json' output_annotation_file = 'evenly.json' # Load the local COCO-format dataset from your JSON file with open(original_annotation_file, 'r') as f: coco_data = json.load(f) class_counts = defaultdict(int) target_class_counts = defaultdict(int) # Calculate the target count for each class for ann in coco_data['annotations']: class_id = ann['category_id'] class_counts[class_id] += 1 for class_id, count in class_counts.items(): target_class_counts[class_id] = count // 3 # Create a list to hold the selected annotations selected_annotations = [] # Iterate through the annotations and select the subset for ann in coco_data['annotations']: class_id = ann['category_id'] # Only include this annotation if we haven't reached the target count for this class if class_counts[class_id] <= target_class_counts[class_id]: selected_annotations.append(ann) # Update the count for this class class_counts[class_id] += 1 # Create a new COCO-format JSON data structure for the subset subset_data = { 'info': coco_data['info'], 'licenses': coco_data['licenses'], 'categories': coco_data['categories'], 'images': coco_data['images'], 'annotations': selected_annotations } # Shuffle the selected annotations to mix up the classes if desired random.shuffle(subset_data['annotations']) # Write the subset data to a new JSON file with open(output_annotation_file, 'w') as f: json.dump(subset_data, f)
代码问题分析
- 逻辑判断完全错误:原代码用初始为该类总标注数的
class_counts和目标数target_class_counts比较,总标注数远大于目标数,导致几乎不会选中任何标注,完全不符合需求。 - 未处理冗余图像:保留了所有原始图像,但很多图像没有被选中的标注,不符合COCO数据集规范。
- 效率低下:逐个遍历所有标注判断是否选中,不如按类别分组后直接抽取高效。
修正后的代码
import json from collections import defaultdict import random # 路径配置 original_annotation_file = 'coco/annotations/instances_train2017.json' output_annotation_file = 'evenly_subset.json' # 加载原始标注数据 with open(original_annotation_file, 'r') as f: coco_data = json.load(f) # 按类别分组所有标注 annotations_by_class = defaultdict(list) for ann in coco_data['annotations']: annotations_by_class[ann['category_id']].append(ann) # 对每个类别随机抽取1/3的标注 selected_annotations = [] for class_id, ann_list in annotations_by_class.items(): target_num = len(ann_list) // 3 # 随机抽取目标数量的标注,保证类别比例均匀 selected = random.sample(ann_list, target_num) selected_annotations.extend(selected) # 收集所有选中标注对应的图像ID,过滤无标注的图像 selected_image_ids = {ann['image_id'] for ann in selected_annotations} selected_images = [img for img in coco_data['images'] if img['id'] in selected_image_ids] # 构建符合COCO格式的子集数据 subset_data = { 'info': coco_data['info'], 'licenses': coco_data['licenses'], 'categories': coco_data['categories'], 'images': selected_images, 'annotations': selected_annotations } # 可选:打乱标注顺序,混合不同类别 random.shuffle(subset_data['annotations']) # 保存结果到新文件 with open(output_annotation_file, 'w') as f: json.dump(subset_data, f, indent=2)
修正说明
- 按类别分组抽取:先将标注按类别分组,再用
random.sample直接抽取目标数量,既保证每个类别严格按1/3比例保留,又提升了运行效率。 - 过滤冗余图像:只保留包含选中标注的图像,避免子集里出现无标注的无效图像。
- 逻辑修正:替换了原代码错误的判断逻辑,确保选中的标注数量符合预期。
内容的提问来源于stack exchange,提问作者mrconcerned
相关产品推荐
相关产品推荐

