Разделение датасета на тренировочные и валидационные данные.
This commit is contained in:
@@ -1,5 +1,6 @@
|
|||||||
import os
|
|
||||||
import cv2
|
import cv2
|
||||||
|
import os
|
||||||
|
import random
|
||||||
from PIL import Image
|
from PIL import Image
|
||||||
|
|
||||||
def rename_files(directory_path: str):
|
def rename_files(directory_path: str):
|
||||||
@@ -57,10 +58,32 @@ def resize_to_height(dir_path: str, target_height=48):
|
|||||||
resized_img = img.resize((new_width, target_height), Image.LANZOS)
|
resized_img = img.resize((new_width, target_height), Image.LANZOS)
|
||||||
resized_img.save(os.path.join(dir_path, filename))
|
resized_img.save(os.path.join(dir_path, filename))
|
||||||
|
|
||||||
|
def split_dataset(dir_path: str):
|
||||||
|
images_dir = os.path.join(dir_path, "images")
|
||||||
|
filenames = os.listdir(images_dir)
|
||||||
|
dataset_length = len(filenames)
|
||||||
|
random.shuffle(filenames)
|
||||||
|
train_ratio = 0.8
|
||||||
|
val_ratio = 0.2
|
||||||
|
train_files_len = round(dataset_length * train_ratio)
|
||||||
|
val_files_len = round(dataset_length * val_ratio)
|
||||||
|
train_files = filenames[:train_files_len]
|
||||||
|
val_files = filenames[train_files_len:train_files_len+val_files_len]
|
||||||
|
train_file_path = os.path.join(dir_path, "train.txt")
|
||||||
|
val_file_path = os.path.join(dir_path, "val.txt")
|
||||||
|
with open(train_file_path, "w") as train_file:
|
||||||
|
for filename in train_files:
|
||||||
|
label, _ = filename.split(".")
|
||||||
|
train_file.write(f"{filename}\t{label}\n")
|
||||||
|
|
||||||
|
with open(val_file_path, "w") as val_file:
|
||||||
|
for filename in val_files:
|
||||||
|
label, _ = filename.split(".")
|
||||||
|
val_file.write(f"{filename}\t{label}\n")
|
||||||
|
|
||||||
# rename_files("train_data/images/")
|
# rename_files("train_data/images/")
|
||||||
# check_images("train_data/images/")
|
# check_images("train_data/images/")
|
||||||
# check_labels("train_data/images/")
|
# check_labels("train_data/images/")
|
||||||
# check_symbols("train_data/")
|
# check_symbols("train_data/")
|
||||||
# max_height("train_data/images")
|
# max_height("train_data/images")
|
||||||
|
split_dataset("train_data/")
|
||||||
|
|||||||
Reference in New Issue
Block a user