{"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":"### Model Architecture\n\n<center>\n<img src='https://i.postimg.cc/ZYV9RwXt/2-D-model-architecture.png' width=800>\n</center>\n\n<br>\n\n### Links\n\n1. [RSNA Fracture Detection - in-depth EDA](https://www.kaggle.com/code/samuelcortinhas/rsna-fracture-detection-in-depth-eda)\n2. [RSNA 2022 Spine Fracture Detection - Metadata](https://www.kaggle.com/datasets/samuelcortinhas/rsna-2022-spine-fracture-detection-metadata)\n3. [RNSA - 2D model [Train] [PyTorch]](https://www.kaggle.com/code/samuelcortinhas/rnsa-2d-model-train-pytorch)\n4. [RNSA - 2D model [Re-Train] [PyTorch]](https://www.kaggle.com/code/samuelcortinhas/rnsa-2d-model-re-train-pytorch)\n5. [RSNA - Trained 2D models [PyTorch]](https://www.kaggle.com/datasets/samuelcortinhas/rsna-trained-2d-models-pytorch)\n6. [RNSA - 2D model [Validate] [PyTorch]](https://www.kaggle.com/code/samuelcortinhas/rnsa-2d-model-validate-pytorch)\n7. [RNSA - 2D model [Infer] [PyTorch]](https://www.kaggle.com/code/samuelcortinhas/rnsa-2d-model-infer-pytorch)\n\n<hr>\n\nMuch of my work is inspired by https://www.kaggle.com/code/vslaykovsky/train-pytorch-effnetv2-baseline-cv-0-49","metadata":{}},{"cell_type":"markdown","source":"The purpose of this notebook is to evaluate a trained model against the competition metric. It is a separate notebook because of run-time/GPU constraints.","metadata":{}},{"cell_type":"markdown","source":"# Libraries","metadata":{}},{"cell_type":"code","source":"!pip install -qU ../input/for-pydicom/python_gdcm-3.0.14-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl ../input/for-pydicom/pylibjpeg-1.4.0-py3-none-any.whl --find-links frozen_packages --no-index","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-04T18:22:36.616974Z","iopub.execute_input":"2022-10-04T18:22:36.617562Z","iopub.status.idle":"2022-10-04T18:22:49.992411Z","shell.execute_reply.started":"2022-10-04T18:22:36.617452Z","shell.execute_reply":"2022-10-04T18:22:49.991047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints/\n!cp ../input/rsna-2022-whl/efficientnet_v2_s-dd5fe13b.pth  /root/.cache/torch/hub/checkpoints/\n!pip install /kaggle/input/rsna-2022-whl/{pydicom-2.3.0-py3-none-any.whl,pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}\n!pip install /kaggle/input/rsna-2022-whl/{torch-1.12.1-cp37-cp37m-manylinux1_x86_64.whl,torchvision-0.13.1-cp37-cp37m-manylinux1_x86_64.whl}","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-04T18:22:49.994923Z","iopub.execute_input":"2022-10-04T18:22:49.995764Z","iopub.status.idle":"2022-10-04T18:24:16.052643Z","shell.execute_reply.started":"2022-10-04T18:22:49.995721Z","shell.execute_reply":"2022-10-04T18:24:16.051559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n%matplotlib inline\nimport matplotlib.patches as patches\nimport seaborn as sns\nsns.set(style='darkgrid', font_scale=1.6)\nimport cv2\nimport os\nfrom os import listdir\nimport re\nimport gc\nimport random\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nfrom tqdm.auto import tqdm\nfrom pprint import pprint\nfrom time import time\nimport itertools\nfrom skimage import measure\nfrom mpl_toolkits.mplot3d.art3d import Poly3DCollection\nimport nibabel as nib\nfrom glob import glob\nimport warnings\n#warnings.filterwarnings(\"ignore\", category=DeprecationWarning)\n#warnings.filterwarnings(\"ignore\", category=UserWarning)\n#warnings.filterwarnings(\"ignore\", category=FutureWarning)\nimport zipfile\nfrom scipy import ndimage\nfrom sklearn.model_selection import train_test_split, GroupKFold\nfrom joblib import Parallel, delayed\nfrom PIL import Image\nfrom dipy.denoise.nlmeans import nlmeans\nfrom dipy.denoise.noise_estimate import estimate_sigma\n#from kaggle_volclassif.utils import interpolate_volume\nfrom skimage import exposure\n\n# Pytorch\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision as tv\nimport torchvision.transforms as transforms\nimport torch.nn.functional as F\nimport kornia\nimport kornia.augmentation as augmentation\nimport albumentations as A\n\nfrom torch.cuda.amp import GradScaler, autocast\nfrom torchvision.models.feature_extraction import create_feature_extractor\n#from tqdm.notebook import tqdm\n#import wandb","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-04T18:24:16.054833Z","iopub.execute_input":"2022-10-04T18:24:16.055528Z","iopub.status.idle":"2022-10-04T18:24:19.365237Z","shell.execute_reply.started":"2022-10-04T18:24:16.055487Z","shell.execute_reply":"2022-10-04T18:24:19.363941Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Reproducibility","metadata":{}},{"cell_type":"code","source":"# Set random seeds\ndef set_seed(seed=0):\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\nset_seed()","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:19.368593Z","iopub.execute_input":"2022-10-04T18:24:19.36941Z","iopub.status.idle":"2022-10-04T18:24:19.376176Z","shell.execute_reply.started":"2022-10-04T18:24:19.369361Z","shell.execute_reply":"2022-10-04T18:24:19.374942Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"# Hyperparameters\nif torch.cuda.is_available():\n    BATCH_SIZE = 32\nelse:\n    BATCH_SIZE = 4\n\nEXPERIMENTAL = False\nN_FOLDS = 5\nFOLD = 0     # 0,1,2,3,4\nIMG_SIZE = (256,256)\n\n# Config device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:19.378345Z","iopub.execute_input":"2022-10-04T18:24:19.379183Z","iopub.status.idle":"2022-10-04T18:24:19.403513Z","shell.execute_reply.started":"2022-10-04T18:24:19.379115Z","shell.execute_reply":"2022-10-04T18:24:19.402327Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data","metadata":{}},{"cell_type":"markdown","source":"### Load tables","metadata":{}},{"cell_type":"code","source":"# Load tables\ntrain_df = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n\n# Print dataframe shapes\nprint('train shape:', train_df.shape)\n\n# Show first few entries\ntrain_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:19.405301Z","iopub.execute_input":"2022-10-04T18:24:19.407663Z","iopub.status.idle":"2022-10-04T18:24:19.450499Z","shell.execute_reply.started":"2022-10-04T18:24:19.407612Z","shell.execute_reply":"2022-10-04T18:24:19.449319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Slice dataframe with vertebrae predictions","metadata":{}},{"cell_type":"code","source":"# rsna-2022-spine-fracture-detection-metadata contains inference of C1-C7 vertebrae for all training sample\ntrain_df_slices = pd.read_csv('../input/rsna-2022-spine-fracture-detection-metadata/train_segmented.csv')\nc1c7 = [f'C{i}' for i in range(1,8)]\ntrain_df_slices[c1c7] = (train_df_slices[c1c7] > 0.5).astype(int)\n\n# Merge dfs\ntrain_df_vert = train_df_slices.set_index('StudyInstanceUID').join(train_df.set_index('StudyInstanceUID'), rsuffix='_fracture').reset_index().copy()\ntrain_df_vert.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:19.452077Z","iopub.execute_input":"2022-10-04T18:24:19.453329Z","iopub.status.idle":"2022-10-04T18:24:22.906032Z","shell.execute_reply.started":"2022-10-04T18:24:19.453278Z","shell.execute_reply":"2022-10-04T18:24:22.904929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Drop bad scans","metadata":{}},{"cell_type":"code","source":"# https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/344862\nbad_scans = ['1.2.826.0.1.3680043.20574','1.2.826.0.1.3680043.29952']\n\nfor uid in bad_scans:\n    train_df.drop(train_df[train_df['StudyInstanceUID']==uid].index, axis=0, inplace=True)\n    train_df_vert.drop(train_df_vert[train_df_vert['StudyInstanceUID']==uid].index, axis=0, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:22.907302Z","iopub.execute_input":"2022-10-04T18:24:22.907632Z","iopub.status.idle":"2022-10-04T18:24:23.248889Z","shell.execute_reply.started":"2022-10-04T18:24:22.907601Z","shell.execute_reply":"2022-10-04T18:24:23.247622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Feature engineering","metadata":{}},{"cell_type":"code","source":"# Calculate slice ratio\nslice_max = train_df_vert.groupby('StudyInstanceUID')['Slice'].max().to_dict()\ntrain_df_vert['SliceRatio'] = 0\ntrain_df_vert['SliceRatio'] = train_df_vert['Slice']/train_df_vert['StudyInstanceUID'].map(slice_max)\n\n# Number of slices\ntrain_df_vert['SliceTotal'] = 0\ntrain_df_vert['SliceTotal'] = train_df_vert['StudyInstanceUID'].map(slice_max)\n\n# Reversed indicator\nz_reversed = ((train_df_vert.groupby('StudyInstanceUID')['ImagePositionPatient_z'].first()-train_df_vert.groupby('StudyInstanceUID')['ImagePositionPatient_z'].last())<0).astype('int')\ntrain_df_vert['Reversed'] = 0\ntrain_df_vert['Reversed'] = train_df_vert['StudyInstanceUID'].map(z_reversed)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-04T18:24:23.250389Z","iopub.execute_input":"2022-10-04T18:24:23.250771Z","iopub.status.idle":"2022-10-04T18:24:23.595477Z","shell.execute_reply.started":"2022-10-04T18:24:23.250738Z","shell.execute_reply":"2022-10-04T18:24:23.594487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{}},{"cell_type":"code","source":"def load_dicom(path):\n    \"\"\"\n    This supports loading both regular and compressed JPEG images.\n    See the first sell with `pip install` commands for the necessary dependencies\n    \"\"\"\n    img = pydicom.dcmread(path)\n    img.PhotometricInterpretation = 'YBR_FULL'\n    data = img.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return cv2.cvtColor(data, cv2.COLOR_GRAY2RGB), img","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:23.598854Z","iopub.execute_input":"2022-10-04T18:24:23.599231Z","iopub.status.idle":"2022-10-04T18:24:23.605932Z","shell.execute_reply.started":"2022-10-04T18:24:23.599195Z","shell.execute_reply":"2022-10-04T18:24:23.604951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Torch dataset","metadata":{}},{"cell_type":"code","source":"# Dataset for train/valid sets only\nclass RSNADataset(Dataset):\n    # Initialise\n    def __init__(self, df_table = train_df_vert, transform=None):\n        super().__init__()\n        self.df_table = df_table.reset_index(drop=True)\n        self.transform = transform\n        self.train_path = '../input/rsna-2022-cervical-spine-fracture-detection/train_images'\n        self.meta_cols = ['SliceRatio','SliceTotal','SliceThickness','ImagePositionPatient_x','ImagePositionPatient_y','ImagePositionPatient_z','Reversed']\n        \n    # Get item in position given by index\n    def __getitem__(self, index):\n        # Load image\n        path = os.path.join(self.train_path, self.df_table.iloc[index].StudyInstanceUID, f'{self.df_table.iloc[index].Slice}.dcm')\n        img = load_dicom(path)[0]\n        \n        # Data augmentations\n        if self.transform is not None:\n            img = self.transform(image=img)['image']\n        \n        # Pytorch uses (batch, channel, height, width) order. Converting (height, width, channel) -> (channel, height, width)\n        img = np.transpose(img, (2, 0, 1))\n        \n        # Convert to tensor\n        img = torch.from_numpy(img.astype('float32'))\n        \n        # Targets\n        frac_targets = torch.as_tensor(self.df_table.iloc[index][[f'C{i}_fracture' for i in range(1,8)]].astype('float32').values)\n        vert_targets = torch.as_tensor(self.df_table.iloc[index][[f'C{i}' for i in range(1,8)]].astype('float32').values)\n        frac_targets = frac_targets * vert_targets  # we only enable targets that are visible on the current slice\n        \n        # Metadata\n        meta = torch.as_tensor(self.df_table.iloc[index][self.meta_cols].astype('float32').values)\n        \n        return img, meta, frac_targets, vert_targets\n\n    # Length of dataset\n    def __len__(self):\n        return len(self.df_table)","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:23.60728Z","iopub.execute_input":"2022-10-04T18:24:23.607636Z","iopub.status.idle":"2022-10-04T18:24:23.621495Z","shell.execute_reply.started":"2022-10-04T18:24:23.607603Z","shell.execute_reply":"2022-10-04T18:24:23.619873Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Cross-Validation strategy","metadata":{}},{"cell_type":"code","source":"# Train/valid split (grouped by patients)\nif EXPERIMENTAL:\n    train_patients, valid_patients = train_test_split(train_df['StudyInstanceUID'].values, train_size=0.04, test_size=0.01, random_state=0)\n    \n    # Select train/valid tables\n    train_table = train_df_vert[train_df_vert['StudyInstanceUID'].isin(train_patients)]\n    valid_table = train_df_vert[train_df_vert['StudyInstanceUID'].isin(valid_patients)]\n    \n    # Define train/valid dataset\n    train_dataset = RSNADataset(df_table = train_table, transform=A.Resize(*IMG_SIZE, interpolation=cv2.INTER_LINEAR))\n    valid_dataset = RSNADataset(df_table = valid_table, transform=A.Resize(*IMG_SIZE, interpolation=cv2.INTER_LINEAR))\nelse:\n    gkf = GroupKFold(N_FOLDS)\n    for k, (train_idx, valid_idx) in enumerate(gkf.split(train_df_vert, groups=train_df_vert.StudyInstanceUID)):\n        if k==FOLD:\n            # Select train/valid tables\n            train_table = train_df_vert.iloc[train_idx]\n            valid_table = train_df_vert.iloc[valid_idx]\n\n    # Define train/valid dataset\n    train_dataset = RSNADataset(df_table = train_table, transform=A.Resize(*IMG_SIZE, interpolation=cv2.INTER_LINEAR))\n    valid_dataset = RSNADataset(df_table = valid_table, transform=A.Resize(*IMG_SIZE, interpolation=cv2.INTER_LINEAR))","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:23.623576Z","iopub.execute_input":"2022-10-04T18:24:23.624185Z","iopub.status.idle":"2022-10-04T18:24:24.332188Z","shell.execute_reply.started":"2022-10-04T18:24:23.624138Z","shell.execute_reply":"2022-10-04T18:24:24.331281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Torch dataloaders","metadata":{}},{"cell_type":"code","source":"# Dataloaders\ntrain_loader = DataLoader(dataset=train_dataset, batch_size=BATCH_SIZE, shuffle=True)\nvalid_loader = DataLoader(dataset=valid_dataset, batch_size=BATCH_SIZE, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:24.333508Z","iopub.execute_input":"2022-10-04T18:24:24.334612Z","iopub.status.idle":"2022-10-04T18:24:24.340545Z","shell.execute_reply.started":"2022-10-04T18:24:24.334574Z","shell.execute_reply":"2022-10-04T18:24:24.339428Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class ImageMetaModel(torch.nn.Module):\n    def __init__(self):\n        super().__init__()\n        # Image model layers\n        effnet = tv.models.efficientnet_v2_s()\n        self.effnet = create_feature_extractor(effnet, ['flatten'])\n        self.img_linear = nn.Linear(1280, 64)\n        self.swish = nn.SiLU()\n        \n        # Metadata model layers\n        self.tab_linear1 = nn.Linear(7, 128)\n        self.tab_linear2 = nn.Linear(128, 64)\n        self.relu = nn.ReLU()\n        \n        # Combined model layers\n        self.drop = nn.Dropout(p=0.3)\n        self.linear1 = nn.Linear(128, 256)\n        self.linear2 = nn.Linear(256, 7)\n        self.linear3 = nn.Linear(256, 7)\n        \n        \n    # Forward pass\n    def forward(self, img, meta):\n        # Image model\n        x_img = self.effnet(img)['flatten']\n        x_img = self.img_linear(x_img)\n        x_img = self.swish(x_img)\n        \n        # Metadata model\n        x_meta = self.tab_linear1(meta)\n        x_meta = self.relu(x_meta)\n        x_meta = self.tab_linear2(x_meta)\n        x_meta = self.relu(x_meta)\n        \n        # Concatenate\n        x = torch.cat([x_img, x_meta], dim=1)\n        \n        # Combined model\n        x = self.drop(x)\n        x = self.linear1(x)\n        #x = self.relu(x)\n        \n        # Split\n        x_frac = self.linear2(x)\n        x_vert = self.linear3(x)\n        \n        return x_frac, x_vert\n\n    def predict(self, img, meta):\n        frac, vert = self.forward(img, meta)\n        return torch.sigmoid(frac), torch.sigmoid(vert)\n\nmodel = ImageMetaModel().to(device)","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:24.342107Z","iopub.execute_input":"2022-10-04T18:24:24.342551Z","iopub.status.idle":"2022-10-04T18:24:25.073071Z","shell.execute_reply.started":"2022-10-04T18:24:24.342516Z","shell.execute_reply":"2022-10-04T18:24:25.071535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load model","metadata":{}},{"cell_type":"code","source":"# Load checkpoint\nPATH='../input/rsna-trained-2d-models-pytorch/EffNet_model_fold0_epoch6g.pt'\ncheckpoint = torch.load(PATH, map_location=torch.device(device))\n\n# Load states\nmodel.load_state_dict(checkpoint['model_state_dict'])\nepoch = checkpoint['epoch']\nloss = checkpoint['loss']\n\n# Evaluation mode\nmodel.eval()\nprint('')","metadata":{"_kg_hide-output":false,"execution":{"iopub.status.busy":"2022-10-04T18:24:25.074806Z","iopub.execute_input":"2022-10-04T18:24:25.075512Z","iopub.status.idle":"2022-10-04T18:24:27.933294Z","shell.execute_reply.started":"2022-10-04T18:24:25.075466Z","shell.execute_reply":"2022-10-04T18:24:27.932185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Print final loss and epoch\nprint('Final epoch:', epoch)\nprint('Final loss:', loss)","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:27.934554Z","iopub.execute_input":"2022-10-04T18:24:27.934892Z","iopub.status.idle":"2022-10-04T18:24:27.94226Z","shell.execute_reply.started":"2022-10-04T18:24:27.934861Z","shell.execute_reply":"2022-10-04T18:24:27.941096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss","metadata":{}},{"cell_type":"code","source":"#loss_fn = nn.BCEWithLogitsLoss(reduction='none')\nloss_fn = nn.BCELoss(reduction=\"none\")\n\ncompetition_weights = {\n    '-' : torch.tensor([1, 1, 1, 1, 1, 1, 1], dtype=torch.float, device=device),\n    '+' : torch.tensor([2, 2, 2, 2, 2, 2, 2], dtype=torch.float, device=device),\n}\n\n# y_hat.shape = (batch_size, num_classes)\n# y.shape = (batch_size, num_classes)\n\ndef FracLoss(y_hat, y):\n    '''Weighted Multilabel Log Loss'''\n    loss = loss_fn(y_hat, y.to(y_hat.dtype))\n    weights = y * competition_weights['+'] + (1 - y) * competition_weights['-']\n    return (loss * weights).sum(axis=1)\n    \ndef VertLoss(y_hat, y):\n    '''BCE Loss'''\n    loss = loss_fn(y_hat, y.to(y_hat.dtype))\n    return loss.sum(axis=1)\n\ndef CombinedLoss(y_hat_frac, y_frac, y_hat_vert, y_vert):\n    # Loss for fracture detection\n    L1 = FracLoss(y_hat_frac, y_frac)\n    \n    # Loss for vertebrae detection\n    L2 = VertLoss(y_hat_vert, y_vert)\n    \n    return (L1+L2).mean()","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:24:27.943902Z","iopub.execute_input":"2022-10-04T18:24:27.944218Z","iopub.status.idle":"2022-10-04T18:24:27.957283Z","shell.execute_reply.started":"2022-10-04T18:24:27.944189Z","shell.execute_reply":"2022-10-04T18:24:27.955969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make predictions\n\nPredictions on validation set.","metadata":{}},{"cell_type":"code","source":"def get_predictions(model, data_loader):\n    '''Make model predictions'''\n    val_loss = 0\n    with torch.no_grad():\n        predictions = []\n        for idx, (img, meta, frac_targets, vert_targets) in enumerate(tqdm(data_loader)):\n            y1, y2 = model.predict(img.to(device), meta.to(device))\n            \n            # Validation loss (different to comp. metric)\n            L = CombinedLoss(y1, frac_targets.to(device), y2, vert_targets.to(device))\n            val_loss += L.detach().item()\n            \n            # Save predictions\n            pred = torch.cat([y1, y2], dim=1)\n            predictions.append(pred)\n            \n            # Print iteration\n            if (idx%1000)==0:\n                print(f'Minibatch iteration {idx}/{len(data_loader)}')\n        \n        print('Combined validation loss:', val_loss/len(data_loader))\n        \n        return torch.cat(predictions).cpu().numpy()","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-04T18:27:07.566122Z","iopub.execute_input":"2022-10-04T18:27:07.566679Z","iopub.status.idle":"2022-10-04T18:27:07.578551Z","shell.execute_reply.started":"2022-10-04T18:27:07.566621Z","shell.execute_reply":"2022-10-04T18:27:07.577151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = get_predictions(model, valid_loader)","metadata":{"execution":{"iopub.status.busy":"2022-10-04T18:27:13.101228Z","iopub.execute_input":"2022-10-04T18:27:13.102458Z","iopub.status.idle":"2022-10-04T18:28:17.533521Z","shell.execute_reply.started":"2022-10-04T18:27:13.102406Z","shell.execute_reply":"2022-10-04T18:28:17.531403Z"},"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_preds = pd.DataFrame(data=preds, columns=[f'C{i}_frac' for i in range(1, 8)] + [f'C{i}_vert' for i in range(1, 8)])\ndf_preds_concat = pd.concat([valid_table.reset_index(drop=True).loc[:,['StudyInstanceUID','Slice']], df_preds], axis=1)\ndf_preds_concat.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-10-03T20:04:42.277625Z","iopub.execute_input":"2022-10-03T20:04:42.278089Z","iopub.status.idle":"2022-10-03T20:04:42.325938Z","shell.execute_reply.started":"2022-10-03T20:04:42.278004Z","shell.execute_reply":"2022-10-03T20:04:42.324566Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualise predictions","metadata":{}},{"cell_type":"code","source":"patient = np.random.choice(valid_table.query('patient_overall > 0').StudyInstanceUID)\ndef plot_sample_patient(df_pred, patient):\n    df = df_pred.query('StudyInstanceUID == @patient').reset_index(drop=True)\n    \n    plt.figure(figsize=(24,5))\n    df[[f'C{i}_frac' for i in range(1, 8)]].plot(\n        title=f'{patient}, fracture prediction',\n        ax=(plt.subplot(1, 2, 1)))\n    \n    df[[f'C{i}_vert' for i in range(1, 8)]].plot(\n        title=f'{patient}, vertebrae prediction',\n        ax=plt.subplot(1, 2, 2)\n    )\n\nplot_sample_patient(df_preds_concat, patient)","metadata":{"execution":{"iopub.status.busy":"2022-10-03T20:04:56.017552Z","iopub.execute_input":"2022-10-03T20:04:56.01798Z","iopub.status.idle":"2022-10-03T20:04:56.85108Z","shell.execute_reply.started":"2022-10-03T20:04:56.017946Z","shell.execute_reply":"2022-10-03T20:04:56.849542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# E.g. ground truth\ntrain_df[train_df['StudyInstanceUID']==patient]","metadata":{"execution":{"iopub.status.busy":"2022-09-21T22:29:41.022383Z","iopub.execute_input":"2022-09-21T22:29:41.023032Z","iopub.status.idle":"2022-09-21T22:29:41.039412Z","shell.execute_reply.started":"2022-09-21T22:29:41.022985Z","shell.execute_reply":"2022-09-21T22:29:41.037983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Patient overall","metadata":{}},{"cell_type":"code","source":"def patient_prediction(df):\n    c1c7 = np.average(df[[f'C{i}_frac' for i in range(1, 8)]].values, axis=0, weights=df[[f'C{i}_vert' for i in range(1, 8)]].values)\n    pred_patient_overall = 1 - np.prod(1 - c1c7)\n    return pd.Series(data=np.concatenate([[pred_patient_overall], c1c7]), index=['patient_overall'] + [f'C{i}' for i in range(1, 8)])","metadata":{"execution":{"iopub.status.busy":"2022-09-21T22:29:41.041089Z","iopub.execute_input":"2022-09-21T22:29:41.041459Z","iopub.status.idle":"2022-09-21T22:29:41.053306Z","shell.execute_reply.started":"2022-09-21T22:29:41.041427Z","shell.execute_reply":"2022-09-21T22:29:41.051913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_patient_pred = df_preds_concat.groupby('StudyInstanceUID').apply(lambda df: patient_prediction(df))\ndf_patient_pred.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T22:29:41.055093Z","iopub.execute_input":"2022-09-21T22:29:41.05607Z","iopub.status.idle":"2022-09-21T22:29:41.119846Z","shell.execute_reply.started":"2022-09-21T22:29:41.056032Z","shell.execute_reply":"2022-09-21T22:29:41.118578Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Evaluate with competition metric","metadata":{}},{"cell_type":"code","source":"# Replicate competition metric (https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/341854)\n#loss_fn = nn.BCEWithLogitsLoss(reduction='none')\nloss_fn = nn.BCELoss(reduction='none')\n\ncompetition_weights = {\n    '-' : torch.tensor([7, 1, 1, 1, 1, 1, 1, 1], dtype=torch.float, device=device),\n    '+' : torch.tensor([14, 2, 2, 2, 2, 2, 2, 2], dtype=torch.float, device=device),\n}\n\n# y_hat.shape = (batch_size, num_classes)\n# y.shape = (batch_size, num_classes)\n\n# with row-wise weights normalization (https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/344565)\ndef competiton_loss_row_norm(y_hat, y):\n    loss = loss_fn(y_hat, y.to(y_hat.dtype))\n    weights = y * competition_weights['+'] + (1 - y) * competition_weights['-']\n    loss = (loss * weights).sum(axis=1)\n    w_sum = weights.sum(axis=1)\n    loss = torch.div(loss, w_sum)\n    return loss.mean()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T22:36:10.9085Z","iopub.execute_input":"2022-09-21T22:36:10.908972Z","iopub.status.idle":"2022-09-21T22:36:10.918752Z","shell.execute_reply.started":"2022-09-21T22:36:10.908938Z","shell.execute_reply":"2022-09-21T22:36:10.917516Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"y_hat = df_patient_pred\ny = train_df.merge(df_patient_pred.reset_index()['StudyInstanceUID'], how='right', on='StudyInstanceUID')\nL = competiton_loss_row_norm(torch.from_numpy(y_hat.values).to(device), torch.from_numpy(y.iloc[:,1:].values).to(device)).item()\nprint(f'Competition loss on validation set predictions: {L:.6f}')","metadata":{"execution":{"iopub.status.busy":"2022-09-21T22:36:11.136783Z","iopub.execute_input":"2022-09-21T22:36:11.137239Z","iopub.status.idle":"2022-09-21T22:36:11.153841Z","shell.execute_reply.started":"2022-09-21T22:36:11.1372Z","shell.execute_reply":"2022-09-21T22:36:11.151979Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Post-processing","metadata":{}},{"cell_type":"code","source":"# https://stats.stackexchange.com/questions/214877/is-there-a-formula-for-an-s-shaped-curve-with-domain-and-range-0-1\ndef squish(x, beta=3):\n    return 1/(1+(x/(1-x))**(-beta))","metadata":{"execution":{"iopub.status.busy":"2022-09-21T22:36:13.282453Z","iopub.execute_input":"2022-09-21T22:36:13.282922Z","iopub.status.idle":"2022-09-21T22:36:13.289113Z","shell.execute_reply.started":"2022-09-21T22:36:13.28288Z","shell.execute_reply":"2022-09-21T22:36:13.287763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"xx = np.linspace(0.01,0.99,100)\nyy = squish(xx, beta=0.7)\nplt.plot(xx, yy)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T22:36:13.983046Z","iopub.execute_input":"2022-09-21T22:36:13.983468Z","iopub.status.idle":"2022-09-21T22:36:14.168122Z","shell.execute_reply.started":"2022-09-21T22:36:13.983434Z","shell.execute_reply":"2022-09-21T22:36:14.16687Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"L_beta = []\nxx = np.linspace(0.1,3,300)\nfor BETA in xx:\n    y_hat_copy = y_hat.copy()\n    y_hat_copy['patient_overall'] = y_hat_copy['patient_overall'].apply(lambda x : squish(x, beta=BETA))\n    L_beta.append(competiton_loss_row_norm(torch.from_numpy(y_hat_copy.values).to(device), torch.from_numpy(y.iloc[:,1:].values).to(device)).item())\n\nprint('Optimal value of beta:', xx[np.argmin(L_beta)])","metadata":{"execution":{"iopub.status.busy":"2022-09-21T22:36:16.52425Z","iopub.execute_input":"2022-09-21T22:36:16.524694Z","iopub.status.idle":"2022-09-21T22:36:16.845021Z","shell.execute_reply.started":"2022-09-21T22:36:16.524657Z","shell.execute_reply":"2022-09-21T22:36:16.843735Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(10,5))\nplt.plot(xx, L_beta)\nplt.xlabel('Beta')\nplt.ylabel('Comp. metric on valid set')\nplt.title('Beta vs comp. loss')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T22:36:17.894073Z","iopub.execute_input":"2022-09-21T22:36:17.89452Z","iopub.status.idle":"2022-09-21T22:36:18.102387Z","shell.execute_reply.started":"2022-09-21T22:36:17.894483Z","shell.execute_reply":"2022-09-21T22:36:18.100891Z"},"trusted":true},"execution_count":null,"outputs":[]}]}