{"cells":[{"metadata":{},"cell_type":"markdown","source":"# Cassava Leaf Disease Classification"},{"metadata":{"trusted":true},"cell_type":"code","source":"from tensorflow.keras import models, layers\nfrom tensorflow import keras\nfrom tensorflow.keras.applications import ResNet50, DenseNet121, EfficientNetB0\nimport tensorflow as tf\nimport tensorflow_hub as hub\nfrom tensorflow.keras import layers\n# tf.enable_eager_execution()\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport tensorflow_hub as hub","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"TRAIN_DIR = '../input/cassava-leaf-disease-classification/train_images/'\nlabels = pd.read_csv(\"../input/cassava-leaf-disease-classification/train.csv\")\nlabels.label = labels.label.astype('str')\n\nlabels.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"BATCH_SIZE = 16\nIMAGE_SIZE = [512, 512]\nAUTOTUNE = tf.data.experimental.AUTOTUNE\n\nBATCH_SIZE = 8\nSTEPS_PER_EPOCH = len(labels)*0.8 / BATCH_SIZE\nVALIDATION_STEPS = len(labels)*0.2 / BATCH_SIZE\nEPOCHS = 50\nTARGET_SIZE = 512","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"generator_train = keras.preprocessing.image.ImageDataGenerator(rotation_range=90,\n                                                               shear_range=0.2, \n                                                               zoom_range=0.2, \n                                                               horizontal_flip=True,\n                                                               vertical_flip=True,\n                                                               validation_split=0.2)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_gen = generator_train.flow_from_dataframe(labels,\n                                          directory = TRAIN_DIR,\n                                          subset='training',\n                                          x_col = \"image_id\",\n                                          y_col = \"label\",\n                                          batch_size = BATCH_SIZE,\n                                          class_mode = \"sparse\",\n                                          shuffle=True)\nval_gen   = generator_train.flow_from_dataframe(labels,\n                                          directory = TRAIN_DIR,\n                                          batch_size = BATCH_SIZE,\n                                          x_col = \"image_id\",\n                                          y_col = \"label\",\n                                          class_mode = \"sparse\",\n                                          subset='validation')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# from internet\naug_images = [train_gen[0][0][0]/255 for i in range(10)]\nfig, axes = plt.subplots(2, 5, figsize = (20, 10))\naxes = axes.flatten()\nfor img, ax in zip(aug_images, axes):\n    ax.imshow(img)\n    ax.axis('off')\nplt.tight_layout()\nplt.show()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"num_classes = 5\nimg_height = img_width = 512","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"feature_extractor_model1 = \"https://tfhub.dev/google/imagenet/resnet_v2_50/feature_vector/4\"\nfeature_extractor_model2 = 'https://tfhub.dev/tensorflow/efficientnet/b7/feature-vector/1'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# first not trainable\nfeature_extractor_layer1 = hub.KerasLayer(\n    feature_extractor_model1, input_shape=(img_width, img_height, 3), trainable=False)\n\nfeature_extractor_layer2 = hub.KerasLayer(\n    feature_extractor_model1, input_shape=(img_width, img_height, 3), trainable=True)\n\nfeature_extractor_layer3 = hub.KerasLayer(\n    feature_extractor_model2, input_shape=(img_width, img_height, 3), trainable=True)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"resnet_not = tf.keras.Sequential([\n  layers.experimental.preprocessing.Rescaling(1./255, input_shape=(img_height, img_width, 3)),\n  feature_extractor_layer1,\n  tf.keras.layers.Dense(100, activation = \"relu\"),  \n  tf.keras.layers.Dropout(0.2),  \n  tf.keras.layers.Dense(100, activation = \"relu\"),  \n  tf.keras.layers.Dropout(0.5),  \n  tf.keras.layers.Dense(5, activation = \"softmax\")\n])\n\nresnet_not.compile(\n  optimizer=tf.keras.optimizers.Adam(learning_rate = 0.0001),\n  loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),\n  metrics=['acc'])\n\nresnet = tf.keras.Sequential([\n  layers.experimental.preprocessing.Rescaling(1./255, input_shape=(img_height, img_width, 3)),\n  feature_extractor_layer2,\n  tf.keras.layers.Dense(100, activation = \"relu\"),  \n  tf.keras.layers.Dropout(0.2),  \n  tf.keras.layers.Dense(5, activation = \"softmax\")\n])\n\nresnet.compile(\n  optimizer=tf.keras.optimizers.Adam(learning_rate = 0.0001),\n  loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),\n  metrics=['acc'])\n\nefficient = tf.keras.Sequential([\n  layers.experimental.preprocessing.Rescaling(1./255, input_shape=(img_height, img_width, 3)),\n  feature_extractor_layer3,\n  tf.keras.layers.Dense(100, activation = \"relu\"),  \n  tf.keras.layers.Dropout(0.2),  \n  tf.keras.layers.Dense(5, activation = \"softmax\")\n])\n\nefficient.compile(\n  optimizer=tf.keras.optimizers.Adam(learning_rate = 0.0001),\n  loss=tf.keras.losses.SparseCategoricalCrossentropy(from_logits=True),\n  metrics=['acc'])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"# from internet\nmodel1_save = tf.keras.callbacks.ModelCheckpoint('./resnet_not_trained_v2_50.h5', \n                             save_best_only = True, \n                             save_weights_only = True,\n                             monitor = 'val_loss', \n                             mode = 'min', verbose = 1)\n\nmodel2_save = tf.keras.callbacks.ModelCheckpoint('./resnet_trained_v2_50.h5', \n                             save_best_only = True, \n                             save_weights_only = True,\n                             monitor = 'val_loss', \n                             mode = 'min', verbose = 1)\n\n\nmodel3_save = tf.keras.callbacks.ModelCheckpoint('./efficientnet.h5', \n                             save_best_only = True, \n                             save_weights_only = True,\n                             monitor = 'val_loss', \n                             mode = 'min', verbose = 1)\n\n\n\n\nearly_stop = tf.keras.callbacks.EarlyStopping(monitor = 'val_loss', min_delta = 0.001, \n                           patience = 5, mode = 'min', verbose = 1,\n                           restore_best_weights = True)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# ResNet2 non trainable"},{"metadata":{"trusted":true},"cell_type":"code","source":"try: \n    with tf.device('/gpu:0'):\n        history_resnet_not = resnet_not.fit(train_gen,\n                                            validation_data = val_gen,\n                                            steps_per_epoch = 1800,\n                                            validation_steps = VALIDATION_STEPS,\n                                            epochs = EPOCHS,\n                                            callbacks=[model1_save, early_stop])\nexcept RuntimeError as e:\n    print(e)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"plt.plot(history_resnet_not.history['acc'])\nplt.plot(history_resnet_not.history['val_acc'])\nplt.title('model accuracy')\nplt.ylabel('accuracy')\nplt.xlabel('epoch')\nplt.legend(['train', 'val'], loc='upper left')\nplt.show()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# ResNet2 trainable"},{"metadata":{"trusted":true},"cell_type":"code","source":"try: \n    with tf.device('/gpu:0'):\n        history_resnet = resnet.fit(train_gen,\n                                            validation_data = val_gen,\n                                            steps_per_epoch = 1800,\n                                            validation_steps = VALIDATION_STEPS,\n                                            epochs = EPOCHS,\n                                            callbacks=[model2_save, early_stop])\nexcept RuntimeError as e:\n    print(e)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"plt.plot(history_resnet.history['acc'])\nplt.plot(history_resnet.history['val_acc'])\nplt.title('model accuracy')\nplt.ylabel('accuracy')\nplt.xlabel('epoch')\nplt.legend(['train', 'val'], loc='upper left')\nplt.show()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"# EfficientNet trainable"},{"metadata":{"trusted":true},"cell_type":"code","source":"#try: \n#    with tf.device('/gpu:0'):\n#        history_efficient = efficient.fit(train_gen,\n#                                        validation_data = val_gen,\n#                                        steps_per_epoch = 1800,\n#                                        validation_steps = VALIDATION_STEPS,\n#                                        epochs = EPOCHS,\n#                                        callbacks=[model3_save, early_stop])\n#except RuntimeError as e:\n#    print(e)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"#plt.plot(history_efficient.history['acc'])\n#plt.plot(history_efficient.history['val_acc'])\n#plt.title('model accuracy')\n#plt.ylabel('accuracy')\n#plt.xlabel('epoch')\n#plt.legend(['train', 'val'], loc='upper left')\n#plt.show()","execution_count":null,"outputs":[]}],"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat":4,"nbformat_minor":4}