{"cells":[{"metadata":{},"cell_type":"markdown","source":"# Cassava Leaf Disease Classification"},{"metadata":{"trusted":true,"_kg_hide-input":true,"_kg_hide-output":true},"cell_type":"code","source":"import os\nimport json\nimport pandas as pd\n\n%matplotlib inline\n%config InlineBackend.figure_format = 'retina'\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\nimport seaborn as sns\nimport cv2\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import accuracy_score\nimport tensorflow as tf\nfrom tensorflow.keras import models, layers\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.callbacks import ReduceLROnPlateau\nfrom keras.optimizers import RMSprop\nfrom keras.layers.normalization import BatchNormalization\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"path = '../input/cassava-leaf-disease-classification/'","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"os.listdir(path)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"print('No of Train images: ' + str(len(os.listdir(path + 'train_images'))))\nprint('No of Test images: ' + str(len(os.listdir(path + 'test_images'))))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train = pd.read_csv(path + 'train.csv')\ntrain.head()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"with open(os.path.join(path + 'label_num_to_disease_map.json')) as f:\n    label_name = json.loads(f.read())\n    \nprint(json.dumps(label_name, indent = 1))","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train['label'] = train['label'].astype(str)\ntrain['label_name'] = train['label'].map(label_name)\ntrain.head()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Exploratory Data Analysis"},{"metadata":{"trusted":true},"cell_type":"code","source":"plt.figure(figsize = (12,6))\nsns.countplot(y = 'label_name', data = train, order = pd.value_counts(train['label_name']).index, palette = 'muted', edgecolor = 'black')\n\nplt.xlabel(\"\")\nplt.ylabel(\"\")\nplt.yticks(fontsize = 12)\nplt.show()","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train['label_name'].value_counts()","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"There are:\n\n- <b>13158</b> leaf images having Cassava Mosaic Disease (CMD)\n- <b>2577</b> healthy leaf images \n- <b>2386</b> leaf images having Cassava Green Mottle (CGM)\n- <b>2189</b> leaf images having Cassava Brown Streak Disease (CBSD)\n- <b>1087</b> leaf images having Cassava Bacterial Blight (CBB)"},{"metadata":{"_kg_hide-input":true,"trusted":true},"cell_type":"code","source":"##Credits to https://www.kaggle.com/ihelon/cassava-leaf-disease-exploratory-data-analysis for this function\n\ndef get_image(image_id, labels):\n    \n    plt.figure(figsize=(20, 18))\n    \n    for i, (image_id, label_name) in enumerate(zip(image_id, labels)):\n        plt.subplot(4, 3, i + 1)\n        image = cv2.imread(os.path.join(path, 'train_images', image_id))\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        \n        plt.imshow(image)\n        plt.title(f\"{label_name}\", fontweight='bold', fontsize=12)\n        plt.axis(\"off\")\n    \n    plt.show()","execution_count":null,"outputs":[]},{"metadata":{"_kg_hide-input":true,"trusted":true},"cell_type":"code","source":"sample = train.sample(12)\nimage_ids = sample['image_id'].values\nlabels = sample['label_name'].values\n\nget_image(image_ids, labels)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Cassava Mosaic Disease (CMD)**"},{"metadata":{"_kg_hide-input":true,"trusted":true},"cell_type":"code","source":"##Cassava Mosaic Disease (CMD)\ncmd_sample = train[train['label'] == '3'].sample(12)\nimage_ids = cmd_sample['image_id'].values\nlabels = cmd_sample['label_name'].values\n\nget_image(image_ids, labels)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Healthy**"},{"metadata":{"_kg_hide-input":true,"trusted":true},"cell_type":"code","source":"##healthy\nhealthy_sample = train[train['label'] == '4'].sample(12)\nimage_ids = healthy_sample['image_id'].values\nlabels = healthy_sample['label_name'].values\n\nget_image(image_ids, labels)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Cassava Green Mottle (CGM)**"},{"metadata":{"_kg_hide-input":true,"trusted":true},"cell_type":"code","source":"##Cassava Green Mottle (CGM)\ncgm_sample = train[train['label'] == '2'].sample(12)\nimage_ids = cgm_sample['image_id'].values\nlabels = cgm_sample['label_name'].values\n\nget_image(image_ids, labels)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Cassava Brown Streak Disease (CBSD)**"},{"metadata":{"_kg_hide-input":true,"trusted":true},"cell_type":"code","source":"##Cassava Brown Streak Disease (CBSD)\ncbsd_sample = train[train['label'] == '1'].sample(12)\nimage_ids = cbsd_sample['image_id'].values\nlabels = cbsd_sample['label_name'].values\n\nget_image(image_ids, labels)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"**Cassava Bacterial Blight (CBB)**"},{"metadata":{"_kg_hide-input":true,"trusted":true},"cell_type":"code","source":"##Cassava Bacterial Blight (CBB)\ncbb_sample = train[train['label'] == '0'].sample(12)\nimage_ids = cbb_sample['image_id'].values\nlabels = cbb_sample['label_name'].values\n\nget_image(image_ids, labels)","execution_count":null,"outputs":[]},{"metadata":{},"cell_type":"markdown","source":"## Modelling"},{"metadata":{"trusted":true},"cell_type":"code","source":"train, validation = train_test_split(train, train_size = 0.8, shuffle = True, random_state = 8)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-output":true},"cell_type":"code","source":"model = tf.keras.Sequential([\n    tf.keras.layers.Conv2D(32, (5, 5), activation='relu', input_shape=(150, 150, 3)),\n    tf.keras.layers.MaxPooling2D(2, 2),\n    tf.keras.layers.Conv2D(64, (5, 5), activation='relu'),\n    tf.keras.layers.MaxPooling2D(2, 2),\n    tf.keras.layers.Conv2D(128, (5, 5), activation='relu'),\n    tf.keras.layers.MaxPooling2D(2, 2),\n    tf.keras.layers.Conv2D(128, (5, 5), activation='relu'),\n    tf.keras.layers.MaxPooling2D(2, 2),\n    tf.keras.layers.Flatten(),\n    tf.keras.layers.Dense(512, activation='relu'),\n    tf.keras.layers.Dense(5, activation='softmax')\n])\n\nmodel.compile(optimizer = RMSprop(), loss='categorical_crossentropy', metrics=['acc'])\n\ncallbacks = ReduceLROnPlateau(monitor='val_acc', \n                              factor=0.5, \n                              patience=5, \n                              verbose=1, \n                              min_lr=0.0001)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"train_datagen = ImageDataGenerator(rescale=1/255,\n                                   rotation_range=40,\n                                   width_shift_range=0.2,\n                                   height_shift_range=0.2,\n                                   shear_range=0.2,\n                                   zoom_range=0.2,\n                                   horizontal_flip=True,\n                                   vertical_flip=True)\n\nvalidation_datagen = ImageDataGenerator(rescale=1/255)","execution_count":null,"outputs":[]},{"metadata":{"trusted":true},"cell_type":"code","source":"BATCH_SIZE = 256\nSTEPS_PER_EPOCH = train.shape[0]/BATCH_SIZE\nVALIDATION_STEPS = validation.shape[0]/BATCH_SIZE\nEPOCHS = 20\n\ntrain_generator = train_datagen.flow_from_dataframe(train, \n                                                    directory = os.path.join(path, 'train_images'),\n                                                    x_col = 'image_id',\n                                                    y_col = 'label',\n                                                    target_size = (150, 150),\n                                                    batch_size = BATCH_SIZE,\n                                                    class_mode = 'categorical')\n\nvalidation_generator = validation_datagen.flow_from_dataframe(validation, \n                                                    directory = os.path.join(path, 'train_images'),\n                                                    x_col = 'image_id',\n                                                    y_col = 'label',\n                                                    target_size = (150,150),\n                                                    batch_size = BATCH_SIZE,\n                                                    class_mode = 'categorical')","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-output":true},"cell_type":"code","source":"history = model.fit_generator(\n            train_generator,\n            steps_per_epoch = STEPS_PER_EPOCH,\n            epochs = EPOCHS,\n            validation_data = validation_generator,\n            validation_steps = VALIDATION_STEPS,\n            verbose = 1,\n            callbacks = [callbacks])","execution_count":null,"outputs":[]},{"metadata":{"trusted":true,"_kg_hide-input":true},"cell_type":"code","source":"epochs = range(1, EPOCHS + 1)\n\nacc = history.history['acc']\nval_acc = history.history['val_acc']\nloss = history.history['loss']\nval_loss = history.history['val_loss']\n\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize=(20, 6))\nax1.plot(epochs, acc, label = 'Training Accuracy')\nax1.plot(epochs, val_acc, label = 'Validation Accuracy')\nax1.set_title('Training & Validation Accuracy', fontweight='bold', fontsize=16)\nax1.legend()\n\nax2.plot(epochs, loss, label = 'Training loss')\nax2.plot(epochs, val_loss, label = 'Validation loss')\nax2.set_title('Training & Validation Loss', fontweight='bold', fontsize=16)\nax2.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}