{"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":"# RSNA Fracture Detection | Grad-CAM PoC ✍️\n\n![Competition image](https://i.imgur.com/3iLsS6i.png)\n\n- *Author: Mariusz Wiśniewski*\n- *Competition: [RSNA 2022 Cervical Spine Fracture Detection](https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection)*\n\n## Overview\n\nIn this notebook, we will train a 3D convolutional neural network to predict the **probability of a fracture** for each of the **seven cervical vertebrae** denoted by *C1-C7*, as well as the **overall probability** of any fracture in the cervical spine from volumetric computer tomography (CT) scans. After that, we will see how to generate a class activation heatmap for our 3D image classification model.\n\n### Libraries Used\n\n- [Tensorflow 🔥](https://www.tensorflow.org)\n- [Weights&Biases 📈](https://wandb.ai/)\n\n### References\n\n- [[RSNA_22] Dicom to NumPy 3D 📓](https://www.kaggle.com/code/vmuzhichenko/rsna-22-dicom-to-numpy-3d)\n- [[RSNA_22] ResNet 50 3D Train 📓](https://www.kaggle.com/code/vmuzhichenko/rsna-22-resnet-50-3d-train)\n- [[RSNA_22] ResNet 50 3D Inference 📓](https://www.kaggle.com/code/vmuzhichenko/rsna-22-resnet-50-3d-inference)\n- [RNSA - 3D model [Train] [PyTorch] 📓](https://www.kaggle.com/code/samuelcortinhas/rnsa-3d-model-train-pytorch/notebook)\n- [Uniformizing Techniques to Process CT scans with 3D CNNs for Tuberculosis Prediction 📃](https://arxiv.org/abs/2007.13224)\n- [VoxNet: A 3D Convolutional Neural Network for Real-Time ObjectRecognition 📃](https://ieeexplore.ieee.org/document/7353481)\n- [3D image classification from CT scans 📝](https://keras.io/examples/vision/3D_image_classification/)","metadata":{}},{"cell_type":"markdown","source":"# Project Setup","metadata":{}},{"cell_type":"markdown","source":"## Import Statements","metadata":{}},{"cell_type":"code","source":"import os\nimport random\n\nos.environ['TF_CPP_MIN_LOG_LEVEL'] = '3'\n\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport tensorflow as tf\nimport tensorflow_addons as tfa\nimport wandb\nfrom ipywidgets import IntSlider, interact\nfrom kaggle_secrets import UserSecretsClient\nfrom matplotlib import animation, rc\nfrom matplotlib.patches import PathPatch, Rectangle\nfrom matplotlib.path import Path\nfrom scipy import ndimage\nfrom sklearn.model_selection import StratifiedShuffleSplit\nfrom tensorflow import keras\nfrom tensorflow.keras import layers as L\nfrom tensorflow.keras.callbacks import EarlyStopping, ReduceLROnPlateau, ModelCheckpoint\nfrom wandb.keras import WandbCallback","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:07.471391Z","iopub.execute_input":"2022-12-09T13:36:07.471891Z","iopub.status.idle":"2022-12-09T13:36:15.647698Z","shell.execute_reply.started":"2022-12-09T13:36:07.471775Z","shell.execute_reply":"2022-12-09T13:36:15.646496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Random Seed","metadata":{}},{"cell_type":"code","source":"# Random seed for reproducibility\nseed = 27\n\nrandom.seed(seed)\nos.environ['PYTHONHASHSEED'] = str(seed)\nnp.random.seed(seed)\ntf.random.set_seed(seed)\ntf.compat.v1.set_random_seed(seed)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:15.649735Z","iopub.execute_input":"2022-12-09T13:36:15.65048Z","iopub.status.idle":"2022-12-09T13:36:15.662376Z","shell.execute_reply.started":"2022-12-09T13:36:15.650446Z","shell.execute_reply":"2022-12-09T13:36:15.658528Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Weights & Biases Setup","metadata":{}},{"cell_type":"code","source":"user_secrets = UserSecretsClient()\n\n# wandb.login(key=user_secrets.get_secret('WANDB_API_KEY'))\n# run = wandb.init(\n#     name='12-mw-3D-CNN3-128x128x64-ess',\n#     project=user_secrets.get_secret('WANDB_PROJECT'),\n#     entity=user_secrets.get_secret('WANDB_ENTITY'),\n#     id='12-mw-3D-CNN3-v3',\n# )","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:15.663904Z","iopub.execute_input":"2022-12-09T13:36:15.664397Z","iopub.status.idle":"2022-12-09T13:36:15.681795Z","shell.execute_reply.started":"2022-12-09T13:36:15.664367Z","shell.execute_reply":"2022-12-09T13:36:15.679974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Project Configuration","metadata":{}},{"cell_type":"code","source":"wandb.config = {\n    'learning_rate': 1e-3,\n    'min_learning_rate': 1e-7,\n    'epochs': 200,\n    'batch_size': 4,\n    'test_batch_size': 4,\n    'img_size': 128,\n    'depth': 64,\n}\n\nconfig = wandb.config","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:15.683235Z","iopub.execute_input":"2022-12-09T13:36:15.684Z","iopub.status.idle":"2022-12-09T13:36:15.694518Z","shell.execute_reply.started":"2022-12-09T13:36:15.683958Z","shell.execute_reply":"2022-12-09T13:36:15.693131Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Preparation","metadata":{}},{"cell_type":"markdown","source":"## Loading and Preprocessing","metadata":{}},{"cell_type":"code","source":"SIZE_STRING = f\"{config['img_size']}x{config['img_size']}x{config['depth']}\"\nIMG_PATH_TRAIN = (\n    f'../input/rsna-fd-3d-numpy-{SIZE_STRING}-ess/train_volumes_{SIZE_STRING}_ess/'\n)\nIMG_PATH_TEST = (\n    f'../input/rsna-fd-3d-numpy-{SIZE_STRING}-ess/test_volumes_{SIZE_STRING}_ess/'\n)\nTRAIN_CSV_PATH = '../input/rsna-2022-cervical-spine-fracture-detection/train.csv'\nTEST_CSV_PATH = '../input/rsna-2022-cervical-spine-fracture-detection/test.csv'\n\ntrain_images = [os.path.splitext(filename)[0] for filename in os.listdir(IMG_PATH_TRAIN)]\ntest_images = [os.path.splitext(filename)[0] for filename in os.listdir(IMG_PATH_TEST)]\ntrain_df = pd.read_csv(TRAIN_CSV_PATH)\ntest_df = train_df[train_df['StudyInstanceUID'].isin(test_images)]\n# We want our dataframe to contain only the train_images\ntrain_df = train_df[train_df['StudyInstanceUID'].isin(train_images)]\n\ntrain_df['numpy_path'] = train_df['StudyInstanceUID'].apply(\n    lambda x: f'{IMG_PATH_TRAIN}{x}.npz'\n)\ntest_df['numpy_path'] = test_df['StudyInstanceUID'].apply(\n    lambda x: f'{IMG_PATH_TEST}{x}.npz'\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:15.699709Z","iopub.execute_input":"2022-12-09T13:36:15.700191Z","iopub.status.idle":"2022-12-09T13:36:16.172182Z","shell.execute_reply.started":"2022-12-09T13:36:15.700153Z","shell.execute_reply":"2022-12-09T13:36:16.171408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Removal of Bad Scans\n\nScan `1.2.826.0.1.3680043.20574` does not include a full cervical spine, whereas scan `1.2.826.0.1.3680043.29952` includes slices obtained not in the axial, but in the coronal plane. We will simply ignore these scans. See [this post](https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/344862) for more.","metadata":{}},{"cell_type":"code","source":"bad_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)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:16.173222Z","iopub.execute_input":"2022-12-09T13:36:16.17387Z","iopub.status.idle":"2022-12-09T13:36:16.182883Z","shell.execute_reply.started":"2022-12-09T13:36:16.173841Z","shell.execute_reply":"2022-12-09T13:36:16.181833Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Augmentation\n\nThere are numerous preprocessing and augmentation techniques available, this example demonstrates only a very basic one to get started.","metadata":{}},{"cell_type":"code","source":"def random_rotate(volume):\n    # define some rotation angles\n    angles = [-20, -10, -5, 5, 10, 20]\n    # pick angles at random\n    angle = random.choice(angles)\n    # rotate volume\n    volume = ndimage.rotate(volume, angle, reshape=False)\n    return volume","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:16.184197Z","iopub.execute_input":"2022-12-09T13:36:16.184694Z","iopub.status.idle":"2022-12-09T13:36:16.192945Z","shell.execute_reply.started":"2022-12-09T13:36:16.18466Z","shell.execute_reply":"2022-12-09T13:36:16.192076Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Augmented CT Scan Visualization","metadata":{}},{"cell_type":"code","source":"example_volume = np.load(train_df['numpy_path'].iloc[0])\nexample_volume = example_volume['data']\nexample_volume = random_rotate(example_volume)\nplt.imshow(np.squeeze(example_volume[:, :, 0]), cmap='bone')","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:16.1942Z","iopub.execute_input":"2022-12-09T13:36:16.194695Z","iopub.status.idle":"2022-12-09T13:36:16.557714Z","shell.execute_reply.started":"2022-12-09T13:36:16.194668Z","shell.execute_reply":"2022-12-09T13:36:16.556704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Generator","metadata":{}},{"cell_type":"code","source":"class SampleGenerator(tf.keras.utils.Sequence):\n    def __init__(\n        self,\n        df: pd.DataFrame,\n        batch_size,\n        resample_rate: float = None,\n        steps_per_epoch: int = 10000,\n        is_train=True,\n        shuffle=True,\n        augment=False\n    ):\n        self.is_train = is_train\n        self.numpy_path = df.numpy_path\n        self.df = df\n        self.batch_size = batch_size\n        self.length = len(df)\n        self.resample = resample_rate\n        self.shuffle = shuffle\n        self.steps_per_epoch = steps_per_epoch\n        self.augment = augment\n\n    def __len__(self):\n        return min(\n            int(np.ceil(self.length / float(self.batch_size))), self.steps_per_epoch\n        )\n\n    def on_epoch_end(self):\n        if self.shuffle:\n            self.df = self.df.sample(frac=1, random_state=seed).reset_index(drop=True)\n            self.numpy_path = self.df.numpy_path\n\n    def __getitem__(self, index):\n        batch_x = []\n        if self.is_train:\n            batch_y = []\n\n            targets = self.df[\n                ['patient_overall', 'C1', 'C2', 'C3', 'C4', 'C5', 'C6', 'C7']\n            ]\n\n            for i in range(self.batch_size):\n                cur_ind = self.batch_size * index + i\n                if cur_ind < self.length:\n                    volume = np.load(self.numpy_path.iloc[cur_ind])\n                    if type(volume) is np.lib.npyio.NpzFile:\n                        volume = volume['data']\n                    batch_x.append(volume)\n                    batch_y.append(targets.iloc[cur_ind])\n\n            if self.resample is not None:\n                n_images = batch_x[0].shape[0]\n                im_ids = sorted(\n                    np.random.choice(\n                        list(range(n_images)),\n                        int(n_images * self.resample),\n                        replace=False,\n                    )\n                )\n                batch_x = np.array(batch_x)[:, im_ids]\n\n            if self.augment:\n                batch_x = random_rotate(batch_x)\n\n            return np.array(batch_x), np.array(batch_y).astype(np.float32)\n\n        else:\n            for i in range(self.batch_size):\n                cur_ind = self.batch_size * index + i\n                if cur_ind < self.length:\n                    volume = np.load(self.numpy_path.iloc[cur_ind])\n                    if type(volume) is np.lib.npyio.NpzFile:\n                        volume = volume['data']\n                    batch_x.append(volume)\n            if self.augment:\n                batch_x = random_rotate(batch_x)\n            \n            return np.array(batch_x)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:16.558989Z","iopub.execute_input":"2022-12-09T13:36:16.559345Z","iopub.status.idle":"2022-12-09T13:36:16.576111Z","shell.execute_reply.started":"2022-12-09T13:36:16.559317Z","shell.execute_reply":"2022-12-09T13:36:16.57466Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Build Train and Validation Datasets\n\nWe use a stratified split based on `patient_overall`.","metadata":{}},{"cell_type":"code","source":"pred_columns = [\n    'patient_overall_pred',\n    'C1_pred',\n    'C2_pred',\n    'C3_pred',\n    'C4_pred',\n    'C5_pred',\n    'C6_pred',\n    'C7_pred',\n]\ntrain_df[pred_columns] = 0\n\ntrain_idx, val_idx = next(\n    StratifiedShuffleSplit(1, train_size=0.8, random_state=seed).split(\n        train_df, train_df.patient_overall\n    )\n)\nX_train = train_df.iloc[train_idx]\nX_valid = train_df.iloc[val_idx]\n\ntrain_data = SampleGenerator(\n    X_train,\n    config['batch_size'],\n    steps_per_epoch=min(len(X_train), config['batch_size']),\n    shuffle=True,\n    augment=True,\n)\nvalidation_data = SampleGenerator(X_valid, config['test_batch_size'], shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:16.577578Z","iopub.execute_input":"2022-12-09T13:36:16.5779Z","iopub.status.idle":"2022-12-09T13:36:16.616071Z","shell.execute_reply.started":"2022-12-09T13:36:16.577844Z","shell.execute_reply":"2022-12-09T13:36:16.614263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3D Convolutional Neural Network\n\n## 3D Convolutions\n\nRGB images that consist of 3 channels are typically processed using 2D CNNs. A 3D CNN is essentially the 3D equivalent: it takes as input a 3D volume or a sequence of 2D frames (e.g. slices in a CT scan). 3D CNNs are a powerful model for learning representations for volumetric data. In contrary to the 2D convolutions, here in 3D convolution, the filter depth is smaller than the input layer depth (*kernel size < channel size*). As a result, the 3D filter has the ability to move in all three directions (height, width, channel of the image). At each position, the element-wise multiplication and addition provide one number. Since the filter slides through a 3D space, the output numbers are arranged in a 3D space as well. The output is then a 3D data.\n\n<center>\n<figure>\n    <img src='https://i.imgur.com/2nJzE83.gif' alt='3D Convolution' style='width: 680px;'/>\n    <figcaption>Visualization of a 3D convolution of <i>5x5x5</i> volume with <i>3x3x3</i> kernel, no padding, no strides. It results in a <i>3x3x3</i> output volume.</figcaption>\n</figure>\n</center>\n    \nSimilarly to 2D convolutions, which encode spatial relationships of objects in a 2D domain, 3D convolutions can describe the spatial relationships of objects in the 3D space.","metadata":{}},{"cell_type":"markdown","source":"## Loss Function","metadata":{}},{"cell_type":"code","source":"def competition_loss(y_true, y_pred):\n    \"\"\"\n    Source: https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/341854#1884562\n    \"\"\"\n    competition_weights = {\n        '-': tf.constant([7, 1, 1, 1, 1, 1, 1, 1], dtype=tf.float32),\n        '+': tf.constant([14, 2, 2, 2, 2, 2, 2, 2], dtype=tf.float32),\n    }\n\n    loss = tf.keras.losses.BinaryCrossentropy(reduction=tf.keras.losses.Reduction.NONE)(\n        tf.expand_dims(y_true, -1), tf.expand_dims(y_pred, -1)\n    )\n    weights = (\n        y_true * competition_weights['+'] + (1 - y_true) * competition_weights['-']\n    )\n\n    loss = tf.reduce_mean(tf.reduce_sum(loss * weights, axis=1)) / tf.reduce_sum(\n        weights\n    )\n    return loss","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:16.618016Z","iopub.execute_input":"2022-12-09T13:36:16.618576Z","iopub.status.idle":"2022-12-09T13:36:16.62992Z","shell.execute_reply.started":"2022-12-09T13:36:16.618514Z","shell.execute_reply":"2022-12-09T13:36:16.628222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## About the Model\n\nTo make the model easier to understand, we structure it into blocks. The architecture of the 3D CNN used in this example is based on [Uniformizing Techniques to Process CT scans with 3D CNNs for Tuberculosis Prediction](https://arxiv.org/abs/2007.13224) paper.","metadata":{}},{"cell_type":"code","source":"def create_model(width=128, height=128, depth=64, num_classes=8):\n    inputs = L.Input((width, height, depth, 1), name='inputs')\n    x = L.Rescaling(1.0 / 255)(inputs)\n\n    x = L.Conv3D(\n        filters=16, kernel_size=7, strides=(1, 1, 1), padding='same', activation='relu'\n    )(x)\n    x = L.MaxPool3D(pool_size=(2, 2, 2), strides=(2, 2, 2), padding='same')(x)\n    x = L.BatchNormalization()(x)\n\n    x = L.Conv3D(\n        filters=32, kernel_size=3, strides=(1, 1, 1), padding='same', activation='relu'\n    )(x)\n    x = L.MaxPool3D(pool_size=(2, 2, 2), strides=(2, 2, 2), padding='same')(x)\n    x = L.BatchNormalization()(x)\n\n    x = L.Conv3D(\n        filters=64, kernel_size=3, strides=(1, 1, 1), padding='same', activation='relu'\n    )(x)\n    x = L.MaxPool3D(pool_size=(2, 2, 2), strides=(2, 2, 2), padding='same')(x)\n    x = L.BatchNormalization()(x)\n\n    x = L.GlobalAveragePooling3D()(x)\n    x = L.Dense(units=32, activation='relu')(x)\n    x = L.Dropout(0.5)(x)\n\n    outputs = L.Dense(units=num_classes, activation='sigmoid')(x)\n    model = tf.keras.Model(\n        inputs=inputs, outputs=outputs, name=f'3D-CNN_SIZ_{SIZE_STRING}'\n    )\n\n    model.compile(\n        loss=competition_loss,  # 'binary_crossentropy',\n        optimizer=keras.optimizers.Adam(learning_rate=config['learning_rate']),\n        metrics=[\n            'acc',\n            'AUC',\n            tfa.metrics.F1Score(\n                num_classes=num_classes, threshold=0.25, average='macro'\n            ),\n        ],\n    )\n    return model\n\n\nmodel = create_model(\n    width=config['img_size'],\n    height=config['img_size'],\n    depth=config['depth'],\n    num_classes=8,\n)\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:16.631686Z","iopub.execute_input":"2022-12-09T13:36:16.632119Z","iopub.status.idle":"2022-12-09T13:36:16.835019Z","shell.execute_reply.started":"2022-12-09T13:36:16.632085Z","shell.execute_reply":"2022-12-09T13:36:16.833044Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Training","metadata":{}},{"cell_type":"code","source":"# model.fit(\n#     train_data,\n#     epochs=config['epochs'],\n#     callbacks=[\n#         EarlyStopping(\n#             monitor='val_loss', mode='min', patience=50, restore_best_weights=True\n#         ),\n#         ReduceLROnPlateau(\n#             monitor='val_loss',\n#             mode='min',\n#             patience=13,\n#             factor=0.5,\n#             min_lr=config['min_learning_rate'],\n#         ),\n#         ModelCheckpoint('best.h5', mode='min', save_best_only=True),\n# #         WandbCallback(mode='min'),\n#     ],\n#     validation_data=(validation_data),\n# )","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-09T13:36:16.836847Z","iopub.execute_input":"2022-12-09T13:36:16.837269Z","iopub.status.idle":"2022-12-09T13:36:16.842244Z","shell.execute_reply.started":"2022-12-09T13:36:16.837229Z","shell.execute_reply":"2022-12-09T13:36:16.841086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualizing Training History\n\nThe training and validation sets' model accuracy and loss are displayed here. Accuracy gives an impartial picture of the model's performance since the validation set is class-balanced.","metadata":{}},{"cell_type":"code","source":"# fig, ax = plt.subplots(1, 2, figsize=(20, 3))\n# ax = ax.ravel()\n\n# for i, metric in enumerate(['acc', 'loss']):\n#     ax[i].plot(model.history.history[metric])\n#     ax[i].plot(model.history.history[f'val_{metric}'])\n#     ax[i].set_title(f'Model {metric}')\n#     ax[i].set_xlabel('epochs')\n#     ax[i].set_ylabel(metric)\n#     ax[i].legend(['train', 'val'])","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:16.846914Z","iopub.execute_input":"2022-12-09T13:36:16.847289Z","iopub.status.idle":"2022-12-09T13:36:16.860244Z","shell.execute_reply.started":"2022-12-09T13:36:16.84726Z","shell.execute_reply":"2022-12-09T13:36:16.859384Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Cleanup","metadata":{}},{"cell_type":"code","source":"del X_train, train_data\n# wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:16.861391Z","iopub.execute_input":"2022-12-09T13:36:16.86271Z","iopub.status.idle":"2022-12-09T13:36:16.875753Z","shell.execute_reply.started":"2022-12-09T13:36:16.862636Z","shell.execute_reply":"2022-12-09T13:36:16.874093Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Example Prediction","metadata":{}},{"cell_type":"code","source":"model = tf.keras.models.load_model(\n    f'../input/rsna22fracturedetectionmodels/12-mw-3D-CNN3-128x128x64-ess.h5',\n    custom_objects={\n        'competition_loss': competition_loss,\n        'f1_score': tfa.metrics.F1Score,\n    },\n)\ntest_data = SampleGenerator(test_df, 1, shuffle=False)\ninput_volume = test_data.__getitem__(0)[0][0]\ninput_label = test_df['StudyInstanceUID'].iloc[0]\nmodel(np.expand_dims(input_volume, axis=0))","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:16.877189Z","iopub.execute_input":"2022-12-09T13:36:16.878638Z","iopub.status.idle":"2022-12-09T13:36:17.489159Z","shell.execute_reply.started":"2022-12-09T13:36:16.878573Z","shell.execute_reply":"2022-12-09T13:36:17.487911Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(input_volume.shape)\nplt.imshow(np.squeeze(input_volume[:, :, 30]), cmap='bone')","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:17.490494Z","iopub.execute_input":"2022-12-09T13:36:17.491028Z","iopub.status.idle":"2022-12-09T13:36:17.653068Z","shell.execute_reply.started":"2022-12-09T13:36:17.490998Z","shell.execute_reply":"2022-12-09T13:36:17.652182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Grad-CAM 3D Visualizations\n\nNow let us obtain a class activation heatmap for our image classification model. A detailed description of the procedure can be found in [Grad-CAM: Visual Explanations from Deep Networks via Gradient-based Localization](https://arxiv.org/abs/1610.02391) paper.\n\n**Gradient-weighted Class Activation Mapping (Grad-CAM)** employs the gradients of any target concept (for example, \"tiger\" in a classification network or a sequence of words in a captioning network) flowing into the final convolutional layer to generate a coarse localization map highlighting the important regions in the image for predicting the concept.","metadata":{}},{"cell_type":"markdown","source":"## Configurable Parameters\n\nSeveral prior studies claim that deeper representations in a CNN capture higher-level visual constructs. Furthermore, because convolutional layers naturally preserve spatial information that is lost in fully-connected layers, we may anticipate the **last** convolutional layers to provide the best compromise of high-level semantics and detailed spatial information.\n\nUse `model.summary()` to see the names of all layers in the model. These are necessary to get the value for `last_conv_layer_name`.","metadata":{}},{"cell_type":"code","source":"volume_size = input_volume.shape\nlast_conv_layer_name = 'conv3d_2'","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:17.654504Z","iopub.execute_input":"2022-12-09T13:36:17.65585Z","iopub.status.idle":"2022-12-09T13:36:17.661989Z","shell.execute_reply.started":"2022-12-09T13:36:17.65579Z","shell.execute_reply":"2022-12-09T13:36:17.66023Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Grad-CAM Algorithm\n\n*Grad-CAM* is **class-discriminative**, meaning it can produce a separate visualization for every class present in the image. This is the reason why we introduce the `pred_index` argument. Keep in mind that when we do not pass any value to our `pred_index`, the generated heatmap will correspond to the class with the highest probability. It uses the gradient information flowing into the last convolutional layer of the CNN to assign importance values to each neuron for a particular decision of interest.","metadata":{}},{"cell_type":"code","source":"def make_gradcam_heatmap(img_array, model, last_conv_layer_name, pred_index=None):\n    \"\"\"Generate class activation heatmap\"\"\"\n    # First, we create a model that maps the input image to the activations\n    # of the last conv layer as well as the output predictions\n    grad_model = tf.keras.Model(\n        [model.inputs], [model.get_layer(\n            last_conv_layer_name).output, model.output]\n    )\n\n    # Then, we compute the gradient of the top predicted class for our input image\n    # with respect to the activations of the last conv layer\n    with tf.GradientTape() as tape:\n        last_conv_layer_output, preds = grad_model(img_array)\n        if pred_index is None:\n            pred_index = tf.argmax(preds[0])\n        class_channel = preds[:, pred_index]\n\n    # This is the gradient of the output neuron (top predicted or chosen)\n    # with regard to the output feature map of the last conv layer\n    grads = tape.gradient(class_channel, last_conv_layer_output)\n\n    # This is a vector where each entry is the mean intensity of the gradient\n    # over a specific feature map channel (equivalent to global average pooling)\n    pooled_grads = tf.reduce_mean(grads, axis=(0, 1, 2, 3))\n\n    # We multiply each channel in the feature map array\n    # by 'how important this channel is' with regard to the top predicted class\n    # then sum all the channels to obtain the heatmap class activation\n    last_conv_layer_output = last_conv_layer_output[0]\n    heatmap = last_conv_layer_output @ pooled_grads[..., tf.newaxis]\n    heatmap = tf.squeeze(heatmap)\n\n    # For visualization purpose, we will also normalize the heatmap between 0 & 1\n    # Notice that we clip the heatmap values, which is equivalent to applying ReLU\n    heatmap = tf.maximum(heatmap, 0) / tf.math.reduce_max(heatmap)\n    return heatmap.numpy()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:17.663784Z","iopub.execute_input":"2022-12-09T13:36:17.664088Z","iopub.status.idle":"2022-12-09T13:36:17.674743Z","shell.execute_reply.started":"2022-12-09T13:36:17.66406Z","shell.execute_reply":"2022-12-09T13:36:17.672888Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Heatmap Generation","metadata":{}},{"cell_type":"code","source":"# Remove last layer's activation\nmodel.layers[-1].activation = None\n\n# Print what the top predicted class is\nimg_array = np.expand_dims(input_volume, axis=0)\n\npreds = model.predict(img_array)\nprint('Predicted:', preds[0])\n\n# Generate class activation heatmap\nheatmaps = [make_gradcam_heatmap(img_array, model, last_conv_layer_name, pred_index=idx) for idx in range(8)]\n# heatmap = make_gradcam_heatmap(img_array, model, last_conv_layer_name)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:17.677408Z","iopub.execute_input":"2022-12-09T13:36:17.677818Z","iopub.status.idle":"2022-12-09T13:36:53.860748Z","shell.execute_reply.started":"2022-12-09T13:36:17.677782Z","shell.execute_reply":"2022-12-09T13:36:53.859395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for heatmap in heatmaps:\n    print(f'Heatmap shape: {heatmap.shape}')\n    plt.matshow(np.squeeze(heatmap[:, :, 1]))\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:53.862634Z","iopub.execute_input":"2022-12-09T13:36:53.86299Z","iopub.status.idle":"2022-12-09T13:36:55.149199Z","shell.execute_reply.started":"2022-12-09T13:36:53.862959Z","shell.execute_reply":"2022-12-09T13:36:55.148047Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Expanding Heatmap Dimensions\n\nNotice that similarly to resizing the input volume, expanding the heatmap dimensions is based on the *spline interpolated zoom*.","metadata":{}},{"cell_type":"code","source":"def get_resized_heatmap(heatmap, shape):\n    \"\"\"Resize heatmap to shape\"\"\"\n    # Rescale heatmap to a range 0-255\n    upscaled_heatmap = np.uint8(255 * heatmap)\n\n    upscaled_heatmap = ndimage.zoom(\n        upscaled_heatmap,\n        (\n            shape[0] / upscaled_heatmap.shape[0],\n            shape[1] / upscaled_heatmap.shape[1],\n            shape[2] / upscaled_heatmap.shape[2],\n        ),\n    )\n\n    return upscaled_heatmap\n\n\n# resized_heatmap = get_resized_heatmap(heatmap, input_volume.shape)\nresized_heatmaps = [get_resized_heatmap(heatmap, input_volume.shape) for heatmap in heatmaps]","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:55.150882Z","iopub.execute_input":"2022-12-09T13:36:55.151182Z","iopub.status.idle":"2022-12-09T13:36:57.875318Z","shell.execute_reply.started":"2022-12-09T13:36:55.151154Z","shell.execute_reply":"2022-12-09T13:36:57.873918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualizations\n\nNow it is time for us to graphically visualize the results obtained by overlaying the heatmap with the image. We utilize the `jet` colormap for this.","metadata":{}},{"cell_type":"code","source":"for resized_heatmap in resized_heatmaps:\n    fig, ax = plt.subplots(1, 2, figsize=(10, 20))\n\n    ax[0].imshow(np.squeeze(input_volume[:, :, 30]), cmap='bone')\n    img0 = ax[1].imshow(np.squeeze(input_volume[:, :, 30]), cmap='bone')\n    img1 = ax[1].imshow(np.squeeze(resized_heatmap[:, :, 30]),\n                        cmap='jet', alpha=0.3, extent=img0.get_extent())\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:36:57.87717Z","iopub.execute_input":"2022-12-09T13:36:57.877617Z","iopub.status.idle":"2022-12-09T13:37:00.42995Z","shell.execute_reply.started":"2022-12-09T13:36:57.877578Z","shell.execute_reply":"2022-12-09T13:37:00.428634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Bounding Boxes\n\nHere we prepare some functions that we will use to annotate images by drawing bounding boxes around the regions of interest. The bounding boxes are drawn on the images using the coordinates of the obtained from the heatmap in the following format: `(x_center, y_center, width, height)`.\n\nThe process of obtaining the bounding boxes is as follows:\n\n1. Obtain the coordinates of the heatmap in places where its values are above a certain threshold (optionally, we can exploit Otsu's method to automatically determine the threshold).\n2. Find the connected components in the binary image obtained from the heatmap. Each connected component corresponds to a region of interest.\n3. For each connected component, obtain the bounding box coordinates using the minimal up-right rectangle technique.\n4. Draw the bounding boxes on the image.","metadata":{}},{"cell_type":"code","source":"def get_bounding_boxes(heatmap, threshold=0.15, otsu=False):\n    \"\"\"Get bounding boxes from heatmap\"\"\"\n    p_heatmap = np.copy(heatmap)\n\n    if otsu:\n        # Otsu's thresholding method to find the bounding boxes\n        threshold, p_heatmap = cv2.threshold(\n            heatmap, 0, 255, cv2.THRESH_BINARY + cv2.THRESH_OTSU\n        )\n    else:\n        # Using a fixed threshold\n        p_heatmap[p_heatmap < threshold * 255] = 0\n        p_heatmap[p_heatmap >= threshold * 255] = 1\n\n    # find the contours in the thresholded heatmap\n    contours = cv2.findContours(p_heatmap, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n    contours = contours[0] if len(contours) == 2 else contours[1]\n\n    # get the bounding boxes from the contours\n    bboxes = []\n    for c in contours:\n        x, y, w, h = cv2.boundingRect(c)\n        bboxes.append([x, y, x + w, y + h])\n\n    return bboxes\n\n\ndef get_bbox_patches(bboxes, color='r', linewidth=2):\n    \"\"\"Get patches for bounding boxes\"\"\"\n    patches = []\n    for bbox in bboxes:\n        x1, y1, x2, y2 = bbox\n        patches.append(\n            Rectangle(\n                (x1, y1),\n                x2 - x1,\n                y2 - y1,\n                edgecolor=color,\n                facecolor='none',\n                linewidth=linewidth,\n            )\n        )\n    return patches","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:37:00.431327Z","iopub.execute_input":"2022-12-09T13:37:00.431734Z","iopub.status.idle":"2022-12-09T13:37:00.446369Z","shell.execute_reply.started":"2022-12-09T13:37:00.431701Z","shell.execute_reply":"2022-12-09T13:37:00.444823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for resized_heatmap in resized_heatmaps:\n    # show the bounding boxes on the original image\n    fig, ax = plt.subplots(1, 2, figsize=(10, 20))\n\n    ax[0].imshow(np.squeeze(input_volume[:, :, 30]), cmap='bone')\n    img0 = ax[1].imshow(np.squeeze(input_volume[:, :, 30]), cmap='bone')\n    img1 = ax[1].imshow(np.squeeze(resized_heatmap[:, :, 30]),\n                        cmap='jet', alpha=0.3, extent=img0.get_extent())\n\n    bboxes = get_bounding_boxes(np.squeeze(resized_heatmap[:, :, 30]))\n    patches = get_bbox_patches(bboxes)\n\n    for patch in patches:\n        ax[1].add_patch(patch)\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:37:00.447951Z","iopub.execute_input":"2022-12-09T13:37:00.448305Z","iopub.status.idle":"2022-12-09T13:37:02.809619Z","shell.execute_reply.started":"2022-12-09T13:37:00.448264Z","shell.execute_reply":"2022-12-09T13:37:02.807762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Interactive Slice Viewers\n\n_**Note:** unfortunately, *kaggle* does not currently handle preserving the widget state and embedding it into the static notebook preview, thus **dragging the slider has no effect on the displayed figures**. See the following discussions for more information: [#33450](https://www.kaggle.com/questions-and-answers/33450), [#42782](https://www.kaggle.com/product-feedback/42782), [#2360](https://github.com/jupyter-widgets/ipywidgets/issues/2360), [#13637754](https://stackoverflow.com/a/63575304/13637754). To see the interactive visualizations, simply copy the notebook and run it yourself._","metadata":{}},{"cell_type":"code","source":"def _draw_line(ax, coords, clr='g'):\n    line = Path(coords, [Path.MOVETO, Path.LINETO])\n    pp = PathPatch(line, linewidth=3, edgecolor=clr, facecolor='none')\n    ax.add_patch(pp)\n\n\ndef _set_axes_labels(ax, axes_x, axes_y):\n    ax.set_xlabel(axes_x)\n    ax.set_ylabel(axes_y)\n    ax.set_aspect('equal', 'box')\n\n\ndef _draw_bboxes(ax, heatmap):\n    bboxes = get_bounding_boxes(heatmap, otsu=True)\n    patches = get_bbox_patches(bboxes)\n    for patch in patches:\n        ax.add_patch(patch)\n\n\n_rec_prop = dict(linewidth=5, facecolor='none')\n\n\ndef show_volume(vol, z, y, x, heatmap=None, alpha=0.3, fig_size=(6, 6)):\n    \"\"\"Show a slice of a volume with optional heatmap\"\"\"\n    fig, axarr = plt.subplots(nrows=2, ncols=2, figsize=fig_size)\n    v_z, v_y, v_x = vol.shape\n\n    img0 = axarr[0, 0].imshow(vol[z, :, :], cmap='bone')\n    if heatmap is not None:\n        axarr[0, 0].imshow(\n            heatmap[z, :, :], cmap='jet', alpha=alpha, extent=img0.get_extent()\n        )\n        _draw_bboxes(axarr[0, 0], heatmap[z, :, :])\n\n    axarr[0, 0].add_patch(Rectangle((-1, -1), v_x, v_y, edgecolor='r', **_rec_prop))\n    _draw_line(axarr[0, 0], [(x, 0), (x, v_y)], 'g')\n    _draw_line(axarr[0, 0], [(0, y), (v_x, y)], 'b')\n    _set_axes_labels(axarr[0, 0], 'X', 'Y')\n\n    img1 = axarr[0, 1].imshow(vol[:, :, x].T, cmap='bone')\n    if heatmap is not None:\n        axarr[0, 1].imshow(\n            heatmap[:, :, x].T, cmap='jet', alpha=alpha, extent=img1.get_extent()\n        )\n        _draw_bboxes(axarr[0, 1], heatmap[:, :, x].T)\n\n    axarr[0, 1].add_patch(Rectangle((-1, -1), v_z, v_y, edgecolor='g', **_rec_prop))\n    _draw_line(axarr[0, 1], [(z, 0), (z, v_y)], 'r')\n    _draw_line(axarr[0, 1], [(0, y), (v_x, y)], \"b\")\n    _set_axes_labels(axarr[0, 1], 'Z', 'Y')\n\n    img2 = axarr[1, 0].imshow(vol[:, y, :], cmap='bone')\n    if heatmap is not None:\n        axarr[1, 0].imshow(\n            heatmap[:, y, :], cmap='jet', alpha=alpha, extent=img2.get_extent()\n        )\n        _draw_bboxes(axarr[1, 0], heatmap[:, y, :])\n\n    axarr[1, 0].add_patch(Rectangle((-1, -1), v_x, v_z, edgecolor='b', **_rec_prop))\n    _draw_line(axarr[1, 0], [(0, z), (v_x, z)], 'r')\n    _draw_line(axarr[1, 0], [(x, 0), (x, v_y)], 'g')\n    _set_axes_labels(axarr[1, 0], 'X', 'Z')\n    axarr[1, 1].set_axis_off()\n    fig.tight_layout()\n\n\ndef interactive_show(volume, heatmap=None):\n    \"\"\"Show a volume interactively\"\"\"\n    # transpose volume from (x, y, z) to (z, x, y)\n    volume = np.transpose(volume, (2, 0, 1))\n    if heatmap is not None:\n        heatmap = np.transpose(heatmap, (2, 0, 1))\n    vol_shape = volume.shape\n\n    interact(\n        lambda x, y, z: plt.show(show_volume(volume, z, y, x, heatmap)),\n        z=IntSlider(min=0, max=vol_shape[0] - 1, step=1, value=int(vol_shape[0] / 2)),\n        y=IntSlider(min=0, max=vol_shape[1] - 1, step=1, value=int(vol_shape[1] / 2)),\n        x=IntSlider(min=0, max=vol_shape[2] - 1, step=1, value=int(vol_shape[2] / 2)),\n    )","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-09T13:37:02.811572Z","iopub.execute_input":"2022-12-09T13:37:02.811995Z","iopub.status.idle":"2022-12-09T13:37:02.835117Z","shell.execute_reply.started":"2022-12-09T13:37:02.811963Z","shell.execute_reply":"2022-12-09T13:37:02.833427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# interactive_show(input_volume)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:37:02.8366Z","iopub.execute_input":"2022-12-09T13:37:02.836958Z","iopub.status.idle":"2022-12-09T13:37:02.849153Z","shell.execute_reply.started":"2022-12-09T13:37:02.836927Z","shell.execute_reply":"2022-12-09T13:37:02.848071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# interactive_show(input_volume, resized_heatmap)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:37:02.850396Z","iopub.execute_input":"2022-12-09T13:37:02.851814Z","iopub.status.idle":"2022-12-09T13:37:02.861412Z","shell.execute_reply.started":"2022-12-09T13:37:02.851748Z","shell.execute_reply":"2022-12-09T13:37:02.859956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Animations","metadata":{}},{"cell_type":"code","source":"rc('animation', html='jshtml')\n\n\ndef create_animation(array, case, heatmap=None, alpha=0.3):\n    \"\"\"Create an animation of a volume\"\"\"\n    # transpose volume from (x, y, z) to (z, x, y)\n    array = np.transpose(array, (2, 0, 1))\n    if heatmap is not None:\n        heatmap = np.transpose(heatmap, (2, 0, 1))\n    fig = plt.figure(figsize=(6, 6))\n    images = []\n    for idx, image in enumerate(array):\n        # plot image without notifying animation\n        image_plot = plt.imshow(image, animated=True, cmap='bone')\n        aux = [image_plot]\n        if heatmap is not None:\n            image_plot2 = plt.imshow(\n                heatmap[idx], animated=True, cmap='jet', alpha=alpha, extent=image_plot.get_extent())\n            aux.append(image_plot2)\n\n            # add bounding boxes to the heatmap image as animated patches\n            bboxes = get_bounding_boxes(heatmap[idx], otsu=True)\n            patches = get_bbox_patches(bboxes)\n            aux.extend(image_plot2.axes.add_patch(patch) for patch in patches)\n        images.append(aux)\n\n    plt.axis('off')\n    plt.tight_layout()\n    plt.subplots_adjust(top=0.90)\n    plt.title(f'Patient ID: {case}', fontsize=16)\n    ani = animation.ArtistAnimation(\n        fig, images, interval=5000//len(array), blit=False, repeat_delay=1000)\n    plt.show()\n    return ani","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-09T13:37:53.20359Z","iopub.execute_input":"2022-12-09T13:37:53.204045Z","iopub.status.idle":"2022-12-09T13:37:53.220125Z","shell.execute_reply.started":"2022-12-09T13:37:53.204007Z","shell.execute_reply":"2022-12-09T13:37:53.218752Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"create_animation(input_volume, input_label)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:37:56.243697Z","iopub.execute_input":"2022-12-09T13:37:56.244179Z","iopub.status.idle":"2022-12-09T13:38:02.589034Z","shell.execute_reply.started":"2022-12-09T13:37:56.244142Z","shell.execute_reply":"2022-12-09T13:38:02.588332Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for resized_heatmap in resized_heatmaps:\n    create_animation(input_volume, input_label, heatmap=resized_heatmap)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:38:02.590269Z","iopub.execute_input":"2022-12-09T13:38:02.590767Z","iopub.status.idle":"2022-12-09T13:38:29.447128Z","shell.execute_reply.started":"2022-12-09T13:38:02.590734Z","shell.execute_reply":"2022-12-09T13:38:29.446098Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Testing\n\n*To be done*","metadata":{}},{"cell_type":"code","source":"test_bbox_df = pd.read_csv(\n    '../input/rsna-2022-cervical-spine-fracture-detection/train_bounding_boxes.csv'\n)\n\nprint(len(test_bbox_df[test_bbox_df['StudyInstanceUID'] == test_df['StudyInstanceUID'].iloc[0]]))","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:37:10.182569Z","iopub.execute_input":"2022-12-09T13:37:10.183212Z","iopub.status.idle":"2022-12-09T13:37:10.216369Z","shell.execute_reply.started":"2022-12-09T13:37:10.183166Z","shell.execute_reply":"2022-12-09T13:37:10.214977Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# preds = model.predict(test_data)","metadata":{"execution":{"iopub.status.busy":"2022-12-09T13:37:10.21826Z","iopub.execute_input":"2022-12-09T13:37:10.218652Z","iopub.status.idle":"2022-12-09T13:37:10.225045Z","shell.execute_reply.started":"2022-12-09T13:37:10.218619Z","shell.execute_reply":"2022-12-09T13:37:10.224267Z"},"trusted":true},"execution_count":null,"outputs":[]}]}