{"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":"# Interpreting 2-Stage 3D CSN with Grad-CAM, HiResCAM and Grad-CAM++ ✍️\n\n![CT Image](https://storage.googleapis.com/kaggle-datasets-images/674071/1185670/0d449dc88ac1318321ae8f7de974f0fa/dataset-cover.jpg?t=2020-05-25-13-21-33)\n\n- **Author:** *Mariusz Wiśniewski*\n- **Date created:** *December 13th, 2022*\n- **Last modified:** *December 28th, 2022*\n\n## Overview\n\nIn this notebook we will see how to generate class activation heatmaps for our 2-stage 3D image segmentation + classification model.\n\n### Libraries Used\n\n- [PyTorch 🔥](https://pytorch.org)\n- [NiBabel 🩻](https://nipy.org/nibabel/)\n- [SciPy 🔬](https://scipy.org)\n- [OpenCV 🖼️](https://opencv.org)\n- [pytorch-gradcam 🗾](https://jacobgil.github.io/pytorch-gradcam-book/introduction.html)\n- [torchinfo 💁‍♂️](https://github.com/TylerYep/torchinfo)\n\n### References\n\n- [RSNA CSN segmentor & classifier (by Selim Seferbekov) 📝](https://www.kaggle.com/code/selimsef/rsna-csn-segmentor-classifier)\n- [Video Classification with Channel-Separated Convolutional Networks 📃](https://arxiv.org/abs/1904.02811)\n- [Large-scale weakly-supervised pre-training for video action recognition 📃](https://arxiv.org/abs/1905.00561)\n- [Grad-CAM: Visual Explanations from Deep Networks via Gradient-based Localization 📃](https://arxiv.org/abs/1610.02391)\n- [Use HiResCAM instead of Grad-CAM for faithful explanations of convolutional neural networks 📃](https://arxiv.org/abs/2011.08891)\n- [Grad-CAM++: Improved Visual Explanations for Deep Convolutional Networks 📃](https://arxiv.org/abs/1710.11063)","metadata":{}},{"cell_type":"markdown","source":"# Notebook Setup","metadata":{}},{"cell_type":"markdown","source":"## Package Installation","metadata":{}},{"cell_type":"code","source":"!pip install --user torchinfo grad-cam mmcv -q","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:22:57.974153Z","iopub.execute_input":"2022-12-23T01:22:57.975053Z","iopub.status.idle":"2022-12-23T01:23:36.722685Z","shell.execute_reply.started":"2022-12-23T01:22:57.974907Z","shell.execute_reply":"2022-12-23T01:23:36.721265Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Import Statements","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('/kaggle/input/rsnazoopublic')","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:23:36.725252Z","iopub.execute_input":"2022-12-23T01:23:36.725655Z","iopub.status.idle":"2022-12-23T01:23:36.733437Z","shell.execute_reply.started":"2022-12-23T01:23:36.725616Z","shell.execute_reply":"2022-12-23T01:23:36.732395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport random\nimport re\nfrom dataclasses import dataclass\nfrom typing import Dict, List\n\nimport albumentations\nimport cv2\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport scipy.ndimage as ndimage\nimport tifffile\nimport torch\nimport torch.hub\nfrom albumentations import ReplayCompose\nfrom ipywidgets import IntSlider, interact\nfrom matplotlib import animation, rc\nfrom matplotlib.patches import PathPatch, Rectangle\nfrom matplotlib.path import Path\nfrom pytorch_grad_cam import GradCAM, GradCAMPlusPlus, HiResCAM\nfrom pytorch_grad_cam.utils.model_targets import ClassifierOutputTarget\nfrom skimage import measure\nfrom torch import nn\nfrom torch.functional import Tensor\nfrom torch.utils.data import DataLoader, Dataset\nfrom tqdm.notebook import tqdm\n\nimport zoo\nfrom zoo import ResNet3dCSN\nfrom zoo.utils import _initialize_weights","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:23:36.734763Z","iopub.execute_input":"2022-12-23T01:23:36.735064Z","iopub.status.idle":"2022-12-23T01:23:41.548793Z","shell.execute_reply.started":"2022-12-23T01:23:36.735036Z","shell.execute_reply":"2022-12-23T01:23:41.547481Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Global Variables","metadata":{}},{"cell_type":"code","source":"CROP_INFO = {}\nDEBUG = False","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:23:41.552134Z","iopub.execute_input":"2022-12-23T01:23:41.553039Z","iopub.status.idle":"2022-12-23T01:23:41.559602Z","shell.execute_reply.started":"2022-12-23T01:23:41.553001Z","shell.execute_reply":"2022-12-23T01:23:41.558361Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"markdown","source":"## Data Loading","metadata":{}},{"cell_type":"code","source":"@dataclass\nclass BatchSlice:\n    i_from: int\n    i_to: int\n    i_start: int\n\n\ndef get_slices(\n    batch: Tensor, dim=1, window: int = 16, overlap: int = 8\n) -> List[BatchSlice]:\n    num_imgs = batch.size(dim)\n    if num_imgs <= window:\n        print(f'Num imgs ({num_imgs}) <= window ({window})')\n        return [BatchSlice(0, num_imgs, 0)]\n    stride = window - overlap\n    result = []\n    current_idx = 0\n    while True:\n        next_idx = current_idx + window\n\n        if next_idx >= num_imgs:\n            current_idx = num_imgs - window\n            offset = overlap // 2 if current_idx > 0 else 0\n            next_idx = num_imgs\n            result.append(BatchSlice(current_idx, next_idx, offset))\n            break\n        else:\n            offset = overlap // 2 if current_idx > 0 else 0\n            result.append(BatchSlice(current_idx, next_idx, offset))\n        current_idx += stride\n    return result","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-23T01:23:41.561785Z","iopub.execute_input":"2022-12-23T01:23:41.5623Z","iopub.status.idle":"2022-12-23T01:23:41.576767Z","shell.execute_reply.started":"2022-12-23T01:23:41.562234Z","shell.execute_reply":"2022-12-23T01:23:41.575543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Preprocessing","metadata":{}},{"cell_type":"code","source":"def combine_scan(scan_dir: str, size=512, fix_monochrome: bool = True) -> np.ndarray:\n    num_files = len(os.listdir(scan_dir))\n    images = []\n    offset = 0\n    first = None\n    last = None\n    files = []\n    for i in range(num_files):\n        dpath = os.path.join(scan_dir, f'{i + offset}.dcm')\n        if i == 0:\n            while not os.path.exists(dpath):\n                offset += 1\n                dpath = os.path.join(scan_dir, f'{i + offset}.dcm')\n        files.append(dpath)\n\n    for dpath in files[::2]:\n        ds = pydicom.dcmread(dpath)\n        if not first:\n            first = ds\n        last = ds\n\n        data = ds.pixel_array\n        data = cv2.resize(data, (size, size))\n        if fix_monochrome and ds.PhotometricInterpretation == 'MONOCHROME1':\n            data = np.amax(data) - data\n        images.append(data)\n\n    if first and last and last.ImagePositionPatient[2] > first.ImagePositionPatient[2]:\n        images = images[::-1]\n    return np.array(images)\n\n\ndef preprocess_volume_for_segmentation(dataset_dir: str, case: str):\n    volume = combine_scan(os.path.join(dataset_dir, case), size=256)\n    volume_mean = volume.mean()\n    volume_std = volume.std()\n    h = volume.shape[0]\n\n    if h % 32 > 0:\n        tmp = np.zeros(((h // 32 + 1) * 32, 256, 256))\n        tmp[:h] = volume\n        volume = tmp\n\n    volume = (volume - volume_mean) / volume_std\n    volume = np.expand_dims(volume, 0)\n    return {'image': torch.from_numpy(volume).float(), 'cube_id': case, 'h': h}\n\n\ncrop_augs = albumentations.ReplayCompose(\n    [\n        albumentations.LongestMaxSize(256),\n        albumentations.PadIfNeeded(256, 256, border_mode=cv2.BORDER_CONSTANT),\n    ]\n)\n\n\ndef process_bbox(bbox, area, img_volume, mask_volume, n_slices, case, vertebra_idx):\n    z1, z2 = bbox[0], bbox[3]\n    y1, y2 = max(bbox[1] - 16, 0), min(bbox[4] + 16, 256)\n    x1, x2 = max(bbox[2] - 16, 0), min(bbox[5] + 16, 256)\n\n    if DEBUG:\n        print(\n            f'z1: {z1}, z2: {z2}, y1: {y1*2}, y2: {y2*2}, x1: {x1*2}, x2: {x2*2}')\n\n    if case not in CROP_INFO.keys():\n        CROP_INFO[case] = []\n    CROP_INFO[case].append((x1, x2, y1, y2, z1, z2))\n    images = img_volume[z1:z2, y1 * 2: y2 * 2, x1 * 2: x2 * 2].copy()\n    masks = mask_volume[z1:z2, y1:y2, x1:x2].copy()\n    return images, masks\n\n\ndef stack_cropped_volume_back_to_original_volume(\n    cropped_volume, original_volume, case\n):\n    replace_volume = np.zeros(original_volume.shape)\n    for vertebra_idx in range(len(cropped_volume)):\n        # get crop coordinates\n        if CROP_INFO[case][vertebra_idx] is not None:\n            x1, x2, y1, y2, z1, z2 = CROP_INFO[case][vertebra_idx]\n        else:\n            continue\n\n        # element-wise max\n        cropped_vertebra_volume = cropped_volume[vertebra_idx][0][0]\n\n        if DEBUG:\n            print(\n                f'Replace volume shape: {replace_volume[z1:z2, y1*2:y2*2, x1*2:x2*2].shape}')\n            print(f'Cropped volume shape: {cropped_vertebra_volume.shape}')\n\n        replace_volume[z1:z2, y1*2:y2*2, x1*2:x2*2] = np.maximum(replace_volume[z1:z2, y1*2:y2*2, x1*2:x2*2],\n                                                                 cropped_vertebra_volume)\n        if DEBUG:\n            print(\n                f'Cropped_volume min: {np.min(cropped_vertebra_volume.numpy())}, max: {np.max(cropped_vertebra_volume.numpy())}')\n\n    if DEBUG:\n        print(f'Min original_volume: {np.min(original_volume)}')\n        print(f'Max original_volume: {np.max(original_volume)}')\n\n    original_volume = np.where(\n        replace_volume != 0, replace_volume, original_volume)\n    original_volume = torch.from_numpy(original_volume)\n\n    return original_volume if type(original_volume) != torch.Tensor else original_volume.detach().clone()\n\n\ndef crop_image(image, mask, replay=None, transforms=crop_augs):\n    (\n        h,\n        w,\n    ) = mask.shape\n\n    mask = cv2.resize(mask, (w * 2, h * 2), interpolation=cv2.INTER_NEAREST)\n    if replay is None:\n        sample = transforms(image=image, mask=mask)\n        replay = sample['replay']\n    else:\n        sample = ReplayCompose.replay(replay, image=image, mask=mask)\n    return sample['image'], sample['mask']\n\n\ndef crop_volume(images, masks, transforms=crop_augs):\n    image_crops = []\n    mask_crops = []\n    for i in range(images.shape[0]):\n        img_crop, mask_crop = crop_image(\n            images[i], masks[i], replay=None, transforms=transforms\n        )\n        image_crops.append(img_crop)\n        mask_crops.append(mask_crop)\n\n    return image_crops, mask_crops\n\n\ndef stack_images_and_masks(\n    image_crops, mask_crops, image_volume_mean, image_volume_std\n):\n    images = np.array(image_crops).astype(np.float32)\n    masks = np.array(mask_crops).astype(np.float32)\n    images = np.expand_dims(images, -1)\n    masks = np.expand_dims(masks, -1)\n    images = (images - image_volume_mean) / image_volume_std\n\n    return np.concatenate([images, images, masks], axis=-1)\n\n\ndef resize_volume(images, n_slices=40):\n    return [np.moveaxis(images, -1, 0)]\n\n\ndef preprocess_volumes_for_classification(\n    dataset_dir: str, case: str, n_slices=40, transforms=crop_augs\n):\n    mask_volume = tifffile.imread(os.path.join('seg_preds', f'{case}.tif'))\n    print(np.array(mask_volume).shape)\n    image_volume = combine_scan(os.path.join(dataset_dir, case), size=512)\n    image_volume_mean = image_volume.mean()\n    image_volume_std = image_volume.std()\n    n_slices = image_volume.shape[0]\n\n    if DEBUG:\n        print(f'Image volume shape: {np.array(image_volume.shape)}')\n        print(f'Mask volume shape: {np.array(mask_volume.shape)}')\n\n    boxes = {\n        rprop.label: (rprop.bbox, rprop.area)\n        for rprop in measure.regionprops(mask_volume)\n    }\n\n    labels = np.zeros((8,))\n    all_images = []\n    volume_shape = []\n\n    for li in range(1, 8):\n        if li not in boxes:\n            all_images.append(np.zeros((3, n_slices, 256, 256)))\n            CROP_INFO[case].append(None)\n            volume_shape.append(0)\n        else:\n            bbox, area = boxes[li]\n            images, masks = process_bbox(\n                bbox, area, image_volume, mask_volume, n_slices, case, li\n            )\n\n            resized_masks = [\n                cv2.resize(\n                    mask,\n                    (mask.shape[1] * 2, mask.shape[0] * 2),\n                    interpolation=cv2.INTER_NEAREST,\n                )\n                for mask in masks\n            ]\n            masks = np.array(resized_masks)\n\n            if DEBUG:\n                print(\n                    f'Non transformed cropped images: {np.array(images).shape}')\n                print(\n                    f'Non transformed cropped masks: {np.array(masks).shape}')\n\n            images = stack_images_and_masks(\n                images, masks, image_volume_mean, image_volume_std\n            )\n\n            if DEBUG:\n                print(f'Volumes before resizing: {np.array(images).shape}')\n\n            resized_volumes = resize_volume(images, n_slices)\n\n            if DEBUG:\n                print(f'Resized volumes: {np.array(resized_volumes).shape}')\n\n            volume_shape.append(np.array(resized_volumes).shape)\n            all_images.extend(resized_volumes)\n\n    for idx, _ in enumerate(all_images):\n        if volume_shape and np.expand_dims(np.array(all_images[idx]), 0).shape != volume_shape[idx]:\n\n            if DEBUG:\n                print(f'Volume shape: {volume_shape[idx]}')\n                print(f'All images shape: {np.array(all_images[idx]).shape}')\n\n            all_images[idx] = np.zeros(volume_shape[idx])\n\n    all_images = [torch.from_numpy(np.array(all_images[idx])).unsqueeze(\n        0).float() for idx in range(len(all_images))]\n    return {\n        'image': all_images,\n        'label': torch.from_numpy(labels).float(),\n        'cube_id': case,\n    }","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-23T01:23:41.579105Z","iopub.execute_input":"2022-12-23T01:23:41.579691Z","iopub.status.idle":"2022-12-23T01:23:41.626494Z","shell.execute_reply.started":"2022-12-23T01:23:41.579591Z","shell.execute_reply":"2022-12-23T01:23:41.625344Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Models","metadata":{}},{"cell_type":"markdown","source":"## Loading Weights","metadata":{}},{"cell_type":"code","source":"def load_checkpoint(model, checkpoint_path, strict=False, verbose=True):\n    if verbose:\n        print(f\"=> loading checkpoint '{checkpoint_path}'\")\n    checkpoint = torch.load(checkpoint_path, map_location='cpu')\n    if 'state_dict' in checkpoint:\n        state_dict = checkpoint['state_dict']\n        state_dict = {re.sub('^module.', '', k): w for k,\n                      w in state_dict.items()}\n        orig_state_dict = model.state_dict()\n        mismatched_keys = []\n        for k, v in state_dict.items():\n            ori_size = orig_state_dict[k].size(\n            ) if k in orig_state_dict else None\n            if v.size() != ori_size:\n                if verbose:\n                    print(\n                        f'SKIPPING!!! Shape of {k} changed from {v.size()} to {ori_size}'\n                    )\n                mismatched_keys.append(k)\n        for k in mismatched_keys:\n            del state_dict[k]\n        model.load_state_dict(state_dict, strict=strict)\n        del state_dict\n        del orig_state_dict\n        print(\n            f\"=> loaded checkpoint '{checkpoint_path}' (epoch {checkpoint['epoch']})\")\n\n    else:\n        model.load_state_dict(checkpoint)\n    del checkpoint\n\n\ndef load_model(conf: Dict, checkpoint: str):\n    model = conf['network'](**conf['encoder_params'])\n    # model = model.cuda()\n    load_checkpoint(model, checkpoint)\n    return model.eval()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-23T01:23:41.628Z","iopub.execute_input":"2022-12-23T01:23:41.628394Z","iopub.status.idle":"2022-12-23T01:23:41.644149Z","shell.execute_reply.started":"2022-12-23T01:23:41.628361Z","shell.execute_reply":"2022-12-23T01:23:41.643353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config_seg = {\n    'network': zoo.ResNet3dCSN2P1D,\n    'encoder_params': {\n        'encoder': 'r50ir'\n    }\n}\ntest_dataset_dir = '/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images'\npreds_dir = './seg_preds'\ncases = os.listdir(test_dataset_dir)\nseg_model = load_model(\n    config_seg, '/kaggle/input/rsna-weights/256_ResNet3dCSN2P1D_r50ir_0_dice')","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:23:41.645751Z","iopub.execute_input":"2022-12-23T01:23:41.646062Z","iopub.status.idle":"2022-12-23T01:23:43.682217Z","shell.execute_reply.started":"2022-12-23T01:23:41.646035Z","shell.execute_reply":"2022-12-23T01:23:43.681238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# use bbox annotated data for testing\ntest_bbox_df = pd.read_csv(\n    '../input/rsna-2022-cervical-spine-fracture-detection/train_bounding_boxes.csv'\n)\n\ncases = list(set(test_bbox_df['StudyInstanceUID'].unique()).intersection(cases))","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:23:43.686448Z","iopub.execute_input":"2022-12-23T01:23:43.688529Z","iopub.status.idle":"2022-12-23T01:23:43.725691Z","shell.execute_reply.started":"2022-12-23T01:23:43.688484Z","shell.execute_reply":"2022-12-23T01:23:43.724417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ClassifierResNet3dCSN2P1D(nn.Module):\n    def __init__(\n        self, encoder='r50ir', pool='avg', norm_eval=False, num_classes=1\n    ) -> None:\n        super().__init__()\n\n        self.final = nn.Linear(2048, out_features=num_classes)\n        self.avg_pool = (\n            nn.AdaptiveAvgPool3d((1, 1, 1))\n            if pool == 'avg'\n            else nn.AdaptiveMaxPool3d((1, 1, 1))\n        )\n        self.dropout = nn.Dropout(0.5)\n        _initialize_weights(self)\n\n        self.backbone = ResNet3dCSN(\n            pretrained2d=False,\n            pretrained=None,\n            depth=int(encoder[1:-2]),\n            with_pool2=False,\n            bottleneck_mode=encoder[-2:],\n            norm_eval=norm_eval,\n            zero_init_residual=False,\n        )\n\n    def forward(self, x):\n        if x.size(1) == 1:\n            x = x.repeat(1, 3, 1, 1, 1)[:, :, :, :, :]\n        x = self.backbone(x)[-1]\n        x = self.avg_pool(x)\n        x = self.dropout(x)\n        x = x.flatten(1)\n        x = self.final(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:23:43.730246Z","iopub.execute_input":"2022-12-23T01:23:43.730661Z","iopub.status.idle":"2022-12-23T01:23:43.741443Z","shell.execute_reply.started":"2022-12-23T01:23:43.730626Z","shell.execute_reply":"2022-12-23T01:23:43.740242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"config_cls = {\n    'network': ClassifierResNet3dCSN2P1D,\n    'encoder_params': {\n        'encoder': 'r152ip',\n        'num_classes': 8,\n        'pool': 'max'\n    }\n}\ncls_model = load_model(\n    config_cls, '/kaggle/input/rsna-weights/swa_3_best_full_frozen_ClassifierResNet3dCSN2P1D_r152ip_1.pth'\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:23:43.743009Z","iopub.execute_input":"2022-12-23T01:23:43.743988Z","iopub.status.idle":"2022-12-23T01:23:47.146074Z","shell.execute_reply.started":"2022-12-23T01:23:43.743954Z","shell.execute_reply":"2022-12-23T01:23:47.144261Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Stage 1: Segmentation","metadata":{}},{"cell_type":"code","source":"sample = preprocess_volume_for_segmentation(test_dataset_dir, cases[0])\ninput_volume = sample['image'].unsqueeze(0)\n\nif DEBUG:\n    print(f'Input volume shape: {input_volume.shape}')\n\nplt.imshow(np.squeeze(input_volume[0, 0, 0, :, :]), cmap='bone')\nplt.show()\n\nh = int(sample['h'])\n\nif DEBUG:\n    print(f'H: {h}')\n\ncube_id = sample['cube_id']\nimgs = input_volume.cpu().float()\ncase_preds = np.zeros((imgs.shape[2], 256, 256), dtype=np.float32)\n\nwith torch.no_grad():\n    slices = get_slices(imgs, dim=2, window=256, overlap=128)\n    for slice in slices:\n        batch = imgs[:, :, slice.i_from: slice.i_to].float()\n        with torch.cuda.amp.autocast(enabled=True):\n            preds = torch.softmax(seg_model(batch)['mask'], dim=1)[0]\n            preds = torch.argmax(preds, dim=0)\n        preds = preds.cpu().numpy()\n        for pred_idx in range(slice.i_start, preds.shape[0]):\n            idx = slice.i_from + pred_idx\n            y_pred = preds[pred_idx]\n            case_preds[idx] = y_pred[:, :]\n        torch.cuda.empty_cache()\n\ncase_preds = np.array(case_preds)[:h]\ncase_preds = case_preds.astype(np.uint8)\nos.makedirs(preds_dir, exist_ok=True)\ntifffile.imwrite(os.path.join(preds_dir, f'{cube_id}.tif'), case_preds)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:23:47.147687Z","iopub.execute_input":"2022-12-23T01:23:47.148473Z","iopub.status.idle":"2022-12-23T01:27:10.295348Z","shell.execute_reply.started":"2022-12-23T01:23:47.148425Z","shell.execute_reply":"2022-12-23T01:27:10.294169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Stage 2: Classification","metadata":{}},{"cell_type":"code","source":"def predict_model(model, volume):\n    preds = []\n    with torch.no_grad():\n        for i in range(len(volume)):\n            with torch.cuda.amp.autocast():\n                output = model(volume[i])[0]\n            pred_slice = torch.sigmoid(\n                output.float()).cpu().numpy().astype(np.float32)\n            with torch.cuda.amp.autocast():\n                output = model(torch.flip(volume[i], dims=(-1,)))[0]\n            pred_slice += torch.sigmoid(output.float()\n                                        ).cpu().numpy().astype(np.float32)\n            pred_slice /= 2\n            preds.append(pred_slice)\n    preds = np.max(np.array(preds), axis=0)\n    preds[np.isnan(preds)] = 0.01\n    return preds\n\n\noutput_list = []\nsample = preprocess_volumes_for_classification(test_dataset_dir, cases[0])\ninput_volume_crop = sample['image']\n\nn_rows, n_cols = 2, 4\nfig, axs = plt.subplots(n_rows, n_cols, figsize=(20, 10))\nfig.suptitle('Input vertebrae', fontsize=16, weight='bold')\n\nfor i in range(n_rows):\n    for j in range(n_cols):\n        if i * n_cols + j >= len(input_volume_crop):\n            axs[i, j].axis('off')\n            break\n        axs[i, j].set_title(f'Vertebra {i * n_cols + j + 1}')\n        axs[i, j].imshow(\n            np.squeeze(input_volume_crop[i * n_cols + j][0, 0, 0, :, :].cpu()), cmap='bone'\n        )\nplt.show()\n\nwith torch.no_grad():\n    preds = predict_model(cls_model, input_volume_crop)\n    preds = np.clip(preds, 0.01, 0.99)\n    output_list.append([sample['cube_id'], preds])\n\nprint(f'Predictions: {output_list}')","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:27:10.297487Z","iopub.execute_input":"2022-12-23T01:27:10.297841Z","iopub.status.idle":"2022-12-23T01:30:27.470393Z","shell.execute_reply.started":"2022-12-23T01:27:10.297809Z","shell.execute_reply":"2022-12-23T01:30:27.469189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Stacking Cropped Volumes Back into Place","metadata":{}},{"cell_type":"code","source":"# given boxes used to obtain input_volume_crop from input_volume\n# put input_volume_crop back to input_volume\ninput_volume = combine_scan(os.path.join(test_dataset_dir, cases[0]), size=512)\ninput_volume = (input_volume - input_volume.mean()) / input_volume.std()\n\nre_input_volume = stack_cropped_volume_back_to_original_volume(\n    input_volume_crop,\n    input_volume,\n    cases[0]\n)\n\nfig, axes = plt.subplots(1, 3, figsize=(15, 5))\ndepth = 56\naxes[0].set_title('Cropped Input Volume')\naxes[0].imshow(np.squeeze(input_volume_crop[0][0, 0, 5, :, :]), cmap='bone')\naxes[1].set_title('Original Input Volume')\naxes[1].imshow(np.squeeze(input_volume[depth, :, :]), cmap='bone')\naxes[2].set_title('Recovered Input Volume')\naxes[2].imshow(np.squeeze(re_input_volume[depth, :, :]), cmap='bone')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:30:27.47225Z","iopub.execute_input":"2022-12-23T01:30:27.472732Z","iopub.status.idle":"2022-12-23T01:30:30.59881Z","shell.execute_reply.started":"2022-12-23T01:30:27.472687Z","shell.execute_reply":"2022-12-23T01:30:30.59742Z"},"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":"## Model Summary","metadata":{}},{"cell_type":"code","source":"from torchinfo import summary\n\nprint(input_volume_crop[0].shape)\n# save summary to file\nwith open('summary.txt', 'w') as f:\n    print(summary(cls_model, input_volume_crop[0].shape, verbose=0), file=f)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:30:30.600461Z","iopub.execute_input":"2022-12-23T01:30:30.601734Z","iopub.status.idle":"2022-12-23T01:30:40.429296Z","shell.execute_reply.started":"2022-12-23T01:30:30.601694Z","shell.execute_reply":"2022-12-23T01:30:40.42819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Class Activation Mapping\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 `summary()` to see the names of all layers in the model. These are necessary to get the value for `target_layers`.","metadata":{}},{"cell_type":"code","source":"# implement 3D Grad-CAM for ClassifierResNet3dCSN2P1D and plot the heatmap\ntarget_layers = [cls_model.backbone.layer4[-1]]  # last layer of the backbone\ncams = [\n    GradCAM(model=cls_model, target_layers=target_layers),\n    HiResCAM(model=cls_model, target_layers=target_layers),\n    GradCAMPlusPlus(model=cls_model, target_layers=target_layers),\n]\n\ncam_names = ['Grad-CAM', 'HiResCAM', 'Grad-CAM++']\n\n# each vertebra volume is passed to the model separately\ngrayscale_cams = [\n    [\n        cam(\n            input_tensor=input_volume_crop[i],\n            targets=[ClassifierOutputTarget(i)],\n        )[0]\n        for i in range(len(input_volume_crop))\n    ] for cam in cams]\n\nprint(f'CAMs shape: {np.array(grayscale_cams).shape}')","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:38:55.235106Z","iopub.execute_input":"2022-12-23T01:38:55.237528Z","iopub.status.idle":"2022-12-23T01:39:26.875676Z","shell.execute_reply.started":"2022-12-23T01:38:55.23743Z","shell.execute_reply":"2022-12-23T01:39:26.874137Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualizing the Heatmaps","metadata":{}},{"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    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        order=3,\n    )\n\n    return upscaled_heatmap\n\n\ngrayscale_cams = [\n    [\n        np.transpose(grayscale_cams[cam_idx][vertebra_idx], (2, 0, 1))\n        for vertebra_idx in range(len(grayscale_cams[cam_idx]))\n    ] for cam_idx in range(len(cam_names))\n]\n\nresized_cams = [\n    [\n        get_resized_heatmap(grayscale_cams[cam_idx][vertebra_idx],\n                            input_volume_crop[vertebra_idx][0, 0, :, :, :].shape)\n        for vertebra_idx in range(len(grayscale_cams[cam_idx]))\n    ] for cam_idx in range(len(cam_names))\n]","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:30:40.483197Z","iopub.status.idle":"2022-12-23T01:30:40.483741Z","shell.execute_reply.started":"2022-12-23T01:30:40.483523Z","shell.execute_reply":"2022-12-23T01:30:40.483544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"input_volume = combine_scan(os.path.join(test_dataset_dir, cases[0]), size=512)\n\nproper_cams = {\n    cam_name: stack_cropped_volume_back_to_original_volume(\n        resized_cams[cam_idx],\n        np.zeros(input_volume.shape),\n        cases[0]\n    ).numpy().astype(np.uint8) for cam_idx, cam_name in enumerate(cam_names)\n}","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:30:40.485037Z","iopub.status.idle":"2022-12-23T01:30:40.485464Z","shell.execute_reply.started":"2022-12-23T01:30:40.485238Z","shell.execute_reply":"2022-12-23T01:30:40.485256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualizations","metadata":{}},{"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    # p_heatmap = cv2.cvtColor(p_heatmap, cv2.COLOR_BGR2GRAY)\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(\n        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":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-23T01:30:40.487261Z","iopub.status.idle":"2022-12-23T01:30:40.487963Z","shell.execute_reply.started":"2022-12-23T01:30:40.487755Z","shell.execute_reply":"2022-12-23T01:30:40.487776Z"},"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(\n        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(\n        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(\n        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    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,\n                    step=1, value=int(vol_shape[0] / 2)),\n        y=IntSlider(min=0, max=vol_shape[1] - 1,\n                    step=1, value=int(vol_shape[1] / 2)),\n        x=IntSlider(min=0, max=vol_shape[2] - 1,\n                    step=1, value=int(vol_shape[2] / 2)),\n    )","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-23T01:30:40.489592Z","iopub.status.idle":"2022-12-23T01:30:40.49016Z","shell.execute_reply.started":"2022-12-23T01:30:40.489869Z","shell.execute_reply":"2022-12-23T01:30:40.489895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"interactive_show(input_volume, proper_cams[cam_names[0]])","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:30:40.491545Z","iopub.status.idle":"2022-12-23T01:30:40.492101Z","shell.execute_reply.started":"2022-12-23T01:30:40.491819Z","shell.execute_reply":"2022-12-23T01:30:40.491845Z"},"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, bboxes=False):\n    \"\"\"Create an animation of a volume\"\"\"\n    fig = plt.figure(figsize=(8, 8))\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            if bboxes:\n                # add bounding boxes to the heatmap image as animated patches\n                discovered_bboxes = get_bounding_boxes(heatmap[idx], otsu=True)\n                patches = get_bbox_patches(discovered_bboxes)\n                aux.extend(image_plot.axes.add_patch(patch)\n                           for patch in patches)\n            else:\n                image_plot2 = plt.imshow(\n                    heatmap[idx],\n                    animated=True,\n                    cmap='jet',\n                    alpha=alpha,\n                    extent=image_plot.get_extent(),\n                )\n                aux.append(image_plot2)\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=12)\n    ani = animation.ArtistAnimation(\n        fig, images, interval=5000 // len(array), blit=False, repeat_delay=1000\n    )\n    plt.close()\n    return ani","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-23T01:30:40.494351Z","iopub.status.idle":"2022-12-23T01:30:40.49522Z","shell.execute_reply.started":"2022-12-23T01:30:40.494923Z","shell.execute_reply":"2022-12-23T01:30:40.494951Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"create_animation(\n    input_volume,\n    f'{cases[0]} {cam_names[0]}',\n    heatmap=proper_cams[cam_names[0]],\n    bboxes=True,\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:30:40.498686Z","iopub.status.idle":"2022-12-23T01:30:40.499247Z","shell.execute_reply.started":"2022-12-23T01:30:40.498949Z","shell.execute_reply":"2022-12-23T01:30:40.498975Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"create_animation(\n    input_volume,\n    f'{cases[0]} {cam_names[1]}',\n    heatmap=proper_cams[cam_names[1]],\n    bboxes=True,\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:30:40.501049Z","iopub.status.idle":"2022-12-23T01:30:40.501854Z","shell.execute_reply.started":"2022-12-23T01:30:40.501534Z","shell.execute_reply":"2022-12-23T01:30:40.501571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"create_animation(\n    input_volume,\n    f'{cases[0]} {cam_names[2]}',\n    heatmap=proper_cams[cam_names[2]],\n    bboxes=True,\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:30:40.503796Z","iopub.status.idle":"2022-12-23T01:30:40.504489Z","shell.execute_reply.started":"2022-12-23T01:30:40.504126Z","shell.execute_reply":"2022-12-23T01:30:40.504153Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"data = []\nfor cube_id, preds in output_list:\n    data.append([f'{cube_id}_patient_overall', preds[0]])\n    data.extend([f'{cube_id}_C{i}', preds[i]] for i in range(1, 8))\npred_df = pd.DataFrame(data, columns=['row_id', 'fractured'])\npred_df.to_csv('submission.csv', index=False)\npred_df.head(8)","metadata":{"execution":{"iopub.status.busy":"2022-12-23T01:30:40.506183Z","iopub.status.idle":"2022-12-23T01:30:40.507069Z","shell.execute_reply.started":"2022-12-23T01:30:40.506849Z","shell.execute_reply":"2022-12-23T01:30:40.506871Z"},"trusted":true},"execution_count":null,"outputs":[]}]}