{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":11848,"databundleVersionId":862157,"sourceType":"competition"},{"sourceId":212102090,"sourceType":"kernelVersion"}],"dockerImageVersionId":30804,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"##### Histopathologic Cancer Detection - Model Training and Metrics\n##### Author: Aaron Storey\n##### Date: Decemeber 09, 2024\n##### Version: 1.0\n\nThis notebook implements the training pipeline for the histopathologic cancer detection model,\nincluding model architecture, training loops, metrics tracking, and visualization.\nBest performing version (V13: 0.8682) using EfficientNet-B0, modified in order to fine tune model using pseudo-labeling approach.","metadata":{"execution":{"iopub.status.busy":"2024-12-09T16:14:46.488612Z","iopub.execute_input":"2024-12-09T16:14:46.489009Z","iopub.status.idle":"2024-12-09T16:14:46.519045Z","shell.execute_reply.started":"2024-12-09T16:14:46.488971Z","shell.execute_reply":"2024-12-09T16:14:46.517802Z"}}},{"cell_type":"markdown","source":"##### 1. Import Required Libraries","metadata":{}},{"cell_type":"code","source":"!pip install efficientnet_pytorch wandb albumentations --quiet","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:20.178092Z","iopub.execute_input":"2024-12-10T14:27:20.178441Z","iopub.status.idle":"2024-12-10T14:27:28.940905Z","shell.execute_reply.started":"2024-12-10T14:27:20.178412Z","shell.execute_reply":"2024-12-10T14:27:28.939726Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader, ConcatDataset\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score\nfrom albumentations import (Compose, Normalize, Resize, HorizontalFlip, VerticalFlip, \n                            ShiftScaleRotate, RandomBrightnessContrast)\nfrom albumentations.pytorch import ToTensorV2\nfrom efficientnet_pytorch import EfficientNet\nfrom PIL import Image\nfrom tqdm import tqdm\nimport random","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:28.943605Z","iopub.execute_input":"2024-12-10T14:27:28.944044Z","iopub.status.idle":"2024-12-10T14:27:28.950469Z","shell.execute_reply.started":"2024-12-10T14:27:28.943998Z","shell.execute_reply":"2024-12-10T14:27:28.949597Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"SEED = 42\nrandom.seed(SEED)\nnp.random.seed(SEED)\ntorch.manual_seed(SEED)\nif torch.cuda.is_available():\n    torch.cuda.manual_seed(SEED)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:28.951739Z","iopub.execute_input":"2024-12-10T14:27:28.952117Z","iopub.status.idle":"2024-12-10T14:27:28.965884Z","shell.execute_reply.started":"2024-12-10T14:27:28.952075Z","shell.execute_reply":"2024-12-10T14:27:28.965031Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### 2. Configuration","metadata":{}},{"cell_type":"code","source":"DATA_DIR = \"/kaggle/input/histopathologic-cancer-detection\"\nTRAIN_DIR = f\"{DATA_DIR}/train\"\nLABELS_FILE = f\"{DATA_DIR}/train_labels.csv\"\n\nTARGET_SIZE = (96, 96)\nBATCH_SIZE = 64\nEPOCHS = 7\nLR = 1e-3\nVAL_RATIO = 0.2\nMODEL_SAVE_PATH = \"model_best.pth\"\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(f\"Using device: {device}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:28.968221Z","iopub.execute_input":"2024-12-10T14:27:28.968550Z","iopub.status.idle":"2024-12-10T14:27:28.975815Z","shell.execute_reply.started":"2024-12-10T14:27:28.968513Z","shell.execute_reply":"2024-12-10T14:27:28.974943Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### 3. Load and Preprocess Data","metadata":{}},{"cell_type":"code","source":"labels_df = pd.read_csv(LABELS_FILE)\nimg_ids = labels_df['id'].values\nlabels = labels_df['label'].values\n\ntrain_ids, val_ids, train_labels, val_labels = train_test_split(\n    img_ids, labels, test_size=VAL_RATIO, random_state=SEED, stratify=labels\n)\n\nprint(f\"Train size: {len(train_ids)}, Validation size: {len(val_ids)}\")\nprint(labels_df['label'].value_counts())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:28.976807Z","iopub.execute_input":"2024-12-10T14:27:28.977148Z","iopub.status.idle":"2024-12-10T14:27:29.299266Z","shell.execute_reply.started":"2024-12-10T14:27:28.977116Z","shell.execute_reply":"2024-12-10T14:27:29.298277Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### 4. Dataset and Transforms","metadata":{}},{"cell_type":"code","source":"# 4.1 Dataset Class\n\nclass HistologyDataset(Dataset):\n    def __init__(self, img_ids, labels, img_dir, transform=None):\n        self.img_ids = img_ids\n        self.labels = labels\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.img_ids)\n\n    def __getitem__(self, idx):\n        img_id = self.img_ids[idx]\n        label = self.labels[idx]\n        img_path = os.path.join(self.img_dir, f\"{img_id}.tif\")\n        img = np.array(Image.open(img_path).convert(\"RGB\"))\n\n        if self.transform:\n            img = self.transform(image=img)['image']\n\n        return img, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:29.300295Z","iopub.execute_input":"2024-12-10T14:27:29.300577Z","iopub.status.idle":"2024-12-10T14:27:29.307562Z","shell.execute_reply.started":"2024-12-10T14:27:29.300549Z","shell.execute_reply":"2024-12-10T14:27:29.306487Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 4.2 Albumentations Transforms\ntrain_transforms = Compose([\n    Resize(*TARGET_SIZE),\n    HorizontalFlip(p=0.5),\n    VerticalFlip(p=0.5),\n    ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, p=0.5),\n    RandomBrightnessContrast(p=0.5),\n    Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2(),\n])\n\nval_transforms = Compose([\n    Resize(*TARGET_SIZE),\n    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":"2024-12-10T14:27:29.308902Z","iopub.execute_input":"2024-12-10T14:27:29.309964Z","iopub.status.idle":"2024-12-10T14:27:29.331057Z","shell.execute_reply.started":"2024-12-10T14:27:29.309918Z","shell.execute_reply":"2024-12-10T14:27:29.330147Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 4.3 Create Datasets and Dataloaders\n\ntrain_dataset = HistologyDataset(train_ids, train_labels, TRAIN_DIR, transform=train_transforms)\nval_dataset = HistologyDataset(val_ids, val_labels, TRAIN_DIR, transform=val_transforms)\n\ntrain_loader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=4)\nval_loader = DataLoader(val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:29.331933Z","iopub.execute_input":"2024-12-10T14:27:29.332189Z","iopub.status.idle":"2024-12-10T14:27:29.337537Z","shell.execute_reply.started":"2024-12-10T14:27:29.332165Z","shell.execute_reply":"2024-12-10T14:27:29.336610Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### 5. Visual Exploration of Data\nDisplaying examples of transformations outside of the dataset class to visualize results","metadata":{}},{"cell_type":"code","source":"# Show some random training images without augmentation\nsample_ids = np.random.choice(train_ids, 4, replace=False)\nplt.figure(figsize=(10, 10))\nfor i, img_id in enumerate(sample_ids, 1):\n    img_path = os.path.join(TRAIN_DIR, f\"{img_id}.tif\")\n    img = Image.open(img_path).convert(\"RGB\")\n    plt.subplot(2, 2, i)\n    plt.imshow(img)\n    plt.title(f\"ID: {img_id}\")\n    plt.axis('off')\nplt.suptitle(\"Sample Training Images (Original)\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:29.338658Z","iopub.execute_input":"2024-12-10T14:27:29.338999Z","iopub.status.idle":"2024-12-10T14:27:29.875477Z","shell.execute_reply.started":"2024-12-10T14:27:29.338962Z","shell.execute_reply":"2024-12-10T14:27:29.874478Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Show some random training images with augmentation\naugmented_images = []\nimg_path = os.path.join(TRAIN_DIR, f\"{sample_ids[0]}.tif\")\nimg = np.array(Image.open(img_path).convert(\"RGB\"))\n\nfor _ in range(4):\n    transformed = train_transforms(image=img)['image']\n    # Convert tensor back to numpy for display\n    img_np = transformed.permute(1, 2, 0).cpu().numpy()\n    # Denormalize for visualization\n    img_np = (img_np * np.array([0.229, 0.224, 0.225])) + np.array([0.485, 0.456, 0.406])\n    img_np = np.clip(img_np, 0, 1)\n    augmented_images.append(img_np)\n\nplt.figure(figsize=(10,10))\nfor i, im in enumerate(augmented_images, 1):\n    plt.subplot(2, 2, i)\n    plt.imshow(im)\n    plt.title(\"Augmented Image\")\n    plt.axis('off')\nplt.suptitle(\"Augmented Transformations\")\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:29.877917Z","iopub.execute_input":"2024-12-10T14:27:29.878228Z","iopub.status.idle":"2024-12-10T14:27:30.350637Z","shell.execute_reply.started":"2024-12-10T14:27:29.878201Z","shell.execute_reply":"2024-12-10T14:27:30.349746Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### 6. Model Definition","metadata":{}},{"cell_type":"code","source":"class CancerClassifier(nn.Module):\n    def __init__(self, num_classes=2):\n        super(CancerClassifier, self).__init__()\n        self.model = EfficientNet.from_pretrained(\"efficientnet-b0\")\n        self.model._fc = nn.Linear(self.model._fc.in_features, num_classes)\n\n    def forward(self, x):\n        return self.model(x)\n\nmodel = CancerClassifier().to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:30.351818Z","iopub.execute_input":"2024-12-10T14:27:30.352248Z","iopub.status.idle":"2024-12-10T14:27:31.633552Z","shell.execute_reply.started":"2024-12-10T14:27:30.352196Z","shell.execute_reply":"2024-12-10T14:27:31.632534Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### 7. Loss, Optimizer","metadata":{}},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(), lr=LR)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:31.634654Z","iopub.execute_input":"2024-12-10T14:27:31.634959Z","iopub.status.idle":"2024-12-10T14:27:32.587232Z","shell.execute_reply.started":"2024-12-10T14:27:31.634920Z","shell.execute_reply":"2024-12-10T14:27:32.586537Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### 8. Training and Validation Functions","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, loader, optimizer, criterion, device):\n    model.train()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n\n    for images, labels in loader:\n        images = images.to(device)\n        labels = 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() * images.size(0)\n        _, predicted = outputs.max(1)\n        correct += predicted.eq(labels).sum().item()\n        total += labels.size(0)\n\n    epoch_loss = running_loss / total\n    epoch_acc = correct / total\n    return epoch_loss, epoch_acc\n\ndef validate(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    correct = 0\n    total = 0\n    all_probs = []\n    all_targets = []\n\n    with torch.no_grad():\n        for images, labels in loader:\n            images = images.to(device)\n            labels = labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n\n            running_loss += loss.item() * images.size(0)\n            _, predicted = outputs.max(1)\n            correct += predicted.eq(labels).sum().item()\n            total += labels.size(0)\n\n            # Probabilities for positive class\n            probs = torch.softmax(outputs, dim=1)[:, 1].cpu().numpy()\n            all_probs.extend(probs)\n            all_targets.extend(labels.cpu().numpy())\n\n    epoch_loss = running_loss / total\n    epoch_acc = correct / total\n    roc_auc = roc_auc_score(all_targets, all_probs)\n    return epoch_loss, epoch_acc, roc_auc","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:32.588193Z","iopub.execute_input":"2024-12-10T14:27:32.588588Z","iopub.status.idle":"2024-12-10T14:27:32.597904Z","shell.execute_reply.started":"2024-12-10T14:27:32.588561Z","shell.execute_reply":"2024-12-10T14:27:32.596941Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### 9. Training Loop","metadata":{}},{"cell_type":"code","source":"train_losses = []\ntrain_accuracies = []\nval_losses = []\nval_accuracies = []\nval_aucs = []\n\nbest_val_auc = 0.0\n\nfor epoch in range(EPOCHS):\n    train_loss, train_acc = train_one_epoch(model, train_loader, optimizer, criterion, device)\n    val_loss, val_acc, val_auc = validate(model, val_loader, criterion, device)\n\n    train_losses.append(train_loss)\n    train_accuracies.append(train_acc)\n    val_losses.append(val_loss)\n    val_accuracies.append(val_acc)\n    val_aucs.append(val_auc)\n\n    print(f\"Epoch [{epoch+1}/{EPOCHS}]\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f}\")\n    print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f} | Val AUC: {val_auc:.4f}\")\n\n    # Save the best model based on validation AUC\n    if val_auc > best_val_auc:\n        best_val_auc = val_auc\n        torch.save(model.state_dict(), MODEL_SAVE_PATH)\n        print(\"Saved Best Model\\n\")\n\nprint(\"Training complete.\")\nprint(f\"Best Validation AUC: {best_val_auc:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-10T14:27:32.598904Z","iopub.execute_input":"2024-12-10T14:27:32.599178Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### 10. Visualization of Training Curves","metadata":{}},{"cell_type":"code","source":"# Plot Loss\nplt.figure(figsize=(10,5))\nplt.plot(train_losses, label='Train Loss')\nplt.plot(val_losses, label='Val Loss')\nplt.title('Loss Curve')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.show()\n\n# Plot Accuracy\nplt.figure(figsize=(10,5))\nplt.plot(train_accuracies, label='Train Acc')\nplt.plot(val_accuracies, label='Val Acc')\nplt.title('Accuracy Curve')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy')\nplt.legend()\nplt.show()\n\n# Plot AUC\nplt.figure(figsize=(10,5))\nplt.plot(val_aucs, label='Val AUC')\nplt.title('Validation AUC Curve')\nplt.xlabel('Epoch')\nplt.ylabel('AUC')\nplt.legend()\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### 11. Load Best Model and Inspect Predictions","metadata":{}},{"cell_type":"code","source":"model.load_state_dict(torch.load(MODEL_SAVE_PATH))\nmodel.eval()\nprint(\"Best model loaded.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"val_sample_ids = np.random.choice(val_ids, 4, replace=False)\nval_sample_imgs = []\nval_sample_labels = []\nval_sample_preds = []\nval_sample_probs = []\n\nwith torch.no_grad():\n    for img_id in val_sample_ids:\n        img_path = os.path.join(TRAIN_DIR, f\"{img_id}.tif\")\n        img = np.array(Image.open(img_path).convert(\"RGB\"))\n        \n        # Apply validation transforms\n        transformed = val_transforms(image=img)['image']\n        transformed = transformed.unsqueeze(0).to(device)\n        \n        output = model(transformed)\n        prob = torch.softmax(output, dim=1)[:,1].item()\n        pred = (prob > 0.5)*1  # Binary prediction\n        label = labels_df[labels_df['id'] == img_id]['label'].values[0]\n\n        val_sample_imgs.append(img)\n        val_sample_labels.append(label)\n        val_sample_preds.append(pred)\n        val_sample_probs.append(prob)\n\n# Visualize predictions\nplt.figure(figsize=(10,10))\nfor i, (im, lb, pr, pb) in enumerate(zip(val_sample_imgs, val_sample_labels, val_sample_preds, val_sample_probs), 1):\n    plt.subplot(2,2,i)\n    plt.imshow(im)\n    plt.title(f\"True: {lb}, Pred: {pr}, Prob: {pb:.2f}\")\n    plt.axis('off')\nplt.suptitle(\"Validation Samples - Predictions\")\nplt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"##### 12. Pseudo-Labeling and Fine Tuning\nPseudo-labeling approach:\n\n    Identify validation images with high-confidence predictions (e.g., probability < 0.1 for class 0 or > 0.9 for class 1).\n    Add these high-confidence samples with predicted labels to the training set.\n    Fine-tune the model on this augmented training set to potentially improve performance.","metadata":{}},{"cell_type":"code","source":"# Extract predictions for entire validation set\nval_probs = []\nval_targets = []\nval_images_full = []\nwith torch.no_grad():\n    for img_id, lbl in zip(val_ids, val_labels):\n        img_path = os.path.join(TRAIN_DIR, f\"{img_id}.tif\")\n        img = np.array(Image.open(img_path).convert(\"RGB\"))\n        transformed = val_transforms(image=img)['image'].unsqueeze(0).to(device)\n\n        output = model(transformed)\n        prob = torch.softmax(output, dim=1)[:,1].item()\n        val_probs.append(prob)\n        val_targets.append(lbl)\n        val_images_full.append((img_id, lbl))\n\nval_probs = np.array(val_probs)\nval_targets = np.array(val_targets)\n\n# Define thresholds for confident predictions\nlower_threshold = 0.1\nupper_threshold = 0.9\n\npseudo_ids = []\npseudo_labels = []\n\nfor (img_id, true_lbl), prob in zip(val_images_full, val_probs):\n    if prob < lower_threshold:\n        pseudo_ids.append(img_id)\n        pseudo_labels.append(0)  # confident prediction of class 0\n    elif prob > upper_threshold:\n        pseudo_ids.append(img_id)\n        pseudo_labels.append(1)  # confident prediction of class 1\n\nprint(f\"Selected {len(pseudo_ids)} pseudo-labeled samples.\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create a pseudo-labeled dataset and combine with original training set\n\npseudo_dataset = HistologyDataset(pseudo_ids, pseudo_labels, TRAIN_DIR, transform=train_transforms)\naugmented_dataset = ConcatDataset([train_dataset, pseudo_dataset])\n\naugmented_loader = DataLoader(augmented_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=4)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Fine-tune the model\n# We can re-initialize optimizer if needed\noptimizer = torch.optim.Adam(model.parameters(), lr=LR * 0.5)  # slightly lower LR for fine-tuning\nFINE_TUNE_EPOCHS = 2\n\nfor epoch in range(FINE_TUNE_EPOCHS):\n    train_loss, train_acc = train_one_epoch(model, augmented_loader, optimizer, criterion, device)\n    val_loss, val_acc, val_auc = validate(model, val_loader, criterion, device)\n\n    print(f\"Fine-tune Epoch [{epoch+1}/{FINE_TUNE_EPOCHS}]\")\n    print(f\"Train Loss: {train_loss:.4f} | Train Acc: {train_acc:.4f}\")\n    print(f\"Val   Loss: {val_loss:.4f} | Val   Acc: {val_acc:.4f} | Val AUC: {val_auc:.4f}\")\n\n    # Save improved model if any\n    if val_auc > best_val_auc:\n        best_val_auc = val_auc\n        torch.save(model.state_dict(), MODEL_SAVE_PATH)\n        print(\"Saved Improved Model after Fine-tuning\\n\")\n\nprint(\"Fine-tuning complete.\")\nprint(f\"Best Validation AUC after Fine-tuning: {best_val_auc:.4f}\")","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}