{"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":"👋 In this kernel we will perform Geras et al. (2018) BIRADS classifier inference on 16-bit versions of the competition images. This notebook borrows from @hengck23's tensorRT example and Kaggle community research on reading `DICOM` files.","metadata":{}},{"cell_type":"markdown","source":"<div class=\"alert alert-block alert-success\">\n    <b>Data source:</b> Radiological Society of North America. (2022, Nov 29). RSNA Screening Mammography Breast Cancer Detection, Version 1. Retrieved 2023 Feb 9 from [https://www.kaggle.com/competitions/rsna-breast-cancer-detection/data].\n</div>","metadata":{}},{"cell_type":"markdown","source":"# Initial imports","metadata":{}},{"cell_type":"code","source":"import os\nos.environ['CUDA_MODULE_LOADING']='LAZY'\nimport sys\nimport glob\n\nimport gc\n\nfrom pathlib import Path\n\nfrom timeit import default_timer as timer\nfrom joblib import Parallel, delayed\n\nfrom tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:05.352243Z","iopub.execute_input":"2023-02-18T00:34:05.352673Z","iopub.status.idle":"2023-02-18T00:34:05.539841Z","shell.execute_reply.started":"2023-02-18T00:34:05.352589Z","shell.execute_reply":"2023-02-18T00:34:05.538708Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\n\nimport matplotlib\nimport matplotlib.pyplot as plt\n\nimport cv2\n\nfrom sklearn import metrics","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:05.542128Z","iopub.execute_input":"2023-02-18T00:34:05.54257Z","iopub.status.idle":"2023-02-18T00:34:06.781605Z","shell.execute_reply.started":"2023-02-18T00:34:05.542528Z","shell.execute_reply":"2023-02-18T00:34:06.780337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nprint('torch.__version__ = %s'%torch.__version__)\nprint('CUDA:',torch.version.cuda)\nfrom torch.utils.data.dataset import Dataset\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data.sampler import *\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.cuda.amp as amp\nprint('torch.cuda.device_count() = %d'%torch.cuda.device_count())\nprint('torch.cuda.get_device_properties() = %s'%str(torch.cuda.get_device_properties(0))[21:])\nprint('torch.backends.cudnn.version() = %s'%str(torch.backends.cudnn.version()))\n\nimport torchvision.transforms as transforms","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:06.788071Z","iopub.execute_input":"2023-02-18T00:34:06.790836Z","iopub.status.idle":"2023-02-18T00:34:09.282913Z","shell.execute_reply.started":"2023-02-18T00:34:06.790783Z","shell.execute_reply":"2023-02-18T00:34:09.281613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Install dependencies","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/rsna-2022-whl/pylibjpeg-1.4.0-py3-none-any.whl\n!pip install /kaggle/input/rsna-2022-whl/python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl\n!pip install /kaggle/input/rsnamodules/dicomsdl-0.109.1-cp37-cp37m-manylinux_2_12_x86_64.manylinux2010_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:09.288627Z","iopub.execute_input":"2023-02-18T00:34:09.291923Z","iopub.status.idle":"2023-02-18T00:34:47.596887Z","shell.execute_reply.started":"2023-02-18T00:34:09.291877Z","shell.execute_reply":"2023-02-18T00:34:47.595547Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"https://gitlab.com/drj11/pypng is licensed with the MIT licence.","metadata":{}},{"cell_type":"code","source":"!pip install /kaggle/input/pypng-0202207150-py3-none-any/wheelhouse/pypng-0.20220715.0-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:47.598813Z","iopub.execute_input":"2023-02-18T00:34:47.599783Z","iopub.status.idle":"2023-02-18T00:34:57.940889Z","shell.execute_reply.started":"2023-02-18T00:34:47.599743Z","shell.execute_reply":"2023-02-18T00:34:57.939597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"https://github.com/louis-she/nvjpeg2k-python is published with a MIT license.","metadata":{}},{"cell_type":"code","source":"!cp /kaggle/input/nvjpeg2k/nvjpeg2k.so ./","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:57.94417Z","iopub.execute_input":"2023-02-18T00:34:57.944651Z","iopub.status.idle":"2023-02-18T00:34:59.076842Z","shell.execute_reply.started":"2023-02-18T00:34:57.9446Z","shell.execute_reply":"2023-02-18T00:34:59.07535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sys.path.append('/kaggle/input/nvjpeg2k/')","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:59.079215Z","iopub.execute_input":"2023-02-18T00:34:59.079695Z","iopub.status.idle":"2023-02-18T00:34:59.08516Z","shell.execute_reply.started":"2023-02-18T00:34:59.079651Z","shell.execute_reply":"2023-02-18T00:34:59.084003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Initialize utilities","metadata":{}},{"cell_type":"code","source":"def time_to_str(t, mode='min'):\n    if mode=='min':\n        t  = int(t)/60\n        hr = t//60\n        min = t%60\n        return '%2d hr %02d min'%(hr,min)\n\n    elif mode=='sec':\n        t   = int(t)\n        min = t//60\n        sec = t%60\n        return '%2d min %02d sec'%(min,sec)\n\n    else:\n        raise NotImplementedError","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:59.087021Z","iopub.execute_input":"2023-02-18T00:34:59.087914Z","iopub.status.idle":"2023-02-18T00:34:59.097315Z","shell.execute_reply.started":"2023-02-18T00:34:59.087874Z","shell.execute_reply":"2023-02-18T00:34:59.096195Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data processing","metadata":{}},{"cell_type":"markdown","source":"Adapted from https://www.kaggle.com/code/hengck23/3hr-tensorrt-nextvit-example","metadata":{}},{"cell_type":"code","source":"import pydicom\nimport dicomsdl\nimport nvjpeg2k\nimport png","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:59.099047Z","iopub.execute_input":"2023-02-18T00:34:59.099786Z","iopub.status.idle":"2023-02-18T00:34:59.461354Z","shell.execute_reply.started":"2023-02-18T00:34:59.099749Z","shell.execute_reply":"2023-02-18T00:34:59.460366Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Load metadata and images","metadata":{}},{"cell_type":"code","source":"dcm_loc = Path('/kaggle/input/rsna-breast-cancer-detection/train_images/')","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:59.465714Z","iopub.execute_input":"2023-02-18T00:34:59.466024Z","iopub.status.idle":"2023-02-18T00:34:59.470969Z","shell.execute_reply.started":"2023-02-18T00:34:59.465996Z","shell.execute_reply":"2023-02-18T00:34:59.469803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"csv_file = Path('/kaggle/input/rsna-breast-cancer-detection/train.csv')\ntrain_meta = pd.read_csv(csv_file)\ntrain_meta.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:59.47268Z","iopub.execute_input":"2023-02-18T00:34:59.473347Z","iopub.status.idle":"2023-02-18T00:34:59.608627Z","shell.execute_reply.started":"2023-02-18T00:34:59.47331Z","shell.execute_reply":"2023-02-18T00:34:59.607577Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_meta['prediction_id'] = train_meta['patient_id'].astype(str) + \"_\" + train_meta['laterality'].astype(str)\ntrain_meta.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:59.610224Z","iopub.execute_input":"2023-02-18T00:34:59.610884Z","iopub.status.idle":"2023-02-18T00:34:59.678029Z","shell.execute_reply.started":"2023-02-18T00:34:59.610841Z","shell.execute_reply":"2023-02-18T00:34:59.676974Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Add column `projection`\n\nFor convinience:","metadata":{}},{"cell_type":"code","source":"train_meta['projection'] = train_meta[['laterality', 'view']].astype(str).apply('-'.join, axis=1)\ntrain_meta.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:34:59.681522Z","iopub.execute_input":"2023-02-18T00:34:59.681852Z","iopub.status.idle":"2023-02-18T00:35:00.139446Z","shell.execute_reply.started":"2023-02-18T00:34:59.68182Z","shell.execute_reply":"2023-02-18T00:35:00.138246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.DataFrame()\ntrain_df = train_meta.loc[(train_meta['projection'] == 'R-MLO') | (train_meta['projection'] == 'L-MLO') | (train_meta['projection'] == 'R-CC') | (train_meta['projection'] == 'L-CC')]\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:35:00.141059Z","iopub.execute_input":"2023-02-18T00:35:00.141448Z","iopub.status.idle":"2023-02-18T00:35:00.195051Z","shell.execute_reply.started":"2023-02-18T00:35:00.141415Z","shell.execute_reply":"2023-02-18T00:35:00.193806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Initialize `nvjpeg2k` decoder","metadata":{}},{"cell_type":"code","source":"j2k_decoder = nvjpeg2k.Decoder()","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:35:00.197041Z","iopub.execute_input":"2023-02-18T00:35:00.19752Z","iopub.status.idle":"2023-02-18T00:35:03.737689Z","shell.execute_reply.started":"2023-02-18T00:35:00.197473Z","shell.execute_reply":"2023-02-18T00:35:03.736468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define `DICOM` processing functions","metadata":{}},{"cell_type":"code","source":"def _process_j2k(df, dcm_dir, image_dir, is_voi_lut=True):\n    dcm_file = f'{dcm_dir}/{d.patient_id.values[0]}/{d.image_id.values[0]}.dcm'\n    dc = pydicom.dcmread(dcm_file)\n        \n    height = dc.Rows\n    width = dc.Columns\n        \n    offset = dc.PixelData.find(b'\\x00\\x00\\x00\\x0C')\n    jpeg_stream = bytearray(dc.PixelData[offset:])\n    image = j2k_decoder.decode(jpeg_stream)\n\n    if is_voi_lut:\n        image = pydicom.pixel_data_handlers.util.apply_voi_lut(image, dc)\n            \n    image = image.reshape((height, width))  # (height, width)\n        \n    if dc.PhotometricInterpretation == 'MONOCHROME1':  # ranges from bright to dark with ascending pixel values\n        image = image.max() - image\n    elif dc.PhotometricInterpretation == 'MONOCHROME2':  # ranges from dark to bright with ascending pixel values\n        pass\n        \n    laterality = dc.ImageLaterality\n    view = d.view.values[0]\n        \n    bitdepth = dc.BitsAllocated\n    # bitsstored = dc.BitsStored\n\n    image = image.astype(np.float64)\n    image /= image.max()\n    image *= pow(2, bitdepth) - 1\n    image = image.astype(np.uint16)\n        \n    if laterality == 'R':\n        image = np.fliplr(image)\n        \n    return image","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:50:55.113034Z","iopub.execute_input":"2023-02-18T00:50:55.113477Z","iopub.status.idle":"2023-02-18T00:50:55.126107Z","shell.execute_reply.started":"2023-02-18T00:50:55.11344Z","shell.execute_reply":"2023-02-18T00:50:55.12491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dicomsdl_to_numpy_image(ds, index=0):\n    # https://stackoverflow.com/questions/44659924/returning-numpy-arrays-via-pybind11\n    info = ds.getPixelDataInfo()\n    if info['SamplesPerPixel'] != 1:\n        raise RuntimeError('SamplesPerPixel != 1')\n\n    shape = [info['Rows'], info['Cols']]\n    dtype = info['dtype']\n    outarr = np.empty(shape, dtype=dtype)\n    ds.copyFrameData(index, outarr)\n    return outarr","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:35:03.753288Z","iopub.execute_input":"2023-02-18T00:35:03.753893Z","iopub.status.idle":"2023-02-18T00:35:03.762934Z","shell.execute_reply.started":"2023-02-18T00:35:03.753855Z","shell.execute_reply":"2023-02-18T00:35:03.761755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def _process_non_j2k(d, dcm_dir, image_dir, is_voi_lut=True):\n    dcm_file = f'{dcm_dir}/{d.patient_id.values[0]}/{d.image_id.values[0]}.dcm'\n    ds = dicomsdl.open(dcm_file)\n    \n    height = ds.Rows\n    width = ds.Columns\n    \n    image = dicomsdl_to_numpy_image(ds)\n\n    if is_voi_lut:\n        dc = pydicom.dcmread(dcm_file)\n        image = pydicom.pixel_data_handlers.util.apply_voi_lut(image, dc)\n\n    image = image.reshape((height, width))  # (height, width)\n        \n    if ds.PhotometricInterpretation == 'MONOCHROME1':  # ranges from bright to dark with ascending pixel values\n        image = image.max() - image\n    elif ds.PhotometricInterpretation == 'MONOCHROME2':  # ranges from dark to bright with ascending pixel values\n        pass\n        \n    laterality = dc.ImageLaterality\n    view = d.view.values[0]\n        \n    bitdepth = ds.BitsAllocated\n    # bitsstored = ds.BitsStored\n\n    image = image.astype(np.float64)\n    image /= image.max()\n    image *= pow(2, bitdepth) - 1\n    image = image.astype(np.uint16)\n    \n    if laterality == 'R':\n        image = np.fliplr(image)\n        \n    return image","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:50:57.9803Z","iopub.execute_input":"2023-02-18T00:50:57.980664Z","iopub.status.idle":"2023-02-18T00:50:57.989356Z","shell.execute_reply.started":"2023-02-18T00:50:57.980634Z","shell.execute_reply":"2023-02-18T00:50:57.988045Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Define function to retrieve `TransferSyntaxUID`","metadata":{}},{"cell_type":"code","source":"def _get_transfer_syntax_uid(df, dcm_dir):\n    f = f'{dcm_dir}/{d.patient_id.values[0]}/{d.image_id.values[0]}.dcm'\n    ds = pydicom.dcmread(f)\n    transfer_uid = ds.file_meta.TransferSyntaxUID\n    return transfer_uid","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:35:03.77935Z","iopub.execute_input":"2023-02-18T00:35:03.779627Z","iopub.status.idle":"2023-02-18T00:35:03.78812Z","shell.execute_reply.started":"2023-02-18T00:35:03.779601Z","shell.execute_reply":"2023-02-18T00:35:03.786762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create `exam_list`\n\nCreate a new `DataFrame` with `patient_id` and number of images (i.e., standard view representatives) for that particular examination.","metadata":{}},{"cell_type":"code","source":"exam_list = pd.DataFrame()\nexam_list = train_df[\"patient_id\"].value_counts().rename_axis('patient_id').reset_index(name='file_count')\nexam_list.head()","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:35:44.653356Z","iopub.execute_input":"2023-02-18T00:35:44.653738Z","iopub.status.idle":"2023-02-18T00:35:44.673827Z","shell.execute_reply.started":"2023-02-18T00:35:44.653705Z","shell.execute_reply":"2023-02-18T00:35:44.672858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"exam_list_sel = exam_list.loc[exam_list['file_count'] >= 4]\nexam_list_sel","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:35:48.09677Z","iopub.execute_input":"2023-02-18T00:35:48.097465Z","iopub.status.idle":"2023-02-18T00:35:48.120307Z","shell.execute_reply.started":"2023-02-18T00:35:48.097414Z","shell.execute_reply":"2023-02-18T00:35:48.119096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Implement PyTorch `Dataset`","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Dataset","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:35:51.432214Z","iopub.execute_input":"2023-02-18T00:35:51.432621Z","iopub.status.idle":"2023-02-18T00:35:51.438316Z","shell.execute_reply.started":"2023-02-18T00:35:51.432586Z","shell.execute_reply":"2023-02-18T00:35:51.436816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"INPUT_SIZE_H = 2600\nINPUT_SIZE_W = 2000","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:35:53.536657Z","iopub.execute_input":"2023-02-18T00:35:53.537073Z","iopub.status.idle":"2023-02-18T00:35:53.542019Z","shell.execute_reply.started":"2023-02-18T00:35:53.537038Z","shell.execute_reply":"2023-02-18T00:35:53.540886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def pad_right_bottom(img):\n    h, w = img.shape[:2]\n    res = np.zeros((INPUT_SIZE_H, INPUT_SIZE_W))\n    res[:h, :w] = img\n    return res","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:35:55.332768Z","iopub.execute_input":"2023-02-18T00:35:55.333259Z","iopub.status.idle":"2023-02-18T00:35:55.33985Z","shell.execute_reply.started":"2023-02-18T00:35:55.333188Z","shell.execute_reply":"2023-02-18T00:35:55.338744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from PIL import Image","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:35:57.612648Z","iopub.execute_input":"2023-02-18T00:35:57.613103Z","iopub.status.idle":"2023-02-18T00:35:57.618312Z","shell.execute_reply.started":"2023-02-18T00:35:57.613059Z","shell.execute_reply":"2023-02-18T00:35:57.617383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    def __init__(self, split, metadata, phase='test'):\n        self.split = split  # train_split, val_split, or hidden_test_split\n        self.length = len(split)\n        self.metadata = metadata\n        self.phase = phase  # train, val, or test\n\n    def __getitem__(self, index):\n        if isinstance(index, torch.Tensor):\n            index = index.item()\n        \n        s = self.split.iloc[index]\n        \n        patient_id = s.patient_id\n        \n        file_df = self.metadata.loc[self.metadata['patient_id'] == int(patient_id)]\n        file_dict = dict(zip(file_df['projection'], file_df['image_id']))\n        \n        # We only deal with examinations with four standard views\n        if not all(k in file_dict for k in ('R-MLO', 'L-MLO', 'R-CC', 'L-CC')):\n            return {'inputs': {}, 'targets': {}, 'file_dict': file_dict, 'index': index, 'patient_id': patient_id}\n                \n        if not self.phase == 'test':\n            tg_dict = dict(zip(file_df['projection'], file_df['cancer'])) \n        else:\n            tg_dict = {}\n        \n        projections = ['R-MLO', 'L-MLO', 'R-CC', 'L-CC']\n        \n        im_dict = {projection: [] for projection in projections}\n        for projection in projections:\n            image_id = file_dict[projection]\n            \n            # Get TransferSyntaxUID\n            try:\n                trfx_uid = _get_transfer_syntax_uid(file_df, dcm_loc)\n            except:\n                trfx_uid = None\n            \n            im_dir = dcm_loc / str(patient_id)  # dcm_loc is pathlib.Path()\n            \n            # Read DICOM\n            if trfx_uid is not None and '1.2.840.10008.1.2.4.70' == str(trfx_uid):\n                im = _process_jk2(file_df.loc[file_df['image_id'] == image_id], dcm_loc, im_dir, is_voi_lut=True)\n            else:\n                im = _process_non_j2k(file_df.loc[file_df['image_id'] == image_id], dcm_loc, im_dir, is_voi_lut=True)\n\n            # Resize\n            im = transforms.functional.to_pil_image(im.astype(np.float32))\n            im = transforms.functional.resize(im, size=[INPUT_SIZE_H, INPUT_SIZE_W,])\n            im = np.asarray(im)\n            \n            # Pad\n            im = pad_right_bottom(im)\n            \n            # Standardize image\n            im -= np.mean(im)\n            im /= np.maximum(np.std(im), 10 ** (-5))\n            \n            im_dict[projection].append(torch.from_numpy(np.expand_dims(im, axis=0)))\n            \n        return {'inputs': im_dict, 'targets': tg_dict, 'file_dict': file_dict, 'index': index, 'patient_id': patient_id}\n    \n    def __len__(self):\n        return self.length","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:49:47.68055Z","iopub.execute_input":"2023-02-18T00:49:47.680925Z","iopub.status.idle":"2023-02-18T00:49:47.695282Z","shell.execute_reply.started":"2023-02-18T00:49:47.680894Z","shell.execute_reply":"2023-02-18T00:49:47.694275Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Clone `BIRADS_classifier` repository","metadata":{}},{"cell_type":"code","source":"if not os.path.isdir('/kaggle/working/BIRADS_classifier'):\n    !git clone https://github.com/nyukat/BIRADS_classifier.git","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:36:07.400195Z","iopub.execute_input":"2023-02-18T00:36:07.400573Z","iopub.status.idle":"2023-02-18T00:36:16.415431Z","shell.execute_reply.started":"2023-02-18T00:36:07.400541Z","shell.execute_reply":"2023-02-18T00:36:16.41413Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BIRADS_classifier_loc = '/kaggle/working/BIRADS_classifier'\nsys.path.append(BIRADS_classifier_loc)  # According to https://www.kaggle.com/general/135988#1985330","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:36:18.552811Z","iopub.execute_input":"2023-02-18T00:36:18.553251Z","iopub.status.idle":"2023-02-18T00:36:18.558899Z","shell.execute_reply.started":"2023-02-18T00:36:18.553192Z","shell.execute_reply":"2023-02-18T00:36:18.557729Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Run inference","metadata":{}},{"cell_type":"markdown","source":"## Prepare model","metadata":{}},{"cell_type":"code","source":"import utils\nimport models_torch as models","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:36:21.500297Z","iopub.execute_input":"2023-02-18T00:36:21.501027Z","iopub.status.idle":"2023-02-18T00:36:21.5122Z","shell.execute_reply.started":"2023-02-18T00:36:21.50099Z","shell.execute_reply":"2023-02-18T00:36:21.511167Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"parameters = {\n    'model_path': '/kaggle/working/BIRADS_classifier/saved_models/model.p',\n    'device_type': 'cuda',\n    \"input_size\": (INPUT_SIZE_H, INPUT_SIZE_W),\n}","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:36:35.766806Z","iopub.execute_input":"2023-02-18T00:36:35.767324Z","iopub.status.idle":"2023-02-18T00:36:35.773388Z","shell.execute_reply.started":"2023-02-18T00:36:35.767275Z","shell.execute_reply":"2023-02-18T00:36:35.772309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\n    parameters[\"device_type\"]\n)","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:36:37.668413Z","iopub.execute_input":"2023-02-18T00:36:37.668796Z","iopub.status.idle":"2023-02-18T00:36:37.67406Z","shell.execute_reply.started":"2023-02-18T00:36:37.668762Z","shell.execute_reply":"2023-02-18T00:36:37.672734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = models.BaselineBreastModel(device, nodropout_probability=1.0, gaussian_noise_std=0.0).to(device)\nmodel.load_state_dict(torch.load(parameters[\"model_path\"]), strict=False)\nmodel.eval()","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:52:20.940207Z","iopub.execute_input":"2023-02-18T00:52:20.940615Z","iopub.status.idle":"2023-02-18T00:52:21.053442Z","shell.execute_reply.started":"2023-02-18T00:52:20.940582Z","shell.execute_reply":"2023-02-18T00:52:21.052235Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Initialize session","metadata":{}},{"cell_type":"code","source":"val_ds = RSNADataset(split=exam_list_sel.loc[exam_list_sel['patient_id'] == 44047], \n                     metadata=train_meta.loc[train_meta['patient_id'] == 44047], \n                     phase='validation')","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:55:29.512888Z","iopub.execute_input":"2023-02-18T00:55:29.51331Z","iopub.status.idle":"2023-02-18T00:55:29.521451Z","shell.execute_reply.started":"2023-02-18T00:55:29.513272Z","shell.execute_reply":"2023-02-18T00:55:29.520294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_bs = 1\nn_threads = 0","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:43:58.072661Z","iopub.execute_input":"2023-02-18T00:43:58.073303Z","iopub.status.idle":"2023-02-18T00:43:58.078964Z","shell.execute_reply.started":"2023-02-18T00:43:58.073263Z","shell.execute_reply":"2023-02-18T00:43:58.077309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_loader = DataLoader(dataset=val_ds,\n                        batch_size=val_bs,\n                        num_workers=n_threads,\n                        drop_last=False,\n                        shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:55:32.03226Z","iopub.execute_input":"2023-02-18T00:55:32.032973Z","iopub.status.idle":"2023-02-18T00:55:32.039306Z","shell.execute_reply.started":"2023-02-18T00:55:32.032933Z","shell.execute_reply":"2023-02-18T00:55:32.037108Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Predict","metadata":{}},{"cell_type":"code","source":"start_timer = timer()\n\nfor i, batch in enumerate(val_loader):\n    input_batch = batch['inputs']\n    target_batch = batch['targets']\n    files_batch = batch['file_dict']\n    patient_id = batch['patient_id']\n    \n    inputs = {\n        \"L-CC\": input_batch['L-CC'][0].float().to(device),\n        \"L-MLO\": input_batch['L-MLO'][0].float().to(device),\n        \"R-CC\": input_batch['R-CC'][0].float().to(device),\n        \"R-MLO\": input_batch['R-MLO'][0].float().to(device),\n    }\n    \n    with torch.no_grad():\n        with amp.autocast(enabled=True):\n            preds = model(inputs)\n    \n    fig = plt.figure(figsize=(7, 7))\n    ax = fig.add_subplot(2,2,1)\n    ax.imshow(input_batch['L-CC'][0].cpu().detach().numpy()[0][0], cmap='gray')\n    ax = fig.add_subplot(2,2,2)\n    ax.imshow(np.fliplr(input_batch['R-CC'][0].cpu().detach().numpy()[0][0]), cmap='gray')\n    ax = fig.add_subplot(2,2,3)\n    ax.imshow(input_batch['L-MLO'][0].cpu().detach().numpy()[0][0], cmap='gray')\n    ax = fig.add_subplot(2,2,4)\n    ax.imshow(np.fliplr(input_batch['R-MLO'][0].cpu().detach().numpy()[0][0]), cmap='gray')\n    plt.show()\n    \n    print(f'Probabilities: BI-RADS 0: {preds[0][0].cpu().detach().numpy()}, BI-RADS 1: {preds[0][1].cpu().detach().numpy()}, BI-RADS 2: {preds[0][2].cpu().detach().numpy()}')\n    \nprint(time_to_str(timer() - start_timer, 'sec'), flush=True)","metadata":{"execution":{"iopub.status.busy":"2023-02-18T01:16:17.101094Z","iopub.execute_input":"2023-02-18T01:16:17.101832Z","iopub.status.idle":"2023-02-18T01:16:22.05886Z","shell.execute_reply.started":"2023-02-18T01:16:17.101792Z","shell.execute_reply":"2023-02-18T01:16:22.057658Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_meta.loc[train_meta['patient_id'] == 44047]","metadata":{"execution":{"iopub.status.busy":"2023-02-18T00:55:51.91261Z","iopub.execute_input":"2023-02-18T00:55:51.913021Z","iopub.status.idle":"2023-02-18T00:55:51.940499Z","shell.execute_reply.started":"2023-02-18T00:55:51.912985Z","shell.execute_reply":"2023-02-18T00:55:51.939348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"According to the metadata, the reference class for `this patient is 'BIRADS 1'.","metadata":{}},{"cell_type":"markdown","source":"# Conclusions\n\nThe method is somewhat dependent on the pre-processing (here I have skipped some of the steps and forced the images to have the expected `(2600, 2000)` size).","metadata":{}},{"cell_type":"markdown","source":"# References\n\nKrzysztof J. Geras, Stacey Wolfson, Yiqiu Shen, Nan Wu, S. Gene Kim, Eric Kim, Laura Heacock, Ujas Parikh, Linda Moy and Kyunghyun Cho. \"High-resolution breast cancer screening with multi-view deep convolutional neural networks.\" URL: https://github.com/nyukat/BIRADS_classifier.","metadata":{}}]}