MCPcopy Create free account
hub / github.com/Sirwenhao/Deep-Learning-Notes / main

Function main

CV/data_set/split_data.py:13–65  ·  view source on GitHub ↗
()

Source from the content-addressed store, hash-verified

11
12
13def main():
14 # 保证随机可复现
15 random.seed(0)
16
17 #将数据集中的10%数据划分到验证集中
18 split_rate = 0.1
19
20 #指向解压后的flower_photos文件夹
21 cwd = os.getcwd()
22 data_root = os.path.join(cwd, "flower_data")
23 origin_flower_path = os.path.join(data_root, "flower_photos")
24 assert os.path.exists(origin_flower_path), "path '{}' does not exist.".format(origin_flower_path)
25
26 flower_class = [cla for cla in os.listdir(origin_flower_path)
27 if os.path.isdir(os.path.join(origin_flower_path, cla))]
28
29 # 建立保存训练集的文件夹
30 train_root = os.path.join(data_root, "train")
31 mk_file(train_root)
32 for cla in flower_class:
33 # 建立每个类别对应的文件夹
34 mk_file(os.path.join(train_root, cla))
35
36 # 建立保存验证集的文件夹
37 val_root = os.path.join(data_root, "val")
38 mk_file(val_root)
39 for cla in flower_class:
40
41 # 2022/6/29补充,author:WH
42 # 建立保存验证集中对应每个类别的文件夹
43 mk_file(os.path.join(val_root, cla))
44
45
46 cla_path = os.path.join(origin_flower_path, cla)
47 images = os.listdir(cla_path)
48 num = len(images)
49 # 随机采样验证机的索引
50 eval_index =random.sample(images, k = int(num*split_rate))
51 for index, image in enumerate(images):
52 if image in eval_index:
53 # 将分配至验证集中文件复制到相应的目录
54 image_path = os.path.join(cla_path, image)
55 new_path = os.path.join(val_root, cla)
56 copy(image_path, new_path)
57 else:
58 # 将分配至训练集中的文件复制到相应的目录
59 image_path = os.path.join(cla_path, image)
60 new_path = os.path.join(train_root, cla)
61 copy(image_path, new_path)
62 print("\r[{}] processing [{}/{}]".format(cla, index+1, num), end = "")
63 print()
64
65 print("processing done!")
66
67if __name__ == '__main__':
68 main()

Callers 1

split_data.pyFile · 0.70

Calls 1

mk_fileFunction · 0.85

Tested by

no test coverage detected