{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"name":"python","version":"3.11.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":31703,"databundleVersionId":2871752,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":3140514,"sourceType":"datasetVersion","datasetId":1912529}],"dockerImageVersionId":31192,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## 1. Setup and Imports","metadata":{}},{"cell_type":"code","source":"# Core libraries\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom pathlib import Path\nimport os\nimport random\nfrom PIL import Image\n\n# TensorFlow and Keras\nimport tensorflow as tf\nfrom tensorflow import keras\nfrom tensorflow.keras import layers, models\nfrom tensorflow.keras.preprocessing.image import ImageDataGenerator\nfrom tensorflow.keras.applications import (\n    EfficientNetB0, \n    ResNet50, \n    MobileNetV2,\n    VGG16\n)\nfrom tensorflow.keras.callbacks import (\n    EarlyStopping, \n    ReduceLROnPlateau, \n    ModelCheckpoint,\n    TensorBoard\n)\n\n# Metrics\nfrom sklearn.metrics import (\n    classification_report, \n    confusion_matrix, \n    roc_curve, \n    auc,\n    precision_recall_curve\n)\n\n# Visualization\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# Set style\nsns.set_style('whitegrid')\nplt.rcParams['figure.figsize'] = (12, 8)\n\n# Set random seeds for reproducibility\nSEED = 42\nnp.random.seed(SEED)\ntf.random.set_seed(SEED)\nrandom.seed(SEED)\n\nprint(f\"TensorFlow version: {tf.__version__}\")\nprint(f\"GPU Available: {tf.config.list_physical_devices('GPU')}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:14.187473Z","iopub.execute_input":"2026-06-17T00:44:14.187716Z","iopub.status.idle":"2026-06-17T00:44:32.764709Z","shell.execute_reply.started":"2026-06-17T00:44:14.187693Z","shell.execute_reply":"2026-06-17T00:44:32.763834Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Data Loading and Exploration","metadata":{}},{"cell_type":"code","source":"# Set data directory path\nDATA_DIR = '/kaggle/input/binary-cropped-crown-of-thorns-dataset'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:32.766259Z","iopub.execute_input":"2026-06-17T00:44:32.766853Z","iopub.status.idle":"2026-06-17T00:44:32.770483Z","shell.execute_reply.started":"2026-06-17T00:44:32.766832Z","shell.execute_reply":"2026-06-17T00:44:32.769679Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Ddefine class names and paths\nCLASSES = ['cots_crops', 'notcots_crops']\n\n# Count images per class\nclass_counts = {}\nfor class_name in CLASSES:\n    class_path = os.path.join(DATA_DIR, class_name)\n    if os.path.exists(class_path):\n        image_files = [f for f in os.listdir(class_path) if f.endswith(('.jpg', '.jpeg', '.png'))]\n        class_counts[class_name] = len(image_files)\n        print(f\"{class_name}: {class_counts[class_name]} images\")\n    else:\n        print(f\"Warning: {class_path} does not exist\")\n\n# Visualize class distribution\nplt.figure(figsize=(8, 6))\nplt.bar(class_counts.keys(), class_counts.values(), color=['#FF6B6B', '#4ECDC4'])\nplt.title('Class Distribution', fontsize=16, fontweight='bold')\nplt.xlabel('Class', fontsize=12)\nplt.ylabel('Number of Images', fontsize=12)\nplt.grid(axis='y', alpha=0.3)\nfor i, (k, v) in enumerate(class_counts.items()):\n    plt.text(i, v + 50, str(v), ha='center', fontsize=12, fontweight='bold')\nplt.tight_layout()\nplt.show()\n\ntotal_images = sum(class_counts.values())\nprint(f\"\\nTotal images: {total_images}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:32.772050Z","iopub.execute_input":"2026-06-17T00:44:32.772363Z","iopub.status.idle":"2026-06-17T00:44:33.334055Z","shell.execute_reply.started":"2026-06-17T00:44:32.772332Z","shell.execute_reply":"2026-06-17T00:44:33.333311Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize sample images from each class\ndef display_sample_images(data_dir, classes, samples_per_class=5):\n    fig, axes = plt.subplots(len(classes), samples_per_class, figsize=(15, 6))\n    \n    for i, class_name in enumerate(classes):\n        class_path = os.path.join(data_dir, class_name)\n        image_files = [f for f in os.listdir(class_path) if f.endswith(('.jpg', '.jpeg', '.png'))]\n        sample_files = random.sample(image_files, min(samples_per_class, len(image_files)))\n        \n        for j, img_file in enumerate(sample_files):\n            img_path = os.path.join(class_path, img_file)\n            img = Image.open(img_path)\n            \n            if len(classes) > 1:\n                ax = axes[i, j]\n            else:\n                ax = axes[j]\n            \n            ax.imshow(img)\n            ax.axis('off')\n            if j == 0:\n                ax.set_title(f'{class_name}\\n{img.size}', fontsize=10, fontweight='bold')\n            else:\n                ax.set_title(f'{img.size}', fontsize=9)\n    \n    plt.tight_layout()\n    plt.show()\n\ndisplay_sample_images(DATA_DIR, CLASSES, samples_per_class=5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:33.334907Z","iopub.execute_input":"2026-06-17T00:44:33.335178Z","iopub.status.idle":"2026-06-17T00:44:34.856627Z","shell.execute_reply.started":"2026-06-17T00:44:33.335158Z","shell.execute_reply":"2026-06-17T00:44:34.855742Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Analyze image dimensions\ndef analyze_image_dimensions(data_dir, classes, sample_size=100):\n    widths = []\n    heights = []\n    \n    for class_name in classes:\n        class_path = os.path.join(data_dir, class_name)\n        image_files = [f for f in os.listdir(class_path) if f.endswith(('.jpg', '.jpeg', '.png'))]\n        sample_files = random.sample(image_files, min(sample_size, len(image_files)))\n        \n        for img_file in sample_files:\n            img_path = os.path.join(class_path, img_file)\n            img = Image.open(img_path)\n            widths.append(img.size[0])\n            heights.append(img.size[1])\n    \n    fig, axes = plt.subplots(1, 3, figsize=(15, 4))\n    \n    # Width distribution\n    axes[0].hist(widths, bins=30, color='skyblue', edgecolor='black')\n    axes[0].set_title('Image Width Distribution')\n    axes[0].set_xlabel('Width (pixels)')\n    axes[0].set_ylabel('Frequency')\n    \n    # Height distribution\n    axes[1].hist(heights, bins=30, color='lightcoral', edgecolor='black')\n    axes[1].set_title('Image Height Distribution')\n    axes[1].set_xlabel('Height (pixels)')\n    axes[1].set_ylabel('Frequency')\n    \n    # Scatter plot\n    axes[2].scatter(widths, heights, alpha=0.5, color='green')\n    axes[2].set_title('Width vs Height')\n    axes[2].set_xlabel('Width (pixels)')\n    axes[2].set_ylabel('Height (pixels)')\n    \n    plt.tight_layout()\n    plt.show()\n    \n    print(f\"Width - Mean: {np.mean(widths):.1f}, Median: {np.median(widths):.1f}, Std: {np.std(widths):.1f}\")\n    print(f\"Height - Mean: {np.mean(heights):.1f}, Median: {np.median(heights):.1f}, Std: {np.std(heights):.1f}\")\n    \n    return widths, heights\n\nwidths, heights = analyze_image_dimensions(DATA_DIR, CLASSES)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:34.858668Z","iopub.execute_input":"2026-06-17T00:44:34.858918Z","iopub.status.idle":"2026-06-17T00:44:36.267499Z","shell.execute_reply.started":"2026-06-17T00:44:34.858898Z","shell.execute_reply":"2026-06-17T00:44:36.266836Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Data Preprocessing Configuration","metadata":{}},{"cell_type":"code","source":"# Configuration parameters\nIMG_SIZE = (224, 224)  # Standard size for most pre-trained models\nBATCH_SIZE = 32\nEPOCHS = 50\nLEARNING_RATE = 0.001\nVALIDATION_SPLIT = 0.2\n\n# Data augmentation for training\ntrain_datagen = ImageDataGenerator(\n    rescale=1./255,\n    validation_split=VALIDATION_SPLIT,\n    rotation_range=30,\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    brightness_range=[0.8, 1.2]\n)\n\n# Validation data generator (only rescaling)\nvalidation_datagen = ImageDataGenerator(\n    rescale=1./255,\n    validation_split=VALIDATION_SPLIT\n)\n\nprint(f\"Image size: {IMG_SIZE}\")\nprint(f\"Batch size: {BATCH_SIZE}\")\nprint(f\"Validation split: {VALIDATION_SPLIT * 100}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:36.268191Z","iopub.execute_input":"2026-06-17T00:44:36.268381Z","iopub.status.idle":"2026-06-17T00:44:36.274068Z","shell.execute_reply.started":"2026-06-17T00:44:36.268365Z","shell.execute_reply":"2026-06-17T00:44:36.273142Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create data generators\ntrain_generator = train_datagen.flow_from_directory(\n    DATA_DIR,\n    target_size=IMG_SIZE,\n    batch_size=BATCH_SIZE,\n    class_mode='binary',\n    subset='training',\n    shuffle=True,\n    seed=SEED\n)\n\nvalidation_generator = validation_datagen.flow_from_directory(\n    DATA_DIR,\n    target_size=IMG_SIZE,\n    batch_size=BATCH_SIZE,\n    class_mode='binary',\n    subset='validation',\n    shuffle=False,\n    seed=SEED\n)\n\nprint(f\"\\nTraining samples: {train_generator.samples}\")\nprint(f\"Validation samples: {validation_generator.samples}\")\nprint(f\"\\nClass indices: {train_generator.class_indices}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:36.274958Z","iopub.execute_input":"2026-06-17T00:44:36.275189Z","iopub.status.idle":"2026-06-17T00:44:47.082413Z","shell.execute_reply.started":"2026-06-17T00:44:36.275172Z","shell.execute_reply":"2026-06-17T00:44:47.081826Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize augmented images\ndef show_augmented_images(generator, num_images=5):\n    fig, axes = plt.subplots(1, num_images, figsize=(15, 3))\n    \n    batch = next(generator)\n    images = batch[0]\n    labels = batch[1]\n    \n    for i in range(min(num_images, len(images))):\n        axes[i].imshow(images[i])\n        axes[i].set_title(f'Label: {int(labels[i])}')\n        axes[i].axis('off')\n    \n    plt.suptitle('Sample Augmented Training Images', fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.show()\n\nshow_augmented_images(train_generator)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:47.083211Z","iopub.execute_input":"2026-06-17T00:44:47.083532Z","iopub.status.idle":"2026-06-17T00:44:48.220955Z","shell.execute_reply.started":"2026-06-17T00:44:47.083511Z","shell.execute_reply":"2026-06-17T00:44:48.220177Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Model Building","metadata":{}},{"cell_type":"code","source":"# Function to create model with transfer learning\ndef create_transfer_model(base_model_name='EfficientNetB0', img_size=IMG_SIZE):\n    \"\"\"\n    Create a transfer learning model with a pre-trained base.\n    \n    Args:\n        base_model_name: Name of the base model architecture\n        img_size: Input image size\n    \n    Returns:\n        Compiled Keras model\n    \"\"\"\n    # Select base model\n    if base_model_name == 'EfficientNetB0':\n        base_model = EfficientNetB0(weights='imagenet', include_top=False, input_shape=(*img_size, 3))\n    elif base_model_name == 'ResNet50':\n        base_model = ResNet50(weights='imagenet', include_top=False, input_shape=(*img_size, 3))\n    elif base_model_name == 'MobileNetV2':\n        base_model = MobileNetV2(weights='imagenet', include_top=False, input_shape=(*img_size, 3))\n    elif base_model_name == 'VGG16':\n        base_model = VGG16(weights='imagenet', include_top=False, input_shape=(*img_size, 3))\n    else:\n        raise ValueError(f\"Unknown base model: {base_model_name}\")\n    \n    # Freeze base model layers\n    base_model.trainable = False\n    \n    # Build the model\n    model = models.Sequential([\n        base_model,\n        layers.GlobalAveragePooling2D(),\n        layers.BatchNormalization(),\n        layers.Dropout(0.5),\n        layers.Dense(256, activation='relu'),\n        layers.BatchNormalization(),\n        layers.Dropout(0.3),\n        layers.Dense(1, activation='sigmoid')\n    ])\n    \n    # Compile model\n    model.compile(\n        optimizer=keras.optimizers.Adam(learning_rate=LEARNING_RATE),\n        loss='binary_crossentropy',\n        metrics=['accuracy', keras.metrics.Precision(), keras.metrics.Recall(), keras.metrics.AUC()]\n    )\n    \n    return model\n\n# Create model\nMODEL_NAME = 'EfficientNetB0'  # Change to 'ResNet50', 'MobileNetV2', or 'VGG16' as needed\nmodel = create_transfer_model(MODEL_NAME)\n\n# Display model summary\nmodel.summary()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:48.221597Z","iopub.execute_input":"2026-06-17T00:44:48.221864Z","iopub.status.idle":"2026-06-17T00:44:51.220287Z","shell.execute_reply.started":"2026-06-17T00:44:48.221845Z","shell.execute_reply":"2026-06-17T00:44:51.219665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Visualize model architecture\nkeras.utils.plot_model(model, show_shapes=True, show_layer_names=True, dpi=70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:51.220977Z","iopub.execute_input":"2026-06-17T00:44:51.221236Z","iopub.status.idle":"2026-06-17T00:44:51.338993Z","shell.execute_reply.started":"2026-06-17T00:44:51.221209Z","shell.execute_reply":"2026-06-17T00:44:51.337973Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Training","metadata":{}},{"cell_type":"code","source":"# Define callbacks\ncallbacks = [\n    EarlyStopping(\n        monitor='val_loss',\n        patience=10,\n        restore_best_weights=True,\n        verbose=1\n    ),\n    ReduceLROnPlateau(\n        monitor='val_loss',\n        factor=0.5,\n        patience=5,\n        min_lr=1e-7,\n        verbose=1\n    ),\n    ModelCheckpoint(\n        'best_cots_model.keras',\n        monitor='val_accuracy',\n        save_best_only=True,\n        verbose=1\n    )\n]\n\nprint(\"Callbacks configured\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:51.341123Z","iopub.execute_input":"2026-06-17T00:44:51.341431Z","iopub.status.idle":"2026-06-17T00:44:51.347568Z","shell.execute_reply.started":"2026-06-17T00:44:51.341409Z","shell.execute_reply":"2026-06-17T00:44:51.346678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Train the model\nhistory = model.fit(\n    train_generator,\n    epochs=EPOCHS,\n    validation_data=validation_generator,\n    callbacks=callbacks,\n    verbose=1,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-06-17T00:44:51.348604Z","iopub.execute_input":"2026-06-17T00:44:51.349198Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Training Visualization","metadata":{}},{"cell_type":"code","source":"# Plot training history\ndef plot_training_history(history):\n    fig, axes = plt.subplots(2, 2, figsize=(15, 10))\n    \n    # Accuracy\n    axes[0, 0].plot(history.history['accuracy'], label='Train Accuracy', linewidth=2)\n    axes[0, 0].plot(history.history['val_accuracy'], label='Val Accuracy', linewidth=2)\n    axes[0, 0].set_title('Model Accuracy', fontsize=14, fontweight='bold')\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('Accuracy')\n    axes[0, 0].legend()\n    axes[0, 0].grid(alpha=0.3)\n    \n    # Loss\n    axes[0, 1].plot(history.history['loss'], label='Train Loss', linewidth=2)\n    axes[0, 1].plot(history.history['val_loss'], label='Val Loss', linewidth=2)\n    axes[0, 1].set_title('Model Loss', fontsize=14, fontweight='bold')\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('Loss')\n    axes[0, 1].legend()\n    axes[0, 1].grid(alpha=0.3)\n    \n    # Precision\n    axes[1, 0].plot(history.history['precision'], label='Train Precision', linewidth=2)\n    axes[1, 0].plot(history.history['val_precision'], label='Val Precision', linewidth=2)\n    axes[1, 0].set_title('Model Precision', fontsize=14, fontweight='bold')\n    axes[1, 0].set_xlabel('Epoch')\n    axes[1, 0].set_ylabel('Precision')\n    axes[1, 0].legend()\n    axes[1, 0].grid(alpha=0.3)\n    \n    # Recall\n    axes[1, 1].plot(history.history['recall'], label='Train Recall', linewidth=2)\n    axes[1, 1].plot(history.history['val_recall'], label='Val Recall', linewidth=2)\n    axes[1, 1].set_title('Model Recall', fontsize=14, fontweight='bold')\n    axes[1, 1].set_xlabel('Epoch')\n    axes[1, 1].set_ylabel('Recall')\n    axes[1, 1].legend()\n    axes[1, 1].grid(alpha=0.3)\n    \n    plt.tight_layout()\n    plt.show()\n\nplot_training_history(history)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Model Evaluation","metadata":{}},{"cell_type":"code","source":"# Evaluate on validation set\nprint(\"Evaluating model on validation set...\\n\")\nresults = model.evaluate(validation_generator, verbose=1)\n\nprint(\"\\n\" + \"=\"*50)\nprint(\"VALIDATION RESULTS\")\nprint(\"=\"*50)\nmetric_names = model.metrics_names\nfor name, value in zip(metric_names, results):\n    print(f\"{name.capitalize()}: {value:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Get predictions\nvalidation_generator.reset()\ny_pred_proba = model.predict(validation_generator, verbose=1)\ny_pred = (y_pred_proba > 0.5).astype(int)\ny_true = validation_generator.classes\n\nprint(f\"Predictions shape: {y_pred.shape}\")\nprint(f\"True labels shape: {y_true.shape}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Confusion Matrix\ncm = confusion_matrix(y_true, y_pred)\n\nplt.figure(figsize=(8, 6))\nsns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n            xticklabels=CLASSES, yticklabels=CLASSES,\n            cbar_kws={'label': 'Count'})\nplt.title('Confusion Matrix', fontsize=16, fontweight='bold')\nplt.ylabel('True Label', fontsize=12)\nplt.xlabel('Predicted Label', fontsize=12)\nplt.tight_layout()\nplt.show()\n\nprint(\"\\nConfusion Matrix:\")\nprint(cm)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ROC Curve and AUC\nfpr, tpr, thresholds = roc_curve(y_true, y_pred_proba)\nroc_auc = auc(fpr, tpr)\n\nplt.figure(figsize=(10, 8))\nplt.plot(fpr, tpr, color='darkorange', lw=2, label=f'ROC curve (AUC = {roc_auc:.4f})')\nplt.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--', label='Random Classifier')\nplt.xlim([0.0, 1.0])\nplt.ylim([0.0, 1.05])\nplt.xlabel('False Positive Rate', fontsize=12)\nplt.ylabel('True Positive Rate', fontsize=12)\nplt.title('Receiver Operating Characteristic (ROC) Curve', fontsize=14, fontweight='bold')\nplt.legend(loc='lower right', fontsize=12)\nplt.grid(alpha=0.3)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Precision-Recall Curve\nprecision, recall, pr_thresholds = precision_recall_curve(y_true, y_pred_proba)\n\nplt.figure(figsize=(10, 8))\nplt.plot(recall, precision, color='blue', lw=2, label='Precision-Recall curve')\nplt.xlabel('Recall', fontsize=12)\nplt.ylabel('Precision', fontsize=12)\nplt.title('Precision-Recall Curve', fontsize=14, fontweight='bold')\nplt.legend(loc='lower left', fontsize=12)\nplt.grid(alpha=0.3)\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Prediction Visualization","metadata":{}},{"cell_type":"code","source":"# Visualize predictions on validation samples\ndef plot_predictions(generator, model, num_images=16):\n    generator.reset()\n    batch = next(generator)\n    images = batch[0]\n    true_labels = batch[1]\n    \n    predictions = model.predict(images, verbose=0)\n    pred_labels = (predictions > 0.5).astype(int).flatten()\n    \n    n = min(num_images, len(images))\n    rows = int(np.sqrt(n))\n    cols = int(np.ceil(n / rows))\n    \n    fig, axes = plt.subplots(rows, cols, figsize=(16, 12))\n    axes = axes.flatten() if n > 1 else [axes]\n    \n    class_names_list = list(CLASSES)\n    \n    for i in range(n):\n        axes[i].imshow(images[i])\n        \n        true_class = class_names_list[int(true_labels[i])]\n        pred_class = class_names_list[pred_labels[i]]\n        confidence = predictions[i][0] if pred_labels[i] == 1 else 1 - predictions[i][0]\n        \n        color = 'green' if pred_labels[i] == int(true_labels[i]) else 'red'\n        \n        axes[i].set_title(\n            f'True: {true_class}\\nPred: {pred_class} ({confidence:.2%})',\n            color=color,\n            fontsize=10,\n            fontweight='bold'\n        )\n        axes[i].axis('off')\n    \n    # Hide unused subplots\n    for i in range(n, len(axes)):\n        axes[i].axis('off')\n    \n    plt.suptitle('Model Predictions on Validation Set', fontsize=16, fontweight='bold')\n    plt.tight_layout()\n    plt.show()\n\nplot_predictions(validation_generator, model, num_images=16)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 10. Save and Export Model","metadata":{}},{"cell_type":"code","source":"# Save the final model\nmodel.save('cots_classifier_final.keras')\nprint(\"Model saved as 'cots_classifier_final.keras'\")\n\n# Also save in TensorFlow SavedModel format\nmodel.save('cots_classifier_savedmodel', save_format='tf')\nprint(\"Model saved in SavedModel format: 'cots_classifier_savedmodel'\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save training history\nhistory_df = pd.DataFrame(history.history)\nhistory_df.to_csv('training_history.csv', index=False)\nprint(\"Training history saved as 'training_history.csv'\")\n\n# Display final metrics\nprint(\"\\n\" + \"=\"*60)\nprint(\"FINAL MODEL PERFORMANCE\")\nprint(\"=\"*60)\nprint(f\"Best Validation Accuracy: {max(history.history['val_accuracy']):.4f}\")\nprint(f\"Best Validation Loss: {min(history.history['val_loss']):.4f}\")\nprint(f\"Best Validation Precision: {max(history.history['val_precision']):.4f}\")\nprint(f\"Best Validation Recall: {max(history.history['val_recall']):.4f}\")\nprint(f\"Best Validation AUC: {max(history.history['val_auc']):.4f}\")\nprint(\"=\"*60)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}