{"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":"# RNSA_2022: ConvNext_LSTM Inference\n## Imports","metadata":{}},{"cell_type":"code","source":"try:\n    import pylibjpeg\nexcept:\n    # Offline dependencies:\n    !mkdir -p /root/.cache/torch/hub/checkpoints/\n\n    !pip install /kaggle/input/rsna-2022-whl/{pydicom-2.3.0-py3-none-any.whl,pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}\n    !pip install /kaggle/input/rsna-2022-whl/{torch-1.12.1-cp37-cp37m-manylinux1_x86_64.whl,torchvision-0.13.1-cp37-cp37m-manylinux1_x86_64.whl}","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:28:55.016687Z","iopub.execute_input":"2022-10-28T19:28:55.01752Z","iopub.status.idle":"2022-10-28T19:30:51.283489Z","shell.execute_reply.started":"2022-10-28T19:28:55.017396Z","shell.execute_reply":"2022-10-28T19:30:51.282318Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport pydicom\n\nimport cv2\nimport os\nfrom tqdm import tqdm\nimport glob\nimport pickle\nfrom albumentations import *\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision import transforms\nfrom torchvision.models.convnext import convnext_tiny, convnext_small, convnext_base\nfrom torchvision.models.efficientnet import efficientnet_v2_l\n\nfrom torch.utils.data import Dataset, DataLoader\nimport torch\nimport sys\nimport time\n","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:30:51.286184Z","iopub.execute_input":"2022-10-28T19:30:51.286564Z","iopub.status.idle":"2022-10-28T19:30:53.737146Z","shell.execute_reply.started":"2022-10-28T19:30:51.286523Z","shell.execute_reply":"2022-10-28T19:30:53.736192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config + Dataset","metadata":{}},{"cell_type":"code","source":"config = {'seq_len': 150,\n          'feature_size': 1024,\n          'lstm_size': 128,\n          'target_size': 512,\n          'crop_size': 384,\n          'num_classes':7,\n          'batch_size_image_level': 32,\n          'batch_size_patient_level': 8\n         }","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:30:53.738631Z","iopub.execute_input":"2022-10-28T19:30:53.739227Z","iopub.status.idle":"2022-10-28T19:30:53.747533Z","shell.execute_reply.started":"2022-10-28T19:30:53.73919Z","shell.execute_reply":"2022-10-28T19:30:53.746711Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#https://www.kaggle.com/code/vslaykovsky/infer-pytorch-effnetv2-single-model-lb-0-49/notebook?scriptVersionId=104434908\ndef load_df_test():\n    df_test = pd.read_csv(f'../input/rsna-2022-cervical-spine-fracture-detection/test.csv')\n\n    if df_test.iloc[0].row_id == '1.2.826.0.1.3680043.10197_C1':\n        # test_images and test.csv are inconsistent in the dev dataset, fixing labels for the dev run.\n        df_test = pd.DataFrame({\n            \"row_id\": ['1.2.826.0.1.3680043.22327_C1', '1.2.826.0.1.3680043.25399_C1', '1.2.826.0.1.3680043.5876_C1'],\n            \"StudyInstanceUID\": ['1.2.826.0.1.3680043.22327', '1.2.826.0.1.3680043.25399', '1.2.826.0.1.3680043.5876'],\n            \"prediction_type\": [\"C1\", \"C1\", \"patient_overall\"]}\n        )\n    return df_test","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:30:53.750377Z","iopub.execute_input":"2022-10-28T19:30:53.75082Z","iopub.status.idle":"2022-10-28T19:30:53.760065Z","shell.execute_reply.started":"2022-10-28T19:30:53.750783Z","shell.execute_reply":"2022-10-28T19:30:53.758968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = load_df_test()\nTEST_PATH = '../input/rsna-2022-cervical-spine-fracture-detection/test_images'\nstudy_id_list = list(test_df.StudyInstanceUID.unique()) #uids\nstudy_id_list","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:30:53.761945Z","iopub.execute_input":"2022-10-28T19:30:53.762385Z","iopub.status.idle":"2022-10-28T19:30:53.792741Z","shell.execute_reply.started":"2022-10-28T19:30:53.76235Z","shell.execute_reply":"2022-10-28T19:30:53.791767Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"selected_image_dict = {}\nfor uid in study_id_list:\n    dicom_files = glob.glob(os.path.join(f'{TEST_PATH}/{uid}', '*.dcm'))\n    # print(len(dicom_files))\n    middle_slice = int(len(dicom_files)/2)\n    num_left_images = num_right_images = int(0.15*len(dicom_files)) # select 15% to the left, 15% to the right\n    # print(middle_slice, num_left_images)\n    selected_image_dict[uid] = list(np.arange(middle_slice-num_left_images, middle_slice, 1)) +\\\n                               list(np.arange(middle_slice+1, middle_slice+num_right_images+1,1))","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:30:53.793799Z","iopub.execute_input":"2022-10-28T19:30:53.794125Z","iopub.status.idle":"2022-10-28T19:30:54.037967Z","shell.execute_reply.started":"2022-10-28T19:30:53.794087Z","shell.execute_reply":"2022-10-28T19:30:54.037069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(selected_image_dict[study_id_list[0]])","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:30:54.039337Z","iopub.execute_input":"2022-10-28T19:30:54.039677Z","iopub.status.idle":"2022-10-28T19:30:54.04622Z","shell.execute_reply.started":"2022-10-28T19:30:54.039642Z","shell.execute_reply":"2022-10-28T19:30:54.045079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Helper functions","metadata":{}},{"cell_type":"code","source":"# adapted from https://www.kaggle.com/code/sparkyjunior/windowing-in-ct-scans\ndef window(data, WL=400, WW=1800):\n    data.PhotometricInterpretation = 'YBR_FULL'\n    slope     = data.RescaleSlope\n    intercept = data.RescaleIntercept\n    img = data.pixel_array\n    img = (img*slope +intercept)\n    upper, lower = WL+WW//2, WL-WW//2\n    X = np.clip(img.copy(), lower, upper)\n    X = X - np.min(X)\n    X = X / np.max(X)\n    X = (X*255.0).astype('uint8')\n    return X\n\n# https://www.kaggle.com/code/iafoss/hubmap-pytorch-fast-ai-starter\ndef img2tensor(img,dtype:np.dtype=np.float32):\n    if img.ndim==2 : img = np.expand_dims(img,2)\n    img = np.transpose(img,(2,0,1)) # since numpy array has [H,W,C] -> we want [C,H,W] for torch tensor\n    return torch.from_numpy(img.astype(dtype, copy=False))","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:30:54.047647Z","iopub.execute_input":"2022-10-28T19:30:54.048326Z","iopub.status.idle":"2022-10-28T19:30:54.077592Z","shell.execute_reply.started":"2022-10-28T19:30:54.048239Z","shell.execute_reply":"2022-10-28T19:30:54.076499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"mean=np.array([0.456, 0.456, 0.456])\nstd=np.array([0.224, 0.224, 0.224])\n\nclass CSFImageDataset(Dataset):\n    def __init__(self, uid, image_list, target_size, crop_size):\n        self.uid = uid\n        self.image_list  = image_list\n        self.target_size = target_size\n        self.crop_size   = crop_size\n        \n    def __len__(self):\n        return len(self.image_list)\n    \n    def __getitem__(self, index):\n        \n        PATH = os.path.join(TEST_PATH, self.uid)\n        data_list  = [pydicom.dcmread(f'{PATH}/{str(self.image_list[index]-1)}.dcm'), \\\n                      pydicom.dcmread(f'{PATH}/{str(self.image_list[index])}.dcm'), \\\n                      pydicom.dcmread(f'{PATH}/{str(self.image_list[index]+1)}.dcm')]\n        # print(data_list)\n                \n        imgs  = [window(data) for data in data_list]\n        \n        stacked_img = np.stack(imgs, axis=-1)\n        stacked_img = cv2.resize(stacked_img, (self.target_size, self.target_size))\n        \n        inference_transform = Compose([CenterCrop(self.crop_size,self.crop_size)])\n        stacked_img = inference_transform(image=stacked_img)\n        # transformed image\n        X = stacked_img['image']\n        X = img2tensor((X/255.0 - mean)/std)\n        \n        return X","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:30:54.078973Z","iopub.execute_input":"2022-10-28T19:30:54.079625Z","iopub.status.idle":"2022-10-28T19:30:54.09053Z","shell.execute_reply.started":"2022-10-28T19:30:54.079591Z","shell.execute_reply":"2022-10-28T19:30:54.089479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CSFInstanceDataset(Dataset):\n    def __init__(self, feature_array_dict, study_id_list, seq_len):\n        self.feature_array_dict = feature_array_dict\n        self.study_id_list = study_id_list\n        self.seq_len = seq_len\n        \n    def __len__(self):\n        return len(self.study_id_list)\n    \n    def __getitem__(self, index):\n        uid = self.study_id_list[index]\n        feature_array = self.feature_array_dict[uid]\n        # if a study has more slices than seq_len\n        # resize features\n        if len(feature_array) > self.seq_len: \n            x = cv2.resize(feature_array, (feature_array.shape[1], self.seq_len), interpolation = cv2.INTER_LINEAR)\n        else:\n            # a study has less slices than seq_len\n            # pad with zeros   \n            x = np.pad(feature_array, pad_width=[(0,self.seq_len-feature_array.shape[0]), (0,0)], constant_values=0) \n            \n        X = torch.tensor(x, dtype=torch.float32) #(seq_len,1024)\n        return X, uid","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:30:54.094023Z","iopub.execute_input":"2022-10-28T19:30:54.094409Z","iopub.status.idle":"2022-10-28T19:30:54.10553Z","shell.execute_reply.started":"2022-10-28T19:30:54.094346Z","shell.execute_reply":"2022-10-28T19:30:54.104571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model architecture","metadata":{}},{"cell_type":"code","source":"class ConvNextCNN_B_Feature(nn.Module):\n    def __init__(self):\n        super(ConvNextCNN_B_Feature, self).__init__()\n        m = convnext_base() #extract the output layer\n        in_features = m.classifier[-1].in_features\n        self.features = m.features\n        # self.relu = nn.ReLU()\n        self.avgpool = nn.AdaptiveAvgPool2d(1) \n        self.drop = nn.Dropout(p=0.5)\n        self.fc = nn.Linear(in_features=in_features, out_features=7)\n        \n    # Progresses data across layers    \n    def forward(self, x):\n        out = self.features(x)\n        # out = self.relu(out)\n        out = self.avgpool(out)\n        out = self.drop(out)\n        feature = out.view(x.size(0), -1) #layer before the fc\n        # shape after conv and max_pool layer [B, n_out_channel, H, W]\n        # out = out.reshape(out.size(0), -1)\n        \n        out = self.fc(feature)\n        \n        return feature, out","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:30:54.107228Z","iopub.execute_input":"2022-10-28T19:30:54.107611Z","iopub.status.idle":"2022-10-28T19:30:54.118881Z","shell.execute_reply.started":"2022-10-28T19:30:54.107577Z","shell.execute_reply":"2022-10-28T19:30:54.117769Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CSFNet(nn.Module):\n    def __init__(self, input_len, lstm_size):\n        super().__init__()\n        self.lstm1 = nn.GRU(input_len, lstm_size, bidirectional=True, batch_first=True)\n        self.last_linear = nn.Linear(lstm_size*2, 1) #(*4)\n        # self.attention = Attention(lstm_size*2, config['seq_len'])\n        \n    def forward(self, x):\n        h_lstm1, _ = self.lstm1(x)\n\n        max_pool, _ = torch.max(h_lstm1, 1)\n        # att_pool = self.attention(h_lstm1)\n        # conc = torch.cat((max_pool, att_pool), 1)  \n\n        logits = self.last_linear(max_pool)\n        return logits","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:30:54.120033Z","iopub.execute_input":"2022-10-28T19:30:54.121095Z","iopub.status.idle":"2022-10-28T19:30:54.129442Z","shell.execute_reply.started":"2022-10-28T19:30:54.121046Z","shell.execute_reply":"2022-10-28T19:30:54.128349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Load model from `state_dict`\n# Image level\nlv1_model = ConvNextCNN_B_Feature()\nlv1_model.load_state_dict(torch.load('../input/cnn-lstm-oct-25/run_3/run_3/model_3.pth'))\nlv1_model = lv1_model.cuda()\nlv1_model.eval()\n\n# Exam level\nlv2_model = CSFNet(input_len=config['feature_size'], lstm_size=config['lstm_size'])\nlv2_model.load_state_dict(torch.load('../input/cnn-lstm-oct-25/run_3/run_3/model_lstm_3.pth'))\nlv2_model = lv2_model.cuda()\nlv2_model.eval()","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:31:32.397014Z","iopub.execute_input":"2022-10-28T19:31:32.39738Z","iopub.status.idle":"2022-10-28T19:31:41.294099Z","shell.execute_reply.started":"2022-10-28T19:31:32.397344Z","shell.execute_reply":"2022-10-28T19:31:41.292276Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Inference","metadata":{}},{"cell_type":"markdown","source":"### Stage 1: Feature extraction","metadata":{}},{"cell_type":"code","source":"submission_dict = {'row_id':[], 'fractured':[]}\nfeature_array_dict = {}\nfor uid in study_id_list:\n    image_list = selected_image_dict[uid]\n    dataset = CSFImageDataset(uid=uid, image_list=image_list,\\\n                              target_size=config['target_size'], \\\n                              crop_size=config['crop_size'])\n    generator = DataLoader(dataset, batch_size=config['batch_size_image_level'],\\\n                           shuffle=False, pin_memory=True, drop_last=False)\n    \n    preds_uid = []\n    feature_array = np.zeros((len(dataset), config['feature_size']), dtype=np.float32)\n    #image level prediction\n    for i, images in tqdm(enumerate(generator), total=len(generator)):\n        with torch.no_grad():\n            start = i*config['batch_size_image_level']\n            end = start+config['batch_size_image_level']\n            \n            if i == len(generator) - 1: # last batch\n                end = len(generator.dataset)\n            \n            images = images.cuda()\n            \n            features, preds = lv1_model(images)\n                \n            feature_array[start:end] = np.squeeze(features.cpu().data.numpy())\n            preds = preds.sigmoid().detach()\n            preds_uid.append(preds)\n            \n    feature_array_dict[uid] = feature_array\n    \n    mean_preds = torch.mean(torch.cat(preds_uid), dim=0).cpu().data.numpy()\n    \n    \n    submission_dict['row_id'].append(f'{uid}_C1')\n    submission_dict['fractured'].append(mean_preds[0])\n        \n    submission_dict['row_id'].append(f'{uid}_C2')\n    submission_dict['fractured'].append(mean_preds[1])\n        \n    submission_dict['row_id'].append(f'{uid}_C3')\n    submission_dict['fractured'].append(mean_preds[2])\n        \n    submission_dict['row_id'].append(f'{uid}_C4')\n    submission_dict['fractured'].append(mean_preds[3])\n        \n    submission_dict['row_id'].append(f'{uid}_C5')\n    submission_dict['fractured'].append(mean_preds[4])\n        \n    submission_dict['row_id'].append(f'{uid}_C6')\n    submission_dict['fractured'].append(mean_preds[5])\n        \n    submission_dict['row_id'].append(f'{uid}_C7')\n    submission_dict['fractured'].append(mean_preds[6])\n   ","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:31:41.299103Z","iopub.execute_input":"2022-10-28T19:31:41.30142Z","iopub.status.idle":"2022-10-28T19:32:11.175413Z","shell.execute_reply.started":"2022-10-28T19:31:41.301377Z","shell.execute_reply":"2022-10-28T19:32:11.174414Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Stage 2: Instance level prediction","metadata":{}},{"cell_type":"code","source":"# Exam level prediction\ndataset = CSFInstanceDataset(feature_array_dict=feature_array_dict,\n                             study_id_list=study_id_list,\n                             seq_len=config['seq_len']\n                            )\ngenerator = DataLoader(dataset=dataset,\n                       batch_size=config['batch_size_patient_level'],\n                       shuffle=False,\n                       pin_memory=True)\n\nfor features, list_uid in tqdm(generator, total=len(generator)):\n    with torch.no_grad():\n        features = features.cuda()\n        preds = lv2_model(features)\n        preds = np.squeeze(preds.sigmoid().cpu().data.numpy())\n    \n    for j in range(len(features)):# n features extracted from n patients from stage 1\n        submission_dict['row_id'].append(f'{list_uid[j]}_patient_overall')\n        submission_dict['fractured'].append(preds[j])","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:32:11.176985Z","iopub.execute_input":"2022-10-28T19:32:11.177438Z","iopub.status.idle":"2022-10-28T19:32:11.215681Z","shell.execute_reply.started":"2022-10-28T19:32:11.177401Z","shell.execute_reply":"2022-10-28T19:32:11.214653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Submission","metadata":{}},{"cell_type":"code","source":"sub_df = pd.DataFrame.from_dict(submission_dict)\nsub_df = sub_df.sort_values(by='row_id', axis=0, ascending=True).reset_index(drop=True)\nsub_df","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:32:11.217847Z","iopub.execute_input":"2022-10-28T19:32:11.218279Z","iopub.status.idle":"2022-10-28T19:32:11.247113Z","shell.execute_reply.started":"2022-10-28T19:32:11.218243Z","shell.execute_reply":"2022-10-28T19:32:11.246042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub_df.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2022-10-28T19:32:11.248459Z","iopub.execute_input":"2022-10-28T19:32:11.248826Z","iopub.status.idle":"2022-10-28T19:32:11.257873Z","shell.execute_reply.started":"2022-10-28T19:32:11.248792Z","shell.execute_reply":"2022-10-28T19:32:11.256898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}