{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.6","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":13836,"databundleVersionId":1718836,"sourceType":"competition"}],"dockerImageVersionId":30034,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **Cassava Leaf Disease Classification: CNN Keras Baseline**\n![Cassava](https://scx2.b-cdn.net/gfx/news/2019/3-geneeditingt.jpg)\n\n\n","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport datetime\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 keras.callbacks import ModelCheckpoint, EarlyStopping, ReduceLROnPlateau\nfrom tensorflow.keras.applications import ResNet50, DenseNet121, EfficientNetB0\nfrom keras.optimizers import Adam\n\n# ignoring warnings\nimport warnings\nwarnings.simplefilter(\"ignore\")\n\nimport os, cv2, json\nfrom PIL import Image","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-04-15T10:38:25.125920Z","iopub.execute_input":"2024-04-15T10:38:25.126317Z","iopub.status.idle":"2024-04-15T10:38:30.673686Z","shell.execute_reply.started":"2024-04-15T10:38:25.126281Z","shell.execute_reply":"2024-04-15T10:38:30.672967Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Work directory","metadata":{}},{"cell_type":"code","source":"WORK_DIR = '../input/cassava-leaf-disease-classification'\nos.listdir(WORK_DIR)","metadata":{"_uuid":"d629ff2d2480ee46fbb7e2d37f6b5fab8052498a","_cell_guid":"79c7e3d0-c299-4dcb-8224-4455121ee9b0","execution":{"iopub.status.busy":"2024-04-15T10:38:30.675763Z","iopub.execute_input":"2024-04-15T10:38:30.676062Z","iopub.status.idle":"2024-04-15T10:38:30.683979Z","shell.execute_reply.started":"2024-04-15T10:38:30.676032Z","shell.execute_reply":"2024-04-15T10:38:30.683230Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-one\"></a>\n# First look at the data","metadata":{}},{"cell_type":"code","source":"print('Train images: %d' %len(os.listdir(\n    os.path.join(WORK_DIR, \"train_images\"))))","metadata":{"execution":{"iopub.status.busy":"2024-04-15T10:38:30.685212Z","iopub.execute_input":"2024-04-15T10:38:30.685607Z","iopub.status.idle":"2024-04-15T10:38:31.005587Z","shell.execute_reply.started":"2024-04-15T10:38:30.685569Z","shell.execute_reply":"2024-04-15T10:38:31.004764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"with open(os.path.join(WORK_DIR, \"label_num_to_disease_map.json\")) as file:\n    print(json.dumps(json.loads(file.read()), indent=4))","metadata":{"execution":{"iopub.status.busy":"2024-04-15T10:38:31.006584Z","iopub.execute_input":"2024-04-15T10:38:31.006845Z","iopub.status.idle":"2024-04-15T10:38:31.016206Z","shell.execute_reply.started":"2024-04-15T10:38:31.006819Z","shell.execute_reply":"2024-04-15T10:38:31.015431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels = pd.read_csv(os.path.join(WORK_DIR, \"train.csv\"))\ntrain_labels.head()","metadata":{"execution":{"iopub.status.busy":"2024-04-15T10:38:31.020154Z","iopub.execute_input":"2024-04-15T10:38:31.020445Z","iopub.status.idle":"2024-04-15T10:38:31.065788Z","shell.execute_reply.started":"2024-04-15T10:38:31.020392Z","shell.execute_reply":"2024-04-15T10:38:31.064923Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sns.countplot(train_labels.label, edgecolor = 'black',\n              palette = sns.color_palette(\"viridis\", 5))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-04-15T10:38:31.067972Z","iopub.execute_input":"2024-04-15T10:38:31.068333Z","iopub.status.idle":"2024-04-15T10:38:31.242538Z","shell.execute_reply.started":"2024-04-15T10:38:31.068295Z","shell.execute_reply":"2024-04-15T10:38:31.241756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We have imbalanced data with domination of third class: \"Cassava Mosaic Disease (CMD)\"","metadata":{}},{"cell_type":"markdown","source":"## Some photos of dominant class (with CMD)","metadata":{}},{"cell_type":"code","source":"sample = train_labels[train_labels.label == 3].sample(6)\nplt.figure(figsize=(16, 8))\nfor ind, (image_id, label) in enumerate(zip(sample.image_id, sample.label)):\n    plt.subplot(2, 3, ind + 1)\n    image = cv2.imread(os.path.join(WORK_DIR, \"train_images\", image_id))\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    plt.imshow(image)\n    plt.axis(\"off\")\n    \nplt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-04-15T10:38:31.243666Z","iopub.execute_input":"2024-04-15T10:38:31.243913Z","iopub.status.idle":"2024-04-15T10:38:32.018095Z","shell.execute_reply.started":"2024-04-15T10:38:31.243888Z","shell.execute_reply":"2024-04-15T10:38:32.017200Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Some photos of healthy plants","metadata":{}},{"cell_type":"code","source":"sample = train_labels[train_labels.label == 4].sample(6)\nplt.figure(figsize=(16, 8))\nfor ind, (image_id, label) in enumerate(zip(sample.image_id, sample.label)):\n    plt.subplot(2, 3, ind + 1)\n    image = cv2.imread(os.path.join(WORK_DIR, \"train_images\", image_id))\n    image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n    plt.imshow(image)\n    plt.axis(\"off\")\n    \nplt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-04-15T10:38:32.019361Z","iopub.execute_input":"2024-04-15T10:38:32.019686Z","iopub.status.idle":"2024-04-15T10:38:32.794006Z","shell.execute_reply.started":"2024-04-15T10:38:32.019656Z","shell.execute_reply":"2024-04-15T10:38:32.792857Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-two\"></a>\n# The baseline level of accuracy","metadata":{}},{"cell_type":"code","source":"y_pred = [3] * len(train_labels.label)\nprint('The baseline accuracy: %.3f' \n      %accuracy_score(y_pred, train_labels.label))","metadata":{"execution":{"iopub.status.busy":"2024-04-15T10:38:32.795233Z","iopub.execute_input":"2024-04-15T10:38:32.795594Z","iopub.status.idle":"2024-04-15T10:38:32.816714Z","shell.execute_reply.started":"2024-04-15T10:38:32.795562Z","shell.execute_reply":"2024-04-15T10:38:32.815922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Our future model must have accuracy better than 0.615.","metadata":{}},{"cell_type":"markdown","source":"<a id=\"section-three\"></a>\n# Modeling","metadata":{}},{"cell_type":"code","source":"# The TRAIN/VALID split is performing in the generator directly.\n\n#train, valid = train_test_split(train_labels, train_size = 0.8, shuffle = True,\n#                                random_state = 0)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-04-15T10:38:32.818159Z","iopub.execute_input":"2024-04-15T10:38:32.818647Z","iopub.status.idle":"2024-04-15T10:38:32.822748Z","shell.execute_reply.started":"2024-04-15T10:38:32.818596Z","shell.execute_reply":"2024-04-15T10:38:32.821747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(15, 5))\n#sns.set_style(\"white\")\n#plt.suptitle('Train vs Valid labels', size = 15)\n#\n#sns.countplot(train.label, edgecolor = 'black', ax = ax1,\n#              palette = sns.color_palette(\"viridis\", 5))\n#sns.countplot(valid.label, edgecolor = 'black', ax = ax2,\n#              palette = sns.color_palette(\"viridis\", 5))\n#plt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-04-15T10:38:32.823845Z","iopub.execute_input":"2024-04-15T10:38:32.824132Z","iopub.status.idle":"2024-04-15T10:38:32.834045Z","shell.execute_reply.started":"2024-04-15T10:38:32.824104Z","shell.execute_reply":"2024-04-15T10:38:32.833313Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 16\nSTEPS_PER_EPOCH = len(train_labels)*0.8 / BATCH_SIZE\nVALIDATION_STEPS = len(train_labels)*0.2 / BATCH_SIZE\nEPOCHS = 20\nTARGET_SIZE = 224","metadata":{"execution":{"iopub.status.busy":"2024-04-15T10:38:32.835127Z","iopub.execute_input":"2024-04-15T10:38:32.835382Z","iopub.status.idle":"2024-04-15T10:38:32.843464Z","shell.execute_reply.started":"2024-04-15T10:38:32.835357Z","shell.execute_reply":"2024-04-15T10:38:32.842736Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels.label = train_labels.label.astype('str')\n\ntrain_generator = ImageDataGenerator(validation_split = 0.2,\n                                     preprocessing_function = None,\n                                     zoom_range = 0.2,\n                                     cval = 0.2,\n                                     horizontal_flip = True,\n                                     vertical_flip = True,\n                                     fill_mode = 'nearest',\n                                     shear_range = 0.2,\n                                     height_shift_range = 0.2,\n                                     width_shift_range = 0.2) \\\n    .flow_from_dataframe(train_labels,\n                         directory = os.path.join(WORK_DIR, \"train_images\"),\n                         subset = \"training\",\n                         x_col = \"image_id\",\n                         y_col = \"label\",\n                         target_size = (TARGET_SIZE, TARGET_SIZE),\n                         batch_size = BATCH_SIZE,\n                         class_mode = \"sparse\")\n\nvalidation_generator = ImageDataGenerator(validation_split = 0.2) \\\n    .flow_from_dataframe(train_labels,\n                         directory = os.path.join(WORK_DIR, \"train_images\"),\n                         subset = \"validation\",\n                         x_col = \"image_id\",\n                         y_col = \"label\",\n                         target_size = (TARGET_SIZE, TARGET_SIZE),\n                         batch_size = BATCH_SIZE,\n                         class_mode = \"sparse\")","metadata":{"execution":{"iopub.status.busy":"2024-04-15T10:38:32.844694Z","iopub.execute_input":"2024-04-15T10:38:32.845058Z","iopub.status.idle":"2024-04-15T10:39:14.772222Z","shell.execute_reply.started":"2024-04-15T10:38:32.845022Z","shell.execute_reply":"2024-04-15T10:39:14.771269Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"After some experiments with various pre-trained networks, I've stopped on [EfficientNetB0](https://www.tensorflow.org/api_docs/python/tf/keras/applications/EfficientNetB0). It will be my baseline NN for future improvements.","metadata":{}},{"cell_type":"code","source":"def create_model():\n    model = models.Sequential()\n\n    model.add(EfficientNetB0(include_top = False, weights = 'imagenet',\n                             input_shape = (TARGET_SIZE, TARGET_SIZE, 3)))\n    \n    model.add(layers.GlobalAveragePooling2D())\n    model.add(layers.Dense(5, activation = \"softmax\"))\n\n    model.compile(optimizer = Adam(lr = 0.001),\n                  loss = \"sparse_categorical_crossentropy\",\n                  metrics = [\"acc\"])\n    return model","metadata":{"execution":{"iopub.status.busy":"2024-04-15T10:39:14.773337Z","iopub.execute_input":"2024-04-15T10:39:14.773616Z","iopub.status.idle":"2024-04-15T10:39:14.780562Z","shell.execute_reply.started":"2024-04-15T10:39:14.773588Z","shell.execute_reply":"2024-04-15T10:39:14.779506Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = create_model()\nmodel.summary()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-04-15T10:39:14.781918Z","iopub.execute_input":"2024-04-15T10:39:14.782198Z","iopub.status.idle":"2024-04-15T10:39:20.459986Z","shell.execute_reply.started":"2024-04-15T10:39:14.782170Z","shell.execute_reply":"2024-04-15T10:39:20.459085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_save = ModelCheckpoint('./best_baseline_model.h5', \n                             save_best_only = True, \n                             save_weights_only = True,\n                             monitor = 'val_loss', \n                             mode = 'min', verbose = 1)\nearly_stop = EarlyStopping(monitor = 'val_loss', min_delta = 0.001, \n                           patience = 5, mode = 'min', verbose = 1,\n                           restore_best_weights = True)\nreduce_lr = ReduceLROnPlateau(monitor = 'val_loss', factor = 0.3, \n                              patience = 2, min_delta = 0.001, \n                              mode = 'min', verbose = 1)\n\n\nhistory = 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    callbacks = [model_save, early_stop, reduce_lr]\n)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-04-15T10:39:20.461165Z","iopub.execute_input":"2024-04-15T10:39:20.461471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"acc = history.history['acc']\nval_acc = history.history['val_acc']\nloss = history.history['loss']\nval_loss = history.history['val_loss']\n\nepochs = range(1, len(acc) + 1)\n\nfig, (ax1, ax2) = plt.subplots(1, 2, figsize=(15, 5))\nsns.set_style(\"white\")\nplt.suptitle('Train history', size = 15)\n\nax1.plot(epochs, acc, \"bo\", label = \"Training acc\")\nax1.plot(epochs, val_acc, \"b\", label = \"Validation acc\")\nax1.set_title(\"Training and validation acc\")\nax1.legend()\n\nax2.plot(epochs, loss, \"bo\", label = \"Training loss\", color = 'red')\nax2.plot(epochs, val_loss, \"b\", label = \"Validation loss\", color = 'red')\nax2.set_title(\"Training and validation loss\")\nax2.legend()\n\nplt.show()","metadata":{"_kg_hide-input":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.save('./baseline_model.h5')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<a id=\"section-four\"></a>\n# Prediction","metadata":{}},{"cell_type":"code","source":"ss = pd.read_csv(os.path.join(WORK_DIR, \"sample_submission.csv\"))\nss","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = []\n\nfor image_id in ss.image_id:\n    image = Image.open(os.path.join(WORK_DIR,  \"test_images\", image_id))\n    image = image.resize((TARGET_SIZE, TARGET_SIZE))\n    image = np.expand_dims(image, axis = 0)\n    preds.append(np.argmax(model.predict(image)))\n\nss['label'] = preds\nss","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ss.to_csv('submission.csv', index = False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}}]}