{"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":"## **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 simple solution**:\n\n* For each patient, we consider 4 sequences (FLAIR, T1w, T1Gd, T2), and for each of those sequences take a slice randomly. Idea from [https://github.com/zabir-nabil/Fibro-CoSANet](https://github.com/zabir-nabil/Fibro-CoSANet)\n\n* Construct a 4-channel image out of these 4 sequences.\n\n* Design a 4 channel pytorch model.\n\n* Perform binary classification.\n\n### **Updates** v3:\n\n* Added augmentation. Check out the augmentation notebook: https://www.kaggle.com/furcifer/mri-data-augmentation-pipeline\n* Added few heuristics to avoid black/empty scans.\n* Added full training script with weights saving.\n* Added modified efficient-net (4 channels) as model.\n\n### **Updates** v4:\n\n* Experimenting with first training without any augmentation, and after few epochs adding augmentation.\n* Training longer\n\n## **Check out my other kernels**\n\n### ⚡ **Training kernel:** https://www.kaggle.com/furcifer/torch-efficientnet3d-for-mri-no-train/\n\n### ⚡ **Inference kernel:** https://www.kaggle.com/furcifer/torch-effnet3d-for-mri-no-inference/\n\n","metadata":{}},{"cell_type":"markdown","source":"## **CNN with 4 channels**","metadata":{}},{"cell_type":"code","source":"from IPython.display import Image\nImage(\"../input/diagram/diagram.png\")","metadata":{"execution":{"iopub.status.busy":"2021-07-27T17:02:08.560486Z","iopub.execute_input":"2021-07-27T17:02:08.560911Z","iopub.status.idle":"2021-07-27T17:02:08.596502Z","shell.execute_reply.started":"2021-07-27T17:02:08.560824Z","shell.execute_reply":"2021-07-27T17:02:08.593982Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport glob\nfrom tqdm import tqdm_notebook as tqdm\nimport random\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\nimport imgaug as ia\nimport imgaug.augmenters as iaa\nfrom sklearn.metrics import roc_auc_score\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-07-25T05:21:13.866883Z","iopub.execute_input":"2021-07-25T05:21:13.86727Z","iopub.status.idle":"2021-07-25T05:21:17.452395Z","shell.execute_reply.started":"2021-07-25T05:21:13.867239Z","shell.execute_reply":"2021-07-25T05:21:17.451575Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nsys.path.append('../input/efficientnet-pytorch')\nsys.path.append('../input/efficientnet/EfficientNet-PyTorch-master/')","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:21:17.455654Z","iopub.execute_input":"2021-07-25T05:21:17.455966Z","iopub.status.idle":"2021-07-25T05:21:17.461984Z","shell.execute_reply.started":"2021-07-25T05:21:17.455937Z","shell.execute_reply":"2021-07-25T05:21:17.461026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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))","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:21:17.463624Z","iopub.execute_input":"2021-07-25T05:21:17.464186Z","iopub.status.idle":"2021-07-25T05:21:17.484825Z","shell.execute_reply.started":"2021-07-25T05:21:17.464105Z","shell.execute_reply":"2021-07-25T05:21:17.484011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.head()","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:21:17.486283Z","iopub.execute_input":"2021-07-25T05:21:17.486642Z","iopub.status.idle":"2021-07-25T05:21:17.506734Z","shell.execute_reply.started":"2021-07-25T05:21:17.486605Z","shell.execute_reply":"2021-07-25T05:21:17.505874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Augmentation**","metadata":{}},{"cell_type":"code","source":"sometimes = lambda aug: iaa.Sometimes(0.1, aug)\n\nseq = iaa.Sequential(\n    [\n        # apply the following augmenters to most images\n        iaa.Fliplr(0.5), # horizontally flip 50% of all images\n        iaa.Flipud(0.5), # vertically flip 20% of all images\n        # crop images by -5% to 10% of their height/width\n        sometimes(iaa.CropAndPad(\n            percent=(-0.05, 0.05),\n            pad_mode=ia.ALL,\n            pad_cval=(0, 255)\n        )),\n        sometimes(iaa.Affine(\n            scale={\"x\": (0.8, 1.2), \"y\": (0.8, 1.2)}, # scale images to 80-120% of their size, individually per axis\n            translate_percent={\"x\": (-0.2, 0.2), \"y\": (-0.2, 0.2)}, # translate by -20 to +20 percent (per axis)\n            rotate=(-45, 45), # rotate by -45 to +45 degrees\n            shear=(-16, 16), # shear by -16 to +16 degrees\n            order=[0, 1], # use nearest neighbour or bilinear interpolation (fast)\n            cval=(0, 255), # if mode is constant, use a cval between 0 and 255\n            mode=ia.ALL # use any of scikit-image's warping modes (see 2nd image from the top for examples)\n        )),\n        # execute 0 to 5 of the following (less important) augmenters per image\n        # don't execute all of them, as that would often be way too strong\n        iaa.SomeOf((0, 5),\n            [\n                sometimes(iaa.Superpixels(p_replace=(0, 1.0), n_segments=(20, 200))), # convert images into their superpixel representation\n                iaa.OneOf([\n                    iaa.GaussianBlur((0, 3.0)), # blur images with a sigma between 0 and 3.0\n                    iaa.AverageBlur(k=(2, 7)), # blur image using local means with kernel sizes between 2 and 7\n                    iaa.MedianBlur(k=(3, 11)), # blur image using local medians with kernel sizes between 2 and 7\n                ]),\n                iaa.Sharpen(alpha=(0, 1.0), lightness=(0.75, 1.5)), # sharpen images\n                iaa.Emboss(alpha=(0, 1.0), strength=(0, 2.0)), # emboss images\n                # search either for all edges or for directed edges,\n                # blend the result with the original image using a blobby mask\n                iaa.SimplexNoiseAlpha(iaa.OneOf([\n                    iaa.EdgeDetect(alpha=(0.5, 1.0)),\n                    iaa.DirectedEdgeDetect(alpha=(0.5, 1.0), direction=(0.0, 1.0)),\n                ])),\n                iaa.AdditiveGaussianNoise(loc=0, scale=(0.0, 0.05*255), per_channel=0.5), # add gaussian noise to images\n                iaa.OneOf([\n                    iaa.Dropout((0.01, 0.1), per_channel=0.5), # randomly remove up to 10% of the pixels\n                    iaa.CoarseDropout((0.03, 0.15), size_percent=(0.02, 0.05), per_channel=0.2),\n                ]),\n                iaa.Invert(0.05, per_channel=True), # invert color channels\n                iaa.Add((-10, 10), per_channel=0.5), # change brightness of images (by -10 to 10 of original value)\n                \n                # either change the brightness of the whole image (sometimes\n                # per channel) or change the brightness of subareas\n                iaa.OneOf([\n                    iaa.Multiply((0.5, 1.5), per_channel=0.5),\n                    iaa.FrequencyNoiseAlpha(\n                        exponent=(-4, 0),\n                        first=iaa.Multiply((0.5, 1.5), per_channel=True),\n                        second=iaa.LinearContrast((0.5, 2.0))\n                    )\n                ]),\n                iaa.LinearContrast((0.5, 2.0), per_channel=0.5), # improve or worsen the contrast\n                sometimes(iaa.ElasticTransformation(alpha=(0.5, 3.5), sigma=0.25)), # move pixels locally around (with random strengths)\n                sometimes(iaa.PiecewiseAffine(scale=(0.01, 0.05))), # sometimes move parts of the image around\n                sometimes(iaa.PerspectiveTransform(scale=(0.01, 0.1)))\n            ],\n            random_order=True\n        )\n    ],\n    random_order=True\n)","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:21:20.941107Z","iopub.execute_input":"2021-07-25T05:21:20.941456Z","iopub.status.idle":"2021-07-25T05:21:20.965308Z","shell.execute_reply.started":"2021-07-25T05:21:20.941426Z","shell.execute_reply":"2021-07-25T05:21:20.964317Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dicom2array(paths, voi_lut=True, fix_monochrome=True, remove_black_boundary=True, aug = False):\n    \n    for path in paths:\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        if data.max() > 0.0: # avoiding black images (if possible)\n            break\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    if remove_black_boundary: # we get slightly more details\n        (x, y) = np.where(data > 0)\n        if len(x) > 0 and len(y) > 0:\n            x_mn = np.min(x)\n            x_mx = np.max(x)\n            y_mn = np.min(y)\n            y_mx = np.max(y)\n            if (x_mx - x_mn) > 10 and (y_mx - y_mn) > 10:\n                data = data[:,np.min(y):np.max(y)]\n    data = cv2.resize(data, (512, 512))\n    if aug and random.randint(0,1) == 1: # augmenting only 50% of the time\n        data = seq(images=data)\n    return data\n\ndef load_rand_dicom_images(scan_id, split = \"train\", aug = False):\n    \"\"\"\n    send 4 random slices of each modality\n    \"\"\"\n    if split != \"train\" or split != \"test\":\n        split = \"train\"\n    flair = sorted(glob.glob(f\"{path}/{split}/{scan_id}/FLAIR/*.dcm\"))\n    flair_img = dicom2array(random.sample(flair, max(len(flair)//2, 1)), aug = aug)\n    t1w = sorted(glob.glob(f\"{path}/{split}/{scan_id}/T1w/*.dcm\"))\n    t1w_img = dicom2array(random.sample(t1w, max(len(t1w)//2, 1)), aug = aug)\n    t1wce = sorted(glob.glob(f\"{path}/{split}/{scan_id}/T1wCE/*.dcm\"))\n    t1wce_img = dicom2array(random.sample(t1wce, max(len(t1wce)//2, 1)), aug = aug)\n    t2w = sorted(glob.glob(f\"{path}/{split}/{scan_id}/T2w/*.dcm\"))\n    t2w_img = dicom2array(random.sample(t2w, max(len(t2w)//2, 1)), aug = aug)\n    \n    return np.array((flair_img, t1w_img, t1wce_img, t2w_img)).T","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:22:53.030396Z","iopub.execute_input":"2021-07-25T05:22:53.030713Z","iopub.status.idle":"2021-07-25T05:22:53.044836Z","shell.execute_reply.started":"2021-07-25T05:22:53.030679Z","shell.execute_reply":"2021-07-25T05:22:53.044017Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"load_rand_dicom_images(\"00000\", aug = True).shape","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:22:53.475067Z","iopub.execute_input":"2021-07-25T05:22:53.475424Z","iopub.status.idle":"2021-07-25T05:22:53.890286Z","shell.execute_reply.started":"2021-07-25T05:22:53.475393Z","shell.execute_reply":"2021-07-25T05:22:53.889525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_imgs(imgs, cols=4, size=7, is_rgb=True, title=\"\", cmap='gray', img_size=(512,512)):\n    rows = len(imgs)//cols + 1\n    fig = plt.figure(figsize=(cols*size, rows*size))\n    for i in range(4):\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()","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:22:55.451077Z","iopub.execute_input":"2021-07-25T05:22:55.451403Z","iopub.status.idle":"2021-07-25T05:22:55.45749Z","shell.execute_reply.started":"2021-07-25T05:22:55.451377Z","shell.execute_reply":"2021-07-25T05:22:55.456476Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"slices = load_rand_dicom_images(\"00000\")\nplot_imgs(slices)\naug_slices = load_rand_dicom_images(\"00000\", aug = True)\nplot_imgs(aug_slices)","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:22:58.047726Z","iopub.execute_input":"2021-07-25T05:22:58.048063Z","iopub.status.idle":"2021-07-25T05:23:00.001149Z","shell.execute_reply.started":"2021-07-25T05:22:58.048034Z","shell.execute_reply":"2021-07-25T05:23:00.000309Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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.2):\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\"/train/\" + \"/*\"))]\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_rand_dicom_images(self.ids[idx], self.split, aug = True)\n        \n        transform = transforms.Compose([transforms.ToTensor()]) # transforms.Normalize((0.5, 0.5, 0.5, 0.5), (0.5, 0.5, 0.5, 0.5))\n        imgs = transform(imgs)\n        \n        imgs = imgs - imgs.min()\n        imgs = (imgs + 1e-5) / (imgs.max() - imgs.min() + 1e-5)\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-25T05:31:35.836069Z","iopub.execute_input":"2021-07-25T05:31:35.836412Z","iopub.status.idle":"2021-07-25T05:31:35.852279Z","shell.execute_reply.started":"2021-07-25T05:31:35.836384Z","shell.execute_reply":"2021-07-25T05:31:35.851407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# testing the dataloader\n# testing the dataloader\ntrain_bs = 10\nval_bs = 2\n\ntrain_dataset = BrainTumor()\ntrain_loader = DataLoader(train_dataset, batch_size=train_bs, shuffle=True, num_workers = 8)\n\nval_dataset = BrainTumor(split=\"valid\")\nval_loader = DataLoader(val_dataset, batch_size=val_bs, shuffle=True, num_workers = 8)","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:31:36.701386Z","iopub.execute_input":"2021-07-25T05:31:36.701717Z","iopub.status.idle":"2021-07-25T05:31:36.730266Z","shell.execute_reply.started":"2021-07-25T05:31:36.701688Z","shell.execute_reply":"2021-07-25T05:31:36.729504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for img, label in train_loader:\n    print(img.shape)\n    print(img.min())\n    print(img.mean())\n    print(img.max())\n    print(label.shape)\n    break\n\nfor img, label in val_loader:\n    print(img.shape)\n    print(label.shape)\n    break","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:31:37.570418Z","iopub.execute_input":"2021-07-25T05:31:37.570734Z","iopub.status.idle":"2021-07-25T05:31:54.947656Z","shell.execute_reply.started":"2021-07-25T05:31:37.570706Z","shell.execute_reply":"2021-07-25T05:31:54.946725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **CNN (dummy) + EfficientNet**","metadata":{}},{"cell_type":"code","source":"# let's write our simplest cnn\n\nclass SimpleCNN(nn.Module): \n    def __init__(self):\n        super(SimpleCNN, self).__init__()\n        self.conv1 = nn.Conv2d(in_channels=4, out_channels=64, kernel_size=9) # 4 incoming channels\n        self.conv2 = nn.Conv2d(64, 32, kernel_size=7)\n        self.conv2_drop = nn.Dropout2d()\n        self.conv3 = nn.Conv2d(32, 16, kernel_size=5)\n        self.conv3_drop = nn.Dropout2d()\n        self.conv4 = nn.Conv2d(16, 8, kernel_size=3)\n        self.conv4_drop = nn.Dropout2d()\n        self.fc1 = nn.Linear(200, 256)\n        self.fc2 = nn.Linear(256, 2)\n\n    def forward(self, x):\n        x = F.relu(F.max_pool2d(self.conv1(x), 4))\n        x = F.relu(F.max_pool2d(self.conv2_drop(self.conv2(x)), 4))\n        x = F.relu(F.max_pool2d(self.conv3_drop(self.conv3(x)), 2))\n        x = F.relu(F.max_pool2d(self.conv4_drop(self.conv4(x)), 2))\n        x = x.view(x.shape[0],-1)\n        x = F.relu(self.fc1(x))\n        x = F.dropout(x, training=self.training)\n        x = self.fc2(x)\n        return x","metadata":{"execution":{"iopub.status.busy":"2021-07-25T00:36:00.372997Z","iopub.execute_input":"2021-07-25T00:36:00.373296Z","iopub.status.idle":"2021-07-25T00:36:00.382657Z","shell.execute_reply.started":"2021-07-25T00:36:00.373267Z","shell.execute_reply":"2021-07-25T00:36:00.381785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from efficientnet_pytorch import EfficientNet\nfrom efficientnet_pytorch.utils import Conv2dStaticSamePadding\n\nPATH = \"../input/efficientnet-pytorch/efficientnet-b1-dbc7070a.pth\"\nmodel = EfficientNet.from_name('efficientnet-b1')\nmodel.load_state_dict(torch.load(PATH))\n\n# augment model with 4 channels\n\nmodel._conv_stem = Conv2dStaticSamePadding(4, 32, kernel_size = (3,3), stride = (2,2), \n                                                             bias = False, image_size = 512)\n\nmodel._fc = torch.nn.Linear(in_features=1280, out_features=2, bias=True)\n\n# https://github.com/zabir-nabil/Fibro-CoSANet    # will update later        ","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:28:53.564787Z","iopub.execute_input":"2021-07-25T05:28:53.565164Z","iopub.status.idle":"2021-07-25T05:28:54.935918Z","shell.execute_reply.started":"2021-07-25T05:28:53.56513Z","shell.execute_reply":"2021-07-25T05:28:54.934991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"criterion = nn.CrossEntropyLoss()\noptimizer = torch.optim.AdamW(model.parameters(),lr = 0.001, weight_decay=0.02)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=3, factor=0.5)\nn_epochs = 2","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:29:49.717656Z","iopub.execute_input":"2021-07-25T05:29:49.71801Z","iopub.status.idle":"2021-07-25T05:29:49.727725Z","shell.execute_reply.started":"2021-07-25T05:29:49.717978Z","shell.execute_reply":"2021-07-25T05:29:49.726898Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(model)\nmodel(torch.randn(1, 4, 512, 512))","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:29:51.021515Z","iopub.execute_input":"2021-07-25T05:29:51.021846Z","iopub.status.idle":"2021-07-25T05:29:52.114788Z","shell.execute_reply.started":"2021-07-25T05:29:51.021817Z","shell.execute_reply":"2021-07-25T05:29:52.113774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Training**","metadata":{}},{"cell_type":"code","source":"# helper\ndef one_hot(arr):\n    return [[1, 0] if a_i == 0 else [0, 1] for a_i in arr]","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:30:31.30424Z","iopub.execute_input":"2021-07-25T05:30:31.304571Z","iopub.status.idle":"2021-07-25T05:30:31.308583Z","shell.execute_reply.started":"2021-07-25T05:30:31.304543Z","shell.execute_reply":"2021-07-25T05:30:31.307642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# let's train\ngpu = torch.device(f\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\nmodel.to(gpu)\n\ntrain_loss = []\nval_loss = []\ntrain_roc = []\nval_roc = []\nbest_roc = 0.0\n\nfor epoch in range(n_epochs):  # loop over the dataset multiple times\n    y_all = []\n    outputs_all = []\n    running_loss = 0.0\n    roc = 0.0\n    \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        running_loss += loss.item()\n        y_all.extend(y.tolist())\n        outputs_all.extend(outputs.tolist())\n    \n    roc += roc_auc_score(one_hot(y_all), outputs_all) / train_bs\n    print(f\"epoch {epoch+1} train loss: {running_loss} train roc: {roc}\")\n    \n    train_loss.append(running_loss)\n    train_roc.append(roc)\n\n    y_all = []\n    outputs_all = []\n    running_loss = 0.0\n    roc = 0.0 \n    \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 += loss.item()\n        y_all.extend(y.tolist())\n        outputs_all.extend(outputs.tolist())\n    \n    roc += roc_auc_score(one_hot(y_all), outputs_all) / val_bs\n    scheduler.step(running_loss)\n        \n    print(f\"epoch {epoch+1} val loss: {running_loss} val roc: {roc}\")\n    \n    val_loss.append(running_loss)\n    val_roc.append(roc)\n    \n    if roc > best_roc:\n        best_roc = roc\n        torch.save(model.state_dict(), f'best_roc_{round(roc, 2)}_loss_{round(running_loss, 2)}.pt')","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:33:28.732713Z","iopub.execute_input":"2021-07-25T05:33:28.733066Z","iopub.status.idle":"2021-07-25T05:35:37.220051Z","shell.execute_reply.started":"2021-07-25T05:33:28.733033Z","shell.execute_reply":"2021-07-25T05:35:37.217331Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.plot(train_loss, label = 'train loss')\nplt.plot(val_loss, label = 'val loss')\nplt.xlabel('epochs')\nplt.ylabel('loss')\nplt.legend(['train loss', 'val loss'])\nplt.show()\n\nplt.plot(train_roc, label = 'train roc')\nplt.plot(val_roc, label = 'val roc')\nplt.xlabel('epochs')\nplt.ylabel('roc auc')\nplt.legend(['train roc', 'val roc'])\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-07-25T05:35:42.436625Z","iopub.execute_input":"2021-07-25T05:35:42.437005Z","iopub.status.idle":"2021-07-25T05:35:42.738233Z","shell.execute_reply.started":"2021-07-25T05:35:42.436958Z","shell.execute_reply":"2021-07-25T05:35:42.737392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv(\"../input/rsna-miccai-brain-tumor-radiogenomic-classification/sample_submission.csv\")\nsubmission.to_csv(\"submission.csv\", index=False)\nsubmission","metadata":{"execution":{"iopub.status.busy":"2021-07-14T15:34:28.176834Z","iopub.execute_input":"2021-07-14T15:34:28.177359Z","iopub.status.idle":"2021-07-14T15:34:28.229467Z","shell.execute_reply.started":"2021-07-14T15:34:28.177317Z","shell.execute_reply":"2021-07-14T15:34:28.228385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}