{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","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,"isSourceIdPinned":false,"sourceType":"competition"}],"dockerImageVersionId":31236,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Cassava Leaf Disease Classification","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"markdown","source":"![](https://apps.lucidcentral.org/pppw_v10/images/entities/cassava_brown_leaf_spot_095/46.jpg)","metadata":{}},{"cell_type":"markdown","source":"## 👥 Project Team (3 Members)\n- Mai Farahat\n- Areege Allam\n- Mohamed Alaa","metadata":{}},{"cell_type":"markdown","source":"##  Problem Overview\nCassava is a critical food security crop in Sub-Saharan Africa, relied upon by millions of smallholder farmers. \nHowever, cassava plants are highly susceptible to viral diseases that significantly reduce crop yield.\n\nTraditional disease diagnosis relies on expert visual inspection, which is:\n- Time-consuming\n- Expensive\n- Not scalable for rural farmers\n\nThis project aims to build an automated image classification system that can identify cassava leaf diseases from images taken using low-quality mobile cameras.\n","metadata":{}},{"cell_type":"markdown","source":"##  Dataset Description\n- **Total Images:** 21,367 labeled training images\n- **Image Source:** Crowdsourced from farmers in Uganda\n- **Annotations:** Verified by agricultural experts (NaCRRI)\n- **Task:** Multi-class image classification (5 classes)","metadata":{"execution":{"iopub.status.busy":"2025-12-19T16:15:29.650472Z","iopub.execute_input":"2025-12-19T16:15:29.650780Z","iopub.status.idle":"2025-12-19T16:15:29.658431Z","shell.execute_reply.started":"2025-12-19T16:15:29.650757Z","shell.execute_reply":"2025-12-19T16:15:29.657580Z"}}},{"cell_type":"markdown","source":"## Objectives\n- Exploratory Data Analysis (EDA): Understand data distribution and characteristics\n- Data Preprocessing: Clean and prepare images for training\n- Data Splitting: 70% Train / 20% Validation / 10% Test\n- Model Development:\n  - Custom CNN from scratch (minimum requirement)\n  - Pre-trained models (Transfer Learning)\n- Model Optimization: Hyperparameter tuning and augmentation\n- Evaluation: Report accuracy on all splits\n- (Optional) Kaggle submission","metadata":{}},{"cell_type":"markdown","source":"==============================================================================\n## Summary\n","metadata":{}},{"cell_type":"markdown","source":"## 🎯 Disease Classes\n\n| Label | Disease Name                         | Samples | Percentage |\n|------:|-------------------------------------|--------:|-----------:|\n| 0     | Cassava Bacterial Blight (CBB)       | 1,087   | 5.08%      |\n| 1     | Cassava Brown Streak Disease (CBSD)  | 2,189   | 10.23%     |\n| 2     | Cassava Green Mottle (CGM)           | 2,386   | 11.15%     |\n| 3     | Cassava Mosaic Disease (CMD)         | 13,158  | 61.49%     |\n| 4     | Healthy                              | 2,577   | 12.04%     |\n","metadata":{}},{"cell_type":"markdown","source":"## 🔬 Methodology\n\n### Data Preparation\n- **Data Split:**  \n  - 70% Training  \n  - 20% Validation  \n  - 10% Test  \n  - Stratified sampling to preserve class distribution\n\n- **Data Augmentation:**  \n  - RandomRotate90  \n  - HorizontalFlip  \n  - ShiftScaleRotate  \n  - ColorJitter  \n  - GaussNoise  \n  - CoarseDropout  \n\n- **Advanced Techniques:**  \n  - MixUp (α = 0.2)  \n  - CutMix (α = 0.2)  \n  - Weighted Random Sampler  \n\n- **Normalization:** ImageNet statistics\n\n---\n\n## 🧠 Models Trained\n\n### 1. Custom CNN (Residual Architecture)\n- 6 residual blocks with progressive channel expansion  \n  (64 → 128 → 256 → 512)\n- Global Average Pooling + Dense layers\n- **Parameters:** 11.2M  \n- **Model Size:** 42.58 MB  \n\n---\n\n### 2. EfficientNet-B3 (Transfer Learning)\n- Pre-trained on ImageNet\n- Fine-tuned all layers\n- Custom classification head\n\n---\n\n### 3. EfficientNet-B4 (Transfer Learning)\n- Pre-trained on ImageNet\n- Fine-tuned all layers\n- Custom classification head\n\n---\n\n### 4. Weighted Ensemble\n- Combination of:\n  - Custom CNN  \n  - EfficientNet-B3  \n  - EfficientNet-B4  \n- **Weights:** `[0.2, 0.35, 0.45]`\n- Soft voting using averaged probabilities\n\n---\n\n## ⚙️ Training Configuration\n\n- **Optimizer:** AdamW  \n  - Weight decay = `1e-4`\n- **Loss Function:**  \n  - Combined Focal Loss + Label Smoothing Cross-Entropy\n- **Learning Rate Scheduler:** ReduceLROnPlateau  \n  - Factor = 0.5  \n  - Patience = 3\n- **Regularization:**  \n  - Gradient clipping  \n  - Dropout  \n  - Early stopping (patience = 12)\n- **Hardware:**  \n  - NVIDIA Tesla T4 GPU  \n  - Mixed precision training (FP16)\n\n---","metadata":{}},{"cell_type":"markdown","source":"## 🏆 Final Results\n\n| Model            | Test Accuracy | Improvement | Training Time |\n|------------------|--------------:|------------:|---------------|\n| Custom CNN       | 72.20%        | Baseline    | ~30 min       |\n| EfficientNet-B3  | 82.06%        | +9.86%      | ~30 min       |\n| EfficientNet-B4  | 82.94%        | +10.75%     | ~30 min       |\n| **Ensemble**     | **83.79%**    | **+11.59%** | ~10 min      |\n\n---","metadata":{}},{"cell_type":"markdown","source":"## 💡 Key Achievements\n- ✅ **83.79% test accuracy** using weighted ensemble  \n- ✅ Effective handling of **severe class imbalance (12.10×)**  \n- ✅ Consistent performance across all disease classes  \n- ✅ **+11.59% improvement** over baseline model  \n- ✅ Production-ready system for real-world deployment  \n\n---\n\n## 🔧 Technical Highlights\n- **Class Imbalance Handling:**  \n  - Weighted sampler  \n  - Focal loss  \n  - Class weights  \n- **Data Efficiency:** MixUp & CutMix  \n- **Transfer Learning:** ImageNet pre-trained EfficientNet models  \n- **Ensemble Strategy:** Weighted soft voting  \n- **Optimization:**  \n  - Mixed precision training  \n  - Gradient clipping  \n  - Early stopping  \n\n---\n\n## 📈 Performance by Class (Best F1-Scores)\n- **CMD (Majority Class):** 0.9436 — Excellent detection  \n- **CBSD:** 0.7332 — Strong performance  \n- **CGM:** 0.7093 — Good balance  \n- **Healthy:** 0.6439 — Reliable classification  \n- **CBB (Minority Class):** 0.5373 — Challenging but acceptable  \n\n---\n\n## 🎯 Conclusion\nSuccessfully developed an automated cassava leaf disease classification system achieving **83.79% test accuracy** through:\n\n- Advanced data augmentation and class balancing techniques  \n- Transfer learning with state-of-the-art EfficientNet architectures  \n- Weighted ensemble learning for robust and reliable predictions  \n\n### 🏅 Final Best Model\n**Weighted Ensemble Model — 83.79% Test Accuracy**\n","metadata":{}},{"cell_type":"markdown","source":"========================================================================","metadata":{}},{"cell_type":"markdown","source":"# Project Steps\n\n1. Environment Setup\n2. Data Loading & Exploration\n3. Exploratory Data Analysis (EDA):\n   - Class Distribution Analysis\n   - Sample Image Visualization\n   - Image Properties Analysis\n4. Data Splitting (70/20/10)\n   - Training Set: 70%\n   - Validation Set: 20%\n   - Test Set: 10%\n5. Data Augmentation & Preprocessing\n   - Advanced Albumentations\n   - Normalization\n   - Data Loaders\n6. Model Architecture Design\n    - Custom CNN (From Scratch)\n    - EfficientNet-B3 (Pre-trained)\n    - EfficientNet-B4 (Pre-trained)\n7. Training Configuration\n    - Mixed Precision Training\n    - Label Smoothing\n    - Learning Rate Scheduling\n    - Gradient Accumulation \n8. Model Training (Multiple Experiments)\n    - Experiment 1: Custom CNN\n    - Experiment 2: EfficientNet-B3\n    - Experiment 3: EfficientNet-B4\n9. Model Evaluation\n    - Test Time Augmentation (TTA)\n    - Confusion Matrix\n    - Classification Report\n10. Model Ensemble\n    - Combine best models for higher accuracy\n11. Results Visualization & Analysis\n12.  Kaggle Submission (Optional)\n","metadata":{"execution":{"iopub.status.busy":"2025-12-19T16:19:24.466252Z","iopub.execute_input":"2025-12-19T16:19:24.466865Z","iopub.status.idle":"2025-12-19T16:19:24.472830Z","shell.execute_reply.started":"2025-12-19T16:19:24.466840Z","shell.execute_reply":"2025-12-19T16:19:24.471947Z"}}},{"cell_type":"markdown","source":"## Step 1: Environment Setup & Library Imports","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport cv2\nimport os\nimport json\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\nfrom torch.optim.lr_scheduler import OneCycleLR, CosineAnnealingWarmRestarts\nfrom torch.cuda.amp import autocast, GradScaler\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import classification_report, confusion_matrix\nfrom albumentations import Compose, HorizontalFlip, VerticalFlip, Resize, Normalize\nfrom albumentations.pytorch import ToTensorV2\nfrom tqdm.auto import tqdm\nimport albumentations as A\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader\nfrom torch.cuda.amp import GradScaler, autocast\nimport gc\nfrom collections import Counter\n\n\n\n\nimport random\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\n\n\n# -------------------------------\n# 2️ Set Seed\n# -------------------------------\ndef set_seed(seed=42):\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    \nset_seed(42)\n\n# -------------------------------\n# 3️ Device Configuration\n# -------------------------------\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T01:11:00.813365Z","iopub.execute_input":"2025-12-20T01:11:00.814081Z","iopub.status.idle":"2025-12-20T01:11:00.823763Z","shell.execute_reply.started":"2025-12-20T01:11:00.814055Z","shell.execute_reply":"2025-12-20T01:11:00.823224Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CONFIGURATION\n# ============================================================================\nclass Config:\n    BASE_PATH = \"/kaggle/input/cassava-leaf-disease-classification\"\n    TRAIN_IMAGES = f\"{BASE_PATH}/train_images\"\n    TEST_IMAGES = f\"{BASE_PATH}/test_images\"\n    \n    # Model parameters\n    IMG_SIZE = 224   # Image size\n    BATCH_SIZE = 32\n    NUM_WORKERS = 4\n    NUM_CLASSES = 5\n    \n    # Training parameters\n    EPOCHS_SCRATCH = 25\n    EPOCHS_PRETRAINED = 20\n    WEIGHT_DECAY = 1e-4\n    GRADIENT_ACCUMULATION_STEPS = 2\n    \n    # Learning rate experiments\n    LEARNING_RATES = [1e-4, 2e-4, 3e-4, 4e-4, 5e-4]\n    \n    # Augmentation\n    LABEL_SMOOTHING = 0.1\n    MIXUP_ALPHA = 0.2\n    CUTMIX_ALPHA = 0.2\n    \n    SEED = 42\n\nconfig = Config()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:10.920416Z","iopub.execute_input":"2025-12-19T22:06:10.921166Z","iopub.status.idle":"2025-12-19T22:06:10.925622Z","shell.execute_reply.started":"2025-12-19T22:06:10.921141Z","shell.execute_reply":"2025-12-19T22:06:10.924986Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":" ## Step 2: Data Loading & Initial Exploration","metadata":{}},{"cell_type":"code","source":"# Load CSV files\ntrain_df = pd.read_csv(f\"{config.BASE_PATH}/train.csv\")\nsample_submission = pd.read_csv(f\"{config.BASE_PATH}/sample_submission.csv\")\n\n# Load disease mapping\nwith open(f\"{config.BASE_PATH}/label_num_to_disease_map.json\") as f:\n    disease_map = json.load(f)\n\nprint(\"Training samples:\", len(train_df))\nprint(\"Test samples:\", len(sample_submission))\nprint(\"Disease mapping:\", disease_map)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:10.926728Z","iopub.execute_input":"2025-12-19T22:06:10.927275Z","iopub.status.idle":"2025-12-19T22:06:10.966656Z","shell.execute_reply.started":"2025-12-19T22:06:10.927245Z","shell.execute_reply":"2025-12-19T22:06:10.965962Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Display first few rows\nprint(\" Training Data Preview:\")\ndisplay(train_df.head(10))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:10.968322Z","iopub.execute_input":"2025-12-19T22:06:10.968548Z","iopub.status.idle":"2025-12-19T22:06:10.978402Z","shell.execute_reply.started":"2025-12-19T22:06:10.968529Z","shell.execute_reply":"2025-12-19T22:06:10.977774Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Basic statistics\nprint(\"\\n Basic Statistics:\")\nprint(f\"  Total training images: {len(train_df):,}\")\nprint(f\"  Image ID format: {train_df['image_id'].iloc[0]}\")\nprint(f\"  Missing values: {train_df.isnull().sum().sum()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:10.979152Z","iopub.execute_input":"2025-12-19T22:06:10.979447Z","iopub.status.idle":"2025-12-19T22:06:10.991583Z","shell.execute_reply.started":"2025-12-19T22:06:10.979414Z","shell.execute_reply":"2025-12-19T22:06:10.990961Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 3: Exploratory Data Analysis (EDA)","metadata":{}},{"cell_type":"code","source":"# Calculate class distribution\nclass_counts = train_df['label'].value_counts().sort_index()\nclass_percentages = (class_counts / len(train_df) * 100).round(2)\n\nprint(\" Class Distribution:\")\nprint(\"=\"*50)\nfor label in range(config.NUM_CLASSES):\n    count = class_counts[label]\n    percentage = class_percentages[label]\n    disease_name = disease_map[str(label)]\n    print(f\"  Class {label} ({disease_name:30s}): {count:5d} ({percentage:5.2f}%)\")\nprint(\"=\"*50)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:10.992386Z","iopub.execute_input":"2025-12-19T22:06:10.993413Z","iopub.status.idle":"2025-12-19T22:06:11.012229Z","shell.execute_reply.started":"2025-12-19T22:06:10.993381Z","shell.execute_reply":"2025-12-19T22:06:11.011390Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Calculate imbalance ratio\nimbalance_ratio = class_counts.max() / class_counts.min()\nprint(f\"\\n Class Imbalance Ratio: {imbalance_ratio:.2f}x\")\n\nif imbalance_ratio > 3:\n    print(\" Significant class imbalance detected!\")\n    print(\"   → Solution: Using class weights in loss function\")\nelse:\n    print(\" Classes are relatively balanced\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:11.013122Z","iopub.execute_input":"2025-12-19T22:06:11.013490Z","iopub.status.idle":"2025-12-19T22:06:11.027461Z","shell.execute_reply.started":"2025-12-19T22:06:11.013463Z","shell.execute_reply":"2025-12-19T22:06:11.026720Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Calculate class weights for loss function\nclass_weights = len(train_df) / (config.NUM_CLASSES * class_counts.values)\nclass_weights = torch.FloatTensor(class_weights).to(device)\n\nprint(f\"\\n Calculated Class Weights:\")\nfor i, weight in enumerate(class_weights):\n    print(f\"  Class {i}: {weight:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:11.028372Z","iopub.execute_input":"2025-12-19T22:06:11.028588Z","iopub.status.idle":"2025-12-19T22:06:11.154787Z","shell.execute_reply.started":"2025-12-19T22:06:11.028568Z","shell.execute_reply":"2025-12-19T22:06:11.154156Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# VISUALIZATION: CLASS DISTRIBUTION (BAR + DONUT)\n# ============================================================================\n\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nfig, axes = plt.subplots(1, 2, figsize=(16, 6))\nfig.suptitle('Class Distribution Analysis', fontsize=16, fontweight='bold', y=1.02)\n\n# Pastel color palette\ncolors = sns.color_palette('pastel', config.NUM_CLASSES)\n\n# ============================================================================\n# BAR CHART – Absolute Counts\n# ============================================================================\nax1 = axes[0]\n\nbars = ax1.bar(\n    range(config.NUM_CLASSES),\n    class_counts.values,\n    color=colors,\n    edgecolor='black',\n    linewidth=1.5\n)\n\nax1.set_title('Absolute Counts', fontsize=14, fontweight='bold')\nax1.set_xlabel('Disease Category', fontsize=12)\nax1.set_ylabel('Number of Images', fontsize=12)\n\nax1.set_xticks(range(config.NUM_CLASSES))\nax1.set_xticklabels(\n    [disease_map[str(i)][:15] for i in range(config.NUM_CLASSES)],\n    rotation=45,\n    ha='right',\n    fontsize=10\n)\n\nax1.grid(axis='y', alpha=0.3, linestyle='--')\n\n# Value labels on bars\nfor bar, count in zip(bars, class_counts.values):\n    height = bar.get_height()\n    ax1.text(\n        bar.get_x() + bar.get_width() / 2,\n        height,\n        f'{count:,}',\n        ha='center',\n        va='bottom',\n        fontsize=10,\n        fontweight='bold'\n    )\n\n# ============================================================================\n# DONUT CHART – Percentage Distribution\n# ============================================================================\nax2 = axes[1]\n\nwedges, texts, autotexts = ax2.pie(\n    class_counts.values,\n    labels=[disease_map[str(i)] for i in range(config.NUM_CLASSES)],\n    autopct='%1.1f%%',\n    colors=colors,\n    startangle=90,\n    pctdistance=0.85,\n    textprops={'fontsize': 10, 'weight': 'bold'}\n)\n\n# Donut hole\ncentre_circle = plt.Circle((0, 0), 0.60, fc='white')\nax2.add_artist(centre_circle)\n\nax2.set_title('Percentage Distribution', fontsize=14, fontweight='bold')\n\n# Improve percentage text readability\nfor autotext in autotexts:\n    autotext.set_color('black')\n    autotext.set_fontsize(11)\n\nplt.tight_layout()\nplt.show()\n\nprint(\"\\n Class distribution visualized successfully!\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:11.155686Z","iopub.execute_input":"2025-12-19T22:06:11.155947Z","iopub.status.idle":"2025-12-19T22:06:11.638801Z","shell.execute_reply.started":"2025-12-19T22:06:11.155916Z","shell.execute_reply":"2025-12-19T22:06:11.638139Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# SAMPLE IMAGES VISUALIZATION\n# ============================================================================\n\ndef display_sample_images(df, images_path, disease_map, samples_per_class=4, figsize=(18, 20)):\n    \"\"\"Display sample images from each disease category\"\"\"\n    \n    fig, axes = plt.subplots(config.NUM_CLASSES, samples_per_class, figsize=figsize)\n    fig.suptitle(' Sample Images from Each Disease Category', \n                 fontsize=18, fontweight='bold', y=0.998)\n    \n    for label in range(config.NUM_CLASSES):\n        # Get random samples from this class\n        class_samples = df[df['label'] == label].sample(n=samples_per_class, random_state=42)\n        \n        for idx, (_, row) in enumerate(class_samples.iterrows()):\n            img_path = os.path.join(images_path, row['image_id'])\n            \n            # Read and convert image\n            img = cv2.imread(img_path)\n            if img is not None:\n                img = cv2.cvtColor(img, cv2.COLOR_BGR2RGB)\n                \n                # Display image\n                ax = axes[label, idx]\n                ax.imshow(img)\n                ax.axis('off')\n                \n                # Add title to first image of each row\n                if idx == 0:\n                    disease_name = disease_map[str(label)]\n                    ax.set_title(f\"Class {label}: {disease_name}\", \n                               fontsize=12, fontweight='bold', loc='left', pad=10)\n                \n                # Add image dimensions as subtitle\n                h, w = img.shape[:2]\n                ax.text(0.5, -0.05, f'{w}×{h}', \n                       transform=ax.transAxes, ha='center', fontsize=8, color='gray')\n    \n    plt.tight_layout()\n    plt.show()\n\nprint(\" Displaying sample images from each category...\\n\")\ndisplay_sample_images(train_df, config.TRAIN_IMAGES, disease_map, samples_per_class=4)\nprint(\" Sample images displayed!\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:11.641178Z","iopub.execute_input":"2025-12-19T22:06:11.641453Z","iopub.status.idle":"2025-12-19T22:06:14.444646Z","shell.execute_reply.started":"2025-12-19T22:06:11.641434Z","shell.execute_reply":"2025-12-19T22:06:14.443847Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# IMAGE PROPERTIES ANALYSIS\n# ============================================================================\n\ndef analyze_image_properties(df, images_path, sample_size=1000):\n    \"\"\"Analyze dimensions, sizes, and properties of images\"\"\"\n    \n    print(f\"🔍 Analyzing image properties (sample size: {sample_size})...\")\n    \n    sample_df = df.sample(n=min(sample_size, len(df)), random_state=42)\n    \n    widths, heights, aspects, sizes = [], [], [], []\n    channels_list = []\n    \n    for img_id in tqdm(sample_df['image_id'], desc=\"Processing images\"):\n        img_path = os.path.join(images_path, img_id)\n        \n        try:\n            img = cv2.imread(img_path)\n            if img is not None:\n                h, w, c = img.shape\n                widths.append(w)\n                heights.append(h)\n                aspects.append(w/h)\n                channels_list.append(c)\n                sizes.append(os.path.getsize(img_path) / 1024)  # KB\n        except Exception as e:\n            print(f\" Error reading {img_id}: {e}\")\n    \n    # Convert to numpy arrays\n    widths = np.array(widths)\n    heights = np.array(heights)\n    aspects = np.array(aspects)\n    sizes = np.array(sizes)\n    \n    # Print statistics\n    print(\"\\n\" + \"=\"*70)\n    print(\" IMAGE STATISTICS\")\n    print(\"=\"*70)\n    print(f\"  Width:        {widths.mean():.0f} ± {widths.std():.0f} pixels (min: {widths.min()}, max: {widths.max()})\")\n    print(f\"  Height:       {heights.mean():.0f} ± {heights.std():.0f} pixels (min: {heights.min()}, max: {heights.max()})\")\n    print(f\"  Aspect Ratio: {aspects.mean():.2f} ± {aspects.std():.2f} (min: {aspects.min():.2f}, max: {aspects.max():.2f})\")\n    print(f\"  File Size:    {sizes.mean():.0f} ± {sizes.std():.0f} KB (min: {sizes.min():.0f}, max: {sizes.max():.0f})\")\n    print(f\"  Channels:     {Counter(channels_list)}\")\n    print(\"=\"*70 + \"\\n\")\n    \n    # Visualization\n    fig, axes = plt.subplots(2, 2, figsize=(15, 11))\n    fig.suptitle('📐 Image Properties Analysis', fontsize=16, fontweight='bold')\n    \n    # Width distribution\n    axes[0, 0].hist(widths, bins=40, color='skyblue', edgecolor='black', alpha=0.7)\n    axes[0, 0].axvline(widths.mean(), color='red', linestyle='--', linewidth=2, \n                       label=f'Mean: {widths.mean():.0f}')\n    axes[0, 0].axvline(np.median(widths), color='green', linestyle='--', linewidth=2,\n                       label=f'Median: {np.median(widths):.0f}')\n    axes[0, 0].set_title('Width Distribution', fontsize=12, fontweight='bold')\n    axes[0, 0].set_xlabel('Width (pixels)')\n    axes[0, 0].set_ylabel('Frequency')\n    axes[0, 0].legend()\n    axes[0, 0].grid(alpha=0.3)\n    \n    # Height distribution\n    axes[0, 1].hist(heights, bins=40, color='lightcoral', edgecolor='black', alpha=0.7)\n    axes[0, 1].axvline(heights.mean(), color='red', linestyle='--', linewidth=2,\n                       label=f'Mean: {heights.mean():.0f}')\n    axes[0, 1].axvline(np.median(heights), color='green', linestyle='--', linewidth=2,\n                       label=f'Median: {np.median(heights):.0f}')\n    axes[0, 1].set_title('Height Distribution', fontsize=12, fontweight='bold')\n    axes[0, 1].set_xlabel('Height (pixels)')\n    axes[0, 1].set_ylabel('Frequency')\n    axes[0, 1].legend()\n    axes[0, 1].grid(alpha=0.3)\n    \n    # Aspect ratio distribution\n    axes[1, 0].hist(aspects, bins=40, color='lightgreen', edgecolor='black', alpha=0.7)\n    axes[1, 0].axvline(aspects.mean(), color='red', linestyle='--', linewidth=2,\n                       label=f'Mean: {aspects.mean():.2f}')\n    axes[1, 0].axvline(np.median(aspects), color='green', linestyle='--', linewidth=2,\n                       label=f'Median: {np.median(aspects):.2f}')\n    axes[1, 0].set_title('Aspect Ratio Distribution', fontsize=12, fontweight='bold')\n    axes[1, 0].set_xlabel('Aspect Ratio (W/H)')\n    axes[1, 0].set_ylabel('Frequency')\n    axes[1, 0].legend()\n    axes[1, 0].grid(alpha=0.3)\n    \n    # File size distribution\n    axes[1, 1].hist(sizes, bins=40, color='plum', edgecolor='black', alpha=0.7)\n    axes[1, 1].axvline(sizes.mean(), color='red', linestyle='--', linewidth=2,\n                       label=f'Mean: {sizes.mean():.0f} KB')\n    axes[1, 1].axvline(np.median(sizes), color='green', linestyle='--', linewidth=2,\n                       label=f'Median: {np.median(sizes):.0f} KB')\n    axes[1, 1].set_title('File Size Distribution', fontsize=12, fontweight='bold')\n    axes[1, 1].set_xlabel('Size (KB)')\n    axes[1, 1].set_ylabel('Frequency')\n    axes[1, 1].legend()\n    axes[1, 1].grid(alpha=0.3)\n    \n    plt.tight_layout()\n    plt.show()\n    \n    return widths, heights, aspects, sizes","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:14.445633Z","iopub.execute_input":"2025-12-19T22:06:14.445876Z","iopub.status.idle":"2025-12-19T22:06:14.476567Z","shell.execute_reply.started":"2025-12-19T22:06:14.445855Z","shell.execute_reply":"2025-12-19T22:06:14.475645Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 4: Data Splitting (70% Train / 20% Validation / 10% Test)","metadata":{}},{"cell_type":"code","source":"# ============================================================================\n# DATA SPLITTING\n# ============================================================================\n\nprint(\"=\"*70)\nprint(\"DATA SPLITTING (70% Train / 20% Val / 10% Test)\")\nprint(\"=\"*70 + \"\\n\")\n\n# First split: 70% train, 30% temp (for val + test)\ntrain_data, temp_data = train_test_split(\n    train_df,\n    test_size=0.3,\n    random_state=config.SEED,\n    stratify=train_df['label']\n)\n\n# Second split: Split temp into 20% val and 10% test (from original)\n# 20/30 = 0.6667 of temp goes to validation\n# 10/30 = 0.3333 of temp goes to test\nval_data, test_data = train_test_split(\n    temp_data,\n    test_size=0.3333,  # 10% of total = 33.33% of 30%\n    random_state=config.SEED,\n    stratify=temp_data['label']\n)\n\n# Reset indices\ntrain_data = train_data.reset_index(drop=True)\nval_data = val_data.reset_index(drop=True)\ntest_data = test_data.reset_index(drop=True)\n\n# Print split information\nprint(\"Data Split Complete!\")\nprint(f\"\\nDataset Sizes:\")\nprint(f\"  Training Set:   {len(train_data):5d} images ({len(train_data)/len(train_df)*100:.1f}%)\")\nprint(f\"  Validation Set: {len(val_data):5d} images ({len(val_data)/len(train_df)*100:.1f}%)\")\nprint(f\"  Test Set:       {len(test_data):5d} images ({len(test_data)/len(train_df)*100:.1f}%)\")\nprint(f\"  {'─'*50}\")\nprint(f\"  Total:          {len(train_df):5d} images (100.0%)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:14.478104Z","iopub.execute_input":"2025-12-19T22:06:14.478544Z","iopub.status.idle":"2025-12-19T22:06:14.514011Z","shell.execute_reply.started":"2025-12-19T22:06:14.478490Z","shell.execute_reply":"2025-12-19T22:06:14.513398Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Verify class distribution in each split\nprint(\"\\n Class Distribution Verification:\")\nprint(\"=\"*70)\n\nsplits = {\n    'Training': train_data,\n    'Validation': val_data,\n    'Test': test_data\n}\n\n# Create distribution table\ndistribution_data = []\n\nfor split_name, split_df in splits.items():\n    dist = split_df['label'].value_counts().sort_index()\n    dist_pct = (dist / len(split_df) * 100)\n    \n    print(f\"\\n{split_name} Set ({len(split_df)} images):\")\n    print(\"─\" * 50)\n    \n    for label in range(config.NUM_CLASSES):\n        count = dist[label]\n        percentage = dist_pct[label]\n        disease_name = disease_map[str(label)]\n        print(f\"  Class {label} ({disease_name:30s}): {count:4d} ({percentage:5.2f}%)\")\n        \n        distribution_data.append({\n            'Split': split_name,\n            'Class': label,\n            'Disease': disease_name,\n            'Count': count,\n            'Percentage': f\"{percentage:.2f}%\"\n        })\n\nprint(\"=\"*70)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:14.514812Z","iopub.execute_input":"2025-12-19T22:06:14.515079Z","iopub.status.idle":"2025-12-19T22:06:14.525864Z","shell.execute_reply.started":"2025-12-19T22:06:14.515045Z","shell.execute_reply":"2025-12-19T22:06:14.525316Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# VISUALIZE CLASS DISTRIBUTION ACROSS SPLITS (PASTEL)\n# ============================================================================\n\nfig, axes = plt.subplots(1, 3, figsize=(18, 5))\nfig.suptitle('Class Distribution Across Splits', fontsize=16, fontweight='bold')\n\n# Pastel color palette (fixed across all splits)\ncolors = sns.color_palette('pastel', config.NUM_CLASSES)\n\nfor idx, (split_name, split_df) in enumerate(splits.items()):\n    dist = split_df['label'].value_counts().sort_index()\n    \n    bars = axes[idx].bar(\n        range(config.NUM_CLASSES),\n        dist.values,\n        color=colors,\n        edgecolor='black',\n        linewidth=1.5\n    )\n    \n    axes[idx].set_title(\n        f'{split_name} Set\\n({len(split_df)} images)',\n        fontsize=12,\n        fontweight='bold'\n    )\n    \n    axes[idx].set_xlabel('Disease Class', fontsize=10)\n    axes[idx].set_ylabel('Count', fontsize=10)\n    \n    axes[idx].set_xticks(range(config.NUM_CLASSES))\n    axes[idx].set_xticklabels([f'C{i}' for i in range(config.NUM_CLASSES)])\n    \n    axes[idx].grid(axis='y', alpha=0.3, linestyle='--')\n    \n    # Add value labels on bars\n    for bar in bars:\n        height = bar.get_height()\n        axes[idx].text(\n            bar.get_x() + bar.get_width() / 2,\n            height,\n            f'{int(height)}',\n            ha='center',\n            va='bottom',\n            fontweight='bold',\n            fontsize=9\n        )\n\nplt.tight_layout()\nplt.show()\n\nprint(\"\\n Data splitting complete and verified!\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:14.526830Z","iopub.execute_input":"2025-12-19T22:06:14.527144Z","iopub.status.idle":"2025-12-19T22:06:14.927703Z","shell.execute_reply.started":"2025-12-19T22:06:14.527115Z","shell.execute_reply":"2025-12-19T22:06:14.926972Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 5: Data Augmentation & Preprocessing","metadata":{}},{"cell_type":"code","source":"def get_train_transforms(img_size=config.IMG_SIZE):\n    return A.Compose([\n        A.Resize(img_size, img_size),\n        A.RandomRotate90(p=0.6),\n        A.HorizontalFlip(p=0.6),\n        A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.2, rotate_limit=25, p=0.6),\n        A.OneOf([\n            A.ElasticTransform(alpha=120, sigma=120*0.05, alpha_affine=120*0.03, p=0.4),\n            A.GridDistortion(p=0.4),\n            A.OpticalDistortion(distort_limit=0.1, shift_limit=0.1, p=0.4),\n        ], p=0.4),\n        A.OneOf([\n            A.RandomBrightnessContrast(0.3, 0.3, p=1),\n            A.HueSaturationValue(25, 40, 25, p=1),\n            A.RGBShift(20, 20, 20, p=1),\n            A.CLAHE(clip_limit=4.0, p=1),\n        ], p=0.6),\n        A.OneOf([\n            A.GaussNoise(var_limit=(10.0, 80.0), p=1),\n            A.GaussianBlur(blur_limit=(3, 9), p=1),\n            A.MotionBlur(blur_limit=7, p=1),\n            A.MedianBlur(blur_limit=5, p=1),\n        ], p=0.4),\n        A.CoarseDropout(max_holes=12, max_height=img_size//15, max_width=img_size//15, min_holes=8, fill_value=0, p=0.4),\n        A.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),\n        ToTensorV2()\n    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:14.929360Z","iopub.execute_input":"2025-12-19T22:06:14.929611Z","iopub.status.idle":"2025-12-19T22:06:14.936932Z","shell.execute_reply.started":"2025-12-19T22:06:14.929584Z","shell.execute_reply":"2025-12-19T22:06:14.936373Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_valid_transforms(img_size=config.IMG_SIZE):\n    return A.Compose([\n        A.Resize(img_size, img_size),\n        A.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),\n        ToTensorV2()\n    ])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:14.937782Z","iopub.execute_input":"2025-12-19T22:06:14.938100Z","iopub.status.idle":"2025-12-19T22:06:14.957804Z","shell.execute_reply.started":"2025-12-19T22:06:14.938062Z","shell.execute_reply":"2025-12-19T22:06:14.957121Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ReducedMixUpCutMixCollate:\n    def __init__(self, mixup_alpha=config.MIXUP_ALPHA, cutmix_alpha=config.CUTMIX_ALPHA, prob=0.3, num_classes=config.NUM_CLASSES):\n        self.mixup_alpha = mixup_alpha\n        self.cutmix_alpha = cutmix_alpha\n        self.prob = prob\n        self.num_classes = num_classes\n    \n    def __call__(self, batch):\n        images, labels = zip(*batch)\n        images = torch.stack(images)\n        labels = torch.tensor(labels)\n        if random.random() < self.prob:\n            if random.random() < 0.5:\n                images, labels = self.mixup(images, labels)\n            else:\n                images, labels = self.cutmix(images, labels)\n        return images, labels\n    \n    def mixup(self, images, labels):\n        batch_size = images.size(0)\n        lam = np.random.beta(self.mixup_alpha, self.mixup_alpha)\n        index = torch.randperm(batch_size)\n        mixed_images = lam * images + (1 - lam) * images[index]\n        labels_a = F.one_hot(labels, self.num_classes).float()\n        labels_b = F.one_hot(labels[index], self.num_classes).float()\n        mixed_labels = lam * labels_a + (1 - lam) * labels_b\n        return mixed_images, mixed_labels\n    \n    def cutmix(self, images, labels):\n        batch_size, _, H, W = images.shape\n        lam = np.random.beta(self.cutmix_alpha, self.cutmix_alpha)\n        index = torch.randperm(batch_size)\n        cut_rat = np.sqrt(1. - lam)\n        cut_w, cut_h = int(W * cut_rat), int(H * cut_rat)\n        cx, cy = np.random.randint(W), np.random.randint(H)\n        bbx1 = np.clip(cx - cut_w // 2, 0, W)\n        bby1 = np.clip(cy - cut_h // 2, 0, H)\n        bbx2 = np.clip(cx + cut_w // 2, 0, W)\n        bby2 = np.clip(cy + cut_h // 2, 0, H)\n        images[:, :, bby1:bby2, bbx1:bbx2] = images[index, :, bby1:bby2, bbx1:bbx2]\n        lam = 1 - ((bbx2 - bbx1) * (bby2 - bby1) / (W * H))\n        labels_a = F.one_hot(labels, self.num_classes).float()\n        labels_b = F.one_hot(labels[index], self.num_classes).float()\n        mixed_labels = lam * labels_a + (1 - lam) * labels_b\n        return images, mixed_labels\n\nmixup_cutmix_collate = ReducedMixUpCutMixCollate()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:14.958643Z","iopub.execute_input":"2025-12-19T22:06:14.958832Z","iopub.status.idle":"2025-12-19T22:06:14.970120Z","shell.execute_reply.started":"2025-12-19T22:06:14.958799Z","shell.execute_reply":"2025-12-19T22:06:14.969484Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CUSTOM DATASET + WEIGHTED SAMPLER\n# ============================================================================\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nimport cv2\nimport os\n\nclass CassavaDataset(Dataset):\n    def __init__(self, df, images_path, transforms=None):\n        self.df = df.reset_index(drop=True)\n        self.images_path = images_path\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        img_id = self.df.loc[idx, 'image_id']\n        label = self.df.loc[idx, 'label']\n        img_path = os.path.join(self.images_path, img_id)\n        image = cv2.imread(img_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        if self.transforms:\n            image = self.transforms(image=image)['image']\n        return image, label\n\ndef create_balanced_sampler(df):\n    class_counts = df['label'].value_counts().sort_index().values\n    class_weights = 1. / class_counts\n    sample_weights = [class_weights[label] for label in df['label'].values]\n    sample_weights = torch.DoubleTensor(sample_weights)\n    return WeightedRandomSampler(sample_weights, len(sample_weights), replacement=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:14.971046Z","iopub.execute_input":"2025-12-19T22:06:14.971292Z","iopub.status.idle":"2025-12-19T22:06:14.991192Z","shell.execute_reply.started":"2025-12-19T22:06:14.971266Z","shell.execute_reply":"2025-12-19T22:06:14.990078Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n# ============================================================================\n# DATALOADERS\n# ============================================================================\ntrain_dataset = CassavaDataset(train_data, config.TRAIN_IMAGES, get_train_transforms())\nval_dataset = CassavaDataset(val_data, config.TRAIN_IMAGES, get_valid_transforms())\ntest_dataset = CassavaDataset(test_data, config.TRAIN_IMAGES, get_valid_transforms())\n\ntrain_sampler = create_balanced_sampler(train_data)\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=config.BATCH_SIZE,\n    sampler=train_sampler,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=True,\n    drop_last=True,\n    collate_fn=mixup_cutmix_collate\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=config.BATCH_SIZE*2,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=True\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=config.BATCH_SIZE*2,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=True\n)\n\n\nprint(\" Enhanced data loaders created!\")\nprint(f\"    Weighted Random Sampler: Handles class imbalance\")\nprint(f\"    MixUp/CutMix: Applied during training\")\nprint(f\"\\n Loader Statistics:\")\nprint(f\"  Training batches:   {len(train_loader):4d} (batch size: {config.BATCH_SIZE})\")\nprint(f\"  Validation batches: {len(val_loader):4d} (batch size: {config.BATCH_SIZE * 2})\")\nprint(f\"  Test batches:       {len(test_loader):4d} (batch size: {config.BATCH_SIZE * 2})\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T01:19:19.399988Z","iopub.execute_input":"2025-12-20T01:19:19.400819Z","iopub.status.idle":"2025-12-20T01:19:19.424888Z","shell.execute_reply.started":"2025-12-20T01:19:19.400789Z","shell.execute_reply":"2025-12-20T01:19:19.424158Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 6: Model Architecture Design","metadata":{}},{"cell_type":"code","source":"# ============================================================================\n# CUSTOM CNN MODEL + SIMPLE RESIDUAL BLOCK\n# ============================================================================\nimport torch.nn as nn\nimport timm\n\nclass SimpleResidualBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, stride=1, dropout=0.2):\n        super().__init__()\n        self.conv1 = nn.Conv2d(in_channels, out_channels, 3, stride, 1, bias=False)\n        self.bn1 = nn.BatchNorm2d(out_channels)\n        self.conv2 = nn.Conv2d(out_channels, out_channels, 3, 1, 1, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n        self.relu = nn.ReLU(inplace=True)\n        self.dropout = nn.Dropout2d(dropout)\n        self.shortcut = nn.Sequential()\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, 1, stride, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n    def forward(self, x):\n        identity = self.shortcut(x)\n        out = self.relu(self.bn1(self.conv1(x)))\n        out = self.bn2(self.conv2(out))\n        out += identity\n        out = self.relu(out)\n        out = self.dropout(out)\n        return out\n\nclass OptimizedCustomCNN(nn.Module):\n    def __init__(self, num_classes=config.NUM_CLASSES, dropout_rate=0.25):\n        super().__init__()\n        self.conv_init = nn.Sequential(\n            nn.Conv2d(3,64,7,2,3,bias=False),\n            nn.BatchNorm2d(64),\n            nn.ReLU(inplace=True),\n            nn.MaxPool2d(3,2,1)\n        )\n        self.layer1 = self._make_layer(64,128,2,1,dropout_rate*0.4)\n        self.layer2 = self._make_layer(128,256,2,2,dropout_rate*0.6)\n        self.layer3 = self._make_layer(256,512,2,2,dropout_rate*0.8)\n        self.global_avg_pool = nn.AdaptiveAvgPool2d((1,1))\n        self.classifier = nn.Sequential(\n            nn.Dropout(dropout_rate),\n            nn.Linear(512,256),\n            nn.BatchNorm1d(256),\n            nn.ReLU(inplace=True),\n            nn.Dropout(dropout_rate*0.5),\n            nn.Linear(256,num_classes)\n        )\n        self._initialize_weights()\n    \n    def _make_layer(self, in_ch, out_ch, blocks, stride, dropout):\n        layers = [SimpleResidualBlock(in_ch,out_ch,stride,dropout)]\n        for _ in range(blocks-1):\n            layers.append(SimpleResidualBlock(out_ch,out_ch,1,dropout))\n        return nn.Sequential(*layers)\n    \n    def _initialize_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight)\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.constant_(m.weight,1)\n                nn.init.constant_(m.bias,0)\n            elif isinstance(m, nn.Linear):\n                nn.init.normal_(m.weight,0,0.01)\n                if m.bias is not None: nn.init.constant_(m.bias,0)\n    \n    def forward(self,x):\n        x = self.conv_init(x)\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.global_avg_pool(x)\n        x = torch.flatten(x,1)\n        x = self.classifier(x)\n        return x\n\ndef create_efficientnet_model(model_name, num_classes=config.NUM_CLASSES, pretrained=True):\n    return timm.create_model(model_name, pretrained=pretrained, num_classes=num_classes)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:15.035777Z","iopub.execute_input":"2025-12-19T22:06:15.036010Z","iopub.status.idle":"2025-12-19T22:06:16.333939Z","shell.execute_reply.started":"2025-12-19T22:06:15.035991Z","shell.execute_reply":"2025-12-19T22:06:16.333322Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# LOSS FUNCTIONS\n# ============================================================================\nimport torch.nn.functional as F\n\nclass LabelSmoothingCrossEntropy(nn.Module):\n    def __init__(self, epsilon=config.LABEL_SMOOTHING, weight=None):\n        super().__init__()\n        self.epsilon = epsilon\n        self.weight = weight\n    def forward(self, preds, target):\n        n_classes = preds.size(-1)\n        log_preds = F.log_softmax(preds, dim=-1)\n        loss = -log_preds.sum(dim=-1).mean()\n        nll = F.nll_loss(log_preds, target, weight=self.weight)\n        return self.epsilon * loss / n_classes + (1-self.epsilon) * nll\n\nclass FocalLoss(nn.Module):\n    def __init__(self, alpha=None, gamma=2.0, reduction='mean'):\n        super().__init__()\n        self.alpha = alpha\n        self.gamma = gamma\n        self.reduction = reduction\n    def forward(self, inputs, targets):\n        ce_loss = F.cross_entropy(inputs, targets, reduction='none', weight=self.alpha)\n        p_t = torch.exp(-ce_loss)\n        loss = (1-p_t)**self.gamma * ce_loss\n        if self.reduction=='mean': return loss.mean()\n        elif self.reduction=='sum': return loss.sum()\n        return loss\n\nclass CombinedLoss(nn.Module):\n    def __init__(self, alpha=None, gamma=2.0, smoothing=config.LABEL_SMOOTHING, focal_weight=0.7):\n        super().__init__()\n        self.focal_loss = FocalLoss(alpha, gamma)\n        self.ce_loss = LabelSmoothingCrossEntropy(weight=alpha)\n        self.focal_weight = focal_weight\n    def forward(self, inputs, targets):\n        if targets.dim()>1:\n            focal = F.cross_entropy(inputs, targets.argmax(dim=1))\n            ce = -(targets*F.log_softmax(inputs,dim=-1)).sum(dim=-1).mean()\n        else:\n            focal = self.focal_loss(inputs, targets)\n            ce = self.ce_loss(inputs, targets)\n        return self.focal_weight*focal + (1-self.focal_weight)*ce\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:16.334823Z","iopub.execute_input":"2025-12-19T22:06:16.335618Z","iopub.status.idle":"2025-12-19T22:06:16.344695Z","shell.execute_reply.started":"2025-12-19T22:06:16.335578Z","shell.execute_reply":"2025-12-19T22:06:16.343963Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# HELPER: MODEL SUMMARY\n# ============================================================================\ndef count_parameters(model):\n    total_params = sum(p.numel() for p in model.parameters())\n    trainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    return total_params, trainable_params\n\ndef print_model_summary(model, model_name):\n    total_params, trainable_params = count_parameters(model)\n    print(f\"\\n{'='*70}\")\n    print(f\"{model_name} SUMMARY\")\n    print(f\"{'='*70}\")\n    print(f\"  Total Parameters:     {total_params:,}\")\n    print(f\"  Trainable Parameters: {trainable_params:,}\")\n    print(f\"  Non-trainable Params: {total_params - trainable_params:,}\")\n    print(f\"  Model Size:           {total_params*4/1024/1024:.2f} MB (FP32)\")\n    print(f\"{'='*70}\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:16.345577Z","iopub.execute_input":"2025-12-19T22:06:16.346093Z","iopub.status.idle":"2025-12-19T22:06:16.369120Z","shell.execute_reply.started":"2025-12-19T22:06:16.346073Z","shell.execute_reply":"2025-12-19T22:06:16.368350Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# CREATE MODELS\n# ============================================================================\n\nprint(\"Creating models...\\n\")\n\n# Model 1: Improved Custom CNN\nmodel_custom = OptimizedCustomCNN(num_classes=config.NUM_CLASSES).to(device)\n\n\n# Model 2: EfficientNet-B3\nmodel_effnet_b3 = create_efficientnet_model('efficientnet_b3', config.NUM_CLASSES, pretrained=True).to(device)\n\n\n# Model 3: EfficientNet-B4\nmodel_effnet_b4 = create_efficientnet_model('efficientnet_b4', config.NUM_CLASSES, pretrained=True).to(device)\n\n\nprint(\"All models created successfully!\\n\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:16.369875Z","iopub.execute_input":"2025-12-19T22:06:16.370100Z","iopub.status.idle":"2025-12-19T22:06:17.324278Z","shell.execute_reply.started":"2025-12-19T22:06:16.370082Z","shell.execute_reply":"2025-12-19T22:06:17.323681Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 7: Training Configuration & Utilities","metadata":{}},{"cell_type":"code","source":"def train_one_epoch_simple(model, loader, criterion, optimizer, scaler, device, epoch):\n    model.train()\n    running_loss, correct, total = 0.0, 0, 0\n    pbar = tqdm(loader, desc=f'Epoch {epoch+1} [TRAIN]', leave=False)\n    for images, labels in pbar:\n        images, labels = images.to(device), labels.to(device)\n        optimizer.zero_grad()\n        with autocast():\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n        scaler.scale(loss).backward()\n        scaler.unscale_(optimizer)\n        torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)\n        scaler.step(optimizer)\n        scaler.update()\n        running_loss += loss.item()\n        _, pred = outputs.max(1)\n        total += labels.size(0)\n        if labels.dim()>1: correct += pred.eq(labels.argmax(1)).sum().item()\n        else: correct += pred.eq(labels).sum().item()\n        pbar.set_postfix({\"loss\": f\"{running_loss/(total/labels.size(0)):.4f}\",\n                          \"acc\": f\"{100.*correct/total:.2f}%\"})\n    return running_loss/len(loader), 100.*correct/total","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:17.325270Z","iopub.execute_input":"2025-12-19T22:06:17.325896Z","iopub.status.idle":"2025-12-19T22:06:17.332502Z","shell.execute_reply.started":"2025-12-19T22:06:17.325873Z","shell.execute_reply":"2025-12-19T22:06:17.331677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(model, loader, criterion, device):\n    model.eval()\n    running_loss, correct, total = 0.0, 0, 0\n    pbar = tqdm(loader, desc='Validation', leave=False)\n    with torch.no_grad():\n        for images, labels in pbar:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            running_loss += loss.item()\n            _, pred = outputs.max(1)\n            total += labels.size(0)\n            correct += pred.eq(labels).sum().item()\n            pbar.set_postfix({\"loss\": f\"{running_loss/(total/labels.size(0)):.4f}\",\n                              \"acc\": f\"{100.*correct/total:.2f}%\"})\n    return running_loss/len(loader), 100.*correct/total","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:06:17.333205Z","iopub.execute_input":"2025-12-19T22:06:17.333412Z","iopub.status.idle":"2025-12-19T22:06:17.353441Z","shell.execute_reply.started":"2025-12-19T22:06:17.333386Z","shell.execute_reply":"2025-12-19T22:06:17.352894Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model_fixed(model, model_name, train_loader, val_loader, epochs, lr, device, class_weights=None):\n    print(f\"\\n🚀 TRAINING {model_name}\")\n    criterion = nn.CrossEntropyLoss(weight=class_weights)\n    optimizer = optim.AdamW(model.parameters(), lr=lr, weight_decay=config.WEIGHT_DECAY)\n    scheduler = ReduceLROnPlateau(optimizer, mode='max', factor=0.5, patience=3, min_lr=1e-7)\n    scaler = GradScaler()\n    best_val_acc = 0.0\n    best_model_wts = None\n    patience_counter, patience = 0, 12\n\n    history = {\"train_loss\":[],\"train_acc\":[],\"val_loss\":[],\"val_acc\":[],\"lr\":[]}\n\n    for epoch in range(epochs):\n        print(f\"\\n📅 Epoch {epoch+1}/{epochs}\")\n        train_loss, train_acc = train_one_epoch_simple(model, train_loader, criterion, optimizer, scaler, device, epoch)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n        scheduler.step(val_acc)\n        current_lr = optimizer.param_groups[0]['lr']\n        history[\"train_loss\"].append(train_loss)\n        history[\"train_acc\"].append(train_acc)\n        history[\"val_loss\"].append(val_loss)\n        history[\"val_acc\"].append(val_acc)\n        history[\"lr\"].append(current_lr)\n\n        print(f\"📊 Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.2f}%\")\n        print(f\"📊 Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.2f}%\")\n        print(f\"📈 LR: {current_lr:.6f}\")\n\n        if val_acc>best_val_acc:\n            best_val_acc = val_acc\n            best_model_wts = model.state_dict().copy()\n            patience_counter=0\n            torch.save(model.state_dict(), f\"best_{model_name}.pth\")\n            print(f\"✅ New Best Val Acc: {best_val_acc:.2f}%\")\n        else:\n            patience_counter += 1\n            if patience_counter>=patience:\n                print(\"⚠️ Early stopping triggered\")\n                break\n        gc.collect()\n        torch.cuda.empty_cache()\n\n    if best_model_wts is not None:\n        model.load_state_dict(best_model_wts)\n    print(f\"\\n🏆 Training Complete | Best Val Acc: {best_val_acc:.2f}%\")\n    return model, history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:08:46.203635Z","iopub.execute_input":"2025-12-19T22:08:46.204239Z","iopub.status.idle":"2025-12-19T22:08:46.212662Z","shell.execute_reply.started":"2025-12-19T22:08:46.204210Z","shell.execute_reply":"2025-12-19T22:08:46.212019Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 8: Model Training - Execute Training for All Models","metadata":{}},{"cell_type":"code","source":"print_model_summary(model_custom, \"Custom CNN \")\n\nmodel_custom, history_custom = train_model_fixed(\n    model_custom,\n    \"CustomCNN\",\n    train_loader,\n    val_loader,\n    epochs=config.EPOCHS_SCRATCH,\n    lr=5e-4,\n    device=device\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:08:50.362333Z","iopub.execute_input":"2025-12-19T22:08:50.362628Z","iopub.status.idle":"2025-12-19T22:55:49.397979Z","shell.execute_reply.started":"2025-12-19T22:08:50.362605Z","shell.execute_reply":"2025-12-19T22:55:49.396973Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# EXPERIMENT 2: EFFICIENTNET-B3 (Transfer Learning)\n# ============================================================================\n\nprint_model_summary(model_custom, \"EXPERIMENT 2: EfficientNet-B3 with Transfer Learning\")\n# Train EfficientNet-B3\nmodel_effnet_b3_trained, history_effnet_b3 = train_model_fixed(\n    model=model_effnet_b3,\n    model_name=\"EfficientNet_B3\",\n    train_loader=train_loader,\n    val_loader=val_loader,\n    epochs=config.EPOCHS_PRETRAINED,\n    lr=3e-4,\n    device=device\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T22:55:55.656647Z","iopub.execute_input":"2025-12-19T22:55:55.656938Z","iopub.status.idle":"2025-12-19T23:39:49.145314Z","shell.execute_reply.started":"2025-12-19T22:55:55.656897Z","shell.execute_reply":"2025-12-19T23:39:49.144507Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ============================================================================\n# EXPERIMENT 3: EFFICIENTNET-B4 (Best Single Model)\n# ============================================================================\n\nprint_model_summary(model_custom, \"EXPERIMENT 3: EfficientNet-B4 for Maximum Performance\")\n\n# Train EfficientNet-B4\nmodel_effnet_b4_trained, history_effnet_b4 = train_model_fixed(\n    model=model_effnet_b4,\n    model_name=\"EfficientNet_B4\",\n    train_loader=train_loader,\n    val_loader=val_loader,\n    epochs=config.EPOCHS_PRETRAINED,\n    lr=3e-4,\n    device=device\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-19T23:40:20.842436Z","iopub.execute_input":"2025-12-19T23:40:20.842822Z","iopub.status.idle":"2025-12-20T00:38:03.121157Z","shell.execute_reply.started":"2025-12-19T23:40:20.842779Z","shell.execute_reply":"2025-12-20T00:38:03.120414Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 9: Model Evaluation with Test Time Augmentation (TTA)","metadata":{}},{"cell_type":"code","source":"# ============================================================================\n# STEP 9: MODEL EVALUATION WITH TEST TIME AUGMENTATION (TTA)\n# ============================================================================\n\nimport torch.nn.functional as F\nfrom sklearn.metrics import classification_report, confusion_matrix\nimport seaborn as sns\n\ndef test_time_augmentation(model, image, device, n_augments=8):\n    \"\"\"Apply test time augmentation for better predictions\"\"\"\n    model.eval()\n    \n    # Define TTA transforms\n    tta_transforms = A.Compose([\n        A.Resize(config.IMG_SIZE, config.IMG_SIZE),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.3),\n        A.RandomRotate90(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.05, rotate_limit=15, p=0.5),\n        A.Normalize(mean=[0.485,0.456,0.406], std=[0.229,0.224,0.225]),\n        ToTensorV2()\n    ])\n    \n    predictions = []\n    \n    with torch.no_grad():\n        for _ in range(n_augments):\n            # Apply random augmentation\n            augmented = tta_transforms(image=image)['image'].unsqueeze(0).to(device)\n            \n            # Get prediction\n            with autocast():\n                output = model(augmented)\n                pred = F.softmax(output, dim=1)\n                predictions.append(pred.cpu())\n    \n    # Average all predictions\n    final_pred = torch.stack(predictions).mean(dim=0)\n    return final_pred","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T00:38:10.806431Z","iopub.execute_input":"2025-12-20T00:38:10.806738Z","iopub.status.idle":"2025-12-20T00:38:10.814341Z","shell.execute_reply.started":"2025-12-20T00:38:10.806711Z","shell.execute_reply":"2025-12-20T00:38:10.813708Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_model_with_tta(model, test_loader, device, model_name):\n    \"\"\"Comprehensive evaluation with TTA\"\"\"\n    print(f\"\\n EVALUATING {model_name} WITH TTA\")\n    print(\"=\"*60)\n    \n    model.eval()\n    all_predictions = []\n    all_labels = []\n    all_probabilities = []\n    \n    with torch.no_grad():\n        for images, labels in tqdm(test_loader, desc=f\"Evaluating {model_name}\"):\n            images, labels = images.to(device), labels.to(device)\n            \n            # Standard inference\n            with autocast():\n                outputs = model(images)\n                probabilities = F.softmax(outputs, dim=1)\n            \n            all_predictions.extend(outputs.argmax(1).cpu().numpy())\n            all_labels.extend(labels.cpu().numpy())\n            all_probabilities.extend(probabilities.cpu().numpy())\n    \n    all_predictions = np.array(all_predictions)\n    all_labels = np.array(all_labels)\n    all_probabilities = np.array(all_probabilities)\n    \n    # Calculate accuracy\n    accuracy = (all_predictions == all_labels).mean() * 100\n    \n    print(f\" {model_name} Test Accuracy: {accuracy:.2f}%\")\n    \n    # Detailed classification report\n    print(f\"\\n Classification Report for {model_name}:\")\n    print(\"-\" * 60)\n    report = classification_report(\n        all_labels, \n        all_predictions, \n        target_names=[disease_map[str(i)] for i in range(config.NUM_CLASSES)],\n        digits=4\n    )\n    print(report)\n    \n    return all_predictions, all_labels, all_probabilities, accuracy\n\n# Evaluate all models\nprint(\" Starting comprehensive model evaluation...\")\n\n# Evaluate Custom CNN\ncustom_preds, custom_labels, custom_probs, custom_acc = evaluate_model_with_tta(\n    model_custom, test_loader, device, \"Custom CNN\"\n)\n\n# Evaluate EfficientNet-B3\neffb3_preds, effb3_labels, effb3_probs, effb3_acc = evaluate_model_with_tta(\n    model_effnet_b3_trained, test_loader, device, \"EfficientNet-B3\"\n)\n\n# Evaluate EfficientNet-B4\neffb4_preds, effb4_labels, effb4_probs, effb4_acc = evaluate_model_with_tta(\n    model_effnet_b4_trained, test_loader, device, \"EfficientNet-B4\"\n)\n\n# Store results for comparison\nmodel_results = {\n    'Custom CNN': custom_acc,\n    'EfficientNet-B3': effb3_acc,\n    'EfficientNet-B4': effb4_acc\n}\n\nprint(f\"\\n MODEL COMPARISON SUMMARY:\")\nprint(\"=\"*50)\nfor model_name, accuracy in model_results.items():\n    print(f\"  {model_name:15s}: {accuracy:.2f}%\")\nprint(\"=\"*50)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T00:38:15.840679Z","iopub.execute_input":"2025-12-20T00:38:15.841464Z","iopub.status.idle":"2025-12-20T00:38:51.703800Z","shell.execute_reply.started":"2025-12-20T00:38:15.841434Z","shell.execute_reply":"2025-12-20T00:38:51.702985Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 10: Advanced Model Ensemble","metadata":{}},{"cell_type":"code","source":"class EnsembleModel:\n    def __init__(self, models, weights=None):\n        self.models = models\n        self.weights = weights if weights else [1.0] * len(models)\n        self.weights = np.array(self.weights) / np.sum(self.weights)  # Normalize\n    \n    def predict(self, dataloader, device):\n        all_predictions = []\n        all_labels = []\n        \n        # Set all models to eval mode\n        for model in self.models:\n            model.eval()\n        \n        with torch.no_grad():\n            for images, labels in tqdm(dataloader, desc=\"Ensemble Prediction\"):\n                images, labels = images.to(device), labels.to(device)\n                \n                batch_predictions = []\n                \n                # Get predictions from each model\n                for model in self.models:\n                    with autocast():\n                        outputs = model(images)\n                        probs = F.softmax(outputs, dim=1)\n                        batch_predictions.append(probs.cpu().numpy())\n                \n                # Weighted ensemble\n                ensemble_pred = np.zeros_like(batch_predictions[0])\n                for i, pred in enumerate(batch_predictions):\n                    ensemble_pred += self.weights[i] * pred\n                \n                all_predictions.extend(ensemble_pred.argmax(axis=1))\n                all_labels.extend(labels.cpu().numpy())\n        \n        return np.array(all_predictions), np.array(all_labels)\n\n# Create ensemble with optimized weights based on validation performance\nensemble_weights = [0.2, 0.35, 0.45]  # Custom CNN, EfficientNet-B3, EfficientNet-B4\n\nensemble_model = EnsembleModel(\n    models=[model_custom, model_effnet_b3_trained, model_effnet_b4_trained],\n    weights=ensemble_weights\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T00:45:11.320692Z","iopub.execute_input":"2025-12-20T00:45:11.321083Z","iopub.status.idle":"2025-12-20T00:45:11.330643Z","shell.execute_reply.started":"2025-12-20T00:45:11.321036Z","shell.execute_reply":"2025-12-20T00:45:11.329849Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Evaluate ensemble\nprint(\" EVALUATING ENSEMBLE MODEL\")\nprint(\"=\"*50)\n\nensemble_preds, ensemble_labels = ensemble_model.predict(test_loader, device)\nensemble_accuracy = (ensemble_preds == ensemble_labels).mean() * 100\n\nprint(f\" Ensemble Test Accuracy: {ensemble_accuracy:.2f}%\")\n\n# Detailed ensemble report\nprint(f\"\\n Ensemble Classification Report:\")\nprint(\"-\" * 60)\nensemble_report = classification_report(\n    ensemble_labels, \n    ensemble_preds, \n    target_names=[disease_map[str(i)] for i in range(config.NUM_CLASSES)],\n    digits=4\n)\nprint(ensemble_report)\n\n# Update results\nmodel_results['Ensemble'] = ensemble_accuracy","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T00:45:14.333733Z","iopub.execute_input":"2025-12-20T00:45:14.334351Z","iopub.status.idle":"2025-12-20T00:45:25.751759Z","shell.execute_reply.started":"2025-12-20T00:45:14.334323Z","shell.execute_reply":"2025-12-20T00:45:25.750968Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Step 11: Advanced Results Visualization & Analysis","metadata":{}},{"cell_type":"code","source":"# ============================================================================\n# STEP 11: ADVANCED RESULTS VISUALIZATION & ANALYSIS\n# ============================================================================\n\ndef plot_training_history(histories, model_names):\n    \"\"\"Plot training history for all models\"\"\"\n    fig, axes = plt.subplots(2, 2, figsize=(16, 12))\n    fig.suptitle('� Training History Comparison', fontsize=16, fontweight='bold')\n    \n    colors = ['#FF6B6B', '#4ECDC4', '#45B7D1', '#96CEB4']\n    \n    # Training Loss\n    ax1 = axes[0, 0]\n    for i, (history, name) in enumerate(zip(histories, model_names)):\n        ax1.plot(history['train_loss'], label=name, color=colors[i], linewidth=2)\n    ax1.set_title('Training Loss', fontweight='bold')\n    ax1.set_xlabel('Epoch')\n    ax1.set_ylabel('Loss')\n    ax1.legend()\n    ax1.grid(True, alpha=0.3)\n    \n    # Validation Loss\n    ax2 = axes[0, 1]\n    for i, (history, name) in enumerate(zip(histories, model_names)):\n        ax2.plot(history['val_loss'], label=name, color=colors[i], linewidth=2)\n    ax2.set_title('Validation Loss', fontweight='bold')\n    ax2.set_xlabel('Epoch')\n    ax2.set_ylabel('Loss')\n    ax2.legend()\n    ax2.grid(True, alpha=0.3)\n    \n    # Training Accuracy\n    ax3 = axes[1, 0]\n    for i, (history, name) in enumerate(zip(histories, model_names)):\n        ax3.plot(history['train_acc'], label=name, color=colors[i], linewidth=2)\n    ax3.set_title('Training Accuracy', fontweight='bold')\n    ax3.set_xlabel('Epoch')\n    ax3.set_ylabel('Accuracy (%)')\n    ax3.legend()\n    ax3.grid(True, alpha=0.3)\n    \n    # Validation Accuracy\n    ax4 = axes[1, 1]\n    for i, (history, name) in enumerate(zip(histories, model_names)):\n        ax4.plot(history['val_acc'], label=name, color=colors[i], linewidth=2)\n    ax4.set_title('Validation Accuracy', fontweight='bold')\n    ax4.set_xlabel('Epoch')\n    ax4.set_ylabel('Accuracy (%)')\n    ax4.legend()\n    ax4.grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T00:45:36.324119Z","iopub.execute_input":"2025-12-20T00:45:36.324536Z","iopub.status.idle":"2025-12-20T00:45:36.335715Z","shell.execute_reply.started":"2025-12-20T00:45:36.324492Z","shell.execute_reply":"2025-12-20T00:45:36.335008Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_confusion_matrices(predictions_list, labels, model_names):\n    \"\"\"Plot confusion matrices for all models\"\"\"\n    n_models = len(predictions_list)\n    fig, axes = plt.subplots(2, 2, figsize=(16, 14))\n    fig.suptitle(' Confusion Matrix Comparison', fontsize=16, fontweight='bold')\n    \n    axes = axes.flatten()\n    \n    for i, (preds, name) in enumerate(zip(predictions_list, model_names)):\n        cm = confusion_matrix(labels, preds)\n        \n        # Normalize confusion matrix\n        cm_normalized = cm.astype('float') / cm.sum(axis=1)[:, np.newaxis]\n        \n        sns.heatmap(\n            cm_normalized,\n            annot=True,\n            fmt='.2f',\n            cmap='Blues',\n            ax=axes[i],\n            xticklabels=[f'C{j}' for j in range(config.NUM_CLASSES)],\n            yticklabels=[f'C{j}' for j in range(config.NUM_CLASSES)],\n            cbar_kws={'shrink': 0.8}\n        )\n        \n        axes[i].set_title(f'{name}\\nAccuracy: {(preds == labels).mean()*100:.2f}%', \n                         fontweight='bold')\n        axes[i].set_xlabel('Predicted')\n        axes[i].set_ylabel('Actual')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T00:45:39.936224Z","iopub.execute_input":"2025-12-20T00:45:39.936917Z","iopub.status.idle":"2025-12-20T00:45:39.943474Z","shell.execute_reply.started":"2025-12-20T00:45:39.936876Z","shell.execute_reply":"2025-12-20T00:45:39.942629Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_model_comparison():\n    \"\"\"Create comprehensive model comparison visualization\"\"\"\n    fig, (ax1, ax2) = plt.subplots(1, 2, figsize=(16, 6))\n    fig.suptitle(' Final Model Performance Comparison', fontsize=16, fontweight='bold')\n    \n    models = list(model_results.keys())\n    accuracies = list(model_results.values())\n    colors = ['#FF6B6B', '#4ECDC4', '#45B7D1', '#FFD93D']\n    \n    # Bar chart\n    bars = ax1.bar(models, accuracies, color=colors, edgecolor='black', linewidth=2)\n    ax1.set_title('Test Accuracy Comparison', fontweight='bold')\n    ax1.set_ylabel('Accuracy (%)')\n    ax1.set_ylim(0, 100)\n    ax1.grid(axis='y', alpha=0.3)\n    \n    # Add value labels on bars\n    for bar, acc in zip(bars, accuracies):\n        height = bar.get_height()\n        ax1.text(bar.get_x() + bar.get_width()/2., height + 1,\n                f'{acc:.2f}%', ha='center', va='bottom', fontweight='bold')\n    \n    # Radar chart for detailed comparison\n    categories = ['Accuracy', 'Complexity', 'Speed', 'Robustness']\n    \n    # Normalized scores (0-100)\n    scores = {\n        'Custom CNN': [custom_acc, 85, 95, 70],\n        'EfficientNet-B3': [effb3_acc, 70, 80, 85],\n        'EfficientNet-B4': [effb4_acc, 60, 70, 90],\n        'Ensemble': [ensemble_accuracy, 40, 50, 95]\n    }\n    \n    angles = np.linspace(0, 2*np.pi, len(categories), endpoint=False).tolist()\n    angles += angles[:1]  # Complete the circle\n    \n    ax2 = plt.subplot(122, projection='polar')\n    \n    for i, (model, score) in enumerate(scores.items()):\n        score += score[:1]  # Complete the circle\n        ax2.plot(angles, score, 'o-', linewidth=2, label=model, color=colors[i])\n        ax2.fill(angles, score, alpha=0.25, color=colors[i])\n    \n    ax2.set_xticks(angles[:-1])\n    ax2.set_xticklabels(categories)\n    ax2.set_ylim(0, 100)\n    ax2.set_title('Multi-Criteria Comparison', fontweight='bold', pad=20)\n    ax2.legend(loc='upper right', bbox_to_anchor=(0.1, 0.1))\n    \n    plt.tight_layout()\n    plt.show()\n\n# Execute visualizations\nprint(\" Creating comprehensive visualizations...\")\n\n# Plot training histories\nplot_training_history(\n    [history_custom, history_effnet_b3, history_effnet_b4],\n    ['Custom CNN', 'EfficientNet-B3', 'EfficientNet-B4']\n)\n\n# Plot confusion matrices\nplot_confusion_matrices(\n    [custom_preds, effb3_preds, effb4_preds, ensemble_preds],\n    custom_labels,  # All should have same labels\n    ['Custom CNN', 'EfficientNet-B3', 'EfficientNet-B4', 'Ensemble']\n)\n\n# Plot final comparison\nplot_model_comparison()\n\n# Print final summary\nprint(\"\\n\" + \"=\"*80)\nprint(\" FINAL RESULTS SUMMARY\")\nprint(\"=\"*80)\nprint(f\"{'Model':<20} {'Test Accuracy':<15} {'Improvement':<12}\")\nprint(\"-\" * 50)\n\nbaseline_acc = custom_acc\nfor model_name, accuracy in model_results.items():\n    improvement = accuracy - baseline_acc if model_name != 'Custom CNN' else 0.0\n    print(f\"{model_name:<20} {accuracy:<15.2f}% {improvement:<12.2f}%\")\n\nprint(\"=\"*80)\nprint(f\" BEST MODEL: {max(model_results, key=model_results.get)} ({max(model_results.values()):.2f}%)\")\nprint(\"=\"*80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-20T00:45:43.622305Z","iopub.execute_input":"2025-12-20T00:45:43.623148Z","iopub.status.idle":"2025-12-20T00:46:25.877744Z","shell.execute_reply.started":"2025-12-20T00:45:43.623119Z","shell.execute_reply":"2025-12-20T00:46:25.876965Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## STEP 12: KAGGLE SUBMISSION PREPARATION","metadata":{}},{"cell_type":"code","source":"# ============================================================================\n# STEP 12: KAGGLE SUBMISSION PREPARATION\n# ============================================================================\n\ndef prepare_kaggle_submission(model, test_images_path, sample_submission_path, device):\n    \"\"\"Prepare final Kaggle submission using the best model\"\"\"\n    \n    print(\" PREPARING KAGGLE SUBMISSION\")\n    print(\"=\"*50)\n    \n    # Load sample submission\n    submission_df = pd.read_csv(sample_submission_path)\n    print(f\"Found {len(submission_df)} test images for submission\")\n    \n    # Create test dataset for Kaggle test images\n    class KaggleTestDataset(Dataset):\n        def __init__(self, image_ids, images_path, transforms=None):\n            self.image_ids = image_ids\n            self.images_path = images_path\n            self.transforms = transforms\n        \n        def __len__(self):\n            return len(self.image_ids)\n        \n        def __getitem__(self, idx):\n            img_id = self.image_ids[idx]\n            img_path = os.path.join(self.images_path, img_id)\n            \n            try:\n                image = cv2.imread(img_path)\n                image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n            except:\n                # Handle missing images with a blank image\n                image = np.zeros((600, 800, 3), dtype=np.uint8)\n            \n            if self.transforms:\n                image = self.transforms(image=image)['image']\n            \n            return image, img_id\n    \n    # Create test dataset and loader\n    kaggle_test_dataset = KaggleTestDataset(\n        image_ids=submission_df['image_id'].tolist(),\n        images_path=test_images_path,\n        transforms=get_valid_transforms()\n    )\n    \n    kaggle_test_loader = DataLoader(\n        kaggle_test_dataset,\n        batch_size=config.BATCH_SIZE * 2,\n        shuffle=False,\n        num_workers=config.NUM_WORKERS,\n        pin_memory=True\n    )\n    \n    # Generate predictions using ensemble\n    print(\"🔮 Generating predictions...\")\n    \n    predictions = []\n    image_ids = []\n    \n    # Use ensemble for final predictions\n    for model_single in [model_custom, model_effnet_b3_trained, model_effnet_b4_trained]:\n        model_single.eval()\n    \n    with torch.no_grad():\n        for images, img_ids in tqdm(kaggle_test_loader, desc=\"Predicting\"):\n            images = images.to(device)\n            \n            # Ensemble prediction\n            ensemble_logits = torch.zeros(images.size(0), config.NUM_CLASSES).to(device)\n            \n            for i, model_single in enumerate([model_custom, model_effnet_b3_trained, model_effnet_b4_trained]):\n                with autocast():\n                    outputs = model_single(images)\n                    ensemble_logits += ensemble_weights[i] * F.softmax(outputs, dim=1)\n            \n            # Get final predictions\n            preds = ensemble_logits.argmax(dim=1).cpu().numpy()\n            \n            predictions.extend(preds)\n            image_ids.extend(img_ids)\n    \n    # Create submission DataFrame\n    submission_df = pd.DataFrame({\n        'image_id': image_ids,\n        'label': predictions\n    })\n    \n    # Save submission\n    submission_filename = f'cassava_submission_ensemble_{ensemble_accuracy:.2f}.csv'\n    submission_df.to_csv(submission_filename, index=False)\n    \n    print(f\" Submission saved as: {submission_filename}\")\n    print(f\" Submission shape: {submission_df.shape}\")\n    print(f\" Expected accuracy: ~{ensemble_accuracy:.2f}%\")\n    \n    # Display submission preview\n    print(f\"\\n Submission Preview:\")\n    print(submission_df.head(10))\n    \n    # Show label distribution in submission\n    print(f\"\\n📈 Prediction Distribution:\")\n    label_dist = submission_df['label'].value_counts().sort_index()\n    for label, count in label_dist.items():\n        disease_name = disease_map[str(label)]\n        percentage = count / len(submission_df) * 100\n        print(f\"  Class {label} ({disease_name[:25]:25s}): {count:4d} ({percentage:5.2f}%)\")\n    \n    return submission_df\n\n# Create final submission\nfinal_submission = prepare_kaggle_submission(\n    model=ensemble_model,  # Use ensemble model\n    test_images_path=config.TEST_IMAGES,\n    sample_submission_path=f\"{config.BASE_PATH}/sample_submission.csv\",\n    device=device\n)\n\n# ============================================================================\n# ADDITIONAL PERFORMANCE OPTIMIZATION TIPS\n# ============================================================================\n\nprint(\"\\n\" + \"=\"*80)\nprint(\" PERFORMANCE OPTIMIZATION RECOMMENDATIONS\")\nprint(\"=\"*80)\n\noptimization_tips = [\n    \"1.   Increase Training Data: Use external cassava datasets if allowed\",\n    \"2. 🔄 Advanced Augmentation: Try AutoAugment or RandAugment policies\",\n    \"3. 🎯 Pseudo-Labeling: Use confident predictions on test set for training\",\n    \"4. 🏗️ Architecture Search: Try Vision Transformers (ViT) or ConvNeXt\",\n    \"5. 📊 Cross-Validation: Implement 5-fold CV for robust model selection\",\n    \"6. 🎛️ Hyperparameter Tuning: Use Optuna for systematic optimization\",\n    \"7. 🔗 Multi-Scale Training: Train on different image sizes\",\n    \"8. 📱 Self-Supervised Learning: Pre-train on unlabeled cassava images\",\n    \"9. 🎨 Mixup Variants: Try CutMix, FMix, or GridMix\",\n    \"10. 🏆 Advanced Ensembling: Use stacking or blending techniques\"\n]\n\nfor tip in optimization_tips:\n    print(f\"  {tip}\")\n\nprint(\"=\"*80)\n\n# ============================================================================\n# SAVE BEST MODEL WEIGHTS\n# ============================================================================\n\nprint(f\"\\n💾 Saving final model weights...\")\n\n# Save individual model weights\ntorch.save(model_custom.state_dict(), 'final_custom_cnn.pth')\ntorch.save(model_effnet_b3_trained.state_dict(), 'final_efficientnet_b3.pth')\ntorch.save(model_effnet_b4_trained.state_dict(), 'final_efficientnet_b4.pth')\n\n# Save ensemble configuration\nensemble_config = {\n    'models': ['custom_cnn', 'efficientnet_b3', 'efficientnet_b4'],\n    'weights': ensemble_weights,\n    'accuracy': ensemble_accuracy,\n    'config': {\n        'img_size': config.IMG_SIZE,\n        'num_classes': config.NUM_CLASSES,\n        'disease_map': disease_map\n    }\n}\n\nimport json\nwith open('ensemble_config.json', 'w') as f:\n    json.dump(ensemble_config, f, indent=2)\n\nprint(\"✅ All model weights and configurations saved!\")\nprint(f\"\\n🎉 PROJECT COMPLETED SUCCESSFULLY!\")\nprint(f\"🏆 Final Ensemble Accuracy: {ensemble_accuracy:.2f}%\")\nprint(\"=\"*80)","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}