{"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 sys \nimport gc\nimport glob\nimport random \nimport cv2\nimport numpy as np \nimport pandas as pd \nfrom sklearn import metrics\nimport torch\nfrom torch import nn\nimport torch.nn.functional as F\nfrom tqdm import tqdm_notebook as tqdm\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\n\ngc.enable()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2021-08-09T20:33:00.770478Z","iopub.execute_input":"2021-08-09T20:33:00.770896Z","iopub.status.idle":"2021-08-09T20:33:03.315175Z","shell.execute_reply.started":"2021-08-09T20:33:00.770809Z","shell.execute_reply":"2021-08-09T20:33:03.314184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"package_path = \"../input/efficientnet-pytorch/EfficientNet-PyTorch/EfficientNet-PyTorch-master/\"\nsys.path.append(package_path)\n\nimport efficientnet_pytorch","metadata":{"execution":{"iopub.status.busy":"2021-08-09T20:33:03.317830Z","iopub.execute_input":"2021-08-09T20:33:03.318126Z","iopub.status.idle":"2021-08-09T20:33:03.363475Z","shell.execute_reply.started":"2021-08-09T20:33:03.318085Z","shell.execute_reply":"2021-08-09T20:33:03.362489Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 48 \nNUM_FOLDS = 5\nNUM_EPOCHS = 3\nDEVICE = 'cuda'\nLEARNING_RATE = 1e-5\n\ntrain_df = pd.read_csv('../input/rsna-brain-folds/train_folds.csv')","metadata":{"execution":{"iopub.status.busy":"2021-08-09T20:33:03.365538Z","iopub.execute_input":"2021-08-09T20:33:03.365960Z","iopub.status.idle":"2021-08-09T20:33:03.383084Z","shell.execute_reply.started":"2021-08-09T20:33:03.365916Z","shell.execute_reply":"2021-08-09T20:33:03.382180Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def set_random_seed(random_seed):\n    random.seed(random_seed)\n    np.random.seed(random_seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(random_seed)\n\n    torch.manual_seed(random_seed)\n    torch.cuda.manual_seed(random_seed)\n    torch.cuda.manual_seed_all(random_seed)\n\n    torch.backends.cudnn.deterministic = True\n    \nset_random_seed(1729)","metadata":{"execution":{"iopub.status.busy":"2021-08-09T20:33:03.384950Z","iopub.execute_input":"2021-08-09T20:33:03.385373Z","iopub.status.idle":"2021-08-09T20:33:03.396670Z","shell.execute_reply.started":"2021-08-09T20:33:03.385329Z","shell.execute_reply":"2021-08-09T20:33:03.395544Z"},"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-08-09T20:33:03.398199Z","iopub.execute_input":"2021-08-09T20:33:03.398596Z","iopub.status.idle":"2021-08-09T20:33:03.404841Z","shell.execute_reply.started":"2021-08-09T20:33:03.398556Z","shell.execute_reply":"2021-08-09T20:33:03.403811Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Dataset:\n    def __init__(self, paths, targets):\n        self.paths = paths\n        self.targets = targets\n    \n    def __len__(self):\n        return len(self.paths)\n    \n    def __getitem__(self, index, inference_only=False):\n        _id = self.paths[index]\n        patient_path = f\"../input/rsna-miccai-brain-tumor-radiogenomic-classification/train/{str(_id).zfill(5)}/\"\n        channels = []\n        for t in (\"FLAIR\", \"T1w\", \"T1wCE\"): # \"T2w\"\n            t_paths = sorted(\n                glob.glob(os.path.join(patient_path, t, \"*\")), \n                key=lambda x: int(x[:-4].split(\"-\")[-1]),\n            )\n            # start, end = int(len(t_paths) * 0.475), int(len(t_paths) * 0.525)\n            x = len(t_paths)\n            if x < 10:\n                r = range(x)\n            else:\n                d = x // 10\n                r = range(d, x - d, d)\n                \n            channel = []\n            # for i in range(start, end + 1):\n            for i in r:\n                channel.append(cv2.resize(load_dicom(t_paths[i]), (256, 256)) / 255)\n            channel = np.mean(channel, axis=0)\n            channels.append(channel)\n        \n        if inference_only:\n            return {\n                'X': torch.tensor(channels).float()\n            }\n        \n        return {\n            \"X\": torch.tensor(channels).float(), \n            \"y\": torch.tensor(self.targets[index], dtype=torch.float),\n        }","metadata":{"execution":{"iopub.status.busy":"2021-08-09T20:33:03.406911Z","iopub.execute_input":"2021-08-09T20:33:03.407734Z","iopub.status.idle":"2021-08-09T20:33:03.421342Z","shell.execute_reply.started":"2021-08-09T20:33:03.407690Z","shell.execute_reply":"2021-08-09T20:33:03.420293Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Model(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.net = efficientnet_pytorch.EfficientNet.from_name(\"efficientnet-b0\")\n        checkpoint = torch.load(\"../input/efficientnet-pytorch/efficientnet-b0-08094119.pth\")\n        self.net.load_state_dict(checkpoint)\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 out","metadata":{"execution":{"iopub.status.busy":"2021-08-09T20:33:03.422996Z","iopub.execute_input":"2021-08-09T20:33:03.423664Z","iopub.status.idle":"2021-08-09T20:33:03.434452Z","shell.execute_reply.started":"2021-08-09T20:33:03.423613Z","shell.execute_reply":"2021-08-09T20:33:03.433375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_fn(data_loader, model, optimizer, device, valid_data, best_score, fold):\n    \n    for iteration, data in enumerate(data_loader):\n        # print('chk1')\n        features = data['X']\n        target = data['y']\n        features = features.to(device, dtype=torch.float)\n        target = target.to(device, dtype=torch.float)\n        \n        optimizer.zero_grad()\n        \n        # print('chk2')\n        predictions = model(features)\n        loss = F.binary_cross_entropy_with_logits(predictions.flatten(), target)\n        loss.backward()\n        optimizer.step()\n        \n        # print('chk3')\n        if len(data_loader) == iteration + 1:\n            current_score = eval_fn(valid_data, model, device)\n            if current_score > best_score: \n                best_score = current_score\n                torch.save(model.state_dict(), f'model_{fold}.pth')\n            print(f'Step: {iteration}, Current Score: {current_score}, Best Score: {best_score}')\n                \n    return best_score\n        \ndef eval_fn(data_loader, model, device):\n    final_predictions = []\n    final_targets = []\n    \n    model.eval()\n    \n    with torch.no_grad():\n        for data in data_loader:\n            features = data['X']\n            target = data['y']\n\n            features = features.to(device, dtype=torch.float)\n            target = target.to(device, dtype=torch.float)\n            \n            predictions = model(features).squeeze()\n            predictions = torch.sigmoid(predictions).cpu().detach().numpy().tolist()\n            final_predictions.extend(predictions)\n            \n            target = target.cpu().detach().numpy().tolist()\n            final_targets.extend(target)\n        \n        score = metrics.roc_auc_score(final_targets, final_predictions)\n    \n        return score","metadata":{"execution":{"iopub.status.busy":"2021-08-09T20:33:03.437354Z","iopub.execute_input":"2021-08-09T20:33:03.437767Z","iopub.status.idle":"2021-08-09T20:33:03.452837Z","shell.execute_reply.started":"2021-08-09T20:33:03.437723Z","shell.execute_reply":"2021-08-09T20:33:03.451510Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run(data, fold):\n    \n    print(f'Fold: {fold}')\n    \n    train_data = data[data['kfold'] != fold].reset_index(drop=True)\n    val_data = data[data['kfold'] == fold].reset_index(drop=True)\n    \n    train_dataset = Dataset(\n        train_data.BraTS21ID.values,\n        train_data.MGMT_value.values\n    )\n    val_dataset = Dataset(\n        val_data.BraTS21ID.values,\n        val_data.MGMT_value.values\n    )\n    \n    train_loader = torch.utils.data.DataLoader(\n        train_dataset, \n        batch_size=BATCH_SIZE\n    )\n    val_loader = torch.utils.data.DataLoader(\n        val_dataset,\n        batch_size=BATCH_SIZE\n    )\n    \n    \n    DEVICE = torch.device('cuda')\n    model = Model()\n    model.to(DEVICE)\n    \n    optimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE)\n    \n    best_score = -1\n    for epoch in tqdm(range(NUM_EPOCHS)):\n        print(f'Epoch: {epoch + 1}/{NUM_EPOCHS}')\n        best_score = train_fn(train_loader, model, optimizer, DEVICE, val_loader, best_score, fold)\n        print(f'Best Score for epoch {epoch + 1}: {best_score}')\n        \n    del model\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2021-08-09T20:33:03.454486Z","iopub.execute_input":"2021-08-09T20:33:03.455144Z","iopub.status.idle":"2021-08-09T20:33:03.467618Z","shell.execute_reply.started":"2021-08-09T20:33:03.455082Z","shell.execute_reply":"2021-08-09T20:33:03.466404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for fold in range(NUM_FOLDS):\n    run(train_df, fold)","metadata":{"execution":{"iopub.status.busy":"2021-08-09T20:33:03.469469Z","iopub.execute_input":"2021-08-09T20:33:03.469891Z","iopub.status.idle":"2021-08-09T20:43:07.177897Z","shell.execute_reply.started":"2021-08-09T20:33:03.469845Z","shell.execute_reply":"2021-08-09T20:43:07.175253Z"},"trusted":true},"execution_count":null,"outputs":[]}]}