{"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":"# Notebook 2 - Train segmentation model on sagittal images and predict","metadata":{}},{"cell_type":"markdown","source":"This is the second notebook of two. The first notebook creates 16 jpgs and masks for each of the segmented cases.\n\n- The images are all sagittal\n- This notebook imports that data in order to do training and prediction\n- It also uses the CSpine helper notebook to import libraries if the internet is off\n- Training is done using fastai and the unet_learner\n- The base model is resnet34, my next experiment would like to try it with convnext from the timm library\n- Please note that only a small sample of cases were used for prediction in this initial experiment\n- I chose to let it train for all the vertebrae segmented included the thoracic vertebrae in the hopes that that would improve segmentation of the c-spine\n- Predictions are done in 2D on each of the 16 sagittal images for each case sampled","metadata":{}},{"cell_type":"code","source":"from fastai.vision.all import *\nfrom fastai.medical.imaging import *\nfrom fastcore.all import *\nimport pandas as pd\nimport pydicom\nimport numpy as np\n\nimport matplotlib.image as mpimg\nimport cv2","metadata":{"execution":{"iopub.status.busy":"2022-09-22T14:33:45.720494Z","iopub.execute_input":"2022-09-22T14:33:45.720942Z","iopub.status.idle":"2022-09-22T14:33:50.355083Z","shell.execute_reply.started":"2022-09-22T14:33:45.720854Z","shell.execute_reply":"2022-09-22T14:33:50.35403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\n\n!{sys.executable} -m pip install '../input/cspine-helper/fastai-2.7.9-py3-none-any.whl' --upgrade --no-deps --ignore-installed -q \n!{sys.executable} -m pip install '../input/cspine-helper/fastcore-1.5.21-py3-none-any.whl' -Uqq\n!{sys.executable} -m pip install '../input/cspine-helper/pylibjpeg-1.4.0-py3-none-any.whl' -q\n!{sys.executable} -m pip install '../input/cspine-helper/pylibjpeg_libjpeg-1.3.1-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl' -q\n!{sys.executable} -m pip install '../input/cspine-helper/python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl' -q\n#!{sys.executable} -m pip install '../input/cspine-helper/timm-0.6.7-py3-none-any.whl' -q","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:16:50.304315Z","iopub.execute_input":"2022-09-21T17:16:50.304926Z","iopub.status.idle":"2022-09-21T17:17:32.8731Z","shell.execute_reply.started":"2022-09-21T17:16:50.304889Z","shell.execute_reply":"2022-09-21T17:17:32.87188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\ntest_df = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/test.csv\")\ntrain_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:17:32.876209Z","iopub.execute_input":"2022-09-21T17:17:32.876601Z","iopub.status.idle":"2022-09-21T17:17:32.91782Z","shell.execute_reply.started":"2022-09-21T17:17:32.876558Z","shell.execute_reply":"2022-09-21T17:17:32.916913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# list of study_ids\n\nseg_path = Path('../input/rsna-2022-cervical-spine-fracture-detection/segmentations')\n#orig_path = Path('../input/rsna-2022-cervical-spine-fracture-detection/train_images')\n\nsegmentations = list(seg_path.iterdir())\nstudy_ids = [o.stem for o in segmentations]\n\nlen(study_ids), study_ids[0]","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:17:32.920546Z","iopub.execute_input":"2022-09-21T17:17:32.920927Z","iopub.status.idle":"2022-09-21T17:17:32.936379Z","shell.execute_reply.started":"2022-09-21T17:17:32.920889Z","shell.execute_reply":"2022-09-21T17:17:32.935358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# list of train images\nimg_dir = Path('../input/segment-cspine')\nimage_list = []\nn = 16\nfor study_id in study_ids:\n    for i in range(n):\n        image_list.append(f'{img_dir/study_id}_{i}.jpg') \n        \nlen(image_list), image_list[0]","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:17:32.938051Z","iopub.execute_input":"2022-09-21T17:17:32.938415Z","iopub.status.idle":"2022-09-21T17:17:32.954708Z","shell.execute_reply.started":"2022-09-21T17:17:32.93838Z","shell.execute_reply":"2022-09-21T17:17:32.953642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"labels = {}\nfor i in range(1,8):\n    labels[f'C{i}'] = i\nfor i in range(8,20):\n    labels[f'T{i-7}'] = i\nlabels","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:17:32.956157Z","iopub.execute_input":"2022-09-21T17:17:32.9572Z","iopub.status.idle":"2022-09-21T17:17:32.967243Z","shell.execute_reply.started":"2022-09-21T17:17:32.957162Z","shell.execute_reply":"2022-09-21T17:17:32.966298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Fastai Datablock with ImageBlock and MaskBlock using the unet_learner","metadata":{}},{"cell_type":"code","source":"#unet\n\ndblock = DataBlock(blocks    = (ImageBlock, MaskBlock(labels)),\n                   get_x     = lambda o: Path(o),\n                   get_y     = lambda o: np.load(Path(o).with_suffix('.npy')),\n                   item_tfms=[Resize(224, method='pad',pad_mode=PadMode.Zeros)],\n                   batch_tfms=[*aug_transforms(size=(224)), Normalize.from_stats(*imagenet_stats)])\ndls = dblock.dataloaders(image_list, bs = 16)\n\nlearn = unet_learner(dls, resnet34)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:17:32.96893Z","iopub.execute_input":"2022-09-21T17:17:32.969457Z","iopub.status.idle":"2022-09-21T17:17:44.061115Z","shell.execute_reply.started":"2022-09-21T17:17:32.969402Z","shell.execute_reply":"2022-09-21T17:17:44.060116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load prior trained model if available","metadata":{}},{"cell_type":"code","source":"!mkdir  'models'\n!cp  '../input/sagsegmodel/segment.h5.pth' 'models'","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:41:16.996771Z","iopub.execute_input":"2022-09-21T17:41:16.997196Z","iopub.status.idle":"2022-09-21T17:41:24.843746Z","shell.execute_reply.started":"2022-09-21T17:41:16.99716Z","shell.execute_reply":"2022-09-21T17:41:24.835138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_fn = Path('../input/sagsegmodel/segment.h5.pth')\nif model_fn.exists():\n    learn.load('segment.h5')\nelse:\n    learn.fine_tune(20)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:41:36.155372Z","iopub.execute_input":"2022-09-21T17:41:36.155964Z","iopub.status.idle":"2022-09-21T17:41:36.798763Z","shell.execute_reply.started":"2022-09-21T17:41:36.155919Z","shell.execute_reply":"2022-09-21T17:41:36.797718Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"learn.save('/kaggle/working/segment.h5')","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Make predictions on the non segmented cases to test model","metadata":{}},{"cell_type":"code","source":"#test - take a subset of the non segmented cases\n# test_amount = 5\n\n# test_ids = train_df[~train_df.StudyInstanceUID.isin(study_ids)].StudyInstanceUID.unique()[:test_amount]","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:17:44.236927Z","iopub.status.idle":"2022-09-21T17:17:44.237275Z","shell.execute_reply.started":"2022-09-21T17:17:44.237095Z","shell.execute_reply":"2022-09-21T17:17:44.237111Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#all train that are not in segmentation\ntest_ids = train_df[~train_df.StudyInstanceUID.isin(study_ids)].StudyInstanceUID.unique()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:41:41.67505Z","iopub.execute_input":"2022-09-21T17:41:41.675491Z","iopub.status.idle":"2022-09-21T17:41:41.697032Z","shell.execute_reply.started":"2022-09-21T17:41:41.675453Z","shell.execute_reply":"2022-09-21T17:41:41.695688Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Convert the cases to be segmented\n\n#### Functions to create a 3D array from the dicom image directory\n\n- Reads Dicom and applies meta data values Intercept/Slope and Window Center/Width\n- Creates a 3D array ","metadata":{}},{"cell_type":"code","source":"## Original\n\n#Adapted from Pydicom: 'Load CT slices and plot axial, sagittal and coronal images'\ndef create_3d(case):\n    files = []\n    fns = get_dicom_files(case)\n        \n    for fn in fns:\n        files.append(pydicom.dcmread(fn))\n        \n    # skip files with no InstanceNumber (eg. Scout)\n    slices = []\n    skipcount = 0\n    for f in files:\n        if hasattr(f, 'InstanceNumber'):\n            slices.append(f)\n        else:\n            skipcount = skipcount + 1\n\n    if skipcount > 0:\n        print(\"skipped, no InstanceNumber: {}\".format(skipcount))\n\n    # ensure they are in the correct order\n    slices = sorted(slices, key=lambda s: s.ImagePositionPatient[2])\n    dcm = slices[0]\n\n    # pixel aspects, assuming all slices are the same\n    ps = dcm.PixelSpacing\n    ss = dcm.SliceThickness\n    #ax_aspect = ps[1]/ps[0]\n    sag_aspect = ps[1]/ss\n    #cor_aspect = ss/ps[0]\n\n    # create 3D array\n    img_shape = list(dcm.pixel_array.shape)\n    img_shape.append(len(slices))\n    img3d = np.zeros(img_shape)\n    \n    # fill 3D array with the images from the files\n    for i, s in enumerate(slices):\n        img2d = dcm_apply_windows(s)\n        img3d[:, :, i] = img2d\n        \n    return img3d, sag_aspect, fn.parent.name\n\ndef get_window_from_dicom(dcm, default = (2000, 500)):\n    \"\"\"\n    Returns window width and window center values or first example if MultiValue\n    Strips comma from value if present (seen in a different dataset)\n    If no window width/level is provided or available, returns default.\n    \"\"\"\n    width, level = default\n\n    if \"WindowWidth\" in dcm:\n        width = dcm.WindowWidth\n        if isinstance(width, pydicom.multival.MultiValue):\n            width = float(width[0])\n        else:\n            width = float(str(width).replace(',', ''))\n\n    if \"WindowCenter\" in dcm:\n        level = dcm.WindowCenter\n        if isinstance(level, pydicom.multival.MultiValue):\n            level = float(level[0])\n        else:\n            level = float(str(level).replace(',', ''))\n            \n    return width, level\n\n\ndef dcm_apply_windows(dcm):\n    \"\"\"\n    Applies Intercept/Slope and Window Center/Width\n    \"\"\"\n    arr = dcm.pixel_array\n    #slope, intercept\n    slope = 1\n    intercept = 0\n    if \"RescaleIntercept\" in dcm and \"RescaleSlope\" in dcm:\n        intercept = int(dcm.RescaleIntercept)\n        slope = int(dcm.RescaleSlope)\n        \n    arr = arr * slope + intercept\n    \n    #window\n    width,level = get_window_from_dicom(dcm)\n    if width is not None and level is not None:\n        arr = np.clip(arr, level - width // 2, level + width // 2)\n        \n     #scale\n    arr = (arr - np.min(arr)) / np.max(arr)\n        \n    return arr","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:41:43.036945Z","iopub.execute_input":"2022-09-21T17:41:43.0374Z","iopub.status.idle":"2022-09-21T17:41:43.061255Z","shell.execute_reply.started":"2022-09-21T17:41:43.037362Z","shell.execute_reply":"2022-09-21T17:41:43.060155Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Create and save sagittal images","metadata":{}},{"cell_type":"code","source":"orig_path = Path('../input/rsna-2022-cervical-spine-fracture-detection/train_images')\nsave_dir = Path('/kaggle/working/')\n\ndef get_sag_for_predict(study_id, n = 16, spread = 160):\n    thickness = int(spread/n)\n    orig = []\n    dirname = orig_path/study_id\n\n    img, sag_aspect, folder_name = create_3d(dirname)\n    slice = int(img.shape[1]/2 - spread/2)\n    for i in range(n):\n        arr = img[:, slice + (i*thickness),:] \n        #flip and rotate so that C1 is at the top\n        arr = np.transpose(arr)\n        arr = np.flip(arr, 0)\n        orig.append(arr)    \n    return orig\n\ndef save_sag_for_predict(study_id):\n    #use half the images to prevent running out of memory on prediction\n    orig = get_sag_for_predict(study_id, n = 8, spread = 96)\n    orig = np.array(orig)  * 256\n    slices = orig.shape[0]\n\n    for i in range(slices):\n        cv2.imwrite(f'{save_dir/study_id}_{i}.jpg', orig[i,:,:])","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:41:44.209776Z","iopub.execute_input":"2022-09-21T17:41:44.210224Z","iopub.status.idle":"2022-09-21T17:41:44.226445Z","shell.execute_reply.started":"2022-09-21T17:41:44.210186Z","shell.execute_reply":"2022-09-21T17:41:44.225288Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Convert with Parallel (2 hours for prediction)","metadata":{}},{"cell_type":"code","source":"from tqdm.auto import tqdm\nfrom joblib import Parallel, delayed\n\n_= Parallel(n_jobs=4)(delayed(save_sag_for_predict)(study_id) for study_id in tqdm(test_ids))","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:41:46.609604Z","iopub.execute_input":"2022-09-21T17:41:46.610086Z","iopub.status.idle":"2022-09-21T17:43:09.368158Z","shell.execute_reply.started":"2022-09-21T17:41:46.610044Z","shell.execute_reply":"2022-09-21T17:43:09.36523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#get list of converted images\ntest_image_list = []\n#n = 16\nn = 8\nfor study_id in test_ids:\n    for i in range(n):\n        test_image_list.append(save_dir/f'{study_id}_{i}.jpg') \n        \nlen(test_image_list)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:17:44.248708Z","iopub.status.idle":"2022-09-21T17:17:44.249438Z","shell.execute_reply.started":"2022-09-21T17:17:44.24918Z","shell.execute_reply":"2022-09-21T17:17:44.249205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#preds = []\nfor  fn in test_image_list:\n    pred = learn.predict(fn)\n    #preds.append(pred[0])\n    np.save(f'{fn.stem}.npy', pred[0]) \n    del(pred)","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:17:44.250746Z","iopub.status.idle":"2022-09-21T17:17:44.251466Z","shell.execute_reply.started":"2022-09-21T17:17:44.251218Z","shell.execute_reply":"2022-09-21T17:17:44.251243Z"},"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# n = 30\n# fn = test_image_list[n]\n# im = cv2.imread(str(fn))\n# fig, ax = plt.subplots(1,2, figsize=(10,10))\n\n# ax[0].imshow(im)\n# ax[0].set_aspect(2)\n# mask = np.load(f'{fn.stem}.npy')\n# ax[1].imshow(mask)\n# ax[1].set_aspect(2)\n# plt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:17:44.252835Z","iopub.status.idle":"2022-09-21T17:17:44.253579Z","shell.execute_reply.started":"2022-09-21T17:17:44.253334Z","shell.execute_reply":"2022-09-21T17:17:44.253359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Look at some more samples","metadata":{}},{"cell_type":"code","source":"n = 30\nr = 4\nc = 4\n\nfig, axs = plt.subplots(r,c, figsize=(15,15))\naxs = axs.flatten()\n\nfor i in range(0, r*c, 2):\n    fn = test_image_list[i+n]\n    im = cv2.imread(str(fn))\n\n    axs[i].imshow(im)\n    axs[i].set_aspect(2)\n    axs[i].axis('off')\n    mask = np.load(f'{fn.stem}.npy')\n    axs[i+1].imshow(mask)\n    axs[i+1].set_aspect(2)\n    axs[i+1].axis('off')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-21T17:17:44.254911Z","iopub.status.idle":"2022-09-21T17:17:44.255636Z","shell.execute_reply.started":"2022-09-21T17:17:44.25538Z","shell.execute_reply":"2022-09-21T17:17:44.255405Z"},"trusted":true},"execution_count":null,"outputs":[]}]}