{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","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":13836,"databundleVersionId":1718836,"sourceType":"competition"}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Cassava Leaf Disease Classification\n\n## Members \n- Omar Mohamed Mostafa Dawoud\n- Nidal Gaber\n- Ali Abdelrehiem\n\n**Objective:** Classify cassava leaf images into 5 categories:\n- Cassava Bacterial Blight (CBB)\n- Cassava Brown Streak Disease (CBSD)\n- Cassava Green Mottle (CGM)\n- Cassava Mosaic Disease (CMD)\n- Healthy\n\n## Dataset Split:\n\n- Training: 70%\n- Validation: 20%\n- Test: 10%\n\n","metadata":{}},{"cell_type":"code","source":"import os\nimport json\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom PIL import Image\nimport cv2\nfrom pathlib import Path\nimport random\nfrom tqdm.auto import tqdm\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, random_split\nimport torchvision.transforms as transforms\nimport torchvision.models as models\nfrom torchvision.models import resnet50, efficientnet_b0, ResNet50_Weights, EfficientNet_B0_Weights\n\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nset_seed(42)\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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-14T07:38:28.998585Z","iopub.execute_input":"2025-12-14T07:38:28.998898Z","iopub.status.idle":"2025-12-14T07:38:29.007921Z","shell.execute_reply.started":"2025-12-14T07:38:28.998872Z","shell.execute_reply":"2025-12-14T07:38:29.007184Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from pathlib import Path\n\nDATA_DIR = Path('/kaggle/input/cassava-leaf-disease-classification')\nTRAIN_DIR = DATA_DIR / 'train_images'\nTEST_DIR = DATA_DIR / 'test_images'\nTRAIN_CSV = DATA_DIR / 'train.csv'\n\ntrain_df = pd.read_csv(TRAIN_CSV)\n\nprint(f\"Total training images: {len(train_df)}\")\ntrain_df.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T07:38:32.474492Z","iopub.execute_input":"2025-12-14T07:38:32.475208Z","iopub.status.idle":"2025-12-14T07:38:32.495975Z","shell.execute_reply.started":"2025-12-14T07:38:32.475185Z","shell.execute_reply":"2025-12-14T07:38:32.495270Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## EDA ","metadata":{}},{"cell_type":"code","source":"class_names = {\n    0: 'Cassava Bacterial Blight (CBB)',\n    1: 'Cassava Brown Streak Disease (CBSD)',\n    2: 'Cassava Green Mottle (CGM)',\n    3: 'Cassava Mosaic Disease (CMD)',\n    4: 'Healthy'\n}\n\ntrain_df['class_name'] = train_df['label'].map(class_names)\n\nplt.figure(figsize=(12, 6))\nsns.countplot(data=train_df, x='class_name', palette='viridis')\nplt.title('Class Distribution in Training Data', fontsize=16, fontweight='bold')\nplt.xlabel('Disease Class', fontsize=12)\nplt.ylabel('Count', fontsize=12)\nplt.xticks(rotation=45, ha='right')\nplt.tight_layout()\nplt.show()\n\nprint(\"\\nClass Distribution:\")\nprint(train_df['class_name'].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T07:38:34.970029Z","iopub.execute_input":"2025-12-14T07:38:34.970305Z","iopub.status.idle":"2025-12-14T07:38:35.185442Z","shell.execute_reply.started":"2025-12-14T07:38:34.970286Z","shell.execute_reply":"2025-12-14T07:38:35.184488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig, axes = plt.subplots(5, 5, figsize=(15, 15))\nfig.suptitle('Sample Images from Each Class', fontsize=16, fontweight='bold')\n\nfor idx, (label, name) in enumerate(class_names.items()):\n    class_samples = train_df[train_df['label'] == label].sample(5)\n\n    for i, (_, row) in enumerate(class_samples.iterrows()):\n        img_path = TRAIN_DIR / row['image_id']\n        img = Image.open(img_path)\n        axes[idx, i].imshow(img)\n        axes[idx, i].axis('off')\n        if i == 0:\n            axes[idx, i].set_title(f'{name}', fontsize=10, fontweight='bold')\n\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T07:38:37.689269Z","iopub.execute_input":"2025-12-14T07:38:37.689548Z","iopub.status.idle":"2025-12-14T07:38:40.334875Z","shell.execute_reply.started":"2025-12-14T07:38:37.689525Z","shell.execute_reply":"2025-12-14T07:38:40.334107Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\n\ntrain_data, temp_data = train_test_split(\n    train_df,\n    test_size=0.3,\n    random_state=42,\n    stratify=train_df['label']\n)\n\nval_data, test_data = train_test_split(\n    temp_data,\n    test_size=0.333,  # 0.333 of 30% = 10% of total\n    random_state=42,\n    stratify=temp_data['label']\n)\n\nprint(f\"Total samples: {len(train_df)}\")\nprint(f\"Training samples: {len(train_data)} ({len(train_data)/len(train_df)*100:.1f}%)\")\nprint(f\"Validation samples: {len(val_data)} ({len(val_data)/len(train_df)*100:.1f}%)\")\nprint(f\"Test samples: {len(test_data)} ({len(test_data)/len(train_df)*100:.1f}%)\")\n\nprint(\"\\nClass distribution:\")\nprint(\"Train:\", train_data['label'].value_counts().sort_index().tolist())\nprint(\"Val:  \", val_data['label'].value_counts().sort_index().tolist())\nprint(\"Test: \", test_data['label'].value_counts().sort_index().tolist())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T07:38:49.560270Z","iopub.execute_input":"2025-12-14T07:38:49.560544Z","iopub.status.idle":"2025-12-14T07:38:49.584152Z","shell.execute_reply.started":"2025-12-14T07:38:49.560522Z","shell.execute_reply":"2025-12-14T07:38:49.583416Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Custom Dataset Class with Augmentation","metadata":{}},{"cell_type":"code","source":"class CassavaDataset(Dataset):\n    def __init__(self, dataframe, img_dir, transform=None):\n        self.df = dataframe.reset_index(drop=True)\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        img_path = os.path.join(self.img_dir, self.df.loc[idx, 'image_id'])\n        image = cv2.imread(str(img_path))\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)\n        label = self.df.loc[idx, 'label']\n\n        if self.transform:\n            augmented = self.transform(image=image)\n            image = augmented['image']\n\n        return image, label\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T07:38:52.445215Z","iopub.execute_input":"2025-12-14T07:38:52.445486Z","iopub.status.idle":"2025-12-14T07:38:52.451412Z","shell.execute_reply.started":"2025-12-14T07:38:52.445466Z","shell.execute_reply":"2025-12-14T07:38:52.450648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IMG_SIZE = 224\n\ntrain_transform = A.Compose([\n    A.RandomResizedCrop(size=(IMG_SIZE, IMG_SIZE), scale=(0.8, 1.0)),\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.Rotate(limit=30, p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5),\n    A.OneOf([\n        A.GaussNoise(var_limit=(10.0, 50.0)),\n        A.GaussianBlur(blur_limit=3),\n        A.MotionBlur(blur_limit=3),\n    ], p=0.3),\n    A.OneOf([\n        A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2),\n        A.HueSaturationValue(hue_shift_limit=20, sat_shift_limit=30, val_shift_limit=20),\n    ], p=0.3),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ToTensorV2(),\n])\n\nval_transform = A.Compose([\n    A.Resize(height=IMG_SIZE, width=IMG_SIZE),\n    A.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),\n    ToTensorV2(),\n])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T07:38:54.802038Z","iopub.execute_input":"2025-12-14T07:38:54.802781Z","iopub.status.idle":"2025-12-14T07:38:54.815307Z","shell.execute_reply.started":"2025-12-14T07:38:54.802754Z","shell.execute_reply":"2025-12-14T07:38:54.814639Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dataset = CassavaDataset(train_data, TRAIN_DIR, transform=train_transform)\nval_dataset = CassavaDataset(val_data, TRAIN_DIR, transform=val_transform)\ntest_dataset = CassavaDataset(test_data, TRAIN_DIR, transform=val_transform)\n\nprint(f\"Train dataset: {len(train_dataset)} images\")\nprint(f\"Validation dataset: {len(val_dataset)} images\")\nprint(f\"Test dataset: {len(test_dataset)} images\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T06:25:48.056978Z","iopub.execute_input":"2025-12-14T06:25:48.057277Z","iopub.status.idle":"2025-12-14T06:25:48.065650Z","shell.execute_reply.started":"2025-12-14T06:25:48.057257Z","shell.execute_reply":"2025-12-14T06:25:48.064827Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"BATCH_SIZE = 32\n\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=True,\n    num_workers=0,  \n    pin_memory=True\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=0,  \n    pin_memory=True\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=0,  \n    pin_memory=True\n)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T05:45:06.203432Z","iopub.execute_input":"2025-12-14T05:45:06.204056Z","iopub.status.idle":"2025-12-14T05:45:06.210330Z","shell.execute_reply.started":"2025-12-14T05:45:06.204023Z","shell.execute_reply":"2025-12-14T05:45:06.209612Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model Architectures - Experiments","metadata":{}},{"cell_type":"code","source":"# Experiment 1: Custom CNN from Scratch\nclass CustomCNN(nn.Module):\n    def __init__(self, num_classes=5):\n        super(CustomCNN, self).__init__()\n\n        # Convolutional layers\n        self.conv1 = nn.Conv2d(3, 64, kernel_size=3, padding=1)\n        self.bn1 = nn.BatchNorm2d(64)\n        self.conv2 = nn.Conv2d(64, 128, kernel_size=3, padding=1)\n        self.bn2 = nn.BatchNorm2d(128)\n        self.conv3 = nn.Conv2d(128, 256, kernel_size=3, padding=1)\n        self.bn3 = nn.BatchNorm2d(256)\n        self.conv4 = nn.Conv2d(256, 512, kernel_size=3, padding=1)\n        self.bn4 = nn.BatchNorm2d(512)\n\n        self.pool = nn.MaxPool2d(2, 2)\n        self.dropout = nn.Dropout(0.5)\n        self.global_avg_pool = nn.AdaptiveAvgPool2d(1)\n\n        # Fully connected layers\n        self.fc1 = nn.Linear(512, 256)\n        self.fc2 = nn.Linear(256, num_classes)\n\n    def forward(self, x):\n        x = self.pool(F.relu(self.bn1(self.conv1(x))))\n        x = self.pool(F.relu(self.bn2(self.conv2(x))))\n        x = self.pool(F.relu(self.bn3(self.conv3(x))))\n        x = self.pool(F.relu(self.bn4(self.conv4(x))))\n\n        x = self.global_avg_pool(x)\n        x = x.view(x.size(0), -1)\n\n        x = self.dropout(F.relu(self.fc1(x)))\n        x = self.fc2(x)\n\n        return x\n\n# Experiment 2: ResNet50 with Transfer Learning\nclass ResNet50Model(nn.Module):\n    def __init__(self, num_classes=5, pretrained=True):\n        super(ResNet50Model, self).__init__()\n        if pretrained:\n            self.model = resnet50(weights=ResNet50_Weights.IMAGENET1K_V2)\n        else:\n            self.model = resnet50(weights=None)\n\n        num_features = self.model.fc.in_features\n        self.model.fc = nn.Sequential(\n            nn.Dropout(0.5),\n            nn.Linear(num_features, 512),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(512, num_classes)\n        )\n\n    def forward(self, x):\n        return self.model(x)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T07:38:57.279772Z","iopub.execute_input":"2025-12-14T07:38:57.280055Z","iopub.status.idle":"2025-12-14T07:38:57.290529Z","shell.execute_reply.started":"2025-12-14T07:38:57.280035Z","shell.execute_reply":"2025-12-14T07:38:57.289719Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training Functions","metadata":{}},{"cell_type":"code","source":"def train_epoch(model, loader, criterion, optimizer, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    pbar = tqdm(loader, desc='Training')\n    for images, labels in pbar:\n        images, labels = images.to(device), labels.to(device)\n\n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n\n        pbar.set_postfix({'loss': f'{running_loss/(pbar.n+1):.4f}',\n                         'acc': f'{100.*correct/total:.2f}%'})\n\n    return running_loss / len(loader), 100. * correct / total\n\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    with torch.no_grad():\n        pbar = tqdm(loader, desc='Validation')\n        for images, labels in pbar:\n            images, labels = images.to(device), labels.to(device)\n\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n\n            pbar.set_postfix({'loss': f'{running_loss/(pbar.n+1):.4f}',\n                             'acc': f'{100.*correct/total:.2f}%'})\n\n    return running_loss / len(loader), 100. * correct / total\n\ndef train_model(model, train_loader, val_loader, criterion, optimizer, scheduler,\n                num_epochs, device, model_name='model'):\n    best_val_acc = 0.0\n    history = {\n        'train_loss': [], 'train_acc': [],\n        'val_loss': [], 'val_acc': []\n    }\n\n    for epoch in range(num_epochs):\n        print(f'\\nEpoch {epoch+1}/{num_epochs}')\n        print('-' * 50)\n\n        train_loss, train_acc = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_acc = validate(model, val_loader, criterion, device)\n\n        if scheduler:\n            scheduler.step()\n\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\n        print(f'\\nTrain Loss: {train_loss:.4f} | Train Acc: {train_acc:.2f}%')\n        print(f'Val Loss: {val_loss:.4f} | Val Acc: {val_acc:.2f}%')\n\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            torch.save(model.state_dict(), f'{model_name}_best.pth')\n            print(f'✓ Best model saved with validation accuracy: {best_val_acc:.2f}%')\n\n    return history","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T07:39:07.534056Z","iopub.execute_input":"2025-12-14T07:39:07.534741Z","iopub.status.idle":"2025-12-14T07:39:07.544342Z","shell.execute_reply.started":"2025-12-14T07:39:07.534714Z","shell.execute_reply":"2025-12-14T07:39:07.543637Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Experiment 1: Custom CNN from Scratch","metadata":{}},{"cell_type":"code","source":"model1 = CustomCNN(num_classes=5).to(device)\nprint(f\"Total parameters: {sum(p.numel() for p in model1.parameters()):,}\")\n\n#Loss and optimizer\ncriterion = nn.CrossEntropyLoss()\noptimizer1 = torch.optim.Adam(model1.parameters(), lr=0.001)\nscheduler1 = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer1, T_max=10)\n\n# Train\nhistory1 = train_model(\n    model1, train_loader, val_loader, criterion, optimizer1, scheduler1,\n    num_epochs=10, device=device, model_name='custom_cnn'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T07:39:14.673294Z","iopub.execute_input":"2025-12-14T07:39:14.673567Z","iopub.status.idle":"2025-12-14T08:09:46.124138Z","shell.execute_reply.started":"2025-12-14T07:39:14.673546Z","shell.execute_reply":"2025-12-14T08:09:46.123251Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Experiment 2: ResNet50 Transfer Learning","metadata":{}},{"cell_type":"code","source":"model2 = ResNet50Model(num_classes=5, pretrained=True).to(device)\nprint(f\"Total parameters: {sum(p.numel() for p in model2.parameters()):,}\")\n\noptimizer2 = torch.optim.Adam(model2.parameters(), lr=0.0001)\nscheduler2 = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer2, T_max=15)\n\nhistory2 = train_model(\n    model2, train_loader, val_loader, criterion, optimizer2, scheduler2,\n    num_epochs=10, device=device, model_name='resnet50'\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T08:09:47.524298Z","iopub.execute_input":"2025-12-14T08:09:47.524977Z","iopub.status.idle":"2025-12-14T08:47:54.054180Z","shell.execute_reply.started":"2025-12-14T08:09:47.524955Z","shell.execute_reply":"2025-12-14T08:47:54.053430Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Compare Experiments","metadata":{}},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 2, figsize=(16, 5))\n\n#Loss comparison\naxes[0].plot(history1['train_loss'], label='Custom CNN - Train', marker='o')\naxes[0].plot(history1['val_loss'], label='Custom CNN - Val', marker='o')\n\naxes[0].plot(history2['train_loss'], label='ResNet50 - Train', marker='s')\naxes[0].plot(history2['val_loss'], label='ResNet50 - Val', marker='s')\n\naxes[0].set_xlabel('Epoch', fontsize=12)\naxes[0].set_ylabel('Loss', fontsize=12)\naxes[0].set_title('Training and Validation Loss Comparison', fontsize=14, fontweight='bold')\naxes[0].legend()\naxes[0].grid(True, alpha=0.3)\n\n#Accuracy comparison\naxes[1].plot(history1['train_acc'], label='Custom CNN - Train', marker='o')\naxes[1].plot(history1['val_acc'], label='Custom CNN - Val', marker='o')\n\naxes[1].plot(history2['train_acc'], label='ResNet50 - Train', marker='s')\naxes[1].plot(history2['val_acc'], label='ResNet50 - Val', marker='s')\n\naxes[1].set_xlabel('Epoch', fontsize=12)\naxes[1].set_ylabel('Accuracy (%)', fontsize=12)\naxes[1].set_title('Training and Validation Accuracy Comparison', fontsize=14, fontweight='bold')\naxes[1].legend()\naxes[1].grid(True, alpha=0.3)\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T08:48:52.598278Z","iopub.execute_input":"2025-12-14T08:48:52.599049Z","iopub.status.idle":"2025-12-14T08:48:53.033601Z","shell.execute_reply.started":"2025-12-14T08:48:52.599014Z","shell.execute_reply":"2025-12-14T08:48:53.032801Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report, confusion_matrix\n\ndef evaluate_model(model, loader, device, model_name):\n    model.eval()\n    all_preds = []\n    all_labels = []\n\n    with torch.no_grad():\n        for images, labels in tqdm(loader, desc=f'Evaluating {model_name}'):\n            images = images.to(device)\n            outputs = model(images)\n            _, predicted = outputs.max(1)\n\n            all_preds.extend(predicted.cpu().numpy())\n            all_labels.extend(labels.numpy())\n\n    acc = 100. * np.mean(np.array(all_preds) == np.array(all_labels))\n    return all_preds, all_labels, acc\n\n\n#Load best models and evaluate\nmodels_to_eval = [\n    ('Custom CNN', model1, 'custom_cnn_best.pth'),\n    ('ResNet50', model2, 'resnet50_best.pth')\n]\n\nresults = {}\n\nfor model_name, model, checkpoint_path in models_to_eval:\n    if os.path.exists(checkpoint_path):\n        model.load_state_dict(torch.load(checkpoint_path))\n\n    # Evaluate on train set\n    train_preds, train_labels, train_acc = evaluate_model(\n        model, train_loader, device, f'{model_name} (Train)'\n    )\n\n    # Evaluate on val set\n    val_preds, val_labels, val_acc = evaluate_model(\n        model, val_loader, device, f'{model_name} (Val)'\n    )\n\n    # Evaluate on test set\n    test_preds, test_labels, test_acc = evaluate_model(\n        model, test_loader, device, f'{model_name} (Test)'\n    )\n\n    results[model_name] = {\n        'train_acc': train_acc,\n        'val_acc': val_acc,\n        'test_acc': test_acc,\n        'test_preds': test_preds,\n        'test_labels': test_labels\n    }\n\n    print(f\"\\n{'='*60}\")\n    print(f\"{model_name} Results:\")\n    print(f\"{'='*60}\")\n    print(f\"Train Accuracy: {train_acc:.2f}%\")\n    print(f\"Validation Accuracy: {val_acc:.2f}%\")\n    print(f\"Test Accuracy: {test_acc:.2f}%\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T08:49:04.246345Z","iopub.execute_input":"2025-12-14T08:49:04.246607Z","iopub.status.idle":"2025-12-14T08:55:09.466482Z","shell.execute_reply.started":"2025-12-14T08:49:04.246590Z","shell.execute_reply":"2025-12-14T08:55:09.465591Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"summary_df = pd.DataFrame([\n    {\n        'Model': name,\n        'Train Accuracy (%)': f\"{res['train_acc']:.2f}\",\n        'Validation Accuracy (%)': f\"{res['val_acc']:.2f}\",\n        'Test Accuracy (%)': f\"{res['test_acc']:.2f}\"\n    }\n    for name, res in results.items()\n])\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"FINAL RESULTS SUMMARY\")\nprint(\"=\"*80)\nprint(summary_df.to_string(index=False))\nprint(\"=\"*80)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-14T08:58:42.416624Z","iopub.execute_input":"2025-12-14T08:58:42.417293Z","iopub.status.idle":"2025-12-14T08:58:42.424639Z","shell.execute_reply.started":"2025-12-14T08:58:42.417269Z","shell.execute_reply":"2025-12-14T08:58:42.423836Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}