{"cells":[{"metadata":{"id":"MXT7mxydn5x1","trusted":true},"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport matplotlib.pyplot as plt\nimport cv2\nimport tensorflow as tf\nfrom tensorflow import keras","execution_count":null,"outputs":[]},{"metadata":{"id":"2EqnKuq7r9_R","trusted":true},"cell_type":"code","source":"keras.backend.clear_session()\ntf.random.set_seed(42)\nnp.random.seed(42)","execution_count":null,"outputs":[]},{"metadata":{"id":"7Gzav3H6W-sX","outputId":"8588d6e4-dd10-4dac-9f02-584d9ddf5e99","trusted":true},"cell_type":"code","source":"PATH =  '../input/cassava-leaf-disease-classification/'\nos.listdir(PATH)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"id":"UWFt-_AQWiQc","outputId":"f409083a-9254-4920-97a3-906054678195"},"cell_type":"code","source":"print('Train images: %d' %len(os.listdir(os.path.join(PATH, 'train_images'))))","execution_count":null,"outputs":[]},{"metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","trusted":true,"id":"EJZCEKeIWiQd"},"cell_type":"code","source":"train_full_df=pd.read_csv(PATH + 'train.csv')\nlabel_js=pd.read_json(PATH + 'label_num_to_disease_map.json', typ='series')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"id":"Pp1WfTdhWiQd"},"cell_type":"code","source":"train_full_df[\"label\"] = train_full_df[\"label\"].astype(\"str\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"id":"R5fTpMQ8WiQd","outputId":"8b9b90c9-c595-4bd5-f30e-e30a958b5f8f"},"cell_type":"code","source":"train_full_df.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"id":"8uNrp4QUWiQd","outputId":"ae88310e-2772-4b20-c95f-ad9cfe8b73f7"},"cell_type":"code","source":"label_js","execution_count":null,"outputs":[]},{"metadata":{"id":"wNVoBKxZWiQe"},"cell_type":"markdown","source":"## Cassava Bacterial Blight"},{"metadata":{"trusted":true,"id":"Fi1OWsO0WiQf","outputId":"82d53646-cc3f-4f8d-ffb9-113ef7df6223"},"cell_type":"code","source":"samples = train_full_df[train_full_df.label == \"0\"].sample(6)\n\nplt.figure(figsize=(16, 8))\n\nfor ind, (image_id, label) in enumerate(zip(samples.image_id, samples.label)):\n    plt.subplot(2, 3, ind + 1)\n    image = cv2.imread(os.path.join(PATH, \"train_images\", image_id))\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    plt.imshow(image)\n    \nplt.show()","execution_count":null,"outputs":[]},{"metadata":{"id":"kQ4dReTlWiQf"},"cell_type":"markdown","source":"## Cassava Brown Streak Disease"},{"metadata":{"trusted":true,"id":"5XrfaU_LWiQf","outputId":"f3e009dc-9936-4f59-ed1d-739f1ea287cd"},"cell_type":"code","source":"samples = train_full_df[train_full_df.label == \"1\"].sample(6)\n\nplt.figure(figsize=(16, 8))\n\nfor ind, (image_id, label) in enumerate(zip(samples.image_id, samples.label)):\n    plt.subplot(2, 3, ind + 1)\n    image = cv2.imread(os.path.join(PATH, \"train_images\", image_id))\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    plt.imshow(image)\n    \nplt.show()","execution_count":null,"outputs":[]},{"metadata":{"id":"EWaJqLRDWiQg"},"cell_type":"markdown","source":"## Cassava Green Mottle"},{"metadata":{"trusted":true,"id":"XdoEJxE5WiQg","outputId":"dc80394b-c684-43b7-8afe-47c708920626"},"cell_type":"code","source":"samples = train_full_df[train_full_df.label == \"2\"].sample(6)\n\nplt.figure(figsize=(16, 8))\n\nfor ind, (image_id, label) in enumerate(zip(samples.image_id, samples.label)):\n    plt.subplot(2, 3, ind + 1)\n    image = cv2.imread(os.path.join(PATH, \"train_images\", image_id))\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    plt.imshow(image)\n    \nplt.show()","execution_count":null,"outputs":[]},{"metadata":{"id":"TbWz2V8-WiQg"},"cell_type":"markdown","source":"## Cassava Mosaic Disease"},{"metadata":{"trusted":true,"id":"OxmnRKQvWiQh","outputId":"3b587084-1c69-4c73-9499-df4f7c42de5b"},"cell_type":"code","source":"samples = train_full_df[train_full_df.label == \"3\"].sample(6)\n\nplt.figure(figsize=(16, 8))\n\nfor ind, (image_id, label) in enumerate(zip(samples.image_id, samples.label)):\n    plt.subplot(2, 3, ind + 1)\n    image = cv2.imread(os.path.join(PATH, \"train_images\", image_id))\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    plt.imshow(image)\n    \nplt.show()","execution_count":null,"outputs":[]},{"metadata":{"id":"ISpqqvrHWiQh"},"cell_type":"markdown","source":"## Healthy"},{"metadata":{"trusted":true,"id":"4MeUi5kjWiQi","outputId":"f9074ff2-17c0-41ad-eed3-792bf09fda6a"},"cell_type":"code","source":"samples = train_full_df[train_full_df.label == \"4\"].sample(6)\n\nplt.figure(figsize=(16, 8))\n\nfor ind, (image_id, label) in enumerate(zip(samples.image_id, samples.label)):\n    plt.subplot(2, 3, ind + 1)\n    image = cv2.imread(os.path.join(PATH, \"train_images\", image_id))\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    plt.imshow(image)\n    \nplt.show()","execution_count":null,"outputs":[]},{"metadata":{"id":"lhTXshOxWiQi"},"cell_type":"markdown","source":"## Load Data & Modeling"},{"metadata":{"trusted":true,"id":"dHzgprp1WiQi"},"cell_type":"code","source":"train_size = 0.8\n\ntrain_samples_count = int(len(train_full_df) * train_size)\nvalidation_item_count = len(train_full_df) - train_samples_count \n\ntrain_df = train_full_df[:train_samples_count]\nvalid_df = train_full_df[train_samples_count:]","execution_count":null,"outputs":[]},{"metadata":{"id":"5yPvbelBdbsk","trusted":true},"cell_type":"code","source":"train_batch_size = 8\nvalid_batch_size = 8\n\nimage_size = 512\n\ntarget_size = (image_size, image_size)","execution_count":null,"outputs":[]},{"metadata":{"id":"ez2u7PJA5vZj","outputId":"6d46a9af-1300-4b86-f86c-9516b62ad2c2","trusted":true},"cell_type":"code","source":"!pip install -q git+https://github.com/albu/albumentations --no-cache-dir","execution_count":null,"outputs":[]},{"metadata":{"id":"tn0vw_qEeZS5","trusted":true},"cell_type":"code","source":"from albumentations import (Compose, RandomCrop,Transpose, HorizontalFlip, \n                            VerticalFlip, ShiftScaleRotate, HueSaturationValue,   \n                            RandomBrightness,RandomContrast, CenterCrop, ToFloat, Normalize)\n\ntrain_transforms = Compose([\n            RandomCrop(image_size, image_size),\n            HorizontalFlip(p=0.5),\n            VerticalFlip(p=0.5),\n            ShiftScaleRotate(p=0.5),\n            RandomBrightness(limit=0.1, p=0.5),\n            HueSaturationValue(hue_shift_limit=20, sat_shift_limit=0.2, val_shift_limit=0.2, p=0.5),\n            RandomContrast(limit=0.2, p=0.5),\n            Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0)\n        ])\n\nvalid_transforms = Compose([\n            CenterCrop(image_size, image_size),\n            Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225], max_pixel_value=255.0, p=1.0)\n        ])","execution_count":null,"outputs":[]},{"metadata":{"id":"rhARrvy8Ar-r","outputId":"3d2833c8-a59e-4756-9cad-44433b8ccd3f","trusted":true},"cell_type":"code","source":"!pip install -q git+https://github.com/mjkvaak/ImageDataAugmentor","execution_count":null,"outputs":[]},{"metadata":{"id":"dok-VfLVlm0G","outputId":"b3edcfce-53e3-4d68-cd77-53981280bac0","trusted":true},"cell_type":"code","source":"from ImageDataAugmentor.image_data_augmentor import *\n\ntrain_datagen = ImageDataAugmentor(augment=train_transforms)\nvalid_datagen = ImageDataAugmentor(augment=valid_transforms)\n\ntrain_generator = train_datagen.flow_from_dataframe(\n    dataframe = train_df,\n    x_col='image_id',\n    y_col='label',\n    directory=PATH + 'train_images/',\n    target_size=target_size,\n    batch_size=train_batch_size,\n    shuffle=True,\n    class_mode='categorical')\n\nvalid_generator = valid_datagen.flow_from_dataframe(\n    dataframe = valid_df,\n    x_col='image_id',\n    y_col='label',\n    directory=PATH + 'train_images/',\n    target_size=target_size,\n    shuffle=False,\n    batch_size=valid_batch_size,\n    class_mode='categorical')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"id":"eh9OP71nWiQk","outputId":"9f3117d9-8a43-4f85-f7fa-1fa6a982749f"},"cell_type":"code","source":"from tensorflow.keras.applications import EfficientNetB4\nfrom keras.models import Model\nfrom keras.layers import Input, Dense, Dropout, GlobalAveragePooling2D, BatchNormalization\n\n\nefficientnet = EfficientNetB4(weights=\"imagenet\", include_top=False,\n                                  drop_connect_rate=0.2, input_shape=(image_size, image_size, 3))\n\n#efficientnet.trainable = False\n\ninputs = Input(shape=(image_size, image_size, 3))\nefficientnet = efficientnet(inputs, training=False)\npooling = GlobalAveragePooling2D()(efficientnet)\nnormalization = BatchNormalization()(pooling)\ndropout = Dropout(0.4)(normalization)\noutputs = Dense(5, activation=\"softmax\")(dropout)\nmodel = Model(inputs=inputs, outputs=outputs)\n    \nmodel.summary()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"id":"cQH5hsl_WiQk"},"cell_type":"code","source":"from keras.callbacks import ReduceLROnPlateau\nfrom keras.callbacks import EarlyStopping\nfrom keras.callbacks import ModelCheckpoint\n\nlr_reduce=ReduceLROnPlateau(monitor='val_loss',\n                            factor=.5,\n                            patience=4,\n                            mode='min',\n                            min_lr=.000001,\n                            verbose=1)\n\nes_monitor=EarlyStopping(monitor='val_loss',\n                         patience=8,\n                         mode = 'min',\n                         verbose = 1,\n                         restore_best_weights = True)\n\nmdl_check = ModelCheckpoint('EffNetB4_512_8_best_weights.h5',\n                            monitor='val_loss',\n                            save_weights_only = False, \n                            verbose=1, \n                            save_best_only=True, \n                            mode='min')","execution_count":null,"outputs":[]},{"metadata":{"id":"sM8JifK6f-DM","trusted":true},"cell_type":"code","source":"model.compile(loss=\"categorical_crossentropy\", \n              optimizer=tf.keras.optimizers.Adam(learning_rate=0.00015), \n              metrics=[\"accuracy\"])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"history = model.fit_generator(train_generator,\n                     epochs = 18,\n                     verbose = 1,\n                     callbacks = [es_monitor, mdl_check, lr_reduce],\n                     validation_data = valid_generator)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"id":"PCPY5QU_WiQn"},"cell_type":"code","source":"h = history.history\n\noffset = 0\nepochs = range(offset, len(h['loss']))\n\nplt.figure(1, figsize=(12, 12))\n\nplt.subplot(211)\nplt.xlabel('epochs')\nplt.ylabel('loss')\nplt.plot(epochs, h['loss'][offset:], label='train')\nplt.plot(epochs, h['val_loss'][offset:], label='val')\nplt.legend()\n\nplt.subplot(212)\nplt.xlabel('epochs')\nplt.ylabel('accuracy')\nplt.plot(h[f'accuracy'], label='train')\nplt.plot(h[f'val_accuracy'], label='val')\nplt.legend()\n\nplt.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}