{"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":"<h1 style=\"padding: 10px;\n           background-color:#5642C5;\n           color:whitesmoke;\">Vertebrae Classification</h1>","metadata":{}},{"cell_type":"markdown","source":"<img style=\"float:right\" src=\"https://www.sci-info-pages.com/wp-content/media/spinal-cord-segments.png\" width=\"250\">\n\n<h3><b>Task:</b></h3>\n\n<div>\nIn this notebook we will create a model that will classifies the cervical vertebraes from CT scans. Images are stored in DICOM format and vertebrae labels we have to extract from NIFTI segmentation masks. For each patient we have batch of CT scans from different depth levels. Single DICOM image may contain more than one vertebrae.\n</div>\n\n<h3><b>Solution:</b></h3>\n\n<div>\nFor vertebrae classification we will use <a href=\"https://monai.io/\">Monai</a> framework. This tool help as easily process DICOM images and create pre-trained model for 2D image classification task. We will handle training and validation loop using <a href=\"https://www.pytorchlightning.ai/\">PyTorch Lightning</a>.\n</div>\n\n","metadata":{}},{"cell_type":"code","source":"! pip -q install python-gdcm\n! pip -q install pylibjpeg pylibjpeg-libjpeg pydicom\n!pip -q install monai","metadata":{"execution":{"iopub.status.busy":"2022-09-15T11:58:04.492164Z","iopub.execute_input":"2022-09-15T11:58:04.492707Z","iopub.status.idle":"2022-09-15T11:58:16.832262Z","shell.execute_reply.started":"2022-09-15T11:58:04.492588Z","shell.execute_reply":"2022-09-15T11:58:16.830903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport re\nimport glob\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport numpy as np \nimport pandas as pd \nimport nibabel as nib\nimport pytorch_lightning as pl\nimport matplotlib.pyplot as plt\nfrom tqdm import notebook\nfrom typing import Optional\nfrom monai.networks.nets import DenseNet121\nfrom monai.transforms import (   \n    EnsureChannelFirst,\n    Compose,\n    LoadImage, \n    ScaleIntensity\n)\n\n\nROOT_DIR = \"../input/rsna-2022-cervical-spine-fracture-detection\"\nTRAIN_DIR = os.path.join(ROOT_DIR, \"train_images\")\nSEGMENTATION_DIR = os.path.join(ROOT_DIR, \"segmentations\")\nLABEL_NAMES = [\"C1\", \"C2\", \"C3\", \"C4\", \"C5\", \"C6\", \"C7\"]\nnum_class = len(LABEL_NAMES)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-15T11:58:23.613726Z","iopub.execute_input":"2022-09-15T11:58:23.614125Z","iopub.status.idle":"2022-09-15T11:58:30.842406Z","shell.execute_reply.started":"2022-09-15T11:58:23.614091Z","shell.execute_reply":"2022-09-15T11:58:30.841356Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"img_path = os.path.join(TRAIN_DIR, os.listdir(TRAIN_DIR)[0])\nimg_name = glob.glob(os.path.join(img_path, \"*.dcm\"))[0]\nimg_sample = pydicom.dcmread(img_name)\nprint(img_name + \"\\n\")\nprint(img_sample)\ndel img_sample","metadata":{"execution":{"iopub.status.busy":"2022-09-09T10:44:43.912467Z","iopub.execute_input":"2022-09-09T10:44:43.913189Z","iopub.status.idle":"2022-09-09T10:44:43.933948Z","shell.execute_reply.started":"2022-09-09T10:44:43.913149Z","shell.execute_reply":"2022-09-09T10:44:43.932732Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Now we will load vertebrae classes from NIFTI files. We have labels only for subset of our dataset. Vertebrae type is stored in segmentation masks so we extract all unique ids that are in segemntation mask. These ids represent vertebrae type.","metadata":{}},{"cell_type":"code","source":"def read_nifti(file_path): \n    \"\"\"Load segmentation masks from NIFTI file\"\"\"\n    mask = nib.load(file_path)\n    mask = mask.get_fdata()  \n    return mask[:, ::-1, ::-1].transpose(2, 1, 0)\n\n\ndef mask_vertebraes(mask):  \n    \"\"\"Extract vertebrae class id from segmentation mask\"\"\"\n    vertebraes = np.unique(mask)[1:] - 1 # Id == 0 is for background\n    vertebraes = vertebraes.astype(int) \n    vertebraes = vertebraes[vertebraes < num_class]       \n    sample_labels = np.zeros(num_class)\n\n    if len(vertebraes) > 0:            \n        sample_labels[vertebraes] = 1\n\n    return sample_labels.tolist()\n     \n\ndef reed_sample_meta(mask, img_path):\n    \"\"\"Read labels and slice number for DICOM image\"\"\"\n    img_meta = {}\n    img_sample = pydicom.dcmread(img_path)\n    img_meta[\"StudyInstanceUID\"] = img_sample[0x0020, 0x000d].value\n    img_meta[\"slice\"] = img_sample[0x0020, 0x0013].value \n    for cls_name, cls_label in zip(LABEL_NAMES, mask_vertebraes(mask)):\n        img_meta[cls_name] = cls_label        \n    return img_meta\n\n\ndef segmentation_data():\n    train_segmentation = []\n    for patient_id in notebook.tqdm(os.listdir(SEGMENTATION_DIR)):  \n        mask_path = os.path.join(SEGMENTATION_DIR, patient_id)\n        mask = read_nifti(mask_path)   \n       \n        patient_id = patient_id[:-4]\n        dicom_files = os.listdir(os.path.join(TRAIN_DIR, patient_id))\n        dicom_files = sorted(dicom_files, key=lambda x: int(re.sub(\"\\D\", \"\", x))) # Sort patient scans by slice\n        for i, img_name in enumerate(dicom_files): \n            img_path = os.path.join(TRAIN_DIR, patient_id, img_name)\n            sample_meta = reed_sample_meta(mask[i], img_path)\n            train_segmentation.append(sample_meta)\n            \n    return train_segmentation\n    \n              \nif not os.path.isfile(\"train_segmentation.csv\"):           \n    train_segmentation = pd.DataFrame(segmentation_data())\n    train_segmentation.set_index(\"StudyInstanceUID\", inplace=True)\n    train_segmentation.to_csv(\"train_segmentation.csv\")\nelse:\n    train_segmentation = pd.read_csv(\"train_segmentation.csv\", index_col=\"StudyInstanceUID\")","metadata":{"execution":{"iopub.status.busy":"2022-09-09T10:44:44.003772Z","iopub.execute_input":"2022-09-09T10:44:44.004414Z","iopub.status.idle":"2022-09-09T10:44:44.05413Z","shell.execute_reply.started":"2022-09-09T10:44:44.004361Z","shell.execute_reply":"2022-09-09T10:44:44.052512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_segmentation[LABEL_NAMES].sum(axis=1).value_counts().plot(kind='bar')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-09T10:44:44.057207Z","iopub.execute_input":"2022-09-09T10:44:44.058264Z","iopub.status.idle":"2022-09-09T10:44:44.267414Z","shell.execute_reply.started":"2022-09-09T10:44:44.058207Z","shell.execute_reply":"2022-09-09T10:44:44.266213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Most of DICOM images in our training set contain zero or two type of vertebrae but there are also images with four vertebrae type. \n\nWe can use slice relative position as extra input for our model. It is unlikely that vertebrae C2 will be placed in scan with hight level depth.","metadata":{}},{"cell_type":"code","source":"segmentation_meta = train_segmentation.groupby(\"StudyInstanceUID\").slice.max()\ntrain_segmentation[\"slice_ratio\"] =  train_segmentation.slice /  train_segmentation.index.map(segmentation_meta)\ntrain_segmentation[\"img_path\"] = train_segmentation.apply(lambda x: os.path.join(TRAIN_DIR, x.name, str(int(x.slice)) + \".dcm\"), axis=1)\ndel segmentation_meta\ntrain_segmentation.sample(5)","metadata":{"execution":{"iopub.status.busy":"2022-09-09T10:44:44.270346Z","iopub.execute_input":"2022-09-09T10:44:44.270939Z","iopub.status.idle":"2022-09-09T10:44:44.865572Z","shell.execute_reply.started":"2022-09-09T10:44:44.270885Z","shell.execute_reply":"2022-09-09T10:44:44.863937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def display_images(samples, masks=True):    \n    n_rows = 2 if masks else 1\n    n_cols=len(samples)\n    read_dicom = Compose([LoadImage(image_only=True, swap_ij=False), ScaleIntensity()])\n    fig, axs = plt.subplots(nrows=n_rows, ncols=n_cols, figsize=(16, 4 * n_rows))\n    for i, sample in enumerate(samples):\n        patient_id, slc = sample\n        img = read_dicom(os.path.join(TRAIN_DIR, patient_id, str(slc) + \".dcm\"))\n        axs.flat[i].set_title(f\"Slice: {slc}\")\n        axs.flat[i].imshow(img, cmap=plt.cm.bone)\n        \n        if masks:\n            # We must convert img slice to mask index in nii file\n            dicom_files = os.listdir(os.path.join(TRAIN_DIR, patient_id))\n            dicom_files = sorted(dicom_files, key=lambda x: int(re.sub(\"\\D\", \"\", x)))\n            mask_idx = dicom_files.index(str(slc) + \".dcm\")\n            mask = read_nifti(os.path.join(SEGMENTATION_DIR, patient_id) + \".nii\")\n            slice_mask = mask[mask_idx, :, :]            \n            axs.flat[i + n_cols].set_title(f\"Vertebraes: {np.unique(slice_mask)[1:]}\")\n            axs.flat[i + n_cols].imshow(slice_mask, cmap=plt.cm.bone)\n    plt.tight_layout()\n    plt.show()\n     \n        \ntrain_batch = train_segmentation.groupby(\"StudyInstanceUID\").get_group(train_segmentation.index[0])\ntrain_batch = train_batch.sample(5).sort_values(\"slice\")\ntrain_batch = list(zip(train_batch.index, train_batch[\"slice\"]))\ndisplay_images(train_batch)\ndel train_batch","metadata":{"execution":{"iopub.status.busy":"2022-09-09T10:44:44.867178Z","iopub.execute_input":"2022-09-09T10:44:44.867693Z","iopub.status.idle":"2022-09-09T10:44:47.784458Z","shell.execute_reply.started":"2022-09-09T10:44:44.867649Z","shell.execute_reply":"2022-09-09T10:44:47.783147Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style=\"padding: 10px;\n           background-color:#5642C5;\n           color:whitesmoke;\">Dataset</h1>","metadata":{}},{"cell_type":"code","source":"class SpineSegmentationDataset(torch.utils.data.Dataset):\n    def __init__(self, meta_df, segmentation_dir, num_class=7):\n        self.meta_df = meta_df\n        self.segmentation_dir = segmentation_dir\n        self.num_class = num_class\n        self.read_image = Compose([LoadImage(image_only=True, swap_ij=False), \n                                   EnsureChannelFirst(), \n                                   ScaleIntensity()])\n       \n    def __getitem__(self, idx):\n        sample_meta = self.meta_df.loc[idx]        \n        image_path = os.path.join(TRAIN_DIR, sample_meta.StudyInstanceUID, str(sample_meta.slice) + \".dcm\")\n        image = self.read_image(image_path)\n        labels = sample_meta[LABEL_NAMES]\n        slice_ratio = torch.tensor(sample_meta.slice_ratio).float()\n        return image, torch.unsqueeze(slice_ratio, 0), torch.tensor(labels)\n    \n    def __len__(self):\n        return self.meta_df.shape[0]\n    \n    def collate_fn(self, batch):\n        return tuple(map(torch.stack, zip(*batch)))        ","metadata":{"execution":{"iopub.status.busy":"2022-09-09T10:44:47.788142Z","iopub.execute_input":"2022-09-09T10:44:47.788597Z","iopub.status.idle":"2022-09-09T10:44:47.80122Z","shell.execute_reply.started":"2022-09-09T10:44:47.788555Z","shell.execute_reply":"2022-09-09T10:44:47.79903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import DataLoader \n\nclass SpineDataModule(pl.LightningDataModule):\n    def __init__(self, train_full: pd.DataFrame, \n                 data_dir: str, \n                 valid_frac: float= 0.2, \n                 batch_size: int = 8,\n                 num_workers: int = 2):\n        super().__init__()\n        self.data_dir = data_dir\n        self.valid_frac = valid_frac\n        self.batch_size = batch_size\n        self.num_workers = num_workers\n        self.n_samples = train_full.shape[0]\n        self.train_full = train_full\n        \n    def setup(self, stage: Optional[str] = None):   \n        \"\"\"Shuffle and split samples into training and validation partitions\"\"\"\n        if not self.train_full.index.is_numeric():\n            self.train_full = self.train_full.reset_index() \n            \n        valid_size = int(self.valid_frac * self.n_samples)\n        perm = np.random.permutation(self.n_samples)\n\n        train_meta = self.train_full.iloc[perm[valid_size:]]\n        valid_meta = self.train_full.iloc[perm[:valid_size]]\n        self.train_df = train_meta.reset_index(drop=True) \n        self.valid_df = valid_meta.reset_index(drop=True)\n        \n    def train_dataloader(self):\n        train_ds = SpineSegmentationDataset(self.train_df, self.data_dir)\n        return DataLoader(train_ds, \n                          num_workers=self.num_workers,\n                          batch_size=self.batch_size,                                                  \n                          collate_fn=train_ds.collate_fn,\n                          shuffle=True)\n                            \n    def val_dataloader(self):\n        valid_ds = SpineSegmentationDataset(self.valid_df, self.data_dir)\n        return DataLoader(valid_ds, \n                          num_workers=self.num_workers,\n                          batch_size=self.batch_size,\n                          collate_fn=valid_ds.collate_fn,\n                          shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-09-09T10:44:47.803263Z","iopub.execute_input":"2022-09-09T10:44:47.803629Z","iopub.status.idle":"2022-09-09T10:44:47.820715Z","shell.execute_reply.started":"2022-09-09T10:44:47.803597Z","shell.execute_reply":"2022-09-09T10:44:47.819236Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style=\"padding: 10px;\n           background-color:#5642C5;\n           color:whitesmoke;\">Model training</h1>","metadata":{}},{"cell_type":"code","source":"from torchmetrics import Accuracy\n\nclass SpineClassifier(pl.LightningModule):\n    def __init__(self, in_channels: int=1, n_class: int=7, treshold: int=0.5):\n        super().__init__()\n        self.treshold = treshold\n        self.n_class = n_class\n        self.backbone = DenseNet121(spatial_dims=2, \n                                 in_channels=in_channels,\n                                 pretrained=True,\n                                 out_channels=50) \n        \n        self.fc1 = nn.Linear(in_features=1, out_features=50)\n        self.fc2 = nn.Linear(in_features=100, out_features=num_class)      \n        self.sigm = nn.Sigmoid()\n        self.loss_fn = nn.BCELoss()\n        self.metric = Accuracy(average=None, num_classes=n_class)\n        \n    def forward(self, x):     \n        images, ratio = x\n        out1 = self.backbone(images)\n        out2 = self.fc1(ratio)\n        out = torch.cat((out1,out2),1)\n        out = nn.functional.relu(out) # out\n        out = self.fc2(out)\n               \n        out = self.sigm(out)        \n        return out\n    \n    def training_step(self, batch, batch_idx):\n        images, ratio, labels = batch  \n        preds = self.forward((images, ratio))\n        \n        loss = self.loss_fn(preds, labels.float())      \n        self.log(\"train_loss\", loss,  on_step=False, on_epoch=True, prog_bar=True)\n        return loss\n    \n    def validation_step(self, batch, batch_idx):\n        images, ratio, labels = batch \n        preds = self.forward((images, ratio))\n        \n        loss = self.loss_fn(preds, labels.float())        \n        self.metric.update(preds, labels.long())\n        self.log(\"val_loss\", loss, on_step=False, on_epoch=True, prog_bar=True)\n        return loss\n    \n    def predict_step(self, batch, batch_idx): \n        inputs = (batch[0], batch[1])    \n        scores = self.forward(inputs).detach().cpu().numpy()\n        labels = np.array(scores > self.treshold, np.int8) \n        output = {\"preds\": labels, \"scores\": scores}\n        \n        if len(batch) == 3:\n            output[\"labels\"] = batch[2]\n        \n        return output\n    \n    def validation_epoch_end(self, outputs):\n        accuracy = self.metric.compute()\n        \n        for i, CLS_NAME in enumerate(LABEL_NAMES):\n            self.log(f\"val_acc_{CLS_NAME}\", accuracy[i], on_step=False, on_epoch=True)\n                   \n    def configure_optimizers(self):\n        return torch.optim.Adam(self.parameters(), lr=1e-5)   ","metadata":{"execution":{"iopub.status.busy":"2022-09-09T10:44:47.822625Z","iopub.execute_input":"2022-09-09T10:44:47.823048Z","iopub.status.idle":"2022-09-09T10:44:47.844199Z","shell.execute_reply.started":"2022-09-09T10:44:47.823009Z","shell.execute_reply":"2022-09-09T10:44:47.842685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from pytorch_lightning.loggers import CSVLogger\n\ndevice = \"gpu\" if torch.cuda.is_available() else \"cpu\"\nlogger = CSVLogger(\"logs\", name=\"spine_classification\")\ndataset = SpineDataModule(train_segmentation, TRAIN_DIR)\nmodel = SpineClassifier()\n\ntrainer = pl.Trainer(accelerator=device, max_epochs=10, logger=logger)\ntrainer.fit(model, dataset)\ndel train_segmentation","metadata":{"execution":{"iopub.status.busy":"2022-09-09T10:44:47.846119Z","iopub.execute_input":"2022-09-09T10:44:47.846647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"fig, axs = plt.subplots(1, 2, figsize=(16, 4))\n\nmetrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\nmetrics.drop(\"step\", axis=1, inplace=True)\nmetrics = metrics.set_index(\"epoch\").stack()\nmetrics = metrics.unstack()\n\nmetrics[['train_loss', 'val_loss']].plot(ax=axs[0], xticks=metrics.index)\nmetrics[\"mean_accuracy\"] = metrics[[f\"val_acc_{cls_name}\" \n                                    for cls_name in LABEL_NAMES]].mean(axis=1)\nmetrics[\"mean_accuracy\"].plot(ax=axs[1], xticks=metrics.index, legend=True)\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<h1 style=\"padding: 10px;\n           background-color:#5642C5;\n           color:whitesmoke;\">Evaluation</h1>","metadata":{}},{"cell_type":"code","source":"pd.DataFrame(metrics.iloc[-1]).T","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"outputs = trainer.predict(model, dataset.val_dataloader())\n\npreds = []\nscores = []\nlabels = []\n\nfor output in outputs:\n    preds.append(output[\"preds\"])\n    scores.append(output[\"scores\"])\n    labels.append(output[\"labels\"])\n    \npreds = np.concatenate(preds, axis=0)\nscores = np.concatenate(scores, axis=0)\nlabels =  np.concatenate(labels, axis=0)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix\nimport seaborn as sns\n\n\nplt.figure(figsize=(16, 7))\n\nfor i in range(7):\n    plt.subplot(2, 4, i + 1)\n    sns.heatmap(confusion_matrix(preds[:, i], labels[:, i]), annot=True, fmt=\"d\")\n    plt.title(LABEL_NAMES[i])\n    plt.xlabel(\"Preds\")\n    plt.ylabel(\"True\")\n    \nplt.tight_layout()\n\ndel preds, scores, labels, outputs","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_preds(imgs, labels, preds, scores):\n    n_rows = len(imgs) // 4\n    fig, axs = plt.subplots(n_rows, 4, figsize=(18, 5 * n_rows))\n    for i, ax in enumerate(axs.flat):\n        cls = [str(int(c)) for c in labels[i]]\n        prs = [str(int(p)) for p in preds[i]]\n        scr = [str(round(p, 1)) for p in scores[i]]\n\n        title = f\"Labels: {', '.join(cls)} \\nPreds: {', '.join(prs)} \\nScr: {', '.join(scr)}\"\n        ax.set_title(title)\n        ax.imshow(imgs[i], cmap=plt.cm.bone)\n        \n        \nimgs, ratio, labels = next(iter(dataset.val_dataloader()))\nscores = model((imgs, ratio))\nplot_preds(imgs.permute(0, 2, 3, 1), \n           labels, \n           np.array(scores > 0.5, np.int8), \n           scores.detach().numpy())","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model, \"checkpoints\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}