holmeshoo/beans_sorting
0
1import tensorflow as tf2from tensorflow import keras3from tensorflow.keras import layers4from tensorflow.keras.models import Sequential5from tensorflow.keras.preprocessing.image import ImageDataGenerator6from tensorflow.keras.optimizers import Adam7import japanize_matplotlib8import numpy as np9import matplotlib.pyplot as plt10import os11import json12 13TRAINING_DIR = "./beens/white/training"14VALIDATION_DIR = "./beens/white/validation"15BATCH_SIZE = 1616EPOCHS = 10017IMG_HEIGHT = 8618IMG_WIDTH = 8619class_names = ['broken', 'ng', 'ok', 'purple', 'seed']20 21 22def train_val_generators(TRAINING_DIR, VALIDATION_DIR):23 # 訓練用データジェネレータの設定24 train_datagen = tf.keras.preprocessing.image.ImageDataGenerator(25 rescale=1.0 / 255.0,26 rotation_range=45, # 画像をランダムに回転する回転範囲27 horizontal_flip=True, # 水平方向に入力をランダムに反転します28 )29 30 train_generator = train_datagen.flow_from_directory(31 directory=TRAINING_DIR,32 batch_size=BATCH_SIZE,33 target_size=(IMG_HEIGHT, IMG_WIDTH),34 class_mode="categorical",35 classes=class_names,36 save_format="jpeg"37 )38 39 # 検証用データジェネレータの設定40 validation_datagen = tf.keras.preprocessing.image.ImageDataGenerator(rescale=1.0 / 255.0)41 42 validation_generator = validation_datagen.flow_from_directory(43 directory=VALIDATION_DIR,44 batch_size=BATCH_SIZE,45 target_size=(IMG_HEIGHT, IMG_WIDTH),46 class_mode="categorical",47 classes=class_names,48 save_format="jpeg"49 )50 return train_generator, validation_generator51 52 53train_generator, validation_generator = train_val_generators(TRAINING_DIR, VALIDATION_DIR)54 55 56def create_model():57 # モデルの作成58 model = tf.keras.models.Sequential()59 model.add(tf.keras.layers.Conv2D(32, (3, 3), padding='same', activation='relu', input_shape=(IMG_HEIGHT, IMG_WIDTH, 3)))60 model.add(tf.keras.layers.MaxPooling2D()),61 model.add(tf.keras.layers.Dropout(0.2)),62 model.add(tf.keras.layers.Conv2D(64, 3, padding='same', activation='relu')),63 model.add(tf.keras.layers.MaxPooling2D()),64 model.add(tf.keras.layers.Dropout(0.2)),65 model.add(tf.keras.layers.Conv2D(128, 3, padding='same', activation='relu')),66 model.add(tf.keras.layers.MaxPooling2D()),67 model.add(tf.keras.layers.Dropout(0.2)),68 model.add(tf.keras.layers.Conv2D(256, 3, padding='same', activation='relu')),69 model.add(tf.keras.layers.MaxPooling2D()),70 model.add(tf.keras.layers.Dropout(0.2)),71 model.add(tf.keras.layers.Flatten()),72 model.add(tf.keras.layers.Dense(512, activation='relu')),73 model.add(tf.keras.layers.Dense(len(class_names), name="outputs", activation='softmax'))74 75 model.compile(optimizer=tf.keras.optimizers.Adam(),76 loss='categorical_crossentropy',77 metrics=['accuracy'])78 79 return model80 81 82def training():83 model = create_model()84 history = model.fit(train_generator,85 epochs=EPOCHS,86 verbose=1,87 validation_data=validation_generator)88 return model, history89 90 91def saveModelData(file_path, label):92 if os.path.isfile(file_path):93 os.remove(file_path)94 95 data = {"label": label, "data_form": [IMG_HEIGHT, IMG_WIDTH, "RGB"]}96 97 with open(file_path, "w") as file:98 json.dump(data, file)99 100 101if __name__ == "__main__":102 saveModelData("./model/white_beans.json", class_names)103 model, history = training()104 acc = history.history['accuracy']105 val_acc = history.history['val_accuracy']106 loss = history.history['loss']107 val_loss = history.history['val_loss']108 epochs = range(len(acc))109 plt.plot(epochs, acc, 'r', label='学習用データの正解率')110 plt.plot(epochs, val_acc, 'b', label='評価用データの正解率')111 plt.title('正解率')112 plt.legend()113 plt.figure()114 plt.plot(epochs, loss, 'r', label='学習用データの誤差')115 plt.plot(epochs, val_loss, 'b', label='評価用データの誤差')116 plt.title('誤差')117 plt.legend()118 plt.show()119 model.save('./model/white_beans.h5')120 121# モデルの保存122 