{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","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":[{"sourceType":"datasetVersion","sourceId":23812,"datasetId":17810,"databundleVersionId":23851}],"dockerImageVersionId":31329,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Pneumonia Detection from Chest X-Ray Images","metadata":{}},{"cell_type":"markdown","source":"## 1. Import Libraries & Setup\n\nThis section imports the required packages, sets the random seed for reproducibility, checks GPU availability, and defines the global configuration values used later in training and evaluation.","metadata":{}},{"cell_type":"code","source":"# Standard library utilities\nimport os\nimport random\nfrom pathlib import Path\nfrom collections import defaultdict\n\n# Data handling and visualization\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm.auto import tqdm\n\n# Scikit-learn utilities for splitting and evaluation\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import (\n    accuracy_score,\n    precision_score,\n    recall_score,\n    f1_score,\n    roc_auc_score,\n    confusion_matrix,\n    roc_curve,\n    classification_report,\n)\n\n# Albumentations for image augmentation and tensor conversion\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\n\n# PyTorch core modules\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader, WeightedRandomSampler\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\n# timm is used to load pretrained ResNet50 and EfficientNet-B3\nimport timm\n\n# Plotting theme for cleaner visuals in presentation\nsns.set_theme(style=\"whitegrid\", context=\"notebook\")\n\n# Global configuration values\nSEED = 42\nIMG_SIZE = 224\nBATCH_SIZE = 32\nNUM_EPOCHS = 4\nNUM_WORKERS = 2\nLEARNING_RATE = 1e-4\nWEIGHT_DECAY = 1e-4\nEARLY_STOPPING_PATIENCE = 2\nCLASS_NAMES = [\"NORMAL\", \"PNEUMONIA\"]\nCLASS_TO_IDX = {name: idx for idx, name in enumerate(CLASS_NAMES)}\nIMAGENET_MEAN = (0.485, 0.456, 0.406)\nIMAGENET_STD = (0.229, 0.224, 0.225)\n\n# This function fixes randomness as much as possible for reproducible results.\ndef seed_everything(seed: int = 42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(SEED)\n\n# Select GPU if available; otherwise fall back to CPU.\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nUSE_AMP = False\nprint(f\"Using device: {device}\")\nif device.type == \"cuda\":\n    gpu_name = torch.cuda.get_device_name(0)\n    gpu_capability = torch.cuda.get_device_capability(0)\n    print(gpu_name)\n    print(f\"CUDA capability: {gpu_capability}\")\n    # Enable mixed precision only on GPUs where it is reliably supported.\n    USE_AMP = gpu_capability[0] >= 7\n    print(f\"Mixed precision enabled: {USE_AMP}\")\n\nprint(f\"Training image size: {IMG_SIZE}\")\nprint(f\"Epochs per model: {NUM_EPOCHS}\")\n\n# Support both Kaggle path and local path for flexibility.\nKAGGLE_DATASET_ROOT = Path(\"/kaggle/input/datasets/paultimothymooney/chest-xray-pneumonia/chest_xray\")\nLOCAL_DATASET_ROOT = Path(\"./chest_xray\")\nDATASET_ROOT = KAGGLE_DATASET_ROOT if KAGGLE_DATASET_ROOT.exists() else LOCAL_DATASET_ROOT\n\nassert DATASET_ROOT.exists(), (\n    \"Dataset path not found. Download the Kaggle dataset and place the `chest_xray` folder locally, \"\n    \"or run this notebook on Kaggle with the dataset attached.\"\n)\n\nprint(f\"Dataset root: {DATASET_ROOT.resolve()}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T20:31:07.479698Z","iopub.execute_input":"2026-09-02T20:31:07.480383Z","iopub.status.idle":"2026-09-02T20:31:21.550675Z","shell.execute_reply.started":"2026-09-02T20:31:07.480352Z","shell.execute_reply":"2026-09-02T20:31:21.550052Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Load and Explore Dataset (brief EDA)\n\nThe Kaggle dataset provides `train`, `val`, and `test` folders. We inspect the original structure first, because later we will merge `train + val` and create a fresh stratified split for the experiment.","metadata":{}},{"cell_type":"code","source":"# Convert one folder split into a DataFrame containing image paths and labels.\ndef build_split_dataframe(split_name: str) -> pd.DataFrame:\n    split_dir = DATASET_ROOT / split_name\n    records = []\n\n    for class_name in CLASS_NAMES:\n        class_dir = split_dir / class_name\n        for img_path in sorted(class_dir.glob(\"*\")):\n            if img_path.suffix.lower() in {\".jpeg\", \".jpg\", \".png\"}:\n                records.append(\n                    {\n                        \"split\": split_name,\n                        \"label_name\": class_name,\n                        \"label\": CLASS_TO_IDX[class_name],\n                        \"path\": str(img_path),\n                    }\n                )\n\n    return pd.DataFrame(records)\n\n# Read the three original Kaggle folders into DataFrames.\noriginal_train_df = build_split_dataframe(\"train\")\noriginal_val_df = build_split_dataframe(\"val\")\noriginal_test_df = build_split_dataframe(\"test\")\n\n# Combine them temporarily for a quick overview.\nall_original_df = pd.concat([original_train_df, original_val_df, original_test_df], ignore_index=True)\n\ndisplay(all_original_df.head())\n\n# Count images per split and class for a quick EDA summary.\nsummary_df = (\n    all_original_df.groupby([\"split\", \"label_name\"])\n    .size()\n    .reset_index(name=\"count\")\n    .sort_values([\"split\", \"label_name\"])\n)\n\ndisplay(summary_df)\n\n# Visualize the original class distribution.\nplt.figure(figsize=(10, 4))\nsns.countplot(data=all_original_df, x=\"split\", hue=\"label_name\", palette=\"Set2\")\nplt.title(\"Original Dataset Distribution\")\nplt.xlabel(\"Original Folder Split\")\nplt.ylabel(\"Number of Images\")\nplt.show()\n\n# Show a few sample X-ray images from each class.\ndef show_sample_images(df: pd.DataFrame, samples_per_class: int = 4):\n    fig, axes = plt.subplots(len(CLASS_NAMES), samples_per_class, figsize=(14, 6))\n\n    for row_idx, class_name in enumerate(CLASS_NAMES):\n        sample_paths = df[df[\"label_name\"] == class_name][\"path\"].sample(samples_per_class, random_state=SEED).tolist()\n        for col_idx, img_path in enumerate(sample_paths):\n            image = cv2.imread(img_path, cv2.IMREAD_GRAYSCALE)\n            axes[row_idx, col_idx].imshow(image, cmap=\"gray\")\n            axes[row_idx, col_idx].set_title(class_name)\n            axes[row_idx, col_idx].axis(\"off\")\n\n    plt.suptitle(\"Sample Chest X-Ray Images\", fontsize=14)\n    plt.tight_layout()\n    plt.show()\n\nshow_sample_images(all_original_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T20:31:21.551925Z","iopub.execute_input":"2026-09-02T20:31:21.55231Z","iopub.status.idle":"2026-09-02T20:31:24.830696Z","shell.execute_reply.started":"2026-09-02T20:31:21.552286Z","shell.execute_reply":"2026-09-02T20:31:24.829937Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 3. Stratified Split + Handle Imbalance\n\nWe combine the original `train` and `val` folders, then create a new **stratified 85:15 split**. The original `test` folder remains untouched and serves as the final hold-out test set.\n\nTo address class imbalance, we use:\n\n1. **Class weights** inside the loss function\n2. **WeightedRandomSampler** in the training DataLoader","metadata":{}},{"cell_type":"code","source":"# Merge Kaggle's original train and validation folders.\ntrain_val_df = pd.concat([original_train_df, original_val_df], ignore_index=True)\ntest_df = original_test_df.copy().reset_index(drop=True)\n\n# Create a new stratified split so class proportions stay balanced.\ntrain_df, valid_df = train_test_split(\n    train_val_df,\n    test_size=0.15,\n    stratify=train_val_df[\"label\"],\n    random_state=SEED,\n)\n\ntrain_df = train_df.reset_index(drop=True)\nvalid_df = valid_df.reset_index(drop=True)\n\nprint(f\"Train size: {len(train_df)}\")\nprint(f\"Validation size: {len(valid_df)}\")\nprint(f\"Test size: {len(test_df)}\")\n\n# Summarize the new split for presentation.\nresplit_summary_df = pd.concat(\n    [\n        train_df.assign(split=\"train_resplit\"),\n        valid_df.assign(split=\"valid_resplit\"),\n        test_df.assign(split=\"test_holdout\"),\n    ],\n    ignore_index=True,\n)\n\ndisplay(resplit_summary_df.groupby([\"split\", \"label_name\"]).size().reset_index(name=\"count\"))\n\nplt.figure(figsize=(10, 4))\nsns.countplot(data=resplit_summary_df, x=\"split\", hue=\"label_name\", palette=\"Set1\")\nplt.title(\"Distribution After Stratified Re-Split\")\nplt.xlabel(\"Working Split\")\nplt.ylabel(\"Number of Images\")\nplt.show()\n\n# Compute inverse-frequency class weights for the training set.\nclass_counts = train_df[\"label\"].value_counts().sort_index()\nclass_weights = len(train_df) / (len(CLASS_NAMES) * class_counts)\nclass_weights_tensor = torch.tensor(class_weights.values, dtype=torch.float32, device=device)\n\nprint(\"Class counts in the training split:\")\nprint(class_counts)\nprint(\"\\nClass weights used in CrossEntropyLoss:\")\nprint(class_weights.to_dict())\n\n# WeightedRandomSampler makes minority-class samples appear more often during training.\nsample_weight_lookup = {class_idx: class_weights[class_idx] for class_idx in class_counts.index}\nsample_weights = train_df[\"label\"].map(sample_weight_lookup).values.astype(np.float32)\ntrain_sampler = WeightedRandomSampler(\n    weights=sample_weights,\n    num_samples=len(sample_weights),\n    replacement=True,\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T20:31:24.831954Z","iopub.execute_input":"2026-09-02T20:31:24.832214Z","iopub.status.idle":"2026-09-02T20:31:25.284649Z","shell.execute_reply.started":"2026-09-02T20:31:24.832192Z","shell.execute_reply":"2026-09-02T20:31:25.284016Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 4. Data Augmentation & DataLoader\n\nWe use **Albumentations** to improve generalization through image augmentation. The training pipeline includes:\n\n- `Resize`\n- `RandomBrightnessContrast`\n- `HorizontalFlip`\n- `A.Affine` for scaling, translation, and rotation\n\nThis section also defines the custom dataset class and the PyTorch DataLoaders.","metadata":{}},{"cell_type":"code","source":"# Augmentations for the training set.\ntrain_transforms = A.Compose(\n    [\n        A.Resize(IMG_SIZE, IMG_SIZE),\n        A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.5),\n        A.HorizontalFlip(p=0.5),\n        A.Affine(\n            scale=(0.90, 1.10),\n            translate_percent=(-0.05, 0.05),\n            rotate=(-10, 10),\n            border_mode=cv2.BORDER_REPLICATE,\n            p=0.5,\n        ),\n        A.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n        ToTensorV2(),\n    ]\n)\n\n# Validation and test sets use only deterministic preprocessing.\neval_transforms = A.Compose(\n    [\n        A.Resize(IMG_SIZE, IMG_SIZE),\n        A.Normalize(mean=IMAGENET_MEAN, std=IMAGENET_STD),\n        ToTensorV2(),\n    ]\n)\n\n# Custom dataset: reads image path, loads the image, applies transforms, returns image and label.\nclass ChestXRayDataset(Dataset):\n    def __init__(self, dataframe: pd.DataFrame, transforms=None):\n        self.dataframe = dataframe.reset_index(drop=True)\n        self.transforms = transforms\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, idx):\n        row = self.dataframe.iloc[idx]\n        image = cv2.imread(row.path, cv2.IMREAD_GRAYSCALE)\n        image = cv2.cvtColor(image, cv2.COLOR_GRAY2RGB)\n\n        if self.transforms is not None:\n            image = self.transforms(image=image)[\"image\"]\n\n        label = int(row.label)\n        return image, label\n\n# Build datasets for the three working splits.\ntrain_dataset = ChestXRayDataset(train_df, transforms=train_transforms)\nvalid_dataset = ChestXRayDataset(valid_df, transforms=eval_transforms)\ntest_dataset = ChestXRayDataset(test_df, transforms=eval_transforms)\n\n# Build DataLoaders.\n# The training loader uses the weighted sampler to reduce class imbalance.\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=BATCH_SIZE,\n    sampler=train_sampler,\n    num_workers=NUM_WORKERS,\n    pin_memory=(device.type == \"cuda\"),\n)\n\nvalid_loader = DataLoader(\n    valid_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=(device.type == \"cuda\"),\n)\n\ntest_loader = DataLoader(\n    test_dataset,\n    batch_size=BATCH_SIZE,\n    shuffle=False,\n    num_workers=NUM_WORKERS,\n    pin_memory=(device.type == \"cuda\"),\n)\n\nprint(f\"Training batches: {len(train_loader)}\")\nprint(f\"Validation batches: {len(valid_loader)}\")\nprint(f\"Test batches: {len(test_loader)}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T20:31:25.286496Z","iopub.execute_input":"2026-09-02T20:31:25.286881Z","iopub.status.idle":"2026-09-02T20:31:25.302456Z","shell.execute_reply.started":"2026-09-02T20:31:25.286853Z","shell.execute_reply":"2026-09-02T20:31:25.301827Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 5. Load Two Models (ResNet50 & EfficientNet-B3)\n\nBoth models are loaded with **ImageNet pretrained weights** using `timm.create_model(..., pretrained=True, num_classes=2)`. This directly implements transfer learning while adapting each model to binary classification.","metadata":{}},{"cell_type":"code","source":"# Create one of the requested models and move it to the chosen device.\ndef build_model(model_name: str):\n    if model_name == \"resnet50\":\n        model = timm.create_model(\"resnet50\", pretrained=True, num_classes=2)\n    elif model_name == \"efficientnet_b3\":\n        model = timm.create_model(\"efficientnet_b3\", pretrained=True, num_classes=2)\n    else:\n        raise ValueError(f\"Unsupported model name: {model_name}\")\n\n    return model.to(device)\n\n# Quick model-loading check before training.\nresnet50_model = build_model(\"resnet50\")\nefficientnet_b3_model = build_model(\"efficientnet_b3\")\n\nprint(\"Loaded timm model: resnet50\")\nprint(\"Loaded timm model: efficientnet_b3\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T20:31:25.303328Z","iopub.execute_input":"2026-09-02T20:31:25.303526Z","iopub.status.idle":"2026-09-02T20:31:31.740078Z","shell.execute_reply.started":"2026-09-02T20:31:25.303505Z","shell.execute_reply":"2026-09-02T20:31:31.739318Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Training Function\n\nThe training utilities below define the full training and validation workflow. The pipeline includes:\n\n- **CrossEntropyLoss** with class weights\n- **AdamW** optimizer\n- **CosineAnnealingLR** scheduler\n- **Automatic mixed precision** when supported by the GPU\n- **Early stopping** based on validation AUC-ROC\n\nProgress bars are shown at the model, epoch, and batch levels to make notebook execution easier to follow during presentation.","metadata":{}},{"cell_type":"code","source":"# Compute the main binary classification metrics from true labels and predicted probabilities.\ndef compute_metrics(y_true, y_prob, threshold: float = 0.5):\n    y_pred = (y_prob >= threshold).astype(int)\n    return {\n        \"accuracy\": accuracy_score(y_true, y_pred),\n        \"precision\": precision_score(y_true, y_pred, zero_division=0),\n        \"recall\": recall_score(y_true, y_pred, zero_division=0),\n        \"f1\": f1_score(y_true, y_pred, zero_division=0),\n        \"auc_roc\": roc_auc_score(y_true, y_prob),\n    }\n\n# Run one full validation pass and return loss + metrics.\ndef run_validation(model, loader, criterion):\n    model.eval()\n    running_loss = 0.0\n    all_labels = []\n    all_probs = []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            probs = torch.softmax(outputs, dim=1)[:, 1]\n            running_loss += loss.item() * images.size(0)\n            all_labels.extend(labels.cpu().numpy())\n            all_probs.extend(probs.cpu().numpy())\n\n    avg_loss = running_loss / len(loader.dataset)\n    metrics = compute_metrics(np.array(all_labels), np.array(all_probs))\n    return avg_loss, metrics\n\n# Main training loop for one model.\ndef train_model(model, model_name, train_loader, valid_loader, class_weights_tensor, num_epochs=NUM_EPOCHS):\n    criterion = nn.CrossEntropyLoss(weight=class_weights_tensor)\n    optimizer = AdamW(model.parameters(), lr=LEARNING_RATE, weight_decay=WEIGHT_DECAY)\n    scheduler = CosineAnnealingLR(optimizer, T_max=num_epochs)\n    scaler = torch.amp.GradScaler(\"cuda\", enabled=USE_AMP) if device.type == \"cuda\" else None\n\n    history = defaultdict(list)\n    best_state = None\n    best_auc = -np.inf\n    patience_counter = 0\n\n    # Outer tqdm bar: progress across epochs.\n    epoch_iterator = tqdm(range(num_epochs), desc=f\"{model_name} epochs\", leave=True)\n\n    for epoch in epoch_iterator:\n        model.train()\n        running_loss = 0.0\n\n        # Inner tqdm bar: progress across batches inside the current epoch.\n        batch_iterator = tqdm(train_loader, desc=f\"{model_name} epoch {epoch + 1}\", leave=False)\n\n        for images, labels in batch_iterator:\n            images = images.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n\n            optimizer.zero_grad(set_to_none=True)\n\n            # Use AMP only when the GPU can handle it reliably.\n            if USE_AMP:\n                with torch.amp.autocast(\"cuda\", enabled=True):\n                    outputs = model(images)\n                    loss = criterion(outputs, labels)\n\n                scaler.scale(loss).backward()\n                scaler.step(optimizer)\n                scaler.update()\n            else:\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                loss.backward()\n                optimizer.step()\n\n            running_loss += loss.item() * images.size(0)\n            batch_iterator.set_postfix(batch_loss=f\"{loss.item():.4f}\")\n\n        scheduler.step()\n\n        train_loss = running_loss / len(train_loader.dataset)\n        val_loss, val_metrics = run_validation(model, valid_loader, criterion)\n\n        # Save metrics for later plots and comparison.\n        history[\"train_loss\"].append(train_loss)\n        history[\"val_loss\"].append(val_loss)\n        history[\"val_accuracy\"].append(val_metrics[\"accuracy\"])\n        history[\"val_precision\"].append(val_metrics[\"precision\"])\n        history[\"val_recall\"].append(val_metrics[\"recall\"])\n        history[\"val_f1\"].append(val_metrics[\"f1\"])\n        history[\"val_auc_roc\"].append(val_metrics[\"auc_roc\"])\n        history[\"lr\"].append(optimizer.param_groups[0][\"lr\"])\n\n        # Update the epoch progress bar with the most important summary values.\n        epoch_iterator.set_postfix(\n            train_loss=f\"{train_loss:.4f}\",\n            val_loss=f\"{val_loss:.4f}\",\n            val_f1=f\"{val_metrics['f1']:.4f}\",\n            val_auc=f\"{val_metrics['auc_roc']:.4f}\",\n        )\n\n        # Also print a readable epoch summary line.\n        print(\n            f\"[{model_name}] Epoch {epoch + 1:02d}/{num_epochs} | \"\n            f\"Train Loss: {train_loss:.4f} | Val Loss: {val_loss:.4f} | \"\n            f\"Val F1: {val_metrics['f1']:.4f} | Val AUC: {val_metrics['auc_roc']:.4f}\"\n        )\n\n        # Keep the model weights that give the best validation AUC.\n        if val_metrics[\"auc_roc\"] > best_auc:\n            best_auc = val_metrics[\"auc_roc\"]\n            patience_counter = 0\n            best_state = {\n                \"model_state_dict\": model.state_dict(),\n                \"history\": dict(history),\n                \"best_auc\": best_auc,\n            }\n        else:\n            patience_counter += 1\n\n        # Stop early if performance has stopped improving.\n        if patience_counter >= EARLY_STOPPING_PATIENCE:\n            print(f\"Early stopping triggered for {model_name}.\")\n            break\n\n    if best_state is not None:\n        model.load_state_dict(best_state[\"model_state_dict\"])\n\n    return model, dict(history), best_auc\n\n# Generate final probabilities for one trained model on a given loader.\ndef predict_probabilities(model, loader):\n    model.eval()\n    all_labels = []\n    all_probs = []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device, non_blocking=True)\n            outputs = model(images)\n            probs = torch.softmax(outputs, dim=1)[:, 1]\n\n            all_labels.extend(labels.numpy())\n            all_probs.extend(probs.cpu().numpy())\n\n    return np.array(all_labels), np.array(all_probs)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T20:31:31.740985Z","iopub.execute_input":"2026-09-02T20:31:31.741284Z","iopub.status.idle":"2026-09-02T20:31:31.75813Z","shell.execute_reply.started":"2026-09-02T20:31:31.74126Z","shell.execute_reply":"2026-09-02T20:31:31.757284Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 7. Train Both Models\n\nWe train **ResNet50** and **EfficientNet-B3** independently under exactly the same split, augmentation, optimizer, scheduler, and imbalance-handling strategy so the comparison remains fair.","metadata":{}},{"cell_type":"code","source":"trained_artifacts = {}\n\n# Train the two required models one after the other.\nfor model_name in tqdm([\"resnet50\", \"efficientnet_b3\"], desc=\"Models\", leave=True):\n    model = build_model(model_name)\n    model, history, best_auc = train_model(\n        model=model,\n        model_name=model_name,\n        train_loader=train_loader,\n        valid_loader=valid_loader,\n        class_weights_tensor=class_weights_tensor,\n        num_epochs=NUM_EPOCHS,\n    )\n\n    # Save each trained model and its training history for later analysis.\n    trained_artifacts[model_name] = {\n        \"model\": model,\n        \"history\": history,\n        \"best_val_auc\": best_auc,\n    }\n\nprint(\"Training complete.\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T20:31:31.75904Z","iopub.execute_input":"2026-09-02T20:31:31.759342Z","iopub.status.idle":"2026-09-02T20:38:20.252899Z","shell.execute_reply.started":"2026-09-02T20:31:31.759309Z","shell.execute_reply":"2026-09-02T20:38:20.251668Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 8. Evaluation & Comparison (tables + plots)\n\nThis section evaluates **ResNet50** and **EfficientNet-B3** individually on the hold-out test set. The goal is to compare the two models directly using the same set of final metrics and visualizations.","metadata":{}},{"cell_type":"code","source":"evaluation_rows = []\ntest_predictions = {}\n\n# Evaluate each trained model on the untouched hold-out test set.\nfor model_name, artifact in trained_artifacts.items():\n    y_true, y_prob = predict_probabilities(artifact[\"model\"], test_loader)\n    metrics = compute_metrics(y_true, y_prob)\n    y_pred = (y_prob >= 0.5).astype(int)\n\n    evaluation_rows.append(\n        {\n            \"model\": model_name,\n            \"accuracy\": metrics[\"accuracy\"],\n            \"precision\": metrics[\"precision\"],\n            \"recall\": metrics[\"recall\"],\n            \"f1\": metrics[\"f1\"],\n            \"auc_roc\": metrics[\"auc_roc\"],\n        }\n    )\n\n    test_predictions[model_name] = {\n        \"y_true\": y_true,\n        \"y_prob\": y_prob,\n        \"y_pred\": y_pred,\n    }\n\n# Create a compact comparison table.\nevaluation_df = pd.DataFrame(evaluation_rows).sort_values(by=\"auc_roc\", ascending=False).reset_index(drop=True)\ndisplay(evaluation_df.style.format({col: \"{:.4f}\" for col in evaluation_df.columns if col != \"model\"}))\n\n# Plot training curves for both models.\nfig, axes = plt.subplots(1, 2, figsize=(14, 5))\n\nfor model_name, artifact in trained_artifacts.items():\n    history_df = pd.DataFrame(artifact[\"history\"])\n    axes[0].plot(history_df[\"train_loss\"], label=f\"{model_name} train\")\n    axes[0].plot(history_df[\"val_loss\"], linestyle=\"--\", label=f\"{model_name} val\")\n    axes[1].plot(history_df[\"val_auc_roc\"], label=f\"{model_name} val_auc\")\n\naxes[0].set_title(\"Training and Validation Loss\")\naxes[0].set_xlabel(\"Epoch\")\naxes[0].set_ylabel(\"Loss\")\naxes[0].legend()\n\naxes[1].set_title(\"Validation AUC-ROC\")\naxes[1].set_xlabel(\"Epoch\")\naxes[1].set_ylabel(\"AUC-ROC\")\naxes[1].legend()\n\nplt.tight_layout()\nplt.show()\n\n# Plot a bar chart of final test metrics.\nplt.figure(figsize=(10, 5))\nmetrics_for_plot = evaluation_df.melt(id_vars=\"model\", var_name=\"metric\", value_name=\"score\")\nsns.barplot(data=metrics_for_plot, x=\"metric\", y=\"score\", hue=\"model\")\nplt.title(\"Test Set Metric Comparison\")\nplt.ylim(0.0, 1.0)\nplt.xticks(rotation=20)\nplt.show()\n\n# Plot ROC curves for both models.\nplt.figure(figsize=(7, 6))\nfor model_name, preds in test_predictions.items():\n    fpr, tpr, _ = roc_curve(preds[\"y_true\"], preds[\"y_prob\"])\n    plt.plot(fpr, tpr, label=f\"{model_name} (AUC={roc_auc_score(preds['y_true'], preds['y_prob']):.4f})\")\n\nplt.plot([0, 1], [0, 1], linestyle=\"--\", color=\"gray\")\nplt.title(\"ROC Curves on the Test Set\")\nplt.xlabel(\"False Positive Rate\")\nplt.ylabel(\"True Positive Rate\")\nplt.legend()\nplt.show()\n\n# Plot one confusion matrix per model.\nfig, axes = plt.subplots(1, 2, figsize=(12, 5))\n\nfor ax, (model_name, preds) in zip(axes, test_predictions.items()):\n    cm = confusion_matrix(preds[\"y_true\"], preds[\"y_pred\"])\n    sns.heatmap(cm, annot=True, fmt=\"d\", cmap=\"Blues\", cbar=False, ax=ax, xticklabels=CLASS_NAMES, yticklabels=CLASS_NAMES)\n    ax.set_title(f\"Confusion Matrix: {model_name}\")\n    ax.set_xlabel(\"Predicted\")\n    ax.set_ylabel(\"Actual\")\n\nplt.tight_layout()\nplt.show()\n\n# Print the full classification report for each model.\nfor model_name, preds in test_predictions.items():\n    print(f\"Classification report for {model_name}:\")\n    print(classification_report(preds[\"y_true\"], preds[\"y_pred\"], target_names=CLASS_NAMES, zero_division=0))\n    print(\"-\" * 80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T20:38:20.254839Z","iopub.execute_input":"2026-09-02T20:38:20.25517Z","iopub.status.idle":"2026-09-02T20:38:33.196081Z","shell.execute_reply.started":"2026-09-02T20:38:20.255138Z","shell.execute_reply":"2026-09-02T20:38:33.195328Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Export the trained model weights\nimport torch\nfor model_name, artifact in trained_artifacts.items():\n    torch.save(artifact['model'].state_dict(), f\"{model_name}.pth\")\n    print(f\"Saved {model_name}.pth\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-09-02T20:38:33.197219Z","iopub.execute_input":"2026-09-02T20:38:33.197726Z","iopub.status.idle":"2026-09-02T20:38:33.434786Z","shell.execute_reply.started":"2026-09-02T20:38:33.197694Z","shell.execute_reply":"2026-09-02T20:38:33.433911Z"}},"outputs":[],"execution_count":null}]}