{"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":"# 1. Setup","metadata":{}},{"cell_type":"markdown","source":"### Links","metadata":{}},{"cell_type":"markdown","source":"Jirka's links:\n* [Spine🦴Fracture: EDA🔎 & loading DICOM & 3D browse](https://www.kaggle.com/code/jirkaborovec/spine-fracture-eda-loading-dicom-3d-browse)\n* [Spine🦴Fracture: convert🤖 DICOM imgs -> 3D volume](https://www.kaggle.com/code/jirkaborovec/spine-fracture-convert-dicom-imgs-3d-volume)\n* [Spine🦴Fracture: convert🤖 DICOM -> equalized PNG](https://www.kaggle.com/code/jirkaborovec/spine-fracture-convert-dicom-equalized-png)\n* [SpineFrac🦴Classif: e2e ~ Lightning⚡MONAI⚕️3D](https://www.kaggle.com/code/jirkaborovec/spinefrac-classif-e2e-lightning-monai-3d/notebook)\n* [Cervical Spine Fracture Detection: 3D volumes](https://www.kaggle.com/datasets/jirkaborovec/cervical-spine-fracture-detection-npz-3d-volumes)\n* [Cervical Spine Fracture Detection: equalized PNG](https://www.kaggle.com/datasets/jirkaborovec/cervical-spine-fracture-detection-equalized-png)\n\nSam's links:\n\n* [🦴 RSNA Fracture Detection - in-depth EDA](https://www.kaggle.com/code/samuelcortinhas/rsna-fracture-detection-in-depth-eda)\n* [Extracting Vertebrae C1, ..., C7](https://www.kaggle.com/code/samuelcortinhas/extracting-vertebrae-c1-c7)\n* [RSNA - CT gifs](https://www.kaggle.com/code/samuelcortinhas/rsna-ct-gifs)\n* [RSNA 2022 Spine Fracture Detection - Metadata](https://www.kaggle.com/datasets/samuelcortinhas/rsna-2022-spine-fracture-detection-metadata)\n* [RSNA - 3D train tensors [first half]](https://www.kaggle.com/datasets/samuelcortinhas/rsna-3d-train-tensors-first-half)\n* [RSNA - 3D train tensors [second half]](https://www.kaggle.com/datasets/samuelcortinhas/rsna-3d-train-tensors-second-half)\n* [RNSA - 3D model [Train] [PyTorch]](https://www.kaggle.com/code/samuelcortinhas/rnsa-3d-model-train-pytorch)\n* [RSNA - Trained 3D model weights [PyTorch]](https://www.kaggle.com/datasets/samuelcortinhas/rsna-trained-3d-model-weights-pytorch)","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-11T11:35:01.443766Z","iopub.execute_input":"2022-09-11T11:35:01.444247Z","iopub.status.idle":"2022-09-11T11:35:15.644668Z","shell.execute_reply.started":"2022-09-11T11:35:01.444156Z","shell.execute_reply":"2022-09-11T11:35:15.643255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install -q kaggle_vol3d_classify -f ../input/cervical-spine-fracture-detection-npz-3d-volumes/frozen_packages --no-index\n# !pip install -qU \"pytorch-lightning>1.5.0\" --no-index\n#!pip uninstall -y torchtext\n#!pip list | grep -e lightning -e kaggle -e monai","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-09-11T11:35:15.647346Z","iopub.execute_input":"2022-09-11T11:35:15.64793Z","iopub.status.idle":"2022-09-11T11:35:31.54386Z","shell.execute_reply.started":"2022-09-11T11:35:15.64789Z","shell.execute_reply":"2022-09-11T11:35:31.542269Z"},"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\nfrom joblib import Parallel, delayed\nfrom PIL import Image\nfrom dipy.denoise.nlmeans import nlmeans\nfrom dipy.denoise.noise_estimate import estimate_sigma\nfrom 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.transforms as transforms\nimport torch.nn.functional as F","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-09-11T11:35:31.546071Z","iopub.execute_input":"2022-09-11T11:35:31.546499Z","iopub.status.idle":"2022-09-11T11:35:33.940767Z","shell.execute_reply.started":"2022-09-11T11:35:31.546458Z","shell.execute_reply":"2022-09-11T11:35:33.939777Z"},"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-11T11:35:33.945606Z","iopub.execute_input":"2022-09-11T11:35:33.946258Z","iopub.status.idle":"2022-09-11T11:35:33.954269Z","shell.execute_reply.started":"2022-09-11T11:35:33.946222Z","shell.execute_reply":"2022-09-11T11:35:33.953322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Config","metadata":{}},{"cell_type":"code","source":"# Hyperparameters\nBATCH_SIZE = 1\n\n# Config device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2022-09-11T11:35:33.955664Z","iopub.execute_input":"2022-09-11T11:35:33.955974Z","iopub.status.idle":"2022-09-11T11:35:33.969177Z","shell.execute_reply.started":"2022-09-11T11:35:33.955947Z","shell.execute_reply":"2022-09-11T11:35:33.968229Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2. Data","metadata":{}},{"cell_type":"markdown","source":"### Load tables","metadata":{}},{"cell_type":"code","source":"# Load metadata\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-11T11:35:33.970622Z","iopub.execute_input":"2022-09-11T11:35:33.971408Z","iopub.status.idle":"2022-09-11T11:35:34.046043Z","shell.execute_reply.started":"2022-09-11T11:35:33.971376Z","shell.execute_reply":"2022-09-11T11:35:34.044681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Debug","metadata":{}},{"cell_type":"code","source":"debug = False\nif len(ss)==3:\n    debug = True\n    \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-11T11:35:34.047385Z","iopub.execute_input":"2022-09-11T11:35:34.047732Z","iopub.status.idle":"2022-09-11T11:35:34.121448Z","shell.execute_reply.started":"2022-09-11T11:35:34.047703Z","shell.execute_reply":"2022-09-11T11:35:34.120297Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load volumes","metadata":{}},{"cell_type":"code","source":"# Convert dicom images to 3d tensor\ndef convert_volume(dir_path, out_dir = \"test_volumes\", size = (224, 224, 224)):\n    ls_imgs = glob(os.path.join(dir_path, \"*.dcm\"))\n    ls_imgs = sorted(ls_imgs, key=lambda p: int(os.path.splitext(os.path.basename(p))[0]))\n\n    imgs = []\n    for p_img in ls_imgs:\n        dicom = pydicom.dcmread(p_img)\n        img = apply_voi_lut(dicom.pixel_array, dicom)\n        img = cv2.resize(img, size[:2], interpolation=cv2.INTER_LINEAR)\n        imgs.append(img.tolist())\n    vol = torch.tensor(imgs, dtype=torch.float32)\n\n    vol = (vol - vol.min()) / float(vol.max() - vol.min())\n    vol = interpolate_volume(vol, size).numpy()\n    \n    # https://scikit-image.org/docs/stable/auto_examples/color_exposure/plot_adapt_hist_eq_3d.html\n    vol = exposure.equalize_adapthist(vol, kernel_size=np.array([64, 64, 64]), clip_limit=0.01)\n    # vol = exposure.equalize_hist(vol)\n    vol = np.clip(vol * 255, 0, 255).astype(np.uint8)\n    \n    path_pt = os.path.join(out_dir, f\"{os.path.basename(dir_path)}.pt\")\n    torch.save(torch.tensor(vol), path_pt)","metadata":{"execution":{"iopub.status.busy":"2022-09-11T11:35:34.122837Z","iopub.execute_input":"2022-09-11T11:35:34.12317Z","iopub.status.idle":"2022-09-11T11:35:34.13398Z","shell.execute_reply.started":"2022-09-11T11:35:34.123139Z","shell.execute_reply":"2022-09-11T11:35:34.132775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Make directory\nos.mkdir('/kaggle/working/test_volumes')\n\n# Get paths\nls_dirs = [p for p in glob(os.path.join(\"../input/rsna-2022-cervical-spine-fracture-detection\", \"test_images\", \"*\")) if os.path.isdir(p)]\nprint(f\"volumes: {len(ls_dirs)}\")\n\n# Convert volumes\n_= Parallel(n_jobs=3)(delayed(convert_volume)(p_dir, out_dir='/kaggle/working/test_volumes') for p_dir in tqdm(ls_dirs))","metadata":{"execution":{"iopub.status.busy":"2022-09-11T11:35:34.135457Z","iopub.execute_input":"2022-09-11T11:35:34.136105Z","iopub.status.idle":"2022-09-11T11:36:08.218367Z","shell.execute_reply.started":"2022-09-11T11:35:34.13607Z","shell.execute_reply":"2022-09-11T11:36:08.21684Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Torch dataset","metadata":{}},{"cell_type":"code","source":"# Dataset for test set only\nclass RSNADataset(Dataset):\n    # Initialise\n    def __init__(self, subset='test', df_table=test_df):\n        super().__init__()\n        \n        self.subset = subset\n        self.df_table = df_table\n        \n        # Image paths\n        self.volume_dir = '/kaggle/working/test_volumes/'\n        \n    # Get item in position given by index\n    def __getitem__(self, index):\n        \n        # load 3d volume\n        patient = self.df_table.loc[index,'StudyInstanceUID']\n        path = os.path.join(self.volume_dir, f'{patient}.pt')\n        vol = torch.load(path).to(torch.float32)\n        \n        return (vol.unsqueeze(0), patient)\n\n    # Length of dataset\n    def __len__(self):\n        return len(self.df_table)","metadata":{"execution":{"iopub.status.busy":"2022-09-11T11:36:08.222348Z","iopub.execute_input":"2022-09-11T11:36:08.222779Z","iopub.status.idle":"2022-09-11T11:36:08.232576Z","shell.execute_reply.started":"2022-09-11T11:36:08.222739Z","shell.execute_reply":"2022-09-11T11:36:08.230536Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Test dataset\ntest_table = pd.DataFrame(pd.unique(test_df['StudyInstanceUID']),columns=['StudyInstanceUID'])\ntest_dataset = RSNADataset(subset='test', df_table = test_table)","metadata":{"execution":{"iopub.status.busy":"2022-09-11T11:36:08.23659Z","iopub.execute_input":"2022-09-11T11:36:08.23703Z","iopub.status.idle":"2022-09-11T11:36:08.253376Z","shell.execute_reply.started":"2022-09-11T11:36:08.236995Z","shell.execute_reply":"2022-09-11T11:36:08.252507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Torch dataloader","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-11T11:36:08.254972Z","iopub.execute_input":"2022-09-11T11:36:08.255976Z","iopub.status.idle":"2022-09-11T11:36:08.265298Z","shell.execute_reply.started":"2022-09-11T11:36:08.255933Z","shell.execute_reply":"2022-09-11T11:36:08.263953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3. Model\n\nconv output size = floor((W-F+2P)/S + 1)","metadata":{}},{"cell_type":"code","source":"# 3D convolutional neural network\nclass Conv3DNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n        # Layers\n        self.conv1 = nn.Conv3d(in_channels=1, out_channels=16, kernel_size=7, stride=1, padding=0)\n        self.pool = nn.MaxPool3d(kernel_size=2, stride=2, padding=0)\n        self.norm1 = nn.BatchNorm3d(num_features=16)\n        self.conv2 = nn.Conv3d(in_channels=16, out_channels=32, kernel_size=3, stride=1, padding=0)\n        self.norm2 = nn.BatchNorm3d(num_features=32)\n        self.conv3 = nn.Conv3d(in_channels=32, out_channels=64, kernel_size=3, stride=1, padding=0)\n        self.norm3 = nn.BatchNorm3d(num_features=64)\n        self.avg = nn.AdaptiveAvgPool3d((7, 1, 1))\n        self.flat = nn.Flatten()\n        self.relu = nn.ReLU()\n        self.lin1 = nn.Linear(in_features=448, out_features=128)\n        self.lin2 = nn.Linear(in_features=128, out_features=8)\n        \n    def forward(self, x):\n        # Conv block 1\n        out = self.conv1(x)\n        out = self.relu(out)\n        out = self.pool(out)\n        out = self.norm1(out)\n        \n        # Conv block 2\n        out = self.conv2(out)\n        out = self.relu(out)\n        out = self.pool(out)\n        out = self.norm2(out)\n        \n        # Conv block 3\n        out = self.conv3(out)\n        out = self.relu(out)\n        out = self.pool(out)\n        out = self.norm3(out)\n        \n        # Average & flatten\n        out = self.avg(out)\n        out = self.flat(out)\n        \n        # Fully connected layer\n        out = self.lin1(out)\n        out = self.relu(out)\n        \n        # Output layer (no sigmoid needed)\n        out = self.lin2(out)\n        \n        return out\n\nmodel = Conv3DNet().to(device)","metadata":{"execution":{"iopub.status.busy":"2022-09-11T11:36:08.266894Z","iopub.execute_input":"2022-09-11T11:36:08.268013Z","iopub.status.idle":"2022-09-11T11:36:08.294172Z","shell.execute_reply.started":"2022-09-11T11:36:08.267969Z","shell.execute_reply":"2022-09-11T11:36:08.293258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import math\nfrom functools import partial\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\n\ndef get_inplanes():\n    return [64, 128, 256, 512]\n\n\ndef conv3x3x3(in_planes, out_planes, stride=1):\n    return nn.Conv3d(in_planes,\n                     out_planes,\n                     kernel_size=3,\n                     stride=stride,\n                     padding=1,\n                     bias=False)\n\n\ndef conv1x1x1(in_planes, out_planes, stride=1):\n    return nn.Conv3d(in_planes,\n                     out_planes,\n                     kernel_size=1,\n                     stride=stride,\n                     bias=False)\n\n\nclass BasicBlock(nn.Module):\n    expansion = 1\n\n    def __init__(self, in_planes, planes, stride=1, downsample=None):\n        super().__init__()\n\n        self.conv1 = conv3x3x3(in_planes, planes, stride)\n        self.bn1 = nn.BatchNorm3d(planes)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = conv3x3x3(planes, planes)\n        self.bn2 = nn.BatchNorm3d(planes)\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x):\n        residual = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n\n        if self.downsample is not None:\n            residual = self.downsample(x)\n\n        out += residual\n        out = self.relu(out)\n\n        return out\n\n\nclass Bottleneck(nn.Module):\n    expansion = 4\n\n    def __init__(self, in_planes, planes, stride=1, downsample=None):\n        super().__init__()\n\n        self.conv1 = conv1x1x1(in_planes, planes)\n        self.bn1 = nn.BatchNorm3d(planes)\n        self.conv2 = conv3x3x3(planes, planes, stride)\n        self.bn2 = nn.BatchNorm3d(planes)\n        self.conv3 = conv1x1x1(planes, planes * self.expansion)\n        self.bn3 = nn.BatchNorm3d(planes * self.expansion)\n        self.relu = nn.ReLU(inplace=True)\n        self.downsample = downsample\n        self.stride = stride\n\n    def forward(self, x):\n        residual = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n        out = self.relu(out)\n\n        out = self.conv3(out)\n        out = self.bn3(out)\n\n        if self.downsample is not None:\n            residual = self.downsample(x)\n\n        out += residual\n        out = self.relu(out)\n\n        return out\n\n\nclass ResNet(nn.Module):\n\n    def __init__(self,\n                 block,\n                 layers,\n                 block_inplanes,\n                 n_input_channels=1,\n                 conv1_t_size=7,\n                 conv1_t_stride=1,\n                 no_max_pool=False,\n                 shortcut_type='B',\n                 widen_factor=1.0,\n                 n_classes=8):\n        super().__init__()\n\n        block_inplanes = [int(x * widen_factor) for x in block_inplanes]\n\n        self.in_planes = block_inplanes[0]\n        self.no_max_pool = no_max_pool\n\n        self.conv1 = nn.Conv3d(n_input_channels,\n                               self.in_planes,\n                               kernel_size=(conv1_t_size, 7, 7),\n                               stride=(conv1_t_stride, 2, 2),\n                               padding=(conv1_t_size // 2, 3, 3),\n                               bias=False)\n        self.bn1 = nn.BatchNorm3d(self.in_planes)\n        self.relu = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool3d(kernel_size=3, stride=2, padding=1)\n        self.layer1 = self._make_layer(block, block_inplanes[0], layers[0],\n                                       shortcut_type)\n        self.layer2 = self._make_layer(block,\n                                       block_inplanes[1],\n                                       layers[1],\n                                       shortcut_type,\n                                       stride=2)\n        self.layer3 = self._make_layer(block,\n                                       block_inplanes[2],\n                                       layers[2],\n                                       shortcut_type,\n                                       stride=2)\n        self.layer4 = self._make_layer(block,\n                                       block_inplanes[3],\n                                       layers[3],\n                                       shortcut_type,\n                                       stride=2)\n\n        self.avgpool = nn.AdaptiveAvgPool3d((1, 1, 1))\n        self.fc = nn.Linear(block_inplanes[3] * block.expansion, n_classes)\n\n        for m in self.modules():\n            if isinstance(m, nn.Conv3d):\n                nn.init.kaiming_normal_(m.weight,\n                                        mode='fan_out',\n                                        nonlinearity='relu')\n            elif isinstance(m, nn.BatchNorm3d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n\n    def _downsample_basic_block(self, x, planes, stride):\n        out = F.avg_pool3d(x, kernel_size=1, stride=stride)\n        zero_pads = torch.zeros(out.size(0), planes - out.size(1), out.size(2),\n                                out.size(3), out.size(4))\n        if isinstance(out.data, torch.cuda.FloatTensor):\n            zero_pads = zero_pads.cuda()\n\n        out = torch.cat([out.data, zero_pads], dim=1)\n\n        return out\n\n    def _make_layer(self, block, planes, blocks, shortcut_type, stride=1):\n        downsample = None\n        if stride != 1 or self.in_planes != planes * block.expansion:\n            if shortcut_type == 'A':\n                downsample = partial(self._downsample_basic_block,\n                                     planes=planes * block.expansion,\n                                     stride=stride)\n            else:\n                downsample = nn.Sequential(\n                    conv1x1x1(self.in_planes, planes * block.expansion, stride),\n                    nn.BatchNorm3d(planes * block.expansion))\n\n        layers = []\n        layers.append(\n            block(in_planes=self.in_planes,\n                  planes=planes,\n                  stride=stride,\n                  downsample=downsample))\n        self.in_planes = planes * block.expansion\n        for i in range(1, blocks):\n            layers.append(block(self.in_planes, planes))\n\n        return nn.Sequential(*layers)\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        if not self.no_max_pool:\n            x = self.maxpool(x)\n\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n\n        x = self.avgpool(x)\n\n        x = x.view(x.size(0), -1)\n        x = self.fc(x)\n\n        return x\n    \ndef generate_model(model_depth, **kwargs):\n    assert model_depth in [10, 18, 34, 50, 101, 152, 200]\n\n    if model_depth == 10:\n        model = ResNet(BasicBlock, [1, 1, 1, 1], get_inplanes(), **kwargs)\n    elif model_depth == 18:\n        model = ResNet(BasicBlock, [2, 2, 2, 2], get_inplanes(), **kwargs)\n    elif model_depth == 34:\n        model = ResNet(BasicBlock, [3, 4, 6, 3], get_inplanes(), **kwargs)\n    elif model_depth == 50:\n        model = ResNet(Bottleneck, [3, 4, 6, 3], get_inplanes(), **kwargs)\n    elif model_depth == 101:\n        model = ResNet(Bottleneck, [3, 4, 23, 3], get_inplanes(), **kwargs)\n    elif model_depth == 152:\n        model = ResNet(Bottleneck, [3, 8, 36, 3], get_inplanes(), **kwargs)\n    elif model_depth == 200:\n        model = ResNet(Bottleneck, [3, 24, 36, 3], get_inplanes(), **kwargs)\n\n    return model\n\nmodel = generate_model(50).to(device)","metadata":{"execution":{"iopub.status.busy":"2022-09-11T11:36:32.571913Z","iopub.execute_input":"2022-09-11T11:36:32.572322Z","iopub.status.idle":"2022-09-11T11:36:33.423399Z","shell.execute_reply.started":"2022-09-11T11:36:32.572289Z","shell.execute_reply":"2022-09-11T11:36:33.422097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load model","metadata":{}},{"cell_type":"code","source":"# Load checkpoint\nPATH='../input/20220911/Conv3DNet.pt'\nif torch.cuda.is_available():\n    checkpoint = torch.load(PATH)\nelse:\n    checkpoint = torch.load(PATH, map_location=torch.device('cpu'))\n\n# Load states\nmodel.load_state_dict(checkpoint['model_state_dict'])\nepoch = checkpoint['epoch']\nloss = checkpoint['loss']\nval_loss = checkpoint['val_loss']\n\n# Evaluation mode\nmodel.eval()\nmodel.to(device)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-09-11T11:36:33.795052Z","iopub.execute_input":"2022-09-11T11:36:33.795478Z","iopub.status.idle":"2022-09-11T11:36:38.462299Z","shell.execute_reply.started":"2022-09-11T11:36:33.795416Z","shell.execute_reply":"2022-09-11T11:36:38.461135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Print final loss and epoch\nprint('Final epoch:', epoch)\nprint('Final loss:', loss)\nprint('Final valid loss:', val_loss)","metadata":{"execution":{"iopub.status.busy":"2022-09-11T11:36:38.464541Z","iopub.execute_input":"2022-09-11T11:36:38.465833Z","iopub.status.idle":"2022-09-11T11:36:38.473027Z","shell.execute_reply.started":"2022-09-11T11:36:38.465782Z","shell.execute_reply":"2022-09-11T11:36:38.471867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Inference on test set","metadata":{}},{"cell_type":"code","source":"test_df['fractured']=0.5\nwith torch.no_grad():\n    # Loop over batches\n    for i, (imgs, patient) in enumerate(test_loader):\n        #print(f'Iteration {i+1}/{len(test_loader)}')\n        # Send to device\n        imgs = imgs.to(device)\n        \n        # Make predictions\n        preds = model(imgs)\n        print(preds)\n        \n        # Apply sigmoid\n        sig = nn.Sigmoid()\n        preds = sig(preds)\n        preds = preds.to('cpu')\n        print(preds)\n        \n        # Save preds\n        test_df.loc[test_df['StudyInstanceUID']==patient[0],'fractured'] = preds.numpy().squeeze()\n        \nprint('Inference complete!')","metadata":{"execution":{"iopub.status.busy":"2022-09-11T11:36:38.475327Z","iopub.execute_input":"2022-09-11T11:36:38.475715Z","iopub.status.idle":"2022-09-11T11:37:23.501663Z","shell.execute_reply.started":"2022-09-11T11:36:38.47568Z","shell.execute_reply":"2022-09-11T11:37:23.500459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Submission","metadata":{}},{"cell_type":"code","source":"submission = test_df[['row_id','fractured']]\nsubmission.to_csv('submission.csv', index=False)\nsubmission.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-09-11T11:37:23.503927Z","iopub.execute_input":"2022-09-11T11:37:23.504276Z","iopub.status.idle":"2022-09-11T11:37:23.524815Z","shell.execute_reply.started":"2022-09-11T11:37:23.504244Z","shell.execute_reply":"2022-09-11T11:37:23.523939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}