{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","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":30919,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport random\nimport warnings\nimport json\nimport numpy as np\nimport pandas as pd\nfrom collections import Counter\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport tensorflow as tf\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.models import Sequential\nfrom tensorflow.keras.layers import GlobalAveragePooling2D, Flatten, Dense, Dropout\nfrom tensorflow.keras.optimizers import Adam\nfrom tensorflow.keras.callbacks import EarlyStopping, ModelCheckpoint, ReduceLROnPlateau\nfrom tensorflow.keras.applications import EfficientNetB3\nfrom tensorflow.keras.utils import plot_model\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import confusion_matrix, classification_report\n\n# ===================== Configuration =====================\nIMG_SIZE = 512  \nsize = (IMG_SIZE, IMG_SIZE)\nNUM_CLASSES = 5\nBATCH_SIZE = 16  \nEPOCHS_STAGE1 = 5\nEPOCHS_STAGE2 = 15\n\n# ===================== Setup =====================\nprint(\"Num GPUs Available:\", len(tf.config.list_physical_devices('GPU')))\n\n# Enable mixed precision for faster training\ntf.keras.mixed_precision.set_global_policy('mixed_float16')\n\n# Set seed for reproducibility\ndef seed_everything(seed=21):\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    os.environ['TF_DETERMINISTIC_OPS'] = '1'\n\nseed_everything()\nwarnings.filterwarnings('ignore')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T07:36:17.577137Z","iopub.execute_input":"2025-04-06T07:36:17.577419Z","iopub.status.idle":"2025-04-06T07:36:32.189277Z","shell.execute_reply.started":"2025-04-06T07:36:17.577399Z","shell.execute_reply":"2025-04-06T07:36:32.188518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===================== Load Dataset =====================\nwork_dir = '../input/cassava-leaf-disease-classification/'\ntrain_path = os.path.join(work_dir, 'train_images')\ndata = pd.read_csv(os.path.join(work_dir, 'train.csv'))\n\nwith open(os.path.join(work_dir, 'label_num_to_disease_map.json')) as f:\n    real_labels = json.load(f)\n    real_labels = {int(k): v for k, v in real_labels.items()}\n\n# Map label numbers to disease names\ndata['class_name'] = data['label'].map(real_labels)\ntrain_df, test_df = train_test_split(data, test_size=0.1, random_state=42, stratify=data['class_name'])","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T07:37:27.919039Z","iopub.execute_input":"2025-04-06T07:37:27.919343Z","iopub.status.idle":"2025-04-06T07:37:27.986855Z","shell.execute_reply.started":"2025-04-06T07:37:27.919322Z","shell.execute_reply":"2025-04-06T07:37:27.985985Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===================== Visualize Class Distribution =====================\nlabel_map = {\n    0: 'CBB',\n    1: 'CBSD',\n    2: 'CGM',\n    3: 'CMD',\n    4: 'Healthy'\n}\ndata['short_label'] = data['label'].map(label_map)\n\nplt.figure(figsize=(8, 5))\nax = sns.countplot(x='short_label', data=data, palette='viridis', edgecolor='black')\nplt.title(\"Class Distribution in Original Dataset\")\nplt.xlabel(\"Class Label\")\nplt.ylabel(\"Sample Count\")\nfor p in ax.patches:\n    height = int(p.get_height())\n    ax.annotate(f'{height}', (p.get_x() + p.get_width() / 2., height),\n                ha='center', va='center', fontsize=11, xytext=(0, 6), textcoords='offset points')\nplt.tight_layout()\nplt.savefig(\"class_distribution.png\")\nplt.show()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T07:37:37.573090Z","iopub.execute_input":"2025-04-06T07:37:37.573387Z","iopub.status.idle":"2025-04-06T07:37:37.962562Z","shell.execute_reply.started":"2025-04-06T07:37:37.573365Z","shell.execute_reply":"2025-04-06T07:37:37.961708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===================== Sample Images per Class =====================\nimport matplotlib.image as mpimg\nsample_dir = os.path.join(work_dir, 'train_images')\nfig, axes = plt.subplots(2, 3, figsize=(15, 10))\nclasses = sorted(data['label'].unique())\n\nfor i, label in enumerate(classes):\n    img_name = data[data['label'] == label].iloc[0]['image_id']\n    img_path = os.path.join(sample_dir, img_name)\n    img = mpimg.imread(img_path)\n    row, col = divmod(i, 3)\n    axes[row][col].imshow(img)\n    axes[row][col].set_title(f\"{label_map[label]}\", fontsize=14)\n    axes[row][col].axis('off')\nfor j in range(len(classes), 6):\n    row, col = divmod(j, 3)\n    axes[row][col].axis('off')\nplt.suptitle(\"Example Image from Each Class\", fontsize=16)\nplt.tight_layout()\nplt.savefig(\"class_samples.png\")\nplt.show()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T07:37:43.193387Z","iopub.execute_input":"2025-04-06T07:37:43.193707Z","iopub.status.idle":"2025-04-06T07:37:45.503620Z","shell.execute_reply.started":"2025-04-06T07:37:43.193680Z","shell.execute_reply":"2025-04-06T07:37:45.502693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===================== Data Generators =====================\n# Use ImageDataGenerator for augmentation; reduce augmentation intensity if needed\ndatagen_train = ImageDataGenerator(\n    validation_split=0.2,\n    preprocessing_function=tf.keras.applications.efficientnet.preprocess_input,\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    fill_mode='nearest'\n)\n\ntrain_generator = datagen_train.flow_from_dataframe(\n    train_df,\n    directory=train_path,\n    x_col='image_id',\n    y_col='class_name',\n    subset='training',\n    target_size=size,\n    class_mode='categorical',\n    shuffle=True,\n    seed=42,\n    batch_size=BATCH_SIZE\n)\n\nvalidation_datagen = ImageDataGenerator(\n    validation_split=0.2,\n    preprocessing_function=tf.keras.applications.efficientnet.preprocess_input\n)\n\nvalidation_generator = validation_datagen.flow_from_dataframe(\n    train_df,\n    directory=train_path,\n    x_col='image_id',\n    y_col='class_name',\n    subset='validation',\n    target_size=size,\n    class_mode='categorical',\n    shuffle=True,\n    seed=42,\n    batch_size=BATCH_SIZE\n)\n\ntest_datagen = ImageDataGenerator(\n    preprocessing_function=tf.keras.applications.efficientnet.preprocess_input\n)\n\ntest_generator = test_datagen.flow_from_dataframe(\n    test_df,\n    directory=train_path,\n    x_col='image_id',\n    y_col='class_name',\n    target_size=size,\n    class_mode='categorical',\n    shuffle=False,\n    seed=42,\n    batch_size=BATCH_SIZE\n)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T07:37:45.810351Z","iopub.execute_input":"2025-04-06T07:37:45.810641Z","iopub.status.idle":"2025-04-06T07:39:03.391887Z","shell.execute_reply.started":"2025-04-06T07:37:45.810619Z","shell.execute_reply":"2025-04-06T07:39:03.390882Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===================== Data Augmentation Visualization =====================\nsample_aug = train_df.iloc[[0]].copy()\nsample_aug['class_name'] = sample_aug['label'].map(real_labels)\npreview_gen = datagen_train.flow_from_dataframe(\n    sample_aug,\n    directory=train_path,\n    x_col='image_id',\n    y_col='class_name',\n    target_size=size,\n    class_mode='categorical',\n    batch_size=1\n)\naug_imgs = [preview_gen[0][0][0] / 255.0 for _ in range(4)]\nfig, axs = plt.subplots(2, 2, figsize=(8, 8))\nfor i in range(4):\n    axs.flat[i].imshow(aug_imgs[i])\n    axs.flat[i].axis('off')\nplt.suptitle(\"Data Augmentation Examples\", fontsize=16)\nplt.tight_layout()\nplt.savefig(\"augmentation_samples.png\")\nplt.show()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T07:39:03.393245Z","iopub.execute_input":"2025-04-06T07:39:03.393576Z","iopub.status.idle":"2025-04-06T07:39:04.550140Z","shell.execute_reply.started":"2025-04-06T07:39:03.393542Z","shell.execute_reply":"2025-04-06T07:39:04.549240Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===================== Compute Class Weights =====================\n# Compute class weights to counter class imbalance\nclass_counts = train_df['label'].value_counts().sort_index()\ntotal_samples = len(train_df)\nclass_weights = {i: total_samples / (NUM_CLASSES * class_counts[i]) for i in range(NUM_CLASSES)}\nprint(\"Class Weights:\", class_weights)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T07:39:04.551456Z","iopub.execute_input":"2025-04-06T07:39:04.551723Z","iopub.status.idle":"2025-04-06T07:39:04.561387Z","shell.execute_reply.started":"2025-04-06T07:39:04.551694Z","shell.execute_reply":"2025-04-06T07:39:04.560531Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===================== Model Definition =====================\ndef create_model():\n    model = Sequential([\n        EfficientNetB3(\n            input_shape=(IMG_SIZE, IMG_SIZE, 3),\n            include_top=False,\n            weights='imagenet'\n        ),\n        GlobalAveragePooling2D(),\n        Flatten(),\n        Dense(256, activation='relu',\n              kernel_regularizer=tf.keras.regularizers.l2(1e-4)),  # Using L2 regularization\n        Dropout(0.5),\n        Dense(NUM_CLASSES, activation='softmax')\n    ])\n    return model","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T07:39:04.562349Z","iopub.execute_input":"2025-04-06T07:39:04.562580Z","iopub.status.idle":"2025-04-06T07:39:04.580650Z","shell.execute_reply.started":"2025-04-06T07:39:04.562553Z","shell.execute_reply":"2025-04-06T07:39:04.579882Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===================== Two-Stage Training =====================\n# Stage 1: Train the new classifier head with the base model frozen\nmodel_stage1 = create_model()\n# Freeze base model layers\nmodel_stage1.layers[0].trainable = False\n\nmodel_stage1.compile(\n    optimizer=Adam(learning_rate=1e-3),\n    loss=tf.keras.losses.CategoricalCrossentropy(from_logits=False, label_smoothing=0.1),\n    metrics=['accuracy']\n)\n\ncallbacks_stage1 = [\n    EarlyStopping(monitor='val_accuracy', patience=3, restore_best_weights=True, verbose=1),\n    ModelCheckpoint('Cassava_best_model_stage1.keras', save_best_only=True, monitor='val_accuracy', mode='max'),\n    ReduceLROnPlateau(monitor='val_loss', factor=0.2, patience=2, min_lr=1e-6, verbose=1)\n]\n\nhistory_stage1 = model_stage1.fit(\n    train_generator,\n    validation_data=validation_generator,\n    epochs=EPOCHS_STAGE1,\n    callbacks=callbacks_stage1,\n    class_weight=class_weights\n)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T07:39:04.581521Z","iopub.execute_input":"2025-04-06T07:39:04.581733Z","execution_failed":"2025-04-06T07:45:48.125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Stage 2: Fine-tune by unfreezing part of the base model\n# Unfreeze last 60 layers of the base model as an example\nfor layer in model_stage1.layers[0].layers[-60:]:\n    layer.trainable = True\n\nmodel_stage1.compile(\n    optimizer=Adam(learning_rate=5e-5),\n    loss=tf.keras.losses.CategoricalCrossentropy(from_logits=False, label_smoothing=0.1),\n    metrics=['accuracy']\n)\n\ncallbacks_stage2 = [\n    EarlyStopping(monitor='val_accuracy', patience=5, restore_best_weights=True, verbose=1),\n    ModelCheckpoint('Cassava_best_model_stage2.keras', save_best_only=True, monitor='val_accuracy', mode='max'),\n    ReduceLROnPlateau(monitor='val_loss', factor=0.2, patience=2, min_lr=1e-6, verbose=1)\n]\n\nhistory_stage2 = model_stage1.fit(\n    train_generator,\n    validation_data=validation_generator,\n    epochs=EPOCHS_STAGE2,\n    callbacks=callbacks_stage2,\n    class_weight=class_weights\n)\n\n# Save final fine-tuned model\nmodel_stage1.save('Cassava_model_finetuned.keras')\n\n# ===================== Model Architecture Visualization =====================\nplot_model(model_stage1, to_file=\"model_architecture.png\", show_shapes=True, show_layer_names=True)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"execution_failed":"2025-04-06T07:45:48.125Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ===================== Evaluation =====================\n\ndef plot_training_history(history):\n    # Plot Accuracy\n    plt.figure(figsize=(12, 5))\n    plt.subplot(1, 2, 1)\n    plt.plot(history.history['accuracy'], label='Train Accuracy')\n    plt.plot(history.history['val_accuracy'], label='Validation Accuracy')\n    plt.title('Training & Validation Accuracy')\n    plt.xlabel('Epoch')\n    plt.ylabel('Accuracy')\n    plt.legend()\n    \n    # Plot Loss\n    plt.subplot(1, 2, 2)\n    plt.plot(history.history['loss'], label='Train Loss')\n    plt.plot(history.history['val_loss'], label='Validation Loss')\n    plt.title('Training & Validation Loss')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    \n    plt.tight_layout()\n    plt.show()\n\n\nplot_training_history(history_stage2)\n\ntest_loss, test_acc = model_stage1.evaluate(test_generator, verbose=1)\nprint(f\"\\nFinal model evaluation - Loss: {test_loss:.4f}, Accuracy: {test_acc:.4f}\")\n\n# Predict on test set and generate confusion matrix and classification report\ny_true = test_df['label'].values\npred_probs = model_stage1.predict(test_generator, verbose=1)\ny_pred = np.argmax(pred_probs, axis=1)\n\nconf_mat = confusion_matrix(y_true, y_pred)\nplt.figure(figsize=(6, 5))\nsns.heatmap(conf_mat, annot=True, fmt='d', cmap='Blues')\nplt.xlabel(\"Predicted Label\")\nplt.ylabel(\"True Label\")\nplt.title(\"Confusion Matrix\")\nplt.tight_layout()\nplt.savefig(\"confusion_matrix.png\")\nplt.show()\n\nprint(\"\\nClassification Report:\")\ntarget_names = [real_labels[i] for i in range(NUM_CLASSES)]\nprint(classification_report(y_true, y_pred, target_names=target_names))","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-06T07:36:32.463729Z","iopub.status.idle":"2025-04-06T07:36:32.464055Z","shell.execute_reply":"2025-04-06T07:36:32.463946Z"}},"outputs":[],"execution_count":null}]}