{"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":"## Install Packages","metadata":{}},{"cell_type":"code","source":"!cp -r /kaggle/input/python-packages /kaggle/working\n!pip install -q /kaggle/working/python-packages/pylibjpeg_libjpeg-1.3.2-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install -q /kaggle/working/python-packages/pylibjpeg_openjpeg-1.2.1-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install -q /kaggle/working/python-packages/pylibjpeg_rle-1.3.0-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install -q /kaggle/working/python-packages/iopath-0.1.9-py3-none-any.whl\n!pip install -q /kaggle/working/python-packages/av-9.2.0-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install -q /kaggle/working/python-packages/fvcore-0.1.5.post20220512/\n!pip install -q /kaggle/working/python-packages/parameterized-0.8.1-py2.py3-none-any.whl\n!pip install -q /kaggle/working/python-packages/pytorchvideo-0.1.5/\n!pip install -q /kaggle/working/python-packages/timm-0.6.7-py3-none-any.whl\n!pip install -q /kaggle/working/python-packages/antlr4-python3-runtime-4.9.3/\n!pip install -q /kaggle/working/python-packages/omegaconf-2.2.2-py3-none-any.whl\n!pip install -q /kaggle/working/python-packages/monai-0.8.1-202202162213-py3-none-any.whl\n\n!cp /kaggle/input/gdcm-conda-install/gdcm.tar /kaggle/working/\n!tar -xzvf gdcm.tar\n!conda install --offline /kaggle/working/gdcm/gdcm-2.8.9-py37h71b2a6d_0.tar.bz2","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-10-27T01:27:06.261702Z","iopub.execute_input":"2022-10-27T01:27:06.262726Z","iopub.status.idle":"2022-10-27T01:33:31.023609Z","shell.execute_reply.started":"2022-10-27T01:27:06.262621Z","shell.execute_reply":"2022-10-27T01:33:31.022372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Imports","metadata":{}},{"cell_type":"code","source":"import sys ; sys.path.insert(0, \"/kaggle/input/rsna-cspine-src/\")\n\nimport glob\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport os\nimport os.path as osp\nimport pandas as pd\nimport pydicom\nimport time\nimport torch\nimport torch.nn.functional as F\n\nfrom collections import defaultdict\nfrom omegaconf import OmegaConf\nfrom scipy.ndimage.interpolation import zoom \nfrom sklearn.metrics import roc_auc_score\nfrom skp import builder\nfrom tqdm import tqdm\n\ntorch.set_grad_enabled(False)","metadata":{"execution":{"iopub.status.busy":"2022-10-27T01:33:31.027084Z","iopub.execute_input":"2022-10-27T01:33:31.027463Z","iopub.status.idle":"2022-10-27T01:33:39.622044Z","shell.execute_reply.started":"2022-10-27T01:33:31.027428Z","shell.execute_reply":"2022-10-27T01:33:39.621095Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Helper Functions","metadata":{}},{"cell_type":"code","source":"def window(x, WL, WW):\n    upper, lower = WL+WW//2, WL-WW//2\n    x = np.clip(x, lower, upper)\n    x = x - lower\n    x = x / (upper - lower)\n    x = x * 255\n    x = x.astype('uint8')\n    return x\n\n\ndef load_dicom_volume(dicom_folder):\n    dicom_files = glob.glob(osp.join(dicom_folder, \"*.dcm\"))\n    dicoms = [pydicom.dcmread(_) for _ in dicom_files]\n    z_positions = [float(_.ImagePositionPatient[2]) for _ in dicoms]\n    dicom_arrays = [_.pixel_array.astype(\"float32\") for _ in dicoms]\n    rescale_slope = float(dicoms[0].RescaleSlope)\n    rescale_intercept = float(dicoms[0].RescaleIntercept)\n    del dicoms \n    \n    # Deal with potential scenario where not all arrays are the same shape\n    # This assumes that all arrays have the same number of dimensions (2)\n    array_shapes = np.vstack([_.shape for _ in dicom_arrays])\n    h, w = np.median(array_shapes[:,0]), np.median(array_shapes[:,1])\n    for ind, arr in enumerate(dicom_arrays):\n        if arr.shape[0] != h or arr.shape[1] != w:\n            print(\"Mismatched shape, resizing ...\")\n            scale_h, scale_w = float(h) / arr.shape[0], float(w) / arr.shape[1]\n            arr = zoom(arr, [scale_h, scale_w], order=1, prefilter=False)\n            dicom_arrays[ind] = arr\n    \n    array = np.stack(dicom_arrays)\n    del dicom_arrays \n    array = rescale_slope * array + rescale_intercept\n    array = window(array, WL=400, WW=2500)\n    \n    # Sort in DESCENDING order by z-position\n    array = array[np.argsort(z_positions)[::-1]]\n    return array\n\n\ndef plot_volume(img, num_row, num_col, sagittal=True):\n    index = 1\n    img_shape = img.shape[2] if sagittal else img.shape[0]\n    for i in np.linspace(0, img_shape - 1, num_row * num_col).astype(\"int\"):\n        plt.subplot(num_row, num_col, index)\n        plt.imshow(img[:, :, i] if sagittal else img[i], cmap=\"gray\")\n        index += 1\n    plt.show()\n\n        \ndef rescale(x):\n    # Rescale to [-1, 1]\n    x = x / x.max()\n    x = x - 0.5\n    x = x * 2\n    return x\n\n\ndef unscale(x):\n    x = x + 1\n    x = x * 255 / 2\n    return x\n\n\ndef load_models(config_file, checkpoint_folder, model_type=\"classification\", cuda=True, load_indices=None):\n    assert model_type in [\"classification\", \"segmentation\", \"sequence\", \"tdcnn\"]\n    config = OmegaConf.load(config_file)\n    if model_type == \"segmentation\":\n        config.model.params.encoder_params.pretrained = False\n    elif model_type == \"classification\":\n        config.model.params.pretrained = False\n    elif model_type == \"tdcnn\":\n        config.model.params.cnn_params.pretrained = False \n    checkpoints = np.sort(glob.glob(osp.join(checkpoint_folder, \"*\")))\n    if isinstance(load_indices, (list, tuple)):\n        load_indices = list(load_indices)\n        checkpoints = checkpoints[load_indices]\n    models = []\n    for each_checkpoint in checkpoints:\n        _config = config.copy()\n        _config.model.load_pretrained = str(each_checkpoint)\n        _model = builder.build_model(_config).eval()\n        if cuda:\n            _model = _model.cuda()\n        models.append(_model)\n    return models \n\n            \ndef add_buffer(x1, x2, max_dist, buff=0.1):\n    add_dist = int(buff * (x2 - x1))\n    x1, x2 = x1 - add_dist, x2 + add_dist\n    x1 = max(0, x1)\n    x2 = min(max_dist, x2)\n    return x1, x2\n\n\ndef segment_one_stage(volume, inference_shape, segmentation_models, threshold=0.4):\n    orig_shape = volume.shape\n    volume = F.interpolate(volume.unsqueeze(0).unsqueeze(0), size=inference_shape, mode=\"nearest\")\n    segmentation = torch.sigmoid(torch.cat([seg_model(volume.cuda()) for seg_model in segmentation_models])).mean(0)\n    \n    # Create a 1-channel cervical spine map \n    p_spine = segmentation.sum(0)\n    spine_map = torch.argmax(segmentation, dim=0) + 1\n    spine_map[p_spine < threshold] = 0 \n    spine_map[spine_map == 8] = 0 # Get rid of thoracic spine\n    spine_map = (spine_map * 255) / 7 # Rescale to an 8-bit image\n    \n    cspine_coords = {}\n    print(\"Obtaining cervical spine coordinates ...\")\n    for level in range(7):\n        coords = torch.stack(torch.where(segmentation[level] >= threshold)).cpu().numpy()\n        coords[0] = coords[0] * orig_shape[0] / inference_shape[0] \n        coords[1] = coords[1] * orig_shape[1] / inference_shape[1] \n        coords[2] = coords[2] * orig_shape[2] / inference_shape[2] \n        if coords.shape[1] == 0:\n            print(f\"Segmentation for C{level+1} failed !\")\n            cspine_coords[level] = None\n        else:\n            cspine_coords[level] = (coords[0].min(), coords[0].max(), coords[1].min(), coords[1].max(), coords[2].min(), coords[2].max())\n    return F.interpolate(spine_map.unsqueeze(0).unsqueeze(0), size=orig_shape, mode=\"nearest\").squeeze(0).squeeze(0), cspine_coords\n\n\ndef segment_two_stage(X, input_size1, input_size2, model_list1, model_list2, threshold=0.5, plot=False):\n    # Assumes X is rescaled, a torch tensor, and has dimensions (Z, H, W)\n    # X should NOT be resized\n    assert isinstance(model_list1, list) and isinstance(model_list2, list)\n    Z, H, W = X.size()\n    rescale_factors = [input_size1[0] / Z, input_size1[1] / H, input_size1[2] / W]\n    \n    tic = time.time()\n    # Run first model\n    pseg = torch.cat([model(F.interpolate(X.unsqueeze(0).unsqueeze(0), size=input_size1, mode=\"nearest\"))\n                      for model in model_list1]).mean(0)\n    pseg = torch.sigmoid(pseg)\n    # pseg.shape = (8, Z, H, W)\n    print(f\"1st segmentation model took {time.time() - tic:0.2f}s !\")\n    \n    # Get cervical spine boundaries and single-channel cervical spine level map\n    p_cervical = pseg[:7].sum(0)\n    spine_map = torch.argmax(pseg[:7], dim=0) + 1\n    spine_map[p_cervical < 0.5] = 0\n\n    # Get coordinates to crop input for second model\n    coords = np.vstack(np.where(p_cervical.cpu().numpy() >= 0.25))\n    # Get coordinates before rescaling for the spine map\n    # That way, we don't have to resize the spine map back to original and resize it again to\n    # model input size\n    z1, z2, h1, h2, w1, w2 = coords[0].min(), coords[0].max(), coords[1].min(), coords[1].max(),\\\n                             coords[2].min(), coords[2].max()\n    z1, z2 = add_buffer(z1, z2, input_size1[0], 0.1)\n    h1, h2 = add_buffer(h1, h2, input_size1[1], 0.1)\n    w1, w2 = add_buffer(w1, w2, input_size1[2], 0.1)\n\n    spine_map = (spine_map[z1:z2+1, h1:h2+1, w1:w2+1].unsqueeze(0).unsqueeze(0) * 255 / 7).long()\n\n    if plot:\n        plotseg = spine_map.squeeze(0).squeeze(0).cpu().numpy()\n        plot_volume(plotseg, 6, 4, sagittal=True)\n\n    # Then get coordinates after rescaling, to crop the original input\n    coords = coords.astype(\"float\")\n    coords[0] /= rescale_factors[0]\n    coords[1] /= rescale_factors[1]\n    coords[2] /= rescale_factors[2]\n    coords = coords.astype(\"int\")\n    z1, z2, h1, h2, w1, w2 = coords[0].min(), coords[0].max(), coords[1].min(), coords[1].max(),\\\n                             coords[2].min(), coords[2].max()\n    z1, z2 = add_buffer(z1, z2, Z, 0.1)\n    h1, h2 = add_buffer(h1, h2, H, 0.1)\n    w1, w2 = add_buffer(w1, w2, W, 0.1)\n    crop_cervical_coords = (z1, z2, h1, h2, w1, w2)\n\n    X = X[z1:z2+1, h1:h2+1, w1:w2+1].unsqueeze(0).unsqueeze(0)\n    Z, H, W = X.shape[2:]\n\n    if plot:\n        plotseg = X.squeeze(0).squeeze(0).cpu().numpy()\n        plot_volume(plotseg, 6, 4, sagittal=True)\n    \n    X = torch.cat([F.interpolate(X, size=input_size2, mode=\"nearest\"),\n                   F.interpolate(rescale(spine_map.float()), size=input_size2, mode=\"nearest\")], dim=1)\n    # X.shape = (1, 2, Z, H, W)\n    \n    \n    rescale_factors = [input_size2[0] / Z, input_size2[1] / H, input_size2[2] / W]\n    \n    tic = time.time()\n    # Run second model\n    pseg = torch.cat([model(X) for model in model_list2]).mean(0)\n    pseg = torch.sigmoid(pseg).cpu().numpy()\n    print(f\"2nd segmentation model took {time.time() - tic:0.2f}s !\")\n    \n    if plot:\n        plotseg = np.argmax(pseg, axis=0) + 1\n        plotseg[pseg.sum(0) < 0.5] = 0\n        plot_volume(plotseg, 6, 4, sagittal=True)\n\n    coords_dict = {}\n    for level in range(pseg.shape[0]):\n        coords = np.vstack(np.where(pseg[level] >= threshold))\n        if coords.shape[1] == 0:\n            coords_dict[level] = None\n            continue\n        coords = coords.astype(\"float\")\n        coords[0] /= rescale_factors[0]\n        coords[1] /= rescale_factors[1]\n        coords[2] /= rescale_factors[2]\n        coords = coords.astype(\"int\")\n        z1, z2, h1, h2, w1, w2 = coords[0].min(), coords[0].max(),\\\n                                 coords[1].min(), coords[1].max(),\\\n                                 coords[2].min(), coords[2].max()\n        coords_dict[level] = (z1, z2, h1, h2, w1, w2)\n\n    return coords_dict,\\\n           crop_cervical_coords\n\n\ndef center_crop(x, crop_size):\n    h, w = crop_size\n    orig_h, orig_w = x.shape[-2], x.shape[-1]\n    diff_h, diff_w = (orig_h - h) // 2, (orig_w - w) // 2\n    return x[..., diff_h:diff_h+h, diff_w:diff_w+w]","metadata":{"execution":{"iopub.status.busy":"2022-10-27T01:33:39.62363Z","iopub.execute_input":"2022-10-27T01:33:39.62453Z","iopub.status.idle":"2022-10-27T01:33:39.670974Z","shell.execute_reply.started":"2022-10-27T01:33:39.62449Z","shell.execute_reply":"2022-10-27T01:33:39.669885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load Models","metadata":{}},{"cell_type":"code","source":"segmentation_models0 = load_models(\"/kaggle/input/rsna-cspine-src/configs/seg/pseudoseg000.yaml\",\n                                   \"/kaggle/input/rsna-cspine-pseudoseg000/\",\n                                   model_type=\"segmentation\",\n                                   load_indices=[0, 1, 2])\n\nsegmentation_models1 = load_models(\"/kaggle/input/rsna-cspine-src/configs/seg/seg100.yaml\",\n                                   \"/kaggle/input/rsna-cspine-seg100/\",\n                                   model_type=\"segmentation\",\n                                   load_indices=[0, 1, 2])\n\nsegmentation_models2 = load_models(\"/kaggle/input/rsna-cspine-src/configs/seg/seg101.yaml\",\n                                   \"/kaggle/input/rsna-cspine-seg101/\",\n                                   model_type=\"segmentation\",\n                                   load_indices=[0, 1, 2])","metadata":{"execution":{"iopub.status.busy":"2022-10-27T01:33:39.673601Z","iopub.execute_input":"2022-10-27T01:33:39.674118Z","iopub.status.idle":"2022-10-27T01:33:54.09563Z","shell.execute_reply.started":"2022-10-27T01:33:39.674081Z","shell.execute_reply":"2022-10-27T01:33:54.094597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"x3d_feature_extractors = load_models(\"/kaggle/input/rsna-cspine-src/configs/chunk/chunk200.yaml\",\n                                     \"/kaggle/input/rsna-cspine-chunk200/\",\n                                     load_indices=[0, 1, 2])\n\nx3d_sequence_models = load_models(\"/kaggle/input/rsna-cspine-src/configs/chunkseq/chunkseq100.yaml\",\n                                  \"/kaggle/input/rsna-cspine-chunkseq100/\",\n                                  model_type=\"sequence\",\n                                  load_indices=[0, 1, 2])\n\n\ntdcnn_feature_extractors = load_models(\"/kaggle/input/rsna-cspine-src/configs/chunk/chunk101.yaml\",\n                                       \"/kaggle/input/rsna-cspine-chunk101/\",\n                                       model_type=\"tdcnn\",\n                                       load_indices=[0, 1, 2])\n\ntdcnn_sequence_models = load_models(\"/kaggle/input/rsna-cspine-src/configs/chunkseq/chunkseq005.yaml\",\n                                    \"/kaggle/input/rsna-cspine-chunkseq005/\",\n                                    model_type=\"sequence\",\n                                    load_indices=[0, 1, 2])\n\n\nfused_sequence_models = load_models(\"/kaggle/input/rsna-cspine-src/configs/chunkseq/chunkseq006.yaml\",\n                                    \"/kaggle/input/rsna-cspine-chunkseq300/\",\n                                    model_type=\"sequence\",\n                                    load_indices=[0, 1, 2])\n\n\ncnn2d_feature_extractors = load_models(\"/kaggle/input/rsna-cspine-src/configs/cas/cas001.yaml\",\n                                       \"/kaggle/input/rsna-cspine-cas001/\",\n                                       load_indices=[3, 4])\ncnn2d_sequence_models = load_models(\"/kaggle/input/rsna-cspine-src/configs/casseq/casseq004.yaml\",\n                                    \"/kaggle/input/rsna-cspine-casseq004/\",\n                                    model_type=\"sequence\",\n                                    load_indices=[3, 4])","metadata":{"execution":{"iopub.status.busy":"2022-10-27T01:33:54.096959Z","iopub.execute_input":"2022-10-27T01:33:54.098289Z","iopub.status.idle":"2022-10-27T01:34:26.078459Z","shell.execute_reply.started":"2022-10-27T01:33:54.098249Z","shell.execute_reply":"2022-10-27T01:34:26.077378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_images = glob.glob(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/test_images/*\")\n# train_folds = pd.read_csv(\"/kaggle/input/rsna-cspine-train-folds/train_kfold.csv\")\n# fold0 = train_folds[train_folds.outer == 0]\n# test_images = [osp.join(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images/\", _) for _ in fold0.StudyInstanceUID]\n# test_images = test_images\n# x3d_feature_extractors = x3d_feature_extractors[:1]\n# x3d_sequence_models = x3d_sequence_models[:1]\n\n# tdcnn_feature_extractors = tdcnn_feature_extractors[:1]\n# tdcnn_sequence_models = tdcnn_sequence_models[:1]\n\n# fused_sequence_models = fused_sequence_models[:1]\n\n# cnn2d_feature_extractors = cnn2d_feature_extractors[:1]\n# cnn2d_sequence_models = cnn2d_sequence_models[:1]\n\n\nthreshold1 = 0.4\nthreshold2 = 0.5\nsegmentation_inference_size = (192, 192, 192)\nx3d_inference_size = (64, 288, 288)\ntdcnn_inference_size = (32, 288, 288)\ncnn2d_inference_size = (640, 640)\ncnn2d_inference_crop_size = (560, 560)\ncnn2d_seq_len = 128 \n\n\nx3d_prediction_dict, tdcnn_prediction_dict, fused_prediction_dict, cnn2d_prediction_dict = {}, {}, {}, {}\n\nfor ind, each_image in tqdm(enumerate(test_images)):\n    total_tic = time.time()\n    \n    # LOAD DICOM VOLUME\n    print(\"Loading DICOM volume ...\")\n    tic = time.time()\n    X = load_dicom_volume(each_image)\n    print(f\"Took {time.time() - tic: 0.2f}s !\")\n    \n    X = rescale(X)\n    X = torch.from_numpy(X).float()\n    # X.shape = (num_images, height, width)\n    \n    # SEGMENT USING ONE-STAGE MODEL\n    print(\"Segmenting cervical spine (one-stage)...\")\n    tic = time.time()\n    spine_map, cspine_coords = segment_one_stage(X.cuda(), segmentation_inference_size, segmentation_models0, threshold=threshold1)\n    print(f\"Took {time.time() - tic: 0.2f}s !\")\n    \n    # EXTRACT CHUNK FEATURES FOR TD-CNN\n    print(\"Extracting vertebra chunk features for TD-CNN ...\")\n    tic = time.time() \n    tdcnn_features = defaultdict(list)\n    for level, coords in cspine_coords.items():\n        if not isinstance(coords, tuple):\n            print(f\"C{level+1} not found ... Using 0-vector ...\")\n            for fold, model in enumerate(tdcnn_feature_extractors):\n                tdcnn_features[fold].append(torch.zeros((1, 256)).float().cuda())\n        else:\n            x1, x2, y1, y2, z1, z2 = coords\n            orig_chunk = X[x1:x2, y1:y2, z1:z2].unsqueeze(0).unsqueeze(0)\n            chunk = F.interpolate(orig_chunk, size=tdcnn_inference_size, mode=\"trilinear\")\n            for fold, model in enumerate(tdcnn_feature_extractors):\n                tdcnn_features[fold].append(model.extract_features(chunk.cuda())) \n    print(f\"Took {time.time() - tic: 0.2f}s !\")\n    \n    # SEGMENT USING TWO-STAGE MODEL\n    print(\"Segmenting cervical spine (two-stage)...\")\n    tic = time.time()\n    cspine_coords, crop_cervical_coords = segment_two_stage(X.cuda(), segmentation_inference_size, segmentation_inference_size,\n                                                            segmentation_models1, segmentation_models2, threshold=threshold2,\n                                                            plot=False)\n    x1, x2, y1, y2, z1, z2 = crop_cervical_coords\n    X_crop = X[x1:x2+1, y1:y2+1, z1:z2+1]\n    print(f\"Took {time.time() - tic: 0.2f}s !\")\n\n    # EXTRACT CHUNK FEATURES FOR X3D\n    print(\"Extracting vertebra chunk features for X3D ...\")\n    tic = time.time()\n    x3d_features = defaultdict(list)\n    for level, coords in cspine_coords.items():\n        if not isinstance(coords, tuple):\n            print(f\"C{level+1} not found ... Using 0-vector ...\")\n            for fold, model in enumerate(x3d_feature_extractors):\n                x3d_features[fold].append(torch.zeros((1, 432)).float().cuda())\n        else:\n            x1, x2, y1, y2, z1, z2 = coords\n            orig_chunk = X_crop[x1:x2, y1:y2, z1:z2].unsqueeze(0).unsqueeze(0)\n            chunk = F.interpolate(orig_chunk, size=x3d_inference_size, mode=\"trilinear\")\n            for fold, model in enumerate(x3d_feature_extractors):\n                x3d_features[fold].append(model.extract_features(chunk.cuda())) \n    print(f\"Took {time.time() - tic: 0.2f}s !\")\n                 \n    # CONCATENATE CHUNK FEATURES FOR SEQUENCE MODELS\n    for fold, features in x3d_features.items():\n        x3d_features[fold] = (torch.cat(features).unsqueeze(0).cuda(), torch.ones((1, 7)).float().cuda())\n    for fold, features in tdcnn_features.items():\n        tdcnn_features[fold] = (torch.cat(features).unsqueeze(0).cuda(), torch.ones((1, 7)).float().cuda())\n        \n    fused_features = {}\n    for fold in [*x3d_features]:\n        fused_features[fold] = (torch.cat([x3d_features[fold][0], tdcnn_features[fold][0]], dim=-1), torch.ones((1, 7)).float().cuda())\n        \n    print(\"Chunk sequence inference ...\")\n    tic = time.time()\n    \n    x3d_pred_list = []\n    for fold, model in enumerate(x3d_sequence_models):\n        x3d_pred_list.append(torch.sigmoid(model(x3d_features[fold])).cpu().numpy())\n    x3d_prediction_dict[each_image] = np.mean(np.stack(x3d_pred_list, axis=0), axis=0)\n    \n    tdcnn_pred_list = []\n    for fold, model in enumerate(tdcnn_sequence_models):\n        tdcnn_pred_list.append(torch.sigmoid(model(tdcnn_features[fold])).cpu().numpy())\n    tdcnn_prediction_dict[each_image] = np.mean(np.stack(tdcnn_pred_list, axis=0), axis=0)\n\n            \n    fused_pred_list = []\n    for fold, model in enumerate(fused_sequence_models):\n        fused_pred_list.append(torch.sigmoid(model(fused_features[fold])).cpu().numpy())\n    fused_prediction_dict[each_image] = np.mean(np.stack(fused_pred_list, axis=0), axis=0)\n    \n    print(f\"Took {time.time() - tic: 0.2f}s !\")\n    \n    # EXTRACT FEATURES FOR 2D CNN\n    print(\"2D CNN + sequence inference ...\")\n    tic = time.time() \n    spine_present_on_slice = torch.where(spine_map.sum((1, 2)) > 0)[0]\n    start_slice, end_slice = spine_present_on_slice.min().item(), spine_present_on_slice.max().item()\n    spine_map = rescale(spine_map[start_slice:end_slice + 1]).unsqueeze(0).unsqueeze(0)\n    X = X.unsqueeze(0).unsqueeze(0).cuda()\n    X = torch.cat([X[:, :, start_slice:end_slice + 1], X[:, :, start_slice:end_slice + 1], spine_map], dim=1)\n    del spine_map\n    slice_indices = np.arange(X.size(2))\n    slice_indices = zoom(slice_indices, [cnn2d_seq_len / len(slice_indices)], order=0, prefilter=False)\n    X = center_crop(F.interpolate(X[:, :, slice_indices], \n                                  size=(cnn2d_seq_len, cnn2d_inference_size[0], cnn2d_inference_size[1]), \n                                  mode=\"trilinear\"), \n                    cnn2d_inference_crop_size)\n    X = X.squeeze(0).transpose(0, 1)\n    cnn2d_features = {}\n    for fold, model in enumerate(cnn2d_feature_extractors):\n        tmp_features = [model.extract_features(X[i:i+16]) for i in range(0, len(X), 16)]\n        tmp_features = torch.cat(tmp_features)\n        cnn2d_features[fold] = (tmp_features.unsqueeze(0).cuda(), torch.ones((1, len(tmp_features))).float().cuda())\n    cnn2d_pred_list = []\n    for fold, model in enumerate(cnn2d_sequence_models):\n        cnn2d_pred_list.append(torch.sigmoid(model(cnn2d_features[fold])).cpu().numpy())\n    cnn2d_prediction_dict[each_image] = np.mean(np.stack(cnn2d_pred_list, axis=0), axis=0)\n    print(f\"Took {time.time() - tic: 0.2f}s !\")\n    \n    print(f\"===TOTAL TIME ({ind + 1} / {len(test_images)}): {time.time() - total_tic:0.2f}s !\")","metadata":{"execution":{"iopub.status.busy":"2022-10-27T01:35:23.27629Z","iopub.execute_input":"2022-10-27T01:35:23.276765Z","iopub.status.idle":"2022-10-27T01:35:38.622061Z","shell.execute_reply.started":"2022-10-27T01:35:23.276725Z","shell.execute_reply":"2022-10-27T01:35:38.620411Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def competition_metric(p, t):\n    # p.shape = t.shape = (N, 8)\n    p = torch.from_numpy(p).float()\n    t = torch.from_numpy(t).float()\n    loss_matrix = F.binary_cross_entropy(p, t, reduction=\"none\")\n    # loss_matrix.shape = (N, 8)\n    columnwise_losses = []\n    for col in range(loss_matrix.shape[1]):\n        weights = t[:, col] + 1 # positives are weighted 2x\n        columnwise_losses.append(((loss_matrix[:, col] * weights).sum() / weights.sum()).item())\n    columnwise_losses[-1] *= 7.0\n    return np.sum(columnwise_losses) / 14.0\n\n\ndef auc(p, t):\n    if len(np.unique(t)) == 1:\n        return 0.5\n    return roc_auc_score(t, p)","metadata":{"execution":{"iopub.status.busy":"2022-10-26T22:45:54.239968Z","iopub.execute_input":"2022-10-26T22:45:54.240392Z","iopub.status.idle":"2022-10-26T22:45:54.250499Z","shell.execute_reply.started":"2022-10-26T22:45:54.24035Z","shell.execute_reply":"2022-10-26T22:45:54.248982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# study_id_list = []\n# predictions_list = []\n\n# for study_id, pred in x3d_prediction_dict.items():\n#     study_id_list.append(study_id.split(\"/\")[-1])\n#     predictions_list.append(pred)\n    \n# pred_df = pd.DataFrame(np.concatenate(predictions_list))\n# pred_df.columns = [f\"C{_+1}_pred\" for _ in range(7)] + [\"patient_overall_pred\"]\n# pred_df[\"StudyInstanceUID\"] = study_id_list\n\n# train_df = pd.read_csv(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n# pred_df = pred_df.merge(train_df, on=\"StudyInstanceUID\")\n# pred_df\n\n# t_columns = [f\"C{i+1}\" for i in range(7)] + [\"patient_overall\"]\n# p_columns = [c + \"_pred\" for c in t_columns]\n\n# print(f\"COMP. METRIC : {competition_metric(pred_df[p_columns].values, pred_df[t_columns].values):0.3f}\")\n\n# for i in range(len(t_columns)):\n#     prefix = \"AUC[overall] : \" if i == len(t_columns) - 1 else f\"AUC[C{i+1}]      : \"\n#     p = pred_df[p_columns[i]].values\n#     t = pred_df[t_columns[i]].values\n#     print(f\"{prefix}{auc(p, t):0.3f}\")","metadata":{"execution":{"iopub.status.busy":"2022-10-26T22:45:54.254062Z","iopub.execute_input":"2022-10-26T22:45:54.254593Z","iopub.status.idle":"2022-10-26T22:45:54.311041Z","shell.execute_reply.started":"2022-10-26T22:45:54.254549Z","shell.execute_reply":"2022-10-26T22:45:54.308926Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# study_id_list = []\n# predictions_list = []\n\n# for study_id, pred in tdcnn_prediction_dict.items():\n#     study_id_list.append(study_id.split(\"/\")[-1])\n#     predictions_list.append(pred)\n    \n# pred_df = pd.DataFrame(np.concatenate(predictions_list))\n# pred_df.columns = [f\"C{_+1}_pred\" for _ in range(7)] + [\"patient_overall_pred\"]\n# pred_df[\"StudyInstanceUID\"] = study_id_list\n\n# train_df = pd.read_csv(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n# pred_df = pred_df.merge(train_df, on=\"StudyInstanceUID\")\n# pred_df\n\n# t_columns = [f\"C{i+1}\" for i in range(7)] + [\"patient_overall\"]\n# p_columns = [c + \"_pred\" for c in t_columns]\n\n# print(f\"COMP. METRIC : {competition_metric(pred_df[p_columns].values, pred_df[t_columns].values):0.3f}\")\n\n# for i in range(len(t_columns)):\n#     prefix = \"AUC[overall] : \" if i == len(t_columns) - 1 else f\"AUC[C{i+1}]      : \"\n#     p = pred_df[p_columns[i]].values\n#     t = pred_df[t_columns[i]].values\n#     print(f\"{prefix}{auc(p, t):0.3f}\")","metadata":{"execution":{"iopub.status.busy":"2022-10-26T22:45:54.312867Z","iopub.execute_input":"2022-10-26T22:45:54.313353Z","iopub.status.idle":"2022-10-26T22:45:54.344771Z","shell.execute_reply.started":"2022-10-26T22:45:54.313311Z","shell.execute_reply":"2022-10-26T22:45:54.343422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# study_id_list = []\n# predictions_list = []\n\n# for study_id, pred in fused_prediction_dict.items():\n#     study_id_list.append(study_id.split(\"/\")[-1])\n#     predictions_list.append(pred)\n    \n# pred_df = pd.DataFrame(np.concatenate(predictions_list))\n# pred_df.columns = [f\"C{_+1}_pred\" for _ in range(7)] + [\"patient_overall_pred\"]\n# pred_df[\"StudyInstanceUID\"] = study_id_list\n\n# train_df = pd.read_csv(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n# pred_df = pred_df.merge(train_df, on=\"StudyInstanceUID\")\n# pred_df\n\n# t_columns = [f\"C{i+1}\" for i in range(7)] + [\"patient_overall\"]\n# p_columns = [c + \"_pred\" for c in t_columns]\n\n# print(f\"COMP. METRIC : {competition_metric(pred_df[p_columns].values, pred_df[t_columns].values):0.3f}\")\n\n# for i in range(len(t_columns)):\n#     prefix = \"AUC[overall] : \" if i == len(t_columns) - 1 else f\"AUC[C{i+1}]      : \"\n#     p = pred_df[p_columns[i]].values\n#     t = pred_df[t_columns[i]].values\n#     print(f\"{prefix}{auc(p, t):0.3f}\")","metadata":{"execution":{"iopub.status.busy":"2022-10-26T22:45:54.347555Z","iopub.execute_input":"2022-10-26T22:45:54.348913Z","iopub.status.idle":"2022-10-26T22:45:54.383359Z","shell.execute_reply.started":"2022-10-26T22:45:54.348843Z","shell.execute_reply":"2022-10-26T22:45:54.382258Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# study_id_list = []\n# predictions_list = []\n\n# for study_id, pred in cnn2d_prediction_dict.items():\n#     study_id_list.append(study_id.split(\"/\")[-1])\n#     predictions_list.append(pred)\n    \n# pred_df = pd.DataFrame(np.concatenate(predictions_list))\n# pred_df.columns = [f\"C{_+1}_pred\" for _ in range(7)] + [\"patient_overall_pred\"]\n# pred_df[\"StudyInstanceUID\"] = study_id_list\n\n# train_df = pd.read_csv(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n# pred_df = pred_df.merge(train_df, on=\"StudyInstanceUID\")\n# pred_df\n\n# t_columns = [f\"C{i+1}\" for i in range(7)] + [\"patient_overall\"]\n# p_columns = [c + \"_pred\" for c in t_columns]\n\n# print(f\"COMP. METRIC : {competition_metric(pred_df[p_columns].values, pred_df[t_columns].values):0.3f}\")\n\n# for i in range(len(t_columns)):\n#     prefix = \"AUC[overall] : \" if i == len(t_columns) - 1 else f\"AUC[C{i+1}]      : \"\n#     p = pred_df[p_columns[i]].values\n#     t = pred_df[t_columns[i]].values\n#     print(f\"{prefix}{auc(p, t):0.3f}\")","metadata":{"execution":{"iopub.status.busy":"2022-10-26T22:45:54.385416Z","iopub.execute_input":"2022-10-26T22:45:54.385929Z","iopub.status.idle":"2022-10-26T22:45:54.421666Z","shell.execute_reply.started":"2022-10-26T22:45:54.385884Z","shell.execute_reply":"2022-10-26T22:45:54.420241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ensemble_pred_dict = {}\n# x3d_weight, tdcnn_weight, fused_weight, cnn2d_weight = 0.25, 0.25, 0.25, 0.25\n# for study_id in [*x3d_prediction_dict]:\n#     ensemble_pred_dict[study_id.split(\"/\")[-1]] = x3d_weight * x3d_prediction_dict[study_id] + \\\n#                                                   tdcnn_weight * tdcnn_prediction_dict[study_id] + \\\n#                                                   fused_weight * fused_prediction_dict[study_id] + \\\n#                                                   cnn2d_weight * cnn2d_prediction_dict[study_id]\n\n# study_id_list = []\n# predictions_list = []\n\n# for study_id, pred in ensemble_pred_dict.items():\n#     study_id_list.append(study_id.split(\"/\")[-1])\n#     predictions_list.append(pred)\n    \n# pred_df = pd.DataFrame(np.concatenate(predictions_list))\n# pred_df.columns = [f\"C{_+1}_pred\" for _ in range(7)] + [\"patient_overall_pred\"]\n# pred_df[\"StudyInstanceUID\"] = study_id_list\n\n# train_df = pd.read_csv(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n# pred_df = pred_df.merge(train_df, on=\"StudyInstanceUID\")\n# pred_df\n\n# t_columns = [f\"C{i+1}\" for i in range(7)] + [\"patient_overall\"]\n# p_columns = [c + \"_pred\" for c in t_columns]\n\n# print(f\"COMP. METRIC : {competition_metric(pred_df[p_columns].values, pred_df[t_columns].values):0.3f}\")\n\n# for i in range(len(t_columns)):\n#     prefix = \"AUC[overall] : \" if i == len(t_columns) - 1 else f\"AUC[C{i+1}]      : \"\n#     p = pred_df[p_columns[i]].values\n#     t = pred_df[t_columns[i]].values\n#     print(f\"{prefix}{auc(p, t):0.3f}\")","metadata":{"execution":{"iopub.status.busy":"2022-10-27T00:34:07.351613Z","iopub.execute_input":"2022-10-27T00:34:07.35384Z","iopub.status.idle":"2022-10-27T00:34:07.853646Z","shell.execute_reply.started":"2022-10-27T00:34:07.353802Z","shell.execute_reply":"2022-10-27T00:34:07.851681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# ensemble_pred_dict = {}\n# x3d_weight, tdcnn_weight, fused_weight, cnn2d_weight = 0.3, 0.3, 0.3, 0.1\n# for study_id in [*x3d_prediction_dict]:\n#     ensemble_pred_dict[study_id.split(\"/\")[-1]] = x3d_weight * x3d_prediction_dict[study_id] + \\\n#                                                   tdcnn_weight * tdcnn_prediction_dict[study_id] + \\\n#                                                   fused_weight * fused_prediction_dict[study_id] + \\\n#                                                   cnn2d_weight * cnn2d_prediction_dict[study_id]\n\n# study_id_list = []\n# predictions_list = []\n\n# for study_id, pred in ensemble_pred_dict.items():\n#     study_id_list.append(study_id.split(\"/\")[-1])\n#     predictions_list.append(pred)\n    \n# pred_df = pd.DataFrame(np.concatenate(predictions_list))\n# pred_df.columns = [f\"C{_+1}_pred\" for _ in range(7)] + [\"patient_overall_pred\"]\n# pred_df[\"StudyInstanceUID\"] = study_id_list\n\n# train_df = pd.read_csv(\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n# pred_df = pred_df.merge(train_df, on=\"StudyInstanceUID\")\n# pred_df\n\n# t_columns = [f\"C{i+1}\" for i in range(7)] + [\"patient_overall\"]\n# p_columns = [c + \"_pred\" for c in t_columns]\n\n# print(f\"COMP. METRIC : {competition_metric(pred_df[p_columns].values, pred_df[t_columns].values):0.3f}\")\n\n# for i in range(len(t_columns)):\n#     prefix = \"AUC[overall] : \" if i == len(t_columns) - 1 else f\"AUC[C{i+1}]      : \"\n#     p = pred_df[p_columns[i]].values\n#     t = pred_df[t_columns[i]].values\n#     print(f\"{prefix}{auc(p, t):0.3f}\")","metadata":{"execution":{"iopub.status.busy":"2022-10-26T22:47:09.14062Z","iopub.execute_input":"2022-10-26T22:47:09.141336Z","iopub.status.idle":"2022-10-26T22:47:09.175958Z","shell.execute_reply.started":"2022-10-26T22:47:09.1413Z","shell.execute_reply":"2022-10-26T22:47:09.174759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Create Submission DataFrame","metadata":{}},{"cell_type":"code","source":"ensemble_pred_dict = {}\nx3d_weight, tdcnn_weight, fused_weight, cnn2d_weight = 0.3, 0.3, 0.3, 0.1\nfor study_id in [*x3d_prediction_dict]:\n    ensemble_pred_dict[study_id.split(\"/\")[-1]] = x3d_weight * x3d_prediction_dict[study_id] + \\\n                                                  tdcnn_weight * tdcnn_prediction_dict[study_id] + \\\n                                                  fused_weight * fused_prediction_dict[study_id] + \\\n                                                  cnn2d_weight * cnn2d_prediction_dict[study_id]\n\nrow_id_list = []\nfractured_list = []\nfor k, v in ensemble_pred_dict.items():\n    for label_ind, label in enumerate(v[0]):\n        row_id = f\"{k}_C{label_ind + 1}\" if label_ind < 7 else f\"{k}_patient_overall\"\n        row_id_list.append(row_id)\n        fractured_list.append(label)\n        \nsub_df = pd.DataFrame({\"row_id\": row_id_list, \"fractured\": fractured_list})\nsub_df","metadata":{"execution":{"iopub.status.busy":"2022-10-27T00:35:08.575586Z","iopub.execute_input":"2022-10-27T00:35:08.576019Z","iopub.status.idle":"2022-10-27T00:35:08.597124Z","shell.execute_reply.started":"2022-10-27T00:35:08.57598Z","shell.execute_reply":"2022-10-27T00:35:08.595722Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!rm -rf /kaggle/working/*\nsub_df.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2022-10-27T00:35:11.824626Z","iopub.execute_input":"2022-10-27T00:35:11.825026Z","iopub.status.idle":"2022-10-27T00:35:13.011389Z","shell.execute_reply.started":"2022-10-27T00:35:11.824986Z","shell.execute_reply":"2022-10-27T00:35:13.010042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#plot_volume(spine_map.numpy(), sagittal=True)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}