{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":29653,"databundleVersionId":2420395,"sourceType":"competition"},{"sourceId":76358302,"sourceType":"kernelVersion"},{"sourceId":44065,"sourceType":"modelInstanceVersion","modelInstanceId":37004}],"dockerImageVersionId":30132,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **Import Libraries**","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nimport torch.nn.functional as F\nimport sklearn\nimport torchvision\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.preprocessing import LabelEncoder\nimport numpy as np\nimport pandas as pd\nimport os\nimport matplotlib.pyplot as plt\nimport PIL\nfrom PIL import Image\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport seaborn as sns\nimport glob\nfrom pathlib import Path\nimport cv2\ntorch.manual_seed(1)\nnp.random.seed(1)\nimport re\nimport pydicom\nimport math\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nfrom sklearn.metrics import precision_score, recall_score\nfrom sklearn.metrics import roc_auc_score","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:12.464303Z","iopub.execute_input":"2024-05-10T20:41:12.464535Z","iopub.status.idle":"2024-05-10T20:41:19.123371Z","shell.execute_reply.started":"2024-05-10T20:41:12.464509Z","shell.execute_reply":"2024-05-10T20:41:19.122555Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"IMAGE_SIZE = 256\nNUM_IMAGES = 64\nBATCH_SIZE= 4","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.124901Z","iopub.execute_input":"2024-05-10T20:41:19.125143Z","iopub.status.idle":"2024-05-10T20:41:19.128942Z","shell.execute_reply.started":"2024-05-10T20:41:19.125115Z","shell.execute_reply":"2024-05-10T20:41:19.128119Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Data Loading and Visualizations**","metadata":{}},{"cell_type":"code","source":"# Loade each individual MRI image\ndef loading_image(path, img_size=IMAGE_SIZE):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    data = apply_voi_lut(dicom.pixel_array, dicom)\n    data = cv2.resize(data, (img_size, img_size))\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.130214Z","iopub.execute_input":"2024-05-10T20:41:19.130487Z","iopub.status.idle":"2024-05-10T20:41:19.139430Z","shell.execute_reply.started":"2024-05-10T20:41:19.130454Z","shell.execute_reply":"2024-05-10T20:41:19.138762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Loads in the Data as 3d slides for feeding into the model\ndef load_3d_image(idx, mri_type, num_imgs=NUM_IMAGES, split='train'):\n    files = sorted(glob.glob(f\"../input/rsna-miccai-brain-tumor-radiogenomic-classification/{split}/{idx}/{mri_type}/*.dcm\"), \n                   key=lambda var:[int(x) if x.isdigit() else x for x in re.findall(r'[^0-9]|[0-9]+', var)])\n    middle = int(len(files) / 2)\n    half_num_imgs = int(num_imgs / 2)\n    start = max(0, middle - half_num_imgs)\n    end = min(len(files) + 1, middle + half_num_imgs)\n    arrays = [loading_image(f) for f in files[start:end]]\n\n    img3d = np.stack(arrays, axis=2)\n    \n    if img3d.shape[-1] < num_imgs:\n        n_zero = np.zeros((IMAGE_SIZE, IMAGE_SIZE, num_imgs - img3d.shape[-1]))\n        img3d = np.concatenate((img3d,  n_zero), axis=-1)\n        \n    if np.min(img3d) < np.max(img3d):\n        img3d = img3d - np.min(img3d)\n        img3d = img3d / np.max(img3d)\n\n    return img3d","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.141323Z","iopub.execute_input":"2024-05-10T20:41:19.141586Z","iopub.status.idle":"2024-05-10T20:41:19.157769Z","shell.execute_reply.started":"2024-05-10T20:41:19.141553Z","shell.execute_reply":"2024-05-10T20:41:19.156831Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels = pd.read_csv('../input/rsna-miccai-brain-tumor-radiogenomic-classification/train_labels.csv')","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.158789Z","iopub.execute_input":"2024-05-10T20:41:19.159092Z","iopub.status.idle":"2024-05-10T20:41:19.179761Z","shell.execute_reply.started":"2024-05-10T20:41:19.159055Z","shell.execute_reply":"2024-05-10T20:41:19.179086Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.hist(train_labels['MGMT_value'], bins=2, edgecolor='black')\n\n# Setting the title and labels\nplt.title(\"MGMT_Value\")\nplt.xlabel(\"Value\")\nplt.ylabel(\"Count\")\n\n# Remove the gaps between the bars\nplt.gca().set_xticks([0, 1])  # Setting x-axis ticks to only 0 and 1\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.180826Z","iopub.execute_input":"2024-05-10T20:41:19.181107Z","iopub.status.idle":"2024-05-10T20:41:19.392855Z","shell.execute_reply.started":"2024-05-10T20:41:19.181080Z","shell.execute_reply":"2024-05-10T20:41:19.392165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.394026Z","iopub.execute_input":"2024-05-10T20:41:19.394574Z","iopub.status.idle":"2024-05-10T20:41:19.410172Z","shell.execute_reply.started":"2024-05-10T20:41:19.394534Z","shell.execute_reply":"2024-05-10T20:41:19.409464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"(train_labels.MGMT_value == 1).sum()","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.411112Z","iopub.execute_input":"2024-05-10T20:41:19.411326Z","iopub.status.idle":"2024-05-10T20:41:19.427652Z","shell.execute_reply.started":"2024-05-10T20:41:19.411301Z","shell.execute_reply":"2024-05-10T20:41:19.426947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data = train_labels['MGMT_value']\n\nd = np.diff(np.unique(data)).min()\nleft_of_first_bin = data.min() - float(d)/2\nright_of_last_bin = data.max() + float(d)/2\nplt.hist(data, np.arange(left_of_first_bin, right_of_last_bin + d, d))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.428556Z","iopub.execute_input":"2024-05-10T20:41:19.428762Z","iopub.status.idle":"2024-05-10T20:41:19.585822Z","shell.execute_reply.started":"2024-05-10T20:41:19.428737Z","shell.execute_reply":"2024-05-10T20:41:19.585075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_files = sorted(os.listdir('../input/rsna-miccai-brain-tumor-radiogenomic-classification/train'))","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.590085Z","iopub.execute_input":"2024-05-10T20:41:19.590673Z","iopub.status.idle":"2024-05-10T20:41:19.652137Z","shell.execute_reply.started":"2024-05-10T20:41:19.590621Z","shell.execute_reply":"2024-05-10T20:41:19.651475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_files)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.653107Z","iopub.execute_input":"2024-05-10T20:41:19.653329Z","iopub.status.idle":"2024-05-10T20:41:19.658142Z","shell.execute_reply.started":"2024-05-10T20:41:19.653302Z","shell.execute_reply":"2024-05-10T20:41:19.657463Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_files = pd.Series(train_files, name='train_files')","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.659059Z","iopub.execute_input":"2024-05-10T20:41:19.659309Z","iopub.status.idle":"2024-05-10T20:41:19.667069Z","shell.execute_reply.started":"2024-05-10T20:41:19.659274Z","shell.execute_reply":"2024-05-10T20:41:19.666420Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels = pd.concat([train_labels, train_files], axis=1)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.667956Z","iopub.execute_input":"2024-05-10T20:41:19.668191Z","iopub.status.idle":"2024-05-10T20:41:19.677376Z","shell.execute_reply.started":"2024-05-10T20:41:19.668163Z","shell.execute_reply":"2024-05-10T20:41:19.676645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.678237Z","iopub.execute_input":"2024-05-10T20:41:19.678450Z","iopub.status.idle":"2024-05-10T20:41:19.694360Z","shell.execute_reply.started":"2024-05-10T20:41:19.678425Z","shell.execute_reply":"2024-05-10T20:41:19.693533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels = train_labels[train_labels['BraTS21ID'] != 109]","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.695288Z","iopub.execute_input":"2024-05-10T20:41:19.695467Z","iopub.status.idle":"2024-05-10T20:41:19.704541Z","shell.execute_reply.started":"2024-05-10T20:41:19.695446Z","shell.execute_reply":"2024-05-10T20:41:19.703766Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_labels = train_labels[train_labels['BraTS21ID'] != 709]","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.705626Z","iopub.execute_input":"2024-05-10T20:41:19.705911Z","iopub.status.idle":"2024-05-10T20:41:19.714065Z","shell.execute_reply.started":"2024-05-10T20:41:19.705850Z","shell.execute_reply":"2024-05-10T20:41:19.713421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"109 and 709 don't have flair images so for this dataset. ","metadata":{}},{"cell_type":"code","source":"train_labels['MGMT_value'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.714994Z","iopub.execute_input":"2024-05-10T20:41:19.715289Z","iopub.status.idle":"2024-05-10T20:41:19.729848Z","shell.execute_reply.started":"2024-05-10T20:41:19.715256Z","shell.execute_reply":"2024-05-10T20:41:19.729059Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Fairly balanced train set.","metadata":{}},{"cell_type":"code","source":"test_data = pd.read_csv('../input/rsna-miccai-brain-tumor-radiogenomic-classification/sample_submission.csv')\ntest_ids = []\nfor f in test_data.itertuples():\n    test_ids.append(f[1])","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.731078Z","iopub.execute_input":"2024-05-10T20:41:19.731371Z","iopub.status.idle":"2024-05-10T20:41:19.742889Z","shell.execute_reply.started":"2024-05-10T20:41:19.731334Z","shell.execute_reply":"2024-05-10T20:41:19.742240Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Dataset and DataLoader**","metadata":{}},{"cell_type":"code","source":"# Formatting and arranging data for model\nclass TumorDataset(torch.utils.data.Dataset):\n    def __init__(self, df=train_labels, transform=transforms.Compose([transforms.ToTensor()]), mri_type=\"FLAIR\", train=True):\n        self.df = df\n        self.transform = transform\n        self.type = mri_type\n        self.train = train\n    \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n            if self.train == True:\n                patient_id = self.df.iloc[idx, 2]\n                \n                image = load_3d_image(str(patient_id), self.type)\n                image = self.transform(image)\n                image = image[None, :, :, :]\n                label = self.df.iloc[idx, 1]\n                label = torch.tensor(label)\n                \n                return image, label\n            \n            else:\n                patient_id = self.df[idx]\n                patient_id = str(patient_id)\n                for i in range(5 - len(patient_id)):\n                    patient_id = '0' + patient_id\n                \n                \n                image = load_3d_image(patient_id, self.type, split='test')\n                image = self.transform(image)\n                image = image[None, :, :, :]\n                \n                return image, idx","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.745603Z","iopub.execute_input":"2024-05-10T20:41:19.747531Z","iopub.status.idle":"2024-05-10T20:41:19.757566Z","shell.execute_reply.started":"2024-05-10T20:41:19.747501Z","shell.execute_reply":"2024-05-10T20:41:19.756823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_dataset = TumorDataset()\ntest_dataset = TumorDataset(df=test_ids, train=False)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.758519Z","iopub.execute_input":"2024-05-10T20:41:19.758728Z","iopub.status.idle":"2024-05-10T20:41:19.765768Z","shell.execute_reply.started":"2024-05-10T20:41:19.758704Z","shell.execute_reply":"2024-05-10T20:41:19.765145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=BATCH_SIZE, num_workers=4)\ntest_loader = torch.utils.data.DataLoader(test_dataset, batch_size=BATCH_SIZE, num_workers=4)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.768515Z","iopub.execute_input":"2024-05-10T20:41:19.768720Z","iopub.status.idle":"2024-05-10T20:41:19.774952Z","shell.execute_reply.started":"2024-05-10T20:41:19.768697Z","shell.execute_reply":"2024-05-10T20:41:19.774252Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\ndevice","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.775883Z","iopub.execute_input":"2024-05-10T20:41:19.776083Z","iopub.status.idle":"2024-05-10T20:41:19.847174Z","shell.execute_reply.started":"2024-05-10T20:41:19.776059Z","shell.execute_reply":"2024-05-10T20:41:19.846222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Simple Model Architecture**","metadata":{}},{"cell_type":"markdown","source":"# 3D CNN","metadata":{}},{"cell_type":"code","source":"class ThreeDNetwork(nn.Module):\n    \n    def conv_layer(self, in_channels, out_channels, kernel_size, stride=2):\n        conv_layer = nn.Sequential(\n            nn.Conv3d(in_channels, out_channels, kernel_size, stride=stride),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((2, 2, 2)),\n            nn.BatchNorm3d(out_channels))\n        return conv_layer\n    \n    def __init__(self, batch_size=BATCH_SIZE):\n        super(ThreeDNetwork, self).__init__()\n        self.batch_size = batch_size\n        self.block1 = nn.Sequential(\n            self.conv_layer(1, 64, 3, 2),\n            self.conv_layer(64, 128, 3, 2))\n        \n        self.fc = nn.Sequential(\n            nn.Linear(86400, 1024),\n            nn.LeakyReLU(),\n            nn.BatchNorm1d(1024),\n            nn.Dropout(0.2),\n            nn.Linear(1024, 1))\n        \n    def forward(self, x):\n        x = self.block1(x)\n        x = x.view(-1, 86400)\n        x = self.fc(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2024-05-06T23:32:55.469972Z","iopub.execute_input":"2024-05-06T23:32:55.470587Z","iopub.status.idle":"2024-05-06T23:32:55.479648Z","shell.execute_reply.started":"2024-05-06T23:32:55.470547Z","shell.execute_reply":"2024-05-06T23:32:55.478862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = ThreeDNetwork()","metadata":{"execution":{"iopub.status.busy":"2024-05-06T00:41:58.613590Z","iopub.execute_input":"2024-05-06T00:41:58.613800Z","iopub.status.idle":"2024-05-06T00:41:59.302920Z","shell.execute_reply.started":"2024-05-06T00:41:58.613776Z","shell.execute_reply":"2024-05-06T00:41:59.302262Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(model)","metadata":{"execution":{"iopub.status.busy":"2024-05-06T00:41:59.306491Z","iopub.execute_input":"2024-05-06T00:41:59.307153Z","iopub.status.idle":"2024-05-06T00:41:59.312382Z","shell.execute_reply.started":"2024-05-06T00:41:59.307120Z","shell.execute_reply":"2024-05-06T00:41:59.311510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### **Training**","metadata":{}},{"cell_type":"code","source":"optimizer = torch.optim.Adam(model.parameters(), lr=0.001)\n\ntrain_criterion = nn.BCELoss()\nlr_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.1, patience=4, cooldown=2, verbose=True)\n\nmodel = model.to(device)\ntrain_criterion = train_criterion.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-05-06T00:41:59.313444Z","iopub.execute_input":"2024-05-06T00:41:59.313668Z","iopub.status.idle":"2024-05-06T00:41:59.420504Z","shell.execute_reply.started":"2024-05-06T00:41:59.313642Z","shell.execute_reply":"2024-05-06T00:41:59.419776Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 25\n\naccumulated_losses = []\naccuracy_history = []  # To store accuracy for each epoch\nrecorded_best_loss = float('inf')\n\nfor current_epoch in range(num_epochs): \n    print(f'Epoch: {current_epoch + 1}')\n    batch_losses = []\n    epoch_loss_avg = 0\n    correct_predictions = 0\n    total_predictions = 0\n    \n    for imgs, labels in train_loader:\n        optimizer.zero_grad()\n        \n        labels_transformed = labels.view(-1, 1).to(torch.float).to(device)\n        imgs = imgs.float().to(device)\n        \n        predictions = model(imgs)\n        predictions = torch.sigmoid(predictions)\n        loss = train_criterion(predictions, labels_transformed)\n        \n        loss.backward()\n        optimizer.step()\n        \n        batch_losses.append(loss.item())\n\n        # For accuracy calculation, assuming binary classification and rounding predictions\n        predicted_classes = predictions.round().to(device)\n        correct_predictions += (predicted_classes == labels_transformed).sum().item()\n        total_predictions += labels_transformed.size(0)\n            \n    epoch_loss_avg = np.mean(batch_losses)\n    epoch_accuracy = (correct_predictions / total_predictions) * 100\n    accuracy_history.append(epoch_accuracy)  # Append accuracy of this epoch to the history\n    \n    print(f'Epoch {current_epoch + 1}, Training Loss: {epoch_loss_avg:.4f}, Training Accuracy: {epoch_accuracy:.2f}%')\n    \n    if epoch_loss_avg < recorded_best_loss:\n        torch.save(model.state_dict(), 'updated_model.pth')\n        print('Model performance improved. Saving the new model.')\n        recorded_best_loss = epoch_loss_avg\n        \n    lr_scheduler.step(epoch_loss_avg)\n    accumulated_losses.append(epoch_loss_avg)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-05-06T00:42:03.620742Z","iopub.execute_input":"2024-05-06T00:42:03.621065Z","iopub.status.idle":"2024-05-06T00:44:12.373120Z","shell.execute_reply.started":"2024-05-06T00:42:03.621028Z","shell.execute_reply":"2024-05-06T00:44:12.372039Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Testing","metadata":{}},{"cell_type":"code","source":"def rounding(num):\n    return math.floor(num + 0.5)","metadata":{"execution":{"iopub.status.busy":"2024-05-06T22:19:59.457516Z","iopub.execute_input":"2024-05-06T22:19:59.457788Z","iopub.status.idle":"2024-05-06T22:19:59.468728Z","shell.execute_reply.started":"2024-05-06T22:19:59.457751Z","shell.execute_reply":"2024-05-06T22:19:59.468079Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model.load_state_dict(torch.load('updated_model.pth'))","metadata":{"execution":{"iopub.status.busy":"2024-05-06T00:44:12.381378Z","iopub.execute_input":"2024-05-06T00:44:12.381610Z","iopub.status.idle":"2024-05-06T00:44:12.682577Z","shell.execute_reply.started":"2024-05-06T00:44:12.381583Z","shell.execute_reply":"2024-05-06T00:44:12.681741Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"acc_correct = 0\nsamples_total = 0\n\ntrue_labels = []\npredictions = []\n\n# Ensure model is in evaluation mode and gradient computation is off\nmodel.eval()\nwith torch.no_grad():\n    for img_batch, labels in train_loader:\n        updated_labels = labels.unsqueeze(1)  # Add an extra dimension\n        updated_labels = updated_labels.to(torch.int32).to(device)\n        img_batch = img_batch.float().to(device)\n        \n        pred_logits = model(img_batch)\n        pred_probs = torch.sigmoid(pred_logits)\n        \n        # Apply a custom rounding function to predictions\n        rounded_preds = torch.tensor([[rounding(prob)] for prob in pred_probs], dtype=torch.int32).to(device)\n        \n        samples_total += len(img_batch)\n        \n        # Calculate number of correct predictions\n        correct_preds = (rounded_preds == updated_labels).sum().item()\n        acc_correct += correct_preds\n        \n        # Store predictions and labels for precision and recall calculation\n        true_labels.extend(updated_labels.cpu().numpy())\n        predictions.extend(rounded_preds.cpu().numpy())\n\n# Calculating overall accuracy\naccuracy_pct = (acc_correct / samples_total) * 100\n\n# Calculate precision and recall\nprecision = precision_score(true_labels, predictions)\nrecall = recall_score(true_labels, predictions)\n\nprint(f'Train Accuracy: {accuracy_pct:.2f} %')\nprint(f'Precision: {precision:.2f}')\nprint(f'Recall: {recall:.2f}')\n","metadata":{"execution":{"iopub.status.busy":"2024-05-06T00:44:12.684120Z","iopub.execute_input":"2024-05-06T00:44:12.684374Z","iopub.status.idle":"2024-05-06T00:46:16.786833Z","shell.execute_reply.started":"2024-05-06T00:44:12.684345Z","shell.execute_reply":"2024-05-06T00:46:16.785900Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Placeholder for true labels and prediction scores\ntrue_labels = []\npred_scores = []\n\nmodel.eval()\nwith torch.no_grad():\n    for img_batch, labels in train_loader:\n        # Ensure labels are in the correct format\n        labels = labels.to(dtype=torch.float32, device=device)\n        \n        img_batch = img_batch.to(dtype=torch.float32, device=device)\n        \n        # Get the prediction scores for the positive class\n        pred_logits = model(img_batch)\n        pred_probs = torch.sigmoid(pred_logits).squeeze().cpu().numpy()  # Assuming binary classification\n        \n        true_labels.extend(labels.squeeze().cpu().numpy())\n        pred_scores.extend(pred_probs)\n\n# Calculate the ROC AUC score\nroc_auc = roc_auc_score(true_labels, pred_scores)\nprint(f'Train ROC AUC: {roc_auc:.2f}')\n","metadata":{"execution":{"iopub.status.busy":"2024-05-05T23:31:24.968308Z","iopub.execute_input":"2024-05-05T23:31:24.968547Z","iopub.status.idle":"2024-05-05T23:33:12.671687Z","shell.execute_reply.started":"2024-05-05T23:31:24.968517Z","shell.execute_reply":"2024-05-05T23:33:12.670837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(12, 5))\n\n# Plotting training loss\nplt.subplot(1, 2, 1)\nplt.plot(accumulated_losses, label='Training Loss', marker='o')\nplt.title('Training Loss Over Epochs Using CNN on Val Set')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.grid(True)\n\n# Plotting training accuracy\nplt.subplot(1, 2, 2)\nplt.plot(accuracy_history, label='Training Accuracy', marker='o', color='orange')\nplt.title('Training Accuracy Over Epochs Using CNN on Val Set')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy (%)')\nplt.legend()\nplt.grid(True)\n\nplt.tight_layout()\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-05-05T23:33:12.674286Z","iopub.execute_input":"2024-05-05T23:33:12.675026Z","iopub.status.idle":"2024-05-05T23:33:13.233959Z","shell.execute_reply.started":"2024-05-05T23:33:12.674983Z","shell.execute_reply":"2024-05-05T23:33:13.233188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3D CNN + RNN","metadata":{}},{"cell_type":"code","source":"class ThreeDNetwork(nn.Module):\n    def conv_layer(self, in_channels, out_channels, kernel_size, stride=2):\n        conv_layer = nn.Sequential(\n            nn.Conv3d(in_channels, out_channels, kernel_size, stride=stride),\n            nn.LeakyReLU(),\n            nn.MaxPool3d((2, 2, 2)),\n            nn.BatchNorm3d(out_channels))\n        return conv_layer\n    \n    def __init__(self, batch_size=BATCH_SIZE):\n        super(ThreeDNetwork, self).__init__()\n        self.batch_size = batch_size\n        self.block1 = nn.Sequential(\n            self.conv_layer(1, 64, 3, 2),\n            self.conv_layer(64, 128, 3, 2))\n        \n        # Output features for RNN\n        self.feature_extractor = nn.AdaptiveAvgPool3d((1, 1, 1))\n        \n    def forward(self, x):\n        x = self.block1(x)\n        x = self.feature_extractor(x)  # Reduce to a single value per feature map\n        x = x.view(-1, 128)  # Flatten the features for RNN\n        return x","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ThreeDCNNtoRNN(nn.Module):\n    def __init__(self, cnn_features_dim, hidden_dim, num_layers, num_classes):\n        super(ThreeDCNNtoRNN, self).__init__()\n        self.cnn = ThreeDNetwork()  # Your existing CNN\n        self.rnn = nn.LSTM(cnn_features_dim, hidden_dim, num_layers, batch_first=True)\n        self.fc = nn.Linear(hidden_dim, num_classes)\n\n    def forward(self, x):\n        batch_size, C, H, W, D = x.size()\n        c_out = self.cnn(x)\n        c_out = c_out.unsqueeze(1)  # Fake a sequence dimension for LSTM\n        r_out, _ = self.rnn(c_out)\n        r_out = r_out[:, -1, :]  # Get the last time step output\n        output = self.fc(r_out)\n        return output\n\n\n# Instantiate the model\nmodel = ThreeDCNNtoRNN(cnn_features_dim=128, hidden_dim=256, num_layers=1, num_classes=1)  # Update dimensions as needed\nprint(model)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-06T23:33:05.726543Z","iopub.execute_input":"2024-05-06T23:33:05.727105Z","iopub.status.idle":"2024-05-06T23:33:06.441348Z","shell.execute_reply.started":"2024-05-06T23:33:05.727072Z","shell.execute_reply":"2024-05-06T23:33:06.440613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Training","metadata":{}},{"cell_type":"code","source":"optimizer = torch.optim.Adam(model.parameters(), lr=0.001)\n\ntrain_criterion = nn.BCELoss()\nlr_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.1, patience=4, cooldown=2, verbose=True)\n\nmodel = model.to(device)\ntrain_criterion = train_criterion.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-05-06T23:33:10.116858Z","iopub.execute_input":"2024-05-06T23:33:10.117153Z","iopub.status.idle":"2024-05-06T23:33:15.804037Z","shell.execute_reply.started":"2024-05-06T23:33:10.117111Z","shell.execute_reply":"2024-05-06T23:33:15.803369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 25\n\naccumulated_losses = []\naccuracy_history = []  # To store accuracy for each epoch\nrecorded_best_loss = float('inf')\n\nfor current_epoch in range(num_epochs): \n    print(f'Epoch: {current_epoch + 1}')\n    batch_losses = []\n    epoch_loss_avg = 0\n    correct_predictions = 0\n    total_predictions = 0\n    \n    for imgs, labels in train_loader:\n        optimizer.zero_grad()\n        \n        labels_transformed = labels.view(-1, 1).to(torch.float).to(device)\n        imgs = imgs.float().to(device)\n        \n        predictions = model(imgs)\n        predictions = torch.sigmoid(predictions)\n        loss = train_criterion(predictions, labels_transformed)\n        \n        loss.backward()\n        optimizer.step()\n        \n        batch_losses.append(loss.item())\n\n        # For accuracy calculation, assuming binary classification and rounding predictions\n        predicted_classes = predictions.round().to(device)\n        correct_predictions += (predicted_classes == labels_transformed).sum().item()\n        total_predictions += labels_transformed.size(0)\n            \n    epoch_loss_avg = np.mean(batch_losses)\n    epoch_accuracy = (correct_predictions / total_predictions) * 100\n    accuracy_history.append(epoch_accuracy)  # Append accuracy of this epoch to the history\n    \n    print(f'Epoch {current_epoch + 1}, Training Loss: {epoch_loss_avg:.4f}, Training Accuracy: {epoch_accuracy:.2f}%')\n    \n    if epoch_loss_avg < recorded_best_loss:\n        torch.save(model.state_dict(), 'updated_model_rnn.pth')\n        print('Model performance improved. Saving the new model.')\n        recorded_best_loss = epoch_loss_avg\n        \n    lr_scheduler.step(epoch_loss_avg)\n    accumulated_losses.append(epoch_loss_avg)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-05-06T23:33:15.805402Z","iopub.execute_input":"2024-05-06T23:33:15.805631Z","iopub.status.idle":"2024-05-06T23:33:29.702636Z","shell.execute_reply.started":"2024-05-06T23:33:15.805604Z","shell.execute_reply":"2024-05-06T23:33:29.701440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Testing","metadata":{}},{"cell_type":"code","source":"model.load_state_dict(torch.load('updated_model_rnn.pth'))","metadata":{"execution":{"iopub.status.busy":"2024-04-30T15:30:53.015226Z","iopub.status.idle":"2024-04-30T15:30:53.015526Z","shell.execute_reply.started":"2024-04-30T15:30:53.015367Z","shell.execute_reply":"2024-04-30T15:30:53.015381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"acc_correct = 0\nsamples_total = 0\n\ntrue_labels = []\npredictions = []\n\n# Ensure model is in evaluation mode and gradient computation is off\nmodel.eval()\nwith torch.no_grad():\n    for img_batch, labels in train_loader:\n        updated_labels = labels.unsqueeze(1)  # Add an extra dimension\n        updated_labels = updated_labels.to(torch.int32).to(device)\n        img_batch = img_batch.float().to(device)\n        \n        pred_logits = model(img_batch)\n        pred_probs = torch.sigmoid(pred_logits)\n        \n        # Apply a custom rounding function to predictions\n        rounded_preds = torch.tensor([[rounding(prob)] for prob in pred_probs], dtype=torch.int32).to(device)\n        \n        samples_total += len(img_batch)\n        \n        # Calculate number of correct predictions\n        correct_preds = (rounded_preds == updated_labels).sum().item()\n        acc_correct += correct_preds\n        \n        # Store predictions and labels for precision and recall calculation\n        true_labels.extend(updated_labels.cpu().numpy())\n        predictions.extend(rounded_preds.cpu().numpy())\n\n# Calculating overall accuracy\naccuracy_pct = (acc_correct / samples_total) * 100\n\n# Calculate precision and recall\nprecision = precision_score(true_labels, predictions)\nrecall = recall_score(true_labels, predictions)\n\nprint(f'Train Accuracy: {accuracy_pct:.2f} %')\nprint(f'Precision: {precision:.2f}')\nprint(f'Recall: {recall:.2f}')\n","metadata":{"execution":{"iopub.status.busy":"2024-04-30T15:30:53.016919Z","iopub.status.idle":"2024-04-30T15:30:53.017261Z","shell.execute_reply.started":"2024-04-30T15:30:53.017094Z","shell.execute_reply":"2024-04-30T15:30:53.017110Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score\n\n# Placeholder for true labels and prediction scores\ntrue_labels = []\npred_scores = []\n\nmodel.eval()\nwith torch.no_grad():\n    for img_batch, labels in train_loader:\n        # Ensure labels are in the correct format\n        labels = labels.to(dtype=torch.float32, device=device)\n        \n        img_batch = img_batch.to(dtype=torch.float32, device=device)\n        \n        # Get the prediction scores for the positive class\n        pred_logits = model(img_batch)\n        pred_probs = torch.sigmoid(pred_logits).squeeze().cpu().numpy()  # Assuming binary classification\n        \n        true_labels.extend(labels.squeeze().cpu().numpy())\n        pred_scores.extend(pred_probs)\n\n# Calculate the ROC AUC score\nroc_auc = roc_auc_score(true_labels, pred_scores)\nprint(f'Train ROC AUC: {roc_auc:.2f}')\n","metadata":{"execution":{"iopub.status.busy":"2024-04-30T15:30:53.018753Z","iopub.status.idle":"2024-04-30T15:30:53.019229Z","shell.execute_reply.started":"2024-04-30T15:30:53.018974Z","shell.execute_reply":"2024-04-30T15:30:53.019021Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(12, 5))\n\n# Plotting training loss\nplt.subplot(1, 2, 1)\nplt.plot(accumulated_losses, label='Training Loss', marker='o')\nplt.title('Training Loss Over Epochs Using CNN+RNN on Val Set')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.grid(True)\n\n# Plotting training accuracy\nplt.subplot(1, 2, 2)\nplt.plot(accuracy_history, label='Training Accuracy', marker='o', color='orange')\nplt.title('Training Accuracy Over Epochs Using CNN+RNN on Val Set')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy (%)')\nplt.legend()\nplt.grid(True)\n\nplt.tight_layout()\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-05-05T23:33:12.674286Z","iopub.execute_input":"2024-05-05T23:33:12.675026Z","iopub.status.idle":"2024-05-05T23:33:13.233959Z","shell.execute_reply.started":"2024-05-05T23:33:12.674983Z","shell.execute_reply":"2024-05-05T23:33:13.233188Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# EfficientNet","metadata":{}},{"cell_type":"code","source":"pytorch3dpath = \"/kaggle/input/efficientnet/pytorch/en/1/EfficientNet-PyTorch-3D-master\"","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:19.848292Z","iopub.execute_input":"2024-05-10T20:41:19.848578Z","iopub.status.idle":"2024-05-10T20:41:19.856534Z","shell.execute_reply.started":"2024-05-10T20:41:19.848538Z","shell.execute_reply":"2024-05-10T20:41:19.855797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys \nsys.path.append(pytorch3dpath)\nfrom efficientnet_pytorch_3d import EfficientNet3D","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:37.363512Z","iopub.execute_input":"2024-05-10T20:41:37.363791Z","iopub.status.idle":"2024-05-10T20:41:37.384114Z","shell.execute_reply.started":"2024-05-10T20:41:37.363760Z","shell.execute_reply":"2024-05-10T20:41:37.383520Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nfrom torch import nn\n\nclass EfficientNet3DWrapper(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # Assuming EfficientNet3D.from_name is a valid constructor that can create 3D EfficientNet models\n        self.net = EfficientNet3D.from_name(\"efficientnet-b0\", override_params={'num_classes': 2}, in_channels=1)\n        n_features = self.net._fc.in_features\n        self.net._fc = nn.Linear(in_features=n_features, out_features=1, bias=True)\n    \n    def forward(self, x):\n        out = self.net(x)\n        return torch.sigmoid(out)\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:38.876093Z","iopub.execute_input":"2024-05-10T20:41:38.876784Z","iopub.status.idle":"2024-05-10T20:41:38.883010Z","shell.execute_reply.started":"2024-05-10T20:41:38.876746Z","shell.execute_reply":"2024-05-10T20:41:38.882202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Train","metadata":{}},{"cell_type":"code","source":"# Setup device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Initialize the model\nmodel = EfficientNet3DWrapper().to(device)\n\n# Print the model summary (optional, requires additional packages like torchsummary)\n# from torchsummary import summary\n# summary(model, input_size=(1, 64, 64, 64))  # Adjust the size according to your input dimensions\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:41:41.248897Z","iopub.execute_input":"2024-05-10T20:41:41.249599Z","iopub.status.idle":"2024-05-10T20:41:46.931119Z","shell.execute_reply.started":"2024-05-10T20:41:41.249561Z","shell.execute_reply":"2024-05-10T20:41:46.930461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.BCEWithLogitsLoss()  # Appropriate for binary classification with a single output unit\n\noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\n\ntrain_criterion = nn.BCELoss()\nlr_scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, factor=0.1, patience=4, cooldown=2, verbose=True)\n\nmodel = model.to(device)\ntrain_criterion = train_criterion.to(device)","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:42:38.651544Z","iopub.execute_input":"2024-05-10T20:42:38.651808Z","iopub.status.idle":"2024-05-10T20:42:38.669164Z","shell.execute_reply.started":"2024-05-10T20:42:38.651781Z","shell.execute_reply":"2024-05-10T20:42:38.668396Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 25\n\naccumulated_losses = []\naccuracy_history = []  # To store accuracy for each epoch\nrecorded_best_loss = float('inf')\n\nfor current_epoch in range(num_epochs): \n    print(f'Epoch: {current_epoch + 1}')\n    batch_losses = []\n    epoch_loss_avg = 0\n    correct_predictions = 0\n    total_predictions = 0\n    \n    for imgs, labels in train_loader:\n        optimizer.zero_grad()\n        \n        labels_transformed = labels.view(-1, 1).to(torch.float).to(device)\n        imgs = imgs.float().to(device)\n        \n        predictions = model(imgs)\n        predictions = torch.sigmoid(predictions)\n        loss = train_criterion(predictions, labels_transformed)\n        \n        loss.backward()\n        optimizer.step()\n        \n        batch_losses.append(loss.item())\n\n        # For accuracy calculation, assuming binary classification and rounding predictions\n        predicted_classes = predictions.round().to(device)\n        correct_predictions += (predicted_classes == labels_transformed).sum().item()\n        total_predictions += labels_transformed.size(0)\n            \n    epoch_loss_avg = np.mean(batch_losses)\n    epoch_accuracy = (correct_predictions / total_predictions) * 100\n    accuracy_history.append(epoch_accuracy)  # Append accuracy of this epoch to the history\n    \n    print(f'Epoch {current_epoch + 1}, Training Loss: {epoch_loss_avg:.4f}, Training Accuracy: {epoch_accuracy:.2f}%')\n    \n    if epoch_loss_avg < recorded_best_loss:\n        torch.save(model.state_dict(), 'updated_model_en.pth')\n        print('Model performance improved. Saving the new model.')\n        recorded_best_loss = epoch_loss_avg\n        \n    lr_scheduler.step(epoch_loss_avg)\n    accumulated_losses.append(epoch_loss_avg)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-05-10T20:42:46.935492Z","iopub.execute_input":"2024-05-10T20:42:46.936159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Test","metadata":{}},{"cell_type":"code","source":"model.load_state_dict(torch.load('updated_model_en.pth'))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\n\n# Necessary imports for precision and recall calculation\nfrom sklearn.metrics import precision_score, recall_score\n\nacc_correct = 0\nsamples_total = 0\n\ntrue_labels = []\npredictions = []\n\n# Ensure model is in evaluation mode and gradient computation is off\nmodel.eval()\nwith torch.no_grad():\n    for img_batch, labels in train_loader:\n        updated_labels = labels.unsqueeze(1)  # Add an extra dimension\n        updated_labels = updated_labels.to(torch.int32).to(device)\n        img_batch = img_batch.float().to(device)\n        \n        pred_logits = model(img_batch)\n        pred_probs = torch.sigmoid(pred_logits)\n        \n        # Apply a custom rounding function to predictions\n        rounded_preds = torch.tensor([[rounding(prob)] for prob in pred_probs], dtype=torch.int32).to(device)\n        \n        samples_total += len(img_batch)\n        \n        # Calculate number of correct predictions\n        correct_preds = (rounded_preds == updated_labels).sum().item()\n        acc_correct += correct_preds\n        \n        # Store predictions and labels for precision and recall calculation\n        true_labels.extend(updated_labels.cpu().numpy())\n        predictions.extend(rounded_preds.cpu().numpy())\n\n# Calculating overall accuracy\naccuracy_pct = (acc_correct / samples_total) * 100\n\n# Calculate precision and recall\nprecision = precision_score(true_labels, predictions)\nrecall = recall_score(true_labels, predictions)\n\nprint(f'Train Accuracy: {accuracy_pct:.2f} %')\nprint(f'Precision: {precision:.2f}')\nprint(f'Recall: {recall:.2f}')\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.metrics import roc_auc_score\n\n# Placeholder for true labels and prediction scores\ntrue_labels = []\npred_scores = []\n\nmodel.eval()\nwith torch.no_grad():\n    for img_batch, labels in train_loader:\n        # Ensure labels are in the correct format\n        labels = labels.to(dtype=torch.float32, device=device)\n        \n        img_batch = img_batch.to(dtype=torch.float32, device=device)\n        \n        # Get the prediction scores for the positive class\n        pred_logits = model(img_batch)\n        pred_probs = torch.sigmoid(pred_logits).squeeze().cpu().numpy()  # Assuming binary classification\n        \n        true_labels.extend(labels.squeeze().cpu().numpy())\n        pred_scores.extend(pred_probs)\n\n# Calculate the ROC AUC score\nroc_auc = roc_auc_score(true_labels, pred_scores)\nprint(f'Train ROC AUC: {roc_auc:.2f}')\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(12, 5))\n\n# Plotting training loss\nplt.subplot(1, 2, 1)\nplt.plot(accumulated_losses, label='Training Loss', marker='o')\nplt.title('Training Loss Over Epochs Using EfficientNet on Val Set')\nplt.xlabel('Epoch')\nplt.ylabel('Loss')\nplt.legend()\nplt.grid(True)\n\n# Plotting training accuracy\nplt.subplot(1, 2, 2)\nplt.plot(accuracy_history, label='Training Accuracy', marker='o', color='orange')\nplt.title('Training Accuracy Over Epochs Using EfficientNet on Val Set')\nplt.xlabel('Epoch')\nplt.ylabel('Accuracy (%)')\nplt.legend()\nplt.grid(True)\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"References:\nhttps://www.kaggle.com/ammarnassanalhajali/brain-tumor-3d-training\nhttps://www.kaggle.com/code/vexxingbanana/simple-pytorch-cnn","metadata":{}}]}