{"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 os\nimport glob\nimport time\nimport shutil\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torchvision.models as models\n\nfrom sklearn import metrics","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:02:10.333217Z","iopub.execute_input":"2021-09-22T20:02:10.334023Z","iopub.status.idle":"2021-09-22T20:02:15.930342Z","shell.execute_reply.started":"2021-09-22T20:02:10.333932Z","shell.execute_reply":"2021-09-22T20:02:15.929477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom(path):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return data","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:02:15.932234Z","iopub.execute_input":"2021-09-22T20:02:15.932444Z","iopub.status.idle":"2021-09-22T20:02:15.942864Z","shell.execute_reply.started":"2021-09-22T20:02:15.932420Z","shell.execute_reply":"2021-09-22T20:02:15.942042Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BrainTumorDataset(Dataset):\n    def __init__(self, root, label_path, transform=None):\n        self.transform = transform\n        self.labels = pd.read_csv(label_path)\n        self.root = root\n        self.types = (\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\")\n\n    \n    def __getitem__(self, idx):\n        brats21id = self.labels.iloc[idx][\"BraTS21ID\"]\n        mgmt_value = self.labels.iloc[idx][\"MGMT_value\"]\n        \n        patient_path = os.path.join(self.root, str(int(brats21id)).zfill(5))\n        t_paths = sorted(\n            glob.glob(os.path.join(patient_path, self.types[0], \"*\")), \n            key=lambda x: int(x[:-4].split(\"-\")[-1]),\n        )\n        \n        t_data = []\n        for i in np.linspace(0,len(t_paths)-1,15):\n            data = load_dicom(t_paths[int(i)])\n            t_data.append(data)\n        t_data = np.stack(t_data, axis=2)\n        if self.transform:\n            t_data = self.transform(t_data)\n        return t_data, float(mgmt_value)\n    \n    \n    def __len__(self):\n        return self.labels.shape[0]","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:14:39.348332Z","iopub.execute_input":"2021-09-22T20:14:39.348592Z","iopub.status.idle":"2021-09-22T20:14:39.357667Z","shell.execute_reply.started":"2021-09-22T20:14:39.348563Z","shell.execute_reply":"2021-09-22T20:14:39.357011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"root = \"../input/rsna-miccai-brain-tumor-radiogenomic-classification\"\n\nnormalize = transforms.Normalize(mean=[0.485, 0.456, 0.406],\n                                     std=[0.229, 0.224, 0.225])\n\ntrain_transform = transforms.Compose([\n    transforms.ToTensor(),\n#     transforms.RandomResizedCrop(224),\n    transforms.Resize(256),\n    transforms.CenterCrop(224),\n    transforms.RandomHorizontalFlip(),\n    transforms.RandomVerticalFlip(),\n#     normalize\n])\n\ntrain_dataset = BrainTumorDataset(os.path.join(root, \"train\"),\n                                  os.path.join(root, \"train_labels.csv\"),\n                                  train_transform)\n\ntrain_loader = DataLoader(\n        train_dataset, batch_size=32, shuffle=True,\n        num_workers=4, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:15:02.076314Z","iopub.execute_input":"2021-09-22T20:15:02.076602Z","iopub.status.idle":"2021-09-22T20:15:02.089928Z","shell.execute_reply.started":"2021-09-22T20:15:02.076571Z","shell.execute_reply":"2021-09-22T20:15:02.088881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize(image3d):\n    fig, ax = plt.subplots(3, 5, figsize=(15, 9))\n    for i, img in enumerate(image3d):\n        ax[i // 5][i % 5].imshow(img, cmap=\"gray\")\n        ax[i // 5][i % 5].axis('off')\n    plt.tight_layout()\n\nexample_batch = next(iter(train_loader))\nfor i in range(5):\n    visualize(example_batch[0][i])","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:02:15.986337Z","iopub.execute_input":"2021-09-22T20:02:15.986546Z","iopub.status.idle":"2021-09-22T20:02:42.309003Z","shell.execute_reply.started":"2021-09-22T20:02:15.986515Z","shell.execute_reply":"2021-09-22T20:02:42.308214Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train(train_loader, model, criterion, optimizer, epoch, device):\n    batch_time = AverageMeter('Time', ':6.3f')\n    data_time = AverageMeter('Data', ':6.3f')\n    losses = AverageMeter('Loss', ':.4e')\n    acc = AverageMeter('Acc', ':6.2f')\n    auc = AverageMeter('AUC', ':6.2f')\n    progress = ProgressMeter(\n        len(train_loader),\n        [batch_time, data_time, losses, acc, auc],\n        prefix=\"Epoch: [{}]\".format(epoch))\n\n    # switch to train mode\n    model.train()\n\n    end = time.time()\n    for i, (images, target) in enumerate(train_loader):\n        # measure data loading time\n        data_time.update(time.time() - end)\n        \n        images = images.to(device)\n        target = target.to(device)\n\n        # compute output\n        output = model(images).squeeze(1)\n        loss = criterion(output, target)\n\n        # measure accuracy and record loss\n        losses.update(loss.item(), images.size(0))\n        acc.update(accuracy(output, target), images.size(0))\n        auc.update(roc_auc(output, target), images.size(0))\n\n        # compute gradient and do SGD step\n        optimizer.zero_grad()\n        loss.backward()\n        optimizer.step()\n\n        # measure elapsed time\n        batch_time.update(time.time() - end)\n        end = time.time()\n\n\n        progress.display(i)\n\n\ndef validate(val_loader, model, criterion, device):\n    batch_time = AverageMeter('Time', ':6.3f')\n    losses = AverageMeter('Loss', ':.4e')\n    acc = AverageMeter('Acc', ':6.2f')\n    auc = AverageMeter('AUC', ':6.2f')\n    progress = ProgressMeter(\n        len(val_loader),\n        [batch_time, losses, acc, auc],\n        prefix='Val: ')\n\n    # switch to evaluate mode\n    model.eval()\n\n    with torch.no_grad():\n        end = time.time()\n        for i, (images, target) in enumerate(val_loader):\n            images = images.to(device)\n            target = target.to(device)\n            \n            # compute output\n            output = model(images).squeeze(1)\n            loss = criterion(output, target)\n\n            # measure accuracy and record loss\n            losses.update(loss.item(), images.size(0))\n            acc.update(accuracy(output, target), images.size(0))\n            auc.update(roc_auc(output, target), images.size(0))\n\n            # measure elapsed time\n            batch_time.update(time.time() - end)\n            end = time.time()\n\n            progress.display(i)\n\n        print(' * Auc {auc.avg:.3f} Acc {acc.avg:.3f}'.format(auc=auc, acc=acc))\n\n    return auc.avg\n\ndef test(test_loader, model, device):\n    batch_time = AverageMeter('Time', ':6.3f')\n    progress = ProgressMeter(\n        len(test_loader),\n        [batch_time],\n        prefix='Test: ')\n\n    # switch to evaluate mode\n    model.eval()\n    \n    scores = []\n\n    with torch.no_grad():\n        end = time.time()\n        for i, (images, _) in enumerate(test_loader):\n            images = images.to(device)\n            \n            # compute output\n            output = model(images).squeeze(1)\n            \n            scores.extend(torch.sigmoid(output).cpu().numpy())\n            \n            # measure elapsed time\n            batch_time.update(time.time() - end)\n            end = time.time()\n\n            progress.display(i)\n            \n    return scores\n\ndef save_checkpoint(state, is_best, filename='checkpoint.pth.tar'):\n    torch.save(state, filename)\n    if is_best:\n        shutil.copyfile(filename, 'model_best.pth.tar')\n\n\nclass AverageMeter(object):\n    \"\"\"Computes and stores the average and current value\"\"\"\n    def __init__(self, name, fmt=':f'):\n        self.name = name\n        self.fmt = fmt\n        self.reset()\n\n    def reset(self):\n        self.val = 0\n        self.avg = 0\n        self.sum = 0\n        self.count = 0\n\n    def update(self, val, n=1):\n        self.val = val\n        self.sum += val * n\n        self.count += n\n        self.avg = self.sum / self.count\n\n    def __str__(self):\n        fmtstr = '{name} {val' + self.fmt + '} ({avg' + self.fmt + '})'\n        return fmtstr.format(**self.__dict__)\n\n\nclass ProgressMeter(object):\n    def __init__(self, num_batches, meters, prefix=\"\"):\n        self.batch_fmtstr = self._get_batch_fmtstr(num_batches)\n        self.meters = meters\n        self.prefix = prefix\n\n    def display(self, batch):\n        entries = [self.prefix + self.batch_fmtstr.format(batch)]\n        entries += [str(meter) for meter in self.meters]\n        print('\\t'.join(entries))\n\n    def _get_batch_fmtstr(self, num_batches):\n        num_digits = len(str(num_batches // 1))\n        fmt = '{:' + str(num_digits) + 'd}'\n        return '[' + fmt + '/' + fmt.format(num_batches) + ']'\n\n\ndef adjust_learning_rate(optimizer, epoch, initial_lr):\n    \"\"\"Sets the learning rate to the initial LR decayed by 10 every 30 epochs\"\"\"\n    lr = initial_lr * (0.1 ** (epoch // 30))\n    for param_group in optimizer.param_groups:\n        param_group['lr'] = lr\n\n\ndef accuracy(output, target):\n    \"\"\"Computes the accuracy for threshold of 0\"\"\"\n    with torch.no_grad():\n        batch_size = target.size(0)\n        pred = (output >= 0).to(torch.float32)\n        correct = pred.eq(target).float().sum()\n        return correct / batch_size\n\n    \ndef roc_auc(output, target):\n    \"\"\"Computes the accuracy for threshold of 0\"\"\"\n    with torch.no_grad():\n        scores = torch.sigmoid(output)\n        return metrics.roc_auc_score(target.cpu().numpy(), scores.cpu().numpy())","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:15:16.164623Z","iopub.execute_input":"2021-09-22T20:15:16.164897Z","iopub.status.idle":"2021-09-22T20:15:16.195275Z","shell.execute_reply.started":"2021-09-22T20:15:16.164869Z","shell.execute_reply":"2021-09-22T20:15:16.194616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = 'cuda' if torch.cuda.is_available() else 'cpu'\nprint(device)\n\nmodel = models.resnet50()\nmodel.conv1 = nn.Conv2d(15, model.conv1.out_channels, model.conv1.kernel_size,\n                        model.conv1.stride, model.conv1.padding, bias=model.conv1.bias)\nmodel.fc = nn.Linear(model.fc.in_features, 1)\n\nmodel = model.to(device)\n# print(model)\n\ncriterion = nn.BCEWithLogitsLoss().to(device)\n\noptimizer = torch.optim.SGD(model.parameters(), 0.1,\n                            momentum=0.9,\n                            weight_decay=1e-4)","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:02:42.344403Z","iopub.execute_input":"2021-09-22T20:02:42.344594Z","iopub.status.idle":"2021-09-22T20:02:42.831926Z","shell.execute_reply.started":"2021-09-22T20:02:42.344572Z","shell.execute_reply":"2021-09-22T20:02:42.831020Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"best_auc = 0\n\nfor epoch in range(10):\n    adjust_learning_rate(optimizer, epoch, 0.1)\n\n    # train for one epoch\n    train(train_loader, model, criterion, optimizer, epoch, device)\n\n    # evaluate on validation set\n    auc = validate(train_loader, model, criterion, device)\n\n    # remember best acc@1 and save checkpoint\n    is_best = auc > best_auc\n    best_auc = max(auc, best_auc)\n\n    save_checkpoint({\n        'epoch': epoch + 1,\n        'arch': 'resnet50',\n        'state_dict': model.state_dict(),\n        'best_auc': best_auc,\n        'optimizer' : optimizer.state_dict(),\n    }, is_best)","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:17:54.725401Z","iopub.execute_input":"2021-09-22T20:17:54.725668Z","iopub.status.idle":"2021-09-22T20:30:02.855580Z","shell.execute_reply.started":"2021-09-22T20:17:54.725638Z","shell.execute_reply":"2021-09-22T20:30:02.854610Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_transform = transforms.Compose([\n    transforms.ToTensor(),\n    transforms.Resize(256),\n    transforms.CenterCrop(224)\n#     normalize\n])\n\ntest_dataset = BrainTumorDataset(os.path.join(root, \"test\"),\n                                  os.path.join(root, \"sample_submission.csv\"),\n                                  test_transform)\n\ntest_loader = DataLoader(\n        test_dataset, batch_size=32,\n        num_workers=4, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:47:50.392177Z","iopub.execute_input":"2021-09-22T20:47:50.392476Z","iopub.status.idle":"2021-09-22T20:47:50.405502Z","shell.execute_reply.started":"2021-09-22T20:47:50.392444Z","shell.execute_reply":"2021-09-22T20:47:50.404630Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"scores = test(test_loader, model, device)\n\nids = test_dataset.labels[\"BraTS21ID\"].tolist()\n\nsubmission = pd.DataFrame({\"BraTS21ID\": ids, \"MGMT_value\": scores})\nsubmission.to_csv(\"submission.csv\", index=False)\n\n","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:47:59.388001Z","iopub.execute_input":"2021-09-22T20:47:59.388298Z","iopub.status.idle":"2021-09-22T20:48:04.214733Z","shell.execute_reply.started":"2021-09-22T20:47:59.388268Z","shell.execute_reply":"2021-09-22T20:48:04.213945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2021-09-22T20:48:11.533474Z","iopub.execute_input":"2021-09-22T20:48:11.533775Z","iopub.status.idle":"2021-09-22T20:48:11.550352Z","shell.execute_reply.started":"2021-09-22T20:48:11.533738Z","shell.execute_reply":"2021-09-22T20:48:11.549568Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}