holmeshoo/beans_sorting
0
1import os2from sklearn.model_selection import train_test_split3import shutil4import numpy as np5import sys6 7 8def split_images(src_dir, training_dest, validation_dest, train_ratio=0.8):9 # 画像ファイルのリストを取得10 image_files = [f for f in os.listdir(src_dir) if f.endswith(('.jpg', '.jpeg', '.png', '.gif'))]11 12 # 画像ファイルをトレーニングセットとバリデーションセットに分割13 train_files, test_files = train_test_split(image_files, train_size=train_ratio, random_state=None)14 15 # トレーニングデータのディレクトリを作成16 if os.path.isdir(training_dest):17 shutil.rmtree(training_dest)18 os.makedirs(training_dest, exist_ok=True)19 20 # トレーニングデータをコピー21 for file in train_files:22 src_path = os.path.join(src_dir, file)23 dst_path = os.path.join(training_dest, file)24 shutil.copy(src_path, dst_path)25 26 # バリデーションデータのディレクトリを作成27 if os.path.isdir(validation_dest):28 shutil.rmtree(validation_dest)29 os.makedirs(validation_dest, exist_ok=True)30 31 # バリデーションデータをコピー32 for file in test_files:33 src_path = os.path.join(src_dir, file)34 dst_path = os.path.join(validation_dest, file)35 shutil.copy(src_path, dst_path)36 37 print(f"Training data num: {len(train_files)}, Validation data num: {len(test_files)}")38 39 40if __name__ == "__main__":41 args = sys.argv42 src_path = args[1]43 training_destination = args[2]44 validation_destination = args[3]45 ratio = float(args[4])46 47 # サブディレクトリのリストを取得してループ48 files_dir = [f for f in os.listdir(src_path) if os.path.isdir(os.path.join(src_path, f))]49 for folder in files_dir:50 print(folder)51 split_images(52 os.path.join(src_path, folder),53 os.path.join(training_destination, folder),54 os.path.join(validation_destination, folder),55 ratio56 )57 58# コマンドラインの例: python ./split_train_test.py ./beens/black/data ./beens/black/training ./beens/black/validation 0.859 