vaishanthr/Image-Classifier-TensorFlow
1
1import gradio as gr2import tensorflow as tf3from tensorflow import keras4from custom_model import ImageClassifier5from resnet_model import ResNetClassifier6from vgg16_model import VGG16Classifier7from inception_v3_model import InceptionV3Classifier8from mobilevet_v2 import MobileNetClassifier9import os10 11CLASS_NAMES =['Airplane', 'Automobile', 'Bird', 'Cat', 'Deer', 'Dog', 'Frog', 'Horse', 'Ship', 'Truck']12 13# models14custom_model = ImageClassifier()15custom_model.load_model("image_classifier_model.h5")16resnet_model = ResNetClassifier()17vgg16_model = VGG16Classifier()18inceptionV3_model = InceptionV3Classifier()19mobilenet_model = MobileNetClassifier()20 21def make_prediction(image, model_type="CNN (Custom)"):22 if "CNN (Custom)" == model_type:23 top_classes, top_probs = custom_model.classify_image(image, top_k=3)24 return {CLASS_NAMES[cls_id]:str(prob) for cls_id, prob in zip(top_classes, top_probs)}25 elif "ResNet50" == model_type:26 predictions = resnet_model.classify_image(image)27 return {class_name:str(prob) for _, class_name, prob in predictions}28 elif "VGG16" == model_type:29 predictions = vgg16_model.classify_image(image)30 return {class_name:str(prob) for _, class_name, prob in predictions}31 elif "Inception v3" == model_type:32 predictions = inceptionV3_model.classify_image(image)33 return {class_name:str(prob) for _, class_name, prob in predictions}34 elif "Mobile Net v2" == model_type:35 predictions = mobilenet_model.classify_image(image)36 return {class_name:str(prob) for _, class_name, prob in predictions}37 else:38 return {"Select a model to classify image"}39 40def train_model(epochs, batch_size, validation_split):41 42 print("Training model")43 44 # Create an instance of the ImageClassifier45 classifier = ImageClassifier()46 47 # Load the dataset48 (x_train, y_train), (x_test, y_test) = classifier.load_dataset()49 50 # Build and train the model51 classifier.build_model(x_train)52 classifier.train_model(x_train, y_train, batch_size=int(batch_size), epochs=int(epochs), validation_split=float(validation_split))53 54 # Evaluate the model55 classifier.evaluate_model(x_test, y_test)56 57 # Save the trained model58 print("Saving model ...")59 classifier.save_model("image_classifier_model.h5")60 61 custom_model = classifier62 63 64def update_train_param_display(model_type):65 if "CNN (Custom)" == model_type:66 return [gr.update(visible=True), gr.update(visible=False)]67 return [gr.update(visible=False), gr.update(visible=True)]68 69if __name__ == "__main__":70 # gradio gui app71 with gr.Blocks() as my_app:72 gr.Markdown("<h1><center>Image Classification using TensorFlow</center></h1>")73 gr.Markdown("<h3><center>This model classifies image using different models.</center></h3>")74 75 with gr.Row():76 with gr.Column(scale=1):77 img_input = gr.Image()78 model_type = gr.Dropdown(79 ["CNN (Custom)", 80 "ResNet50", 81 "VGG16",82 "Inception v3",83 "Mobile Net v2"], 84 label="Model Type", value="CNN (Custom)",85 info="Select the inference model before running predictions!")86 87 with gr.Column() as train_col:88 gr.Markdown("Train Parameters")89 with gr.Row():90 epochs_inp = gr.Textbox(label="Epochs", value="10")91 validation_split = gr.Textbox(label="Validation Split", value="0.1")92 93 with gr.Row():94 batch_size = gr.Textbox(label="Batch Size", value="64")95 96 with gr.Row():97 train_btn = gr.Button(value="Train") 98 predict_btn_1 = gr.Button(value="Predict") 99 100 with gr.Column(visible=False) as no_train_col:101 predict_btn_2 = gr.Button(value="Predict") 102 103 with gr.Column(scale=1):104 output_label = gr.Label()105 106 gr.Markdown("## Sample Images")107 gr.Examples(108 examples=[os.path.join(os.path.dirname(__file__), "assets/dog_2.jpg"),109 os.path.join(os.path.dirname(__file__), "assets/truck.jpg"),110 os.path.join(os.path.dirname(__file__), "assets/car.jpg"),111 os.path.join(os.path.dirname(__file__), "assets/car_32x32.jpg")112 ],113 inputs=img_input,114 outputs=output_label,115 fn=make_prediction,116 cache_examples=True,117 )118 119 120 121 # app logic122 predict_btn_1.click(make_prediction, inputs=[img_input, model_type], outputs=[output_label])123 predict_btn_2.click(make_prediction, inputs=[img_input, model_type], outputs=[output_label])124 model_type.change(update_train_param_display, inputs=model_type, outputs=[train_col, no_train_col])125 train_btn.click(train_model, inputs=[epochs_inp, batch_size, validation_split], outputs=[])126 127my_app.queue(concurrency_count=5, max_size=20).launch(debug=True)