{"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":"# imports","metadata":{}},{"cell_type":"code","source":"import os\nimport cv2\nimport torch\nimport torchvision\n! pip install torchsummary\n! pip install pydicom\nimport torchsummary\nfrom torch.utils.data import Dataset, DataLoader, Subset\nimport torch.nn as nn\nfrom IPython.display import clear_output\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import accuracy_score, roc_auc_score\nfrom tqdm import tqdm\nimport numpy as np\nimport pandas as pd\nfrom PIL import Image\nfrom matplotlib import pyplot as plt\nfrom concurrent.futures import ProcessPoolExecutor \nimport torch.nn.functional as F\nimport pydicom  \nfrom tabulate import tabulate\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5"},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# config","metadata":{}},{"cell_type":"code","source":"class Config:\n    # BASE_PATH = '/kaggle/input/rsna-2023-abdominal-trauma-detection'\n    BASE_PATH = '.'\n    TRAIN_IMG_PATH='rsna-2023-atd-reduced-256-5mm/reduced_256_tickness_5'\n    SEED = 42\n    IMAGE_SIZE = [256, 256]\n    BATCH_SIZE = 25\n    EPOCHS = 3\n    DEPTH=300# median of number of instanceses is 270 ,mean=424\n    NUM_FOLDS=5\n    TARGET_COLS = [\n       'bowel_healthy', 'bowel_injury', 'extravasation_healthy',\n       'extravasation_injury', 'kidney_healthy', 'kidney_low', 'kidney_high',\n       'liver_healthy', 'liver_low', 'liver_high', 'spleen_healthy',\n       'spleen_low', 'spleen_high',\n    ]\n\ntorch.manual_seed(Config.SEED)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# addaptation for depth","metadata":{}},{"cell_type":"code","source":"def standardize_depth(tensor, target_depth=300):\n    # This function assumes tensor shape [depth, channel, height, width]\n    depth = tensor.shape[0]\n\n    if depth == target_depth:\n        return tensor\n    else:\n        if tensor.dim() == 3:\n            tensor = tensor.unsqueeze(0)  # Add a channel dimension\n        return torch.nn.functional.interpolate(\n            tensor.unsqueeze(0), \n            size=(target_depth, tensor.size(2), tensor.size(3)), \n            mode='trilinear', \n            align_corners=False\n        ).squeeze(0).squeeze(0)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# dicom to png ","metadata":{}},{"cell_type":"code","source":"def dicom_to_img(dicom_image):\n    dicom_image = pydicom.dcmread(dicom_image)\n    pixel_array = dicom_image.pixel_array\n    \n    if dicom_image.PixelRepresentation == 1:\n        bit_shift = dicom_image.BitsAllocated - dicom_image.BitsStored\n        new_array = (pixel_array << bit_shift).astype(pixel_array.dtype) >>  bit_shift\n        pixel_array = pydicom.pixel_data_handlers.util.apply_modality_lut(new_array, dicom_image)\n\n    if dicom_image.PhotometricInterpretation == \"MONOCHROME1\":\n        pixel_array = 1 - pixel_array\n\n    # transform to hounsfield units\n    pixel_array = pixel_array * dicom_image.RescaleSlope + dicom_image.RescaleIntercept\n\n    # windowing\n    window_center = int(dicom_image.WindowCenter)\n    window_width = int(dicom_image.WindowWidth)\n    img_min = window_center - window_width // 2\n    img_max = window_center + window_width // 2\n    pixel_array = pixel_array.copy()\n    pixel_array[pixel_array < img_min] = img_min\n    pixel_array[pixel_array > img_max] = img_max\n\n    # normalization\n    pixel_array = np.zeros_like(pixel_array) if pixel_array.max() == pixel_array.min() else (pixel_array - pixel_array.min()) / (pixel_array.max() - pixel_array.min())\n    return pixel_array","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#  utils for image processing","metadata":{}},{"cell_type":"code","source":"train_labels = pd.read_csv(f\"{Config.BASE_PATH}/train.csv\")\n#maby drop lower hu \n\n#input: entry to tree directory of train images\n#output:  list of path to scan ( each scan  is a list with path to the images inside it )\ndef reshapePathsToScan(train_img_path):\n    img_paths = []  \n    for dirpath, dirnames, filenames in os.walk(train_img_path):\n        if not filenames:\n            continue  # skip directories that don't contain files\n        scan = [os.path.join(dirpath, filename) for filename in filenames]\n        img_paths.append(scan)\n    return img_paths\ndef scaleSizeOfDataSet(img_paths,ratio):\n    max_number_of_images=int(len(img_paths)*ratio)\n    return img_paths[:max_number_of_images]\n    \ndef process_image(path):\n   # Open the image\n    image = Image.open(path)\n    \n    # Convert to grayscale\n    grayscale = image.convert('L')\n    \n    # Convert to numpy array and normalize\n    grayscale_np = np.array(grayscale) / 255.0\n    return grayscale_np\n\n\n\nscansPaths=reshapePathsToScan(Config.TRAIN_IMG_PATH)\nscansPaths=scaleSizeOfDataSet(scansPaths,1/30)","metadata":{"editable":true,"scrolled":true,"slideshow":{"slide_type":""},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CTScansDataSet","metadata":{}},{"cell_type":"code","source":"class CTScansDataSet(torch.utils.data.Dataset):\n    def __init__(self, train_labels,scansPaths, current_fold,num_fold=5):\n        self.train_labels=train_labels\n        self.scansPaths  = scansPaths \n        #handle 5 fold validation \n        self.num_fold = num_fold\n        self.current_fold = current_fold\n        self.kfold = KFold(n_splits=num_fold)\n        self.transform= torchvision.transforms.Compose([\n                    torchvision.transforms.Resize((256, 256),antialias=True),\n                    torchvision.transforms.RandomHorizontalFlip(),    # Random horizontal flip\n                    torchvision.transforms.RandomRotation(45),\n                    torchvision.transforms.RandomAffine(degrees=0, scale=(0.8, 1.2)),  # Apply zoom transformation\n                ])\n        \n    def __len__(self):\n        return len(self.scansPaths)\n\n    def __getitem__(self, idx):\n        #      dicom_images = select_elements_with_spacing(self.img_paths[idx],\n#                                                     spacing = 2)\n        #idx to scan-list of images\n        images_paths=self.scansPaths[idx]\n        patient_id = images_paths[0].split('/')[-3]\n        images = [process_image(path) for path in images_paths]\n        images=np.array(images)\n        images=self.rescale_scan_to_group_3(images)\n        # do augmentation to slices- move to 3d becouse we deal with resnet or other-pretraind model.\n        images = self.transform(torch.from_numpy(images).float())\n        images=standardize_depth(images)\n        # get the labels from the df \n        labels = self.train_labels[self.train_labels.patient_id == int(patient_id)].values[0][1:-1]\n        return images,labels,patient_id\n        \n    def rescale_scan_to_group_3(self,scan):\n        padding = np.zeros((1,scan.shape[1], scan.shape[2]))\n        if scan.shape[0]%3 ==1:\n            scan=np.concatenate((padding,scan,padding),axis=0)\n        elif  scan.shape[0]%3 ==2:  \n\n            scan=np.concatenate((padding,scan),axis=0)\n        return scan\n    \n    def split_fold(self):\n        \"\"\"\n        Splits the dataset into training and validation subsets based on the current fold.\n        \n        Returns:\n            tuple: A tuple containing the training and validation subsets.\n        \"\"\"\n        #split across patient *-9\n        fold_data = list(self.kfold.split(self.scansPaths))\n        train_indices, val_indices = fold_data[self.current_fold]\n        train_data = Subset(self, train_indices)\n        val_data = Subset(self,val_indices)\n        return train_data, val_data\n    \n    \n#define k-fold partition for data -defualt 5 -data loaders\n\nfolds=[CTScansDataSet(train_labels,scansPaths, current_fold=i).split_fold() for i in range (Config.NUM_FOLDS)]\ntrain_folds = [fold[0] for fold in folds]\nval_folds = [fold[1] for fold in folds]\ntrain_dataloaders= [DataLoader(train_folds[i],batch_size = Config.BATCH_SIZE, shuffle = True)  for i in range (Config.NUM_FOLDS) ]\nval_dataloaders= [DataLoader(val_folds[i],batch_size = Config.BATCH_SIZE, shuffle = False)  for i in range (Config.NUM_FOLDS) ]","metadata":{"editable":true,"slideshow":{"slide_type":""},"tags":[]},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MultiPathResNet","metadata":{}},{"cell_type":"code","source":"class MultiPathResNet(nn.Module):\n    def __init__(self):\n        super(MultiPathResNet, self).__init__()\n\n#         self.backbone=torchvision.models.mobilenet_v3_large(weights='IMAGENET1K_V1')\n        self.backbone=torchvision.models.resnet18(weights='IMAGENET1K_V1')\n        for param in self.backbone.parameters():\n            param.requires_grad = False\n#         for param in self.backbone.layer4[-1].parameters():\n#             param.requires_grad = True \n        self.backbone = nn.Sequential(*list( self.backbone.children())[:-2])\n        self.ChannelReducer= nn.Conv2d(in_channels=512, out_channels=128, kernel_size=1)\n        self.layer_normalization=nn.Sequential(torch.nn.BatchNorm3d(128),nn.ReLU())\n# avgpool\n\n# avgpool\n# Instantiate the model\n\n    def forward(self, x):\n        batch_size = x.size(0)\n        num_slices=x.size(1)//3\n        outputs = []\n        x=x.view(-1, 3, x.size(2), x.size(3))\n        x=self.backbone(x)\n        x=self.ChannelReducer(x).view(batch_size,128,num_slices,x.shape[-1],x.shape[-1])\n        x=self.layer_normalization(x)\n        return x\n    \nclass Custom3DCNN(nn.Module):\n    def __init__(self):\n        super(Custom3DCNN, self).__init__()\n        # First 3D convolution\n        \n        self.convBlock1=nn.Sequential( nn.Conv3d(128, 32, kernel_size=3,\n                                                stride=(3, 1, 1), padding=(0, 1, 1)),nn.BatchNorm3d(32), nn.ReLU(),nn.Dropout(0.3))\n        \n        # Second 3D convolution\n        self.convBlock2=nn.Sequential( nn.Conv3d(32, 8, kernel_size=3, stride=(3, 1, 1), padding=(0, 1, 1)),\n                                       nn.BatchNorm3d(8), nn.ReLU(),nn.Dropout(0.3))\n        \n        # Fully connected layers\n        self.global_avgpool = nn.AdaptiveAvgPool3d((8, None, None))\n        self.flatten_size = 8 * 8 * 8 * 8  # Adjusted due to global average pooling\n        self.Dense=nn.Sequential(nn.Linear(self.flatten_size, 512),nn.BatchNorm1d(512) ,nn.ReLU())\n        \n    def forward(self, x):\n        batch_size=x.shape[0]\n        x=self.convBlock1(x)\n        x= self.convBlock2(x)\n        x = self.global_avgpool(x)\n        x = x.view(batch_size, self.flatten_size)  # Flatten\n        x =self.Dense(x) \n        return x\n           ","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CT3DModel","metadata":{}},{"cell_type":"code","source":"class CT3DModel(nn.Module):\n    def __init__(self,num_classes=13):\n        super(CT3DModel, self).__init__()\n        self.classifier = nn.Linear(512, num_classes)\n        self.encoder=MultiPathResNet()\n        self.decoder=Custom3DCNN()\n    def forward(self, x):\n            x=self.encoder(x)\n            x=self.decoder(x)\n            x = self.classifier(x)\n            return(x)\n\n# print(model)\n# torchsummary.summary(model,(330 ,256, 256))","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# MetricsCalculator","metadata":{}},{"cell_type":"code","source":"#class to calcilate metrics.\n\nclass MetricsCalculator:\n    def __init__(self, mode = 'binary'):\n        \n        self.probabilities = []\n        self.predictions = []\n        self.targets = []\n        \n        self.mode = mode\n        \n    #logits batch_size*1*n, target batch_size*1*1 \n    def update(self, logits, target):\n        \"\"\"\n        Update the metrics calculator with predicted values and corresponding targets.\n        \n        Args:\n            predicted (torch.Tensor): Predicted values.\n            target (torch.Tensor): Ground truth targets.\n        \"\"\"\n        probabilities = F.softmax(logits, dim = 1)\n        predicted = torch.argmax(probabilities, dim=1)\n        if self.mode == 'binary':\n          #take positive class - for example if has extravision take all probabiltes that patient will have extravation and drop the probabiltes not.\n          probabilities=probabilities[:,1]\n            \n        self.probabilities.extend(probabilities.detach().cpu().numpy())\n        self.predictions.extend(predicted.detach().cpu().numpy())\n        self.targets.extend(target.detach().cpu().numpy())\n    \n    def reset(self):\n        \"\"\"Reset the stored predictions and targets.\"\"\"\n        \n        self.probabilities = []\n        self.predictions = []\n        self.targets = []\n    \n    def compute_accuracy(self):\n        \"\"\"\n        Compute the accuracy metric.\n        \n        Returns:\n            float: Accuracy.\n        \"\"\"\n        return accuracy_score(self.targets, self.predictions)\n    \n    def compute_auc(self):\n        \"\"\"\n        Compute the AUC (Area Under the Curve) metric.\n        \n        Returns:\n            float: AUC.\n        \"\"\"\n        if self.mode == 'multi':\n            return roc_auc_score(self.targets, self.probabilities, multi_class = 'ovo', labels=[0, 1, 2])\n    \n        else:\n            return roc_auc_score(self.targets, self.probabilities)\n\n\n# initialize metrics objects\ntrain_acc_bowel = MetricsCalculator('binary')\ntrain_acc_extravasation = MetricsCalculator('binary')\ntrain_acc_liver = MetricsCalculator('multi')\ntrain_acc_kidney = MetricsCalculator('multi')\ntrain_acc_spleen = MetricsCalculator('multi')\n\nval_acc_bowel = MetricsCalculator('binary')\nval_acc_extravasation = MetricsCalculator('binary')\nval_acc_liver = MetricsCalculator('multi')\nval_acc_kidney = MetricsCalculator('multi')\nval_acc_spleen = MetricsCalculator('multi')\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# main_functionmain_function","metadata":{}},{"cell_type":"code","source":"#main training function \ndef main_function():\n\n    model = CT3DModel(num_classes=13).to(device)\n    #define loss\n    optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\n    scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', patience=5, factor=0.5, verbose=True)\n    loss_two_bowl = torch.nn.CrossEntropyLoss(weight = torch.tensor([1.0, 2.0])).to(device)\n    loss_two_extra = torch.nn.CrossEntropyLoss(weight = torch.tensor([1.0, 6.0])).to(device)\n    loss_three  = torch.nn.CrossEntropyLoss(label_smoothing = 0.05, weight = torch.tensor([1.0, 2.0, 4.0])).to(device)\n    \n    train_losses_ephocs = []\n    val_losses_ephocs = []\n    best_loss=np.inf\n    \n    for epoch in range(Config.EPOCHS):\n        train_lose_ephoc =0\n        val_lose_ephoc = 0\n        train_loss=0\n        val_loss=0\n        print(f'Epoch: [{epoch+1}/{Config.EPOCHS}]')\n        print(f'Fold: {epoch%5}')\n        train_dataloader  = train_dataloaders[epoch%5]\n        val_dataloader    = val_dataloaders[epoch%5]               \n    \n        for batch_idx, (images, labels,patient_id) in enumerate(train_dataloader):\n            #################################### strat train ephoch  #########################################\n            \n            print(len(train_dataloader))\n            print(batch_idx)\n            optimizer.zero_grad()\n            # Move data to GPU           \n            images = images.to(device)  \n            concatenated_categorial=convertOneHotToCategorial(labels).to(device)  \n            bowel_batch_labels,extravasation_batch_labels=concatenated_categorial[:,0],concatenated_categorial[:,1]\n            kidney_batch_labels,liver_batch_labels,spleen_batch_labels=concatenated_categorial[:,2],concatenated_categorial[:,3],concatenated_categorial[:,4]\n\n            ## debug  ### \n            #----plot 10 fist scans\n            # debug_plot_images(images,10,patient_id)\n            # print(images.shape)\n            # print(labels)\n                 \n            \n            outputs = model(images)\n            print(outputs)\n            print(outputs[:, 0:2].dtype)  # Should be torch.float32 or similar      \n            bowel_loss = loss_two_bowl(outputs[:, 0:2], bowel_batch_labels)\n            extravasation_loss = loss_two_extra(outputs[:, 2:4], extravasation_batch_labels)\n            kidney_loss = loss_three(outputs[:, 4:7], kidney_batch_labels)\n            liver_loss = loss_three(outputs[:, 7:10], liver_batch_labels)\n            spleen_loss = loss_three(outputs[:, 10:13], spleen_batch_labels)\n            ###!!!!!!!!!To-do -add anyinjuery loss!!!!!!!!!!!!!!!\n            train_loss = bowel_loss + extravasation_loss + kidney_loss + liver_loss + spleen_loss\n            print(train_loss)\n            # Backward pass\n            \n            train_loss.backward()\n            # Update weights\n            optimizer.step()\n            #aggregate results for metrics\n            train_lose_ephoc += train_loss.item()\n            train_acc_bowel.update(outputs[:, 0:2], bowel_batch_labels)\n            train_acc_extravasation.update(outputs[:, 2:4], extravasation_batch_labels)\n            train_acc_kidney.update(outputs[:, 4:7], kidney_batch_labels)\n            train_acc_liver.update(outputs[:, 7:10], liver_batch_labels)\n            train_acc_spleen.update(outputs[:, 10:13], spleen_batch_labels)\n            \n        ############################################end train ephoch  ################################## \n        train_lose_ephoc=train_lose_ephoc/len(train_dataloader)\n        train_losses_ephocs.append( train_lose_ephoc)\n          \n        model.eval()  # Set the model to evaluation mode    \n        with torch.no_grad():\n            for val_ind, (images, labels,patient_id) in enumerate(val_dataloader):  \n                ##########################  strat eval  ephoch   #############################\n                \n                images = images.to(device)   \n                #parse labels\n                concatenated_categorial=convertOneHotToCategorial(labels).to(device)\n                bowel_batch_labels,extravasation_batch_labels=concatenated_categorial[:,0],concatenated_categorial[:,1]\n                kidney_batch_labels,liver_batch_labels,spleen_batch_labels=concatenated_categorial[:,2],concatenated_categorial[:,3],concatenated_categorial[:,4]\n           \n                outputs = model(images)\n                bowel_loss = loss_two_bowl(outputs[:, 0:2], bowel_batch_labels)\n                extravasation_loss = loss_two_extra(outputs[:, 2:4], extravasation_batch_labels)\n                kidney_loss = loss_three(outputs[:, 4:7], kidney_batch_labels)\n                liver_loss = loss_three(outputs[:, 7:10], liver_batch_labels)\n                spleen_loss = loss_three(outputs[:, 10:13], spleen_batch_labels)\n                \n                ###!!!!!!!!!To-do -add any injuery loss!!!!!!!!!!!!!!!\n                \n                val_loss = bowel_loss + extravasation_loss + kidney_loss + liver_loss + spleen_loss \n                # calculate validation metrics    \n                val_lose_ephoc += val_loss.item()\n                val_acc_bowel.update(outputs[:, 0:2], bowel_batch_labels)\n                val_acc_extravasation.update(outputs[:, 2:4], extravasation_batch_labels)\n                val_acc_kidney.update(outputs[:, 4:7], kidney_batch_labels)\n                val_acc_liver.update(outputs[:, 7:10], liver_batch_labels)\n                val_acc_spleen.update(outputs[:, 10:13], spleen_batch_labels)\n                \n                ################################# end eval ephoch #############################\n                \n        # to get out of plauto         \n        val_lose_ephoc=val_lose_ephoc/len(val_dataloader)\n        scheduler.step(val_lose_ephoc)      \n        val_losses_ephocs.append(val_lose_ephoc)     \n        if val_loss <= best_loss:\n          best_loss = val_loss\n        #   torch.save(model.state_dict(), './drive/MyDrive/model_weights.pth')\n        print(f\"Epoch [{epoch+1}/{Config.EPOCHS}] - train Loss: {train_lose_ephoc:.4f} - Val Loss: {val_lose_ephoc:.4f}\")\n        # accuracy and auc data\n        print(train_acc_bowel.targets)\n        print(train_acc_bowel.predictions)\n        print(train_acc_bowel.probabilities)\n        metrics_data = [\n                    [\"Bowel\", \n                        train_acc_bowel.compute_accuracy(),\n                        val_acc_bowel.compute_accuracy(),\n                        train_acc_bowel.compute_auc(),\n                        val_acc_bowel.compute_auc()],\n                    [\"Extravasation\", \n                        train_acc_extravasation.compute_accuracy(),\n                        val_acc_extravasation.compute_accuracy(),\n                        train_acc_extravasation.compute_auc(),\n                        val_acc_extravasation.compute_auc()],\n                    [\"Liver\", \n                        train_acc_liver.compute_accuracy(),\n                        val_acc_liver.compute_accuracy(),\n                        train_acc_liver.compute_auc(),\n                        val_acc_liver.compute_auc()],\n                    [\"Kidney\", \n                        train_acc_kidney.compute_accuracy(),\n                        val_acc_kidney.compute_accuracy(),\n                        train_acc_kidney.compute_auc(),\n                        val_acc_kidney.compute_auc()],\n                    [\"Spleen\", \n                        train_acc_spleen.compute_accuracy(),\n                        val_acc_spleen.compute_accuracy(),\n                        train_acc_spleen.compute_auc(),\n                        val_acc_spleen.compute_auc()]\n                ]\n        # verbose\n        print(tabulate(metrics_data, headers=[\"\", \"Train Acc\", \"Val Acc\", \"Train AUC\", \"Val AUC\"]))\n        #reset metrics\n        train_acc_bowel.reset()\n        train_acc_extravasation.reset()\n        train_acc_liver.reset()\n        train_acc_kidney.reset()\n        train_acc_spleen.reset()\n        val_acc_bowel.reset()\n        val_acc_extravasation.reset()\n        val_acc_liver.reset()\n        val_acc_kidney.reset()\n        val_acc_spleen.reset()     \n        \ndef convertOneHotToCategorial(labels):\n        bowel_batch_labels = np.argmax(labels[:,0:2],axis=1 ,keepdims = True).squeeze(-1)\n        extravasation_batch_labels = np.argmax(labels[:,2:4],axis=1 ,keepdims = True).squeeze(-1)\n        kidney_batch_labels = np.argmax(labels[:,4:7],axis=1 ,keepdims = True).squeeze(-1)\n        liver_batch_labels = np.argmax(labels[:,7:10],axis=1, keepdims = True).squeeze(-1)\n        spleen_batch_labels = np.argmax(labels[:,10:],axis=1, keepdims = True).squeeze(-1)\n        concatenated_categorial =np.vstack((bowel_batch_labels, extravasation_batch_labels, kidney_batch_labels,liver_batch_labels,spleen_batch_labels)).T\n        concatenated_categorial = torch.tensor(concatenated_categorial)\n        return concatenated_categorial\n\ndef debug_plot_images(images,num_scans,patient_id):\n    \n    first_scans = images[:num_scans]\n            # Plot the firsts scans\n    for scan_idx, scan in enumerate(first_scans):\n                print(f'patient id {patient_id[scan_idx]}')\n                num_images = scan.shape[0]\n                plt.figure(figsize=(20, 20))  # Adjust the figure size as needed\n                plt.suptitle(f\"Scan {scan_idx + 1}\")\n                for image_idx in range(num_images):\n                    plt.subplot(30, 10, image_idx + 1)\n                    plt.imshow(scan[image_idx], cmap='gray')  # Assuming grayscale images\n                    plt.axis('off')  # Turn off axis labels\n\n                plt.show()\n    \nmain_function()\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.cuda.empty_cache()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), './drive/MyDrive/model_ResNet18+3DConv_weights.pth')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"render_window = vtk.vtkRenderWindow()\nrender_window_interactor = vtk.vtkRenderWindowInteractor()\nrender_window_interactor.SetRenderWindow(render_window)","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Plot training and validation progress\nplt.figure(figsize=(12, 4))\nplt.subplot(1, 2, 1)\nplt.plot(train_losses_ephocs, label='Train')\nplt.plot(val_losses_ephocs, label='Validation')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.title('Training and Validation Loss')\n\nplt.tight_layout()\nplt.show()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# model = torch.load('model.pth')\n# model.load_state_dict(torch.load('./drive/MyDrive/DATA/model_weights.pth'))\n\n# trt_model = torch_tensorrt.compile(model, \n#     inputs= [torch_tensorrt.Input((1, 3, 224, 224))],\n#     enabled_precisions= { torch_tensorrt.dtype.half} # Run with FP16\n# )","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}