CoolFace
Apppublic

holmeshoo/beans_sorting

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
traing_white.py122 linesDownload Raw Back to src
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