diff --git a/scripts/dataset_validation.py b/scripts/dataset_validation.py index 7cc6a9e..ec2279b 100644 --- a/scripts/dataset_validation.py +++ b/scripts/dataset_validation.py @@ -1,5 +1,6 @@ -import os import cv2 +import os +import random from PIL import Image def rename_files(directory_path: str): @@ -56,11 +57,33 @@ def resize_to_height(dir_path: str, target_height=48): new_width = int(float(img.width) * width_percent) resized_img = img.resize((new_width, target_height), Image.LANZOS) 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/") # check_images("train_data/images/") # check_labels("train_data/images/") # check_symbols("train_data/") # max_height("train_data/images") +split_dataset("train_data/")