{"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":"\n# Inference Notebook for RSNA 2023\n\nThis notebook outlines the steps to run inference using the pre-trained models developed for the RSNA 2023 Kaggle competition. The models were trained using a 2.5D CNN approach leveraging EfficientNet backbone for feature extraction. \n\nIn the following sections, we will set up the environment, load the pre-trained models, and define utility functions to process the data. Following this, we will load the test data, run the inference to generate predictions, and prepare a submission file.\n","metadata":{}},{"cell_type":"code","source":"!pip install -q /kaggle/input/rsna-atd-whl-ds/python_gdcm-3.0.22-cp310-cp310-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install -q /kaggle/input/rsna-atd-whl-ds/pylibjpeg-1.4.0-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:34:34.272861Z","iopub.execute_input":"2023-09-25T13:34:34.273267Z","iopub.status.idle":"2023-09-25T13:35:42.544556Z","shell.execute_reply.started":"2023-09-25T13:34:34.273237Z","shell.execute_reply":"2023-09-25T13:35:42.542956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n# Import necessary packages\n# import packages\nimport os\nimport pickle\nfrom tqdm.notebook import tqdm\nimport random\nfrom tabulate import tabulate\n\nimport cv2\nimport torch\nimport timm\nfrom glob import glob\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport matplotlib.pyplot as plt\nimport torchvision.transforms.v2 as t\nimport gc\n\nfrom sklearn.model_selection import KFold, StratifiedKFold\nfrom sklearn.metrics import accuracy_score, roc_auc_score\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom PIL import Image\nimport torch.optim as optim\nfrom torchvision import models\nfrom torchvision.transforms.v2 import Resize, Compose, RandomHorizontalFlip, ColorJitter, RandomAffine, RandomErasing, ToTensor\n\n# Set constants\nDEBUG = False\n\n# Define data paths (adjust as necessary)\nDATA_DIR = '/path/to/data'\nTEST_IMG_PATH = os.path.join(DATA_DIR, 'test_images')\n","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:42.547381Z","iopub.execute_input":"2023-09-25T13:35:42.547779Z","iopub.status.idle":"2023-09-25T13:35:42.559852Z","shell.execute_reply.started":"2023-09-25T13:35:42.547747Z","shell.execute_reply":"2023-09-25T13:35:42.558325Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Meta Data","metadata":{}},{"cell_type":"code","source":"def seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True ","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:42.561212Z","iopub.execute_input":"2023-09-25T13:35:42.561596Z","iopub.status.idle":"2023-09-25T13:35:42.573723Z","shell.execute_reply.started":"2023-09-25T13:35:42.561564Z","shell.execute_reply":"2023-09-25T13:35:42.572727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BASE_PATH = '/kaggle/input/rsna-2023-abdominal-trauma-detection'\nIMG_DIR = '/tmp/dataset/rsna-atd'\nNUM_SLICES = 4","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:42.575239Z","iopub.execute_input":"2023-09-25T13:35:42.575541Z","iopub.status.idle":"2023-09-25T13:35:42.592364Z","shell.execute_reply.started":"2023-09-25T13:35:42.575516Z","shell.execute_reply":"2023-09-25T13:35:42.591185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test Paths","metadata":{}},{"cell_type":"markdown","source":"## Test Dataframe","metadata":{}},{"cell_type":"code","source":"test_df = pd.read_csv(f'{BASE_PATH}/test_series_meta.csv')\ntest_df['dicom_folder'] = BASE_PATH + '/' + 'test_images'\\\n                                    + '/' + test_df.patient_id.astype(str)\\\n                                    + '/' + test_df.series_id.astype(str)\ntest_folders = test_df.dicom_folder.tolist()\n\ntest_paths = []\nfor folder in tqdm(test_folders):\n    paths = sorted(glob(os.path.join(folder, '*dcm')),\n                   key=lambda x: int(x.split('/')[-1].split('.')[0]))\n    NUM_DICOM = len(paths)\n    if len(test_folders)>6: # private test; contains all dicom files/folders\n        STRIDE = -(-NUM_DICOM // (NUM_SLICES + 4))\n        test_paths += [paths[STRIDE:NUM_DICOM-3*STRIDE:STRIDE]]\n    else: # we can't access all the test dicom files in public test\n        test_paths += [paths]\n\ntest_df['dicom_paths'] = test_paths\ntest_df = test_df[test_df.dicom_paths.map(len)>0] # in public test not all folder contains dicom file\n\ntest_df['image_path'] = f'{IMG_DIR}/test_images'\\\n                    + '/' + test_df.patient_id.astype(str)\\\n                    + '/' + test_df.series_id.astype(str) +'.png'\n# test_df = test_df.drop_duplicates()\n\ntest_df.head(2)","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:42.596134Z","iopub.execute_input":"2023-09-25T13:35:42.596557Z","iopub.status.idle":"2023-09-25T13:35:42.654884Z","shell.execute_reply.started":"2023-09-25T13:35:42.596522Z","shell.execute_reply":"2023-09-25T13:35:42.653751Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# DICOM to PNG","metadata":{}},{"cell_type":"code","source":"!rm -r /tmp/Dataset/rsna-atd\nos.makedirs('/tmp/dataset/rsna-atd/train_images', exist_ok = True)\nos.makedirs('/tmp/dataset/rsna-atd/test_images', exist_ok = True)","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:42.656272Z","iopub.execute_input":"2023-09-25T13:35:42.656613Z","iopub.status.idle":"2023-09-25T13:35:43.778241Z","shell.execute_reply.started":"2023-09-25T13:35:42.656584Z","shell.execute_reply":"2023-09-25T13:35:43.776632Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dicom Utils","metadata":{}},{"cell_type":"code","source":"import cv2\nimport pydicom\n\ndef standardize_pixel_array(dicom_image):\n    \"\"\"\n    Standardizes a DICOM pixel array by applying various transformations.\n    \n    Args:\n        dicom_path (str): Path to the DICOM image file.\n        \n    Returns:\n        np.ndarray: The standardized pixel array of the DICOM image.\n    \"\"\"\n    pixel_array = dicom_image.pixel_array\n    \n    if dicom_image.PixelRepresentation == 1:\n        bit_shift = dicom_image.BitsAllocated - dicom_image.BitsStored\n        dtype = pixel_array.dtype \n        new_array = (pixel_array << bit_shift).astype(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    intercept = dicom_image.RescaleIntercept\n    slope = dicom_image.RescaleSlope\n    pixel_array = pixel_array * slope + intercept\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    if pixel_array.max() == pixel_array.min():\n        pixel_array = np.zeros_like(pixel_array)  # Handle case of constant array\n    else:\n        pixel_array = (pixel_array - pixel_array.min()) / (pixel_array.max() - pixel_array.min())\n\n    return pixel_array\n\ndef read_xray(path, fix_monochrome = True):\n    dicom = pydicom.dcmread(path)\n    data = standardize_pixel_array(dicom)\n    data = data - np.min(data)\n    data = data / (np.max(data) + 1e-5)\n    if fix_monochrome and dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = 1.0 - data\n    IMG_SIZE = [224, 224]\n    data = cv2.resize(data, IMG_SIZE, cv2.INTER_LINEAR)\n    data = (data * 255).astype(np.uint8)\n    return data\n\ndef load_scan(paths):\n    IMG_SIZE = [224, 224]\n    img = np.empty(shape=(*IMG_SIZE, NUM_SLICES), dtype=np.uint8)\n    for i, path in enumerate(paths):\n        img[...,i] = read_xray(path)\n    return img\n\ndef load_img(path):\n    img = cv2.imread(path, -1)[...,::-1]\n    return img\n    \ndef resize_and_save(paths):\n    img = load_scan(paths)\n    file_path = paths[0]\n    sub_path = file_path.split(\"/\",4)[-1].split('.dcm')[0] + '.png'\n    infos = sub_path.split('/')\n    split = infos[-4]\n    pid = infos[-3]\n    sid = infos[-2]\n    iid = infos[-1]; iid = iid.replace('.png','')\n    new_path = os.path.join(IMG_DIR, split, pid, sid + '.png')\n    os.makedirs(new_path.rsplit('/',1)[0], exist_ok=True)\n    cv2.imwrite(new_path, img[...,::-1])\n    del img; gc.collect()\n    return \n\ndef show_img(img):\n    num_channels = img.shape[-1]\n    fig, axes = plt.subplots(1, num_channels+1, figsize=(num_channels*5, 5))\n    axes[0].imshow(img)\n    axes[0].set_title('Original Image')\n    axes[0].axis('off')\n\n    for i in range(num_channels):\n        axes[i+1].imshow(img[:, :, i], cmap='gray')\n        axes[i+1].set_title(f'Channel: {i:02d}')\n        axes[i+1].axis('off')\n\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:43.780864Z","iopub.execute_input":"2023-09-25T13:35:43.781279Z","iopub.status.idle":"2023-09-25T13:35:43.806138Z","shell.execute_reply.started":"2023-09-25T13:35:43.781242Z","shell.execute_reply":"2023-09-25T13:35:43.804826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.dicom_paths.iloc[0]","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:43.807858Z","iopub.execute_input":"2023-09-25T13:35:43.808251Z","iopub.status.idle":"2023-09-25T13:35:43.829731Z","shell.execute_reply.started":"2023-09-25T13:35:43.80822Z","shell.execute_reply":"2023-09-25T13:35:43.828642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = load_scan(['/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images/10004/21057/1029.dcm',\n                 '/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images/10004/21057/1032.dcm',\n                 '/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images/10004/21057/1045.dcm',\n                 '/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images/10004/21057/1050.dcm'\n                ])\nprint(img.shape)\nshow_img(img)","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:43.83101Z","iopub.execute_input":"2023-09-25T13:35:43.83136Z","iopub.status.idle":"2023-09-25T13:35:44.843822Z","shell.execute_reply.started":"2023-09-25T13:35:43.831331Z","shell.execute_reply":"2023-09-25T13:35:44.842919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Conversion","metadata":{}},{"cell_type":"code","source":"%%time\nfrom joblib import Parallel, delayed\nfile_paths = test_df.dicom_paths.tolist()\n_ = Parallel(n_jobs=-1,backend='loky')(delayed(resize_and_save)(file_path)\\\n                                                  for file_path in tqdm(file_paths))\ndel _; gc.collect()","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:44.845139Z","iopub.execute_input":"2023-09-25T13:35:44.845839Z","iopub.status.idle":"2023-09-25T13:35:46.646121Z","shell.execute_reply.started":"2023-09-25T13:35:44.845803Z","shell.execute_reply":"2023-09-25T13:35:46.645228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img = load_img(f'{IMG_DIR}/test_images/63706/39279.png')\nshow_img(img)","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:46.647333Z","iopub.execute_input":"2023-09-25T13:35:46.648382Z","iopub.status.idle":"2023-09-25T13:35:47.472691Z","shell.execute_reply.started":"2023-09-25T13:35:46.648339Z","shell.execute_reply":"2023-09-25T13:35:47.47135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Pipeline","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"TEST_IMG_PATH = '/kaggle/input/rsna-2023-abdominal-trauma-detection/test_images'","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:47.474303Z","iopub.execute_input":"2023-09-25T13:35:47.474662Z","iopub.status.idle":"2023-09-25T13:35:47.480452Z","shell.execute_reply.started":"2023-09-25T13:35:47.47463Z","shell.execute_reply":"2023-09-25T13:35:47.478763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def preprocess_png(png_path):\n    # Read the image using OpenCV\n    img = cv2.imread(png_path, cv2.IMREAD_UNCHANGED)\n    \n    # Check if the image was loaded successfully\n    if img is None:\n        raise FileNotFoundError(f\"No image found at {png_path}\")\n\n    # Convert the image to greyscale\n    greyscale = cv2.cvtColor(img, cv2.COLOR_BGR2GRAY)\n    \n    # Normalize the pixel values to be between 0 and 1\n    greyscale = greyscale / 255.0\n\n    return greyscale","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:47.482251Z","iopub.execute_input":"2023-09-25T13:35:47.483185Z","iopub.status.idle":"2023-09-25T13:35:47.498292Z","shell.execute_reply.started":"2023-09-25T13:35:47.483133Z","shell.execute_reply":"2023-09-25T13:35:47.496557Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import transforms\nclass AbdominalTestData(Dataset):\n    \"\"\"\n    Custom dataset class for handling abdominal trauma test data classification.\n    \n    Args:\n        img_paths (list of strings): List containing all image paths of a patient\n        target_size (tuple): The target size to resize the images to\n        ext (str): The extension of the image files\n        transform (callable, optional): A function/transform to apply to the images\n    \"\"\"\n    \n    def __init__(self, img_paths,labels=None,target_size=(224, 224), ext='png', transform=None):\n        \n        super().__init__()\n        \n        self.img_paths = img_paths\n        self.labels = labels\n        self.target_size = target_size\n        self.ext = ext\n\n        if transform:\n            self.transform = transform\n        else:\n            self.transform = transforms.Compose([\n                transforms.Resize(target_size, interpolation=Image.BILINEAR), \n                transforms.ToTensor(),\n            ])\n    \n    def __len__(self):\n        \"\"\"\n        Returns the total number of samples in the dataset.\n        \"\"\"\n        return len(self.img_paths)\n    \n    def __getitem__(self, idx):\n        \"\"\"\n        Returns a sample from the dataset at the given index.\n        \n        Args:\n            idx (int): Index\n        \n        Returns:\n            tuple: (image, file_path)\n        \"\"\"\n        \n        file_path = self.img_paths[idx]\n        img = Image.open(file_path)\n        \n        if self.ext == 'png':\n            img = img.convert('RGBA')\n        elif self.ext in ['jpg', 'jpeg']:\n            img = img.convert('RGB')\n        else:\n            raise ValueError(\"Image extension not supported\")\n\n        img = self.transform(img)\n        \n        if self.labels is not None:\n            label = self.labels[idx]\n            return img, label\n        \n        return img","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:47.504775Z","iopub.execute_input":"2023-09-25T13:35:47.505206Z","iopub.status.idle":"2023-09-25T13:35:47.519758Z","shell.execute_reply.started":"2023-09-25T13:35:47.505174Z","shell.execute_reply":"2023-09-25T13:35:47.518319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualization","metadata":{}},{"cell_type":"code","source":"def display_batch(batch, size=2):\n    if isinstance(batch, tuple):\n        imgs, tars = batch\n        tars = torch.cat(tars, dim=-1).numpy()\n    else:\n        imgs = batch\n        tars = None\n\n    plt.figure(figsize=(size*5, 10))\n    for img_idx in range(size):\n        plt.subplot(1, size, img_idx+1)\n        if tars is not None:\n            plt.title(f'{tars[img_idx].round(2)}', fontsize=12)\n        \n        # Move the channel dimension to the end for matplotlib to display the image correctly\n        img = imgs[img_idx].permute(1, 2, 0).numpy()\n        print(img.shape)\n        plt.imshow(img)\n        plt.xticks([]); plt.yticks([])\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:47.521818Z","iopub.execute_input":"2023-09-25T13:35:47.522332Z","iopub.status.idle":"2023-09-25T13:35:47.544201Z","shell.execute_reply.started":"2023-09-25T13:35:47.522284Z","shell.execute_reply":"2023-09-25T13:35:47.541883Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fold_df = test_df.copy()\npaths  = fold_df.image_path.tolist()\nprint(paths)\nlabels = None\ndataset = AbdominalTestData(paths, labels)\ndataloader = DataLoader(dataset, batch_size=20,shuffle=True, num_workers=2)\nfor batch in dataloader:\n    display_batch(batch,size=3)\n    break","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:47.546311Z","iopub.execute_input":"2023-09-25T13:35:47.546831Z","iopub.status.idle":"2023-09-25T13:35:48.192968Z","shell.execute_reply.started":"2023-09-25T13:35:47.546792Z","shell.execute_reply":"2023-09-25T13:35:48.191553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utility","metadata":{}},{"cell_type":"code","source":"def mc_proc(pred):\n    argmax = np.argmax(pred, axis=1).astype('uint8')\n    one_hot = tf.keras.utils.to_categorical(argmax, num_classes=3)\n    return one_hot.astype('uint8')\n\ndef sc_proc(pred, thr=0.5):\n    proc_pred = (pred > thr).astype('uint8')\n    return proc_pred\n\ndef post_proc(pred):\n    proc_pred = np.empty((pred.shape[0], 2 + 2 + 3*3), dtype=np.uint8)\n\n    # bowel, extravasation\n    proc_pred[:, 0] = sc_proc(pred[:, 0])\n    proc_pred[:, 1] = 1 - proc_pred[:, 0]\n    proc_pred[:, 2] = sc_proc(pred[:, 1])\n    proc_pred[:, 3] = 1 - proc_pred[:, 2]\n    \n    # liver, kidney, sneel\n    proc_pred[:, 4:7] = mc_proc(pred[:, 2:5])\n    proc_pred[:, 7:10] = mc_proc(pred[:, 5:8])\n    proc_pred[:, 10:13] = mc_proc(pred[:, 8:11])\n\n    return proc_pred\n\ndef post_proc_v2(pred):\n    proc_pred = np.empty((pred.shape[0], 2*2 + 3*3), dtype='float32')\n\n    # bowel, extravasation\n    proc_pred[:, 0] = 1 - pred[:, 0] # bowel-healthy\n    proc_pred[:, 1] = pred[:, 0] # bowel-injured\n    proc_pred[:, 2] = 1 - pred[:, 1] # extra-healthy\n    proc_pred[:, 3] = pred[:, 1] # extra-injured\n    \n    # liver, kidney, sneel\n    proc_pred[:, 4:7] = pred[:, 2:5]\n    proc_pred[:, 7:10] = pred[:, 5:8]\n    proc_pred[:, 10:13] = pred[:, 8:11]\n\n    return proc_pred","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:48.194774Z","iopub.execute_input":"2023-09-25T13:35:48.195164Z","iopub.status.idle":"2023-09-25T13:35:48.209605Z","shell.execute_reply.started":"2023-09-25T13:35:48.195128Z","shell.execute_reply":"2023-09-25T13:35:48.208582Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Architecture","metadata":{}},{"cell_type":"code","source":"class ChannelSqueezer(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.expand = nn.Sequential(\n            nn.Conv2d(4, 32, kernel_size=3, padding=1),\n            nn.GELU()\n        )\n        self.squeeze = nn.Sequential(\n            nn.Conv2d(32, 3, kernel_size=3, padding=1),\n            nn.GELU()\n        )\n\n    def forward(self, x):\n        x = self.expand(x)\n        x = self.squeeze(x)\n        return x\n\nclass CNNModel(nn.Module):\n    def __init__(self, backbone, pretrained=False):\n        super().__init__()\n        \n        self.channel_squeezer = ChannelSqueezer()\n        self.feature_extractor = timm.create_model(\n            backbone,\n            in_chans=3,\n            pretrained=pretrained\n        )\n        \n#         f = self.feature_extractor.classifier.in_features\n#         self.feature_extractor.classifier = nn.Identity()\n        f = self.feature_extractor.head.in_features\n        self.feature_extractor.head = nn.Identity()\n        self.mlp = nn.Sequential(\n#             nn.BatchNorm2d(f),\n            nn.Linear(f, 64),\n            nn.SiLU(),\n#             nn.BatchNorm2d(64),\n            nn.Linear(64, 32),\n            nn.SiLU(),\n            nn.Dropout(0.4)\n        )\n        # for bowel and extravasation\n        self.logit1 = nn.Linear(32, 1)\n        self.logit2 = nn.Linear(32, 1)\n        \n        # for kidney, liver, spleen\n        self.logit3 = nn.Linear(32, 3)\n        self.logit4 = nn.Linear(32, 3)\n        self.logit5 = nn.Linear(32, 3)\n    \n    def forward(self, x):\n        x = self.channel_squeezer(x)\n        x = self.feature_extractor(x)\n        x = torch.flatten(x, 1)\n        x = self.mlp(x)\n        \n        # output logits\n        bowel = self.logit1(x)\n        extravasation = self.logit2(x)\n        kidney = self.logit3(x)\n        liver = self.logit4(x)\n        spleen = self.logit5(x)\n        \n        return bowel, extravasation, kidney, liver, spleen","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:48.210842Z","iopub.execute_input":"2023-09-25T13:35:48.21134Z","iopub.status.idle":"2023-09-25T13:35:48.233292Z","shell.execute_reply.started":"2023-09-25T13:35:48.211309Z","shell.execute_reply":"2023-09-25T13:35:48.232175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load the Model","metadata":{}},{"cell_type":"code","source":"# Define the model architecture (adjust as necessary)\nbackbone = 'vit_base_patch16_224.augreg_in21k_ft_in1k'\nmodel_dir_cls = '/kaggle/input/vit-b16-pytorch-25d'\n# device = torch.device('cuda')\ndevice = torch.device('cpu')\nREPLICAS = 1\nimg_size = [224,224]\nprint(f'REPLICAS: {REPLICAS}')\n# Load pre-trained models (add paths to your pre-trained model weights)\nmodel_paths = sorted(glob(f'{model_dir_cls}/*.pth'))[:1]\nprint(model_paths)\nmodels = []\nfor model_path in model_paths:\n    try:\n        model = CNNModel(backbone, pretrained=False)\n        model = model.to(device)\n        sd = torch.load(model_path, map_location=torch.device('cpu'))\n\n        # Checking if the loaded object is a dict or a model instance\n        if isinstance(sd, dict):\n            if 'model_state_dict' in sd.keys():\n                sd = sd['model_state_dict']\n            sd = {k[7:] if k.startswith('module.') else k: sd[k] for k in sd.keys()}\n        elif isinstance(sd, torch.nn.Module):\n            # Handling case where the entire model object was saved\n            # (not recommended due to potential compatibility issues)\n            model = sd\n            sd = None\n        else:\n            raise ValueError(\"Unrecognized format for loaded state dict\")\n\n        if sd is not None:\n            model.load_state_dict(sd, strict=True)\n        \n        model.eval()\n        models.append(model)\n    except Exception as e:\n        print(f\"Could not load model from {model_path}: {e}\")\nlen(models)\n","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:48.234672Z","iopub.execute_input":"2023-09-25T13:35:48.235276Z","iopub.status.idle":"2023-09-25T13:35:52.918364Z","shell.execute_reply.started":"2023-09-25T13:35:48.235245Z","shell.execute_reply":"2023-09-25T13:35:52.917024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Getting unique patient IDs from test dataset\nimport torch.nn.functional as F\npatient_ids = test_df['patient_id'].unique()\n\n# Initializing array to store predictions\npatient_preds = np.zeros(shape=(len(patient_ids), 2*2 + 3*3), dtype='float32')\nwith torch.no_grad():\n# Iterating over each patient\n    for pidx, patient_id in tqdm(enumerate(patient_ids), total=len(patient_ids), desc=\"Patients \"):\n        # Query the dataframe for a particular patient\n        patient_df = test_df.query(\"patient_id == @patient_id\", engine='python')\n\n        # Initializing model predictions array\n        model_preds = np.zeros(shape=(1, 11), dtype=np.float32)\n\n        print(\"=\"*25)\n        print(f\"   Patient ID: {patient_id}\")\n        print(\"=\"*25)\n\n        # Iterating over each model\n        for _, model in enumerate(models):\n#             print(model)\n\n            # Getting image paths for a patient\n            patient_paths = patient_df.image_path.tolist()\n            print(patient_paths)\n            # Setting batch size based on number of patient paths and dimension of image\n            dim = np.prod(img_size)**0.5\n            if dim >= 1024:\n                batch_size = REPLICAS * int(4 * 2)\n            elif dim >= 768:\n                batch_size = REPLICAS * int(16 * 2)\n            elif dim >= 640:\n                batch_size = REPLICAS * int(28 * 2)\n            else:\n                batch_size = REPLICAS * int(32 * 2)\n                \n            min_bs = 2**np.floor(np.log2(len(patient_paths)))\n            batch_size = min(min_bs, batch_size)\n            batch_size = int(batch_size)\n            # Building dataset for prediction\n            preds = []\n            test_data = AbdominalTestData(patient_paths)\n            dtest = DataLoader(test_data, batch_size=batch_size, shuffle=False)\n            # Iterating over each fold\n            # Loading a PyTorch model from a fold path\n\n            # Iterating over batches and getting predictions\n            for batch_idx, batch_data in enumerate(tqdm(dtest)):\n                inputs = batch_data.to(device)\n                pred = model(inputs)\n                bowel =F.sigmoid(pred[0].cpu()).numpy().flatten()\n                extra = F.sigmoid(pred[1].cpu()).numpy().flatten()\n                kidney = F.softmax(pred[2].cpu(),dim =1).numpy().flatten()\n                liver = F.softmax(pred[3].cpu(),dim =1).numpy().flatten()\n                spleen = F.softmax(pred[4].cpu(),dim =1).numpy().flatten()\n\n                preds.append(np.concatenate((bowel,extra, kidney, liver, spleen), axis=0))                           \n\n            preds = np.array(preds).astype('float32')\n            preds = preds.reshape(len(patient_paths), 11)\n            pred = np.max(preds, axis=0)\n\n            # Store model's prediction\n            model_preds += pred / (len(models))\n\n                # Deleting variables to free up memory\n            del model, pred\n            gc.collect()\n            print('\\n')\n\n            del dtest, patient_paths; gc.collect()\n\n        # Adding processed predictions to patient_preds\n        # (define the post_proc_v2 function to work with your predictions)\n        patient_preds[pidx, :] += post_proc_v2(model_preds)[0]\n\n        del model_preds\n        gc.collect()\n\nprint(\"Prediction Done!\")\n","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:52.91981Z","iopub.execute_input":"2023-09-25T13:35:52.920173Z","iopub.status.idle":"2023-09-25T13:35:56.003356Z","shell.execute_reply.started":"2023-09-25T13:35:52.920145Z","shell.execute_reply":"2023-09-25T13:35:56.002074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"# Create Submission\ntarget_col  = [\"bowel_healthy\", \"bowel_injury\", \"extravasation_healthy\",\n                   \"extravasation_injury\", \"kidney_healthy\", \"kidney_low\",\n                   \"kidney_high\", \"liver_healthy\", \"liver_low\", \"liver_high\",\n                   \"spleen_healthy\", \"spleen_low\", \"spleen_high\"]\npred_df = pd.DataFrame({'patient_id':patient_ids,})\npred_df[target_col] = patient_preds.astype('float32')\n\n# Align with sample submission\nsub_df = pd.read_csv(f'{BASE_PATH}/sample_submission.csv')\nsub_df = sub_df[['patient_id']]\nsub_df = sub_df.merge(pred_df, on='patient_id', how='left')\n\n# Store submission\n# sub_df.to_csv('submission.csv',index=False)\n# sub_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:56.004943Z","iopub.execute_input":"2023-09-25T13:35:56.005361Z","iopub.status.idle":"2023-09-25T13:35:56.030191Z","shell.execute_reply.started":"2023-09-25T13:35:56.005322Z","shell.execute_reply":"2023-09-25T13:35:56.029137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"scale_by_2 = ['kidney_low','liver_low','spleen_low','bowel_injury']\nscale_by_4 = ['spleen_high','kidney_high','liver_high']\nscale_by_6 = ['extravasation_injury']\nscale_healthy = ['bowel_healthy', 'extravasation_healthy', 'kidney_healthy', 'liver_healthy', 'spleen_healthy']\nsf_2 = 2.8461531332\nsf_4 = 4.841531\nsf_6 = 20.81635153\nscale_h = 0.99519515313","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:56.031446Z","iopub.execute_input":"2023-09-25T13:35:56.031922Z","iopub.status.idle":"2023-09-25T13:35:56.038995Z","shell.execute_reply.started":"2023-09-25T13:35:56.031881Z","shell.execute_reply":"2023-09-25T13:35:56.037757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:56.040244Z","iopub.execute_input":"2023-09-25T13:35:56.040569Z","iopub.status.idle":"2023-09-25T13:35:56.069536Z","shell.execute_reply.started":"2023-09-25T13:35:56.040541Z","shell.execute_reply":"2023-09-25T13:35:56.068435Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Scale each target \nsub_df[scale_by_2] *=sf_2\nsub_df[scale_by_4] *=sf_4\nsub_df[scale_by_6] *=sf_6\nsub_df[scale_healthy] *=scale_h","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:56.070733Z","iopub.execute_input":"2023-09-25T13:35:56.071158Z","iopub.status.idle":"2023-09-25T13:35:56.090192Z","shell.execute_reply.started":"2023-09-25T13:35:56.071118Z","shell.execute_reply":"2023-09-25T13:35:56.088922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Store submission\nsub_df.to_csv('submission.csv',index=False)\nsub_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-09-25T13:35:56.091463Z","iopub.execute_input":"2023-09-25T13:35:56.091798Z","iopub.status.idle":"2023-09-25T13:35:56.123654Z","shell.execute_reply.started":"2023-09-25T13:35:56.091771Z","shell.execute_reply":"2023-09-25T13:35:56.122626Z"},"trusted":true},"execution_count":null,"outputs":[]}]}