{"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":"code","source":"import os\nos.environ[\"CUDA_VISIBLE_DEVICES\"]= \"0\"\n\nimport numpy as np\nimport joblib\nimport pandas as pd\nimport cv2\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport gc\nfrom glob import glob\nimport sys\n\nimport torch\nfrom torch import nn\nfrom torch.nn import functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.utils.data import random_split\nfrom torchvision.datasets import MNIST\nfrom torchvision import transforms\n#import torchmetrics\n#from pytorch_lightning.callbacks import ModelCheckpoint, EarlyStopping\n#import pytorch_lightning as pl\nfrom albumentations.pytorch import ToTensorV2\nimport gc\nfrom tqdm import tqdm\nimport multiprocessing\nfrom ipywidgets import interactive, widgets, fixed\nfrom matplotlib import animation, rc; rc('animation', html='jshtml')\nimport matplotlib.pyplot as plt\n%matplotlib inline\n\ntry:\n    import pylibjpeg\nexcept:\n    # Offline dependencies:\n    !mkdir -p /root/.cache/torch/hub/checkpoints/\n    #!cp ../input/rsna-2022-whl/efficientnet_v2_s-dd5fe13b.pth  /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}\nimport pydicom\n\nimport random\ndef seed_everything(seed=42):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-10-27T23:15:54.710038Z","iopub.execute_input":"2022-10-27T23:15:54.710456Z","iopub.status.idle":"2022-10-27T23:16:30.994974Z","shell.execute_reply.started":"2022-10-27T23:15:54.710368Z","shell.execute_reply":"2022-10-27T23:16:30.993754Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"package_path = '../input/efficientnet/EfficientNet-PyTorch-master/'\nsys.path.append(package_path)\npackage_path = '../input/efficientnet/EfficientNet-PyTorch-master/efficientnet_pytorch'\nsys.path.append(package_path)\nfrom efficientnet_pytorch import EfficientNet\n\n!pip install ../input/timm0412/timm-0.4.12-py3-none-any.whl\n!pip install ../input/pretrainedmodels-wheels/pretrainedmodels-0.7.4-py3-none-any.whl\n\nimport timm\n!pip install ../input/giba-libs/segmentation_models_pytorch-0.3.0-py3-none-any.whl --no-deps\n\nseed_everything()\nprint(torch.__version__)\nprint(timm.__version__)","metadata":{"execution":{"iopub.status.busy":"2022-10-27T23:16:31.000093Z","iopub.execute_input":"2022-10-27T23:16:31.001142Z","iopub.status.idle":"2022-10-27T23:17:53.564156Z","shell.execute_reply.started":"2022-10-27T23:16:31.001101Z","shell.execute_reply":"2022-10-27T23:17:53.562756Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Spine Segmentation Model","metadata":{}},{"cell_type":"code","source":"DEVICE = 'cuda'\nmodel = torch.load('../input/rsna2022weights/model_segmentation_effb2_14.pth')\nmodel.to(DEVICE)\nmodel.eval()","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2022-10-27T23:17:53.567881Z","iopub.execute_input":"2022-10-27T23:17:53.568202Z","iopub.status.idle":"2022-10-27T23:17:59.283402Z","shell.execute_reply.started":"2022-10-27T23:17:53.568169Z","shell.execute_reply":"2022-10-27T23:17:59.282445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dcm(fn):\n    dc = pydicom.dcmread(fn)\n    img = dc.pixel_array\n    img = img.astype('float32')\n    img *= dc.RescaleSlope\n    img += dc.RescaleIntercept\n    img = np.clip(img, -50, 2000)\n    img = (img+50)/(2000+50)\n    return img\n\n\ndef proc_images(fn, PATH='../input/rsna-2022-cervical-spine-fracture-detection/'):\n    ids = fn.split('/')[-1].split('.nii')[0]\n    files = glob(f\"{PATH}test_images/{ids}/*.dcm\")\n    files = [files[i] for i in np.argsort([int(i.split('/')[-1].split('.')[0]) for i in files])]    \n    #print(files)\n    img = np.stack([load_dcm(f) for f in files])\n    return img\n\n\ndef find_draw_mask(res=None, cls=1):\n    kernel = np.ones((5, 5), np.uint8)\n    edged = 255*(res>0).astype('uint8')\n    edged = cv2.dilate(edged, kernel, iterations=1)\n    mask = np.zeros_like(res, dtype='uint8')\n    contours, hierarchy = cv2.findContours(edged, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)\n    for contour in contours:\n        area = cv2.contourArea(contour)\n        if area>500:\n            mask = cv2.drawContours(mask, [contour], 0, (255,255,255), thickness= -1, lineType=cv2.LINE_AA)\n    mask = cv2.dilate(mask, kernel, iterations=1)\n    \n    x,y,w,h = 0,0,0,0\n    contours, hierarchy = cv2.findContours(mask, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)\n    for contour in contours:\n        area = cv2.contourArea(contour)\n        if area>500:\n            x,y,w,h = cv2.boundingRect(contour)\n    \n    return mask, [x,y,w,h]\n\n\ndef find_pos(res=None, cls=1):\n    edged = 255*(res>0).astype('uint8')\n    contours, hierarchy = cv2.findContours(edged, cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_NONE)\n    for contour in contours:\n        area = cv2.contourArea(contour)\n        if area>200:\n            (x,y,w,h) = cv2.boundingRect(contour)\n            #return x-1,y-1,w+1,h+1\n            return y-1, y+h+1\n    return -1, -1\n\n\ndef proc_dicom_segment(fn, model):\n    img = proc_images(fn)\n    print(img.shape)\n    \n    ind = np.linspace(0.05*img.shape[1], 0.95*img.shape[1], 24).astype('int')\n    image = img[:,ind].copy()\n    image = np.moveaxis(image, 1, 2)\n    image = cv2.resize(image, (256, 256))\n    image = np.moveaxis(image, 2, 0)\n    image = np.stack( (image,image,image), axis=1)\n    image = torch.from_numpy(image).float()\n    with torch.no_grad():\n        res = torch.argmax(model(image.to(DEVICE)), 1)\n    res = res.detach().cpu().numpy()\n    \n    # Flip if upside down\n    flip = np.sum( (res[:,:128]>0) & (res[:,:128]<=4) ) < np.sum( (res[:,128:]>0) & (res[:,128:]<=4) )\n    if flip:\n        print('Flip')\n        img = img[-1::-1].copy()\n    \n    ind = np.linspace(0.05*img.shape[1], 0.95*img.shape[1], 64).astype('int')\n    image = img[:,ind].copy()\n    image = np.moveaxis(image, 1, 2)\n    image = cv2.resize(image, (256, 256))\n    image = np.moveaxis(image, 2, 0)\n    image = np.stack( (image,image,image), axis=1)\n    image = torch.from_numpy(image).float()\n    with torch.no_grad():\n        mask = torch.argmax(model(image.to(DEVICE)), 1)\n    mask = mask.detach().cpu().numpy()\n    mask = np.moveaxis(mask, 0, 1)\n    \n    msk0, cntr = find_draw_mask(cv2.resize(np.mean((mask>0)&(mask<=7), 1), (512, 512)))\n    x, y, w, h = cntr\n    if (h>50) and (w>50):\n        img = img[:, y:y+h, x:x+w]\n        #mask = mask[:, y:y+h, x:x+w]\n\n    return img, mask","metadata":{"execution":{"iopub.status.busy":"2022-10-27T23:21:34.611084Z","iopub.execute_input":"2022-10-27T23:21:34.611445Z","iopub.status.idle":"2022-10-27T23:21:34.636244Z","shell.execute_reply.started":"2022-10-27T23:21:34.611414Z","shell.execute_reply":"2022-10-27T23:21:34.635346Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"testfiles = glob('../input/rsna-2022-cervical-spine-fracture-detection/test_images/*')\ntestfiles = [f.split('/')[-1] for f in testfiles]\ntestfiles","metadata":{"execution":{"iopub.status.busy":"2022-10-27T23:21:35.000357Z","iopub.execute_input":"2022-10-27T23:21:35.001043Z","iopub.status.idle":"2022-10-27T23:21:35.009989Z","shell.execute_reply.started":"2022-10-27T23:21:35.001006Z","shell.execute_reply":"2022-10-27T23:21:35.008922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.rcParams['figure.figsize'] = [12, 12]\n\nfor fn in tqdm(testfiles):\n    ids = fn.split('/')[-1]\n    print(ids)\n    image, mask = proc_dicom_segment(fn, model)\n    print(image.shape)\n    \n    print('Original Image')\n    images = np.hstack((\n        cv2.resize(np.mean(image, 0), (256, 256)),\n        cv2.resize(np.mean(image, 1), (256, 256)),\n        cv2.resize(np.mean(image, 2), (256, 256)),\n    ))\n    plt.imshow(images)\n    plt.show()\n\n    print('Segments C1 to C7')\n    masks = np.hstack((\n        cv2.resize(np.mean(mask, 0), (256, 256)),\n        cv2.resize(np.mean(mask, 1), (256, 256)),\n        cv2.resize(np.mean(mask, 2), (256, 256)),\n    ))\n    plt.imshow(masks)\n    plt.show()\n\n    # Select each segment \n    for i in range(8):\n        print(f'Segment C{i+1}')\n        masks = np.hstack((\n            cv2.resize(np.mean(mask==(i+1), 0), (256, 256)),\n            cv2.resize(np.mean(mask==(i+1), 1), (256, 256)),\n            cv2.resize(np.mean(mask==(i+1), 2), (256, 256)),\n        )) \n        plt.imshow(masks)\n        plt.show() \n    print(''.join(['-']*32))\n    print()","metadata":{"execution":{"iopub.status.busy":"2022-10-27T23:29:05.709605Z","iopub.execute_input":"2022-10-27T23:29:05.709998Z","iopub.status.idle":"2022-10-27T23:29:25.364966Z","shell.execute_reply.started":"2022-10-27T23:29:05.709968Z","shell.execute_reply":"2022-10-27T23:29:25.36397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}