{"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 Notebook, I extract sagittal slices from the dcm images and train a conv_next multilabel classification model","metadata":{}},{"cell_type":"code","source":"import numpy as np # linear algebra\nimport pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# Input data files are available in the read-only \"../input/\" directory\n# For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\nimport os","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-18T04:40:45.805439Z","iopub.execute_input":"2022-08-18T04:40:45.80624Z","iopub.status.idle":"2022-08-18T04:40:45.811323Z","shell.execute_reply.started":"2022-08-18T04:40:45.806202Z","shell.execute_reply":"2022-08-18T04:40:45.810274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!conda install '/kaggle/input/pydicom-conda-helper/libjpeg-turbo-2.1.0-h7f98852_0.tar.bz2' -c conda-forge -y\n!conda install '/kaggle/input/pydicom-conda-helper/libgcc-ng-9.3.0-h2828fa1_19.tar.bz2' -c conda-forge -y\n!conda install '/kaggle/input/pydicom-conda-helper/gdcm-2.8.9-py37h500ead1_1.tar.bz2' -c conda-forge -y\n!conda install '/kaggle/input/pydicom-conda-helper/conda-4.10.1-py37h89c1867_0.tar.bz2' -c conda-forge -y\n!conda install '/kaggle/input/pydicom-conda-helper/certifi-2020.12.5-py37h89c1867_1.tar.bz2' -c conda-forge -y\n!conda install '/kaggle/input/pydicom-conda-helper/openssl-1.1.1k-h7f98852_0.tar.bz2' -c conda-forge -y","metadata":{"_kg_hide-input":true,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-08-18T04:40:45.856969Z","iopub.execute_input":"2022-08-18T04:40:45.857529Z","iopub.status.idle":"2022-08-18T04:41:30.441878Z","shell.execute_reply.started":"2022-08-18T04:40:45.857489Z","shell.execute_reply":"2022-08-18T04:41:30.440619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pandas as pd\nimport os\nfrom glob import glob\nfrom PIL import Image\nimport numpy as np\nimport pydicom\nimport cv2\nfrom tqdm import tqdm\nimport sys\nsys.path.append(\"../input/timm-pytorch-image-models/pytorch-image-models-master\")\nimport timm\nfrom time import time\nimport matplotlib.pyplot as plt","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-08-18T04:41:30.444685Z","iopub.execute_input":"2022-08-18T04:41:30.445087Z","iopub.status.idle":"2022-08-18T04:41:30.452015Z","shell.execute_reply.started":"2022-08-18T04:41:30.445043Z","shell.execute_reply":"2022-08-18T04:41:30.45083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"basepath = \"../input/rsna-2022-cervical-spine-fracture-detection\"\ndfte = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/test.csv\")\n\nif glob(os.path.join(basepath,\"test_images\",dfte.iloc[0][\"StudyInstanceUID\"],\"*\")):\n    dfte = dfte.copy()\nelse:\n    dfte = pd.DataFrame({\"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'], \"StudyInstanceUID\": ['1.2.826.0.1.3680043.22327', '1.2.826.0.1.3680043.25399', '1.2.826.0.1.3680043.5876'], \"prediction_type\": [\"C1\", \"C1\", \"C1\"]})  \ndftr = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\ndftr.head()\n","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:41:30.453915Z","iopub.execute_input":"2022-08-18T04:41:30.454703Z","iopub.status.idle":"2022-08-18T04:41:30.483509Z","shell.execute_reply.started":"2022-08-18T04:41:30.454667Z","shell.execute_reply":"2022-08-18T04:41:30.482574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Random sampling for quick iterations\n\n","metadata":{}},{"cell_type":"code","source":"dftr = dftr.sample(500).reset_index(drop=True)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:41:30.486003Z","iopub.execute_input":"2022-08-18T04:41:30.486415Z","iopub.status.idle":"2022-08-18T04:41:30.49169Z","shell.execute_reply.started":"2022-08-18T04:41:30.48638Z","shell.execute_reply":"2022-08-18T04:41:30.490609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dftr.head()\n# 1.2.826.0.1.3680043.2374 special case","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:41:30.493395Z","iopub.execute_input":"2022-08-18T04:41:30.493783Z","iopub.status.idle":"2022-08-18T04:41:30.510525Z","shell.execute_reply.started":"2022-08-18T04:41:30.493718Z","shell.execute_reply":"2022-08-18T04:41:30.509602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Looking at orthogonal slices of the 3d image","metadata":{}},{"cell_type":"code","source":"folder = \"train_images\"\nsid = dftr.loc[0,\"StudyInstanceUID\"]\nslices = glob(os.path.join(basepath,folder,sid,\"*\"))\nslices = [(t,pydicom.read_file(t)) for t in slices]\nslices.sort(key = lambda x: int(x[1].ImagePositionPatient[2]))\nslices = np.array([slice[1].pixel_array for slice in slices])\n# slices = np.clip(slices,-1000,1900)\nslices = slices - slices.min()\nslices = slices/slices.max()\nprint(slices.shape)\ncenter = (np.array(slices.shape)/2).astype(np.int16)\nslice_x = slices[center[0],:,:]\nslice_y = slices[:,center[1],:]\nslice_z = slices[:,:,center[2]]\nfig, (ax1, ax2,ax3) = plt.subplots(1, 3,figsize=(20,5))\nax1.imshow(slice_x)\nax2.imshow(slice_y)\nax3.imshow(slice_z)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:41:30.511818Z","iopub.execute_input":"2022-08-18T04:41:30.512247Z","iopub.status.idle":"2022-08-18T04:41:40.622235Z","shell.execute_reply.started":"2022-08-18T04:41:30.512212Z","shell.execute_reply":"2022-08-18T04:41:40.621211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### We see that most information is present in Z slice (Sagittal view). So I extract images - Z slices close to the center of the voxel","metadata":{}},{"cell_type":"code","source":"slice_z_q1 = slices[:,:,int(0.9*center[2])]\nslice_z_q2 = slices[:,:,center[2]]\nslice_z_q3 = slices[:,:,int(1.1*center[2])]\nfig, (ax1, ax2,ax3) = plt.subplots(1, 3,figsize=(20,5))\nax1.imshow(slice_z_q1)\nax2.imshow(slice_z_q2)\nax3.imshow(slice_z_q3)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:41:40.623799Z","iopub.execute_input":"2022-08-18T04:41:40.624231Z","iopub.status.idle":"2022-08-18T04:41:41.108278Z","shell.execute_reply.started":"2022-08-18T04:41:40.624197Z","shell.execute_reply":"2022-08-18T04:41:41.107367Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### In Below Code, I try out different preprocessing approaches and narrow in on the one that gave me best results . Tried out some approaches mentioned here https://www.kaggle.com/code/allunia/pulmonary-dicom-preprocessing","metadata":{}},{"cell_type":"code","source":"# https://www.kaggle.com/code/allunia/pulmonary-dicom-preprocessing\ndef get_z_slices(sid,folder,display=False):\n    slices = glob(os.path.join(basepath,folder,sid,\"*\"))\n    slices = [(t,pydicom.read_file(t)) for t in slices]\n    slices.sort(key = lambda x: int(x[1].ImagePositionPatient[2]))\n    slices = [slice[1].pixel_array for slice in slices]\n#     slices = np.clip(slices,-1000,1900)\n#     slices = slices - slices.min()\n#     slices = slices/slices.max()\n    slices = [(slice-slice.mean())/slice.std() for slice in slices]\n#     slices = [slice/slice.std() for slice in slices]\n    slices = np.array(slices)\n    slices = (slices - slices.min())\n    slices = slices/slices.max()\n    center = (np.array(slices.shape)/2).astype(np.int16)\n    slice_z = slices[:,:,center[2]]\n    slice_z_q1 = slices[:,:,int(0.9*center[2])]\n    slice_z_q2 = slices[:,:,center[2]]\n    slice_z_q3 = slices[:,:,int(1.1*center[2])]\n    if display:\n        fig, (ax1, ax2,ax3) = plt.subplots(1, 3,figsize=(20,5))\n        ax1.imshow(slice_z_q1)\n        ax2.imshow(slice_z_q2)\n        ax3.imshow(slice_z_q3)\n        plt.show()\n    return slice_z_q1,slice_z_q2,slice_z_q3\nfor i in range(10):\n    _ = get_z_slices(dftr.loc[i,\"StudyInstanceUID\"],\"train_images\",True)","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:41:41.109633Z","iopub.execute_input":"2022-08-18T04:41:41.110266Z","iopub.status.idle":"2022-08-18T04:42:59.467241Z","shell.execute_reply.started":"2022-08-18T04:41:41.110222Z","shell.execute_reply":"2022-08-18T04:42:59.462874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### I convert the slices into squares with padding. I dont want to lose out of information by cropping","metadata":{}},{"cell_type":"code","source":"img = Image.fromarray(np.uint8(_[0]*255),\"L\").convert(\"RGB\")\ndef make_square(im, min_size=256, fill_color=(0, 0, 0)):\n#     https://stackoverflow.com/questions/44231209/resize-rectangular-image-to-square-keeping-ratio-and-fill-background-with-black\n    x, y = im.size\n    size = max(min_size, x, y)\n    new_im = Image.new('RGB', (size, size), fill_color)\n    new_im.paste(im, (int((size - x) / 2), int((size - y) / 2)))\n    return new_im\nmake_square(img)","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:42:59.469391Z","iopub.execute_input":"2022-08-18T04:42:59.470447Z","iopub.status.idle":"2022-08-18T04:42:59.581237Z","shell.execute_reply.started":"2022-08-18T04:42:59.470406Z","shell.execute_reply":"2022-08-18T04:42:59.580422Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_process_image(sid,folder):\n    s1,s2,s3 = get_z_slices(sid,folder,False)\n    \n    s1 = Image.fromarray(np.uint8(s1*255),\"L\").convert(\"RGB\")\n    s2 = Image.fromarray(np.uint8(s2*255),\"L\").convert(\"RGB\")\n    s3 = Image.fromarray(np.uint8(s3*255),\"L\").convert(\"RGB\")\n    s1 = make_square(s1)\n    s2 = make_square(s2)\n    s3 = make_square(s3)\n    return s1,s2,s3","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:42:59.585136Z","iopub.execute_input":"2022-08-18T04:42:59.585739Z","iopub.status.idle":"2022-08-18T04:42:59.592996Z","shell.execute_reply.started":"2022-08-18T04:42:59.58569Z","shell.execute_reply":"2022-08-18T04:42:59.592129Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Build dataloader to get 3 images from the sagittal view for each patient","metadata":{}},{"cell_type":"code","source":"import torch\nfrom torch import nn\nfrom torch.utils import data\nfrom PIL import Image\nimport albumentations as A\nimport albumentations.augmentations.functional as F\nfrom albumentations.pytorch import ToTensorV2\nlabel2id = {'patient_overall': 0,\n 'C1': 1,\n 'C2': 2,\n 'C3': 3,\n 'C4': 4,\n 'C5': 5,\n 'C6': 6,\n 'C7': 7}\nIMG_SIZE = 128\nclass rsnaDataset(data.Dataset):\n    def __init__(self,df,aug,test=False):\n        self.df = df.reset_index(drop=True)\n        self.test = test\n        self.aug = aug\n    def __len__(self):\n        return len(self.df)\n    def __getitem__(self,idx):\n        sid = self.df.loc[idx,\"StudyInstanceUID\"]\n        if self.test:\n            img1,img2,img3 = get_process_image(sid,\"test_images\") \n            img1 = self.aug(image=np.array(img1))[\"image\"]\n            img2 = self.aug(image=np.array(img2))[\"image\"]\n            img3 = self.aug(image=np.array(img3))[\"image\"]\n            return {\"img1\":img1,\"img2\":img2,\"img3\":img3}\n        img1,img2,img3 = get_process_image(sid,\"train_images\") \n        img1 = self.aug(image=np.array(img1))[\"image\"]\n        img2 = self.aug(image=np.array(img2))[\"image\"]\n        img3 = self.aug(image=np.array(img3))[\"image\"]\n        labels = self.df.loc[idx,label2id.keys()].values.astype(np.int16)\n        labels = torch.tensor(labels,dtype=torch.float32)\n        return {\"img1\":img1,\"img2\":img2,\"img3\":img3,\"labels\":labels}\n# https://albumentations.ai/docs/examples/pytorch_semantic_segmentation/\ntrain_transform = A.Compose(\n    [\n        A.Resize(IMG_SIZE, IMG_SIZE),\n        A.ShiftScaleRotate(shift_limit=0.2, scale_limit=0.2, rotate_limit=30, p=0.5),\n        A.RGBShift(r_shift_limit=25, g_shift_limit=25, b_shift_limit=25, p=0.5),\n        A.RandomBrightnessContrast(brightness_limit=0.3, contrast_limit=0.3, p=0.5),\n        A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n        ToTensorV2(),\n    ]\n)    \nval_transform = A.Compose(\n    [A.Resize(IMG_SIZE, IMG_SIZE), A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)), ToTensorV2()]\n)\n# ds = rsnaDataset(dftr,train_transform)\n# test_data = rsnaDataset(dfte,val_transform,True)\n# test_data[0]","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:42:59.594535Z","iopub.execute_input":"2022-08-18T04:42:59.595082Z","iopub.status.idle":"2022-08-18T04:42:59.951277Z","shell.execute_reply.started":"2022-08-18T04:42:59.595022Z","shell.execute_reply":"2022-08-18T04:42:59.949754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class rsnaModel(nn.Module):\n    def __init__(self,):\n        super(rsnaModel,self).__init__()\n        self.model1 = timm.create_model(\"convnext_tiny\",pretrained=False,global_pool=\"avgmax\",)\n        self.model1.load_state_dict(torch.load(\"../input/timm-convnext-tiny-weights/convnext_tiny.pth\"))\n        self.model1.head.drop = nn.Dropout(0.1)\n        self.model1.head.fc = nn.Linear(768,128)\n        \n        self.model2 = timm.create_model(\"convnext_tiny\",pretrained=False,global_pool=\"avgmax\",)\n        self.model2.load_state_dict(torch.load(\"../input/timm-convnext-tiny-weights/convnext_tiny.pth\"))\n        self.model2.head.drop = nn.Dropout(0.1)\n        self.model2.head.fc = nn.Linear(768,128)\n        \n        self.model3 = timm.create_model(\"convnext_tiny\",pretrained=False,global_pool=\"avgmax\",)\n        self.model3.load_state_dict(torch.load(\"../input/timm-convnext-tiny-weights/convnext_tiny.pth\"))\n        self.model3.head.drop = nn.Dropout(0.1)\n        self.model3.head.fc = nn.Linear(768,128)\n        \n        self.dropout = nn.Dropout(0.1)\n        self.outputLayer = nn.Linear(128*3,len(label2id))\n        \n    def forward(self,x1,x2,x3):\n        x1 = self.model1(x1)\n        x2 = self.model2(x2)\n        x3 = self.model3(x3)\n        x = torch.cat([x1,x2,x3],dim=-1)\n        x = self.dropout(x)\n        x = self.outputLayer(x)\n        return x\n# test_m = rsnaModel()     \n# a = torch.randn((1,3,224,224))\n# test_m(a,a,a).shape\n        ","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:42:59.955031Z","iopub.execute_input":"2022-08-18T04:42:59.955339Z","iopub.status.idle":"2022-08-18T04:42:59.968164Z","shell.execute_reply.started":"2022-08-18T04:42:59.955307Z","shell.execute_reply":"2022-08-18T04:42:59.96719Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import StratifiedGroupKFold\nNFOLDS = 5\ndftr = dftr.reset_index(drop=True)\nskf = StratifiedGroupKFold(NFOLDS)\nfor i,(train_idx,val_idx) in enumerate(skf.split(X=dftr,y=dftr[\"patient_overall\"],groups=dftr[\"StudyInstanceUID\"])):\n    dftr.loc[val_idx,\"fold\"]=i\ndftr.fold.value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:42:59.969794Z","iopub.execute_input":"2022-08-18T04:42:59.970076Z","iopub.status.idle":"2022-08-18T04:43:00.174234Z","shell.execute_reply.started":"2022-08-18T04:42:59.97005Z","shell.execute_reply":"2022-08-18T04:43:00.173349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label2id.keys()","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:43:00.17579Z","iopub.execute_input":"2022-08-18T04:43:00.176129Z","iopub.status.idle":"2022-08-18T04:43:00.183855Z","shell.execute_reply.started":"2022-08-18T04:43:00.176095Z","shell.execute_reply":"2022-08-18T04:43:00.182759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE=12","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:43:00.185242Z","iopub.execute_input":"2022-08-18T04:43:00.185662Z","iopub.status.idle":"2022-08-18T04:43:00.192746Z","shell.execute_reply.started":"2022-08-18T04:43:00.185636Z","shell.execute_reply":"2022-08-18T04:43:00.191743Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader\nfold_preds = []\nfor FOLD in range(NFOLDS):\n\n    dftrain = dftr[dftr[\"fold\"]!=FOLD].reset_index(drop=True)#.head(20)\n    dfeval = dftr[dftr[\"fold\"]==FOLD].reset_index(drop=True)#.head(20)\n    train_data = rsnaDataset(dftrain,train_transform)\n    eval_data = rsnaDataset(dfeval,val_transform)\n\n    train_dataloader = DataLoader(train_data,\\\n                        batch_size=BATCH_SIZE,\\\n                        shuffle=True)\n    eval_dataloader = DataLoader(eval_data,\\\n                        batch_size=BATCH_SIZE,\\\n                        shuffle=False)\n\n    \n    break","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:43:00.195095Z","iopub.execute_input":"2022-08-18T04:43:00.195776Z","iopub.status.idle":"2022-08-18T04:43:00.205884Z","shell.execute_reply.started":"2022-08-18T04:43:00.195739Z","shell.execute_reply":"2022-08-18T04:43:00.204696Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def criterion(y_pred,y):\n    y_pred1 = y_pred[:,0]\n    y1 = y[:,0]\n    \n    y_pred2 = y_pred[:,1:]\n    y2 = y[:,1:]\n    loss1 = nn.BCEWithLogitsLoss()(y_pred1,y1)\n    loss2 = nn.BCEWithLogitsLoss()(y_pred2,y2)\n    return (loss1 + loss2)/2\n","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:43:00.207138Z","iopub.execute_input":"2022-08-18T04:43:00.207651Z","iopub.status.idle":"2022-08-18T04:43:00.218256Z","shell.execute_reply.started":"2022-08-18T04:43:00.207617Z","shell.execute_reply":"2022-08-18T04:43:00.217301Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEVICE = \"cpu\"\nLR  = 3e-4\nEPOCHS = 1\nBEST_LOSS_MODEL = \"best_model.pth\"","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:43:00.219996Z","iopub.execute_input":"2022-08-18T04:43:00.220344Z","iopub.status.idle":"2022-08-18T04:43:00.227515Z","shell.execute_reply.started":"2022-08-18T04:43:00.220309Z","shell.execute_reply":"2022-08-18T04:43:00.22661Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from time import time\nmodel = rsnaModel().to(DEVICE)\noptimizer = torch.optim.AdamW(model.parameters(),lr=LR,weight_decay=0.0001)\nbest_loss = 1000\nfor epoch in range(EPOCHS):\n    model.train()\n    train_loss = 0\n    start_time = time()\n    for i, batch in tqdm(enumerate(train_dataloader),total=len(train_dataloader)):\n        img1 = batch[\"img1\"].to(DEVICE)\n        img2 = batch[\"img2\"].to(DEVICE)\n        img3 = batch[\"img3\"].to(DEVICE)\n        y = batch[\"labels\"].to(DEVICE)\n\n        y_pred = model(img1,img2,img3)\n        loss = criterion(y_pred,y)\n        train_loss+=loss.item()\n\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n    train_loss = train_loss/len(train_dataloader)\n    \n    model.eval()\n    eval_loss = 0\n    with torch.no_grad():\n        for i, batch in tqdm(enumerate(eval_dataloader),total=len(eval_dataloader)):\n            img1 = batch[\"img1\"].to(DEVICE)\n            img2 = batch[\"img2\"].to(DEVICE)\n            img3 = batch[\"img3\"].to(DEVICE)\n            y = batch[\"labels\"].to(DEVICE)\n\n            y_pred = model(img1,img2,img3)\n            loss = criterion(y_pred,y)\n            eval_loss+=loss.item()\n\n        eval_loss = eval_loss/len(eval_dataloader)\n    if eval_loss < best_loss:\n        best_loss = eval_loss\n        torch.save(model.state_dict(), BEST_LOSS_MODEL)\n    print(f\"Epoch {epoch}:: Train Loss: {train_loss}; Eval Loss: {eval_loss}; Best: {best_loss} Time taken: {(time() - start_time)/60} secs\")\n    ","metadata":{"execution":{"iopub.status.busy":"2022-08-18T04:43:00.229212Z","iopub.execute_input":"2022-08-18T04:43:00.229584Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfte_deduplicated = dfte[[\"StudyInstanceUID\"]].drop_duplicates().reset_index(drop=True)\ndfte_deduplicated.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = rsnaDataset(dfte_deduplicated,val_transform,True)\ntest_dataloader = DataLoader(test_data,\\\n                        batch_size=BATCH_SIZE,\\\n                        shuffle=False)\nmodel.load_state_dict(torch.load(BEST_LOSS_MODEL))                        \nmodel.eval()\nsize = len(test_dataloader.dataset)\npredictions = np.empty((0,len(label2id)))\nwith torch.no_grad():\n    for i, batch in tqdm(enumerate(test_dataloader),total=len(test_dataloader)):\n        img1 = batch[\"img1\"].to(DEVICE)\n        img2 = batch[\"img2\"].to(DEVICE)\n        img3 = batch[\"img3\"].to(DEVICE)\n        y_pred = model(img1,img2,img3)\n        \n        y_pred = torch.nn.Sigmoid()(y_pred)\n        predictions = np.append(predictions,y_pred.cpu().numpy(),axis=0)\npredictions.shape        ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dfte_deduplicated.loc[:,label2id.keys()] = predictions\ndfte_deduplicated.head()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dicts = dfte_deduplicated.set_index(\"StudyInstanceUID\")[list(label2id.keys())].to_dict()\n# dicts","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub = []\nfor i,row in dfte.iterrows():\n    sub.append([row[\"row_id\"],dicts[row[\"prediction_type\"]][row[\"StudyInstanceUID\"]]])\nsub = pd.DataFrame(sub,columns=[\"row_id\",\"fractured\"])\nsub.head() ","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.to_csv(\"submission.csv\",index=None)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}