{"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":"# RSNA-MICCAI Brain Tumor Radiogenomic Classificationn - **An approach with PyTorch EfficientNet 3D**\n\n## **Problem Description**:\n\nThere are structural multi-parametric MRI (mpMRI) scans for different subjects, in DICOM format. The exact mpMRI scans included are:\n\n* Fluid Attenuated Inversion Recovery (FLAIR)\n* T1-weighted pre-contrast (T1w)\n* T1-weighted post-contrast (T1Gd)\n* T2-weighted (T2)\n\n`train_labels.csv` - file contains the target **MGMT_value** for each subject in the training data **(e.g. the presence of MGMT promoter methylation)**.\n\nSo, it's a binary classification problem.\n\n## **A EfficientNet3D solution**:\n\n* For each patient, we consider 4 sequences (FLAIR, T1w, T1Gd, T2), and for each of those sequences we take 50 slices from the middle, and stack them, to get 50 x 4 = 200 slices. We resize the slices in shape (200, 200).\n\n* Construct an efficientnet-3d in pytorch with input shape (200, 200, 200).\n\n* Perform binary classification.\n","metadata":{}},{"cell_type":"markdown","source":"### **Importing libraries**","metadata":{}},{"cell_type":"code","source":"import os\nimport glob\nfrom tqdm import tqdm_notebook as tqdm\nimport random\nimport numpy as np\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.nn.functional as F\nfrom torchvision import transforms, utils\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport cv2\n\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:49.377637Z","iopub.execute_input":"2021-07-21T00:25:49.378055Z","iopub.status.idle":"2021-07-21T00:25:51.077038Z","shell.execute_reply.started":"2021-07-21T00:25:49.377962Z","shell.execute_reply":"2021-07-21T00:25:51.076024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_seed(seed):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n\n\nset_seed(42)","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:51.078659Z","iopub.execute_input":"2021-07-21T00:25:51.079060Z","iopub.status.idle":"2021-07-21T00:25:51.140736Z","shell.execute_reply.started":"2021-07-21T00:25:51.079021Z","shell.execute_reply":"2021-07-21T00:25:51.139806Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Importing EfficientNet-3D**","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('../input/efficientnetpyttorch3d/EfficientNet-PyTorch-3D')\nfrom efficientnet_pytorch_3d import EfficientNet3D","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:51.142581Z","iopub.execute_input":"2021-07-21T00:25:51.143111Z","iopub.status.idle":"2021-07-21T00:25:51.180719Z","shell.execute_reply.started":"2021-07-21T00:25:51.143073Z","shell.execute_reply":"2021-07-21T00:25:51.179802Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Inspecting Labels**","metadata":{}},{"cell_type":"code","source":"path = '/kaggle/input/rsna-miccai-brain-tumor-radiogenomic-classification'\ntrain_data = pd.read_csv(os.path.join(path, 'train_labels.csv'))\nprint('Num of train samples:', len(train_data))\ntrain_data.head()\nimg_size = 256","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:51.182722Z","iopub.execute_input":"2021-07-21T00:25:51.183143Z","iopub.status.idle":"2021-07-21T00:25:51.204225Z","shell.execute_reply.started":"2021-07-21T00:25:51.183101Z","shell.execute_reply":"2021-07-21T00:25:51.203405Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **MRI Slice Loading/Processing**","metadata":{}},{"cell_type":"code","source":"def dicom2array(path, voi_lut=True, fix_monochrome=True):\n    dicom = pydicom.read_file(path)\n    # VOI LUT (if available by DICOM device) is used to\n    # transform raw DICOM data to \"human-friendly\" view\n    if voi_lut:\n        data = apply_voi_lut(dicom.pixel_array, dicom)\n    else:\n        data = dicom.pixel_array\n    # depending on this value, X-ray may look inverted - fix that:\n    if fix_monochrome and dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = np.amax(data) - data\n    data = data - np.min(data)\n    data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    data = cv2.resize(data, (img_size, img_size))\n    return data\n\ndef load_3d_dicom_images(scan_id, split = \"train\"):\n    \"\"\"\n    we will use some heuristics to choose the slices to avoid any numpy zero matrix (if possible)\n    \"\"\"\n    flair = sorted(glob.glob(f\"{path}/{split}/{scan_id}/FLAIR/*.dcm\"))\n    t1w = sorted(glob.glob(f\"{path}/{split}/{scan_id}/T1w/*.dcm\"))\n    t1wce = sorted(glob.glob(f\"{path}/{split}/{scan_id}/T1wCE/*.dcm\"))\n    t2w = sorted(glob.glob(f\"{path}/{split}/{scan_id}/T2w/*.dcm\"))\n    \n    \n    flair_img = np.array([dicom2array(a) for a in flair[len(flair)//2 - 25:len(flair)//2 + 25]]).T\n    \n    if flair_img.shape[-1] < 50:\n        n_zero = 50 - flair_img.shape[-1]\n        flair_img = np.concatenate((flair_img, np.zeros((img_size, img_size, n_zero))), axis = -1)\n    #print(flair_img.shape)\n        \n    \n    \n    t1w_img = np.array([dicom2array(a) for a in t1w[len(t1w)//2 - 25:len(t1w)//2 + 25]]).T\n    if t1w_img.shape[-1] < 50:\n        n_zero = 50 - t1w_img.shape[-1]\n        t1w_img = np.concatenate((t1w_img, np.zeros((img_size, img_size, n_zero))), axis = -1)\n    #print(t1w_img.shape)\n    \n    \n    t1wce_img = np.array([dicom2array(a) for a in t1wce[len(t1wce)//2 - 25:len(t1wce)//2 + 25]]).T\n    if t1wce_img.shape[-1] < 50:\n        n_zero = 50 - t1wce_img.shape[-1]\n        t1wce_img = np.concatenate((t1wce_img, np.zeros((img_size, img_size, n_zero))), axis = -1)\n    #print(t1wce_img.shape)\n    \n    \n    t2w_img = np.array([dicom2array(a) for a in t2w[len(t2w)//2 - 25:len(t2w)//2 + 25]]).T\n    if t2w_img.shape[-1] < 50:\n        n_zero = 50 - t2w_img.shape[-1]\n        t2w_img = np.concatenate((t2w_img, np.zeros((img_size, img_size, n_zero))), axis = -1)\n    #print(t2w_img.shape)\n    \n    return np.concatenate((flair_img, t1w_img, t1wce_img, t2w_img), axis = -1)","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:51.205629Z","iopub.execute_input":"2021-07-21T00:25:51.206038Z","iopub.status.idle":"2021-07-21T00:25:51.224027Z","shell.execute_reply.started":"2021-07-21T00:25:51.205997Z","shell.execute_reply":"2021-07-21T00:25:51.222779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"load_3d_dicom_images(\"00000\").shape","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:51.225965Z","iopub.execute_input":"2021-07-21T00:25:51.226820Z","iopub.status.idle":"2021-07-21T00:25:53.764347Z","shell.execute_reply.started":"2021-07-21T00:25:51.226777Z","shell.execute_reply":"2021-07-21T00:25:53.763469Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"slices = load_3d_dicom_images(\"00000\")\nprint(slices.shape)","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:53.767263Z","iopub.execute_input":"2021-07-21T00:25:53.767516Z","iopub.status.idle":"2021-07-21T00:25:54.648467Z","shell.execute_reply.started":"2021-07-21T00:25:53.767491Z","shell.execute_reply":"2021-07-21T00:25:54.647499Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Visualization**","metadata":{}},{"cell_type":"code","source":"def plot_imgs(imgs, cols=20, size=7, is_rgb=True, title=\"\", cmap='gray', img_size=(64,64)):\n    rows = len(imgs)//cols + 1\n    fig = plt.figure(figsize=(cols*size, rows*size))\n    for i in range(cols):\n        img = imgs[:,:,i]\n        fig.add_subplot(rows, cols, i+1)\n        plt.imshow(img, cmap=cmap)\n    plt.suptitle(title)\n    plt.show()\n    \nplot_imgs(slices)","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:54.652229Z","iopub.execute_input":"2021-07-21T00:25:54.652504Z","iopub.status.idle":"2021-07-21T00:25:58.168914Z","shell.execute_reply.started":"2021-07-21T00:25:54.652476Z","shell.execute_reply":"2021-07-21T00:25:58.168003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# frames = []\n# for i in range(200):\n#     frames.append(np.array(slices[:,:,i], dtype = np.uint8))\n\n# fourcc = cv2.VideoWriter_fourcc('m', 'p', '4', 'v')\n# out = cv2.VideoWriter('/kaggle/working/out_video.mp4', fourcc, 15, (200,200))\n# for i in range(len(frames)):\n#     c_frame =  cv2.cvtColor(frames[i],cv2.COLOR_GRAY2RGB)\n#     out.write(c_frame)\n    \n# out.release()","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:58.171090Z","iopub.execute_input":"2021-07-21T00:25:58.175073Z","iopub.status.idle":"2021-07-21T00:25:58.179601Z","shell.execute_reply.started":"2021-07-21T00:25:58.175028Z","shell.execute_reply":"2021-07-21T00:25:58.178717Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # the video play doesn't work, you can download it to view\n\n# from IPython.display import HTML\n# from base64 import b64encode\n\n# def play(filename):\n#     html = ''\n#     video = open(filename,'rb').read()\n#     src = 'data:video/mp4;base64,' + b64encode(video).decode()\n#     html += '<video width=1000 controls autoplay loop><source src=\"%s\" type=\"video/mp4\"></video>' % src \n#     return HTML(html)\n\n# play('/kaggle/working/out_video.mp4')","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:58.181489Z","iopub.execute_input":"2021-07-21T00:25:58.182689Z","iopub.status.idle":"2021-07-21T00:25:58.192181Z","shell.execute_reply.started":"2021-07-21T00:25:58.182650Z","shell.execute_reply":"2021-07-21T00:25:58.191190Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Data Loader**","metadata":{}},{"cell_type":"code","source":"# let's write a simple pytorch dataloader\n\n\nclass BrainTumor(Dataset):\n    def __init__(self, path = '/kaggle/input/rsna-miccai-brain-tumor-radiogenomic-classification', split = \"train\", validation_split = 0.0):\n        # labels\n        train_data = pd.read_csv(os.path.join(path, 'train_labels.csv'))\n        self.labels = {}\n        brats = list(train_data[\"BraTS21ID\"])\n        mgmt = list(train_data[\"MGMT_value\"])\n        for b, m in zip(brats, mgmt):\n            self.labels[str(b).zfill(5)] = m\n            \n        if split == \"valid\":\n            self.split = split\n            self.ids = [a.split(\"/\")[-1] for a in sorted(glob.glob(path + f\"/{split}/\" + \"/*\"))]\n            self.ids = self.ids[:int(len(self.ids)* validation_split)] # first 20% as validation\n        elif split == \"train\":\n            self.split = split\n            self.ids = [a.split(\"/\")[-1] for a in sorted(glob.glob(path + f\"/{split}/\" + \"/*\"))]\n            self.ids = self.ids[int(len(self.ids)* validation_split):] # last 80% as train\n        else:\n            self.split = split\n            self.ids = [a.split(\"/\")[-1] for a in sorted(glob.glob(path + f\"/{split}/\" + \"/*\"))]\n            \n    \n    def __len__(self):\n        return len(self.ids)\n    \n    def __getitem__(self, idx):\n        imgs = load_3d_dicom_images(self.ids[idx], self.split)\n        transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,) * 200, (0.5,) * 200)])\n        imgs = transform(imgs)\n        \n        if self.split != \"test\":\n            label = self.labels[self.ids[idx]]\n            return torch.tensor(imgs, dtype = torch.float32), torch.tensor(label, dtype = torch.long)\n        else:\n            return torch.tensor(imgs, dtype = torch.float32)","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:58.193880Z","iopub.execute_input":"2021-07-21T00:25:58.194709Z","iopub.status.idle":"2021-07-21T00:25:58.233473Z","shell.execute_reply.started":"2021-07-21T00:25:58.194664Z","shell.execute_reply":"2021-07-21T00:25:58.232618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# testing the dataloader\ntrain_dataset = BrainTumor()\ntrain_loader = DataLoader(train_dataset, batch_size=2, shuffle=True, num_workers=8)\n# val_dataset = BrainTumor(split=\"valid\")\n# val_loader = DataLoader(val_dataset, batch_size=2, shuffle=False, num_workers=8)","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:58.236497Z","iopub.execute_input":"2021-07-21T00:25:58.237482Z","iopub.status.idle":"2021-07-21T00:25:58.296363Z","shell.execute_reply.started":"2021-07-21T00:25:58.237443Z","shell.execute_reply":"2021-07-21T00:25:58.295343Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for img, label in train_loader:\n    print(img.shape)\n    print(label.shape)\n    break\n\n# for img, label in val_loader:\n#     print(img.shape)\n#     print(label.shape)\n#     break","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:25:58.301237Z","iopub.execute_input":"2021-07-21T00:25:58.303541Z","iopub.status.idle":"2021-07-21T00:26:16.275088Z","shell.execute_reply.started":"2021-07-21T00:25:58.303496Z","shell.execute_reply":"2021-07-21T00:26:16.274040Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Model: EfficientNet-3D B0**","metadata":{}},{"cell_type":"code","source":"model = EfficientNet3D.from_name(\"efficientnet-b0\", override_params={'num_classes': 2}, in_channels=1)\ncriterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.Adam(model.parameters(),lr = 0.0001)\nn_epochs = 2","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:26:16.276781Z","iopub.execute_input":"2021-07-21T00:26:16.277177Z","iopub.status.idle":"2021-07-21T00:26:16.354022Z","shell.execute_reply.started":"2021-07-21T00:26:16.277135Z","shell.execute_reply":"2021-07-21T00:26:16.353145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Training**","metadata":{}},{"cell_type":"code","source":"# let's train\ngpu = torch.device(f\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nmodel.to(gpu)\n\nfor epoch in range(n_epochs):  # loop over the dataset multiple times\n\n    train_loss = []\n    best_pres = 10000\n    model.train()\n    for i, data in tqdm(enumerate(train_loader, 0)):\n        x, y = data\n        \n        x = torch.unsqueeze(x, dim = 1)\n        x = x.to(gpu)\n        y = y.to(gpu)\n\n        # zero the parameter gradients\n        optimizer.zero_grad()\n\n        # forward + backward + optimize\n        outputs = model(x)\n        loss = criterion(outputs, y)\n        loss.backward()\n        optimizer.step()\n\n        # print statistics\n        train_loss.append(loss.item())\n    avg_train = sum(train_loss) / len(train_loss)\n    print(f\"epoch {epoch+1} train: {avg_train}\")\n\n#     running_loss = []\n#     best_pres = 10000\n#     model.eval()\n#     for i, data in tqdm(enumerate(val_loader, 0)):\n\n#         x, y = data\n        \n#         x = torch.unsqueeze(x, dim = 1)\n#         x = x.to(gpu)\n#         y = y.to(gpu)\n\n#         # forward\n#         outputs = model(x)\n#         loss = criterion(outputs, y)\n\n#         # print statistics\n#         running_loss.append(loss.item())\n#     avg_pred = sum(running_loss) / len(running_loss)   \n#    print(f\"epoch {epoch+1} val: {avg_pred}\")\n    if avg_train < best_pres:\n        print('save model...')\n        best_pres = avg_train\n        torch.save(model.state_dict(),'best_loss.pt')","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:26:16.355362Z","iopub.execute_input":"2021-07-21T00:26:16.355723Z","iopub.status.idle":"2021-07-21T00:27:17.558745Z","shell.execute_reply.started":"2021-07-21T00:26:16.355687Z","shell.execute_reply":"2021-07-21T00:27:17.556353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**Inference**","metadata":{}},{"cell_type":"code","source":"# model = EfficientNet3D.from_name(\"efficientnet-b0\", override_params={'num_classes': 2}, in_channels=1)\n# model.to(gpu)\n# checkpoint = torch.load(f\"best_loss.pt\")\n# model.load_state_dict(checkpoint)\n# model.eval()","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:27:17.560221Z","iopub.status.idle":"2021-07-21T00:27:17.560880Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class test_BrainTumor(Dataset):\n#     def __init__(self, path = '/kaggle/input/rsna-miccai-brain-tumor-radiogenomic-classification', split = \"test\"):\n#         # labels\n#         train_data = pd.read_csv(os.path.join(path, 'sample_submission.csv'))\n#         self.labels = {}\n#         brats = list(train_data[\"BraTS21ID\"])  \n#         self.split = split\n#         self.ids = [a.split(\"/\")[-1] for a in sorted(glob.glob(path + f\"/{split}/\" + \"/*\"))]   \n#     def __len__(self):\n#         return len(self.ids)\n    \n#     def __getitem__(self, idx):\n#         imgs = load_3d_dicom_images(self.ids[idx], self.split)\n#         transform = transforms.Compose([transforms.ToTensor(), transforms.Normalize((0.5,) * 200, (0.5,) * 200)])\n#         imgs = transform(imgs)\n#         return torch.tensor(imgs, dtype = torch.float32)","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:27:17.562348Z","iopub.status.idle":"2021-07-21T00:27:17.563068Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# test_dataset = test_BrainTumor(split = \"test\")\n# test_loader = DataLoader(test_dataset, batch_size=2, shuffle=False, num_workers=8)","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:27:17.564501Z","iopub.status.idle":"2021-07-21T00:27:17.565303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# submission = pd.read_csv(\"../input/rsna-miccai-brain-tumor-radiogenomic-classification/sample_submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:27:17.566561Z","iopub.status.idle":"2021-07-21T00:27:17.567193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# y_pred = []\n# ids = []\n\n# for e, batch in enumerate(test_loader):\n#     print(f\"{e}/{len(test_loader)}\", end=\"\\r\")\n#     with torch.no_grad():\n#         tmp_pred = np.zeros((batch.shape[0], ))\n#         tmp_res = torch.sigmoid(model(batch.to(gpu))).cpu().numpy().squeeze()\n#         tmp_pred += tmp_res\n#         y_pred.extend(tmp_pred)\n#         ids.extend(batch[\"id\"].numpy().tolist())","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:27:17.568592Z","iopub.status.idle":"2021-07-21T00:27:17.569349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# submission = pd.DataFrame({\"BraTS21ID\": ids, \"MGMT_value\": y_pred})\n# submission.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:27:17.570601Z","iopub.status.idle":"2021-07-21T00:27:17.571348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# submission","metadata":{"execution":{"iopub.status.busy":"2021-07-21T00:27:17.572603Z","iopub.status.idle":"2021-07-21T00:27:17.573242Z"},"trusted":true},"execution_count":null,"outputs":[]}]}