{"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":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport cv2\nimport torch\n\nimport os\nfrom tqdm.notebook import tqdm\nimport SimpleITK as sitk\n\nimport sys\nsys.path.append('../input/monai-v060-deep-learning-in-healthcare-imaging/')\nfrom monai.transforms import (\n    AddChannel,\n    Compose,\n    RandRotate90,\n    Resize,\n    ScaleIntensity,\n    EnsureType,\n    Randomizable,\n    LoadImaged,\n    EnsureTyped,\n    RandRotate,\n    RandZoom,\n    RandDeformGrid,\n    RandAffine,\n    CenterScaleCrop,\n    Transform\n)\n\nimport monai\n\nfrom monai.data import CacheDataset, DataLoader, ImageDataset\nfrom multiprocessing import Pool","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-10-01T08:29:35.887069Z","iopub.execute_input":"2021-10-01T08:29:35.887542Z","iopub.status.idle":"2021-10-01T08:29:35.895524Z","shell.execute_reply.started":"2021-10-01T08:29:35.887495Z","shell.execute_reply":"2021-10-01T08:29:35.894552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1. Config","metadata":{}},{"cell_type":"code","source":"DICOM_IM_FOLDER = '../input/rsna-miccai-brain-tumor-radiogenomic-classification/test/'\nIM_FOLDER = 'BraTS2021_Testing_Data'\nMRI_TYPES = ['T1wCE', 'T1w', 'T2w', 'FLAIR']\nSHORT_MRI_TYPES = [ 't1ce', 't1', 't2', 'flair']\n\nSEED = 67\nDIM = (240, 240, 155, 1)\nNUM_CLASSES = 1\nNUM_SEG_CLASSES = 0 # whether to use the segment head\nBATCH_SIZE = 6\nROI_SCALE = [0.6, 0.7, 0.7] \nDEVICE = torch.device('cuda:0')\n\nFAST_COMMIT = False\n\nCANDIDATES = [\n\n    {\n        'backbone_name':'densenet121',\n        'model_path':'../input/brain-densenet/v15.7/v15.7/t2_Fold0_densenet121_v15.7_ValidLoss0.576_ValidAUC0.758_Ep49.pth',\n        'mri_type':'t2',\n    },\n]","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:29:35.897533Z","iopub.execute_input":"2021-10-01T08:29:35.898212Z","iopub.status.idle":"2021-10-01T08:29:35.907465Z","shell.execute_reply.started":"2021-10-01T08:29:35.898175Z","shell.execute_reply":"2021-10-01T08:29:35.906504Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_3_planes_sitk(image):\n    voxels = sitk.GetArrayFromImage(image)\n    plt.figure(figsize=(9,3))\n    plt.subplot(1,3,1)\n    plt.imshow(voxels[voxels.shape[0]//2])\n    plt.subplot(1,3,2)\n    plt.imshow(voxels[:, voxels.shape[1]//2, :])\n    plt.subplot(1,3,3)\n    plt.imshow(voxels[:,:,voxels.shape[2]//2])\n    \ndef visualize_3_planes(voxels):\n    plt.figure(figsize=(9,3))\n    plt.subplot(1,3,1)\n    plt.imshow(voxels[voxels.shape[0]//2])\n    plt.subplot(1,3,2)\n    plt.imshow(voxels[:, voxels.shape[1]//2, :])\n    plt.subplot(1,3,3)\n    plt.imshow(voxels[:,:,voxels.shape[2]//2])","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:29:35.909937Z","iopub.execute_input":"2021-10-01T08:29:35.910589Z","iopub.status.idle":"2021-10-01T08:29:35.921677Z","shell.execute_reply.started":"2021-10-01T08:29:35.910526Z","shell.execute_reply":"2021-10-01T08:29:35.920362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# template image\nreader = sitk.ImageFileReader()\nreader.SetImageIO(\"NiftiImageIO\")\nreader.SetFileName('../input/sri24template/atlastImage.nii')\nsri24 = reader.Execute()","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:29:35.923658Z","iopub.execute_input":"2021-10-01T08:29:35.924049Z","iopub.status.idle":"2021-10-01T08:29:35.962349Z","shell.execute_reply.started":"2021-10-01T08:29:35.924012Z","shell.execute_reply":"2021-10-01T08:29:35.961544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualize_3_planes_sitk(sri24)","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:29:35.964292Z","iopub.execute_input":"2021-10-01T08:29:35.964966Z","iopub.status.idle":"2021-10-01T08:29:36.466048Z","shell.execute_reply.started":"2021-10-01T08:29:35.964929Z","shell.execute_reply":"2021-10-01T08:29:36.465143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def register(fixed_image, moving_image):\n    sitk.ProcessObject_SetGlobalDefaultNumberOfThreads(1)\n    fixed_image = sitk.Cast(fixed_image, moving_image.GetPixelID())\n    initial_transform = sitk.CenteredTransformInitializer(fixed_image, \n                                                      moving_image, \n                                                      sitk.Euler3DTransform(), \n                                                      sitk.CenteredTransformInitializerFilter.GEOMETRY)\n    moving_resampled = sitk.Resample(moving_image, fixed_image, initial_transform, sitk.sitkLinear, 0.0, moving_image.GetPixelID())\n    \n    # interact(display_images_with_alpha, image_z=(0,fixed_image.GetSize()[2]-1), alpha=(0.0,1.0,0.05), fixed = fixed(fixed_image), moving=fixed(moving_resampled));\n    registration_method = sitk.ImageRegistrationMethod()\n\n    # Similarity metric settings.\n    registration_method.SetMetricAsMattesMutualInformation(numberOfHistogramBins=50)\n    registration_method.SetMetricSamplingStrategy(registration_method.RANDOM)\n    registration_method.SetMetricSamplingPercentage(0.01, seed=67)\n    registration_method.SetGlobalDefaultNumberOfThreads(1)\n\n    registration_method.SetInterpolator(sitk.sitkLinear)\n    \n    # Optimizer settings.\n    registration_method.SetOptimizerAsGradientDescent(learningRate=1.0, numberOfIterations=100, convergenceMinimumValue=1e-6, convergenceWindowSize=10)\n    registration_method.SetOptimizerScalesFromPhysicalShift()\n\n    # Setup for the multi-resolution framework.            \n    registration_method.SetShrinkFactorsPerLevel(shrinkFactors = [4,2,1])\n    registration_method.SetSmoothingSigmasPerLevel(smoothingSigmas=[2,1,0])\n    registration_method.SmoothingSigmasAreSpecifiedInPhysicalUnitsOn()\n    \n    # Don't optimize in-place, we would possibly like to run this cell multiple times.\n    registration_method.SetInitialTransform(initial_transform, inPlace=False)\n\n    final_transform = registration_method.Execute(sitk.Cast(fixed_image, sitk.sitkFloat32), \n                                                   sitk.Cast(moving_image, sitk.sitkFloat32))\n    \n    return final_transform","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:29:36.467503Z","iopub.execute_input":"2021-10-01T08:29:36.468028Z","iopub.status.idle":"2021-10-01T08:29:36.477663Z","shell.execute_reply.started":"2021-10-01T08:29:36.467989Z","shell.execute_reply":"2021-10-01T08:29:36.47676Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2. SRI24 Registration","metadata":{}},{"cell_type":"code","source":"reader = sitk.ImageSeriesReader()\nreader.LoadPrivateTagsOn()\nwriter = sitk.ImageFileWriter()","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:29:36.48081Z","iopub.execute_input":"2021-10-01T08:29:36.481121Z","iopub.status.idle":"2021-10-01T08:29:36.492259Z","shell.execute_reply.started":"2021-10-01T08:29:36.481087Z","shell.execute_reply":"2021-10-01T08:29:36.491358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"patient_ids = []\nimage_names = []\nmri_types = []\nmetas = []\n\nmri_type_mapping = {\n    'T1w':'t1',\n    'T1wCE':'t1ce',\n    'T2w':'t2',\n    'FLAIR':'flair'\n}\n\ndef error(e):\n    print(e)\n    \ndef update(args):\n    pbar.update()  \n\ndef process_one_patient(patient_id): \n    reader = sitk.ImageSeriesReader()\n    reader.LoadPrivateTagsOn()\n    writer = sitk.ImageFileWriter()\n    \n    patient_dir = os.path.join(DICOM_IM_FOLDER, patient_id) \n    saved_transform = None\n    \n    for mri_type in MRI_TYPES:\n        type_dir = os.path.join(patient_dir, mri_type)\n        try: \n            filenamesDICOM = reader.GetGDCMSeriesFileNames(type_dir)\n            reader.SetFileNames(filenamesDICOM)\n            voxels = reader.Execute()\n\n            moving_image = voxels\n            fixed_image = sri24\n            \n            if(mri_type == 'T1wCE'):\n                saved_transform = register(fixed_image, moving_image)\n            \n            if(saved_transform is None and mri_type != 'T1wCE'):\n                raise ValueError('T1wCE must be registered to SRI24 first')\n\n            registered_voxels = sitk.Resample(moving_image, fixed_image, saved_transform, sitk.sitkLinear, 0.0, moving_image.GetPixelID())\n            \n            outputImageFileName = os.path.join(IM_FOLDER, f'BraTS2021_{patient_id}', \n                                               f'BraTS2021_{patient_id}_{mri_type_mapping[mri_type]}.nii.gz')\n            os.makedirs(os.path.dirname(outputImageFileName), exist_ok=True)\n            writer.SetFileName(outputImageFileName)\n            writer.Execute(registered_voxels)\n            \n        except Exception as ex:\n            print(ex)\n        \n    return ''","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:29:36.495363Z","iopub.execute_input":"2021-10-01T08:29:36.495728Z","iopub.status.idle":"2021-10-01T08:29:36.507452Z","shell.execute_reply.started":"2021-10-01T08:29:36.495694Z","shell.execute_reply":"2021-10-01T08:29:36.506633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if(FAST_COMMIT and len(os.listdir(DICOM_IM_FOLDER)) == 87):\n    iterations = ['00114','00013', '00821']\nelse:\n    iterations = os.listdir(DICOM_IM_FOLDER)\n    \npool = Pool(processes=4)    \npbar = tqdm(total=len(iterations))\n\nfor patient_id in iterations:\n    pool.apply_async(\n        process_one_patient,\n        args=(patient_id, ),\n        callback=update,\n        error_callback=error,\n    )\n    \npool.close()\npool.join()\npbar.close()","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:29:36.509599Z","iopub.execute_input":"2021-10-01T08:29:36.510016Z","iopub.status.idle":"2021-10-01T08:30:36.214404Z","shell.execute_reply.started":"2021-10-01T08:29:36.509979Z","shell.execute_reply":"2021-10-01T08:30:36.213538Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"path = f'{IM_FOLDER}/BraTS2021_00013/BraTS2021_00013_t1ce.nii.gz'\nreader = sitk.ImageFileReader()\nreader.SetImageIO(\"NiftiImageIO\")\nreader.SetFileName(path)\ndemo = reader.Execute()","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:36.21724Z","iopub.execute_input":"2021-10-01T08:30:36.217519Z","iopub.status.idle":"2021-10-01T08:30:36.29031Z","shell.execute_reply.started":"2021-10-01T08:30:36.217488Z","shell.execute_reply":"2021-10-01T08:30:36.289508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"visualize_3_planes_sitk(demo)","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:36.291598Z","iopub.execute_input":"2021-10-01T08:30:36.291969Z","iopub.status.idle":"2021-10-01T08:30:36.603379Z","shell.execute_reply.started":"2021-10-01T08:30:36.291931Z","shell.execute_reply":"2021-10-01T08:30:36.602402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3. Modeling","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.autograd import Variable\nimport math\nfrom functools import partial\n\n__all__ = [\n    'ResNet', 'resnet10', 'resnet18', 'resnet34', 'resnet50', 'resnet101',\n    'resnet152', 'resnet200'\n]\n\n\ndef conv3x3x3(in_planes, out_planes, stride=1, dilation=1):\n    # 3x3x3 convolution with padding\n    return nn.Conv3d(\n        in_planes,\n        out_planes,\n        kernel_size=3,\n        dilation=dilation,\n        stride=stride,\n        padding=dilation,\n        bias=False)\n\n\ndef downsample_basic_block(x, planes, stride, no_cuda=False):\n    out = F.avg_pool3d(x, kernel_size=1, stride=stride)\n    zero_pads = torch.Tensor(\n        out.size(0), planes - out.size(1), out.size(2), out.size(3),\n        out.size(4)).zero_()\n    if not no_cuda:\n        if isinstance(out.data, torch.cuda.FloatTensor):\n            zero_pads = zero_pads.cuda()\n\n    out = Variable(torch.cat([out.data, zero_pads], dim=1))\n\n    return out\n\n\nclass BasicBlock(nn.Module):\n    expansion = 1\n\n    def __init__(self, inplanes, planes, stride=1, dilation=1, downsample=None):\n        super(BasicBlock, self).__init__()\n        self.conv1 = conv3x3x3(inplanes, planes, stride=stride, dilation=dilation)\n        self.bn1 = nn.BatchNorm3d(planes)\n        self.relu = nn.ReLU(inplace=True)\n        self.conv2 = conv3x3x3(planes, planes, dilation=dilation)\n        self.bn2 = nn.BatchNorm3d(planes)\n        self.downsample = downsample\n        self.stride = stride\n        self.dilation = dilation\n\n    def forward(self, x):\n        residual = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n        out = self.conv2(out)\n        out = self.bn2(out)\n\n        if self.downsample is not None:\n            residual = self.downsample(x)\n\n        out += residual\n        out = self.relu(out)\n\n        return out\n\n\nclass Bottleneck(nn.Module):\n    expansion = 4\n\n    def __init__(self, inplanes, planes, stride=1, dilation=1, downsample=None):\n        super(Bottleneck, self).__init__()\n        self.conv1 = nn.Conv3d(inplanes, planes, kernel_size=1, bias=False)\n        self.bn1 = nn.BatchNorm3d(planes)\n        self.conv2 = nn.Conv3d(\n            planes, planes, kernel_size=3, stride=stride, dilation=dilation, padding=dilation, bias=False)\n        self.bn2 = nn.BatchNorm3d(planes)\n        self.conv3 = nn.Conv3d(planes, planes * 4, kernel_size=1, bias=False)\n        self.bn3 = nn.BatchNorm3d(planes * 4)\n        self.relu = nn.ReLU(inplace=True)\n        self.downsample = downsample\n        self.stride = stride\n        self.dilation = dilation\n\n    def forward(self, x):\n        residual = x\n\n        out = self.conv1(x)\n        out = self.bn1(out)\n        out = self.relu(out)\n\n        out = self.conv2(out)\n        out = self.bn2(out)\n        out = self.relu(out)\n\n        out = self.conv3(out)\n        out = self.bn3(out)\n\n        if self.downsample is not None:\n            residual = self.downsample(x)\n\n        out += residual\n        out = self.relu(out)\n\n        return out\n\nclass GAP3D(nn.Module):\n    def __init__(self, feat_dim):\n        super(GAP3D, self).__init__()\n        self.feat_dim = feat_dim\n\n    def forward(self, x):\n        x = F.adaptive_avg_pool3d(x, (1, 1, 1))\n        x = x.view((-1, self.feat_dim))\n        return x\n\nclass ResNet(nn.Module):\n\n    def __init__(self,\n                 block,\n                 layers,\n                 sample_input_D,\n                 sample_input_H,\n                 sample_input_W,\n                 num_classes,\n                 num_seg_classes,\n                 shortcut_type='B',\n                 no_cuda = False):\n        self.inplanes = 64\n        self.no_cuda = no_cuda\n        self.num_seg_classes = num_seg_classes\n        self.num_classes = num_classes\n\n        super(ResNet, self).__init__()\n        self.conv1 = nn.Conv3d(\n            1,\n            64,\n            kernel_size=7,\n            stride=(2, 2, 2),\n            padding=(3, 3, 3),\n            bias=False)\n            \n        self.bn1 = nn.BatchNorm3d(64)\n        self.relu = nn.ReLU(inplace=True)\n        self.maxpool = nn.MaxPool3d(kernel_size=(3, 3, 3), stride=2, padding=1)\n        self.layer1 = self._make_layer(block, 64, layers[0], shortcut_type)\n        self.layer2 = self._make_layer(\n            block, 128, layers[1], shortcut_type, stride=2)\n        self.layer3 = self._make_layer(\n            block, 256, layers[2], shortcut_type, stride=1, dilation=2)\n        self.layer4 = self._make_layer(\n            block, 512, layers[3], shortcut_type, stride=1, dilation=4)\n\n        # classification head\n        self.feat_dim = 512 * block.expansion\n        self.clf_head = nn.Sequential(\n            GAP3D(self.feat_dim),\n            nn.Linear(self.feat_dim, self.num_classes)\n        )\n\n        if(num_seg_classes > 0):\n            self.conv_seg = nn.Sequential(\n                                            nn.ConvTranspose3d(\n                                            512 * block.expansion,\n                                            32,\n                                            2,\n                                            stride=2\n                                            ),\n                                            nn.BatchNorm3d(32),\n                                            nn.ReLU(inplace=True),\n                                            nn.Conv3d(\n                                            32,\n                                            32,\n                                            kernel_size=3,\n                                            stride=(1, 1, 1),\n                                            padding=(1, 1, 1),\n                                            bias=False), \n                                            nn.BatchNorm3d(32),\n                                            nn.ReLU(inplace=True),\n                                            nn.Conv3d(\n                                            32,\n                                            num_seg_classes,\n                                            kernel_size=1,\n                                            stride=(1, 1, 1),\n                                            bias=False) \n                                            )\n\n            for m in self.modules():\n                if isinstance(m, nn.Conv3d):\n                    m.weight = nn.init.kaiming_normal(m.weight, mode='fan_out')\n                elif isinstance(m, nn.BatchNorm3d):\n                    m.weight.data.fill_(1)\n                    m.bias.data.zero_()\n\n    def _make_layer(self, block, planes, blocks, shortcut_type, stride=1, dilation=1):\n        downsample = None\n        # print(planes, stride, self.inplanes, block.expansion)\n        if stride != 1 or self.inplanes != planes * block.expansion:\n            if shortcut_type == 'A':\n                downsample = partial(\n                    downsample_basic_block,\n                    planes=planes * block.expansion,\n                    stride=stride,\n                    no_cuda=self.no_cuda)\n            else:\n                downsample = nn.Sequential(\n                    nn.Conv3d(\n                        self.inplanes,\n                        planes * block.expansion,\n                        kernel_size=1,\n                        stride=stride,\n                        bias=False), \n                    nn.BatchNorm3d(planes * block.expansion))\n\n        layers = []\n        layers.append(block(self.inplanes, planes, stride=stride, dilation=dilation, downsample=downsample))\n        # print(downsample)\n        self.inplanes = planes * block.expansion\n        for i in range(1, blocks):\n            layers.append(block(self.inplanes, planes, dilation=dilation))\n\n        return nn.Sequential(*layers)\n\n    def forward(self, x):\n        x = self.conv1(x)\n        x = self.bn1(x)\n        x = self.relu(x)\n        x = self.maxpool(x)\n        x = self.layer1(x)\n        x = self.layer2(x)\n        x = self.layer3(x)\n        x = self.layer4(x)\n\n        logits = self.clf_head(x)\n\n#         if(self.num_seg_classes > 0):\n#             seg_mask = self.conv_seg(x)\n#             return logits, seg_mask\n        \n        return logits\n\ndef resnet10(**kwargs):\n    \"\"\"Constructs a ResNet-18 model.\n    \"\"\"\n    model = ResNet(BasicBlock, [1, 1, 1, 1], **kwargs)\n    return model\n\n\ndef resnet18(**kwargs):\n    \"\"\"Constructs a ResNet-18 model.\n    \"\"\"\n    model = ResNet(BasicBlock, [2, 2, 2, 2], **kwargs)\n    return model\n\n\ndef resnet34(**kwargs):\n    \"\"\"Constructs a ResNet-34 model.\n    \"\"\"\n    model = ResNet(BasicBlock, [3, 4, 6, 3], **kwargs)\n    return model\n\n\ndef resnet50(**kwargs):\n    \"\"\"Constructs a ResNet-50 model.\n    \"\"\"\n    model = ResNet(Bottleneck, [3, 4, 6, 3], **kwargs)\n    return model\n\n\ndef resnet101(**kwargs):\n    \"\"\"Constructs a ResNet-101 model.\n    \"\"\"\n    model = ResNet(Bottleneck, [3, 4, 23, 3], **kwargs)\n    return model\n\n\ndef resnet152(**kwargs):\n    \"\"\"Constructs a ResNet-101 model.\n    \"\"\"\n    model = ResNet(Bottleneck, [3, 8, 36, 3], **kwargs)\n    return model\n\n\ndef resnet200(**kwargs):\n    \"\"\"Constructs a ResNet-101 model.\n    \"\"\"\n    model = ResNet(Bottleneck, [3, 24, 36, 3], **kwargs)\n    return model\n\nimport torch\nfrom torch import nn\n\ndef get_medicalnet_resnet_model(model_name, inp_w, inp_h, inp_d, short_cut_type='B', num_classes=1, num_seg_classes=1, backbone_pretrained=None):\n    model_func = globals()[model_name]\n    model = model_func(\n                sample_input_W=inp_w,\n                sample_input_H=inp_h,\n                sample_input_D=inp_d,\n                shortcut_type=short_cut_type,\n                no_cuda=False,\n                num_classes = num_classes,\n                num_seg_classes=num_seg_classes)\n    \n    if(backbone_pretrained is not None):\n        print('Load pretrained:', backbone_pretrained)\n        net_dict = model.state_dict()\n        pretrain = torch.load(backbone_pretrained, map_location='cpu')\n        pretrain_dict = {k.replace('module.', ''): v for k, v in pretrain['state_dict'].items() if k.replace('module.', '') in net_dict.keys()}\n        net_dict.update(pretrain_dict)\n        model.load_state_dict(net_dict)\n\n    return model\n\ndef get_model(candidate):\n    dim = candidate.get('dim', DIM)\n    if('resnet' in candidate['backbone_name']):\n        model = get_medicalnet_resnet_model(candidate['backbone_name'], dim[1], dim[0], dim[2], num_classes=NUM_CLASSES,\n                                                num_seg_classes=NUM_SEG_CLASSES, backbone_pretrained=candidate.get('backbone_pretrained'))\n    elif('efficientnet' in candidate['backbone_name']):\n        model = monai.networks.nets.efficientnet.EfficientNetBN(model_name=candidate['backbone_name'],spatial_dims=3, in_channels=1,\n                                                pretrained=False, num_classes=NUM_CLASSES)\n    elif('densenet121' in candidate['backbone_name']):\n        model = monai.networks.nets.DenseNet121(spatial_dims=3, in_channels=1,\n                                                pretrained=False, out_channels=NUM_CLASSES)\n    else:\n        raise ValueError('No such backbone name: '+ candidate['backbone_name'])\n    return model\n\ndef predict_fn(dataloader,model,scaler, device='cuda:0'):\n    model.eval()\n  \n    tk0 = tqdm(enumerate(dataloader), total=len(dataloader))\n    all_predictions = []\n    for i, batch in tk0:\n        # input, gt\n        voxels = batch\n        voxels = voxels.to(device)\n\n        # prediction\n        with torch.cuda.amp.autocast(), torch.no_grad():\n            logits = model(voxels)\n            logits = logits.view(-1)\n            \n            if(torch.isnan(logits.sum())):\n                print(logits)\n                logits[torch.isnan(logits)] = 0\n            \n            probs = logits.sigmoid()\n     \n        # append for metric calculation\n        all_predictions.append(probs.detach().cpu().numpy())\n        \n        del batch, voxels, logits, probs\n        torch.cuda.empty_cache()\n\n    all_predictions = np.concatenate(all_predictions)\n    return all_predictions\n\nclass InstensityOneVolumeNormalization(Transform):\n    \"\"\"\n   Std scaling normalization\n\n    \"\"\"\n    def __init__(self, div_value=255):\n        super(InstensityOneVolumeNormalization, self).__init__()\n\n\n    def __call__(self, volume):\n        \"\"\"\n        normalize the itensity of an nd volume based on the mean and std of nonzeor region\n        inputs:\n            volume: the input nd volume\n        outputs:\n            out: the normalized nd volume\n        \"\"\"\n        # volume = self.__drop_invalid_range__(volume)\n        \n        pixels = volume[volume > 0]\n        if(volume.min() == 0 and volume.max() == 0):\n            print('1 image all zeros')\n            mean = 0\n            std = 1\n        else:\n            mean = pixels.mean()\n            std  = pixels.std()\n            \n        out = (volume - mean)/std\n\n        return out","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:36.604895Z","iopub.execute_input":"2021-10-01T08:30:36.605275Z","iopub.status.idle":"2021-10-01T08:30:36.6623Z","shell.execute_reply.started":"2021-10-01T08:30:36.605238Z","shell.execute_reply":"2021-10-01T08:30:36.661222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.DataFrame(os.listdir(IM_FOLDER), columns=['pfolder'])\ntest_df['BraTS21ID'] = test_df['pfolder'].map(lambda x: x.split('_')[-1])\n\nfor t in SHORT_MRI_TYPES:\n    test_df[f'{t}_data_path'] = test_df.pfolder.map(lambda x: os.path.join(IM_FOLDER, x, x+f'_{t}.nii.gz'))","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:36.663683Z","iopub.execute_input":"2021-10-01T08:30:36.664074Z","iopub.status.idle":"2021-10-01T08:30:36.679082Z","shell.execute_reply.started":"2021-10-01T08:30:36.664032Z","shell.execute_reply":"2021-10-01T08:30:36.678271Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df.head()","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:36.680369Z","iopub.execute_input":"2021-10-01T08:30:36.680753Z","iopub.status.idle":"2021-10-01T08:30:36.698357Z","shell.execute_reply.started":"2021-10-01T08:30:36.680714Z","shell.execute_reply":"2021-10-01T08:30:36.697373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_transforms = Compose([AddChannel(), \n                           CenterScaleCrop(roi_scale=ROI_SCALE),\n                           InstensityOneVolumeNormalization()])\nmri_type = SHORT_MRI_TYPES[0]\n\ntest_dataset = ImageDataset(image_files=test_df[f'{mri_type}_data_path'].tolist(),\n                            transform=test_transforms)","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:36.70064Z","iopub.execute_input":"2021-10-01T08:30:36.701015Z","iopub.status.idle":"2021-10-01T08:30:36.707277Z","shell.execute_reply.started":"2021-10-01T08:30:36.700976Z","shell.execute_reply":"2021-10-01T08:30:36.706464Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# voxels, labels = next(iter(train_loader))\nvoxels  = test_dataset[0]\nvisualize_3_planes(voxels[0])","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:36.708454Z","iopub.execute_input":"2021-10-01T08:30:36.708991Z","iopub.status.idle":"2021-10-01T08:30:37.206751Z","shell.execute_reply.started":"2021-10-01T08:30:36.708953Z","shell.execute_reply":"2021-10-01T08:30:37.205936Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_ensembled_prediction = 0\n    \nfor candidate in CANDIDATES:\n    \n    # create data loader\n    mri_type = candidate.get('mri_type')\n    test_dataset = ImageDataset(image_files=test_df[f'{mri_type}_data_path'].tolist(),\n                            transform=test_transforms)\n\n    batch_size = candidate.get('batch_size', BATCH_SIZE)\n    test_loader = torch.utils.data.DataLoader(test_dataset, batch_size=batch_size, shuffle=False,\n                    num_workers=4, pin_memory=torch.cuda.is_available())\n\n    # Model\n    model = get_model(candidate)\n    print('Load trained model:', candidate['model_path'] )\n    model.load_state_dict(torch.load(candidate['model_path'], map_location='cpu'))\n    model = model.to(DEVICE)\n    print()\n\n    # use amp to accelerate training\n    scaler = torch.cuda.amp.GradScaler()\n\n    test_ensembled_prediction += predict_fn(test_loader, model, scaler, device=DEVICE)\n\ntest_ensembled_prediction /= len(CANDIDATES)","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:37.208076Z","iopub.execute_input":"2021-10-01T08:30:37.208407Z","iopub.status.idle":"2021-10-01T08:30:44.78383Z","shell.execute_reply.started":"2021-10-01T08:30:37.208372Z","shell.execute_reply":"2021-10-01T08:30:44.782922Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df['MGMT_value'] = test_ensembled_prediction","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:44.785266Z","iopub.execute_input":"2021-10-01T08:30:44.785611Z","iopub.status.idle":"2021-10-01T08:30:44.790541Z","shell.execute_reply.started":"2021-10-01T08:30:44.78558Z","shell.execute_reply":"2021-10-01T08:30:44.789592Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df[['BraTS21ID', 'MGMT_value']].head()","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:44.792111Z","iopub.execute_input":"2021-10-01T08:30:44.792757Z","iopub.status.idle":"2021-10-01T08:30:44.810912Z","shell.execute_reply.started":"2021-10-01T08:30:44.792703Z","shell.execute_reply":"2021-10-01T08:30:44.809937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df[['BraTS21ID', 'MGMT_value']].to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:44.812307Z","iopub.execute_input":"2021-10-01T08:30:44.812704Z","iopub.status.idle":"2021-10-01T08:30:44.823523Z","shell.execute_reply.started":"2021-10-01T08:30:44.812666Z","shell.execute_reply":"2021-10-01T08:30:44.822477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ls","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:44.826795Z","iopub.execute_input":"2021-10-01T08:30:44.827039Z","iopub.status.idle":"2021-10-01T08:30:45.509183Z","shell.execute_reply.started":"2021-10-01T08:30:44.827015Z","shell.execute_reply":"2021-10-01T08:30:45.50823Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# clear working dir\n!rm -rf BraTS2021_Testing_Data","metadata":{"execution":{"iopub.status.busy":"2021-10-01T08:30:45.51102Z","iopub.execute_input":"2021-10-01T08:30:45.511406Z","iopub.status.idle":"2021-10-01T08:30:46.195702Z","shell.execute_reply.started":"2021-10-01T08:30:45.511364Z","shell.execute_reply":"2021-10-01T08:30:46.194515Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}