{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":29653,"databundleVersionId":2420395,"sourceType":"competition"}],"dockerImageVersionId":30699,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"NUM_IMAGES_3D = 64\nTRAINING_BATCH_SIZE = 4\nTEST_BATCH_SIZE = 4\nIMAGE_SIZE = 256\nN_EPOCHS = 15\ndo_valid = True\nn_workers = 4\nACCUMULATION_STEPS = 4","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Utils file","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport cv2","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom_image(path, img_size= IMAGE_SIZE, voi_lut=True, rotate=0):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    if voi_lut:\n        data = apply_voi_lut(dicom.pixel_array, dicom)\n    else:\n        data = dicom.pixel_array\n\n    if rotate > 0:\n        rot_choices = [\n            0,\n            cv2.ROTATE_90_CLOCKWISE,\n            cv2.ROTATE_90_COUNTERCLOCKWISE,\n            cv2.ROTATE_180,\n        ]\n        data = cv2.rotate(data, rot_choices[rotate])\n\n    data = cv2.resize(data, (img_size, img_size))\n    data = data - np.min(data)\n    if np.min(data) < np.max(data):\n        data = data / np.max(data)\n    return data\n\n\ndef crop_img(img):\n    rows = np.any(img, axis=1)\n    cols = np.any(img, axis=0)\n    c1, c2 = False, False\n    try:\n        rmin, rmax = np.where(rows)[0][[0, -1]]\n    except:\n        rmin, rmax = 0, img.shape[0]\n        c1 = True\n\n    try:\n        cmin, cmax = np.where(cols)[0][[0, -1]]\n    except:\n        cmin, cmax = 0, img.shape[1]\n        c2 = True\n    bb = (rmin, rmax, cmin, cmax)\n    \n    if c1 and c2:\n        return img[0:0, 0:0]\n    else:\n        return img[bb[0] : bb[1], bb[2] : bb[3]]\n\n\ndef extract_cropped_image_size(path):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    cropped_data = crop_img(data)\n    resolution = cropped_data.shape[0]*cropped_data.shape[1]  \n    return resolution","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset file","metadata":{}},{"cell_type":"code","source":"import glob\nimport os\nimport re\n\nimport joblib\nimport numpy as np\nimport torch\nfrom torch.utils.data import Dataset\nfrom tqdm import tqdm","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BrainRSNADataset(Dataset):\n    def __init__(\n        self, data, transform=None, target=\"MGMT_value\", mri_type=\"FLAIR\", is_train=True, ds_type=\"forgot\", do_load=True\n    ):\n        self.target = target\n        self.data = data\n        self.type = mri_type\n\n        self.transform = transform\n        self.is_train = is_train\n        self.folder = \"train\" if self.is_train else \"test\"\n        self.do_load = do_load\n        self.ds_type = ds_type\n        self.img_indexes = self._prepare_biggest_images()\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, index):\n        row = self.data.loc[index]\n        case_id = int(row.BraTS21ID)\n        target = int(row[self.target])\n        _3d_images = self.load_dicom_images_3d(case_id)\n        _3d_images = torch.tensor(_3d_images).float()\n        if self.is_train:\n            return {\"image\": _3d_images, \"target\": target, \"case_id\": case_id}\n        else:\n            return {\"image\": _3d_images, \"case_id\": case_id}\n\n    def _prepare_biggest_images(self):\n        big_image_indexes = {}\n        if (f\"big_image_indexes_{self.ds_type}.pkl\" in os.listdir(\"/kaggle/working/\")) and (self.do_load) :\n            print(\"Loading the best images indexes for all the cases...\")\n            big_image_indexes = joblib.load(f\"/kaggle/working/big_image_indexes_{self.ds_type}.pkl\")\n            return big_image_indexes\n        else:\n            \n            print(\"Caulculating the best scans for every case...\")\n            for row in tqdm(self.data.iterrows(), total=len(self.data)):\n                case_id = str(int(row[1].BraTS21ID)).zfill(5)\n                path = f\"/kaggle/input/rsna-miccai-brain-tumor-radiogenomic-classification/{self.folder}/{case_id}/{self.type}/*.dcm\"\n                files = sorted(\n                    glob.glob(path),\n                    key=lambda var: [\n                        int(x) if x.isdigit() else x for x in re.findall(r\"[^0-9]|[0-9]+\", var)\n                    ],\n                )\n                resolutions = [extract_cropped_image_size(f) for f in files]\n                middle = np.array(resolutions).argmax()\n                big_image_indexes[case_id] = middle\n\n            joblib.dump(big_image_indexes, f\"/kaggle/working/big_image_indexes_{self.ds_type}.pkl\")\n            return big_image_indexes\n\n\n\n    def load_dicom_images_3d(\n        self,\n        case_id,\n        num_imgs=NUM_IMAGES_3D,\n        img_size=IMAGE_SIZE,\n        rotate=0,\n    ):\n        case_id = str(case_id).zfill(5)\n\n        path = f\"/kaggle/input/rsna-miccai-brain-tumor-radiogenomic-classification/{self.folder}/{case_id}/{self.type}/*.dcm\"\n        files = sorted(\n            glob.glob(path),\n            key=lambda var: [\n                int(x) if x.isdigit() else x for x in re.findall(r\"[^0-9]|[0-9]+\", var)\n            ],\n        )\n\n        if self.is_train:\n            middle = self.img_indexes[case_id]\n        else:\n            middle = len(files) // 2\n\n        num_imgs2 = num_imgs // 2\n        p1 = max(0, middle - num_imgs2)\n        p2 = min(len(files), middle + num_imgs2)\n        image_stack = [load_dicom_image(f, rotate=rotate, voi_lut=True) for f in files[p1:p2]]\n        \n        img3d = np.stack(image_stack).T\n        if img3d.shape[-1] < num_imgs:\n            n_zero = np.zeros((img_size, img_size, num_imgs - img3d.shape[-1]))\n            img3d = np.concatenate((img3d, n_zero), axis=-1)\n\n        return np.expand_dims(img3d, 0)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## create fold","metadata":{}},{"cell_type":"code","source":"import argparse\nimport pandas as pd\nfrom sklearn.model_selection import StratifiedKFold\n\n# Argument parsing\nparser = argparse.ArgumentParser()\nparser.add_argument(\"--n_folds\", default=5, type=int)\nargs, unknown = parser.parse_known_args()  # This will ignore unrecognized arguments\n\n# Load the CSV file\ntrain = pd.read_csv(\"/kaggle/input/rsna-miccai-brain-tumor-radiogenomic-classification/train_labels.csv\")\n\n# Create stratified folds\nskf = StratifiedKFold(n_splits=args.n_folds, shuffle=True, random_state=518)\ntarget = \"MGMT_value\"\n\nfor fold, (trn_idx, val_idx) in enumerate(skf.split(train, train[target])):\n    train.loc[val_idx, \"fold\"] = int(fold)\n\n# Save the modified DataFrame with fold information\ntrain.to_csv(\"/kaggle/working/train.csv\", index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# train file","metadata":{}},{"cell_type":"code","source":"!pip install monai","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import argparse\nimport os\n\nimport monai\nimport numpy as np\nimport pandas as pd\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom sklearn.metrics import roc_auc_score\nfrom torch.optim import lr_scheduler\nfrom tqdm import tqdm","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Ignore unknown arguments in Jupyter Notebook\nimport sys\nif 'ipykernel' in sys.modules:\n    sys.argv = ['']\n\nparser = argparse.ArgumentParser()\nparser.add_argument(\"--fold\", default=0, type=int)\nparser.add_argument(\"--type\", default=\"FLAIR\", type=str)\nparser.add_argument(\"--model_name\", default=\"b0\", type=str)\nargs = parser.parse_args()\n\ndata = pd.read_csv(\"/kaggle/working/train.csv\")\ntrain_df = data[data.fold != args.fold].reset_index(drop=False)\nval_df = data[data.fold == args.fold].reset_index(drop=False)\n\ndevice = torch.device(\"cuda\")\ndevice","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.cuda.amp as amp\nimport torch.utils.checkpoint as checkpoint","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def checkpoint_forward(module, x):\n    def custom_forward(*inputs):\n        return module(*inputs)\n    return checkpoint.checkpoint(custom_forward, x)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\n\n# Specify the folder name\nnew_folder_name = \"weights\"\n\n# Check if the folder already exists, and create it if it doesn't\nif not os.path.exists(new_folder_name):\n    os.makedirs(new_folder_name)\n    print(f\"Folder '{new_folder_name}' created successfully.\")\nelse:\n    print(f\"Folder '{new_folder_name}' already exists.\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(f\"train_{args.type}_{args.fold}\")\ntrain_dataset = BrainRSNADataset(data=train_df, mri_type=args.type, ds_type=f\"train_{args.type}_{args.fold}\")\nvalid_dataset = BrainRSNADataset(data=val_df, mri_type=args.type, ds_type=f\"val_{args.type}_{args.fold}\")\n\ntrain_dl = torch.utils.data.DataLoader(\n    train_dataset,\n    batch_size=TRAINING_BATCH_SIZE,\n    shuffle=True,\n    num_workers=n_workers,\n    drop_last=True,\n    pin_memory=True,\n)\n\nvalidation_dl = torch.utils.data.DataLoader(\n    valid_dataset,\n    batch_size=TEST_BATCH_SIZE,\n    shuffle=False,\n    num_workers=n_workers,\n    pin_memory=True,\n)\n\nmodel = monai.networks.nets.resnet10(spatial_dims=3, n_input_channels=1)\nmodel.fc = nn.Linear(model.fc.in_features, 1)\nmodel.to(device)\n\noptimizer = optim.Adam(model.parameters(), lr=0.0001)\nscheduler = lr_scheduler.MultiStepLR(optimizer, milestones=[10], gamma=0.5, verbose=True)\n\nmodel.zero_grad()\nbest_loss = 9999\nbest_auc = 0\ncriterion = nn.BCEWithLogitsLoss()\n\nfor counter in range(N_EPOCHS):\n    epoch_iterator_train = tqdm(train_dl)\n    tr_loss = 0.0\n    scaler = amp.GradScaler()\n    optimizer.zero_grad()\n\n    for step, batch in enumerate(epoch_iterator_train):\n        torch.cuda.empty_cache()  # Clear CUDA cache\n        model.train()\n        images, targets = batch[\"image\"].to(device), batch[\"target\"].to(device)\n\n        with amp.autocast():\n            # Apply checkpointing to specific parts of the model\n            x = checkpoint_forward(model.conv1, images)\n            x = checkpoint_forward(model.bn1, x)\n            x = model.act(x)\n            x = model.maxpool(x)\n\n            x = checkpoint_forward(model.layer1, x)\n            x = checkpoint_forward(model.layer2, x)\n            x = checkpoint_forward(model.layer3, x)\n            x = checkpoint_forward(model.layer4, x)\n\n            x = model.avgpool(x)\n            x = torch.flatten(x, 1)\n            outputs = model.fc(x)\n\n            targets = targets\n            loss = criterion(outputs.squeeze(1), targets.float())\n\n        scaler.scale(loss).backward()\n\n        if (step + 1) % ACCUMULATION_STEPS == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n\n        tr_loss += loss.item()\n        epoch_iterator_train.set_postfix(\n            batch_loss=(loss.item()), loss=(tr_loss / (step + 1))\n        )\n    scheduler.step()\n\n    if do_valid:\n        with torch.no_grad():\n            val_loss = 0.0\n            preds = []\n            true_labels = []\n            case_ids = []\n            epoch_iterator_val = tqdm(validation_dl)\n            for step, batch in enumerate(epoch_iterator_val):\n                torch.cuda.empty_cache()  # Clear CUDA cache\n                model.eval()\n                images, targets = batch[\"image\"].to(device), batch[\"target\"].to(device)\n\n                with amp.autocast():\n                    x = checkpoint_forward(model.conv1, images)\n                    x = checkpoint_forward(model.bn1, x)\n                    x = model.act(x)\n                    x = model.maxpool(x)\n\n                    x = checkpoint_forward(model.layer1, x)\n                    x = checkpoint_forward(model.layer2, x)\n                    x = checkpoint_forward(model.layer3, x)\n                    x = checkpoint_forward(model.layer4, x)\n\n                    x = model.avgpool(x)\n                    x = torch.flatten(x, 1)\n                    outputs = model.fc(x)\n\n                    targets = targets\n                    loss = criterion(outputs.squeeze(1), targets.float())\n                val_loss += loss.item()\n                epoch_iterator_val.set_postfix(\n                    batch_loss=(loss.item()), loss=(val_loss / (step + 1))\n                )\n                preds.append(outputs.sigmoid().detach().cpu().numpy())\n                true_labels.append(targets.cpu().numpy())\n                case_ids.append(batch[\"case_id\"])\n        preds = np.vstack(preds).T[0].tolist()\n        true_labels = np.hstack(true_labels).tolist()\n        case_ids = np.hstack(case_ids).tolist()\n        auc_score = roc_auc_score(true_labels, preds)\n        auc_score_adj_best = 0\n        for thresh in np.linspace(0, 1, 50):\n            auc_score_adj = roc_auc_score(true_labels, list(np.array(preds) > thresh))\n            if auc_score_adj > auc_score_adj_best:\n                best_thresh = thresh\n                auc_score_adj_best = auc_score_adj\n\n        print(\n            f\"EPOCH {counter}/{N_EPOCHS}: Validation average loss: {val_loss/(step+1)} + AUC SCORE = {auc_score} + AUC SCORE THRESH {best_thresh} = {auc_score_adj_best}\"\n        )\n        if auc_score > best_auc:\n            print(\"Saving the model...\")\n\n            all_files = os.listdir(\"/kaggle/working/weights/\")\n\n            for f in all_files:\n                if f\"{args.model_name}_{args.type}_fold{args.fold}\" in f:\n                    os.remove(f\"/kaggle/working/weights/{f}\")\n\n            best_auc = auc_score\n            torch.save(\n                model.state_dict(),\n                f\"/kaggle/working/weights/3d-{args.model_name}_{args.type}_fold{args.fold}_{round(best_auc,3)}.pth\",\n            )\n\nprint(best_auc)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}