{"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":"![](https://raw.githubusercontent.com/fepegar/torchio/main/docs/source/favicon_io/for_readme_2000x462.png)\n\n[TorchIO](http://torchio.org/) is a Python library for loading, preprocessing, augmentation and sampling for multidimensional medical images in deep learning.\n\nWe can leverage the new [`RSNACervicalSpineFracture`](https://torchio.readthedocs.io/datasets.html#rsnacervicalspinefracture) class to avoid dealing with the complex [DICOM](https://www.dicomstandard.org/) and [NIfTI](https://nifti.nimh.nih.gov/) formats for medical images.","metadata":{}},{"cell_type":"code","source":"%pip install --quiet light-the-torch && ltt install torch\n%pip install --quiet torchio\n%pip install --quiet pytorch-lightning\n%pip install --quiet monai","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.status.busy":"2022-11-28T21:23:54.809995Z","iopub.execute_input":"2022-11-28T21:23:54.81096Z","iopub.status.idle":"2022-11-28T21:24:39.202118Z","shell.execute_reply.started":"2022-11-28T21:23:54.810906Z","shell.execute_reply":"2022-11-28T21:24:39.200461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import copy\nimport multiprocessing\n\nimport torchio as tio\nimport torch\nimport monai\nimport pytorch_lightning as pl\nfrom torch.utils.data import random_split, DataLoader\nfrom tqdm import tqdm\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.optim.lr_scheduler as lr_scheduler\n\nimport matplotlib as mpl\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport pandas as pd\nimport numpy as np\n\nimport time\nfrom pathlib import Path\nfrom datetime import datetime\n\n\ntorch.manual_seed(0)\nnum_workers = multiprocessing.cpu_count()\nmpl.rcParams['figure.figsize'] = 12, 8","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:39.204694Z","iopub.execute_input":"2022-11-28T21:24:39.205082Z","iopub.status.idle":"2022-11-28T21:24:39.214263Z","shell.execute_reply.started":"2022-11-28T21:24:39.205046Z","shell.execute_reply":"2022-11-28T21:24:39.213075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configure","metadata":{}},{"cell_type":"code","source":"BATCH_SIZE = 2\nLEARNING_RATE = 0.0001\nN_EPOCHS = 3\nPATIENCE = 3\n\n# Config device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:39.216175Z","iopub.execute_input":"2022-11-28T21:24:39.216536Z","iopub.status.idle":"2022-11-28T21:24:39.256154Z","shell.execute_reply.started":"2022-11-28T21:24:39.216505Z","shell.execute_reply":"2022-11-28T21:24:39.254875Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset\n\nLet's create an instance of `RSNACervicalSpineFracture`, which inherits directly from [`SubjectsDataset`](https://torchio.readthedocs.io/data/dataset.html) and indirectly from [`torch.utils.data.Dataset`](https://pytorch.org/docs/stable/data.html#torch.utils.data.Dataset).","metadata":{}},{"cell_type":"code","source":"root_dir = '/kaggle/input/rsna-2022-cervical-spine-fracture-detection'\n# root_dir = '/root/input/rsna-2022-cervical-spine-fracture-detection'\ndataset = tio.datasets.RSNACervicalSpineFracture(root_dir, add_segmentations=True, add_bounding_boxes=True)\nprint('Number of subjects:', len(dataset))","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:39.258061Z","iopub.execute_input":"2022-11-28T21:24:39.25911Z","iopub.status.idle":"2022-11-28T21:24:41.701511Z","shell.execute_reply.started":"2022-11-28T21:24:39.259073Z","shell.execute_reply":"2022-11-28T21:24:41.699868Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Subject\n\nLet's inspect the first [`Subject`](https://torchio.readthedocs.io/data/subject.html) in the dataset.","metadata":{}},{"cell_type":"code","source":"# for subject in dataset.dry_iter():\n#     print(subject['StudyInstanceUID'])","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:41.704414Z","iopub.execute_input":"2022-11-28T21:24:41.704752Z","iopub.status.idle":"2022-11-28T21:24:41.710569Z","shell.execute_reply.started":"2022-11-28T21:24:41.704724Z","shell.execute_reply":"2022-11-28T21:24:41.708373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for subject in dataset.dry_iter():\n    if subject['StudyInstanceUID']== '1.2.826.0.1.3680043.581':\n        subject_with_seg = subject\n        break\nsubject_with_seg.plot()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:41.711723Z","iopub.execute_input":"2022-11-28T21:24:41.712118Z","iopub.status.idle":"2022-11-28T21:24:44.624792Z","shell.execute_reply.started":"2022-11-28T21:24:41.712086Z","shell.execute_reply":"2022-11-28T21:24:44.623938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"first_subject = dataset[4]\nfirst_subject","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:44.625863Z","iopub.execute_input":"2022-11-28T21:24:44.626545Z","iopub.status.idle":"2022-11-28T21:24:49.931837Z","shell.execute_reply.started":"2022-11-28T21:24:44.626513Z","shell.execute_reply":"2022-11-28T21:24:49.930867Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for key, value in first_subject.items():\n    print(f'{key}: {value}')","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:49.932989Z","iopub.execute_input":"2022-11-28T21:24:49.933674Z","iopub.status.idle":"2022-11-28T21:24:49.940289Z","shell.execute_reply.started":"2022-11-28T21:24:49.933648Z","shell.execute_reply":"2022-11-28T21:24:49.938937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We can see that the patient's C1 and C2 vertebrae are fractured, according to the labels. We can also see that the CT scan is an instance of [`ScalarImage`](https://torchio.readthedocs.io/data/image.html#torchio.ScalarImage) with $512 \\times 512 \\times 243$ voxels, in-plane pixel spacing of 0.44 mm, slice thickness of 0.80 mm, [LPS+ orientation](https://nipy.org/nibabel/image_orientation.html) data type \"short tensor\" (or signed 16-bit integers) and takes 121.5 MiB of memory.\n\nLet's take a look at the 3D image in the subject.","metadata":{}},{"cell_type":"code","source":"first_subject.plot()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:49.941682Z","iopub.execute_input":"2022-11-28T21:24:49.942157Z","iopub.status.idle":"2022-11-28T21:24:50.933147Z","shell.execute_reply.started":"2022-11-28T21:24:49.942132Z","shell.execute_reply":"2022-11-28T21:24:50.931695Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The [CT scan](https://www.nhs.uk/conditions/ct-scan/) looks black where there is no reconstruction data, and the contrast in the region of interest is not great. Let's look at the intensity distribution.","metadata":{}},{"cell_type":"code","source":"first_subject.ct.hist()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:50.934576Z","iopub.execute_input":"2022-11-28T21:24:50.934966Z","iopub.status.idle":"2022-11-28T21:24:54.432949Z","shell.execute_reply.started":"2022-11-28T21:24:50.9349Z","shell.execute_reply":"2022-11-28T21:24:54.431829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Intensity preprocessing\n\nWe can look at the [Hounsfield scale](https://en.wikipedia.org/wiki/Hounsfield_scale) to keep only intensity values within a sensible range using the [`Clamp`](https://torchio.readthedocs.io/transforms/preprocessing.html#clamp) transform.","metadata":{}},{"cell_type":"code","source":"HOUNSFIELD_AIR, HOUNSFIELD_BONE = -1000, 1900\nclamp = tio.Clamp(out_min=HOUNSFIELD_AIR, out_max=HOUNSFIELD_BONE)\nfirst_subject_clamped = clamp(first_subject)\nfirst_subject_clamped.ct.hist()\nfirst_subject_clamped.plot()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:54.434403Z","iopub.execute_input":"2022-11-28T21:24:54.434755Z","iopub.status.idle":"2022-11-28T21:24:58.00958Z","shell.execute_reply.started":"2022-11-28T21:24:54.434717Z","shell.execute_reply":"2022-11-28T21:24:58.008788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"That looks much better! Let's now normalize the values to [0, 1], a common practice in neural network training. We will also saturate values in the first and last half percentiles, to remove potential outliers. Clamping and rescaling can be concatenated using [`Compose`](https://torchio.readthedocs.io/transforms/augmentation.html#compose) to create an intensity preprocessing transform.","metadata":{}},{"cell_type":"code","source":"rescale = tio.RescaleIntensity(percentiles=(0.5, 99.5))\npreprocess_intensity = tio.Compose([\n    clamp,\n    rescale,\n])","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:58.010819Z","iopub.execute_input":"2022-11-28T21:24:58.011635Z","iopub.status.idle":"2022-11-28T21:24:58.016426Z","shell.execute_reply.started":"2022-11-28T21:24:58.011603Z","shell.execute_reply":"2022-11-28T21:24:58.015464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"first_subject_preprocessed = preprocess_intensity(first_subject)\nfirst_subject_preprocessed.ct.hist(show=False), plt.ylim(0, 1e6)\nfirst_subject_preprocessed.plot()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:24:58.017905Z","iopub.execute_input":"2022-11-28T21:24:58.018394Z","iopub.status.idle":"2022-11-28T21:25:02.569082Z","shell.execute_reply.started":"2022-11-28T21:24:58.01835Z","shell.execute_reply":"2022-11-28T21:25:02.567815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's look for a subject with an associated segmentation. We will use the [`dry_iter` method of `SubjectsDataset`](https://torchio.readthedocs.io/data/dataset.html#torchio.data.SubjectsDataset.dry_iter) as we don't want to load any data here.","metadata":{}},{"cell_type":"code","source":"for subject in dataset.dry_iter():\n    if 'seg' in subject:\n        subject_with_seg = subject\n        break\nsubject_with_seg.plot(reorient=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:02.574466Z","iopub.execute_input":"2022-11-28T21:25:02.574805Z","iopub.status.idle":"2022-11-28T21:25:05.876173Z","shell.execute_reply.started":"2022-11-28T21:25:02.574779Z","shell.execute_reply":"2022-11-28T21:25:05.8749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Spatial preprocessing\n\nThe CT scan and the corresponding segmentation are not aligned! Let's look at their orientations.","metadata":{}},{"cell_type":"code","source":"print('CT orientation:', subject_with_seg.ct.orientation)\nprint('Segmentation orientation:', subject_with_seg.seg.orientation)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:05.877726Z","iopub.execute_input":"2022-11-28T21:25:05.878614Z","iopub.status.idle":"2022-11-28T21:25:05.885525Z","shell.execute_reply.started":"2022-11-28T21:25:05.87857Z","shell.execute_reply":"2022-11-28T21:25:05.8844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"It looks like voxels along the second spatial dimension grow towards the back (Posterior) in the CT and towards the front (Anterior) in the segmentation. That's why the images looked flipped with respect to the [coronal](https://en.wikipedia.org/wiki/Coronal_plane) plane.\n\nWe can normalize the orientation (to RAS+) using [`ToCanonical`](https://torchio.readthedocs.io/transforms/preprocessing.html#tocanonical). Also, as the resolution is quite high, we will downsample to a sensible value (1 mm isotropic) for faster computations using [`Resample`](https://torchio.readthedocs.io/transforms/preprocessing.html#resample). These two transforms are implemented on top of [NiBabel](https://nipy.org/nibabel/) and [SimpleITK](https://simpleitk.org/), respectively.\n\nAs before, we use `Compose` to concatenate preprocessing transforms.","metadata":{}},{"cell_type":"markdown","source":"## ToCanonical","metadata":{}},{"cell_type":"code","source":"normalize_orientation = tio.ToCanonical()\nnew_subject_with_seg = normalize_orientation(subject_with_seg)\nnew_subject_with_seg.plot()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:05.886988Z","iopub.execute_input":"2022-11-28T21:25:05.888088Z","iopub.status.idle":"2022-11-28T21:25:07.145437Z","shell.execute_reply.started":"2022-11-28T21:25:05.888052Z","shell.execute_reply":"2022-11-28T21:25:07.143939Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('CT orientation before canonical:', subject_with_seg.ct.orientation)\nprint('CT orientation after canonical:', new_subject_with_seg.ct.orientation)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:07.147381Z","iopub.execute_input":"2022-11-28T21:25:07.147819Z","iopub.status.idle":"2022-11-28T21:25:07.155872Z","shell.execute_reply.started":"2022-11-28T21:25:07.147779Z","shell.execute_reply":"2022-11-28T21:25:07.154571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Resample","metadata":{}},{"cell_type":"code","source":"downsample = tio.Resample(1)\nnew_subject_with_seg = downsample(new_subject_with_seg)\nnew_subject_with_seg.plot()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:07.157713Z","iopub.execute_input":"2022-11-28T21:25:07.158046Z","iopub.status.idle":"2022-11-28T21:25:10.298092Z","shell.execute_reply.started":"2022-11-28T21:25:07.158022Z","shell.execute_reply":"2022-11-28T21:25:10.297186Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('before resample:', subject_with_seg.ct)\nprint('after resample:', new_subject_with_seg.ct)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:10.299244Z","iopub.execute_input":"2022-11-28T21:25:10.299548Z","iopub.status.idle":"2022-11-28T21:25:10.30823Z","shell.execute_reply.started":"2022-11-28T21:25:10.299521Z","shell.execute_reply":"2022-11-28T21:25:10.307001Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CropOrPad","metadata":{}},{"cell_type":"code","source":"crop = tio.CropOrPad((224,224,224))\nnew_subject_with_seg = crop(new_subject_with_seg)\nnew_subject_with_seg.plot()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:10.309515Z","iopub.execute_input":"2022-11-28T21:25:10.309806Z","iopub.status.idle":"2022-11-28T21:25:11.27461Z","shell.execute_reply.started":"2022-11-28T21:25:10.30978Z","shell.execute_reply":"2022-11-28T21:25:11.27395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print('before CropOrPad:', subject_with_seg.ct)\nprint('after CropOrPad:', new_subject_with_seg.ct)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:11.275569Z","iopub.execute_input":"2022-11-28T21:25:11.276619Z","iopub.status.idle":"2022-11-28T21:25:11.282678Z","shell.execute_reply.started":"2022-11-28T21:25:11.276588Z","shell.execute_reply":"2022-11-28T21:25:11.281614Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Compose together","metadata":{}},{"cell_type":"code","source":"\npreprocess_spatial = tio.Compose([\n    normalize_orientation,\n    downsample,\n    crop\n])\n\npreprocess = tio.Compose([\n    preprocess_intensity,\n    preprocess_spatial,\n])","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:11.284428Z","iopub.execute_input":"2022-11-28T21:25:11.284709Z","iopub.status.idle":"2022-11-28T21:25:11.292473Z","shell.execute_reply.started":"2022-11-28T21:25:11.284682Z","shell.execute_reply":"2022-11-28T21:25:11.291485Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"subject_with_seg_preprocessed = preprocess(subject_with_seg)\nsubject_with_seg_preprocessed.plot(reorient=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:11.293891Z","iopub.execute_input":"2022-11-28T21:25:11.294246Z","iopub.status.idle":"2022-11-28T21:25:15.05023Z","shell.execute_reply.started":"2022-11-28T21:25:11.294213Z","shell.execute_reply":"2022-11-28T21:25:15.049125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"That's better! We actually didn't need to pass `reorient=False` to the `plot` method, as all that does is apply `ToCanonical` internally, which we have already done.","metadata":{}},{"cell_type":"markdown","source":"## Data augmentation\n\nThere are also many [augmentation transforms](https://torchio.readthedocs.io/transforms/augmentation.html) available in TorchIO. Let's compose some appropriate ones. If we were using MRI, we should definitely leverage some of the [MRI k-space artifact simulation transforms](https://torchio.readthedocs.io/transforms/augmentation.html#randommotion). Here, we will use simpler ones.","metadata":{}},{"cell_type":"code","source":"augment = tio.Compose([\n    tio.RandomAnisotropy(p=0.25),              # make images look anisotropic 25% of times\n    tio.RandomAffine(),\n    tio.RandomFlip(),\n    tio.RandomNoise(p=0.25),                   # Gaussian noise 25% of times\n    tio.RandomGamma(p=0.5),                    # Randomly change contrast of an image by raising its values to the power gamma.\n#     tio.OneOf({                                # either\n#         tio.RandomMotion(): 1,                 # random motion artifact\n#         tio.RandomSpike(): 2,                  # or spikes\n#         tio.RandomGhosting(): 2,               # or ghosts\n#     }, p=0.25),                                 # applied to \n])","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:15.051706Z","iopub.execute_input":"2022-11-28T21:25:15.052175Z","iopub.status.idle":"2022-11-28T21:25:15.058515Z","shell.execute_reply.started":"2022-11-28T21:25:15.052141Z","shell.execute_reply":"2022-11-28T21:25:15.057405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"In this case, we always apply an affine (anisotropic scaling + rotation) transform. The others are applied with a probability of 0.25 or 0.5.\n\nInternally, [`RandomFlip`](https://torchio.readthedocs.io/transforms/augmentation.html#randomflip) uses 0.5). By default, flipping is applied only around the sagittal plane, i.e., left-right. But if you feel adventurous, feel free to explore other axes as well!\n\nLet's set our training and validation transforms. We want augmentation during training only; our validation transform is simply our deterministic preprocessing.","metadata":{}},{"cell_type":"code","source":"train_transform = tio.Compose([\n    preprocess,\n    augment,\n])\nval_transform = preprocess","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:15.060236Z","iopub.execute_input":"2022-11-28T21:25:15.060593Z","iopub.status.idle":"2022-11-28T21:25:15.073805Z","shell.execute_reply.started":"2022-11-28T21:25:15.060563Z","shell.execute_reply":"2022-11-28T21:25:15.072834Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's look at some preprocessed and randomly augmented variations of our subject.","metadata":{}},{"cell_type":"code","source":"val_transform(first_subject).plot()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:15.074843Z","iopub.execute_input":"2022-11-28T21:25:15.075183Z","iopub.status.idle":"2022-11-28T21:25:17.956887Z","shell.execute_reply.started":"2022-11-28T21:25:15.075156Z","shell.execute_reply":"2022-11-28T21:25:17.956079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"first_subject.ct","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:17.957835Z","iopub.execute_input":"2022-11-28T21:25:17.95826Z","iopub.status.idle":"2022-11-28T21:25:17.964745Z","shell.execute_reply.started":"2022-11-28T21:25:17.958234Z","shell.execute_reply":"2022-11-28T21:25:17.963849Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform(first_subject).ct","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:17.966149Z","iopub.execute_input":"2022-11-28T21:25:17.966434Z","iopub.status.idle":"2022-11-28T21:25:20.55901Z","shell.execute_reply.started":"2022-11-28T21:25:17.966408Z","shell.execute_reply":"2022-11-28T21:25:20.557746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for _ in range(5):\n    subject_with_seg_augmented = train_transform(subject_with_seg)\n    subject_with_seg_augmented.plot()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:20.560397Z","iopub.execute_input":"2022-11-28T21:25:20.560657Z","iopub.status.idle":"2022-11-28T21:25:41.772526Z","shell.execute_reply.started":"2022-11-28T21:25:20.560631Z","shell.execute_reply":"2022-11-28T21:25:41.77169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Intensity transforms are only applied to instances of `ScalarImage`, whereas spatial transforms are applied also to instances of [`LabelMap`](https://torchio.readthedocs.io/data/image.html#torchio.LabelMap). Of course, the same random parameters for spatial transforms are applied to all images within the same subject!","metadata":{}},{"cell_type":"markdown","source":"# Add label to dataset","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\nno_segs_dataset = tio.datasets.RSNACervicalSpineFracture(root_dir)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:41.77383Z","iopub.execute_input":"2022-11-28T21:25:41.774294Z","iopub.status.idle":"2022-11-28T21:25:43.487496Z","shell.execute_reply.started":"2022-11-28T21:25:41.774257Z","shell.execute_reply":"2022-11-28T21:25:43.486112Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Dataset for train/valid sets only\nclass RSNADataset(Dataset):\n    # Initialise\n    def __init__(self, dataset = no_segs_dataset, df_table = train_df, transform=None):\n        super().__init__()\n        \n        self.df_table = df_table.reset_index(drop=True)\n        self.transform = transform\n        self.set_transform(transform)\n        self.targets = ['C1','C2','C3','C4','C5','C6','C7','patient_overall']\n        # Populate labels\n        self.labels = self.df_table[self.targets].values\n        \n    def set_transform(self, transform= None):\n        \"\"\"Set the :attr:`transform` attribute.\n\n        Args:\n            transform: Callable object, typically an subclass of\n                :class:`torchio.transforms.Transform`.\n        \"\"\"\n        if transform is not None and not callable(transform):\n            message = (\n                'The transform must be a callable object,'\n                f' but it has type {type(transform)}'\n            )\n            raise ValueError(message)\n        self.transform = transform\n        \n    def __getitem__(self, index):\n        vol = dataset[index]['ct']\n        target = torch.tensor(self.labels[index])\n        \n        if self.transform:\n            vol = self.transform(vol)\n        \n        subject = tio.Subject(\n                    image = vol, label = target)\n        return subject\n    \n    # Length of dataset\n    def __len__(self):\n        return len(self.df_table['StudyInstanceUID'])\n","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:43.488927Z","iopub.execute_input":"2022-11-28T21:25:43.489312Z","iopub.status.idle":"2022-11-28T21:25:43.50115Z","shell.execute_reply.started":"2022-11-28T21:25:43.489277Z","shell.execute_reply":"2022-11-28T21:25:43.499772Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"no_seg_set = RSNADataset(no_segs_dataset, train_df)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:43.503402Z","iopub.execute_input":"2022-11-28T21:25:43.503881Z","iopub.status.idle":"2022-11-28T21:25:43.516824Z","shell.execute_reply.started":"2022-11-28T21:25:43.503844Z","shell.execute_reply":"2022-11-28T21:25:43.516087Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"no_seg_set[0]","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:43.518153Z","iopub.execute_input":"2022-11-28T21:25:43.518661Z","iopub.status.idle":"2022-11-28T21:25:48.720153Z","shell.execute_reply.started":"2022-11-28T21:25:43.518634Z","shell.execute_reply":"2022-11-28T21:25:48.719251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's use 80% of our data for training and 20% for validation.","metadata":{}},{"cell_type":"code","source":"train_to_val_ratio = 0.2\nnum_subjects = len(dataset)\nnum_train_subjects = int(train_to_val_ratio * num_subjects)\nnum_val_subjects = num_subjects - num_train_subjects\nprint('Number of subjects for training:  ', num_train_subjects)\nprint('Number of subjects for validation: ', num_val_subjects)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.721644Z","iopub.execute_input":"2022-11-28T21:25:48.722253Z","iopub.status.idle":"2022-11-28T21:25:48.729411Z","shell.execute_reply.started":"2022-11-28T21:25:48.722219Z","shell.execute_reply":"2022-11-28T21:25:48.728071Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We will assume we won't need the segmentations or bounding boxes during training. We create a new instance of `RSNACervicalSpineFracture` and random split it between training and validation sets using PyTorch.","metadata":{}},{"cell_type":"code","source":"train_set, val_set = torch.utils.data.random_split(no_seg_set, [num_train_subjects, num_val_subjects])","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.731242Z","iopub.execute_input":"2022-11-28T21:25:48.731567Z","iopub.status.idle":"2022-11-28T21:25:48.746855Z","shell.execute_reply.started":"2022-11-28T21:25:48.731539Z","shell.execute_reply":"2022-11-28T21:25:48.745535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_set[0]","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.748341Z","iopub.execute_input":"2022-11-28T21:25:48.748856Z","iopub.status.idle":"2022-11-28T21:25:48.757956Z","shell.execute_reply.started":"2022-11-28T21:25:48.748826Z","shell.execute_reply":"2022-11-28T21:25:48.756714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# val_set[0]","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.759192Z","iopub.execute_input":"2022-11-28T21:25:48.760716Z","iopub.status.idle":"2022-11-28T21:25:48.769817Z","shell.execute_reply.started":"2022-11-28T21:25:48.760666Z","shell.execute_reply":"2022-11-28T21:25:48.768813Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"We use the training and validation transforms compute above for our training and validation sets.","metadata":{}},{"cell_type":"code","source":"# Both subsets share the \"dataset\" attribute, so the last statement would\n# overwrite the previous one unless we deep-copy one of the subsets\nval_set = copy.deepcopy(val_set)\ntrain_set.dataset.set_transform(train_transform)\nval_set.dataset.set_transform(val_transform)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.771067Z","iopub.execute_input":"2022-11-28T21:25:48.771421Z","iopub.status.idle":"2022-11-28T21:25:48.783755Z","shell.execute_reply.started":"2022-11-28T21:25:48.771388Z","shell.execute_reply":"2022-11-28T21:25:48.782958Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# train_set[0]","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.785495Z","iopub.execute_input":"2022-11-28T21:25:48.787122Z","iopub.status.idle":"2022-11-28T21:25:48.799228Z","shell.execute_reply.started":"2022-11-28T21:25:48.78706Z","shell.execute_reply":"2022-11-28T21:25:48.79799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# val_set[0]","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.800807Z","iopub.execute_input":"2022-11-28T21:25:48.801198Z","iopub.status.idle":"2022-11-28T21:25:48.807958Z","shell.execute_reply.started":"2022-11-28T21:25:48.801167Z","shell.execute_reply":"2022-11-28T21:25:48.807163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's look at some of the images in our training set. They are randomly transformed by our `train_transform`.","metadata":{}},{"cell_type":"code","source":"# for i in range(1):\n#     train_set[i].plot()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.814084Z","iopub.execute_input":"2022-11-28T21:25:48.814398Z","iopub.status.idle":"2022-11-28T21:25:48.820145Z","shell.execute_reply.started":"2022-11-28T21:25:48.814371Z","shell.execute_reply":"2022-11-28T21:25:48.818824Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Let's now look at some of the images in the validation set, which are not augmented (only preprocessed).","metadata":{}},{"cell_type":"code","source":"# for i in range(1):\n#     val_set[i].plot()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.821524Z","iopub.execute_input":"2022-11-28T21:25:48.822437Z","iopub.status.idle":"2022-11-28T21:25:48.829391Z","shell.execute_reply.started":"2022-11-28T21:25:48.8224Z","shell.execute_reply":"2022-11-28T21:25:48.828693Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Experiment with Small Sample","metadata":{}},{"cell_type":"code","source":"from torch.utils.data import Subset\nnum_subjects = len(dataset)\ndataset_indices = list(range(num_subjects))\nnp.random.shuffle(dataset_indices)\n\nn_train_sample = 20\nn_val_sample = 10\n\ntrain_idx, val_idx = torch.Tensor(dataset_indices[:n_train_sample]), torch.Tensor(dataset_indices[n_train_sample:(n_train_sample+n_val_sample)])\n\ntrain_indices = train_idx.nonzero().reshape(-1)\nval_indices = val_idx.nonzero().reshape(-1)\n\ntrain_subset = Subset(train_set, train_indices)\nval_subset = Subset(val_set, val_indices)\n\ntrain_loader = DataLoader(dataset=train_subset, shuffle=False, batch_size=BATCH_SIZE)\nval_loader = DataLoader(dataset=val_subset, shuffle=False, batch_size=BATCH_SIZE)\n\nprint(len(train_loader))\nprint(len(val_loader))","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.831112Z","iopub.execute_input":"2022-11-28T21:25:48.831519Z","iopub.status.idle":"2022-11-28T21:25:48.848209Z","shell.execute_reply.started":"2022-11-28T21:25:48.831486Z","shell.execute_reply":"2022-11-28T21:25:48.846634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Torch DataLoader","metadata":{}},{"cell_type":"code","source":"# train_loader = DataLoader(dataset=train_set, batch_size=BATCH_SIZE, shuffle=True)\n# valid_loader = DataLoader(dataset=val_set, batch_size=BATCH_SIZE, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.849383Z","iopub.execute_input":"2022-11-28T21:25:48.850093Z","iopub.status.idle":"2022-11-28T21:25:48.855892Z","shell.execute_reply.started":"2022-11-28T21:25:48.850058Z","shell.execute_reply":"2022-11-28T21:25:48.854676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Sanity Check","metadata":{}},{"cell_type":"code","source":"# Experiment with architecture\narr = np.ones((4,1,224,224,224))\nx = torch.tensor(arr, dtype=torch.float32)\n\nconv1 = nn.Conv3d(in_channels=1, out_channels=16, kernel_size=7, stride=1, padding=0)\npool = nn.MaxPool3d(kernel_size=2, stride=2, padding=0)\nnorm1 = nn.BatchNorm3d(num_features=16)\nconv2 = nn.Conv3d(in_channels=16, out_channels=32, kernel_size=3, stride=1, padding=0)\nnorm2 = nn.BatchNorm3d(num_features=32)\nconv3 = nn.Conv3d(in_channels=32, out_channels=64, kernel_size=3, stride=1, padding=0)\nnorm3 = nn.BatchNorm3d(num_features=64)\navg = nn.AdaptiveAvgPool3d((7, 1, 1))\nflat = nn.Flatten()\nrelu = nn.ReLU()\nlin1 = nn.Linear(in_features=32*13*13*13, out_features=256)\nlin2 = nn.Linear(in_features=256, out_features=8)\n\nout = conv1(x)\nout = relu(out)\nout = pool(out)\nout = norm1(out)\nprint(out.shape)\n\nout = conv2(out)\nout = relu(out)\nout = pool(out)\nout = norm2(out)\nprint(out.shape)\n\nout = conv3(out)\nout = relu(out)\nout = pool(out)\nout = norm3(out)\nprint(out.shape)\n\nout = avg(out)\nprint(out.shape)\nout = flat(out)\nprint(out.shape)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:48.857244Z","iopub.execute_input":"2022-11-28T21:25:48.858174Z","iopub.status.idle":"2022-11-28T21:25:57.953432Z","shell.execute_reply.started":"2022-11-28T21:25:48.858141Z","shell.execute_reply":"2022-11-28T21:25:57.952523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"# 3D convolutional neural network\nclass Conv3DNet(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n        # Layers\n        self.conv1 = nn.Conv3d(in_channels=1, out_channels=16, kernel_size=7, stride=1, padding=0)\n        self.pool = nn.MaxPool3d(kernel_size=2, stride=2, padding=0)\n        self.norm1 = nn.BatchNorm3d(num_features=16)\n        self.conv2 = nn.Conv3d(in_channels=16, out_channels=32, kernel_size=3, stride=1, padding=0)\n        self.norm2 = nn.BatchNorm3d(num_features=32)\n        self.conv3 = nn.Conv3d(in_channels=32, out_channels=64, kernel_size=3, stride=1, padding=0)\n        self.norm3 = nn.BatchNorm3d(num_features=64)\n        self.avg = nn.AdaptiveAvgPool3d((7, 1, 1))\n        self.flat = nn.Flatten()\n        self.relu = nn.ReLU()\n        self.lin1 = nn.Linear(in_features=448, out_features=128)\n        self.lin2 = nn.Linear(in_features=128, out_features=8)\n        \n    def forward(self, x):\n        # Conv block 1\n        out = self.conv1(x)\n        out = self.relu(out)\n        out = self.pool(out)\n        out = self.norm1(out)\n        \n        # Conv block 2\n        out = self.conv2(out)\n        out = self.relu(out)\n        out = self.pool(out)\n        out = self.norm2(out)\n        \n        # Conv block 3\n        out = self.conv3(out)\n        out = self.relu(out)\n        out = self.pool(out)\n        out = self.norm3(out)\n        \n        # Average & flatten\n        out = self.avg(out)\n        out = self.flat(out)\n        \n        # Fully connected layer\n        out = self.lin1(out)\n        out = self.relu(out)\n        \n        # Output layer (no sigmoid needed)\n        out = self.lin2(out)\n        \n        return out\n\nmodel = Conv3DNet().to(device)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:57.954742Z","iopub.execute_input":"2022-11-28T21:25:57.955379Z","iopub.status.idle":"2022-11-28T21:25:57.970761Z","shell.execute_reply.started":"2022-11-28T21:25:57.955353Z","shell.execute_reply":"2022-11-28T21:25:57.969636Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"loss_fn = nn.BCEWithLogitsLoss(reduction='none')\n\ncompetition_weights = {\n    '-' : torch.tensor([1, 1, 1, 1, 1, 1, 1, 7], dtype=torch.float, device=device),\n    '+' : torch.tensor([2, 2, 2, 2, 2, 2, 2, 14], dtype=torch.float, device=device),\n}\n\n# y_hat.shape = (batch_size, num_classes)\n# y.shape = (batch_size, num_classes)\n\n# with row-wise weights normalization (https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/344565)\ndef competiton_loss_row_norm(y_hat, y):\n    loss = loss_fn(y_hat, y.to(y_hat.dtype))\n    weights = y * competition_weights['+'] + (1 - y) * competition_weights['-']\n    loss = (loss * weights).sum(axis=1)\n    w_sum = weights.sum(axis=1)\n    loss = torch.div(loss, w_sum)\n    return loss.mean()","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:57.972304Z","iopub.execute_input":"2022-11-28T21:25:57.9726Z","iopub.status.idle":"2022-11-28T21:25:57.984551Z","shell.execute_reply.started":"2022-11-28T21:25:57.972572Z","shell.execute_reply":"2022-11-28T21:25:57.982894Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Adam optimiser\noptimiser = optim.AdamW(params=model.parameters(), lr=LEARNING_RATE)\n\n# Learning rate scheduler\nscheduler = lr_scheduler.CosineAnnealingLR(optimiser, T_max=N_EPOCHS)","metadata":{"execution":{"iopub.status.busy":"2022-11-28T21:25:57.986574Z","iopub.execute_input":"2022-11-28T21:25:57.987284Z","iopub.status.idle":"2022-11-28T21:25:57.999342Z","shell.execute_reply.started":"2022-11-28T21:25:57.987243Z","shell.execute_reply":"2022-11-28T21:25:57.997775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}