{"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":"# 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-09-21T19:26:51.304378Z","iopub.execute_input":"2022-09-21T19:26:51.304748Z","iopub.status.idle":"2022-09-21T19:27:01.897953Z","shell.execute_reply.started":"2022-09-21T19:26:51.304659Z","shell.execute_reply":"2022-09-21T19:27:01.896779Z"},"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-09-21T19:27:01.900193Z","iopub.execute_input":"2022-09-21T19:27:01.900521Z","iopub.status.idle":"2022-09-21T19:28:25.309355Z","shell.execute_reply.started":"2022-09-21T19:27:01.900453Z","shell.execute_reply":"2022-09-21T19:28:25.30808Z"},"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-09-21T19:28:25.311851Z","iopub.execute_input":"2022-09-21T19:28:25.312141Z","iopub.status.idle":"2022-09-21T19:28:30.117291Z","shell.execute_reply.started":"2022-09-21T19:28:25.312104Z","shell.execute_reply":"2022-09-21T19:28:30.116275Z"},"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-09-21T19:28:30.33911Z","iopub.status.idle":"2022-09-21T19:28:30.339472Z","shell.execute_reply.started":"2022-09-21T19:28:30.339287Z","shell.execute_reply":"2022-09-21T19:28:30.339304Z"},"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\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-09-21T19:28:30.340786Z","iopub.status.idle":"2022-09-21T19:28:30.341404Z","shell.execute_reply.started":"2022-09-21T19:28:30.341146Z","shell.execute_reply":"2022-09-21T19:28:30.341175Z"},"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\")\ntrain_bbox = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train_bounding_boxes.csv\")\ntest_df = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/test.csv\")\nss = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/sample_submission.csv\")\n\n# Print dataframe shapes\nprint('train shape:', train_df.shape)\nprint('train bbox shape:', train_bbox.shape)\nprint('test shape:', test_df.shape)\nprint('ss shape:', ss.shape)\nprint('')\n\n# Show first few entries\ntrain_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.344146Z","iopub.status.idle":"2022-09-21T19:28:30.344498Z","shell.execute_reply.started":"2022-09-21T19:28:30.344324Z","shell.execute_reply":"2022-09-21T19:28:30.344347Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Debug","metadata":{}},{"cell_type":"code","source":"if len(ss)==3:\n    # Fix mismatch with test_images folder\n    test_df = pd.DataFrame(columns = ['row_id','StudyInstanceUID','prediction_type'])\n    for i in ['1.2.826.0.1.3680043.22327','1.2.826.0.1.3680043.25399','1.2.826.0.1.3680043.5876']:\n        for j in ['C1','C2','C3','C4','C5','C6','C7','patient_overall']:\n            test_df = test_df.append({'row_id':i+'_'+j,'StudyInstanceUID':i,'prediction_type':j},ignore_index=True)\n    \n    # Sample submission\n    ss = pd.DataFrame(test_df['row_id'])\n    ss['fractured'] = 0.5\n    \n    display(test_df.head(3))","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.346221Z","iopub.status.idle":"2022-09-21T19:28:30.346562Z","shell.execute_reply.started":"2022-09-21T19:28:30.346379Z","shell.execute_reply":"2022-09-21T19:28:30.346395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Test table","metadata":{}},{"cell_type":"code","source":"# Test table\ntest_table = pd.DataFrame(pd.unique(test_df['StudyInstanceUID']), columns=['StudyInstanceUID'])\ntest_table.head()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.347727Z","iopub.status.idle":"2022-09-21T19:28:30.3481Z","shell.execute_reply.started":"2022-09-21T19:28:30.347904Z","shell.execute_reply":"2022-09-21T19:28:30.347921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Extract metadata","metadata":{}},{"cell_type":"code","source":"def get_observation_data(path):\n    '''\n    Get information from the .dcm files\n    '''\n    \n    dataset = pydicom.read_file(path)\n    \n    # Dictionary to store the information from the image\n    observation_data = {\n        \"SOPInstanceUID\" : dataset.get(\"SOPInstanceUID\"),\n        \"InstanceNumber\" : dataset.get(\"InstanceNumber\"),\n        \"SliceThickness\" : dataset.get(\"SliceThickness\"),\n        \"ImagePositionPatient\" : dataset.get(\"ImagePositionPatient\"),\n    }\n\n    return observation_data\n\ndef get_metadata():\n    '''\n    Retrieves the desired metadata from the .dcm files and saves it into dataframe.\n    '''\n    \n    dicts = []\n    \n    for k in tqdm(range(len(test_table))):\n        patient = test_table.loc[k,'StudyInstanceUID']\n\n        # Get all .dcm paths for this Instance\n        dcm_paths = glob(f\"../input/rsna-2022-cervical-spine-fracture-detection/test_images/{patient}/*\")\n\n        for path in dcm_paths:\n            # Get datasets\n            dataset = get_observation_data(path)\n            dicts.append(dataset)\n    \n    return pd.DataFrame(data=dicts, columns=md_example.keys())","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.34912Z","iopub.status.idle":"2022-09-21T19:28:30.349438Z","shell.execute_reply.started":"2022-09-21T19:28:30.349273Z","shell.execute_reply":"2022-09-21T19:28:30.349289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Example\nex_path = \"../input/rsna-2022-cervical-spine-fracture-detection/train_images/1.2.826.0.1.3680043.10001/101.dcm\"\nmd_example = get_observation_data(ex_path)\npprint(md_example)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.350365Z","iopub.status.idle":"2022-09-21T19:28:30.350669Z","shell.execute_reply.started":"2022-09-21T19:28:30.350512Z","shell.execute_reply":"2022-09-21T19:28:30.350527Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Get metadata\ntest_meta = get_metadata()\ntest_meta.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.353546Z","iopub.status.idle":"2022-09-21T19:28:30.354048Z","shell.execute_reply.started":"2022-09-21T19:28:30.353775Z","shell.execute_reply":"2022-09-21T19:28:30.3538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Clean metadata","metadata":{}},{"cell_type":"code","source":"# Change data types\ntest_meta['SOPInstanceUID'] = test_meta['SOPInstanceUID'].astype('str')\ntest_meta['InstanceNumber'] = test_meta['InstanceNumber'].astype('int32')\ntest_meta['SliceThickness'] = test_meta['SliceThickness'].astype('float32')\ntest_meta['ImagePositionPatient'] = test_meta['ImagePositionPatient'].astype('str')\n\n# Patient id\ntest_meta[\"StudyInstanceUID\"] = test_meta[\"SOPInstanceUID\"].apply(lambda x: \".\".join(x.split(\".\")[:-2]))\n\n# Extract x, y, z coordinates of position vector\ntest_meta['ImagePositionPatient_x'] = test_meta['ImagePositionPatient'].apply(lambda x: float(x.replace(',','').replace(']','').replace('[','').split()[0]))\ntest_meta['ImagePositionPatient_y'] = test_meta['ImagePositionPatient'].apply(lambda x: float(x.replace(',','').replace(']','').replace('[','').split()[1]))\ntest_meta['ImagePositionPatient_z'] = test_meta['ImagePositionPatient'].apply(lambda x: float(x.replace(',','').replace(']','').replace('[','').split()[2]))\n\n# Clean metadata\ntest_meta.drop(['SOPInstanceUID','ImagePositionPatient'], axis=1, inplace=True)\ntest_meta.rename(columns={\"InstanceNumber\": \"Slice\"}, inplace=True)\ntest_meta = test_meta[['StudyInstanceUID','Slice','SliceThickness','ImagePositionPatient_x','ImagePositionPatient_y','ImagePositionPatient_z']]\ntest_meta.sort_values(by=['StudyInstanceUID','Slice'], inplace=True)\ntest_meta.reset_index(drop=True, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.355235Z","iopub.status.idle":"2022-09-21T19:28:30.355725Z","shell.execute_reply.started":"2022-09-21T19:28:30.355464Z","shell.execute_reply":"2022-09-21T19:28:30.35549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Feature engineering","metadata":{}},{"cell_type":"code","source":"# Calculate slice ratio\nslice_max = test_meta.groupby('StudyInstanceUID')['Slice'].max().to_dict()\ntest_meta['SliceRatio'] = 0\ntest_meta['SliceRatio'] = test_meta['Slice']/test_meta['StudyInstanceUID'].map(slice_max)\n\n# Number of slices\ntest_meta['SliceTotal'] = 0\ntest_meta['SliceTotal'] = test_meta['StudyInstanceUID'].map(slice_max)\n\n# Reversed indicator\nz_reversed = ((test_meta.groupby('StudyInstanceUID')['ImagePositionPatient_z'].first()-test_meta.groupby('StudyInstanceUID')['ImagePositionPatient_z'].last())<0).astype('int')\ntest_meta['Reversed'] = 0\ntest_meta['Reversed'] = test_meta['StudyInstanceUID'].map(z_reversed)\n\n# Preview\ntest_meta.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.356991Z","iopub.status.idle":"2022-09-21T19:28:30.357502Z","shell.execute_reply.started":"2022-09-21T19:28:30.357226Z","shell.execute_reply":"2022-09-21T19:28:30.357251Z"},"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-09-21T19:28:30.358663Z","iopub.status.idle":"2022-09-21T19:28:30.35907Z","shell.execute_reply.started":"2022-09-21T19:28:30.358821Z","shell.execute_reply":"2022-09-21T19:28:30.358841Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Torch dataset","metadata":{}},{"cell_type":"code","source":"# Dataset for test set only\nclass RSNADataset_test(Dataset):\n    # Initialise\n    def __init__(self, df_table = test_meta, transform=None):\n        super().__init__()\n        self.df_table = df_table.reset_index(drop=True)\n        self.transform = transform\n        self.test_path = '../input/rsna-2022-cervical-spine-fracture-detection/test_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.test_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        # Metadata\n        meta = torch.as_tensor(self.df_table.iloc[index][self.meta_cols].astype('float32').values)\n        \n        return img, meta\n\n    # Length of dataset\n    def __len__(self):\n        return len(self.df_table)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.360302Z","iopub.status.idle":"2022-09-21T19:28:30.360718Z","shell.execute_reply.started":"2022-09-21T19:28:30.360541Z","shell.execute_reply":"2022-09-21T19:28:30.360559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = RSNADataset_test(df_table = test_meta, transform=A.Resize(*IMG_SIZE, interpolation=cv2.INTER_LINEAR))","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.36244Z","iopub.status.idle":"2022-09-21T19:28:30.362946Z","shell.execute_reply.started":"2022-09-21T19:28:30.362679Z","shell.execute_reply":"2022-09-21T19:28:30.362705Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Torch dataloaders","metadata":{}},{"cell_type":"code","source":"# Dataloader\ntest_loader = DataLoader(dataset=test_dataset, batch_size=BATCH_SIZE, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.364198Z","iopub.status.idle":"2022-09-21T19:28:30.366039Z","shell.execute_reply.started":"2022-09-21T19:28:30.365727Z","shell.execute_reply":"2022-09-21T19:28:30.365757Z"},"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-09-21T19:28:30.367135Z","iopub.status.idle":"2022-09-21T19:28:30.367472Z","shell.execute_reply.started":"2022-09-21T19:28:30.367307Z","shell.execute_reply":"2022-09-21T19:28:30.367323Z"},"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-09-21T19:28:30.368474Z","iopub.status.idle":"2022-09-21T19:28:30.368793Z","shell.execute_reply.started":"2022-09-21T19:28:30.368627Z","shell.execute_reply":"2022-09-21T19:28:30.368643Z"},"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-09-21T19:28:30.369655Z","iopub.status.idle":"2022-09-21T19:28:30.369982Z","shell.execute_reply.started":"2022-09-21T19:28:30.3698Z","shell.execute_reply":"2022-09-21T19:28:30.369815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Make predictions","metadata":{}},{"cell_type":"code","source":"def get_predictions(model, data_loader):\n    '''Make model predictions'''\n    with torch.no_grad():\n        predictions = []\n        for idx, (img, meta) in enumerate(tqdm(data_loader)):\n            y1, y2 = model.predict(img.to(device), meta.to(device))\n            pred = torch.cat([y1, y2], dim=1)\n            predictions.append(pred)\n        return torch.cat(predictions).cpu().numpy()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.371314Z","iopub.status.idle":"2022-09-21T19:28:30.371642Z","shell.execute_reply.started":"2022-09-21T19:28:30.371472Z","shell.execute_reply":"2022-09-21T19:28:30.371487Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds = get_predictions(model, test_loader)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.373279Z","iopub.status.idle":"2022-09-21T19:28:30.373786Z","shell.execute_reply.started":"2022-09-21T19:28:30.373531Z","shell.execute_reply":"2022-09-21T19:28:30.373555Z"},"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([test_meta.reset_index(drop=True).loc[:,['StudyInstanceUID','Slice']], df_preds], axis=1)\ndf_preds_concat.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.375236Z","iopub.status.idle":"2022-09-21T19:28:30.37573Z","shell.execute_reply.started":"2022-09-21T19:28:30.375472Z","shell.execute_reply":"2022-09-21T19:28:30.375496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualise predictions","metadata":{}},{"cell_type":"code","source":"def plot_sample_patient(df_pred):\n    patient = '1.2.826.0.1.3680043.22327'\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)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.377424Z","iopub.status.idle":"2022-09-21T19:28:30.377906Z","shell.execute_reply.started":"2022-09-21T19:28:30.377645Z","shell.execute_reply":"2022-09-21T19:28:30.37767Z"},"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-21T19:28:30.380244Z","iopub.status.idle":"2022-09-21T19:28:30.380711Z","shell.execute_reply.started":"2022-09-21T19:28:30.380459Z","shell.execute_reply":"2022-09-21T19:28:30.380482Z"},"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()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.381642Z","iopub.status.idle":"2022-09-21T19:28:30.382146Z","shell.execute_reply.started":"2022-09-21T19:28:30.381864Z","shell.execute_reply":"2022-09-21T19:28:30.381889Z"},"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_count":null,"outputs":[]},{"cell_type":"code","source":"BETA = 0.8\n#df_patient_pred['patient_overall'] = df_patient_pred['patient_overall'].apply(lambda x : squish(x, beta=BETA))\n#df_patient_pred.head()","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"# Melt table\npred_melt = df_patient_pred.reset_index().melt(id_vars='StudyInstanceUID', value_vars=['C1','C2','C3','C4','C5','C6','C7','patient_overall'], value_name='fractured', var_name='prediction_type')\n\n# Merge predictions\nsubmission_df = test_df.merge(pred_melt, how='inner', on=['StudyInstanceUID','prediction_type'])\n\n# Save to csv\nsubmission = submission_df[['row_id','fractured']]\nsubmission.to_csv('submission.csv', index=False)\nsubmission.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T19:28:30.38405Z","iopub.status.idle":"2022-09-21T19:28:30.384385Z","shell.execute_reply.started":"2022-09-21T19:28:30.384221Z","shell.execute_reply":"2022-09-21T19:28:30.384237Z"},"trusted":true},"execution_count":null,"outputs":[]}]}