{"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":"# Goal","metadata":{}},{"cell_type":"markdown","source":"**In this competition you will predict the genetic subtype of glioblastoma using MRI (magnetic resonance imaging) scans to train and test your model to detect for the presence of MGMT promoter methylation.**  \n(このコンペでは膠芽腫のMRI画像から、（脳腫瘍の治療に必要な）膠芽腫の遺伝子バイオマーカーであるMGMTプロモーターのメチル化を予測するモデルを作成する。)   \nThis  competition is not predicting whether brain tumor or not.  \n**(※脳腫瘍かどうかを当てるコンペではない!)**\n\nshow the probability of MGMT promoter methylation or not.  ","metadata":{}},{"cell_type":"markdown","source":"## Ref","metadata":{}},{"cell_type":"markdown","source":"[日本語解説](https://www.kaggle.com/chumajin/brain-tumor-eda-for-starter-version)   \n[Yaroslav Isaienkov](https://www.kaggle.com/ihelon/brain-tumor-eda-with-animations-and-modeling) This notebook almost based from his notebook  \n[Ayush Thakur](https://www.kaggle.com/ayuraj/brain-tumor-eda-and-interactive-viz-with-w-b)\n","metadata":{}},{"cell_type":"markdown","source":"# 🔍EDA","metadata":{}},{"cell_type":"markdown","source":"### Library","metadata":{}},{"cell_type":"code","source":"import os\nimport json\nimport glob\nimport random\nimport collections\n\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom tqdm import tqdm\n\n\nimport pydicom","metadata":{"execution":{"iopub.status.busy":"2021-08-25T06:48:59.822304Z","iopub.execute_input":"2021-08-25T06:48:59.822629Z","iopub.status.idle":"2021-08-25T06:48:59.827926Z","shell.execute_reply.started":"2021-08-25T06:48:59.822602Z","shell.execute_reply":"2021-08-25T06:48:59.827053Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#import wandb\n#wandb.login()","metadata":{"execution":{"iopub.status.busy":"2021-08-25T06:27:38.772906Z","iopub.execute_input":"2021-08-25T06:27:38.773266Z","iopub.status.idle":"2021-08-25T06:27:38.776857Z","shell.execute_reply.started":"2021-08-25T06:27:38.773234Z","shell.execute_reply":"2021-08-25T06:27:38.775784Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Data Vizualization","metadata":{}},{"cell_type":"code","source":"train_df = pd.read_csv(\"../input/rsna-miccai-brain-tumor-radiogenomic-classification/train_labels.csv\")\ntrain_df","metadata":{"execution":{"iopub.status.busy":"2021-08-25T06:21:01.417941Z","iopub.execute_input":"2021-08-25T06:21:01.418311Z","iopub.status.idle":"2021-08-25T06:21:01.452597Z","shell.execute_reply.started":"2021-08-25T06:21:01.418280Z","shell.execute_reply":"2021-08-25T06:21:01.451646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The dataset consist of 585 patients and are given by unique id,BraTS21ID  \n* Class 0 = the presence of MGMT promoter methylation no exist.  \n* Class 1 = the presence of MGMT promoter methylation exist.","metadata":{}},{"cell_type":"code","source":"plt.figure(figsize=(5,5))\nsns.countplot(data=train_df, x=\"MGMT_value\")","metadata":{"execution":{"iopub.status.busy":"2021-08-25T06:27:42.072953Z","iopub.execute_input":"2021-08-25T06:27:42.073306Z","iopub.status.idle":"2021-08-25T06:27:42.219851Z","shell.execute_reply.started":"2021-08-25T06:27:42.073276Z","shell.execute_reply":"2021-08-25T06:27:42.218957Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Look dcm one data ","metadata":{}},{"cell_type":"code","source":"dataset = pydicom.filereader.dcmread(\"../input/rsna-miccai-brain-tumor-radiogenomic-classification/train/00000/FLAIR/Image-110.dcm\")\nimg = dataset.pixel_array # Dataset.pixel_array returns a numpy.ndarray\n\nfig, ax = plt.subplots()\nax.imshow(img, cmap=\"gray\")\nax.set_axis_off()\nprint('Shape of data: ', img.shape)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2021-08-25T06:59:37.068386Z","iopub.execute_input":"2021-08-25T06:59:37.068709Z","iopub.status.idle":"2021-08-25T06:59:37.164193Z","shell.execute_reply.started":"2021-08-25T06:59:37.068680Z","shell.execute_reply":"2021-08-25T06:59:37.163354Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### look at the number of files (.dcm) ","metadata":{}},{"cell_type":"code","source":"filenames = glob.glob('../input/rsna-miccai-brain-tumor-radiogenomic-classification/train/*/*/*')\nprint(f'Total number of files: {len(filenames)}')","metadata":{"execution":{"iopub.status.busy":"2021-08-25T06:29:40.515151Z","iopub.execute_input":"2021-08-25T06:29:40.515491Z","iopub.status.idle":"2021-08-25T06:30:45.684980Z","shell.execute_reply.started":"2021-08-25T06:29:40.515462Z","shell.execute_reply":"2021-08-25T06:30:45.684149Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### look at the distribution of files per scan types","metadata":{}},{"cell_type":"code","source":"label_dict = {\n    'FLAIR': [],\n    'T1w': [],\n    'T1wCE': [],\n    'T2w': []\n}\n\nfor filename in tqdm(filenames):\n    scan = filename.split('/')[-2]\n    if scan==\"FLAIR\":\n        label_dict[\"FLAIR\"].append(filename)\n    elif scan==\"T1w\":\n        label_dict[\"T1w\"].append(filename)\n    elif scan==\"T1wCE\":\n        label_dict[\"T1wCE\"].append(filename)\n    elif scan==\"T2w\":\n        label_dict[\"T2w\"].append(filename)\n        \nprint('Size of FLAIR scan: {}, T1w scan: {}, T1wCE scan: {}, T2W scan: {}'.format(len(label_dict[\"FLAIR\"]),\n                                                                                  len(label_dict[\"T1w\"]),\n                                                                                  len(label_dict[\"T1wCE\"]),\n                                                                                  len(label_dict[\"T2w\"])))","metadata":{"execution":{"iopub.status.busy":"2021-08-25T06:49:20.088208Z","iopub.execute_input":"2021-08-25T06:49:20.088527Z","iopub.status.idle":"2021-08-25T06:49:20.508746Z","shell.execute_reply.started":"2021-08-25T06:49:20.088498Z","shell.execute_reply":"2021-08-25T06:49:20.507803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONFIG = {'IMG_SIZE': 224, \n          'NUM_FRAMES': 14,\n          'competition': 'rsna-miccai-brain', \n          '_wandb_kernel': 'ayut'}","metadata":{"execution":{"iopub.status.busy":"2021-08-25T06:33:05.690254Z","iopub.execute_input":"2021-08-25T06:33:05.690585Z","iopub.status.idle":"2021-08-25T06:33:05.694930Z","shell.execute_reply.started":"2021-08-25T06:33:05.690556Z","shell.execute_reply":"2021-08-25T06:33:05.694114Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"run = wandb.init(project='brain-tumor-viz', config=CONFIG)\ndata = [['FLAIR', 74248], ['T1w', 77627], ['T1wCE' ,96766], ['T2w', 100000]]\ntable = wandb.Table(data=data, columns = [\"Scan Type\", \"Num Files\"])\nwandb.log({\"my_bar_chart_id\" : wandb.plot.bar(table, \"Scan Type\", \"Num Files\", title=\"Scan Types vs Number of Dicom files\")})","metadata":{"execution":{"iopub.status.busy":"2021-08-25T06:49:42.523403Z","iopub.execute_input":"2021-08-25T06:49:42.523732Z","iopub.status.idle":"2021-08-25T06:49:52.844274Z","shell.execute_reply.started":"2021-08-25T06:49:42.523703Z","shell.execute_reply":"2021-08-25T06:49:52.843353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"https://wandb.ai/hondykaito/brain-tumor-viz?workspace=user-hondykaito","metadata":{}},{"cell_type":"markdown","source":"![img](https://i.ibb.co/QNjHyQd/W-B-Chart-7-16-2021-3-49-04-AM.png)","metadata":{}},{"cell_type":"markdown","source":"Look dcm data list","metadata":{}},{"cell_type":"code","source":"train_df.shape","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:29:09.402928Z","iopub.execute_input":"2021-08-24T07:29:09.40344Z","iopub.status.idle":"2021-08-24T07:29:09.412177Z","shell.execute_reply.started":"2021-08-24T07:29:09.403372Z","shell.execute_reply":"2021-08-24T07:29:09.410859Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### type of MRI scan  \nEach case consists of four structural multi-parametric MRI (mpMRI) scans.  \n* Fluid Attenuated Inversion Recovery (FLAIR)  \n    - FLAIR can be roughly thought of as T2, in which the water is also black, making it easier to find the lesion.  \n    (FLAIRは水も黒くすることで、より病変を探しやすくなったT2)\n* T1-weighted pre-contrast (T1w) \n* T1-weighted post-contrast (T1Gd)\n    - T1Gd is T1 imaging with contrast medium and is the method that best reflects the location, size, and shape of the mass.  \n    (T1Gdは造影剤を用いたT1撮影で、腫瘤の位置、大きさ、形が最もよく反映される方法である。)  \n* T2-weighted (T2)\n    - T2 :Water is painted white.Lesions appear white. Suitable for lesion evaluation.  \n    (水が白く描かれる。病変が白く映る。病変の評価に適している。)","metadata":{}},{"cell_type":"markdown","source":"### Visualize dcm images","metadata":{}},{"cell_type":"markdown","source":"### Ref\n[Preprocessing DICOM](https://www.kaggle.com/c/rsna-miccai-brain-tumor-radiogenomic-classification/discussion/253000)","metadata":{}},{"cell_type":"code","source":"# DICOM to PNG dataset (128 GB -> 5.2 GB)\ndef load_dicom(path):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    dara = 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\n\ndef visualize_sample(brats21id, slice_i, mgmt_value, types=(\"FLAIR\", \"T1w\", \"T1wCE\", \"T2w\")):\n    plt.figure(figsize=(16, 5))\n    patient_path = os.path.join(\"../input/rsna-miccai-brain-tumor-radiogenomic-classification/train/\", str(brats21id).zfill(5))\n    for i, t in enumerate(types, 1):\n        t_paths = sorted(glob.glob(os.path.join(patient_path, t, \"*\")),\n                        key=lambda x: int(x[:-4].split(\"-\")[-1]))\n        data = load_dicom(t_paths[int(len(t_paths)* slice_i)]) # why?\n        plt.subplot(1, 4, i)\n        plt.imshow(data, cmap=\"gray\")\n        plt.title(f\"{t}\", fontsize=16)\n        plt.axis(\"off\")\n        \n    plt.suptitle(f\"MGMT_value: {mgmt_value}\", fontsize=18)\n    plt.show","metadata":{"execution":{"iopub.status.busy":"2021-08-25T07:00:20.077207Z","iopub.execute_input":"2021-08-25T07:00:20.077521Z","iopub.status.idle":"2021-08-25T07:00:20.086696Z","shell.execute_reply.started":"2021-08-25T07:00:20.077494Z","shell.execute_reply":"2021-08-25T07:00:20.085885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for i in random.sample(range(train_df.shape[0]), 10):\n    _brats21id = train_df.iloc[i][\"BraTS21ID\"]\n    _mgmt_value = train_df.iloc[i][\"MGMT_value\"]\n    visualize_sample(brats21id = _brats21id, mgmt_value=_mgmt_value, slice_i=0.5)","metadata":{"execution":{"iopub.status.busy":"2021-08-25T07:01:12.076933Z","iopub.execute_input":"2021-08-25T07:01:12.077289Z","iopub.status.idle":"2021-08-25T07:01:16.015992Z","shell.execute_reply.started":"2021-08-25T07:01:12.077258Z","shell.execute_reply":"2021-08-25T07:01:16.015243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Estimation\n#### * Sometimes I cannot detect tumor in images\n    * Just Missing Value or cannnot by only using plt.show (show 3D image)?","metadata":{}},{"cell_type":"markdown","source":"# 🎩Model","metadata":{}},{"cell_type":"markdown","source":"### Libarary","metadata":{}},{"cell_type":"code","source":"package_path =  \"../input/efficientnet-pytorch/EfficientNet-PyTorch/EfficientNet-PyTorch-master/\"","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:29:15.746854Z","iopub.execute_input":"2021-08-24T07:29:15.747179Z","iopub.status.idle":"2021-08-24T07:29:15.753399Z","shell.execute_reply.started":"2021-08-24T07:29:15.747148Z","shell.execute_reply":"2021-08-24T07:29:15.75085Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import sys\nimport time\nsys.path.append(package_path)\n\nimport torch\nfrom torch import nn\nfrom torch.utils import data as torch_data\nfrom sklearn import model_selection as sk_model_selection\nfrom torch.nn import functional as torch_functional\nimport efficientnet_pytorch\n\n# choice \n#from sklearn.model_selection import StratifiedKFold\n#from sklearn.model_selection import KFold","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:34:28.825982Z","iopub.execute_input":"2021-08-24T07:34:28.826502Z","iopub.status.idle":"2021-08-24T07:34:28.833005Z","shell.execute_reply.started":"2021-08-24T07:34:28.826439Z","shell.execute_reply":"2021-08-24T07:34:28.831573Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# maintain Reproducibility\ndef set_seed(seed):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True # deterministic algorithm (決定論的アルゴリズムを使用)\n    \nset_seed(42)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:34:29.804095Z","iopub.execute_input":"2021-08-24T07:34:29.804485Z","iopub.status.idle":"2021-08-24T07:34:29.811996Z","shell.execute_reply.started":"2021-08-24T07:34:29.804456Z","shell.execute_reply":"2021-08-24T07:34:29.810431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.read_csv(\"../input/rsna-miccai-brain-tumor-radiogenomic-classification/train_labels.csv\")\n\ndf_train, df_valid = sk_model_selection.train_test_split(\n    df, \n    test_size=0.2, \n    random_state=42, \n    stratify=train_df[\"MGMT_value\"],\n)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:34:29.958892Z","iopub.execute_input":"2021-08-24T07:34:29.959206Z","iopub.status.idle":"2021-08-24T07:34:29.973767Z","shell.execute_reply.started":"2021-08-24T07:34:29.959176Z","shell.execute_reply":"2021-08-24T07:34:29.972646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataRetriever(torch_data.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    # arrange later \n    def __getitem__(self, index):\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(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            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) # why?\n                \n            channel = np.mean(channel, axis=0) # axis=0 by column\n            channels.append(channel)\n            \n        \n        y = torch.tensor(self.targets[index], dtype=torch.float)\n        \n        return {\"X\": torch.tensor(channels).float(), \"y\": y}    ","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:34:30.110883Z","iopub.execute_input":"2021-08-24T07:34:30.111326Z","iopub.status.idle":"2021-08-24T07:34:30.123388Z","shell.execute_reply.started":"2021-08-24T07:34:30.111286Z","shell.execute_reply":"2021-08-24T07:34:30.1217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_retriever = DataRetriever(\n    df_train[\"BraTS21ID\"].values,\n    df_train[\"MGMT_value\"].values\n)\n\nvalid_data_retriever = DataRetriever(\n    df_valid[\"BraTS21ID\"].values,\n    df_valid[\"MGMT_value\"].values\n)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:30:02.148702Z","iopub.execute_input":"2021-08-24T07:30:02.149086Z","iopub.status.idle":"2021-08-24T07:30:02.157323Z","shell.execute_reply.started":"2021-08-24T07:30:02.149055Z","shell.execute_reply":"2021-08-24T07:30:02.155822Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(16, 6))\nfor i in range(3):\n    plt.subplot(1, 3, i + 1)\n    plt.imshow(train_data_retriever[13][\"X\"].numpy()[i], cmap=\"gray\")","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:30:02.30512Z","iopub.execute_input":"2021-08-24T07:30:02.305535Z","iopub.status.idle":"2021-08-24T07:30:03.520107Z","shell.execute_reply.started":"2021-08-24T07:30:02.305504Z","shell.execute_reply":"2021-08-24T07:30:03.519096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"* margin is black(pixel = 0) and brian is centered in the image","metadata":{}},{"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        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-24T07:30:03.52211Z","iopub.execute_input":"2021-08-24T07:30:03.522592Z","iopub.status.idle":"2021-08-24T07:30:03.530941Z","shell.execute_reply.started":"2021-08-24T07:30:03.522535Z","shell.execute_reply":"2021-08-24T07:30:03.529234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LossMeter:\n    def __init__(self):\n        self.avg = 0\n        self.n = 0\n    \n    def update(self, val):\n        self.n += 1\n        # incremental update\n        self.avg = val / self.n + (self.n - 1) / self.n * self.avg\n        \nclass AccMeter:\n    def __init__(self):\n        self.avg = 0\n        self.n = 0\n    \n    def update(self, y_true, y_pred):\n        y_true = y_true.cpu().numpy().astype(int)\n        y_pred = y_pred.cpu().numpy() >= 0\n        last_n = self.n\n        self.n += len(y_true)\n        true_count = np.sum(y_true == y_pred)\n        # incremental update\n        self.avg = true_count / self.n + last_n / self.n * self.avg","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:38:44.620924Z","iopub.execute_input":"2021-08-24T07:38:44.621369Z","iopub.status.idle":"2021-08-24T07:38:44.629617Z","shell.execute_reply.started":"2021-08-24T07:38:44.621311Z","shell.execute_reply":"2021-08-24T07:38:44.628501Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Trainer\n\nclass Trainer:\n    # do train\n    def __init__(self, model, device, optimizer, criterion, loss_meter, score_meter):\n        self.model = model\n        self.device = device\n        self.optimizer = optimizer\n        self.criterion = criterion\n        self.loss_meter = loss_meter\n        self.score_meter = score_meter\n        \n        self.best_valid_score = -np.inf\n        self.n_patience = 0 # for early_stopping (by 100epoch)\n        \n        self.messages = {\n            \"epoch\": \"[Epoch {}: {}] loss: {:.5f}, score: {:.5f}, time: {} s\",\n            \"checkpoint\": \"The score improved from {:.5f} to {:.5f}. Save model to '{}'\",\n            \"patience\": \"\\nValid score didn't improve last {} epochs.\"\n        }\n        \n    def fit(self, epochs, train_loader, valid_loader, save_path, patience):\n        for n_epoch in range(1, epochs + 1):\n            self.info_message(\"EPOCH: {}\", n_epoch)\n            \n            train_loss, train_score, train_time = self.train_epoch(train_loader)\n            valid_loss, valid_score, valid_time = self.valid_epoch(valid_loader)\n            \n            self.info_message(\n                self.messages[\"epoch\"], \"Train\", n_epoch, train_loss, train_score, train_time\n            )\n            self.info_message(\n                self.messages[\"epoch\"], \"Valid\", n_epoch, valid_loss, valid_score, valid_time\n            )\n            \n            if True:\n                # if self.best_valid_score < valid_score:\n                self.info_message(\n                    self.messages[\"checkpoint\"], self.best_valid_score, valid_score, save_path\n                )\n                self.best_valid_score = valid_score\n                self.save_model(n_epoch, save_path)\n                self.n_patience = 0\n            else:\n                self.n_patience += 1\n            \n            if self.n_patience >= patience:\n                self.info_message(self.messages[\"patience\"], patience)\n                break\n            \n                \n    def train_epoch(self, train_loader):\n        self.model.train()\n        t = time.time()\n        train_loss = self.loss_meter()\n        train_score = self.score_meter()\n        \n        for step, batch in enumerate(train_loader, 1):\n            X = batch[\"X\"].to(self.device)\n            targets = batch[\"y\"].to(self.device)\n            self.optimizer.zero_grad()\n            outputs = self.model(X).squeeze(1)\n            \n            loss = self.criterion(outputs, targets)\n            loss.backward()\n            \n            train_loss.update(loss.detach().item())\n            train_score.update(targets, outputs.detach())\n            \n            self.optimizer.step() # update weight\n            \n            _loss, _score = train_loss.avg, train_score.avg\n            message = 'Train Step {}/{}, train_loss: {:.5f}, train_score: {:.5f}'\n            self.info_message(message, step, len(train_loader), _loss, _score, end=\"\\r\")\n                   \n        return train_loss.avg, train_score.avg, int(time.time() - t)\n            \n    def valid_epoch(self, valid_loader):\n        self.model.eval()\n        t = time.time()\n        valid_loss = self.loss_meter()\n        valid_score = self.score_meter()\n        \n        for step, batch in enumerate(valid_loader, 1):\n            with torch.no_grad():\n                X = batch[\"X\"].to(self.device)\n                targets = batch[\"y\"].to(self.device)\n\n                outputs = self.model(X).squeeze(1)\n                loss = self.criterion(outputs, targets)\n\n                valid_loss.update(loss.detach().item())\n                valid_score.update(targets, outputs)\n                \n            _loss, _score = valid_loss.avg, valid_score.avg\n            message = 'Valid Step {}/{}, valid_loss: {:.5f}, valid_score: {:.5f}' \n            self.info_message(message, step, len(valid_loader), _loss, _score, end=\"\\r\")\n            \n        return valid_loss.avg, valid_score.avg, int(time.time() - t)\n        \n    def save_model(self, n_epoch, save_path):\n        torch.save(\n        {\n            \"model_state_dict\" : self.model.state_dict(),\n            \"optimizer_state_dict\" : self.optimizer.state_dict(),\n            \"best_valid_score\" : self.best_valid_score,\n            \"n_epoch\" : n_epoch,\n        },\n        save_path\n        )\n    @staticmethod        \n    def info_message(message, *args, end=\"\\n\"):\n        print(message.format(*args), end=end)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:30:03.880346Z","iopub.execute_input":"2021-08-24T07:30:03.880718Z","iopub.status.idle":"2021-08-24T07:30:03.902658Z","shell.execute_reply.started":"2021-08-24T07:30:03.880687Z","shell.execute_reply":"2021-08-24T07:30:03.901404Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Training\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Change later\ntrain_data_retriever = DataRetriever(\n    df_train[\"BraTS21ID\"].values,\n    df_train[\"MGMT_value\"].values\n)\n\nvalid_data_retriever = DataRetriever(\n    df_valid[\"BraTS21ID\"].values, \n    df_valid[\"MGMT_value\"].values\n)\n    \ntrain_loader = torch_data.DataLoader(train_data_retriever, batch_size=32, shuffle=True, num_workers=8)\nvalid_loader = torch_data.DataLoader(valid_data_retriever, batch_size=32, shuffle=False, num_workers=8)\n    \nmodel = Model()\nmodel.to(device)\n    \noptimizer = torch.optim.Adam(model.parameters(), lr=0.001)\ncriterion = torch_functional.binary_cross_entropy_with_logits\n    \ntrainer = Trainer(model, device, optimizer, criterion, LossMeter, AccMeter)\n    \nhistory = trainer.fit(10, train_loader, valid_loader, f\"best-model-0.pth\" ,100)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:38:47.507082Z","iopub.execute_input":"2021-08-24T07:38:47.507662Z","iopub.status.idle":"2021-08-24T07:49:48.562674Z","shell.execute_reply.started":"2021-08-24T07:38:47.507628Z","shell.execute_reply":"2021-08-24T07:49:48.561378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\nfor i in range(1):\n    model = Model()\n    model.to(device)\n    \n    checkpoint = torch.load(f\"best-model-{i}.pth\")\n    model.load_state_dict(checkpoint[\"model_state_dict\"])\n    model.eval()\n    \n    models.append(model)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:51:23.59617Z","iopub.execute_input":"2021-08-24T07:51:23.596614Z","iopub.status.idle":"2021-08-24T07:51:23.89214Z","shell.execute_reply.started":"2021-08-24T07:51:23.596582Z","shell.execute_reply":"2021-08-24T07:51:23.891046Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### prepare test data","metadata":{}},{"cell_type":"code","source":"class DataRetriever(torch_data.Dataset):\n    def __init__(self, paths):\n        self.paths = paths\n          \n    def __len__(self):\n        return len(self.paths)\n    \n    def __getitem__(self, index):\n        _id = self.paths[index]\n        patient_path = f\"../input/rsna-miccai-brain-tumor-radiogenomic-classification/test/{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        return {\"X\": torch.tensor(channels).float() ,\"id\": _id }","metadata":{"execution":{"iopub.status.busy":"2021-08-24T07:51:27.373557Z","iopub.execute_input":"2021-08-24T07:51:27.373916Z","iopub.status.idle":"2021-08-24T07:51:27.386581Z","shell.execute_reply.started":"2021-08-24T07:51:27.373866Z","shell.execute_reply":"2021-08-24T07:51:27.384826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = pd.read_csv(\"../input/rsna-miccai-brain-tumor-radiogenomic-classification/sample_submission.csv\")\n\ntest_data_retriever = DataRetriever(\n    submission[\"BraTS21ID\"].values,\n)\n\ntest_loader = torch_data.DataLoader(\n    test_data_retriever,\n    batch_size = 4,\n    shuffle = False,\n    num_workers = 8,\n)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:10:43.59964Z","iopub.execute_input":"2021-08-24T08:10:43.600028Z","iopub.status.idle":"2021-08-24T08:10:43.611448Z","shell.execute_reply.started":"2021-08-24T08:10:43.599997Z","shell.execute_reply":"2021-08-24T08:10:43.610106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Inference(test)","metadata":{}},{"cell_type":"code","source":"y_pred = []\nids = []\n\nfor e, batch in enumerate(test_loader):\n    print(f\"{e}/ {len(test_loader)}\", end=\"\\r\")\n    with torch.no_grad():\n        tmp_pred = np.zeros((batch[\"X\"].shape[0], ))\n        for model in models:\n            tmp_res = torch.sigmoid(model(batch[\"X\"].to(device))).cpu().numpy().squeeze()\n            tmp_pred += tmp_res\n        y_pred.extend(tmp_pred)\n        ids.extend(batch[\"id\"].numpy().tolist()) #?","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:10:54.826561Z","iopub.execute_input":"2021-08-24T08:10:54.8269Z","iopub.status.idle":"2021-08-24T08:11:07.243442Z","shell.execute_reply.started":"2021-08-24T08:10:54.82687Z","shell.execute_reply":"2021-08-24T08:11:07.242103Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### submission","metadata":{}},{"cell_type":"code","source":"### submission(sample)\nsubmission = pd.DataFrame({\"BraTS21ID\": ids, \"MGMT_value\": y_pred})\nsubmission.to_csv(\"submission.csv\", index=False)","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:14:56.272448Z","iopub.execute_input":"2021-08-24T08:14:56.272928Z","iopub.status.idle":"2021-08-24T08:14:56.290329Z","shell.execute_reply.started":"2021-08-24T08:14:56.272855Z","shell.execute_reply":"2021-08-24T08:14:56.288448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(5, 5))\nplt.hist(submission[\"MGMT_value\"]);","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:15:27.520128Z","iopub.execute_input":"2021-08-24T08:15:27.520553Z","iopub.status.idle":"2021-08-24T08:15:27.766041Z","shell.execute_reply.started":"2021-08-24T08:15:27.520521Z","shell.execute_reply":"2021-08-24T08:15:27.76497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission","metadata":{"execution":{"iopub.status.busy":"2021-08-24T08:15:32.168745Z","iopub.execute_input":"2021-08-24T08:15:32.169085Z","iopub.status.idle":"2021-08-24T08:15:32.186141Z","shell.execute_reply.started":"2021-08-24T08:15:32.169053Z","shell.execute_reply":"2021-08-24T08:15:32.184593Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}