{"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":"nvidiaTeslaT4","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\n![Cassava](https://scx2.b-cdn.net/gfx/news/2019/3-geneeditingt.jpg)\n\n## 👥 Team Members\n- **Ahmed Eltokhy**\n- **Ahmed Adel**\n- **Hamza Mahrous**\n\n---\n\n##  Project Overview\n\nAs the second-largest provider of carbohydrates in Africa, cassava is a key food security crop grown by smallholder farmers because it can withstand harsh conditions. At least 80% of household farms in Sub-Saharan Africa grow this starchy root, but viral diseases are major sources of poor yields. With the help of data science, it may be possible to identify common diseases so they can be treated.\n\nExisting methods of disease detection require farmers to solicit the help of government-funded agricultural experts to visually inspect and diagnose the plants. This suffers from being labor-intensive, low-supply and costly. As an added challenge, effective solutions for farmers must perform well under significant constraints, since African farmers may only have access to mobile-quality cameras with low-bandwidth.\n\n---\n\n##  Table of Contents\n\n1. [**First look at the data**](#section-one)\n2. [**The baseline level of accuracy**](#section-two)\n3. [**Modeling**](#section-three)\n4. [**Prediction**](#section-four)\n\n---\n\n##  Dataset Information\n\n### Total Images: 21,397\n\n### Data Split\n- **Training Set:** 70% → 14,978 images\n- **Validation Set:** 20% → 4,279 images\n- **Test Set:** 10% → 2,140 images\n\n---\n\n##  Key Metrics Achieved\n\n### Overall Performance\n- **Overall Accuracy:** 84.91%\n  \n### Model Details\n- **Architecture:** EfficientNetB0\n- **Test Set Size:** 2,140 images\n- **Disease Classes:** 5 (CBB, CBSD, CGM, CMD, Healthy)\n\n### Per-Class Performance\n\n| Disease Class | Accuracy | Test Samples |\n|---------------|----------|---------|\n| CBB (Cassava Bacterial Blight) | High | 109 |\n| CBSD (Cassava Brown Streak Disease) | High | 219 |\n| CGM (Cassava Green Mottle) | High | 239 |\n| CMD (Cassava Mosaic Disease) | High | 1,316 |\n| Healthy | High | 257 |\n\n---\n\n##  Getting Started\n\nThe notebook is organized into four main sections:\n\n1. **First look at the data** - Data exploration and visualization\n2. **The baseline level of accuracy** - Baseline model calculation\n3. **Modeling** - Model architecture and training with EfficientNetB0\n4. **Prediction** - Evaluation metrics and comprehensive analysis\n\nEach section builds upon the previous one to create a complete classification pipeline for cassava leaf disease detection.","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","trusted":true,"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-12-19T19:17:42.435838Z","iopub.execute_input":"2025-12-19T19:17:42.436314Z","iopub.status.idle":"2025-12-19T19:17:48.058197Z","shell.execute_reply.started":"2025-12-19T19:17:42.436269Z","shell.execute_reply":"2025-12-19T19:17:48.057535Z"}},"outputs":[],"execution_count":null},{"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","trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:17:48.059967Z","iopub.execute_input":"2025-12-19T19:17:48.060183Z","iopub.status.idle":"2025-12-19T19:17:48.068818Z","shell.execute_reply.started":"2025-12-19T19:17:48.060163Z","shell.execute_reply":"2025-12-19T19:17:48.068052Z"}},"outputs":[],"execution_count":null},{"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(os.path.join(WORK_DIR, \"train_images\"))))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:17:48.070190Z","iopub.execute_input":"2025-12-19T19:17:48.070525Z","iopub.status.idle":"2025-12-19T19:17:48.288251Z","shell.execute_reply.started":"2025-12-19T19:17:48.070475Z","shell.execute_reply":"2025-12-19T19:17:48.287385Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:17:48.289458Z","iopub.execute_input":"2025-12-19T19:17:48.289784Z","iopub.status.idle":"2025-12-19T19:17:48.298291Z","shell.execute_reply.started":"2025-12-19T19:17:48.289758Z","shell.execute_reply":"2025-12-19T19:17:48.297522Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_labels = pd.read_csv(os.path.join(WORK_DIR, \"train.csv\"))\ntrain_labels.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:17:48.301760Z","iopub.execute_input":"2025-12-19T19:17:48.301966Z","iopub.status.idle":"2025-12-19T19:17:48.337678Z","shell.execute_reply.started":"2025-12-19T19:17:48.301946Z","shell.execute_reply":"2025-12-19T19:17:48.336952Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sns.countplot(train_labels.label, edgecolor = 'black',\n              palette = sns.color_palette(\"viridis\", 5))\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:17:48.339873Z","iopub.execute_input":"2025-12-19T19:17:48.340191Z","iopub.status.idle":"2025-12-19T19:17:48.497642Z","shell.execute_reply.started":"2025-12-19T19:17:48.340159Z","shell.execute_reply":"2025-12-19T19:17:48.496764Z"}},"outputs":[],"execution_count":null},{"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 \"0\": \"Cassava Bacterial Blight (CBB)","metadata":{}},{"cell_type":"code","source":"sample = train_labels[train_labels.label == '0'].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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:35:42.613465Z","iopub.execute_input":"2025-12-19T19:35:42.613825Z","iopub.status.idle":"2025-12-19T19:35:43.353575Z","shell.execute_reply.started":"2025-12-19T19:35:42.613797Z","shell.execute_reply":"2025-12-19T19:35:43.352779Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Some photos of \"1\": \"Cassava Brown Streak Disease (CBSD)\"","metadata":{}},{"cell_type":"code","source":"sample = train_labels[train_labels.label == '1'].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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:35:37.727082Z","iopub.execute_input":"2025-12-19T19:35:37.727365Z","iopub.status.idle":"2025-12-19T19:35:38.387562Z","shell.execute_reply.started":"2025-12-19T19:35:37.727341Z","shell.execute_reply":"2025-12-19T19:35:38.386431Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Some photos of \"2\": \"Cassava Green Mottle (CGM)\"","metadata":{}},{"cell_type":"code","source":"sample = train_labels[train_labels.label == '2'].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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:35:31.280905Z","iopub.execute_input":"2025-12-19T19:35:31.281203Z","iopub.status.idle":"2025-12-19T19:35:31.982033Z","shell.execute_reply.started":"2025-12-19T19:35:31.281179Z","shell.execute_reply":"2025-12-19T19:35:31.981085Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Some photos of \"3\": \"Cassava Mosaic Disease (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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:35:05.194660Z","iopub.execute_input":"2025-12-19T19:35:05.194953Z","iopub.status.idle":"2025-12-19T19:35:05.843058Z","shell.execute_reply.started":"2025-12-19T19:35:05.194929Z","shell.execute_reply":"2025-12-19T19:35:05.842113Z"}},"outputs":[],"execution_count":null},{"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,"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:34:53.696869Z","iopub.execute_input":"2025-12-19T19:34:53.697155Z","iopub.status.idle":"2025-12-19T19:34:54.388554Z","shell.execute_reply.started":"2025-12-19T19:34:53.697131Z","shell.execute_reply":"2025-12-19T19:34:54.387810Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:17:49.926109Z","iopub.execute_input":"2025-12-19T19:17:49.926347Z","iopub.status.idle":"2025-12-19T19:17:49.936903Z","shell.execute_reply.started":"2025-12-19T19:17:49.926322Z","shell.execute_reply":"2025-12-19T19:17:49.936120Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2025-12-19T19:17:49.938237Z","iopub.execute_input":"2025-12-19T19:17:49.938617Z","iopub.status.idle":"2025-12-19T19:17:49.944346Z","shell.execute_reply.started":"2025-12-19T19:17:49.938593Z","shell.execute_reply":"2025-12-19T19:17:49.943372Z"}},"outputs":[],"execution_count":null},{"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,"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:17:49.945601Z","iopub.execute_input":"2025-12-19T19:17:49.945851Z","iopub.status.idle":"2025-12-19T19:17:49.952040Z","shell.execute_reply.started":"2025-12-19T19:17:49.945810Z","shell.execute_reply":"2025-12-19T19:17:49.951336Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 16\nEPOCHS = 20\nTARGET_SIZE = 224\n\nSTEPS_PER_EPOCH = len(train_df) // BATCH_SIZE\nVALIDATION_STEPS = len(val_df) // BATCH_SIZE\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:32:45.863915Z","iopub.execute_input":"2025-12-19T19:32:45.864287Z","iopub.status.idle":"2025-12-19T19:32:45.868710Z","shell.execute_reply.started":"2025-12-19T19:32:45.864238Z","shell.execute_reply":"2025-12-19T19:32:45.867897Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\ntrain_labels.label = train_labels.label.astype(str)\n\ntrain_df, temp_df = train_test_split(train_labels, train_size=0.7, stratify=train_labels.label, random_state=42)\nval_df, test_df = train_test_split(temp_df, test_size=(1/3), stratify=temp_df.label, random_state=42)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:32:49.159447Z","iopub.execute_input":"2025-12-19T19:32:49.159798Z","iopub.status.idle":"2025-12-19T19:32:49.202992Z","shell.execute_reply.started":"2025-12-19T19:32:49.159772Z","shell.execute_reply":"2025-12-19T19:32:49.202353Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_datagen = ImageDataGenerator(\n    zoom_range=0.2,\n    horizontal_flip=True,\n    vertical_flip=True,\n    shear_range=0.2,\n    height_shift_range=0.2,\n    width_shift_range=0.2,\n    fill_mode='nearest'\n)\n\ntrain_generator = train_datagen.flow_from_dataframe(\n    train_df,\n    directory=os.path.join(WORK_DIR, \"train_images\"),\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)\n\nval_datagen = ImageDataGenerator()\n\nvalidation_generator = val_datagen.flow_from_dataframe(\n    val_df,\n    directory=os.path.join(WORK_DIR, \"train_images\"),\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)\n\ntest_datagen = ImageDataGenerator()\n\ntest_generator = test_datagen.flow_from_dataframe(\n    test_df,\n    directory=os.path.join(WORK_DIR, \"train_images\"),\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    shuffle=False\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:32:50.391934Z","iopub.execute_input":"2025-12-19T19:32:50.392225Z","iopub.status.idle":"2025-12-19T19:33:00.182790Z","shell.execute_reply.started":"2025-12-19T19:32:50.392197Z","shell.execute_reply":"2025-12-19T19:33:00.182032Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:35:55.073211Z","iopub.execute_input":"2025-12-19T19:35:55.073545Z","iopub.status.idle":"2025-12-19T19:35:55.079700Z","shell.execute_reply.started":"2025-12-19T19:35:55.073507Z","shell.execute_reply":"2025-12-19T19:35:55.078532Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = create_model()\nmodel.summary()","metadata":{"_kg_hide-output":true,"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:35:56.022549Z","iopub.execute_input":"2025-12-19T19:35:56.022886Z","iopub.status.idle":"2025-12-19T19:35:59.207536Z","shell.execute_reply.started":"2025-12-19T19:35:56.022856Z","shell.execute_reply":"2025-12-19T19:35:59.206766Z"}},"outputs":[],"execution_count":null},{"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(\n    train_generator,\n    steps_per_epoch=len(train_generator),\n    epochs=EPOCHS,\n    validation_data=validation_generator,\n    validation_steps=len(validation_generator),\n    callbacks=[model_save, early_stop, reduce_lr]\n)","metadata":{"_kg_hide-output":true,"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T19:36:01.050267Z","iopub.execute_input":"2025-12-19T19:36:01.050596Z","iopub.status.idle":"2025-12-19T20:57:24.052422Z","shell.execute_reply.started":"2025-12-19T19:36:01.050565Z","shell.execute_reply":"2025-12-19T20:57:24.051419Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_loss, test_acc = model.evaluate(test_generator)\nprint(\"Test accuracy:\", test_acc)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T20:57:24.054628Z","iopub.execute_input":"2025-12-19T20:57:24.054998Z","iopub.status.idle":"2025-12-19T20:57:53.672988Z","shell.execute_reply.started":"2025-12-19T20:57:24.054958Z","shell.execute_reply":"2025-12-19T20:57:53.672143Z"}},"outputs":[],"execution_count":null},{"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":{"iopub.status.busy":"2025-12-19T20:57:53.674025Z","iopub.execute_input":"2025-12-19T20:57:53.674249Z","iopub.status.idle":"2025-12-19T20:57:53.961414Z","shell.execute_reply.started":"2025-12-19T20:57:53.674226Z","shell.execute_reply":"2025-12-19T20:57:53.960570Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model.save('./baseline_model.h5')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T21:15:44.441089Z","iopub.execute_input":"2025-12-19T21:15:44.441376Z","iopub.status.idle":"2025-12-19T21:15:45.062498Z","shell.execute_reply.started":"2025-12-19T21:15:44.441353Z","shell.execute_reply":"2025-12-19T21:15:45.061519Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"<a id=\"section-four\"></a>\n# Predictions - Comprehensive Visualization for Cassava Disease Classification Model\n","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom sklearn.metrics import confusion_matrix, classification_report\nfrom PIL import Image\nimport os\nfrom math import pi\n\ndisease_map = {\n    0: 'CBB',\n    1: 'CBSD',\n    2: 'CGM',\n    3: 'CMD',\n    4: 'Healthy'\n}\n\nWORK_DIR = '../input/cassava-leaf-disease-classification'\nTARGET_SIZE = 224","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T21:44:21.744446Z","iopub.execute_input":"2025-12-19T21:44:21.744785Z","iopub.status.idle":"2025-12-19T21:44:21.749870Z","shell.execute_reply.started":"2025-12-19T21:44:21.744755Z","shell.execute_reply":"2025-12-19T21:44:21.749041Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(f\"\\n Data Info:\")\nprint(f\"   • Total test images: {len(y_true)}\")\nprint(f\"   • Total predictions: {len(y_pred)}\")\n\ny_true = np.array(y_true)\ny_pred = np.array(y_pred)\n\nprint(f\"\\n Test Set Breakdown (by True Label):\")\nfor i in range(5):\n    count = (y_true == i).sum()\n    percentage = (count / len(y_true)) * 100\n    print(f\"   • {disease_map[i]:10s}: {count:4d} images ({percentage:5.2f}%)\")\n\nprint(f\"\\n Predictions Breakdown (by Predicted Label):\")\nfor i in range(5):\n    count = (y_pred == i).sum()\n    percentage = (count / len(y_pred)) * 100\n    print(f\"   • {disease_map[i]:10s}: {count:4d} images ({percentage:5.2f}%)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T21:45:18.213443Z","iopub.execute_input":"2025-12-19T21:45:18.213766Z","iopub.status.idle":"2025-12-19T21:45:18.222763Z","shell.execute_reply.started":"2025-12-19T21:45:18.213733Z","shell.execute_reply":"2025-12-19T21:45:18.221913Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Confusion Matrix Heatmaps\nprint(\"\\n Generating Confusion Matrices...\")\n\nfig, axes = plt.subplots(1, 2, figsize=(18, 6))\n\n# Standard Confusion Matrix\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n            xticklabels=disease_map.values(),\n            yticklabels=disease_map.values(),\n            ax=axes[0], cbar_kws={'label': 'Count'}, \n            annot_kws={'size': 14, 'weight': 'bold'})\naxes[0].set_title('Confusion Matrix (Counts)', fontsize=16, fontweight='bold')\naxes[0].set_ylabel('True Label', fontsize=13)\naxes[0].set_xlabel('Predicted Label', fontsize=13)\n\n# Normalized Confusion Matrix\ncm_normalized = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]\nsns.heatmap(cm_normalized, annot=True, fmt='.2%', cmap='RdYlGn', \n            xticklabels=disease_map.values(),\n            yticklabels=disease_map.values(),\n            ax=axes[1], cbar_kws={'label': 'Percentage'}, \n            annot_kws={'size': 12})\naxes[1].set_title('Normalized Confusion Matrix (%)', fontsize=16, fontweight='bold')\naxes[1].set_ylabel('True Label', fontsize=13)\naxes[1].set_xlabel('Predicted Label', fontsize=13)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T21:45:56.398416Z","iopub.execute_input":"2025-12-19T21:45:56.398762Z","iopub.status.idle":"2025-12-19T21:45:57.089567Z","shell.execute_reply.started":"2025-12-19T21:45:56.398733Z","shell.execute_reply":"2025-12-19T21:45:57.088315Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Per-Class Metrics\nprint(\"\\n Computing Per-Class Metrics...\")\n\nmetrics_data = []\nfor i in range(5):\n    class_mask = np.array(y_true) == i\n    if class_mask.sum() > 0:\n        # Calculate TP, FP, FN\n        tp = np.sum((np.array(y_true) == i) & (np.array(y_pred) == i))\n        fp = np.sum((np.array(y_true) != i) & (np.array(y_pred) == i))\n        fn = np.sum((np.array(y_true) == i) & (np.array(y_pred) != i))\n        \n        precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n        recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n        f1 = 2 * (precision * recall) / (precision + recall) if (precision + recall) > 0 else 0\n        accuracy = np.sum((np.array(y_true)[class_mask] == \n                          np.array(y_pred)[class_mask])) / class_mask.sum()\n        \n        metrics_data.append({\n            'Disease': disease_map[i],\n            'Precision': precision,\n            'Recall': recall,\n            'F1-Score': f1,\n            'Accuracy': accuracy\n        })\n\nmetrics_df = pd.DataFrame(metrics_data)\n\nfig, ax = plt.subplots(figsize=(12, 6))\nmetrics_df.set_index('Disease')[['Precision', 'Recall', 'F1-Score', 'Accuracy']].plot(\n    kind='bar', ax=ax, width=0.8, edgecolor='black', linewidth=1.2)\nax.set_title('Per-Class Performance Metrics', fontsize=16, fontweight='bold')\nax.set_ylabel('Score', fontsize=12)\nax.set_xlabel('Disease Class', fontsize=12)\nax.set_ylim([0, 1.1])\nax.legend(loc='upper right', fontsize=11)\nax.grid(axis='y', alpha=0.3, linestyle='--')\nplt.setp(ax.xaxis.get_majorticklabels(), rotation=45, ha='right')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T21:46:23.584744Z","iopub.execute_input":"2025-12-19T21:46:23.585073Z","iopub.status.idle":"2025-12-19T21:46:23.833360Z","shell.execute_reply.started":"2025-12-19T21:46:23.585045Z","shell.execute_reply":"2025-12-19T21:46:23.832592Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Distribution Comparison\nprint(\"\\n Comparing True vs Predicted Distribution...\")\n\nfig, axes = plt.subplots(1, 2, figsize=(16, 5))\n\n# True labels distribution\nclass_counts_true = pd.Series(y_true).value_counts().sort_index()\naxes[0].bar([disease_map[i] for i in class_counts_true.index], \n            class_counts_true.values,\n            color='steelblue', edgecolor='black', linewidth=1.5, alpha=0.8)\naxes[0].set_title('Test Set Distribution (True Labels)', fontsize=14, fontweight='bold')\naxes[0].set_ylabel('Number of Images', fontsize=12)\naxes[0].set_xlabel('Disease Class', fontsize=12)\naxes[0].grid(axis='y', alpha=0.3, linestyle='--')\n\nfor i, v in enumerate(class_counts_true.values):\n    axes[0].text(i, v + 20, str(v), ha='center', fontweight='bold', fontsize=11)\n\n# Predictions distribution\nclass_counts_pred = pd.Series(y_pred).value_counts().sort_index()\naxes[1].bar([disease_map[i] for i in class_counts_pred.index], \n            class_counts_pred.values,\n            color='coral', edgecolor='black', linewidth=1.5, alpha=0.8)\naxes[1].set_title('Prediction Distribution', fontsize=14, fontweight='bold')\naxes[1].set_ylabel('Number of Images', fontsize=12)\naxes[1].set_xlabel('Disease Class', fontsize=12)\naxes[1].grid(axis='y', alpha=0.3, linestyle='--')\n\nfor i, v in enumerate(class_counts_pred.values):\n    axes[1].text(i, v + 20, str(v), ha='center', fontweight='bold', fontsize=11)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T21:46:49.031613Z","iopub.execute_input":"2025-12-19T21:46:49.031962Z","iopub.status.idle":"2025-12-19T21:46:49.350428Z","shell.execute_reply.started":"2025-12-19T21:46:49.031931Z","shell.execute_reply":"2025-12-19T21:46:49.349596Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Per-Class Accuracy Radar Chart\n# ========================================\nprint(\"\\n Creating Radar Chart for Per-Class Accuracy...\")\n\ncategories = [disease_map[i] for i in range(5)]\nvalues = []\n\nfor i in range(5):\n    class_mask = np.array(y_true) == i\n    if class_mask.sum() > 0:\n        acc = np.sum((np.array(y_true)[class_mask] == \n                     np.array(y_pred)[class_mask])) / class_mask.sum()\n        values.append(acc)\n    else:\n        values.append(0)\n\nvalues += values[:1]\nangles = [n / float(len(categories)) * 2 * pi for n in range(len(categories))]\nangles += angles[:1]\n\nfig, ax = plt.subplots(figsize=(10, 10), subplot_kw=dict(projection='polar'))\nax.plot(angles, values, 'o-', linewidth=2.5, color='#2E86AB', markersize=8)\nax.fill(angles, values, alpha=0.25, color='#2E86AB')\nax.set_xticks(angles[:-1])\nax.set_xticklabels(categories, fontsize=12, fontweight='bold')\nax.set_ylim(0, 1)\nax.set_yticks([0.2, 0.4, 0.6, 0.8, 1.0])\nax.set_yticklabels(['0.2', '0.4', '0.6', '0.8', '1.0'], fontsize=10)\nax.grid(True, linestyle='--', alpha=0.7)\nax.set_title('Per-Class Accuracy (Radar Chart)', fontsize=16, fontweight='bold', pad=20)\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T21:47:28.247984Z","iopub.execute_input":"2025-12-19T21:47:28.248300Z","iopub.status.idle":"2025-12-19T21:47:28.445589Z","shell.execute_reply.started":"2025-12-19T21:47:28.248272Z","shell.execute_reply":"2025-12-19T21:47:28.444754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Sample Predictions Gallery\nprint(\"\\n Generating Sample Predictions Gallery...\")\n\nsample_indices = np.random.choice(len(test_df), min(12, len(test_df)), replace=False)\nfig, axes = plt.subplots(3, 4, figsize=(16, 12))\naxes = axes.flatten()\n\nfor idx, sample_idx in enumerate(sample_indices):\n    img_name = test_df.iloc[sample_idx]['image_id']\n    true_label = int(test_df.iloc[sample_idx]['label'])\n    pred_label = y_pred[sample_idx]\n    \n    try:\n        img_path = os.path.join(WORK_DIR, 'train_images', img_name)\n        img = plt.imread(img_path)\n        \n        axes[idx].imshow(img)\n        \n        # Color code: green if correct, red if wrong\n        color = 'green' if true_label == pred_label else 'red'\n        title = f\"True: {disease_map[true_label]}\\nPred: {disease_map[pred_label]}\"\n        axes[idx].set_title(title, fontsize=11, fontweight='bold', color=color)\n        axes[idx].axis('off')\n    except:\n        axes[idx].text(0.5, 0.5, 'Image not found', ha='center', va='center')\n        axes[idx].axis('off')\n\nplt.suptitle('Sample Predictions (Green=Correct, Red=Wrong)', \n             fontsize=16, fontweight='bold', y=0.995)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T21:47:51.499462Z","iopub.execute_input":"2025-12-19T21:47:51.499778Z","iopub.status.idle":"2025-12-19T21:47:54.429458Z","shell.execute_reply.started":"2025-12-19T21:47:51.499751Z","shell.execute_reply":"2025-12-19T21:47:54.428222Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Summary Statistics\nprint(\"\\n\" + \"=\"*70)\nprint(\"SUMMARY STATISTICS\")\nprint(\"=\"*70)\n\noverall_accuracy = np.mean(np.array(y_true) == np.array(y_pred))\nprint(f\"\\n Overall Accuracy: {overall_accuracy:.4f} ({overall_accuracy*100:.2f}%)\")\n\nprint(f\"\\n Test Set Size: {len(y_true)} images\")\nprint(f\"   • CBB:     {(y_true==0).sum()} images\")\nprint(f\"   • CBSD:    {(y_true==1).sum()} images\")\nprint(f\"   • CGM:     {(y_true==2).sum()} images\")\nprint(f\"   • CMD:     {(y_true==3).sum()} images\")\nprint(f\"   • Healthy: {(y_true==4).sum()} images\")\n\nprint(\"\\n\" + \"=\"*70)\nprint(\"CLASSIFICATION REPORT\")\nprint(\"=\"*70)\nprint(classification_report(y_true, y_pred, \n                          target_names=list(disease_map.values()),\n                          digits=4))\n\nprint(\"\\n\" + \"=\"*70)\nprint(\" Analysis Complete\")\nprint(\"=\"*70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T21:48:38.107087Z","iopub.execute_input":"2025-12-19T21:48:38.107376Z","iopub.status.idle":"2025-12-19T21:48:38.126245Z","shell.execute_reply.started":"2025-12-19T21:48:38.107349Z","shell.execute_reply":"2025-12-19T21:48:38.125568Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Thank You ","metadata":{}}]}