{"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":11848,"databundleVersionId":862157,"sourceType":"competition"}],"dockerImageVersionId":31089,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Histopathologic Cancer Detection — Deep Learning Project\n\n**Course:** DTSA 5511 Introduction to Deep Learning  \n**Date:** 9/20/2025\n\n---\n\n## 1. Overview\n\nThis project is based on the **Kaggle Histopathologic Cancer Detection Competition**. The goal is to build a deep learning model that can detect **metastatic cancer in small patches of histopathology images**.\n\n- **Problem type:** Supervised learning → Binary Image Classification (label = 0 for non-cancer, 1 for cancer).  \n- **Why it matters:** Detecting metastatic tissue quickly and accurately can support pathologists and improve cancer diagnosis workflows.  \n- **Data format:** Images of **96×96 pixels, RGB color**.  \n- **Dataset size:** ~220,000 labeled training patches + unlabeled test set.  \n- **Structure:**  \n  - `train/` — contains training images (`.tif` files).  \n  - `train_labels.csv` — mapping between image IDs and binary labels.  \n  - `test/` — unlabeled images for competition submission (not needed for this deliverable).  \n\n**Task objective:** Develop, train, and evaluate CNN-based architectures to classify patches as cancerous or not, with performance measured primarily by **ROC-AUC**.  ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"# Step 1. Setup\nimport numpy as np\nimport pandas as pd\nimport os\nimport matplotlib.pyplot as plt\nfrom PIL import Image\n\n# Torch / DL\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms, models\n\n# Metrics\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score\n\n# Check environment\nprint(\"Torch:\", torch.__version__)\nprint(\"CUDA available:\", torch.cuda.is_available())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T16:14:36.561806Z","iopub.execute_input":"2025-09-22T16:14:36.562105Z","iopub.status.idle":"2025-09-22T16:14:46.825146Z","shell.execute_reply.started":"2025-09-22T16:14:36.562081Z","shell.execute_reply":"2025-09-22T16:14:46.824369Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define paths (Kaggle auto-mounts competition data under /kaggle/input)\nDATA_DIR = \"/kaggle/input/histopathologic-cancer-detection\"\n\nTRAIN_DIR = os.path.join(DATA_DIR, \"train\")\nTEST_DIR = os.path.join(DATA_DIR, \"test\")\nLABELS = os.path.join(DATA_DIR, \"train_labels.csv\")\n\n# Load labels\ndf = pd.read_csv(LABELS)\nprint(df.head())\nprint(\"Train images:\", len(os.listdir(TRAIN_DIR)))\nprint(\"Test images:\", len(os.listdir(TEST_DIR)))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T16:14:52.776155Z","iopub.execute_input":"2025-09-22T16:14:52.776618Z","iopub.status.idle":"2025-09-22T16:14:58.328429Z","shell.execute_reply.started":"2025-09-22T16:14:52.776591Z","shell.execute_reply":"2025-09-22T16:14:58.327213Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Sanity Check: Image and Label Match\n\nBefore starting preprocessing and modeling, it is important to confirm that:\n1. The images load correctly from disk.  \n2. The IDs in `train_labels.csv` align with the corresponding image files.  \n3. Labels (0 = non-cancer, 1 = cancer) appear as expected.  \n\nThe following cell randomly selects an image and displays it with its label.\n","metadata":{}},{"cell_type":"code","source":"# Randomly select and display one image with its label to verify dataset integrity\n\nsample_id = df.sample(1).iloc[0][\"id\"]\nlabel = df.sample(1).iloc[0][\"label\"]\n\nimg_path = os.path.join(TRAIN_DIR, f\"{sample_id}.tif\")\nimg = Image.open(img_path)\nplt.imshow(img)\nplt.title(f\"Label: {label}\")\nplt.axis(\"off\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T16:15:14.284417Z","iopub.execute_input":"2025-09-22T16:15:14.284749Z","iopub.status.idle":"2025-09-22T16:15:14.675010Z","shell.execute_reply.started":"2025-09-22T16:15:14.284724Z","shell.execute_reply":"2025-09-22T16:15:14.673967Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 2. Exploratory Data Analysis (EDA)\n\nThe purpose of EDA is to:\n- Understand class balance (how many positive vs negative samples).\n- Visually inspect sample images from both classes.\n- Confirm that all images have the expected dimensions and color channels.\n\nThis helps identify potential challenges such as **class imbalance**, low image resolution, or data quality issues.","metadata":{}},{"cell_type":"code","source":"# Plot class balance\nax = df['label'].value_counts().sort_index().plot(\n    kind='bar',\n    color=['steelblue', 'indianred']\n)\nax.set_xticklabels(['Non-cancer (0)', 'Cancer (1)'], rotation=0)\nax.set_title(\"Class Balance in Training Data\")\nax.set_ylabel(\"Count\")\nplt.show()\n\nprint(\"Class distribution:\")\nprint(df['label'].value_counts(normalize=True))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T16:15:26.463821Z","iopub.execute_input":"2025-09-22T16:15:26.464168Z","iopub.status.idle":"2025-09-22T16:15:26.686423Z","shell.execute_reply.started":"2025-09-22T16:15:26.464108Z","shell.execute_reply":"2025-09-22T16:15:26.685617Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Show random samples from each class\ndef show_samples(label, n=6):\n    sample_df = df[df['label'] == label].sample(n)\n    plt.figure(figsize=(12, 6))\n    for i, img_id in enumerate(sample_df['id']):\n        img_path = os.path.join(TRAIN_DIR, f\"{img_id}.tif\")\n        img = Image.open(img_path)\n        plt.subplot(2, 3, i+1)\n        plt.imshow(img)\n        plt.axis(\"off\")\n        plt.title(f\"Label: {label}\")\n    plt.suptitle(f\"Random Samples (label={label})\")\n    plt.show()\n\nshow_samples(0, n=6)  # non-cancer\nshow_samples(1, n=6)  # cancer\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T16:15:31.286691Z","iopub.execute_input":"2025-09-22T16:15:31.287015Z","iopub.status.idle":"2025-09-22T16:15:32.622367Z","shell.execute_reply.started":"2025-09-22T16:15:31.286992Z","shell.execute_reply":"2025-09-22T16:15:32.621305Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### EDA Summary\n\n- **Class balance:** The dataset is imbalanced, with more negative (non-cancer) samples than positive ones.  \n- **Image quality:** All patches are 96×96 pixels in RGB format.  \n- **Visual inspection:** Cancerous patches often show dense, irregular cellular structures, while non-cancer patches appear more uniform.  \n\n**Implication for modeling:**  \nBecause of the imbalance, metrics like **ROC-AUC** and **F1-score** will be more informative than accuracy alone.  \nData augmentation and/or class weighting may be needed to improve generalization.","metadata":{}},{"cell_type":"markdown","source":"## 3. Data Cleaning\n\nThe dataset is already well-structured, but the following checks are important before modeling:\n\n1. **File existence:** Ensure every ID in `train_labels.csv` has a corresponding image file in `train/`.  \n2. **Corrupt files:** Verify images can be opened without error.  \n3. **Label integrity:** Confirm labels are binary (0 or 1) only.  \n\nIf missing or corrupt samples are found, they should be dropped to avoid training issues.\n","metadata":{}},{"cell_type":"code","source":"# Verify that every ID in the labels CSV has a corresponding image file\ndef image_exists(img_id):\n    path = os.path.join(TRAIN_DIR, f\"{img_id}.tif\")\n    return os.path.exists(path)\n\ndf['exists'] = df['id'].apply(image_exists)\n\nmissing = (~df['exists']).sum()\nprint(f\"Missing images: {missing}\")\n\n# Drop any rows without a corresponding image\ndf = df[df['exists']].drop(columns=['exists']).reset_index(drop=True)\nprint(\"Cleaned dataset size:\", len(df))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T16:16:15.561643Z","iopub.execute_input":"2025-09-22T16:16:15.562410Z","iopub.status.idle":"2025-09-22T16:22:35.206283Z","shell.execute_reply.started":"2025-09-22T16:16:15.562379Z","shell.execute_reply":"2025-09-22T16:22:35.205289Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Check unique label values\nprint(\"Unique labels:\", df['label'].unique())\n\n# Quick counts\nprint(df['label'].value_counts())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T16:26:22.055707Z","iopub.execute_input":"2025-09-22T16:26:22.056061Z","iopub.status.idle":"2025-09-22T16:26:22.068002Z","shell.execute_reply.started":"2025-09-22T16:26:22.056033Z","shell.execute_reply":"2025-09-22T16:26:22.066897Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Try opening a small sample of images to confirm no corruption\nfrom tqdm.notebook import tqdm\n\ncorrupt_count = 0\nfor img_id in tqdm(df['id'].sample(500)):  # check random 500 images\n    path = os.path.join(TRAIN_DIR, f\"{img_id}.tif\")\n    try:\n        _ = Image.open(path)\n    except:\n        corrupt_count += 1\n\nprint(\"Corrupt images found:\", corrupt_count)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-09-22T16:26:43.551335Z","iopub.execute_input":"2025-09-22T16:26:43.551652Z","iopub.status.idle":"2025-09-22T16:26:46.928012Z","shell.execute_reply.started":"2025-09-22T16:26:43.551632Z","shell.execute_reply":"2025-09-22T16:26:46.927077Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Data Cleaning Summary\n\n- **File existence:** Verified that all IDs in `train_labels.csv` have corresponding image files. No missing files were found.  \n- **Corrupt files:** A random sample of images was opened successfully; no corrupt images detected.  \n- **Labels:** Confirmed to be binary {0,1} only.  \n\n**Conclusion:** The dataset is clean, consistent, and ready for preprocessing, splitting, and model training. No additional cleaning steps were required.","metadata":{}},{"cell_type":"markdown","source":"## 4. Train/Validation Split & Preprocessing\n\n**Goal:** Create a robust validation strategy and define image transforms.\n\n- **Split:** Stratified 80/20 on `label` to preserve class balance.\n- **Transforms (train):** Light augmentations (H/V flips, small rotations) + normalization for transfer learning.\n- **Transforms (val):** Only normalization.\n- **Normalization:** ImageNet mean/std is used to match pretrained backbones.\n- **Class imbalance:** We compute class counts and (optionally) enable a `WeightedRandomSampler`.","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom torchvision import transforms\n\n# 80/20 stratified split\ntrain_df, val_df = train_test_split(\n    df, test_size=0.2, random_state=42, stratify=df[\"label\"]\n)\nprint(\"Train size:\", len(train_df), \" | Val size:\", len(val_df))\nprint(\"Train class counts:\\n\", train_df[\"label\"].value_counts())\nprint(\"Val class counts:\\n\",   val_df[\"label\"].value_counts())\n\n# Image size (competition tiles are 96x96)\nIMG_SIZE = 96\n\n# ImageNet normalization (for pretrained CNNs)\nIMAGENET_MEAN = (0.485, 0.456, 0.406)\nIMAGENET_STD  = (0.229, 0.224, 0.225)\n\n# Data augmentations for train; only normalization for val\ntrain_tfms = transforms.Compose([\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.RandomRotation(15),\n    transforms.ToTensor(),\n    transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n])\n\nval_tfms = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Normalize(IMAGENET_MEAN, IMAGENET_STD),\n])","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-22T16:13:33.901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nfrom PIL import Image\nfrom pathlib import Path\nimport torch\n\nclass HistoDataset(Dataset):\n    def __init__(self, df, img_dir, tfms=None, return_ids=False):\n        self.df = df.reset_index(drop=True)\n        self.img_dir = Path(img_dir)\n        self.tfms = tfms\n        self.return_ids = return_ids\n\n    def __len__(self):\n        return len(self.df)\n\n    def _img_path(self, _id):\n        # Competition files are .tif; if you ever switch sources, you can add fallbacks here\n        return self.img_dir / f\"{_id}.tif\"\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        img_path = self._img_path(row[\"id\"])\n        img = Image.open(img_path).convert(\"RGB\")\n        if self.tfms:\n            img = self.tfms(img)\n        label = torch.tensor(row[\"label\"], dtype=torch.long)\n        if self.return_ids:\n            return img, label, row[\"id\"]\n        return img, label\n\nBATCH_SIZE = 128\nNUM_WORKERS = 2  # Kaggle usually does fine with 2–4\n\ntrain_ds = HistoDataset(train_df, TRAIN_DIR, tfms=train_tfms)\nval_ds   = HistoDataset(val_df,   TRAIN_DIR, tfms=val_tfms)\n\ntrain_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, shuffle=True,\n                          num_workers=NUM_WORKERS, pin_memory=True)\nval_loader   = DataLoader(val_ds,   batch_size=BATCH_SIZE, shuffle=False,\n                          num_workers=NUM_WORKERS, pin_memory=True)\n\nlen(train_ds), len(val_ds)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-22T16:13:33.901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# WeightedRandomSampler so each batch is more balanced\nfrom torch.utils.data import WeightedRandomSampler\nimport numpy as np\n\nuse_sampler = False  # set to True to enable\n\nif use_sampler:\n    class_counts = train_df[\"label\"].value_counts().to_dict()\n    # weight for each class = 1 / freq\n    weights = train_df[\"label\"].map(lambda y: 1.0 / class_counts[y]).values\n    sampler = WeightedRandomSampler(weights=weights, num_samples=len(weights), replacement=True)\n\n    train_loader = DataLoader(train_ds, batch_size=BATCH_SIZE, sampler=sampler,\n                              num_workers=NUM_WORKERS, pin_memory=True)\n    print(\"Using WeightedRandomSampler\")\nelse:\n    print(\"Using random shuffle (no sampler)\")","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-22T16:13:33.901Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"**Why these choices?**\n- **Stratified split** preserves label proportions in train/val to avoid validation bias.\n- **Augmentations** (flips/rotations) are biologically plausible for histology patches and help reduce overfitting.\n- **ImageNet normalization** aligns input statistics with pretrained CNN backbones.\n- **Weighted sampler (optional)** can mitigate the effect of class imbalance on mini-batch composition.","metadata":{}},{"cell_type":"markdown","source":"## 5. Baseline Model\n\nWe begin with a **transfer learning baseline**:\n\n- **Backbone:** ResNet-18 pretrained on ImageNet.  \n- **Modification:** Replace final fully connected layer with a single output neuron for binary classification.  \n- **Loss:** Binary Cross-Entropy with Logits (`BCEWithLogitsLoss`).  \n- **Optimizer:** Adam with learning rate 1e-3.  \n- **Training:** Early stopping based on validation ROC-AUC.  \n\nThis establishes a strong reference point to evaluate future improvements such as stronger backbones, more augmentation, or class balancing techniques.","metadata":{}},{"cell_type":"code","source":"# Device setup\nimport torch\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", DEVICE)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-22T16:13:33.901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Device setup\nimport torch\nDEVICE = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Using device:\", DEVICE)\n\n# Model setup\nimport torch.nn as nn\nfrom torchvision import models\n\ndef build_model():\n    # No internet in Kaggle => don't request pretrained weights\n    model = models.resnet18(weights=None)  # randomly initialized\n    in_features = model.fc.in_features\n    model.fc = nn.Linear(in_features, 1)   # binary logit\n    return model\n\nmodel = build_model().to(DEVICE)\n\ncriterion = nn.BCEWithLogitsLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-22T16:13:33.901Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## 6. Training & Evaluation Plan\n\n- **Objective:** Optimize ROC-AUC on the validation set.\n- **Training:** Adam (lr=1e-3), mixed precision (if CUDA), ReduceLROnPlateau scheduler.\n- **Early stopping:** Monitor **val ROC-AUC**; stop if no improvement for 3 epochs.\n- **Reported metrics:** ROC-AUC (primary), Accuracy, F1, PR-AUC.\n- **Visualizations:** Training loss curve, ROC curve, PR curve, confusion matrix, metric table.\n","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nimport torch, json, pandas as pd\n\n# Where to save during this run\nCKPT_DIR = Path(\"/kaggle/working/checkpoints\")\nCKPT_DIR.mkdir(parents=True, exist_ok=True)\nLAST_CKPT = CKPT_DIR / \"last.pt\"\nBEST_CKPT = CKPT_DIR / \"best.pt\"\nHIST_CSV  = CKPT_DIR / \"history.csv\"\n\ndef save_checkpoint(epoch, model, optimizer, scheduler, best_auc, history_df):\n    torch.save({\n        \"epoch\": epoch,\n        \"model_state\": model.state_dict(),\n        \"optimizer_state\": optimizer.state_dict(),\n        \"scheduler_state\": scheduler.state_dict() if scheduler is not None else None,\n        \"best_auc\": best_auc,\n    }, LAST_CKPT)\n    # persist history table too\n    history_df.to_csv(HIST_CSV, index=False)\n\ndef save_best(model, best_auc):\n    torch.save({\"model_state\": model.state_dict(), \"best_auc\": best_auc}, BEST_CKPT)\n\ndef try_resume(model, optimizer=None, scheduler=None):\n    # 1) Try local /kaggle/working checkpoint\n    src = None\n    if LAST_CKPT.exists():\n        src = LAST_CKPT\n    else:\n        # 2) Or try prior notebook output you attached via \"Add Data\"\n        #    (edit this path to match the mounted output dataset name)\n        prev = Path(\"/kaggle/input/your-previous-notebook-output/checkpoints/last.pt\")\n        if prev.exists():\n            src = prev\n\n    start_epoch, best_auc = 1, -float(\"inf\")\n    if src is not None:\n        state = torch.load(src, map_location=\"cpu\")\n        model.load_state_dict(state[\"model_state\"])\n        if optimizer is not None and \"optimizer_state\" in state and state[\"optimizer_state\"] is not None:\n            optimizer.load_state_dict(state[\"optimizer_state\"])\n        if scheduler is not None and \"scheduler_state\" in state and state[\"scheduler_state\"] is not None:\n            scheduler.load_state_dict(state[\"scheduler_state\"])\n        best_auc = float(state.get(\"best_auc\", best_auc))\n        start_epoch = int(state.get(\"epoch\", 0)) + 1\n        print(f\"Resumed from {src} at epoch {start_epoch}, best_auc={best_auc:.4f}\")\n    else:\n        print(\"No checkpoint to resume from; starting fresh.\")\n    # load prior history if present\n    hist_path = HIST_CSV if HIST_CSV.exists() else Path(str(CKPT_DIR).replace(\"/working/\",\"/input/your-previous-notebook-output/\"))/\"checkpoints/history.csv\"\n    if hist_path.exists():\n        hist_df = pd.read_csv(hist_path)\n    else:\n        hist_df = pd.DataFrame(columns=[\"epoch\",\"train_loss\",\"val_roc_auc\",\"val_acc\",\"val_f1\",\"val_pr_auc\"])\n    return start_epoch, best_auc, hist_df\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-22T16:13:33.901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import (\n    roc_auc_score, accuracy_score, f1_score, confusion_matrix,\n    precision_recall_curve, roc_curve, auc\n)\nimport torch\n\n@torch.no_grad()\ndef evaluate(model, loader, device=DEVICE):\n    model.eval()\n    all_probs, all_labels = [], []\n    for x, y in loader:\n        x, y = x.to(device), y.to(device)\n        logits = model(x).squeeze(1)\n        probs = torch.sigmoid(logits).detach().cpu().numpy()\n        all_probs.append(probs)\n        all_labels.append(y.detach().cpu().numpy())\n    probs  = np.concatenate(all_probs)\n    labels = np.concatenate(all_labels)\n    preds  = (probs >= 0.5).astype(int)\n\n    roc  = roc_auc_score(labels, probs) if len(np.unique(labels))>1 else np.nan\n    acc  = accuracy_score(labels, preds)\n    f1   = f1_score(labels, preds, zero_division=0)\n    prec, rec, _ = precision_recall_curve(labels, probs)\n    pr_auc = auc(rec, prec)\n\n    return {\n        \"roc_auc\": roc, \"acc\": acc, \"f1\": f1, \"pr_auc\": pr_auc,\n        \"probs\": probs, \"labels\": labels, \"preds\": preds,\n        \"prec\": prec, \"rec\": rec\n    }\n\ndef plot_confusion(cm, labels=(\"Non-cancer\",\"Cancer\"), title=\"Confusion Matrix\"):\n    fig, ax = plt.subplots()\n    im = ax.imshow(cm, cmap=\"Blues\")\n    ax.set_xticks([0,1]); ax.set_yticks([0,1])\n    ax.set_xticklabels(labels); ax.set_yticklabels(labels)\n    ax.set_xlabel(\"Predicted\"); ax.set_ylabel(\"True\"); ax.set_title(title)\n    for (i, j), z in np.ndenumerate(cm):\n        ax.text(j, i, str(z), ha='center', va='center')\n    plt.colorbar(im); plt.show()\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-22T16:13:33.901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm.auto import tqdm\n\nscaler = torch.cuda.amp.GradScaler(enabled=torch.cuda.is_available())\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode=\"max\", factor=0.5, patience=2, verbose=True\n)\n\ndef train_one_epoch(model, loader, optimizer, device=DEVICE):\n    model.train()\n    total = 0.0\n    for x, y in tqdm(loader, leave=False):\n        x, y = x.to(device), y.float().to(device)\n        optimizer.zero_grad(set_to_none=True)\n        with torch.cuda.amp.autocast(enabled=torch.cuda.is_available()):\n            logits = model(x).squeeze(1)\n            loss = criterion(logits, y)\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        total += loss.item() * x.size(0)\n    return total / len(loader.dataset)\n\nEPOCHS = 12\npatience = 3\nbest_auc = -np.inf\nwait = 0\nhistory = []\n\nfor epoch in range(1, EPOCHS+1):\n    train_loss = train_one_epoch(model, train_loader, optimizer)\n    val = evaluate(model, val_loader)\n\n    scheduler.step(val[\"roc_auc\"])  # reduce LR on plateau\n\n    history.append({\"epoch\": epoch, \"train_loss\": train_loss,\n                    \"val_roc_auc\": val[\"roc_auc\"], \"val_acc\": val[\"acc\"],\n                    \"val_f1\": val[\"f1\"], \"val_pr_auc\": val[\"pr_auc\"]})\n    print(f\"Epoch {epoch:02d} | loss {train_loss:.4f} | \"\n          f\"AUC {val['roc_auc']:.4f} | Acc {val['acc']:.4f} | \"\n          f\"F1 {val['f1']:.4f} | PR-AUC {val['pr_auc']:.4f}\")\n\n    # early stopping on ROC-AUC\n    if val[\"roc_auc\"] > best_auc:\n        best_auc = val[\"roc_auc\"]\n        best_state = {k: v.cpu() for k, v in model.state_dict().items()}\n        wait = 0\n    else:\n        wait += 1\n        if wait >= patience:\n            print(\"Early stopping triggered.\")\n            break\n\n# restore best weights\nmodel.load_state_dict({k: v.to(DEVICE) for k, v in best_state.items()})\nval_best = evaluate(model, val_loader)\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-22T16:13:33.901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# History table\nhist_df = pd.DataFrame(history)\ndisplay(hist_df)\n\n# Training loss\nplt.figure()\nplt.plot(hist_df[\"epoch\"], hist_df[\"train_loss\"])\nplt.xlabel(\"Epoch\"); plt.ylabel(\"Train Loss\"); plt.title(\"Training Loss\"); plt.show()\n\n# Validation metrics\nplt.figure()\nplt.plot(hist_df[\"epoch\"], hist_df[\"val_roc_auc\"], label=\"ROC-AUC\")\nplt.plot(hist_df[\"epoch\"], hist_df[\"val_acc\"], label=\"Accuracy\")\nplt.plot(hist_df[\"epoch\"], hist_df[\"val_f1\"], label=\"F1\")\nplt.plot(hist_df[\"epoch\"], hist_df[\"val_pr_auc\"], label=\"PR-AUC\")\nplt.xlabel(\"Epoch\"); plt.title(\"Validation Metrics\"); plt.legend(); plt.show()\n\n# ROC curve\nfpr, tpr, _ = roc_curve(val_best[\"labels\"], val_best[\"probs\"])\nrocA = auc(fpr, tpr)\nplt.figure()\nplt.plot(fpr, tpr, label=f\"AUC={rocA:.3f}\")\nplt.plot([0,1],[0,1],'--')\nplt.xlabel(\"False Positive Rate\"); plt.ylabel(\"True Positive Rate\")\nplt.title(\"ROC Curve\"); plt.legend(); plt.show()\n\n# PR curve\nplt.figure()\nplt.plot(val_best[\"rec\"], val_best[\"prec\"], label=f\"PR-AUC={val_best['pr_auc']:.3f}\")\nplt.xlabel(\"Recall\"); plt.ylabel(\"Precision\"); plt.title(\"Precision-Recall Curve\"); plt.legend(); plt.show()\n\n# Confusion matrix at 0.5 threshold\ncm = confusion_matrix(val_best[\"labels\"], (val_best[\"probs\"]>=0.5).astype(int))\nplot_confusion(cm)\nprint({\n    \"Val ROC-AUC\": float(val_best[\"roc_auc\"]),\n    \"Val PR-AUC\": float(val_best[\"pr_auc\"]),\n    \"Val Accuracy\": float(val_best[\"acc\"]),\n    \"Val F1\": float(val_best[\"f1\"])\n})\n","metadata":{"trusted":true,"execution":{"execution_failed":"2025-09-22T16:13:33.901Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}