{"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":"markdown","source":"# RSNA Abdominal Trauma Detection - Naive Approach ResNet\n\nMW Maddock\n2023-08-08\n\n## 1. Overview\n\n### Goal of the Competition\n\nTraumatic injury is the most common cause of death in the first four decades of life and a major public health problem around the world. There are estimated to be more than 5 million annual deaths worldwide from traumatic injury. Prompt and accurate diagnosis of traumatic injuries is crucial for initiating appropriate and timely interventions, which can significantly improve patient outcomes and survival rates. Computed tomography (CT) has become an indispensable tool in evaluating patients with suspected abdominal injuries due to its ability to provide detailed cross-sectional images of the abdomen.\n\nInterpreting CT scans for abdominal trauma, however, can be a complex and time-consuming task, especially when multiple injuries or areas of subtle active bleeding are present. This challenge seeks to harness the power of artificial intelligence and machine learning to assist medical professionals in rapidly and precisely detecting injuries and grading their severity. The development of advanced algorithms for this purpose has the potential to improve trauma care and patient outcomes worldwide.\n\n### Context\nBlunt force abdominal trauma is among the most common types of traumatic injury, with the most frequent cause being motor vehicle accidents. Abdominal trauma may result in damage and internal bleeding of the internal organs, including the liver, spleen, kidneys, and bowel. Detection and classification of injuries are key to effective treatment and favorable outcomes. A large proportion of patients with abdominal trauma require urgent surgery. Abdominal trauma often cannot be diagnosed clinically by physical exam, patient symptoms, or laboratory tests.\n\nPrompt diagnosis of abdominal trauma using medical imaging is thus critical to patient care. AI tools that assist and expedite diagnosis of abdominal trauma have the potential to substantially improve patient care and health outcomes in the emergency setting.\n\nThe RSNA Abdominal Trauma Detection AI Challenge, organized by the RSNA in collaboration with the American Society of Emergency Radiology (ASER) and the Society for Abdominal Radiology (SAR), gives researchers the task of building models that detect severe injury to the internal abdominal organs, including the liver, kidneys, spleen, and bowel, as well as any active internal bleeding.\n\n## 2. Dataset Description\nThe goal of this competition is to identify several potential injuries in CT scans of trauma patients. Any of these injuries can be fatal on a short time frame if untreated so there is great value in rapid diagnosis.\n\nThis competition uses a hidden test. When your submitted notebook is scored, the actual test data (including a full length sample submission) will be made available to your notebook.\n\n### Files\n**train.csv** Target labels for the train set. Note that patients labeled healthy may still have other medical issues, such as cancer or broken bones, that don't happen to be covered by the competition labels.\n\n* `patient_id` - A unique ID code for each patient.\n* `[bowel/extravasation]_[healthy/injury]` - The two injury types with binary targets.\n* `[kidney/liver/spleen]_[healthy/low/high]` - The three injury types with three target levels.\n* `any_injury` - Whether the patient had any injury at all.\n* `[train/test]_images/[patient_id]/[series_id]/[image_instance_number].dcm` The CT scan data, in DICOM format. Scans from dozens of different CT machines have been reprocessed to use the run length encoded lossless compression format but retain other differences such as the number of bits per pixel, pixel range, and pixel representation. Expect to see roughly 1,100 patients in the test set.\n\n**[train/test]_series_meta.csv** Each patient may have been scanned once or twice. Each scan contains a series of images.\n\n* `patient_id` - A unique ID code for each patient.\n* `series_id` - A unique ID code for each scan.\n* `aortic_hu` - The volume of the aorta in hounsfield units. This acts as a reliable proxy for when the scan was. For a multiphasic CT scan, the higher value indicates the late arterial phase.\n* `incomplete_organ` - True if one or more organs wasn't fully covered by the scan. This label is only provided for the train set.\n\n**sample_submission.csv** A valid sample submission. Only the first few rows are available for download.\n\n**image_level_labels.csv** Train only. Identifies specific images that contain either bowel or extravasation injuries.\n\n* `patient_id` - A unique ID code for each patient.\n* `series_id` - A unique ID code for each scan.\n* `instance_number` - The image number within the scan. The lowest instance number for many series is above zero as the original scans were cropped to the abdomen.\n* `injury_name` - The type of injury visible in the frame.\n\n**segmentations/** Model generated pixel-level annotations of the relevant organs and some major bones for a subset of the scans in the training set. This data is provided in the nifti file format. The filenames are series IDs. You can find a description of the source model (total segmentator) here and the data used to train that model here.\n\nNote that the NIFTI files and DICOM files are not in the same orientation. Use the NIFTI header information along with DICOM metadata to determine the appropriate orientation.\n\n**[train/test]_dicom_tags.parquet** DICOM tags from every image, extracted with Pydicom. Provided for convenience.\n\n\n## 3. Import Modules","metadata":{}},{"cell_type":"code","source":"# Input data files are available in the read-only \"../data/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# Import libraries\nimport json\nimport os\nimport pickle\nimport random\nimport time\nimport glob \nimport pydicom as dicom \n\n# Ignore warnings\nimport warnings\n\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport torch\n\n# PyTorch model\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nimport torchvision.transforms as transforms\nfrom tqdm import tqdm\nfrom skimage import io, transform\nfrom sklearn.metrics import classification_report, confusion_matrix, jaccard_score\nfrom sklearn.model_selection import train_test_split\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.utils.data.sampler import SubsetRandomSampler\n\n# PyTorch dataset\nfrom torchvision import datasets, models, transforms, utils\nfrom torchvision.utils import make_grid\n\nwarnings.filterwarnings(\"ignore\")\n\nplt.ion()  # interactive mode\n\nfrom __future__ import print_function, division\n\n%config InlineBackend.figure_format = 'retina'\n%matplotlib inline\n\ntorch.__version__","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:47:52.848579Z","iopub.execute_input":"2023-08-16T12:47:52.84957Z","iopub.status.idle":"2023-08-16T12:48:00.227496Z","shell.execute_reply.started":"2023-08-16T12:47:52.84953Z","shell.execute_reply":"2023-08-16T12:48:00.226338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# check if CUDA is available\ntrain_on_gpu = torch.cuda.is_available()\n\nif not train_on_gpu:\n    print('CUDA is not available.  Training on CPU ...')\nelse:\n    print('CUDA is available!  Training on GPU ...')","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:48:00.229949Z","iopub.execute_input":"2023-08-16T12:48:00.230957Z","iopub.status.idle":"2023-08-16T12:48:00.304736Z","shell.execute_reply.started":"2023-08-16T12:48:00.230921Z","shell.execute_reply":"2023-08-16T12:48:00.303581Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!nvidia-smi","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:48:00.306562Z","iopub.execute_input":"2023-08-16T12:48:00.307368Z","iopub.status.idle":"2023-08-16T12:48:01.431187Z","shell.execute_reply.started":"2023-08-16T12:48:00.307312Z","shell.execute_reply":"2023-08-16T12:48:01.430023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4. Precrocessing\n\n\nIn this section we do prepare the traing labels and save into the working directory","metadata":{}},{"cell_type":"code","source":"df_train = pd.read_csv('/kaggle/input/rsna-2023-abdominal-trauma-detection/train.csv')\n\ndf_train.head()","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:48:01.435184Z","iopub.execute_input":"2023-08-16T12:48:01.435538Z","iopub.status.idle":"2023-08-16T12:48:01.479787Z","shell.execute_reply.started":"2023-08-16T12:48:01.435507Z","shell.execute_reply":"2023-08-16T12:48:01.478662Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Convert patient_id column from DataFrame to a numpy array and cast it to string data type\npatient_id = df_train.patient_id.to_numpy().astype(str)\n\n# Initialize empty lists for storing image data (X) and corresponding labels (y)\nX, y, paths = [], [], []\n\n# Loop over the first 3 patient IDs\nfor x, p_id in enumerate(tqdm(patient_id)):\n    # Define the directory path for the current patient's images\n    dir = '/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images/' + p_id + '/'\n    \n    # Extract the features (labels) for the current patient using their index\n    labels = df_train.iloc[x].to_list()\n    \n    # Loop through each file in the patient's directory\n    for file in glob.glob(dir + '*'):\n        # Loop through each image file in the current directory\n        \n        \n        for image_path in glob.glob(file + '/*'):\n            # Read the DICOM image and extract the pixel array\n            # X.append(dicom.dcmread(image_path).pixel_array)\n            \n            # Append the features (labels) for this image to the y list\n            y.append(labels)\n            paths.append(image_path)","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:48:01.48232Z","iopub.execute_input":"2023-08-16T12:48:01.483483Z","iopub.status.idle":"2023-08-16T12:53:40.321653Z","shell.execute_reply.started":"2023-08-16T12:48:01.483448Z","shell.execute_reply":"2023-08-16T12:53:40.32065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_labels = pd.DataFrame(y, columns = df_train.columns)\ndf_labels[\"image_path\"] = paths\n\ndf_labels","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:53:40.323257Z","iopub.execute_input":"2023-08-16T12:53:40.323639Z","iopub.status.idle":"2023-08-16T12:53:50.828452Z","shell.execute_reply.started":"2023-08-16T12:53:40.323604Z","shell.execute_reply":"2023-08-16T12:53:50.827362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_labels.shape","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:53:50.829883Z","iopub.execute_input":"2023-08-16T12:53:50.830789Z","iopub.status.idle":"2023-08-16T12:53:50.838463Z","shell.execute_reply.started":"2023-08-16T12:53:50.83075Z","shell.execute_reply":"2023-08-16T12:53:50.837209Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_labels.info()","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:53:50.840248Z","iopub.execute_input":"2023-08-16T12:53:50.84073Z","iopub.status.idle":"2023-08-16T12:53:51.337149Z","shell.execute_reply.started":"2023-08-16T12:53:50.840695Z","shell.execute_reply":"2023-08-16T12:53:51.33613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_labels.to_csv(\"/kaggle/working/df_labels.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:53:51.338607Z","iopub.execute_input":"2023-08-16T12:53:51.339555Z","iopub.status.idle":"2023-08-16T12:54:05.259062Z","shell.execute_reply.started":"2023-08-16T12:53:51.339522Z","shell.execute_reply":"2023-08-16T12:54:05.258035Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3. Create PyTorch `DataLoader`","metadata":{}},{"cell_type":"code","source":"labels = ['bowel_healthy', \n          'bowel_injury', \n          'extravasation_healthy',\n          'extravasation_injury', \n          'kidney_healthy', \n          'kidney_low', \n          'kidney_high',\n          'liver_healthy', \n          'liver_low', \n          'liver_high', \n          'spleen_healthy',\n          'spleen_low', \n          'spleen_high', \n          'any_injury']","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:05.309557Z","iopub.execute_input":"2023-08-16T12:54:05.309948Z","iopub.status.idle":"2023-08-16T12:54:05.31567Z","shell.execute_reply.started":"2023-08-16T12:54:05.309917Z","shell.execute_reply":"2023-08-16T12:54:05.31412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNATraumaDetection (Dataset):\n    \"\"\"Ocular Disease Recognition.\"\"\"\n\n    def __init__(self, csv_file, transform=None):\n        \"\"\"\n        Arguments:\n            csv_file (string): Path to the csv file with labels and image paths.\n            transform (callable, optional): Optional transform to be applied\n                on a sample.\n        \"\"\"\n        self.labels_frame = pd.read_csv(csv_file)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.labels_frame)\n\n    def __getitem__(self, idx):\n        if torch.is_tensor(idx):\n            idx = idx.tolist()\n\n        img_name = self.labels_frame.iloc[idx, :][\"image_path\"]\n        image = dicom.dcmread(image_path).pixel_array\n        target = self.labels_frame.iloc[idx, :][labels].to_list()\n        target = np.array(target)\n        sample = {'image': image, 'labels': target}\n\n        if self.transform:\n            sample = self.transform(sample)\n\n        return sample\n    \nclass Rescale(object):\n    \"\"\"Rescale the image in a sample to a given size.\n\n    Args:\n        output_size (tuple or int): Desired output size. If tuple, output is\n            matched to output_size. If int, smaller of image edges is matched\n            to output_size keeping aspect ratio the same.\n    \"\"\"\n\n    def __init__(self, output_size):\n        assert isinstance(output_size, (int, tuple))\n        self.output_size = output_size\n\n    def __call__(self, sample):\n        image, label = sample['image'], sample['labels']\n\n        h, w = image.shape[:2]\n        if isinstance(self.output_size, int):\n            if h > w:\n                new_h, new_w = self.output_size * h / w, self.output_size\n            else:\n                new_h, new_w = self.output_size, self.output_size * w / h\n        else:\n            new_h, new_w = self.output_size\n\n        new_h, new_w = int(new_h), int(new_w)\n\n        img = transform.resize(image, (new_h, new_w))\n\n        return {'image': img, 'labels': label}\n\n\nclass RandomCrop(object):\n    \"\"\"Crop randomly the image in a sample.\n\n    Args:\n        output_size (tuple or int): Desired output size. If int, square crop\n            is made.\n    \"\"\"\n\n    def __init__(self, output_size):\n        assert isinstance(output_size, (int, tuple))\n        if isinstance(output_size, int):\n            self.output_size = (output_size, output_size)\n        else:\n            assert len(output_size) == 2\n            self.output_size = output_size\n\n    def __call__(self, sample):\n        image, label = sample['image'], sample['labels']\n\n        h, w = image.shape[:2]\n        new_h, new_w = self.output_size\n\n        top = np.random.randint(0, h - new_h)\n        left = np.random.randint(0, w - new_w)\n\n        image = image[top: top + new_h,\n                      left: left + new_w]\n\n\n        return {'image': image, 'labels': label}\n\n\nclass ToTensor(object):\n    \"\"\"Convert ndarrays in sample to Tensors.\"\"\"\n\n    def __call__(self, sample):\n        image, label = sample['image'], sample['labels']\n\n        # swap color axis because\n        # numpy image: H x W x C\n        # torch image: C x H x W\n        return {'image': torch.from_numpy(image),\n                'labels': torch.from_numpy(label)}","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:05.317597Z","iopub.execute_input":"2023-08-16T12:54:05.318337Z","iopub.status.idle":"2023-08-16T12:54:05.338015Z","shell.execute_reply.started":"2023-08-16T12:54:05.318291Z","shell.execute_reply":"2023-08-16T12:54:05.336873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5. Load Image Data\n\n...\n\n### 5.1 Split into Train, Validation and Test Sets","metadata":{}},{"cell_type":"code","source":"# number of subprocesses to use for data loading\nnum_workers = 2\n# how many samples per batch to load\nbatch_size = 32\n# percentage of training set to use as validation\nvalid_size = 0.2\ntest_size = 0.2","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:05.339245Z","iopub.execute_input":"2023-08-16T12:54:05.339651Z","iopub.status.idle":"2023-08-16T12:54:05.354387Z","shell.execute_reply.started":"2023-08-16T12:54:05.33962Z","shell.execute_reply":"2023-08-16T12:54:05.353516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# choose the training and test datasets\nlabels_dir = '/kaggle/working/df_labels.csv'\n\nfull_data  = RSNATraumaDetection(csv_file=labels_dir, \n                                      transform=transforms.Compose([Rescale(512),\n                                                                    transforms.RandomHorizontalFlip(),\n                                                                    ToTensor()])\n                                     )","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:05.355497Z","iopub.execute_input":"2023-08-16T12:54:05.355988Z","iopub.status.idle":"2023-08-16T12:54:08.411507Z","shell.execute_reply.started":"2023-08-16T12:54:05.355957Z","shell.execute_reply":"2023-08-16T12:54:08.410471Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# obtain training indices that will be used for validation\nnum_train = len(full_data)\nindices = list(range(num_train))\n\nnp.random.shuffle(indices)\n\nval_split = int(np.floor(valid_size * num_train))\ntest_split = int(np.floor(valid_size * num_train))\n\n\ntest_idx, valid_idx, train_idx = indices[:test_split], indices[test_split: test_split + val_split], indices[test_split + val_split:]","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:08.412964Z","iopub.execute_input":"2023-08-16T12:54:08.413578Z","iopub.status.idle":"2023-08-16T12:54:08.662921Z","shell.execute_reply.started":"2023-08-16T12:54:08.413545Z","shell.execute_reply":"2023-08-16T12:54:08.661621Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_train","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:08.664604Z","iopub.execute_input":"2023-08-16T12:54:08.665007Z","iopub.status.idle":"2023-08-16T12:54:08.671856Z","shell.execute_reply.started":"2023-08-16T12:54:08.66497Z","shell.execute_reply":"2023-08-16T12:54:08.67076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# define samplers for obtaining training and validation batches\ntrain_sampler = SubsetRandomSampler(train_idx)\nvalid_sampler = SubsetRandomSampler(valid_idx)\ntest_sampler  = SubsetRandomSampler(test_idx)","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:08.673604Z","iopub.execute_input":"2023-08-16T12:54:08.674027Z","iopub.status.idle":"2023-08-16T12:54:08.683695Z","shell.execute_reply.started":"2023-08-16T12:54:08.673992Z","shell.execute_reply":"2023-08-16T12:54:08.681469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# prepare data loaders (combine dataset and sampler)\ntrain_loader = torch.utils.data.DataLoader(full_data, batch_size=batch_size,\n    sampler=train_sampler, num_workers=num_workers)\n\nvalid_loader = torch.utils.data.DataLoader(full_data, batch_size=batch_size, \n    sampler=valid_sampler, num_workers=num_workers)\n\ntest_loader = torch.utils.data.DataLoader(full_data, batch_size=batch_size, \n    sampler=test_sampler, num_workers=num_workers)","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:08.685223Z","iopub.execute_input":"2023-08-16T12:54:08.685653Z","iopub.status.idle":"2023-08-16T12:54:08.697743Z","shell.execute_reply.started":"2023-08-16T12:54:08.685621Z","shell.execute_reply":"2023-08-16T12:54:08.69681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 5.2  Visualize a Batch of Training Data","metadata":{}},{"cell_type":"code","source":"# helper function to un-normalize and display an image\ndef imshow(img):\n    img = img / 2 + 0.5  # unnormalize\n    plt.imshow(img)  # convert from Tensor image","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:08.701027Z","iopub.execute_input":"2023-08-16T12:54:08.701392Z","iopub.status.idle":"2023-08-16T12:54:08.707967Z","shell.execute_reply.started":"2023-08-16T12:54:08.701362Z","shell.execute_reply":"2023-08-16T12:54:08.706864Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# obtain one batch of training images\ndataiter = iter(train_loader)\nsample = next(dataiter)\n\nsample['image'].shape # (number of examples: 20, number of channels: 3, pixel sizes: 256x256)","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:08.709413Z","iopub.execute_input":"2023-08-16T12:54:08.709838Z","iopub.status.idle":"2023-08-16T12:54:13.446533Z","shell.execute_reply.started":"2023-08-16T12:54:08.709797Z","shell.execute_reply":"2023-08-16T12:54:13.445383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# plot the images in the batch, along with the corresponding labels\nfig = plt.figure(figsize=(25, 4))\n# display 20 images\nfor idx in np.arange(batch_size):\n    ax = fig.add_subplot(2, 21, idx+1, xticks=[], yticks=[])\n    imshow(sample['image'][idx])","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:13.451435Z","iopub.execute_input":"2023-08-16T12:54:13.453931Z","iopub.status.idle":"2023-08-16T12:54:17.682863Z","shell.execute_reply.started":"2023-08-16T12:54:13.453881Z","shell.execute_reply":"2023-08-16T12:54:17.681919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 6. Model Training","metadata":{}},{"cell_type":"code","source":"device_name = \"cuda\" if torch.cuda.is_available() else \"cpu\"\ndevice = torch.device(device_name)\n\nprint(device_name)","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:17.683965Z","iopub.execute_input":"2023-08-16T12:54:17.685621Z","iopub.status.idle":"2023-08-16T12:54:17.691393Z","shell.execute_reply.started":"2023-08-16T12:54:17.685586Z","shell.execute_reply":"2023-08-16T12:54:17.69054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_checkpoint(state, is_best, filename='/kaggle/working/rsn_trauma_detection_resnet.pth.tar'):\n    torch.save(state, filename)","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:17.69304Z","iopub.execute_input":"2023-08-16T12:54:17.693689Z","iopub.status.idle":"2023-08-16T12:54:17.703952Z","shell.execute_reply.started":"2023-08-16T12:54:17.693655Z","shell.execute_reply":"2023-08-16T12:54:17.702845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 6.1 Define Model Achitecture","metadata":{}},{"cell_type":"code","source":"# instantiate transfer learning model\nresnet_model = models.resnet50(pretrained=True)\n\n# set all parameters as trainable\nfor param in resnet_model.parameters():\n    param.requires_grad = True\n\nresnet_model.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3,\n                               bias=False)\n# get input of fc layer\nn_inputs = resnet_model.fc.in_features\n\n# redefine fc layer / top layer/ head for our classification problem\nresnet_model.fc = nn.Sequential(nn.Linear(n_inputs, 2048),\n                                nn.ReLU(),\n                                nn.Dropout(p=0.4),\n                                nn.Linear(2048, 2048),\n                                nn.ReLU(),\n                                nn.Dropout(p=0.4),\n                                nn.Linear(2048, len(labels)),\n                                nn.LogSigmoid())\n\n# set all parameters of the model as trainable\nfor name, child in resnet_model.named_children():\n    for name2, params in child.named_parameters():\n        params.requires_grad = True\n\n\n# Disbribute the model to all GPU's\nresnet_model = nn.DataParallel(resnet_model)\n\n# set model to run on GPU or CPU absed on availibility\nresnet_model.to(device)\n\n# print the trasnfer learning NN model's architecture\nresnet_model","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:17.706526Z","iopub.execute_input":"2023-08-16T12:54:17.707124Z","iopub.status.idle":"2023-08-16T12:54:22.235224Z","shell.execute_reply.started":"2023-08-16T12:54:17.707093Z","shell.execute_reply":"2023-08-16T12:54:22.234262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 6.2 Define Criterion & Optimizer","metadata":{}},{"cell_type":"code","source":"# loss function\n# if GPU is available set loss function to use GPU\ncriterion = nn.CrossEntropyLoss().to(device)\n\n# optimizer\noptimizer = torch.optim.SGD(resnet_model.parameters(), momentum=0.9, lr=3e-4)\n\n\n# empty lists to store losses and accuracies\ntrain_losses = []\ntest_losses = []\ntrain_correct = []\ntest_correct = []","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:22.236805Z","iopub.execute_input":"2023-08-16T12:54:22.23717Z","iopub.status.idle":"2023-08-16T12:54:22.245085Z","shell.execute_reply.started":"2023-08-16T12:54:22.237137Z","shell.execute_reply":"2023-08-16T12:54:22.244107Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 6.3 Run the Traing Loop","metadata":{}},{"cell_type":"code","source":"# number of training iterations\nepochs = 4","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:22.246696Z","iopub.execute_input":"2023-08-16T12:54:22.248549Z","iopub.status.idle":"2023-08-16T12:54:22.256234Z","shell.execute_reply.started":"2023-08-16T12:54:22.248516Z","shell.execute_reply":"2023-08-16T12:54:22.255043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:22.257832Z","iopub.execute_input":"2023-08-16T12:54:22.258178Z","iopub.status.idle":"2023-08-16T12:54:22.268287Z","shell.execute_reply.started":"2023-08-16T12:54:22.258146Z","shell.execute_reply":"2023-08-16T12:54:22.26739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# set training start time\nstart_time = time.time()\n\n# set best_prec loss value as 2 for checkpoint threshold\nbest_prec1 = 2\nis_best = False\n\n# empty batch variables\nb = None\ntrain_b = None\ntest_b = None\n\nscaler = GradScaler()\n\nfor i in range(epochs):\n    # empty training correct and test correct counter as 0 during every iteration\n    trn_corr = 0\n    tst_corr = 0\n    trn_loss = 0\n    tst_loss = 0\n    \n    # set epoch's starting time\n    e_start = time.time()\n    \n\n    # train in batches\n    for b, sample in enumerate(tqdm(train_loader)):\n        # set label as cuda if device is cuda\n        X, y = sample['image'].to(device, dtype=torch.float), sample['labels'].to(device, dtype=torch.float)\n        \n        # forward pass image sample\n        y_pred = resnet_model(X.view(-1, 1, 512, 512))\n\n        # calculate loss\n        loss = criterion(y_pred.float(), y.float())\n\n        trn_loss += loss.item()\n        # get argmax of predicted tensor, which is our label\n        predicted = torch.argmax(y_pred, dim=1).data\n        y = torch.argmax(y, dim=1).data\n\n        # if predicted label is correct as true label, calculate the sum for samples\n\n        batch_corr = (predicted == y).sum()\n        # increment train correct with correcly predicted labels per batch\n        trn_corr += batch_corr.item()\n        \n        # set optimizer gradients to zero\n        optimizer.zero_grad()\n        # Backpropagate with autocasting\n        # back propagate with loss\n        scaler.scale(loss).backward()\n        # perform optimizer step\n        scaler.step(optimizer)\n        scaler.update()\n      \n    # set epoch's end time\n    e_end = time.time()\n    \n    # print training metrics\n    print(f'Epoch {(i+1)} Batch {(b+1)}\\nAccuracy: {trn_corr*100/(b*batch_size):2.2f} %  Loss: {trn_loss/len(train_loader):2.4f}  Duration: {((e_end-e_start)/60):.2f} minutes') \n    \n    # some metrics storage for visualization\n    train_b = b\n    train_losses.append(trn_loss)\n    train_correct.append(trn_corr)\n\n    X, y = None, None\n\n    # validate using validation generator\n    # do not perform any gradient updates while validation\n    with torch.no_grad():\n        for b, sample in enumerate(valid_loader):\n            # set label as cuda if device is cuda\n            X, y = sample['image'].to(device, dtype=torch.float), sample['labels'].to(device, dtype=torch.float)\n\n            # forward pass image\n            y_val = resnet_model(X.view(-1, 1, 512, 512))\n\n            # get argmax of predicted tensor, which is our label\n            predicted = torch.argmax(y_val, dim=1).data\n            y = torch.argmax(y, dim=1).data\n\n            # increment test correct with correcly predicted labels per batch\n            tst_corr += (predicted == y).sum().item()\n\n            # get loss of validation set\n            loss = criterion(y_val.float(), y.long())\n            tst_loss += loss.item()\n            \n            \n    # print validation metrics\n    print(f'Validation Accuracy {tst_corr*100/(b*batch_size):2.2f}% Validation Loss: {tst_loss/len(valid_loader):2.4f}\\n')\n\n    # if current validation loss is less than previous iteration's validatin loss create and save a checkpoint\n    is_best = loss < best_prec1\n    best_prec1 = min(loss, best_prec1)\n    \n    if is_best:\n        save_checkpoint({\n                'epoch': i + 1,\n                'state_dict': resnet_model.state_dict(),\n                'best_prec1': best_prec1,\n            }, is_best)\n        \n        is_best = False\n\n    # some metrics storage for visualization\n    test_b  = b\n    test_losses.append(tst_loss)\n    test_correct.append(tst_corr)\n\n# set total training's end time\nend_time = time.time() - start_time    \n\n# print training summary\nprint(\"\\nTraining Duration {:.2f} minutes\".format(end_time/60))\nprint(\"GPU memory used : {} kb\".format(torch.cuda.memory_allocated()))\nprint(\"GPU memory cached : {} kb\".format(torch.cuda.memory_cached()))","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:58:22.513172Z","iopub.execute_input":"2023-08-16T12:58:22.513657Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f'Validation accuracy: {test_correct[-1]*100/(test_b*batch_size):.2f}%')","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:22.291101Z","iopub.status.idle":"2023-08-16T12:54:22.292275Z","shell.execute_reply.started":"2023-08-16T12:54:22.292025Z","shell.execute_reply":"2023-08-16T12:54:22.292051Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot([t/test_b*batch_size for t in torch.tensor(test_correct).cpu()], label='Validation accuracy')\n\nplt.title('Accuracy Metrics')\nplt.ylabel('Accuracy')\nplt.xlabel('Epochs')\n\nplt.legend()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-08-16T12:54:22.29388Z","iopub.status.idle":"2023-08-16T12:54:22.294763Z","shell.execute_reply.started":"2023-08-16T12:54:22.294472Z","shell.execute_reply":"2023-08-16T12:54:22.294496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}