{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.7.12"},"papermill":{"default_parameters":{},"duration":6000.372815,"end_time":"2023-02-05T20:09:04.5213","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2023-02-05T18:29:04.148485","version":"2.3.4"},"widgets":{"application/vnd.jupyter.widget-state+json":{"state":{"0e27e840b2a84f1da5ba0ccbe46fb82a":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HBoxModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HBoxModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HBoxView","box_style":"","children":["IPY_MODEL_1a20ac2be7054c3fb20a857fa8436374","IPY_MODEL_bfc785145ae941c48ac713f4201f6bf0","IPY_MODEL_8e50640c78ee4f14a6a098fe63eaa1bc"],"layout":"IPY_MODEL_369e215917d745a2914431e8d6649511"}},"1a20ac2be7054c3fb20a857fa8436374":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_283aff0630f2418e804e3ea86c43bb00","placeholder":"​","style":"IPY_MODEL_b57d75261f064000a82775566616b334","value":"100%"}},"283aff0630f2418e804e3ea86c43bb00":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"369e215917d745a2914431e8d6649511":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"641a72fff7c046f2a936faaa408aed52":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"8e50640c78ee4f14a6a098fe63eaa1bc":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"HTMLModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"HTMLModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"HTMLView","description":"","description_tooltip":null,"layout":"IPY_MODEL_a2f2f1f1c8c74a919e1eff661d870092","placeholder":"​","style":"IPY_MODEL_641a72fff7c046f2a936faaa408aed52","value":" 44.7M/44.7M [00:05&lt;00:00, 18.0MB/s]"}},"9e11587abdeb4e64a1dc6671570cbfdc":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"ProgressStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"ProgressStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","bar_color":null,"description_width":""}},"a2f2f1f1c8c74a919e1eff661d870092":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}},"b57d75261f064000a82775566616b334":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"DescriptionStyleModel","state":{"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"DescriptionStyleModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"StyleView","description_width":""}},"bfc785145ae941c48ac713f4201f6bf0":{"model_module":"@jupyter-widgets/controls","model_module_version":"1.5.0","model_name":"FloatProgressModel","state":{"_dom_classes":[],"_model_module":"@jupyter-widgets/controls","_model_module_version":"1.5.0","_model_name":"FloatProgressModel","_view_count":null,"_view_module":"@jupyter-widgets/controls","_view_module_version":"1.5.0","_view_name":"ProgressView","bar_style":"success","description":"","description_tooltip":null,"layout":"IPY_MODEL_dd4ad8f7acc34ba9a8924287afff2ad3","max":46830571,"min":0,"orientation":"horizontal","style":"IPY_MODEL_9e11587abdeb4e64a1dc6671570cbfdc","value":46830571}},"dd4ad8f7acc34ba9a8924287afff2ad3":{"model_module":"@jupyter-widgets/base","model_module_version":"1.2.0","model_name":"LayoutModel","state":{"_model_module":"@jupyter-widgets/base","_model_module_version":"1.2.0","_model_name":"LayoutModel","_view_count":null,"_view_module":"@jupyter-widgets/base","_view_module_version":"1.2.0","_view_name":"LayoutView","align_content":null,"align_items":null,"align_self":null,"border":null,"bottom":null,"display":null,"flex":null,"flex_flow":null,"grid_area":null,"grid_auto_columns":null,"grid_auto_flow":null,"grid_auto_rows":null,"grid_column":null,"grid_gap":null,"grid_row":null,"grid_template_areas":null,"grid_template_columns":null,"grid_template_rows":null,"height":null,"justify_content":null,"justify_items":null,"left":null,"margin":null,"max_height":null,"max_width":null,"min_height":null,"min_width":null,"object_fit":null,"object_position":null,"order":null,"overflow":null,"overflow_x":null,"overflow_y":null,"padding":null,"right":null,"top":null,"visibility":null,"width":null}}},"version_major":2,"version_minor":0}},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"sourceType":"competition"},{"sourceId":4904750,"sourceType":"datasetVersion","datasetId":2844493}],"dockerImageVersionId":31193,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"id":"ad7430d6","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":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2023-02-05T18:29:11.855401Z","iopub.status.busy":"2023-02-05T18:29:11.854241Z","iopub.status.idle":"2023-02-05T18:29:16.321028Z","shell.execute_reply":"2023-02-05T18:29:16.319496Z"},"papermill":{"duration":4.483932,"end_time":"2023-02-05T18:29:16.32411","exception":false,"start_time":"2023-02-05T18:29:11.840178","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ccf049c2","cell_type":"code","source":"","metadata":{"papermill":{"duration":0.007473,"end_time":"2023-02-05T18:29:16.340985","exception":false,"start_time":"2023-02-05T18:29:16.333512","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"5b2892c6","cell_type":"code","source":"csvpathtrain = '/kaggle/input/rsna-breast-cancer-detection/train.csv'\n\ndftrain = pd.read_csv(csvpathtrain)\ndftrain.head()","metadata":{"execution":{"iopub.execute_input":"2023-02-05T18:29:16.357287Z","iopub.status.busy":"2023-02-05T18:29:16.356907Z","iopub.status.idle":"2023-02-05T18:29:16.480914Z","shell.execute_reply":"2023-02-05T18:29:16.479523Z"},"papermill":{"duration":0.143548,"end_time":"2023-02-05T18:29:16.491781","exception":false,"start_time":"2023-02-05T18:29:16.348233","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"e252c1ad","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.execute_input":"2023-02-05T18:29:16.525878Z","iopub.status.busy":"2023-02-05T18:29:16.525247Z","iopub.status.idle":"2023-02-05T18:29:16.954142Z","shell.execute_reply":"2023-02-05T18:29:16.953135Z"},"papermill":{"duration":0.451377,"end_time":"2023-02-05T18:29:16.959328","exception":false,"start_time":"2023-02-05T18:29:16.507951","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"1fec419e","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.execute_input":"2023-02-05T18:29:16.986094Z","iopub.status.busy":"2023-02-05T18:29:16.985287Z","iopub.status.idle":"2023-02-05T18:29:17.002815Z","shell.execute_reply":"2023-02-05T18:29:17.001829Z"},"papermill":{"duration":0.034064,"end_time":"2023-02-05T18:29:17.005676","exception":false,"start_time":"2023-02-05T18:29:16.971612","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"b4a5e145","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.execute_input":"2023-02-05T18:29:17.031014Z","iopub.status.busy":"2023-02-05T18:29:17.030414Z","iopub.status.idle":"2023-02-05T18:29:17.0435Z","shell.execute_reply":"2023-02-05T18:29:17.042596Z"},"papermill":{"duration":0.028612,"end_time":"2023-02-05T18:29:17.046135","exception":false,"start_time":"2023-02-05T18:29:17.017523","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"22d8d097","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.execute_input":"2023-02-05T18:29:17.070996Z","iopub.status.busy":"2023-02-05T18:29:17.070429Z","iopub.status.idle":"2023-02-05T18:29:17.083Z","shell.execute_reply":"2023-02-05T18:29:17.081548Z"},"papermill":{"duration":0.027806,"end_time":"2023-02-05T18:29:17.085386","exception":false,"start_time":"2023-02-05T18:29:17.05758","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"168cb4ff","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.execute_input":"2023-02-05T18:29:17.112038Z","iopub.status.busy":"2023-02-05T18:29:17.11167Z","iopub.status.idle":"2023-02-05T18:29:17.200159Z","shell.execute_reply":"2023-02-05T18:29:17.199173Z"},"papermill":{"duration":0.100762,"end_time":"2023-02-05T18:29:17.203276","exception":false,"start_time":"2023-02-05T18:29:17.102514","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"33605a08","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.execute_input":"2023-02-05T18:29:17.229279Z","iopub.status.busy":"2023-02-05T18:29:17.228731Z","iopub.status.idle":"2023-02-05T18:29:17.242761Z","shell.execute_reply":"2023-02-05T18:29:17.241832Z"},"papermill":{"duration":0.029989,"end_time":"2023-02-05T18:29:17.245575","exception":false,"start_time":"2023-02-05T18:29:17.215586","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c9873cbb","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.execute_input":"2023-02-05T18:29:17.271482Z","iopub.status.busy":"2023-02-05T18:29:17.270926Z","iopub.status.idle":"2023-02-05T18:41:16.484714Z","shell.execute_reply":"2023-02-05T18:41:16.483762Z"},"papermill":{"duration":719.229373,"end_time":"2023-02-05T18:41:16.487361","exception":false,"start_time":"2023-02-05T18:29:17.257988","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"2f987527","cell_type":"code","source":"weighted_random_sampler = WeightedRandomSampler(sample_weights, len(sample_weights))","metadata":{"execution":{"iopub.execute_input":"2023-02-05T18:41:17.188631Z","iopub.status.busy":"2023-02-05T18:41:17.188248Z","iopub.status.idle":"2023-02-05T18:41:17.192936Z","shell.execute_reply":"2023-02-05T18:41:17.192002Z"},"papermill":{"duration":0.380244,"end_time":"2023-02-05T18:41:17.195045","exception":false,"start_time":"2023-02-05T18:41:16.814801","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"8989cb98","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.execute_input":"2023-02-05T18:41:17.84599Z","iopub.status.busy":"2023-02-05T18:41:17.845625Z","iopub.status.idle":"2023-02-05T18:41:17.851383Z","shell.execute_reply":"2023-02-05T18:41:17.850338Z"},"papermill":{"duration":0.334363,"end_time":"2023-02-05T18:41:17.853969","exception":false,"start_time":"2023-02-05T18:41:17.519606","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"6cb46fff","cell_type":"code","source":"dataloaders = {'train' : train_dataloader, 'val' : val_dataloader}\ndataset_sizes = {'train': train_size, 'val' : val_size}","metadata":{"execution":{"iopub.execute_input":"2023-02-05T18:41:18.511414Z","iopub.status.busy":"2023-02-05T18:41:18.511028Z","iopub.status.idle":"2023-02-05T18:41:18.516056Z","shell.execute_reply":"2023-02-05T18:41:18.515084Z"},"papermill":{"duration":0.338604,"end_time":"2023-02-05T18:41:18.517983","exception":false,"start_time":"2023-02-05T18:41:18.179379","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"784d7a28","cell_type":"code","source":"print(len(train_dataset), len(val_dataset))\nprint(len(train_dataloader), len(val_dataloader))","metadata":{"execution":{"iopub.execute_input":"2023-02-05T18:41:19.311991Z","iopub.status.busy":"2023-02-05T18:41:19.311431Z","iopub.status.idle":"2023-02-05T18:41:19.317138Z","shell.execute_reply":"2023-02-05T18:41:19.316319Z"},"papermill":{"duration":0.472796,"end_time":"2023-02-05T18:41:19.323252","exception":false,"start_time":"2023-02-05T18:41:18.850456","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"5f85063c","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.execute_input":"2023-02-05T18:41:20.646993Z","iopub.status.busy":"2023-02-05T18:41:20.646258Z","iopub.status.idle":"2023-02-05T18:41:28.599024Z","shell.execute_reply":"2023-02-05T18:41:28.598082Z"},"papermill":{"duration":8.799322,"end_time":"2023-02-05T18:41:28.610266","exception":false,"start_time":"2023-02-05T18:41:19.810944","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"05b39191","cell_type":"code","source":"img.size()","metadata":{"execution":{"iopub.execute_input":"2023-02-05T18:41:29.342637Z","iopub.status.busy":"2023-02-05T18:41:29.342219Z","iopub.status.idle":"2023-02-05T18:41:29.349769Z","shell.execute_reply":"2023-02-05T18:41:29.348754Z"},"papermill":{"duration":0.348234,"end_time":"2023-02-05T18:41:29.351946","exception":false,"start_time":"2023-02-05T18:41:29.003712","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"d4c151c5","cell_type":"code","source":"device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f'Current device is {device}')","metadata":{"execution":{"iopub.execute_input":"2023-02-05T18:41:30.02612Z","iopub.status.busy":"2023-02-05T18:41:30.025741Z","iopub.status.idle":"2023-02-05T18:41:30.031233Z","shell.execute_reply":"2023-02-05T18:41:30.030254Z"},"papermill":{"duration":0.34885,"end_time":"2023-02-05T18:41:30.033446","exception":false,"start_time":"2023-02-05T18:41:29.684596","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"dfef93e6","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.execute_input":"2023-02-05T18:41:30.774911Z","iopub.status.busy":"2023-02-05T18:41:30.774327Z","iopub.status.idle":"2023-02-05T18:41:30.780912Z","shell.execute_reply":"2023-02-05T18:41:30.780022Z"},"papermill":{"duration":0.40685,"end_time":"2023-02-05T18:41:30.782815","exception":false,"start_time":"2023-02-05T18:41:30.375965","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"948c4618","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.execute_input":"2023-02-05T18:41:31.691789Z","iopub.status.busy":"2023-02-05T18:41:31.69141Z","iopub.status.idle":"2023-02-05T18:41:31.697727Z","shell.execute_reply":"2023-02-05T18:41:31.696736Z"},"papermill":{"duration":0.379861,"end_time":"2023-02-05T18:41:31.699706","exception":false,"start_time":"2023-02-05T18:41:31.319845","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"30280b5c","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.execute_input":"2023-02-05T18:41:32.442605Z","iopub.status.busy":"2023-02-05T18:41:32.442233Z","iopub.status.idle":"2023-02-05T18:41:32.448493Z","shell.execute_reply":"2023-02-05T18:41:32.447518Z"},"papermill":{"duration":0.34455,"end_time":"2023-02-05T18:41:32.450566","exception":false,"start_time":"2023-02-05T18:41:32.106016","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"0392f1c1","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.execute_input":"2023-02-05T18:41:33.125558Z","iopub.status.busy":"2023-02-05T18:41:33.125162Z","iopub.status.idle":"2023-02-05T18:41:39.360669Z","shell.execute_reply":"2023-02-05T18:41:39.359449Z"},"papermill":{"duration":6.578766,"end_time":"2023-02-05T18:41:39.36329","exception":false,"start_time":"2023-02-05T18:41:32.784524","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ae409521","cell_type":"code","source":"model","metadata":{"execution":{"iopub.execute_input":"2023-02-05T18:41:40.09378Z","iopub.status.busy":"2023-02-05T18:41:40.092756Z","iopub.status.idle":"2023-02-05T18:41:40.100727Z","shell.execute_reply":"2023-02-05T18:41:40.099783Z"},"papermill":{"duration":0.403036,"end_time":"2023-02-05T18:41:40.102635","exception":false,"start_time":"2023-02-05T18:41:39.699599","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"bc67df49","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.execute_input":"2023-02-05T18:41:40.778199Z","iopub.status.busy":"2023-02-05T18:41:40.777837Z","iopub.status.idle":"2023-02-05T18:41:40.795072Z","shell.execute_reply":"2023-02-05T18:41:40.7942Z"},"papermill":{"duration":0.357425,"end_time":"2023-02-05T18:41:40.797203","exception":false,"start_time":"2023-02-05T18:41:40.439778","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"5d9019bf","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.execute_input":"2023-02-05T18:41:41.531076Z","iopub.status.busy":"2023-02-05T18:41:41.530707Z","iopub.status.idle":"2023-02-05T18:41:41.536402Z","shell.execute_reply":"2023-02-05T18:41:41.535446Z"},"papermill":{"duration":0.345027,"end_time":"2023-02-05T18:41:41.538622","exception":false,"start_time":"2023-02-05T18:41:41.193595","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"7333f96f","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.execute_input":"2023-02-05T18:41:42.454708Z","iopub.status.busy":"2023-02-05T18:41:42.454323Z","iopub.status.idle":"2023-02-05T18:41:42.460794Z","shell.execute_reply":"2023-02-05T18:41:42.459765Z"},"papermill":{"duration":0.537288,"end_time":"2023-02-05T18:41:42.462749","exception":false,"start_time":"2023-02-05T18:41:41.925461","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ef2434f7","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.execute_input":"2023-02-05T18:41:43.190981Z","iopub.status.busy":"2023-02-05T18:41:43.190601Z","iopub.status.idle":"2023-02-05T18:42:49.539327Z","shell.execute_reply":"2023-02-05T18:42:49.538244Z"},"papermill":{"duration":67.049078,"end_time":"2023-02-05T18:42:49.848424","exception":false,"start_time":"2023-02-05T18:41:42.799346","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"9beb8bdf","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.execute_input":"2023-02-05T18:42:50.574716Z","iopub.status.busy":"2023-02-05T18:42:50.574321Z","iopub.status.idle":"2023-02-05T18:42:51.222888Z","shell.execute_reply":"2023-02-05T18:42:51.221734Z"},"papermill":{"duration":0.989301,"end_time":"2023-02-05T18:42:51.225133","exception":false,"start_time":"2023-02-05T18:42:50.235832","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"4bd9995c","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.execute_input":"2023-02-05T18:42:51.914953Z","iopub.status.busy":"2023-02-05T18:42:51.914352Z","iopub.status.idle":"2023-02-05T18:42:51.922041Z","shell.execute_reply":"2023-02-05T18:42:51.921156Z"},"papermill":{"duration":0.359454,"end_time":"2023-02-05T18:42:51.924171","exception":false,"start_time":"2023-02-05T18:42:51.564717","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"77f54873","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.execute_input":"2023-02-05T18:42:52.662698Z","iopub.status.busy":"2023-02-05T18:42:52.662305Z","iopub.status.idle":"2023-02-05T18:42:52.687211Z","shell.execute_reply":"2023-02-05T18:42:52.686223Z"},"papermill":{"duration":0.368532,"end_time":"2023-02-05T18:42:52.689664","exception":false,"start_time":"2023-02-05T18:42:52.321132","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"876feaa4","cell_type":"code","source":"model, train_metrics, val_metrics = train_model(model, criterion, optimizer, scheduler, num_epochs=5)","metadata":{"execution":{"iopub.execute_input":"2023-02-05T18:42:53.380739Z","iopub.status.busy":"2023-02-05T18:42:53.378865Z","iopub.status.idle":"2023-02-05T20:07:34.605066Z","shell.execute_reply":"2023-02-05T20:07:34.603836Z"},"papermill":{"duration":5081.583665,"end_time":"2023-02-05T20:07:34.608105","exception":false,"start_time":"2023-02-05T18:42:53.02444","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a2dabfd6","cell_type":"code","source":"f = plt.subplots(6, 2, figsize = (18, 12))\nkeys = ['loss', 'acc', 'f1', 'precision', 'recall', 'auc']\ni = 0\nfor 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    \nplt.show()","metadata":{"execution":{"iopub.execute_input":"2023-02-05T20:07:36.239253Z","iopub.status.busy":"2023-02-05T20:07:36.238861Z","iopub.status.idle":"2023-02-05T20:07:37.444278Z","shell.execute_reply":"2023-02-05T20:07:37.443407Z"},"papermill":{"duration":1.978471,"end_time":"2023-02-05T20:07:37.446391","exception":false,"start_time":"2023-02-05T20:07:35.46792","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"a81419e9","cell_type":"code","source":"path_to_weights = '/kaggle/working/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.execute_input":"2023-02-05T20:07:39.271099Z","iopub.status.busy":"2023-02-05T20:07:39.270722Z","iopub.status.idle":"2023-02-05T20:08:55.621906Z","shell.execute_reply":"2023-02-05T20:08:55.62073Z"},"papermill":{"duration":77.192849,"end_time":"2023-02-05T20:08:55.624543","exception":false,"start_time":"2023-02-05T20:07:38.431694","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"c8957863","cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nimport seaborn as sns","metadata":{"execution":{"iopub.execute_input":"2023-02-05T20:08:57.723368Z","iopub.status.busy":"2023-02-05T20:08:57.722305Z","iopub.status.idle":"2023-02-05T20:08:57.950309Z","shell.execute_reply":"2023-02-05T20:08:57.949307Z"},"papermill":{"duration":1.102774,"end_time":"2023-02-05T20:08:57.952973","exception":false,"start_time":"2023-02-05T20:08:56.850199","status":"completed"},"tags":[]},"outputs":[],"execution_count":null},{"id":"ba872e24","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.execute_input":"2023-02-05T20:08:59.592569Z","iopub.status.busy":"2023-02-05T20:08:59.592167Z","iopub.status.idle":"2023-02-05T20:08:59.851484Z","shell.execute_reply":"2023-02-05T20:08:59.850523Z"},"papermill":{"duration":1.049896,"end_time":"2023-02-05T20:08:59.853626","exception":false,"start_time":"2023-02-05T20:08:58.80373","status":"completed"},"tags":[]},"outputs":[],"execution_count":null}]}