{"metadata":{"kernelspec":{"name":"python3","display_name":"Python 3","language":"python"},"language_info":{"name":"python","version":"3.11.11","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"}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 🧬 PCam Dataset: Tumor Detection via Binary Image Classification\n\nFor full dataset details, visit the official repository:  \n🔗 [github.com/basveeling/pcam](https://github.com/basveeling/pcam)\n\n\n## 📊 Dataset Overview\n\nhttps://github.com/basveeling/pcam\n\nThe **PatchCamelyon (PCam)** benchmark is a challenging image classification dataset designed for breast cancer detection tasks.\n\n- 📦 **Total images**: 327,680 color patches  \n- 🖼️ **Image size**: 96 × 96 pixels\n- 🧪 **Source**: Histopathologic scans of lymph node sections  \n- 🏷️ **Labels**: Binary — A positive (1) label indicates that the center 32x32px region of a patch contains at least one pixel of tumor tissue. Tumor tissue in the outer region of the patch does not influence the label.\n\n```\nB. S. Veeling, J. Linmans, J. Winkens, T. Cohen, M. Welling. \"Rotation Equivariant CNNs for Digital Pathology\". arXiv:1806.03962\n```\n\n```\nEhteshami Bejnordi et al. Diagnostic Assessment of Deep Learning Algorithms for Detection of Lymph Node Metastases in Women With Breast Cancer. JAMA: The Journal of the American Medical Association, 318(22), 2199–2210. doi:jama.2017.14585\n```\n\nUnder CC0 License\n\n## 🧠 Solution to Implement\n\nIn this notebook, we implement a solution inspired by the following research paper:\n\n> 📄 [**Cancer Image Classification Based on DenseNet Model**](https://arxiv.org/abs/2011.11186)  \n> _by Zhong, Ziliang; Zheng, Muhang; Mai, Huafeng; Zhao, Jianan; Liu, Xinyi_\n\nThis study explores the application of DenseNet architectures to the PCam dataset for accurate cancer classification.\n\n---\n\n## Results\n\nThe submission on kaggle with the model trained on this notebook is \n\n```Public score: 0.9733```\n\n### You can try it now ! With gradio \n\nOn Hugging Face Spaces:\n\n[![Hugging Face Spaces](https://img.shields.io/badge/🤗%20Hugging%20Face-Spaces-blue)](https://huggingface.co/spaces/eloise54/pcam_project)\n\n### You can see the corresponding gitlab repo here \n\n[![GitLab Repo](https://img.shields.io/badge/GitLab-Repository-orange?logo=gitlab)](https://gitlab.com/nn_projects/pcam_project)\n\n","metadata":{}},{"cell_type":"markdown","source":"# 1. Load the dataset\nLoad the training, test and validation datasets from PCAM.\n\nWe are going to use the kaggle version that is a cleaned version of the official PCAM dataset.\n\n```\nThe original PCam dataset contains duplicate images due to its probabilistic sampling, however, the version presented on Kaggle does not contain duplicates. We have otherwise maintained the same data and splits as the PCam benchmark.\n'''\nIn the kaggle version duplicates ar removed and there is no leakage between training and test datasets.","metadata":{}},{"cell_type":"code","source":"import typing as tp\nimport numpy as np\nimport torch\nimport torchvision\nfrom torch import nn\nfrom torch.utils.data import Dataset, DataLoader, ConcatDataset\nfrom torchvision.transforms import ToTensor \nfrom torchvision import datasets\nfrom torch.utils.tensorboard import SummaryWriter","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We need to use GPU if available","metadata":{}},{"cell_type":"code","source":"from torch.optim import Optimizer, lr_scheduler\nfrom torch.optim.lr_scheduler import LRScheduler\n\nif torch.cuda.is_available():\n    device = torch.device(\"cuda\")\nelse:\n    device = torch.device(\"cpu\")\nprint(\"Using device\", device)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Let's download the kaggle dataset.\nFor this you need your credentials.\nIf you did not set already your ```~/.kaggle/kaggle.json``` key:\n - Go to your kaggle account setting and create a new API token if needed.\n - Then feel in this part with your information ```creds = '{\"username\":\"xxxxx\",\"key\":\"xxxxx\"}'```","metadata":{}},{"cell_type":"code","source":"!pip install kaggle\ncreds = '{\"username\":\"xxxxx\",\"key\":\"xxxxx\"}'\nfrom pathlib import Path\n\ncred_path = Path('~/.kaggle/kaggle.json').expanduser()\nif not cred_path.exists():\n    cred_path.parent.mkdir(exist_ok=True)\n    cred_path.write_text(creds)\n    cred_path.chmod(0o600)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport zipfile\n\nroot = \"/kaggle/input/\"\ndataset_dir = \"/kaggle/input/histopathologic-cancer-detection\"\nzip_file = \"histopathologic-cancer-detection.zip\"\ntrain_path = os.path.join(dataset_dir, \"train\")\n\nif not os.path.exists(root):\n  os.mkdir(root)\n\nif not os.path.exists('results'):\n  os.mkdir('results')\n\nif not os.path.exists(train_path):\n    print(\"Downloading Histopathologic Cancer Detection dataset...\")\n    !kaggle competitions download -c histopathologic-cancer-detection -p {root} --force\nelse:\n    print(\"Dataset zip already downloaded.\")\n\nif not os.path.exists(train_path):\n    print(\"Unzipping dataset...\")\n    with zipfile.ZipFile(os.path.join(root, zip_file), 'r') as zip_ref:\n        zip_ref.extractall(dataset_dir)\nelse:\n    print(\"Dataset already unzipped.\")","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Know Let's create our pytorch dataset class.\nI have used train_test_split from sklearn to have a stratified dataset (The kaggle PCAM dataset is unbalanced)","metadata":{}},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom PIL import Image\nimport pandas as pd\n\nclass PcamDatasetKaggle(torchvision.datasets.VisionDataset):\n    def __init__(self, root, split, transform, target_transform = None):\n         super().__init__(root, transform=transform, target_transform=target_transform)\n         self.root = root\n         self.split = split\n         self.transform = transform\n         self.img_path = os.path.join(self.root, \"train\")\n\n         self.full_labels = pd.read_csv(self.root+'/train_labels.csv')\n         X_train, X_test, y_train, y_test = train_test_split(self.full_labels['id'],\n                                                             self.full_labels['label'],\n                                                             test_size = 0.2, \n                                                             train_size = 0.8,\n                                                             random_state=42,\n                                                             shuffle=True,\n                                                             stratify=self.full_labels['label'])\n        \n         if (split == \"train\"):\n             self.imgs = X_train + \".tif\"\n             self.labels = y_train\n         elif (split == \"val\"):\n             self.imgs = X_test + \".tif\"\n             self.labels = y_test\n         else:\n             self.img_path = os.path.join(self.root, self.split)\n             self.imgs = pd.Series(list(sorted(os.listdir(self.img_path))))\n             self.labels = pd.Series(torch.full((len(self.imgs),), -10))      \n         assert len(self.labels) == len(self.imgs)\n         print(\"Split\", split, \"Negative/Positive samples % \" , 100.0*(self.labels.value_counts() / self.labels.shape[0]))\n\n    def __getitem__(self, idx):\n        assert idx < len(self.imgs)\n        img = Image.open(os.path.join(self.img_path, self.imgs.iloc[idx]))\n        if self.transform:\n            img = self.transform(image = np.array(img))\n        label = self.labels.iloc[idx]\n        return img['image'].to(torch.float32), label\n    def __len__(self) :\n        return len(self.imgs)\n\ndef check_dataset_leakage(dataset1, dataset2):\n    duplicates = set(dataset1.imgs) & set(dataset2.imgs)\n    assert len(duplicates) == 0\n    \ndef check_same_imgs(dataset1, dataset2):\n    duplicates = set(dataset1.imgs) & set(dataset2.imgs)\n    assert len(duplicates) == len(dataset1.imgs)\n    assert len(duplicates) == len(dataset2.imgs)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Let's define some transforms for dataloading and data augmentation","metadata":{}},{"cell_type":"code","source":"import torchvision.transforms as transforms\n\ntorch.manual_seed(42)\ntorch.cuda.manual_seed_all(42)\n\n# Preprocess images with transforms\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)), #Match resnet original input size            \n    transforms.ToTensor()\n])\n\ntransform_data_augment = transforms.Compose([ \n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n    transforms.GaussianBlur(kernel_size = (5,5),sigma=(0.2, 0.7)),\n    transforms.RandomRotation(degrees=90),\n    transforms.ColorJitter(\n        brightness=0.4, \n        contrast=0.4, \n        saturation=0.1, \n        hue=0.03\n    ),\n    transforms.RandomResizedCrop(size = (224, 224), scale = (0.7, 1.0)),\n    transforms.ToTensor()\n])\n\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Let's defined a more refined trasformation with [albumentations](https://albumentations.ai/)\n\nI have used the ```transform_data_augment``` from [link](https://github.com/azkalot1/Histopathologic-Cancer-Detection/blob/master/utils.py) and added a normalization per channel layer to improve robustness and lessen overfitting","metadata":{}},{"cell_type":"code","source":"!pip install albumentations","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import albumentations as A\ntransform_data_augment = A.Compose([A.Resize(224, 224), \n                    A.HorizontalFlip(), \n                    A.VerticalFlip(), \n                    A.RandomRotate90(), \n                    A.Transpose(), \n                    A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.50, rotate_limit=60, p=.75),\n                    A.OpticalDistortion(),\n                    A.GridDistortion(), \n                    A.RandomBrightnessContrast(p=0.3), \n                    A.RandomGamma(p=0.3), \n                    A.OneOf([A.HueSaturationValue(hue_shift_limit=20, sat_shift_limit=0.1, val_shift_limit=0.1, p=0.3), \n                            A.ChannelShuffle(p=0.3), A.CLAHE(p=0.3)]),\n                    A.Normalize(normalization=\"image_per_channel\", p=1.0),\n                              A.ToTensorV2()])\n\ntransform = A.Compose([\n                A.Resize(224, 224),\n                A.Normalize(normalization=\"image_per_channel\", p=1.0),\n                A.ToTensorV2()\n            ])","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from copy import deepcopy\n\n\"\"\" PCAM pytorch version but the dataset is not clean \ntraining_set_original = datasets.PCAM(root=\"/kaggle/input\", split=\"train\",download = True, transform = transform) \ntraining_set_augment = datasets.PCAM(root=\"/kaggle/input\", split=\"train\",download = True, transform = transform_data_augment)\nval_set = datasets.PCAM(root=\"/kaggle/input\", split=\"val\", download=True, transform = transform)\ntest_set = datasets.PCAM(root=\"/kaggle/input\", split=\"test\", download=True, transform = transform)\n\"\"\"\n\ntraining_set_original = PcamDatasetKaggle(root=dataset_dir, split=\"train\", transform = deepcopy(transform)) \ntraining_set_augment = PcamDatasetKaggle(root=dataset_dir, split=\"train\", transform = deepcopy(transform_data_augment)) \n\nval_set = PcamDatasetKaggle(root=dataset_dir, split=\"val\", transform = deepcopy(transform)) \nval_set_augment = PcamDatasetKaggle(root=dataset_dir, split=\"val\", transform = deepcopy(transform_data_augment)) \n\ntest_set = PcamDatasetKaggle(root=dataset_dir, split=\"test\", transform = deepcopy(transform))\ntest_set_augment = PcamDatasetKaggle(root=dataset_dir, split=\"test\", transform = deepcopy(transform_data_augment)) #For TTA\n\ncheck_dataset_leakage(training_set_original, val_set)\ncheck_dataset_leakage(training_set_original, test_set)\ncheck_dataset_leakage(val_set, test_set)\ncheck_same_imgs(training_set_original, training_set_augment)\ncheck_same_imgs(val_set, val_set_augment)\ncheck_same_imgs(test_set, test_set_augment)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 2. Plot and visualize original and augmented the data\nEach (3,96,96) image is associated with a binary label indicates the presence of a tumor.\n\nLet's define a function to plot some images with their label.\n\nLet's save the plots in an experiment directory for logging purposes\n","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\ndef plot_training_set_sample(training_set, \n                             file_name = \"results/pcam/data.png\", \n                             rows = 5, \n                             cols = 5, \n                             mean_stdev = torch.Tensor([[0.0, 0.0, 0.0], [1.0, 1.0, 1.0]])):\n    mean = mean_stdev[0].numpy()\n    std  = mean_stdev[1].numpy()\n    fig = plt.figure(figsize=(2*cols, 2*rows))\n    for i in range(1, rows*cols + 1):\n        random_idx = torch.randint(len(training_set), (1,)).item()\n        fig.add_subplot(rows, cols, i)\n        img = training_set[random_idx][0].permute(1,2,0).numpy() \n        img_unnormalized = img*std + mean\n        img_unnormalized = np.clip(img_unnormalized, 0, 1)\n        plt.imshow(img_unnormalized)\n        plt.axis(\"off\")\n        plt.title(training_set[random_idx][1])\n    plt.savefig(file_name)\n    plt.show()\n    ","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nfrom datetime import datetime\nexp_dir = \"results/pcam/\"+datetime.now().strftime(\"%d_%m_%Y_%H_%M_%S\")\nos.makedirs(exp_dir)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Original Training Set\")\nplot_training_set_sample(training_set_original, exp_dir + \"/training_set_original.png\",rows=2, cols=5)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Augmented Training Set\")\nplot_training_set_sample(training_set_augment, exp_dir + \"/training_set_augment.png\",rows=2, cols=5)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3.Normalize and create augmented dataset","metadata":{}},{"cell_type":"markdown","source":"Let's create a function that computes mean, standard deviation and class balance for a pytorch DataLoader.\n\nNormalize the datasets accordingly","metadata":{}},{"cell_type":"code","source":"# Create DataLoader\nbatch_size = 128\ntraining_set_original = PcamDatasetKaggle(root=dataset_dir, split=\"train\", transform = deepcopy(transform)) \ntraining_set_augment = PcamDatasetKaggle(root=dataset_dir, split=\"train\", transform = deepcopy(transform_data_augment)) \nval_set = PcamDatasetKaggle(root=dataset_dir, split=\"val\", transform = deepcopy(transform)) \nval_set_augment = PcamDatasetKaggle(root=dataset_dir, split=\"val\", transform = deepcopy(transform_data_augment)) \ntest_set = PcamDatasetKaggle(root=dataset_dir, split=\"test\", transform = deepcopy(transform))\n\n\n# Create Augmented Training Dataset\ntraining_set = ConcatDataset([training_set_original, training_set_augment])\n\n# Create Final DataLoaders\ntraining_dataloader = DataLoader(training_set, batch_size=batch_size, shuffle=True, pin_memory=True, num_workers=6, persistent_workers = True)\nval_dataloader = DataLoader(val_set, batch_size=batch_size, shuffle=False, pin_memory=True, num_workers=6, persistent_workers = True)\nval_dataloader_augment = DataLoader(val_set_augment, batch_size=batch_size, shuffle=False, pin_memory=True, num_workers=6, persistent_workers = True)\ntest_dataloader = DataLoader(test_set, batch_size=batch_size, shuffle=False, pin_memory=True, num_workers=6, persistent_workers = True)\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(\"Full Training Set Normalized\")\nplot_training_set_sample(training_set, exp_dir + \"/training_set_final.png\", rows = 2, cols = 5)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 3. Defining a training loop over one epoch and a metric\nThe dataset is not balance thus it is better to use roc_auc_score than accuracy","metadata":{}},{"cell_type":"code","source":"def compute_metrics(full_y: torch.Tensor, \n                    full_logits: torch.Tensor,  \n                    full_pred: torch.Tensor,  \n                    sk_learn_metrics_logits: tp.List[tp.Callable],\n                    sk_learn_metrics_pred: tp.List[tp.Callable]) -> tp.Dict:\n    full_y = full_y.detach().cpu().numpy()\n    full_logits = torch.sigmoid(full_logits).detach().cpu().numpy()\n    full_pred = full_pred.detach().cpu().numpy()\n    \n    results = {}\n    for metric in sk_learn_metrics_logits:\n        results[metric.__name__] = metric(full_y, full_logits)\n    for metric in sk_learn_metrics_pred:\n        results[metric.__name__] = metric(full_y, full_pred)\n    return results","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def run_one_epoch(model : nn.Module, \n                   training_dataloader: DataLoader,\n                   optimizer: Optimizer,\n                   loss_function: nn.Module,\n                   scheduler : LRScheduler,\n                   device: torch.cuda.device,\n                   writer: SummaryWriter,\n                   epoch: int,\n                   sk_learn_metrics_logits: tp.List[tp.Callable],\n                   sk_learn_metrics_pred: tp.List[tp.Callable],\n                   threshold: float = 0.5):\n    running_loss = 0.0\n    num_batch = len(training_dataloader)\n    full_y = torch.Tensor([]).to(device)\n    full_logits = torch.Tensor([]).to(device)\n    full_pred = torch.Tensor([]).to(device)\n    \n    model.train()\n    scaler = torch.amp.GradScaler(\"cuda\")\n    for batch, (X, y) in enumerate(training_dataloader):\n        optimizer.zero_grad()\n        X = X.to(device, non_blocking=True)\n        y = y.to(device, non_blocking=True)\n        with torch.amp.autocast(\"cuda\"):\n            logits = model(X).squeeze()\n            loss = loss_function(logits, y.float())\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        \n        with torch.no_grad():\n            preds = (torch.sigmoid(logits) > threshold).float()\n            full_y = torch.cat([full_y, y])\n            full_logits = torch.cat([full_logits, logits])\n            full_pred = torch.cat([full_pred, preds])\n         \n        running_loss += loss.item()\n        avg_loss = running_loss / (batch + 1.)\n        if batch % 250 == 0:\n            writer.add_scalar('Training Loss(avg)', avg_loss, batch + epoch*num_batch)\n            writer.add_scalar('Training Loss (raw)', loss.item(), batch + epoch*num_batch)\n    scheduler.step()\n    writer.flush()\n    return compute_metrics(full_y, full_logits, full_pred, sk_learn_metrics_logits, sk_learn_metrics_pred)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def eval_model(model: nn.Module,\n               dataloader: DataLoader, \n               sk_learn_metrics_logits: tp.List[tp.Callable],\n               sk_learn_metrics_pred: tp.List[tp.Callable],\n               device: torch.cuda.device,\n               threshold: float = 0.5) -> tp.Dict:\n    \n    model.eval()\n    full_y = torch.Tensor([]).to(device)\n    full_logits = torch.Tensor([]).to(device)\n    full_pred = torch.Tensor([]).to(device)\n    \n    with torch.no_grad():\n        for X, y in dataloader:\n            X = X.to(device)\n            y = y.to(device)\n            logits = model(X).squeeze()\n            preds = (torch.sigmoid(logits) > threshold).float()\n\n            full_y = torch.cat([full_y, y])\n            full_logits = torch.cat([full_logits, logits])\n            full_pred = torch.cat([full_pred, preds])\n    return compute_metrics(full_y, full_logits, full_pred, sk_learn_metrics_logits, sk_learn_metrics_pred)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 4. Setup tensorboard for monitoring","metadata":{}},{"cell_type":"code","source":"import threading \nimport tensorboard\nfrom tensorboard import program\n\ndef start_tensorboard(logdir):\n    tb = program.TensorBoard()\n    tb.configure(argv=[None, '--logdir', logdir])\n    url = tb.launch()\n    print(f\"TensorBoard is running at {url}\")\n\n# Replace 'logs' with your actual log directory\nlogdir = exp_dir\ntb_thread = threading.Thread(target=start_tensorboard, args=(logdir,), daemon=True)\n#tb_thread.start() #if you are on a local run uncomment to start tensorboard","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from PIL import Image\n\ndef load_image(path):\n    img = Image.open(path)\n    # Convert to numpy array and add batch dimension (C, H, W)\n    img_array = np.array(img)\n    if len(img_array.shape) == 2:  # Grayscale image\n        img_array = np.expand_dims(img_array, axis=0)  # (1, H, W)\n    else:  # Color image\n        img_array = img_array.transpose(2, 0, 1)  # (C, H, W)\n    return img_array\n    \nwriter = SummaryWriter(exp_dir + '/tensorboard')\nwriter.add_image('training_set_original', load_image(exp_dir + \"/training_set_original.png\"), 0)\nwriter.flush()\nwriter.add_image('training_set_augment',  load_image(exp_dir + \"/training_set_augment.png\"), 0)\nwriter.flush()\nwriter.add_image('training_set_final',  load_image(exp_dir + \"/training_set_final.png\"), 0)\nwriter.flush()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 5. Find best learning rate\n\n> 📄 [**Cancer Image Classification Based on DenseNet Model**](https://arxiv.org/abs/2011.11186)  \n> _by Zhong, Ziliang; Zheng, Muhang; Mai, Huafeng; Zhao, Jianan; Liu, Xinyi_\n\nSuggest to use a learning rate lr = 1e-4 for densenet201. \n\nYou can also plot the loss with respect to the lr evaluated on a few batches.\n\nIt gives insight on which lr to take: between 1e-4 and 1e-3\n\nI have added two fully layers connected layer and two dropout layers to prevent overfitting","metadata":{}},{"cell_type":"code","source":"from torchvision.models import densenet201, DenseNet201_Weights\nmodel = densenet201(weights=DenseNet201_Weights.DEFAULT)\n\nfor params in model.parameters():\n    params.requires_grad = False\n\nmodel.classifier = nn.Sequential(\n    nn.Dropout(0.7),\n    nn.Linear(model.classifier.in_features, 512, bias= True),\n    nn.Dropout(0.5),\n    nn.Linear(512, 1, bias= True))\n\nfor param in model.classifier.parameters():\n    param.requires_grad = True\n\nmodel = model.to(device)\n\ndef custom_lr_find(model : nn.Module, \n                   dataloader: DataLoader,\n                   loss_function: nn.Module,\n                   device: str,\n                   start_lr = 1e-7,\n                   end_lr = 1.0,\n                   num_iteration = 200):\n    rates = []\n    lossses = []\n    model = model.to(device)\n    optimizer = torch.optim.Adam(model.parameters(),lr=start_lr)\n\n    \n    def lr_lambda(iteration):\n        return (end_lr / start_lr) ** (iteration / num_iteration)\n        \n    scheduler = lr_scheduler.LambdaLR(optimizer, lr_lambda)\n    initial_weights = model.state_dict()\n    model.train()\n    \n    X_full = torch.Tensor([]).to(device)\n    y_full = torch.Tensor([]).to(device)\n    \n    for h in range (0, 5):\n        X, y = next(iter(dataloader))\n        X = X.to(device)\n        y = y.to(device)\n        X_full = torch.cat([X_full, X])\n        y_full = torch.cat([y_full, y])\n    \n    for i in range(0, num_iteration):\n        optimizer.zero_grad()\n\n        pred = model(X_full).squeeze()\n        loss = loss_function(pred, y_full.float())\n        lossses.append(loss.item())\n        rates.append(scheduler.get_last_lr()[0])\n        loss.backward()\n        optimizer.step()\n        scheduler.step()\n        model.load_state_dict(initial_weights)\n        if(scheduler.get_last_lr()[0] > end_lr):\n            break\n    return rates, lossses\n        \ndef plot_lr_find(rates, losses, file_name):\n    fig = plt.Figure()\n    plt.plot(rates, losses)\n    plt.xscale('log')\n    plt.xlabel('learning_rate')\n    plt.ylabel('loss')\n    plt.ylim(0.0, 1.0)\n    plt.title('lr_find_results')\n    plt.legend()\n    plt.savefig(file_name)\n    plt.figure()\n    \nrates, losses = custom_lr_find(model, training_dataloader, torch.nn.BCEWithLogitsLoss(), device)\nplot_lr_find(rates, losses, exp_dir + '/lr_find.jpg')\nwriter.add_image('lr_find', load_image(exp_dir + \"/lr_find.jpg\"), 0)\nwriter.flush()","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 6. Using already trained networks: Train the head only\n\nFirst train the head and freeze all other layers","metadata":{}},{"cell_type":"code","source":"from torchvision.models import densenet201, DenseNet201_Weights, densenet121, DenseNet121_Weights\nmodel = densenet201(weights=DenseNet201_Weights.DEFAULT)\n\nfor params in model.parameters():\n    params.requires_grad = False\n\n#Replace the last layer (to output a 1d prediction)\nmodel.classifier = nn.Sequential(\n    nn.Dropout(0.7),\n    nn.Linear(model.classifier.in_features, 512, bias= True),\n    nn.Dropout(0.5),\n    nn.Linear(512, 1, bias= True))\n\nfor param in model.classifier.parameters():\n    param.requires_grad = True\n\nmodel = model.to(device)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#optionnaly load from checkpoint\n'''\nmodel = torch.load('results/pcam/19_06_2025_11_08_15/model_'+str(5)+'.pt', weights_only = False)\nfor params in model.parameters():\n    params.requires_grad = False\nfor param in model.classifier.parameters():\n    param.requires_grad = True\nmodel = model.to(device)\n'''","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"lr = 1e-4\n\noptimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=1e-4)\nloss_func = torch.nn.BCEWithLogitsLoss()\nscheduler = lr_scheduler.StepLR(optimizer, step_size=1000, gamma=0.01)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report, roc_auc_score, f1_score, precision_score, recall_score, accuracy_score, classification_report\nimport time\nepoch_num = 2\nsk_learn_metrics_logits = [roc_auc_score]\nsk_learn_metrics_pred = [f1_score, accuracy_score]\nfor i in range(0, epoch_num):\n    start_time = time.time()\n    train_res = run_one_epoch(model,\n                  training_dataloader,\n                  optimizer,\n                  loss_func,\n                  scheduler,\n                  device,\n                  writer,\n                  i,\n                  sk_learn_metrics_logits,\n                  sk_learn_metrics_pred)\n    end_time = time.time()\n    print(\"epoch n°: \", i, \" training time : \", end_time-start_time, \" sec\")\n    start_time = time.time()\n    val_res = eval_model(model, val_dataloader, sk_learn_metrics_logits, sk_learn_metrics_pred, device)\n    for key in train_res.keys():\n        writer.add_scalars(key, {\"Train \" + key: train_res[key], \"Val \"+ key : val_res[key]}, i*len(training_dataloader))\n    end_time = time.time()\n    print(\"epoch n°: \", i, \" evaluation time : \", end_time-start_time, \" sec\")\n    torch.save(model, exp_dir+\"/model_\" + str(i) + \".pt\")\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 7. Using already trained networks: Fine Tune a few layers\nI did not use it in the end, this is optional","metadata":{}},{"cell_type":"code","source":"'''\nfor name, param in model.features.denseblock4.denselayer32.conv1.named_parameters():\n    param.requires_grad = True\n    \nfor name, param in model.features.denseblock4.denselayer32.conv2.named_parameters():\n    param.requires_grad = True\n'''","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Unfreeze last two blocks (features.6 and features.7)\n'''\nlr = 1e-4\n#optimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=1e-4)\n#loss_func = torch.nn.BCEWithLogitsLoss(pos_weight=pos_weight)\nloss_func = torch.nn.BCEWithLogitsLoss()\n# Use lower LR for fine-tuning\noptimizer = torch.optim.Adam([\n    {\"params\": model.classifier.parameters(), \"lr\": 1e-4},\n     {\"params\": model.features.denseblock4.denselayer32.conv1.parameters(), \"lr\": 1e-5},\n     {\"params\": model.features.denseblock4.denselayer32.conv2.parameters(), \"lr\": 1e-5},\n ])\n'''","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"'''\nfrom sklearn.metrics import classification_report, roc_auc_score, f1_score, precision_score, recall_score, accuracy_score, classification_report\nimport time\nsk_learn_metrics_logits = [roc_auc_score]\nsk_learn_metrics_pred = [f1_score, accuracy_score]\nepoch_num = 2\nfinetune_epoch_num = 6\nfor i in range(epoch_num, epoch_num + finetune_epoch_num):\n    start_time = time.time()\n    train_res = run_one_epoch(model,\n                  training_dataloader,\n                  optimizer,\n                  loss_func,\n                  scheduler,\n                  device,\n                  writer,\n                  i,\n                  sk_learn_metrics_logits,\n                  sk_learn_metrics_pred)\n    end_time = time.time()\n    print(\"epoch n°: \", i, \" training time : \", end_time-start_time, \" sec\")\n    start_time = time.time()\n    val_res = eval_model(model, val_dataloader, sk_learn_metrics_logits, sk_learn_metrics_pred, device)\n    for key in train_res.keys():\n        writer.add_scalars(key, {\"Train \" + key: train_res[key], \"Val \"+ key : val_res[key]}, i*len(training_dataloader))\n    end_time = time.time()\n    print(\"epoch n°: \", i, \" evaluation time : \", end_time-start_time, \" sec\")\n    torch.save(model, exp_dir+\"/model_\" + str(i) + \".pt\")\n\n'''","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 8. Fine tune the entire model","metadata":{}},{"cell_type":"code","source":"for params in model.parameters():\n    params.requires_grad = True\n\nlr = 1e-4\noptimizer = torch.optim.Adam(model.parameters(), lr=lr, weight_decay=1e-4)\nloss_func = torch.nn.BCEWithLogitsLoss()\nscheduler = lr_scheduler.StepLR(optimizer, step_size=1000, gamma=0.01)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report, roc_auc_score, f1_score, precision_score, recall_score, accuracy_score, classification_report\nimport time\nsk_learn_metrics_logits = [roc_auc_score]\nsk_learn_metrics_pred = [f1_score, accuracy_score]\nepoch_num = 2\nfinetune_epoch_num = 5\n    \nfor i in range(epoch_num, epoch_num + finetune_epoch_num):\n    start_time = time.time()\n    train_res = run_one_epoch(model,\n                  training_dataloader,\n                  optimizer,\n                  loss_func,\n                  scheduler,\n                  device,\n                  writer,\n                  i,\n                  sk_learn_metrics_logits,\n                  sk_learn_metrics_pred)\n    end_time = time.time()\n    print(\"epoch n°: \", i, \" training time : \", end_time-start_time, \" sec\")\n    start_time = time.time()\n    val_res = eval_model(model, val_dataloader, sk_learn_metrics_logits, sk_learn_metrics_pred, device)\n    for key in train_res.keys():\n        writer.add_scalars(key, {\"Train \" + key: train_res[key], \"Val \"+ key : val_res[key]}, i*len(training_dataloader))\n    end_time = time.time()\n    print(\"epoch n°: \", i, \" evaluation time : \", end_time-start_time, \" sec\")\n    torch.save(model, exp_dir+\"/model_\" + str(i) + \".pt\")","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 9. Compute test set prediction and submit to kaggle\n\nWe will use TTA (Test Time with Augmentation).\nWe can also optionally use several models to make a prediction and average the results","metadata":{}},{"cell_type":"code","source":"def run_inference(model: nn.Module,\n                     dataloader: DataLoader, \n                     device: torch.cuda.device):\n    \n    model.eval()\n    full_y = torch.Tensor([]).to(device)\n    full_logits = torch.Tensor([]).to(device)\n    \n    with torch.no_grad():\n        for X, y in dataloader:\n            X = X.to(device)\n            y = y.to(device)\n            logits = model(X).squeeze()\n\n            full_y = torch.cat([full_y, y])\n            full_logits = torch.cat([full_logits, logits])\n\n    return full_y, full_logits","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i in range(0, epoch_num + finetune_epoch_num):\n    models_paths = [exp_dir+\"/model_\" + str(i) + \".pt\"]\n    pcam_model = torch.load(models_paths[0], weights_only = False)\n    pcam_model = pcam_model.to(device)\n\n    # First create tta_num augmented dataloaders\n    tta_num = 1\n    logits = []\n    for j in range(0, tta_num):\n        test_set_augment = PcamDatasetKaggle(root=dataset_dir, split=\"test\", transform = deepcopy(transform_data_augment)) #For TTA\n        test_dataloader_augment = DataLoader(test_set_augment, batch_size=batch_size, shuffle=False, pin_memory=True, num_workers=6, persistent_workers = True)\n        for modelp in models_paths:\n            pcam_model = torch.load(modelp, weights_only = False)\n            pcam_model = pcam_model.to(device)\n            test_y, test_logits = run_inference(pcam_model, test_dataloader, device)\n            logits.append(test_logits)\n            test_y_augm, test_logits_aum = run_inference(pcam_model, test_dataloader_augment, device)\n            logits.append(test_logits_aum)\n        \n    # Average logits\n    logits_stacked = torch.stack(logits)\n    mean_logits = torch.mean(logits_stacked, dim = 0, keepdims=True)\n\n    #Create submission file with final predictions\n    image_ids = [img.replace('.tif', '') for img in test_set.imgs.tolist()]\n    test_preds = torch.sigmoid(mean_logits)\n\n    submission_df = pd.DataFrame({\n        'id': image_ids,\n        'label': test_preds.squeeze().detach().cpu().numpy()\n    })\n\n    submission_df.to_csv(exp_dir+'/submission_'+str(i)+'.csv', index=False)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub_path = exp_dir + '/submission_4.csv'\nmodel_path = models_paths[0]\n#you need to update your creds at the top for this\n#!kaggle competitions submit -c histopathologic-cancer-detection -f {sub_path} -m {model_path}","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 11. Find best threshold for prediction on validation set","metadata":{}},{"cell_type":"code","source":"i = 4\nmodels_paths = [exp_dir+\"/model_\" + str(i) + \".pt\"]\npcam_model = torch.load(models_paths[0], weights_only = False)\npcam_model = pcam_model.to(device)\ntest_y, test_logits = run_inference(pcam_model, val_dataloader, device)\ntest_y_augment, test_logits_augment = run_inference(pcam_model, val_dataloader_augment, device)\nfull_y = torch.cat([test_y, test_y_augment])\nfull_logits = torch.cat([test_logits, test_logits_augment])","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import roc_curve, auc\nfpr, tpr, thresholds = roc_curve(full_y.detach().cpu().numpy(), torch.sigmoid(full_logits).detach().cpu().numpy())\nroc_auc = auc(fpr, tpr)","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plt.figure(figsize=(8,6))\nplt.plot(fpr, tpr, color='orange', lw=2, label=f'ROC curve (AUC = {roc_auc})')\nplt.xlim([0.0, 1.0])\nplt.ylim([0.0, 1.0])\nplt.xlabel('False Positive Rate')\nplt.ylabel('True Positive Rate')\nplt.title('Receiver Operating Characteristic')\nplt.grid(alpha=0.3)\nplt.show()\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Find best threshold index (maximize TPR-FPR).\nj_scores = tpr - fpr\nbest_idx = np.argmax(j_scores)\nbest_threshold = thresholds[best_idx]\n","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"best_threshold","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null}]}