{"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":"## pip Install","metadata":{"id":"zJcGEk-YFI27"}},{"cell_type":"code","source":"INTERNET = True\n\nif INTERNET == True:\n    !python --version\n\n    !pip install monai\n    !pip install -q segmentation_models_pytorch\n    !pip install warmup-scheduler","metadata":{"executionInfo":{"elapsed":32581,"status":"ok","timestamp":1676719433752,"user":{"displayName":"구링도구링","userId":"14752850242191720980"},"user_tz":-480},"id":"b-edbDx0SJAY","outputId":"6d4428dc-9412-4bc6-e74a-0cc50fe19420","tags":[],"execution":{"iopub.status.busy":"2023-04-11T12:28:57.544005Z","iopub.execute_input":"2023-04-11T12:28:57.544498Z","iopub.status.idle":"2023-04-11T12:29:42.005297Z","shell.execute_reply.started":"2023-04-11T12:28:57.544457Z","shell.execute_reply":"2023-04-11T12:29:42.003874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Imports","metadata":{"id":"Y1BSVIUxSJAc"}},{"cell_type":"code","source":"import os\nimport sys\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom PIL import Image\nimport cv2\nimport re\nimport gc\nfrom tqdm import tqdm\nfrom pprint import pprint\nimport math\n\nimport matplotlib.pyplot as plt\nfrom matplotlib.patches import Rectangle\n\nimport skimage.transform as skTrans\nfrom skimage import exposure\n\nimport albumentations as alb\nfrom albumentations.pytorch import ToTensorV2\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\nfrom warmup_scheduler import GradualWarmupScheduler\nfrom torch.optim.lr_scheduler import OneCycleLR\n\nimport tensorflow as tf\n\nfrom monai.transforms import Resize\nimport monai.transforms as transforms\n\nimport segmentation_models_pytorch as smp\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.model_selection import KFold, StratifiedKFold\n\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"id":"jGJ3hxE3SJAe","tags":[],"execution":{"iopub.status.busy":"2023-04-11T12:29:42.009055Z","iopub.execute_input":"2023-04-11T12:29:42.009471Z","iopub.status.idle":"2023-04-11T12:29:59.476159Z","shell.execute_reply.started":"2023-04-11T12:29:42.009425Z","shell.execute_reply":"2023-04-11T12:29:59.474909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{"id":"NCM9Y9XQSJAf"}},{"cell_type":"code","source":"SEED = 1927550\nIMG_SIZE = 512\nBATCH = 8\nEPOCH = 5\nCLASS = 9 # from 0 to 8\nWORK = 'kaggle'\nencoder_backbone = 'timm-efficientnet-b5'\nbest_acc = 0\nbest_loss = 1\ndevice = 'cuda' if torch.cuda.is_available() else 'cpu'\nfolds = 5\n\ntrainlosslog = []\ntrainacclog = []\ntrainpreclog = []\nvalidlosslog = []\nvalidacclog = []\nvalidpreclog = []\n\nseg_revert = [\n    '1.2.826.0.1.3680043.1363',\n    '1.2.826.0.1.3680043.20120',\n    '1.2.826.0.1.3680043.2243',\n    '1.2.826.0.1.3680043.24606',\n    '1.2.826.0.1.3680043.32071'\n]","metadata":{"id":"u7dgRAJ7ihtl","execution":{"iopub.status.busy":"2023-04-11T12:57:21.227824Z","iopub.execute_input":"2023-04-11T12:57:21.228289Z","iopub.status.idle":"2023-04-11T12:57:21.240213Z","shell.execute_reply.started":"2023-04-11T12:57:21.228244Z","shell.execute_reply":"2023-04-11T12:57:21.239159Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Preprocessing","metadata":{}},{"cell_type":"code","source":"base_path = '/kaggle/input'\nvert_df = pd.read_csv(f'{base_path}/sagittal-preprocess/vert_list.csv')\nvert_df['StudyInstanceUID'] = 0\n\nfor idx in range(len(vert_df)):\n    vert_id = vert_df.loc[idx]['id']\n    studyuid = vert_id.split('_')[0]\n    vert_df['StudyInstanceUID'][idx] = studyuid","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:29:59.559072Z","iopub.execute_input":"2023-04-11T12:29:59.560875Z","iopub.status.idle":"2023-04-11T12:30:16.782702Z","shell.execute_reply.started":"2023-04-11T12:29:59.56083Z","shell.execute_reply":"2023-04-11T12:30:16.78169Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vert_df","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:30:16.784215Z","iopub.execute_input":"2023-04-11T12:30:16.784697Z","iopub.status.idle":"2023-04-11T12:30:16.808823Z","shell.execute_reply.started":"2023-04-11T12:30:16.784656Z","shell.execute_reply":"2023-04-11T12:30:16.807836Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib import pyplot as plt\nimport nibabel as nib\nimport numpy as np\nfrom glob import glob\n\n# testing for axial viewed image with labels\nbase_path = '/kaggle/input'\nuid_id = '1.2.826.0.1.3680043.32071'\nid = 150\n\nif uid_id in seg_revert:\n    length = len(glob(f'/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images/{uid_id}/*'))\n    length = length - id - 1\n    new_uid_id = f'{uid_id}_{length}'\n\nimg = np.load(f'{base_path}/3-channel-preprocessed-dataset/prep_train/{new_uid_id}.npz')['arr_0']\nmask = nib.load(f'{base_path}/rsna-2022-cervical-spine-fracture-detection/segmentations/{uid_id}.nii')\nmask = mask.get_fdata()\nmask = mask[:, ::-1, ::-1].transpose(2, 1, 0)[id]\nprint(mask)\nprint(np.unique(mask))\nprint(mask.shape)\n\nplt.imshow(img[:,:,0], cmap='bone')\nplt.imshow(mask, interpolation='nearest', alpha=0.5)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:57:26.441696Z","iopub.execute_input":"2023-04-11T12:57:26.44206Z","iopub.status.idle":"2023-04-11T12:57:27.604191Z","shell.execute_reply.started":"2023-04-11T12:57:26.442028Z","shell.execute_reply":"2023-04-11T12:57:27.603189Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"mask = np.where(mask==0, 1, 0)\nplt.imshow(mask, interpolation='nearest')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:42:43.6936Z","iopub.execute_input":"2023-04-11T12:42:43.69479Z","iopub.status.idle":"2023-04-11T12:42:43.901558Z","shell.execute_reply.started":"2023-04-11T12:42:43.694739Z","shell.execute_reply":"2023-04-11T12:42:43.900483Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Drop Bad Scans","metadata":{"id":"Ak8tWMpcSJAi"}},{"cell_type":"markdown","source":"https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/344862\n\n1. 1.2.826.0.1.3680043.20574: does not include a full cervical spine and should be ignored.\n1. 1.2.826.0.1.3680043.29952: the slices are duplicated, meaning that there are 2 scans stiched to each other.","metadata":{"id":"nx4jxfKXSJAi"}},{"cell_type":"code","source":"bad_scans = ['1.2.826.0.1.3680043.20574','1.2.826.0.1.3680043.29952']\n\nfor uid in bad_scans:\n    vert_df.drop(vert_df[vert_df['StudyInstanceUID']==uid].index, axis=0, inplace=True)\n\nvert_df.reset_index(drop=True)","metadata":{"id":"iQ02jh7wSJAi","tags":[],"execution":{"iopub.status.busy":"2023-04-11T12:30:19.679599Z","iopub.execute_input":"2023-04-11T12:30:19.679931Z","iopub.status.idle":"2023-04-11T12:30:19.708354Z","shell.execute_reply.started":"2023-04-11T12:30:19.679902Z","shell.execute_reply":"2023-04-11T12:30:19.706991Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Useful Functions","metadata":{"id":"ZnIzcpx8SJAk"}},{"cell_type":"code","source":"# apply seed\ndef seed_everything(seed):\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:30:19.710234Z","iopub.execute_input":"2023-04-11T12:30:19.710647Z","iopub.status.idle":"2023-04-11T12:30:19.716973Z","shell.execute_reply.started":"2023-04-11T12:30:19.710608Z","shell.execute_reply":"2023-04-11T12:30:19.715403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dataloader_creator(df_train, df_valid):\n    train_dataset = CustomDataset(df=df_train, transform=data_transforms['train'], test=False)\n    valid_dataset = CustomDataset(df=df_valid, transform=data_transforms['valid'], test=False)\n    train_loader = DataLoader(train_dataset, batch_size=BATCH, shuffle=True)\n    valid_loader = DataLoader(valid_dataset, batch_size=BATCH, shuffle=True)\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:30:19.723195Z","iopub.execute_input":"2023-04-11T12:30:19.723615Z","iopub.status.idle":"2023-04-11T12:30:19.729801Z","shell.execute_reply.started":"2023-04-11T12:30:19.723586Z","shell.execute_reply":"2023-04-11T12:30:19.728586Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class GradualWarmupSchedulerV3(GradualWarmupScheduler):\n    def __init__(self, optimizer, multiplier, total_epoch, after_scheduler=None):\n        super(GradualWarmupSchedulerV3, self).__init__(optimizer, multiplier, total_epoch, after_scheduler)\n    def get_lr(self):\n        if self.last_epoch >= self.total_epoch:\n            if self.after_scheduler:\n                if not self.finished:\n                    self.after_scheduler.base_lrs = [base_lr * self.multiplier for base_lr in self.base_lrs]\n                    self.finished = True\n                return self.after_scheduler.get_lr()\n            return [base_lr * self.multiplier for base_lr in self.base_lrs]\n        if self.multiplier == 1.0:\n            return [base_lr * (float(self.last_epoch) / self.total_epoch) for base_lr in self.base_lrs]\n        else:\n            return [base_lr * ((self.multiplier - 1.) * self.last_epoch / self.total_epoch + 1.) for base_lr in self.base_lrs]","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:30:19.731801Z","iopub.execute_input":"2023-04-11T12:30:19.732305Z","iopub.status.idle":"2023-04-11T12:30:19.744247Z","shell.execute_reply.started":"2023-04-11T12:30:19.732269Z","shell.execute_reply":"2023-04-11T12:30:19.74296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Image Transform","metadata":{}},{"cell_type":"code","source":"data_transforms = {\n    'train': alb.Compose([\n                alb.HorizontalFlip(p=0.5),\n                alb.VerticalFlip(p=0.5),\n                alb.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.05, rotate_limit=10, p=0.5),\n                alb.OneOf([\n                    alb.GridDistortion(num_steps=5, distort_limit=0.05, p=1.0),\n                    alb.OpticalDistortion(distort_limit=0.05, shift_limit=0.05, p=1.0),\n                    alb.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=1.0)\n                ], p=0.25),\n                alb.CoarseDropout(max_holes=8, max_height=IMG_SIZE//20, max_width=IMG_SIZE//20, min_holes=5, fill_value=0, mask_fill_value=0, p=0.5),\n             ]),\n    'valid': alb.Compose([])\n}","metadata":{"id":"xoFZZMAgSJAk","tags":[],"execution":{"iopub.status.busy":"2023-04-11T12:30:19.745803Z","iopub.execute_input":"2023-04-11T12:30:19.746706Z","iopub.status.idle":"2023-04-11T12:30:19.758362Z","shell.execute_reply.started":"2023-04-11T12:30:19.746662Z","shell.execute_reply":"2023-04-11T12:30:19.756969Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset","metadata":{"id":"OUh6ciuGSJAk"}},{"cell_type":"markdown","source":"reverted segmentation\n\nhttps://www.kaggle.com/code/itsuki9180/a-segmentation-is-in-reverse-order","metadata":{"id":"j-NlTTBHSJAk"}},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, df=vert_df, transform=None, test=False):\n        super().__init__()\n        self.df        = df\n        self.transform = transform\n        self.test      = test\n    \n    \n    def __getitem__(self, idx):\n        UID = self.df.iloc[idx]\n        uid_id = UID['id'] # 1.2.826.0.1.3680043.{studyuid}_{id}\n        studyuid = UID['StudyInstanceUID']\n        \n        new_uid_id = uid_id\n        if studyuid in seg_revert:\n            length = len(glob(f'/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images/{studyuid}/*'))\n            length = length - int(uid_id.split('_')[1]) - 1\n            new_uid_id = f'{studyuid}_{length}'\n        \n        image = np.load(f'{base_path}/3-channel-preprocessed-dataset/prep_train/{new_uid_id}.npz')['arr_0'] # 512 x 512 x 3\n\n        if self.test == True:\n            if self.transform is not None:\n                trans = self.transform(image=image)\n                image = trans['image']\n                image = np.transpose(image, (2, 0, 1))\n            return UID['id'], torch.from_numpy(np.array(image/255.0, dtype=np.float32)).float()\n        \n        mask_path = f'{base_path}/preprocess-3-channel/{uid_id}.npz' # 512 x 512\n        mask = np.load(mask_path)['arr_0'] # already rotated to sagittal view\n        \n        if self.transform is not None:\n            trans = self.transform(image=image, mask=mask)\n            image = trans['image']\n            mask = trans['mask']\n        \n        # image alignment: (channel, width, height)\n        image = np.transpose(image, (2, 0, 1))\n        \n        real_mask = []\n        # change all vertebraes from T1-T12 to be located in channel 9\n        mask = np.where(mask > 8, 8, mask)\n        \n        # extract which class this image is located,\n        # and place the image at that channel\n        # fill other channels with zeros\n        for channel in range(CLASS):\n            if channel in list(np.unique(mask)): real_mask.append(np.where(mask==channel, 1, 0))\n            else: real_mask.append(np.zeros((IMG_SIZE, IMG_SIZE)))\n        \n        # https://github.com/qubvel/segmentation_models/issues/403\n        train, seg = torch.from_numpy(np.array(image/255.0, dtype=np.float32)).float(), torch.from_numpy(np.array(real_mask, dtype=np.float32)).float()\n        del(real_mask)\n        return train, seg\n    \n    \n    def __len__(self):\n        return self.df.shape[0]","metadata":{"id":"IQmZJWkCSJAl","tags":[],"execution":{"iopub.status.busy":"2023-04-11T12:41:53.380702Z","iopub.execute_input":"2023-04-11T12:41:53.381182Z","iopub.status.idle":"2023-04-11T12:41:53.40666Z","shell.execute_reply.started":"2023-04-11T12:41:53.38114Z","shell.execute_reply":"2023-04-11T12:41:53.40545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{"id":"GbYm4AYqSJAl"}},{"cell_type":"code","source":"# important to have a bigger backbone\n# EfficientNet-B5 backbone + UNet decoder\nclass SegmentationModel(nn.Module):\n    def __init__(self):\n        super(SegmentationModel, self).__init__()\n        self.segmodel = smp.Unet(              # cnn\n            encoder_backbone,                  # efficientnet-b5 backbone\n            encoder_weights='imagenet',        # pretrained-weight = imagenet\n            in_channels=3,                     # in channels of 3\n            classes=CLASS,                     # output channel be 7 from C1 to C7\n            activation=None,\n        )\n        \n    def forward(self, x):\n        return self.segmodel(x)","metadata":{"id":"AWxyCaEJSJAm","tags":[],"execution":{"iopub.status.busy":"2023-04-11T12:36:35.337813Z","iopub.execute_input":"2023-04-11T12:36:35.338188Z","iopub.status.idle":"2023-04-11T12:36:35.347051Z","shell.execute_reply.started":"2023-04-11T12:36:35.338154Z","shell.execute_reply":"2023-04-11T12:36:35.344585Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Loss Function","metadata":{"id":"E0mwH5XvXVT-"}},{"cell_type":"code","source":"def bce_logits(y_pred, y_true):\n    loss = smp.losses.SoftBCEWithLogitsLoss()\n    return loss(y_pred, y_true)","metadata":{"id":"ghQPh0ZCSJAn","tags":[],"execution":{"iopub.status.busy":"2023-04-11T01:34:49.744565Z","iopub.execute_input":"2023-04-11T01:34:49.745242Z","iopub.status.idle":"2023-04-11T01:34:49.753081Z","shell.execute_reply.started":"2023-04-11T01:34:49.745211Z","shell.execute_reply":"2023-04-11T01:34:49.752226Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def iou_coef(y_pred, y_true):\n    jaccard = smp.losses.JaccardLoss(mode='multilabel')\n    return jaccard(y_pred, y_true)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:34:49.754521Z","iopub.execute_input":"2023-04-11T01:34:49.75526Z","iopub.status.idle":"2023-04-11T01:34:49.763745Z","shell.execute_reply.started":"2023-04-11T01:34:49.755217Z","shell.execute_reply":"2023-04-11T01:34:49.762655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def criterion(y_pred, y_true):\n    return bce_logits(y_pred, y_true)*0.5 + iou_coef(y_pred, y_true)*0.5","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:34:49.765314Z","iopub.execute_input":"2023-04-11T01:34:49.765842Z","iopub.status.idle":"2023-04-11T01:34:49.773414Z","shell.execute_reply.started":"2023-04-11T01:34:49.765806Z","shell.execute_reply":"2023-04-11T01:34:49.77213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dice_coef(mask1, mask2, epsilon=1e-7):\n    mask1 = mask1.detach().cpu().numpy()\n    mask2 = mask2.detach().cpu().numpy()\n    mask1 = np.where(mask2>0.5, 1, 0)\n    \n    intersect = np.sum(mask1*mask2)\n    fsum = np.sum(mask1)\n    ssum = np.sum(mask2)\n    dice = (2 * intersect + epsilon) / (fsum + ssum + epsilon)\n    dice = np.mean(dice)\n    \n    return dice","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:34:49.775539Z","iopub.execute_input":"2023-04-11T01:34:49.776139Z","iopub.status.idle":"2023-04-11T01:34:49.78446Z","shell.execute_reply.started":"2023-04-11T01:34:49.776097Z","shell.execute_reply":"2023-04-11T01:34:49.783255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Settings","metadata":{}},{"cell_type":"code","source":"# model setting\nmodel = SegmentationModel()\nmodel.to(device)\n\n# optimizer, scheduler setting\noptimizer = optim.AdamW(model.parameters(), lr=5e-4, weight_decay=0)\nscheduler = CosineAnnealingLR(optimizer, T_max=EPOCH-1, eta_min=1e-6, last_epoch=-1)\nscheduler_warmup = GradualWarmupSchedulerV3(optimizer, multiplier=10, total_epoch=1, after_scheduler=scheduler)","metadata":{"executionInfo":{"elapsed":31031,"status":"ok","timestamp":1676719477794,"user":{"displayName":"구링도구링","userId":"14752850242191720980"},"user_tz":-480},"id":"G6FDrK75Mn6p","outputId":"50d72650-25be-4d53-a623-0da14aad64db","execution":{"iopub.status.busy":"2023-04-11T01:34:49.786502Z","iopub.execute_input":"2023-04-11T01:34:49.786961Z","iopub.status.idle":"2023-04-11T01:35:03.138128Z","shell.execute_reply.started":"2023-04-11T01:34:49.786927Z","shell.execute_reply":"2023-04-11T01:35:03.137073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train and Validation Functions","metadata":{}},{"cell_type":"code","source":"def train(model, dataloader, optimizer):\n    model.train()\n    scaler = GradScaler()\n    scheduler = OneCycleLR(optimizer, max_lr=0.001, epochs=1, steps_per_epoch=len(df_train), pct_start=0.3)\n    \n    train_loss = []\n    train_acc = []\n    train_prec = []\n    \n    for idx, (imgs, masks) in enumerate(tqdm(dataloader)):\n        # set the gradient to 0 at initial\n        optimizer.zero_grad()\n\n        # forward data, making sure the data and model are on the same device\n        with autocast(enabled=True):\n            logits = model(imgs.to(device))\n            loss = criterion(logits, masks.to(device))\n        \n        logits = logits.sigmoid()\n        acc = ((logits>0.5) == masks.to(device)).float().mean()\n        prec = dice_coef(logits, masks.to(device))\n        \n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n        scheduler.step()\n\n        train_loss.append(loss.item())\n        train_acc.append(acc)\n        train_prec.append(prec)\n    \n    train_loss = sum(train_loss) / len(train_loss)\n    train_acc = sum(train_acc) / len(train_acc)\n    train_prec = sum(train_prec) / len(train_prec)\n    \n    return train_loss, train_acc, train_prec","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:35:03.139873Z","iopub.execute_input":"2023-04-11T01:35:03.140581Z","iopub.status.idle":"2023-04-11T01:35:03.151495Z","shell.execute_reply.started":"2023-04-11T01:35:03.14054Z","shell.execute_reply":"2023-04-11T01:35:03.150177Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def validation(model, dataloader, optimizer):\n    model.eval()\n    \n    valid_loss = []\n    valid_acc = []\n    valid_prec = []\n    \n    for imgs, masks in tqdm(dataloader):\n        # no need gradient in validation\n        # use torch.no_grad() accelerates the forward process\n        with torch.no_grad():\n            logits = model(imgs.to(device))\n\n        loss = criterion(logits, masks.to(device))\n        \n        logits = logits.sigmoid()\n        acc = ((logits>0.5) == masks.to(device)).float().mean()\n        prec = dice_coef(logits, masks.to(device))#.detach().cpu().numpy()\n        \n        valid_loss.append(loss.item())\n        valid_acc.append(acc)\n        valid_prec.append(prec)\n    \n    valid_loss = sum(valid_loss) / len(valid_loss)\n    valid_acc = sum(valid_acc) / len(valid_acc)\n    valid_prec = sum(valid_prec) / len(valid_prec)\n    \n    return valid_loss, valid_acc, valid_prec","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:35:03.153337Z","iopub.execute_input":"2023-04-11T01:35:03.153774Z","iopub.status.idle":"2023-04-11T01:35:03.165441Z","shell.execute_reply.started":"2023-04-11T01:35:03.153733Z","shell.execute_reply":"2023-04-11T01:35:03.164355Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train","metadata":{}},{"cell_type":"code","source":"seg_df = pd.read_csv(f'{base_path}/3-channel-preprocessed-dataset/train_df.csv')\ndf = vert_df[np.isin(vert_df['StudyInstanceUID'], seg_df['StudyInstanceUID'])].reset_index(drop=True)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:30:19.784119Z","iopub.execute_input":"2023-04-11T12:30:19.784449Z","iopub.status.idle":"2023-04-11T12:30:33.908916Z","shell.execute_reply.started":"2023-04-11T12:30:19.784419Z","shell.execute_reply":"2023-04-11T12:30:33.907744Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:30:33.910552Z","iopub.execute_input":"2023-04-11T12:30:33.910924Z","iopub.status.idle":"2023-04-11T12:30:33.928861Z","shell.execute_reply.started":"2023-04-11T12:30:33.910884Z","shell.execute_reply":"2023-04-11T12:30:33.927559Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"kf = KFold(folds)\nfor fold_idx, (t_idx, val_idx) in enumerate(kf.split(df, df)):\n    df.loc[val_idx, 'sub_fold'] = fold_idx\n\ndf_train = df[df['sub_fold'] != fold_idx].reset_index(drop=True)\ndf_valid = df[df['sub_fold'] == fold_idx].reset_index(drop=True)\ntrain_loader, valid_loader = dataloader_creator(df_train, df_valid)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:42:00.074354Z","iopub.execute_input":"2023-04-11T12:42:00.075346Z","iopub.status.idle":"2023-04-11T12:42:00.094844Z","shell.execute_reply.started":"2023-04-11T12:42:00.075305Z","shell.execute_reply":"2023-04-11T12:42:00.093714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"seed_everything(SEED)\n\n# initialize the best values to save\nbest_train_loss = 0\nbest_train_acc = 0\nbest_train_prec = 0\nbest_valid_loss = 0\nbest_valid_acc = 0\nbest_valid_prec = 0\n\nbest_epoch = 0\nearly_stop_count = 0\n\n# start training\nfor epoch in range(EPOCH):\n    print(f'### Epoch: {epoch+1} ###')\n    train_loss, train_acc, train_prec = train(model, train_loader, optimizer)\n    print(f'[ Train | {epoch + 1:03d}/{EPOCH:03d} ] loss = {train_loss:.5f}, acc = {train_acc:.5f}, prec = {train_prec:.5f}')\n    valid_loss, valid_acc, valid_prec = validation(model, valid_loader, optimizer)\n    print(f'[ Valid | {epoch + 1:03d}/{EPOCH:03d} ] loss = {valid_loss:.5f}, acc = {valid_acc:.5f}, prec = {valid_prec:.5f}')\n    print()\n\n    scheduler_warmup.step()\n    \n    # save train and valid logs\n    trainlosslog.append(train_loss)\n    trainacclog.append(train_acc.cpu().data.numpy())\n    trainpreclog.append(train_prec)\n    validlosslog.append(valid_loss)\n    validacclog.append(valid_acc.cpu().data.numpy())\n    validpreclog.append(valid_prec)\n\n    # save models\n    if valid_loss < best_loss:\n        # save highest values\n        best_train_loss, best_train_acc, best_train_prec = train_loss, train_acc, train_prec\n        best_valid_loss, best_valid_acc, best_valid_prec = valid_loss, valid_acc, valid_prec\n        best_epoch = epoch\n\n        # save model\n        torch.save(model.state_dict(), \"stage1_sagittal_best.ckpt\") # only save best to prevent output memory exceed error\n        # reset values\n        best_loss = valid_loss\n        early_stop_count = 0\n\n    if early_stop_count > 5:\n        print('Preformance not increasing. Early Stopping...')\n        break\n\n    early_stop_count = early_stop_count + 1\n\nprint()\nprint(f\"[ Best Train | {best_epoch+1:03d} / {EPOCH:03d} ] loss = {best_train_loss:.5f}, acc = {best_train_acc:.5f}, prec = {best_train_prec:.5f}\")\nprint(f\"[ Best Valid | {best_epoch+1:03d} / {EPOCH:03d} ] loss = {best_valid_loss:.5f}, acc = {best_valid_acc:.5f}, prec = {best_valid_prec:.5f}\")","metadata":{"execution":{"iopub.status.busy":"2023-04-11T01:35:16.631034Z","iopub.execute_input":"2023-04-11T01:35:16.631403Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"log_df = pd.DataFrame(\n    {\n        'trainLoss': trainlosslog,\n        'trainAcc': trainacclog,\n        'trainPrec': trainpreclog,\n        'validLoss': validlosslog,\n        'validAcc': validacclog,\n        'validPrec': validpreclog\n    })\n\nlog_df.to_csv('Logs.csv')","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Prediction","metadata":{}},{"cell_type":"code","source":"model = SegmentationModel()\nbest_model = model.to(device)\nbest_model.load_state_dict(torch.load('/kaggle/working/stage1_sagittal_best.ckpt'))\n\ndf_train = df[df['sub_fold'] != fold_idx].reset_index(drop=True)\ndf_valid = df[df['sub_fold'] == fold_idx].reset_index(drop=True)\ntest_dataset = CustomDataset(df=df_valid, transform=data_transforms['valid'], test=True)\ntest_loader = DataLoader(test_dataset, batch_size=5, shuffle=False, pin_memory=True)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T11:54:25.021757Z","iopub.execute_input":"2023-04-11T11:54:25.022282Z","iopub.status.idle":"2023-04-11T11:54:25.986888Z","shell.execute_reply.started":"2023-04-11T11:54:25.022226Z","shell.execute_reply":"2023-04-11T11:54:25.985827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualization","metadata":{}},{"cell_type":"code","source":"def show_img(img, mask=None):\n    plt.imshow(img[:,:,0], cmap='bone')\n\n    if mask is not None:\n        mask = np.argmax(mask, axis=2)\n        plt.imshow(mask, alpha=0.5)\n    else:\n        print('No mask')\n        \n    plt.axis('off')","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:30:33.957055Z","iopub.execute_input":"2023-04-11T12:30:33.957457Z","iopub.status.idle":"2023-04-11T12:30:33.964076Z","shell.execute_reply.started":"2023-04-11T12:30:33.957416Z","shell.execute_reply":"2023-04-11T12:30:33.962764Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_batch(imgs, msks, size=3):\n    plt.figure(figsize=(5*5, 5))\n    for idx in range(size):\n        plt.subplot(1, 5, idx+1)\n        img = imgs[idx,].permute((1, 2, 0)).numpy()*255.0\n        img = img.astype('uint8')\n        msk = msks[idx,].permute((1, 2, 0)).numpy()*255.0\n        show_img(img, msk)\n    plt.tight_layout()\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-11T12:30:33.965988Z","iopub.execute_input":"2023-04-11T12:30:33.966363Z","iopub.status.idle":"2023-04-11T12:30:33.97501Z","shell.execute_reply.started":"2023-04-11T12:30:33.966326Z","shell.execute_reply":"2023-04-11T12:30:33.973783Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_imgs, val_real_mask = next(iter(valid_loader))\n\nwith torch.no_grad():\n    val_msks = best_model(val_imgs.to(device))\n    val_msks = (nn.Sigmoid()(val_msks)>0.5).double()\n    \nplot_batch(val_imgs, val_msks.detach().cpu(), size=5)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T11:54:30.829491Z","iopub.execute_input":"2023-04-11T11:54:30.830393Z","iopub.status.idle":"2023-04-11T11:54:33.050718Z","shell.execute_reply.started":"2023-04-11T11:54:30.830344Z","shell.execute_reply":"2023-04-11T11:54:33.049773Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_batch(val_imgs, val_real_mask, size=5)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T11:54:33.052621Z","iopub.execute_input":"2023-04-11T11:54:33.0533Z","iopub.status.idle":"2023-04-11T11:54:34.417114Z","shell.execute_reply.started":"2023-04-11T11:54:33.053262Z","shell.execute_reply":"2023-04-11T11:54:34.416106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"real_mask_channel = []\npredict_mask_channel = []\n\nfor i in range(5):\n    predict_mask_channel.append(np.unique(np.argmax(val_msks.detach().cpu().numpy()[i], axis=0)))\n    real_mask_channel.append(np.unique(np.argmax(val_real_mask[i], axis=0)))\n\nprint(predict_mask_channel)\nprint(real_mask_channel)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T11:54:34.418603Z","iopub.execute_input":"2023-04-11T11:54:34.426506Z","iopub.status.idle":"2023-04-11T11:54:35.365721Z","shell.execute_reply.started":"2023-04-11T11:54:34.426465Z","shell.execute_reply":"2023-04-11T11:54:35.364545Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"imgs, real_mask = next(iter(train_loader))\nwith torch.no_grad():\n    msks = best_model(imgs.to(device))\n    msks = (nn.Sigmoid()(msks)>0.5).double()\n    \nplot_batch(imgs, msks.detach().cpu(), size=5)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T11:54:35.369654Z","iopub.execute_input":"2023-04-11T11:54:35.369947Z","iopub.status.idle":"2023-04-11T11:54:38.060667Z","shell.execute_reply.started":"2023-04-11T11:54:35.369919Z","shell.execute_reply":"2023-04-11T11:54:38.058788Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_batch(imgs, real_mask, size=5)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T11:54:38.062559Z","iopub.execute_input":"2023-04-11T11:54:38.063152Z","iopub.status.idle":"2023-04-11T11:54:39.218116Z","shell.execute_reply.started":"2023-04-11T11:54:38.063114Z","shell.execute_reply":"2023-04-11T11:54:39.217222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_real_mask_channel = []\ntrain_predict_mask_channel = []\n\nfor i in range(5):\n    train_predict_mask_channel.append(np.unique(np.argmax(msks.detach().cpu().numpy()[i], axis=0)))\n    train_real_mask_channel.append(np.unique(np.argmax(real_mask[i], axis=0)))\n\nprint(train_predict_mask_channel)\nprint(train_real_mask_channel)","metadata":{"execution":{"iopub.status.busy":"2023-04-11T11:54:39.219584Z","iopub.execute_input":"2023-04-11T11:54:39.220178Z","iopub.status.idle":"2023-04-11T11:54:40.003031Z","shell.execute_reply.started":"2023-04-11T11:54:39.220138Z","shell.execute_reply":"2023-04-11T11:54:40.00184Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}