{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"pygments_lexer":"ipython3","nbconvert_exporter":"python","version":"3.6.4","file_extension":".py","codemirror_mode":{"name":"ipython","version":3},"name":"python","mimetype":"text/x-python"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport time\nimport numpy as np\nimport pandas as pd\n# image manipulation\nimport cv2\nimport PIL\nfrom PIL import Image\n\n# visualisation\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\n# helpers\nfrom tqdm import tqdm\nimport time\nimport copy\nimport gc\nfrom enum import Enum\n\n\n# for cnn\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.optim import Adam, SGD\nfrom torch.autograd import Variable\nfrom torch.utils.data import DataLoader, random_split, TensorDataset, Dataset, WeightedRandomSampler\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau, StepLR\nfrom torchvision import models\nfrom torchmetrics.classification import BinaryF1Score, BinaryPrecision, BinaryRecall, BinaryAccuracy, BinaryROC, BinaryAUROC\nfrom torchvision import transforms","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2023-04-04T10:04:12.909877Z","iopub.execute_input":"2023-04-04T10:04:12.910244Z","iopub.status.idle":"2023-04-04T10:04:12.91947Z","shell.execute_reply.started":"2023-04-04T10:04:12.910215Z","shell.execute_reply":"2023-04-04T10:04:12.918468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"csvpathtrain = '/kaggle/input/rsna-breast-cancer-detection/train.csv'\n\ndftrain = pd.read_csv(csvpathtrain)\ndftrain.head()","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:04:12.921181Z","iopub.execute_input":"2023-04-04T10:04:12.922095Z","iopub.status.idle":"2023-04-04T10:04:13.002758Z","shell.execute_reply.started":"2023-04-04T10:04:12.922054Z","shell.execute_reply":"2023-04-04T10:04:13.001578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axes = plt.subplots(1, 2, figsize=(10, 5))\n########## PLOTING CANCER ################\nsplot = sns.countplot(ax = axes[0], x = dftrain['cancer'])\n\ns = dftrain['cancer'].value_counts()\naxes[1].pie(s, autopct=\"%.1f%%\", labels = s.keys())\nfig.suptitle('Cancer distribution')","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:04:13.004394Z","iopub.execute_input":"2023-04-04T10:04:13.005362Z","iopub.status.idle":"2023-04-04T10:04:13.311978Z","shell.execute_reply.started":"2023-04-04T10:04:13.005325Z","shell.execute_reply":"2023-04-04T10:04:13.310745Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"total_samples = len(dftrain['cancer'])\npositive_samples = sum(dftrain['cancer'] == 1)\nnegative_samples = total_samples - positive_samples\nprint(f\"{total_samples}, {positive_samples}, {negative_samples}\")","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:04:13.314325Z","iopub.execute_input":"2023-04-04T10:04:13.314898Z","iopub.status.idle":"2023-04-04T10:04:13.333167Z","shell.execute_reply.started":"2023-04-04T10:04:13.314859Z","shell.execute_reply":"2023-04-04T10:04:13.332292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"samples_weight = torch.Tensor([positive_samples / total_samples, negative_samples / total_samples]).type(dtype = torch.float32)\n\nsamples_weight","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:04:13.334465Z","iopub.execute_input":"2023-04-04T10:04:13.335051Z","iopub.status.idle":"2023-04-04T10:04:13.347991Z","shell.execute_reply.started":"2023-04-04T10:04:13.335016Z","shell.execute_reply":"2023-04-04T10:04:13.346901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNAMamographyDataset(Dataset):\n    def __init__(self, annotations_file, img_dir, transform=None):\n        self.df = pd.read_csv(annotations_file)\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n    \n\n\n    def __getitem__(self, ind):\n        \n        img_path = f\"{self.img_dir}/{self.df.iloc[ind].patient_id}_{self.df.iloc[ind].image_id}.png\"\n        img = Image.open(img_path).convert('RGB')\n        \n        label = self.df.iloc[ind].cancer\n        # there is no need to normalize data, it has already been normalized\n        if self.transform:\n            img = self.transform(img).to(torch.float32) \n        else:\n            default_transform = transforms.Compose([transforms.ToTensor()])\n            img = default_transform(img).to(torch.float32)\n            \n        #sample = {\"image\" : img, \"label\": label}\n        return img, label","metadata":{"execution":{"iopub.status.busy":"2023-04-04T12:18:28.682097Z","iopub.execute_input":"2023-04-04T12:18:28.683302Z","iopub.status.idle":"2023-04-04T12:18:28.694207Z","shell.execute_reply.started":"2023-04-04T12:18:28.683256Z","shell.execute_reply":"2023-04-04T12:18:28.693011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_csv = '/kaggle/input/rsna-breast-cancer-detection/train.csv'\nimgs_dir = '/kaggle/input/rsnamamorgaphybreastcancerrecognition512x512'\n\naugmentator = transforms.Compose([\n    # input for augmentator is always PIL image\n    # transforms.ToPILImage(),\n    transforms.RandomHorizontalFlip(0.5),\n    transforms.RandomVerticalFlip(0.5),\n    transforms.RandomRotation(5),\n    transforms.ToTensor(), # return it as a tensor and transforms it to [0, 1]\n])\ndataset = RSNAMamographyDataset(train_csv, imgs_dir, augmentator)","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:04:13.368004Z","iopub.execute_input":"2023-04-04T10:04:13.368588Z","iopub.status.idle":"2023-04-04T10:04:13.461895Z","shell.execute_reply.started":"2023-04-04T10:04:13.368552Z","shell.execute_reply":"2023-04-04T10:04:13.460815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Use torch.utils.data to create a DataLoader \n# that will take care of creating batches \n\n# TODO, remove using half of dataset\n# dataset, _ = random_split(dataset, [int(len(dataset)*0.02), int(len(dataset)*0.98 + 1)])\n# split training into validation and train\nval_pct = 0.1\nval_size = int(val_pct * len(dataset))\ntrain_size = len(dataset) - val_size\ntrain_dataset, val_dataset = random_split(dataset, [train_size, val_size])\n\n","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:04:13.466559Z","iopub.execute_input":"2023-04-04T10:04:13.468855Z","iopub.status.idle":"2023-04-04T10:04:13.481681Z","shell.execute_reply.started":"2023-04-04T10:04:13.468816Z","shell.execute_reply":"2023-04-04T10:04:13.480763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Class counting...\")\nlabels = dftrain['cancer'].values\nclass_sample_count = np.array([len(np.where(labels == l)[0]) for l in np.unique(labels)])\n\n\n# the trouble with this aproach is that it now has to load all images one by one and label them\n# but it saves RAM memory in training process\n#class_sample_count = np.zeros(2)\n\n#print(\"Class counting...\")\n#for _, label in tqdm(train_dataset):\n#    class_sample_count[label] += 1\n\nprint(class_sample_count)\n\n# This maybe apply, maybe not\n# since there is big class imbalance, we will not sample positive class THAT frequent\n# to be closer to 'reality, every fifth image will be cancer (instead of 50/50 distribution)'\nclass_sample_count[1] *= 5\nclass_weights = 1. / class_sample_count\n\nprint(\"Adding weights to each training sample...\")\nsample_weights = []\nfor _, label in tqdm(train_dataset):\n    sample_weights.append(class_weights[label])\n\nsample_weights = np.array(sample_weights)\nsample_weights = torch.from_numpy(sample_weights)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:04:13.485884Z","iopub.execute_input":"2023-04-04T10:04:13.48815Z","iopub.status.idle":"2023-04-04T10:16:50.492707Z","shell.execute_reply.started":"2023-04-04T10:04:13.488114Z","shell.execute_reply":"2023-04-04T10:16:50.491147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"weighted_random_sampler = WeightedRandomSampler(sample_weights, len(sample_weights))","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:50.496707Z","iopub.execute_input":"2023-04-04T10:16:50.497018Z","iopub.status.idle":"2023-04-04T10:16:50.501974Z","shell.execute_reply.started":"2023-04-04T10:16:50.496991Z","shell.execute_reply":"2023-04-04T10:16:50.50091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nbatch_size = 32\n\n# Applying random sampler just tu train dataset, not for validation, since the validation dataset should be imitation of 'real' DS\ntrain_dataloader = DataLoader(train_dataset, batch_size=batch_size, num_workers = 2, pin_memory = True, sampler = weighted_random_sampler)\nval_dataloader = DataLoader(val_dataset, batch_size=batch_size, shuffle = True, pin_memory = True)\n","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:50.503446Z","iopub.execute_input":"2023-04-04T10:16:50.504001Z","iopub.status.idle":"2023-04-04T10:16:50.514859Z","shell.execute_reply.started":"2023-04-04T10:16:50.503958Z","shell.execute_reply":"2023-04-04T10:16:50.513898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataloaders = {'train' : train_dataloader, 'val' : val_dataloader}\ndataset_sizes = {'train': train_size, 'val' : val_size}","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:50.516236Z","iopub.execute_input":"2023-04-04T10:16:50.516711Z","iopub.status.idle":"2023-04-04T10:16:50.52841Z","shell.execute_reply.started":"2023-04-04T10:16:50.516678Z","shell.execute_reply":"2023-04-04T10:16:50.527499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(train_dataset), len(val_dataset))\nprint(len(train_dataloader), len(val_dataloader))","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:50.529386Z","iopub.execute_input":"2023-04-04T10:16:50.52966Z","iopub.status.idle":"2023-04-04T10:16:50.539149Z","shell.execute_reply.started":"2023-04-04T10:16:50.529636Z","shell.execute_reply":"2023-04-04T10:16:50.538089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"rows = 5\ncols = 5\nplt.subplots(rows, cols, figsize = (20, 20))\n\nbatch_imgs, batch_labels = next(iter(train_dataloader))\ni = 0\nfor img in batch_imgs:\n    if i >= rows*cols:\n        break\n    plt.subplot(rows, cols, i + 1)\n    plt.title(\"Cancer\" if batch_labels[i] == 1 else \"No cancer\")\n    plt.imshow(img.permute(1, 2, 0))\n\n    i += 1\n\nlabels_count = np.zeros(2)\nfor l in batch_labels:\n    labels_count[l] += 1 \n    \nprint(f'There are {labels_count[0]} negative and {labels_count[1]} positive samples in this batch.')","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:50.540534Z","iopub.execute_input":"2023-04-04T10:16:50.540876Z","iopub.status.idle":"2023-04-04T10:16:58.304928Z","shell.execute_reply.started":"2023-04-04T10:16:50.540844Z","shell.execute_reply":"2023-04-04T10:16:58.303788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img.size()","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:58.306137Z","iopub.execute_input":"2023-04-04T10:16:58.306469Z","iopub.status.idle":"2023-04-04T10:16:58.313786Z","shell.execute_reply.started":"2023-04-04T10:16:58.306425Z","shell.execute_reply":"2023-04-04T10:16:58.312905Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'Current device is {device}')","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:58.314985Z","iopub.execute_input":"2023-04-04T10:16:58.316273Z","iopub.status.idle":"2023-04-04T10:16:58.322585Z","shell.execute_reply.started":"2023-04-04T10:16:58.316155Z","shell.execute_reply":"2023-04-04T10:16:58.321747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\nclass CNN(nn.Module):\n    def __init__(self):\n        super(CNN, self).__init__()\n        self.network = models.resnet18(pretrained=True)\n        n_features = self.network.fc.out_features\n        print(n_features)\n        # add additional layer that maps 2048 extracted features from resnet to 1 feature determining the class\n        self.classifier_layer = nn.Sequential(\n            nn.Linear(n_features , 256),\n            nn.Dropout(0.3),\n            nn.Linear(256 , 1)\n        )\n    \n    def forward(self, xb):        \n        xb = self.network(xb)\n        xb = self.classifier_layer(xb)\n        return torch.sigmoid(xb)","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:58.323855Z","iopub.execute_input":"2023-04-04T10:16:58.32493Z","iopub.status.idle":"2023-04-04T10:16:58.332947Z","shell.execute_reply.started":"2023-04-04T10:16:58.324895Z","shell.execute_reply":"2023-04-04T10:16:58.332089Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# create class for earlystopping\nclass EarlyStopper:\n    def __init__(self, patience=1, min_delta=0):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.counter = 0\n        self.min_loss = np.inf\n\n    def early_stop(self, loss):\n        if loss <= self.min_loss:\n            self.min_loss = loss\n            self.counter = 0\n        elif loss > (self.min_loss + self.min_delta):\n            self.counter += 1\n            if self.counter >= self.patience:\n                return True\n        return False","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:58.334095Z","iopub.execute_input":"2023-04-04T10:16:58.335297Z","iopub.status.idle":"2023-04-04T10:16:58.343863Z","shell.execute_reply.started":"2023-04-04T10:16:58.335252Z","shell.execute_reply":"2023-04-04T10:16:58.342779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def BCELoss_class_weighted(weights):\n    \"\"\"\n    weights[0] is weight for class 0 (negative class)\n    weights[1] is weight for class 1 (positive class)\n    \"\"\"\n    def loss(y_pred, target):\n        y_pred = torch.clamp(y_pred,min=1e-7,max=1-1e-7) # for numerical stability\n        bce = - weights[1] * target * torch.log(y_pred) - (1 - target) * weights[0] * torch.log(1 - y_pred)\n        return torch.mean(bce)\n\n    return loss","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:58.345128Z","iopub.execute_input":"2023-04-04T10:16:58.346169Z","iopub.status.idle":"2023-04-04T10:16:58.356443Z","shell.execute_reply.started":"2023-04-04T10:16:58.346136Z","shell.execute_reply":"2023-04-04T10:16:58.355389Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# defining the model for determining LR\nmodel = CNN()\nmodel.to(device)\n# convrt weights to cuda.float if cuda is avaliable\n#if torch.cuda.is_available():\n#    model.cuda()\n\n# defining the optimizer\noptimizer = Adam(model.parameters(), lr=1e-07)\n\n\n# defining the loss function\n# Binary cross entropy is chosen because it is the classification problem\n#labels = dftrain['cancer'].values\n#w_neg = sum(labels == 0) / len(labels)\n#w_pos = sum(labels == 1) / len(labels)\n#print(f\"Class weight: {w_neg}\")\n#criterion = BCELoss_class_weighted(weights = [w_neg, w_pos])\n# criterion = nn.BCEWithLogitsLoss()\n\nw_pos = 3\nw_neg = 1\nprint(f\"Class weight for negative class: {w_neg}, and for positive {w_pos}\")\ncriterion = BCELoss_class_weighted(weights = [w_neg, w_pos])\n\nmetric = BinaryF1Score().to(device)\n\n# print(model)","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:58.35776Z","iopub.execute_input":"2023-04-04T10:16:58.358718Z","iopub.status.idle":"2023-04-04T10:16:59.122934Z","shell.execute_reply.started":"2023-04-04T10:16:58.358685Z","shell.execute_reply":"2023-04-04T10:16:59.121901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:59.124275Z","iopub.execute_input":"2023-04-04T10:16:59.125153Z","iopub.status.idle":"2023-04-04T10:16:59.133758Z","shell.execute_reply.started":"2023-04-04T10:16:59.125123Z","shell.execute_reply":"2023-04-04T10:16:59.132514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def determine_lrs_and_losses(model, criterion, optimizer, metric, num_epochs=25, final_lr = 1e-02, total_batches = 1):\n    since = time.time()\n    \n    lr_list = []\n    loss_list = []\n    \n    train_metrics = {'loss' : [], 'acc' : [], 'f1': []}\n    val_metrics = {'loss' : [], 'acc' : [], 'f1': []}\n    \n    \n    print('Starting training...')\n    print('-' * 20)\n    for epoch in range(num_epochs):\n        \n        # Each epoch has a training and validation phase\n        for phase in ['train']:\n            \n            \n            if phase == 'train':\n                model.train()  # Set model to training mode\n            else:\n                model.eval()   # Set model to evaluate mode\n\n            running_loss = 0.0\n            running_corrects = 0\n            running_f1 = 0.0\n            \n           \n            gc.collect()\n            \n            current_batch = 0\n            # Iterate over data.\n            for inputs, labels in dataloaders[phase]:\n                \n                labels = torch.unsqueeze(labels.to(torch.float32), 1)\n                current_batch += 1\n                if current_batch > total_batches:\n                    break\n\n                inputs = inputs.to(device)\n                labels = labels.to(device)\n\n                # zero the parameter gradients\n                optimizer.zero_grad()\n\n                # forward\n                # track history if only in train\n                with torch.set_grad_enabled(phase == 'train'):\n                    outputs = model(inputs)\n                    # this was different, it took max of output and 1\n                    # output should never be higher than 1, so it is confusing\n                    preds = outputs > 0.5\n                    # _, preds = torch.max(outputs, 1)\n                    loss = criterion(outputs.double(), labels)\n\n                    # backward + optimize only if in training phase\n                    if phase == 'train':\n                        loss.backward()\n                        optimizer.step()  \n\n                    #print(labels.detach().numpy().type,  outputs.detach().numpy().type)\n                #running_f1 += f1_score(labels.detach().numpy(), outputs.detach().numpy())\n                running_f1 += metric(outputs, labels)\n\n                # statistics\n                running_loss += loss.item() \n                #print(f'{phase}, {inputs.size(0)}, {preds.size()} {torch.squeeze(labels.data).size()}')\n                running_corrects += torch.sum(preds == labels.data)\n                \n                gc.collect()\n\n                \n            \n            epoch_loss = running_loss / total_batches \n            epoch_acc = running_corrects.double() / (total_batches * batch_size)\n            epoch_f1 = running_f1 / total_batches\n            if phase == 'train':\n                train_metrics['loss'].append(epoch_loss)\n                train_metrics['acc'].append(epoch_acc)\n                train_metrics['f1'].append(epoch_f1)\n\n            else:\n                val_metrics['loss'].append(epoch_loss)\n                val_metrics['acc'].append(epoch_acc)\n                val_metrics['f1'].append(epoch_f1)\n\n                \n        train_loss_l, train_acc_l, train_f1_l = train_metrics['loss'][-1], train_metrics['acc'][-1], train_metrics['f1'][-1] # cant be formated in string, so should be segregated separately\n        lr = optimizer.param_groups[0]['lr']\n        print(f'Epoch {epoch + 1}/{num_epochs}, Train Loss: {train_loss_l:.4f}, Train Acc: {train_acc_l:.4f}, Train f1: {train_f1_l:.4f}, learning rate: {lr}')\n\n        # set learning rate for optimizer for determining initial learning rate\n        for g in optimizer.param_groups:\n            g['lr'] *= 4\n        \n        \n        loss_list.append(train_loss_l) # the goal is to determine which learning rate results\n        # in steepest training loss difference\n        lr_list.append(optimizer.param_groups[0]['lr'])\n        \n        if optimizer.param_groups[0]['lr'] > final_lr:\n            break\n\n        \n\n    return lr_list, loss_list","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:59.135652Z","iopub.execute_input":"2023-04-04T10:16:59.136091Z","iopub.status.idle":"2023-04-04T10:16:59.153893Z","shell.execute_reply.started":"2023-04-04T10:16:59.136055Z","shell.execute_reply":"2023-04-04T10:16:59.152917Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# function that finds steepest descent in training loss\ndef determine_init_lr(lr_list, loss_list):\n    # find difference beetwen succesive losses\n    diffs = [j-i for i, j in zip(loss_list[:-1], loss_list[1:])]\n    # find where loss change is maximum\n    max_value_ind = np.argmin(diffs) + 1\n    # get learning rate for that change\n    print(f\"Learning rate {lr_list[max_value_ind]} resulted in biggest loss decrease and should be starting learning rate for this neural net\")\n\n    init_lr = lr_list[max_value_ind]\n    return init_lr","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:59.155306Z","iopub.execute_input":"2023-04-04T10:16:59.155805Z","iopub.status.idle":"2023-04-04T10:16:59.167207Z","shell.execute_reply.started":"2023-04-04T10:16:59.155771Z","shell.execute_reply":"2023-04-04T10:16:59.166289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_lr_over_loss(lr_list, loss_list, init_lr):\n    lr_ind = lr_list.index(init_lr)\n    plt.figure(figsize = (15, 7))\n    p1 = plt.plot(lr_list, loss_list)\n    p2 = plt.scatter(lr_list, loss_list)\n    p3 = plt.scatter(lr_list[lr_ind], loss_list[lr_ind], marker = 'D', s = 80, color = 'r')\n    plt.legend((p2, p3), (\"all considered learning rates\", \"best learning rate\"))\n    plt.xlabel(\"Learning rate\")\n    plt.ylabel(\"Loss\")\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:59.16871Z","iopub.execute_input":"2023-04-04T10:16:59.169483Z","iopub.status.idle":"2023-04-04T10:16:59.180181Z","shell.execute_reply.started":"2023-04-04T10:16:59.16944Z","shell.execute_reply":"2023-04-04T10:16:59.179219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lr_list, loss_list = determine_lrs_and_losses(model, criterion, optimizer, metric, num_epochs=25, final_lr = 1e-02, total_batches = 10)\ninit_lr = determine_init_lr(lr_list, loss_list)\nplot_lr_over_loss(lr_list, loss_list, init_lr)","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:16:59.1839Z","iopub.execute_input":"2023-04-04T10:16:59.184159Z","iopub.status.idle":"2023-04-04T10:18:04.503646Z","shell.execute_reply.started":"2023-04-04T10:16:59.184135Z","shell.execute_reply":"2023-04-04T10:18:04.502641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# defining the model\nmodel = CNN()\nmodel.to(device)\n# convrt weights to cuda.float if cuda is avaliable\n#if torch.cuda.is_available():\n#    model.cuda()\n\n# defining the optimizer\noptimizer = Adam(model.parameters(), lr=init_lr)\n# defining learning rate schedualer to fight plateues\n# TODO: figure out how to measure validation loss independently\n# scheduler = ReduceLROnPlateau(optimizer, 'min', patience = 5)\nscheduler = StepLR(optimizer, step_size=5, gamma=0.1)\n# defining the loss function\n# Binary cross entropy is chosen because it is the classification problem\nlabels = dftrain['cancer'].values\n# the weight should be smaller if class count is higher\nneg_count = sum(labels == 0)\npos_count = sum(labels == 1)\nw_pos = 2\nw_neg = 1\nprint(f\"Class weight for negative class: {w_neg}, and for positive {w_pos}\")\ncriterion = BCELoss_class_weighted(weights = [w_neg, w_pos])\n# criterion = nn.BCEWithLogitsLoss()\n# define early stopping\nearlystoper = EarlyStopper(patience = 3)\n\n\ncheckpoint = {'model': CNN(),\n          'state_dict': model.state_dict(),\n          'optimizer' : optimizer.state_dict(),\n             'threshold' : 0.5}\n\n\n# print(model)","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:18:04.505329Z","iopub.execute_input":"2023-04-04T10:18:04.505708Z","iopub.status.idle":"2023-04-04T10:18:05.142256Z","shell.execute_reply.started":"2023-04-04T10:18:04.50567Z","shell.execute_reply":"2023-04-04T10:18:05.141111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def find_optim_thres(fpr, tpr, thresholds):\n    optim_thres = thresholds[0]\n    inx = 0\n    min_dist = 1.0\n    for i in range(len(fpr)):\n        dist = np.linalg.norm(np.array([0.0, 1.0]) - np.array([fpr[i], tpr[i]]))\n        if dist < min_dist:\n            min_dist = dist\n            optim_thres = thresholds[i]\n            inx = i\n            \n    return optim_thres, inx\n        ","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:18:05.14375Z","iopub.execute_input":"2023-04-04T10:18:05.144399Z","iopub.status.idle":"2023-04-04T10:18:05.151278Z","shell.execute_reply.started":"2023-04-04T10:18:05.144361Z","shell.execute_reply":"2023-04-04T10:18:05.150172Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, criterion, optimizer, scheduler, num_epochs=25):\n    since = time.time()\n    \n    metricf1 = BinaryF1Score()\n    precision = BinaryPrecision()\n    recall = BinaryRecall()\n    accuracy = BinaryAccuracy()\n    roc = BinaryROC()\n    auc = BinaryAUROC()\n    \n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_f1 = -1.0\n    \n    train_metrics = {'loss' : [], 'acc' : [], 'f1': [], 'precision': [], 'recall': [], 'auc': []}\n    val_metrics = {'loss' : [], 'acc' : [], 'f1': [], 'precision': [], 'recall': [], 'auc': []}\n    \n    \n    # inital threshold for first epoch, it will change afterwards\n    threshold = 0.5\n    \n    print('Starting training...')\n    print('-' * 20)\n    for epoch in range(num_epochs):\n        \n\n        # Each epoch has a training and validation phase\n        for phase in ['train', 'val']:\n            # empty 'all' tensors for saving\n            # for calculating aoc at the end of epoch, and for calculating new threshold\n            all_outputs = torch.Tensor([])\n            all_labels = torch.Tensor([])\n            if phase == 'train':\n                model.train()  # Set model to training mode\n            else:\n                model.eval()   # Set model to evaluate mode\n\n            running_loss = 0.0\n            n_samples = 0\n            \n            n_correct = 0\n            running_f1 = 0.0\n            # Iterate over data.\n            print(f'{phase} for epoch {epoch + 1}')\n            for inputs, labels in tqdm(dataloaders[phase]):\n                \n                labels = torch.unsqueeze(labels.to(torch.float32), 1)\n                \n                inputs = inputs.to(device)\n                labels = labels.to(device)\n\n                # zero the parameter gradients\n                optimizer.zero_grad()\n\n                # forward\n                # track history if only in train\n                with torch.set_grad_enabled(phase == 'train'):\n                    outputs = model(inputs)\n                    preds = (outputs > threshold).double()\n                    #print(all_outputs)\n                    #print(outputs)\n                    # concatenating all outputs and labels for calculation aoc and new threshold\n                    all_outputs = torch.cat((all_outputs, outputs.to('cpu')))\n                    all_labels = torch.cat((all_labels, labels.to('cpu')))\n                    \n                    #print(labels)\n                    # _, preds = torch.max(outputs, 1)\n                    loss = criterion(outputs, labels)\n\n                    # backward + optimize only if in training phase\n                    if phase == 'train':\n                        loss.backward()\n                        optimizer.step()\n\n                # statistics\n                # n_samples += labels.size(0)\n                running_loss += loss.item()\n                # n_correct += (preds == labels).sum().item()\n                # running_f1 += metric(outputs, labels) \n\n\n                # collect any unused memmory\n                gc.collect()\n                torch.cuda.empty_cache()\n            \n            # statistics\n            epoch_loss = running_loss / len(dataloaders[phase])\n            \n            # find true positive and false positive rates for ROC curve\n            fpr, tpr, thresholds = roc(all_outputs, all_labels)\n            epoch_auc = auc(all_outputs, all_labels)\n            # find new threshold\n            threshold, _ = find_optim_thres(fpr, tpr, thresholds)\n            print(f'New threshold is {threshold}')\n            # calculate metrics using new optimized threshold\n            epoch_f1 = metricf1(all_outputs > threshold, all_labels)\n            epoch_acc = accuracy(all_outputs > threshold, all_labels)\n            epoch_precision = precision(all_outputs > threshold, all_labels)\n            epoch_recall = recall(all_outputs > threshold, all_labels)\n            \n            # save all of the statistics for latter analysis\n            if phase == 'train':\n                scheduler.step()\n                train_metrics['loss'].append(epoch_loss)\n                train_metrics['acc'].append(epoch_acc)\n                train_metrics['f1'].append(epoch_f1)\n                train_metrics['precision'].append(epoch_precision)\n                train_metrics['recall'].append(epoch_recall)\n                train_metrics['auc'].append(epoch_auc)\n\n\n            else:\n                val_metrics['loss'].append(epoch_loss)\n                val_metrics['acc'].append(epoch_acc)\n                val_metrics['f1'].append(epoch_f1)\n                val_metrics['precision'].append(epoch_precision)\n                val_metrics['recall'].append(epoch_recall)\n                val_metrics['auc'].append(epoch_auc)\n\n\n\n            # deep copy the model\n            if phase == 'val' and epoch_f1 > best_f1:\n                best_f1 = epoch_f1\n                best_model_wts = copy.deepcopy(model.state_dict())\n                checkpoint['threshold'] = threshold\n                torch.save(checkpoint, 'checkpoint.pth')\n\n                \n        # cant be formated in string\n        tr_loss, tr_acc, tr_f1, tr_prec, tr_rec, tr_auc = train_metrics['loss'][-1], train_metrics['acc'][-1],  train_metrics['f1'][-1], train_metrics['precision'][-1], train_metrics['recall'][-1], train_metrics['auc'][-1]\n        val_loss, val_acc, val_f1, val_prec, val_rec, val_auc = val_metrics['loss'][-1], val_metrics['acc'][-1], val_metrics['f1'][-1], val_metrics['precision'][-1], val_metrics['recall'][-1], val_metrics['auc'][-1]\n        lr = optimizer.param_groups[0]['lr']\n        print(f'Epoch {epoch + 1}/{num_epochs}, learning rate: {lr}')\n        print(f'Train Loss: {tr_loss:.4f}, Train Acc: {tr_acc:.4f}, Train f1: {tr_f1:.4f}, Train Precision: {tr_prec:.4f}, Train Recall: {tr_rec:.4f}, Train AUC: {tr_auc:.4f}')\n        print(f'Valitadion Loss: {val_loss:.4f}, Validation Acc: {val_acc:.4f}, Vall f1: {val_f1:.4f}, Val Precision: {val_prec:.4f}, Val Recall: {val_rec:.4f}, Val AUC: {val_auc:.4f}')\n        \n        if earlystoper.early_stop(val_loss):\n            break\n        \n        \n    time_elapsed = time.time() - since\n    print(f'Training complete in {time_elapsed // 60:.0f}m {time_elapsed % 60:.0f}s')\n    print(f'Best val f1: {best_f1:4f}')\n\n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    return model, train_metrics, val_metrics","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:18:05.16558Z","iopub.execute_input":"2023-04-04T10:18:05.166213Z","iopub.status.idle":"2023-04-04T10:18:05.196128Z","shell.execute_reply.started":"2023-04-04T10:18:05.166178Z","shell.execute_reply":"2023-04-04T10:18:05.195109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model, train_metrics, val_metrics = train_model(model, criterion, optimizer, scheduler, num_epochs=5)","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:03:52.158436Z","iopub.status.idle":"2023-04-04T10:03:52.159206Z","shell.execute_reply.started":"2023-04-04T10:03:52.158957Z","shell.execute_reply":"2023-04-04T10:03:52.15898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# f = plt.subplots(6, 2, figsize = (18, 12))\n# keys = ['loss', 'acc', 'f1', 'precision', 'recall', 'auc']\n# i = 0\n# for key in keys:\n#     metric = [x for x in train_metrics[key]]\n#     plt.subplot(6, 2, 2*i + 1)\n#     plt.plot(range(1, len(metric) + 1), metric)\n#     plt.xlabel(\"Epoch\")\n#     plt.ylabel(f\"{key}\")\n    \n    \n#     metric = [x for x in val_metrics[key]]\n#     plt.subplot(6, 2, 2*i + 2)\n#     plt.plot(range(1, len(metric) + 1), metric)\n#     plt.xlabel(\"Epoch\")\n#     plt.ylabel(f\"{key}\")\n#     i += 1\n    \n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:39:05.848998Z","iopub.execute_input":"2023-04-04T10:39:05.849381Z","iopub.status.idle":"2023-04-04T10:39:07.00398Z","shell.execute_reply.started":"2023-04-04T10:39:05.849351Z","shell.execute_reply":"2023-04-04T10:39:07.001316Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path_to_weights = '/kaggle/input/resnet18-rnsa-trained/checkpoint.pth'\n\ncheckpoint = torch.load(path_to_weights)\nmodel, best_weights, optimizer, threshold = checkpoint['model'], checkpoint['state_dict'], checkpoint['optimizer'], checkpoint['threshold']\nmodel.load_state_dict(best_weights)\nmodel.to(device)\n\nwith torch.no_grad():\n    n_correct = 0\n    n_samples = 0\n    false_positives = []\n    false_negatives = []\n    y_pred, y_true = [], []\n\n    for images, labels in tqdm(val_dataloader):\n            images = images.to(device)\n            labels = labels.to(device)\n            outputs = model(images)\n\n            predicted = outputs > threshold\n            n_samples += labels.size(0)\n            n_correct += (torch.squeeze(predicted) == labels).sum().item()\n            y_pred.append(np.array(torch.squeeze(predicted.cpu()), dtype = 'int32'))\n            y_true.append(np.array(torch.squeeze(labels.cpu()), dtype = 'int32'))\n            \n\n            #if predicted != labels[i]:\n            #    if predicted == 1:\n            #        false_positives.append(images)\n            #    else:\n            #        false_negatives.append(images)\n            \n    acc = 100.0 * n_correct / n_samples\n    print(f'Accuracy of the network on the {n_samples} test images: {acc} %')\n\n    y_true = np.concatenate(y_true, axis = 0)\n    y_pred = np.concatenate(y_pred, axis = 0)","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:42:23.839155Z","iopub.execute_input":"2023-04-04T10:42:23.839535Z","iopub.status.idle":"2023-04-04T10:43:59.800058Z","shell.execute_reply.started":"2023-04-04T10:42:23.839503Z","shell.execute_reply":"2023-04-04T10:43:59.799054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nimport seaborn as sns","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:44:43.408638Z","iopub.execute_input":"2023-04-04T10:44:43.409328Z","iopub.status.idle":"2023-04-04T10:44:43.536005Z","shell.execute_reply.started":"2023-04-04T10:44:43.409293Z","shell.execute_reply":"2023-04-04T10:44:43.535102Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"cm = confusion_matrix(np.squeeze(np.array(y_true, dtype = 'int32')), np.squeeze(np.array(y_pred, dtype = 'int32')))\ngroup_names = ['True Negatives','False Positives', 'False Negatives','True Positives']\ngroup_counts = [\"{0:0.0f}\".format(value) for value in\n                cm.flatten()]\ngroup_percentages = [\"{0:.2%}\".format(value) for value in\n                     cm.flatten()/np.sum(cm)]\nlabels = [f\"{v1}\\n{v2}\\n{v3}\" for v1, v2, v3 in\n          zip(group_names,group_counts,group_percentages)]\nlabels = np.asarray(labels).reshape(2,2)\nplt.figure(figsize = (12,7))\nsns.heatmap(cm, annot=labels, fmt='', cmap='Blues')\n\n","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:44:47.32961Z","iopub.execute_input":"2023-04-04T10:44:47.330124Z","iopub.status.idle":"2023-04-04T10:44:47.666907Z","shell.execute_reply.started":"2023-04-04T10:44:47.33008Z","shell.execute_reply":"2023-04-04T10:44:47.665918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Grad-CAM Explainability ","metadata":{}},{"cell_type":"code","source":"!pip install captum","metadata":{"execution":{"iopub.status.busy":"2023-04-04T10:48:37.812008Z","iopub.execute_input":"2023-04-04T10:48:37.812398Z","iopub.status.idle":"2023-04-04T10:48:49.03805Z","shell.execute_reply.started":"2023-04-04T10:48:37.812367Z","shell.execute_reply":"2023-04-04T10:48:49.036704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport cv2\nfrom captum.attr import LayerGradCam, visualization as viz","metadata":{"execution":{"iopub.status.busy":"2023-04-04T11:04:54.717951Z","iopub.execute_input":"2023-04-04T11:04:54.718572Z","iopub.status.idle":"2023-04-04T11:04:54.723643Z","shell.execute_reply.started":"2023-04-04T11:04:54.718529Z","shell.execute_reply":"2023-04-04T11:04:54.722569Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Set target for explainer and generate the Grad-CAM attributions for a single input image","metadata":{}},{"cell_type":"code","source":"target_layer = model.network.layer4\ngrad_cam = LayerGradCam(model, target_layer)\nmodel.eval()\n\n# 'images` is a batch of preprocessed input images in the correct format\n# Select the first image from the batch\nimage = images[0].unsqueeze(0).to(device)  # Add a batch dimension and move to the device\n\noutput = model(image)","metadata":{"execution":{"iopub.status.busy":"2023-04-04T11:04:54.725906Z","iopub.execute_input":"2023-04-04T11:04:54.726251Z","iopub.status.idle":"2023-04-04T11:04:54.742632Z","shell.execute_reply.started":"2023-04-04T11:04:54.726216Z","shell.execute_reply":"2023-04-04T11:04:54.741769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Calculate the attributions for the target class","metadata":{}},{"cell_type":"code","source":"target_class = torch.argmax(output, dim=1).item()\nattributions = grad_cam.attribute(image, target=target_class)","metadata":{"execution":{"iopub.status.busy":"2023-04-04T11:04:54.74423Z","iopub.execute_input":"2023-04-04T11:04:54.744869Z","iopub.status.idle":"2023-04-04T11:04:54.756307Z","shell.execute_reply.started":"2023-04-04T11:04:54.744835Z","shell.execute_reply":"2023-04-04T11:04:54.755284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Visualize the Grad-CAM heatmap:\n","metadata":{}},{"cell_type":"code","source":"# Calculate the Grad-CAM attributions\ntarget = output.argmax(dim=1).item()  # Choose the target class as the one with the highest output\nattributions = grad_cam.attribute(image, target=target)\n\n# Normalize the attributions\nattributions = np.squeeze(attributions.cpu().detach().numpy())\nattributions = np.maximum(attributions, 0)\nattributions /= np.max(attributions)\n\n# Load the original image and resize the attributions to match the image size\norig_image = images[0].permute(1, 2, 0).cpu().numpy()  # Change shape from CxHxW to HxWxC\norig_image = (orig_image - orig_image.min()) / (orig_image.max() - orig_image.min())\nheight, width, _ = orig_image.shape\nattributions = cv2.resize(attributions, (width, height))\n\n# Apply a colormap and blend the heatmap with the original image\nheatmap = cv2.applyColorMap(np.uint8(255 * attributions), cv2.COLORMAP_JET)\nheatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\nheatmap = np.float32(heatmap) / 255\noverlay = heatmap * 0.5 + orig_image * 0.5\n\n# Display the images\nplt.figure(figsize=(10, 5))\nplt.subplot(1, 2, 1)\nplt.imshow(orig_image)\nplt.title('Original Image')\nplt.axis('off')\n\nplt.subplot(1, 2, 2)\nplt.imshow(overlay)\nplt.title('Grad-CAM Heatmap Overlay')\nplt.axis('off')\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-04T11:07:13.551Z","iopub.execute_input":"2023-04-04T11:07:13.551586Z","iopub.status.idle":"2023-04-04T11:07:13.885435Z","shell.execute_reply.started":"2023-04-04T11:07:13.551542Z","shell.execute_reply":"2023-04-04T11:07:13.884372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"start_index = 0\nend_index = len(images)  # Or any other desired end_index value\n\nfor i in range(start_index, end_index):\n    image = images[i].unsqueeze(0).to(device)  # Select an image from the batch\n    output = model(image)\n    \n    target_class = torch.argmax(output, dim=1).item()\n    attributions = grad_cam.attribute(image, target=target_class)\n    \n    # Normalize the attributions\n    attributions = np.squeeze(attributions.cpu().detach().numpy())\n    attributions = np.maximum(attributions, 0)\n    attributions /= np.max(attributions)\n    \n    # Load the original image and resize the attributions to match the image size\n    orig_image = images[i].permute(1, 2, 0).cpu().numpy()  # Change shape from CxHxW to HxWxC\n    orig_image = (orig_image - orig_image.min()) / (orig_image.max() - orig_image.min())\n    height, width, _ = orig_image.shape\n    attributions = cv2.resize(attributions, (width, height))\n    \n    # Apply a colormap and blend the heatmap with the original image\n    heatmap = cv2.applyColorMap(np.uint8(255 * attributions), cv2.COLORMAP_JET)\n    heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n    heatmap = np.float32(heatmap) / 255\n    overlay = heatmap * 0.5 + orig_image * 0.5\n    \n    # Display the images\n    plt.figure(figsize=(10, 5))\n    plt.subplot(1, 2, 1)\n    plt.imshow(orig_image)\n    \n    # Display the images\n    \n    class_name = 'Cancer'\n    if y_true[i] == 0:\n        class_name = 'No Cancer'\n    \n    plt.title(f\"Original Image {i}\\nGround Truth: {class_name}\")\n    plt.axis('off')\n    \n    plt.subplot(1, 2, 2)\n    plt.imshow(overlay)\n    plt.title(f'Grad-CAM Heatmap Overlay {i}\\nPrediction: {class_name}')\n    plt.axis('off')\n    \n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-04T12:22:07.854185Z","iopub.execute_input":"2023-04-04T12:22:07.854587Z","iopub.status.idle":"2023-04-04T12:22:17.735691Z","shell.execute_reply.started":"2023-04-04T12:22:07.854549Z","shell.execute_reply":"2023-04-04T12:22:17.734648Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(val_dataset.indices)","metadata":{"execution":{"iopub.status.busy":"2023-04-04T12:04:31.593447Z","iopub.execute_input":"2023-04-04T12:04:31.593832Z","iopub.status.idle":"2023-04-04T12:04:31.600098Z","shell.execute_reply.started":"2023-04-04T12:04:31.5938Z","shell.execute_reply":"2023-04-04T12:04:31.599001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get the indices of positive and negative examples in the validation dataset\npositive_indices = []\nnegative_indices = []\nfor i in range(len(val_dataset)):\n    _, label = val_dataset[i]\n    if label == 1:\n        positive_indices.append(i)\n    else:\n        negative_indices.append(i)\n\n# Create a weighted sampler that balances the number of positive and negative examples\nnum_positives = len(positive_indices)\nnum_negatives = len(negative_indices)\nweights = [1.0 / num_positives if label == 1 else 1.0 / num_negatives for _, label in val_dataset]\nbalanced_val_sampler = WeightedRandomSampler(weights, len(weights))\n\n# Create the dataloader using the balanced sampler\nbalanced_val_dataloader = DataLoader(val_dataset, batch_size=batch_size, sampler=balanced_val_sampler)","metadata":{"execution":{"iopub.status.busy":"2023-04-04T12:29:09.450283Z","iopub.execute_input":"2023-04-04T12:29:09.450616Z","iopub.status.idle":"2023-04-04T12:30:30.995497Z","shell.execute_reply.started":"2023-04-04T12:29:09.450588Z","shell.execute_reply":"2023-04-04T12:30:30.994511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"balanced_val_iter = iter(balanced_val_dataloader)\nbalanced_images, balanced_labels = balanced_val_iter.next()\n\nstart_index = 0\nend_index = len(balanced_images)\n\nfor i in range(start_index, end_index):\n    image = balanced_images[i].unsqueeze(0).to(device)  # Select an image from the batch\n    output = model(image)\n    \n    target_class = torch.argmax(output, dim=1).item()\n    attributions = grad_cam.attribute(image, target=target_class)\n    \n    # Normalize the attributions\n    attributions = np.squeeze(attributions.cpu().detach().numpy())\n    attributions = np.maximum(attributions, 0)\n    attributions /= np.max(attributions)\n    \n    # Load the original image and resize the attributions to match the image size\n    orig_image = balanced_images[i].permute(1, 2, 0).cpu().numpy()  # Change shape from CxHxW to HxWxC\n    orig_image = (orig_image - orig_image.min()) / (orig_image.max() - orig_image.min())\n    height, width, _ = orig_image.shape\n    attributions = cv2.resize(attributions, (width, height))\n    \n    # Apply a colormap and blend the heatmap with the original image\n    heatmap = cv2.applyColorMap(np.uint8(255 * attributions), cv2.COLORMAP_JET)\n    heatmap = cv2.cvtColor(heatmap, cv2.COLOR_BGR2RGB)\n    heatmap = np.float32(heatmap) / 255\n    overlay = heatmap * 0.5 + orig_image * 0.5\n    \n    # Display the images\n    plt.figure(figsize=(10, 5))\n    plt.subplot(1, 2, 1)\n    plt.imshow(orig_image)\n    \n    ground_truth_class_name = 'Cancer'\n    if balanced_labels[i] == 0:\n        ground_truth_class_name = 'No Cancer'\n    \n    plt.title(f\"Original Image {i}\\nGround Truth: {ground_truth_class_name}\")\n    plt.axis('off')\n    \n    plt.subplot(1, 2, 2)\n    plt.imshow(overlay)\n    \n    predicted_class_name = 'Cancer'\n    if target_class == 0:\n        predicted_class_name = 'No Cancer'\n    \n    plt.title(f'Grad-CAM Heatmap Overlay {i}\\nPrediction: {predicted_class_name}')\n    plt.axis('off')\n    \n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2023-04-04T13:04:32.845418Z","iopub.execute_input":"2023-04-04T13:04:32.845895Z","iopub.status.idle":"2023-04-04T13:04:42.60641Z","shell.execute_reply.started":"2023-04-04T13:04:32.845853Z","shell.execute_reply":"2023-04-04T13:04:42.605512Z"},"trusted":true},"execution_count":null,"outputs":[]}]}