"""
拆分数据集为 训练集、测试集、验证集
"""
import os
import random
# 一 . 设置数据集比例
# 训练集的比率
train_percent = 0.8
# 测试集占测试验证集的百分比
test_other_percent = 0.5
# 二 . 获取训练集 , 验证集和测试集的索引列表
# 标注数据地址
voc_annotations_path = 'Annotations'
# 划分数据集的文件位置
division_data_path = 'ImageSets/Main'
# 获取所有的 标注数据 名称列表
voc_annotations_list = os.listdir(voc_annotations_path)
# 所有标注数据的个数
voc_annotations_cnt = len(voc_annotations_list)
# 生成 文件个数大小 的范围 , 可以看成索引列表
list_range = range(voc_annotations_cnt)
# 获取训练集的个数
train_cnt = int(voc_annotations_cnt * train_percent)
# 从文件中随机获取 train_cnt 个训练验证集的索引
train_index = random.sample(list_range, train_cnt)
# 从文件中随机获取训练验证集的索引 ( 全部索引与训练集做差集 )
train_val_index = list(set(list(list_range)).difference(set(train_index)))
# 计算需要获取的测试训练集的个数
test_val_cnt =外汇跟单gendan5.com voc_annotations_cnt - train_cnt
# 计算测试集个数
test_cnt = int(test_val_cnt * test_other_percent)
# 计算验证集个数
val_cnt = test_val_cnt - test_cnt
# 测试集索引列表
test_index = random.sample(train_val_index, test_cnt)
# 验证集索引列表
val_index = list(set(train_val_index).difference(set(test_index)))
# 三 . 将各个数据集名称写入到文件中
# 训练集
train_object = open('%s/train.txt' % division_data_path, 'w')
# 测试集
test_object = open('%s/test.txt' % division_data_path, 'w')
# 验证集
val_object = open('%s/val.txt' % division_data_path, 'w')
train_names = [voc_annotations_list[i][:-4] + '\n' for i in train_index]
train_object.writelines(train_names)
train_object.close()
test_names = [voc_annotations_list[i][:-4] + '\n' for i in test_index]
test_object.writelines(test_names)
test_object.close()
val_names = [voc_annotations_list[i][:-4] + '\n' for i in val_index]
val_object.writelines(val_names)
val_object.close()