{"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-24T08:16:45.83048Z","iopub.execute_input":"2023-04-24T08:16:45.830766Z","iopub.status.idle":"2023-04-24T08:17:25.786479Z","shell.execute_reply.started":"2023-04-24T08:16:45.830737Z","shell.execute_reply":"2023-04-24T08:17:25.785053Z"},"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-24T08:17:25.78917Z","iopub.execute_input":"2023-04-24T08:17:25.789896Z","iopub.status.idle":"2023-04-24T08:17:41.747671Z","shell.execute_reply.started":"2023-04-24T08:17:25.78985Z","shell.execute_reply":"2023-04-24T08:17:41.746554Z"},"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.99901\nbest_loss = 0.07875\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-24T08:17:41.749483Z","iopub.execute_input":"2023-04-24T08:17:41.750087Z","iopub.status.idle":"2023-04-24T08:17:41.818058Z","shell.execute_reply.started":"2023-04-24T08:17:41.750041Z","shell.execute_reply":"2023-04-24T08:17:41.815687Z"},"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-24T08:17:41.821287Z","iopub.execute_input":"2023-04-24T08:17:41.821584Z","iopub.status.idle":"2023-04-24T08:17:58.429446Z","shell.execute_reply.started":"2023-04-24T08:17:41.821555Z","shell.execute_reply":"2023-04-24T08:17:58.42838Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vert_df","metadata":{"execution":{"iopub.status.busy":"2023-04-24T08:17:58.431078Z","iopub.execute_input":"2023-04-24T08:17:58.431432Z","iopub.status.idle":"2023-04-24T08:17:58.450387Z","shell.execute_reply.started":"2023-04-24T08:17:58.431391Z","shell.execute_reply":"2023-04-24T08:17:58.449274Z"},"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-24T08:17:58.451894Z","iopub.execute_input":"2023-04-24T08:17:58.452335Z","iopub.status.idle":"2023-04-24T08:17:59.671477Z","shell.execute_reply.started":"2023-04-24T08:17:58.452297Z","shell.execute_reply":"2023-04-24T08:17:59.670435Z"},"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-24T08:17:59.672669Z","iopub.execute_input":"2023-04-24T08:17:59.673893Z","iopub.status.idle":"2023-04-24T08:17:59.875668Z","shell.execute_reply.started":"2023-04-24T08:17:59.67384Z","shell.execute_reply":"2023-04-24T08:17:59.874601Z"},"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-24T08:17:59.877016Z","iopub.execute_input":"2023-04-24T08:17:59.87861Z","iopub.status.idle":"2023-04-24T08:17:59.907419Z","shell.execute_reply.started":"2023-04-24T08:17:59.878566Z","shell.execute_reply":"2023-04-24T08:17:59.906549Z"},"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-24T08:17:59.909398Z","iopub.execute_input":"2023-04-24T08:17:59.910356Z","iopub.status.idle":"2023-04-24T08:17:59.916458Z","shell.execute_reply.started":"2023-04-24T08:17:59.910317Z","shell.execute_reply":"2023-04-24T08:17:59.915418Z"},"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-24T08:17:59.918595Z","iopub.execute_input":"2023-04-24T08:17:59.919485Z","iopub.status.idle":"2023-04-24T08:17:59.9276Z","shell.execute_reply.started":"2023-04-24T08:17:59.919445Z","shell.execute_reply":"2023-04-24T08:17:59.926598Z"},"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-24T08:17:59.933306Z","iopub.execute_input":"2023-04-24T08:17:59.933587Z","iopub.status.idle":"2023-04-24T08:17:59.942553Z","shell.execute_reply.started":"2023-04-24T08:17:59.93356Z","shell.execute_reply":"2023-04-24T08:17:59.941506Z"},"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-24T08:17:59.944301Z","iopub.execute_input":"2023-04-24T08:17:59.945184Z","iopub.status.idle":"2023-04-24T08:17:59.953883Z","shell.execute_reply.started":"2023-04-24T08:17:59.945144Z","shell.execute_reply":"2023-04-24T08:17:59.953127Z"},"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-24T08:17:59.955752Z","iopub.execute_input":"2023-04-24T08:17:59.956639Z","iopub.status.idle":"2023-04-24T08:17:59.971062Z","shell.execute_reply.started":"2023-04-24T08:17:59.956561Z","shell.execute_reply":"2023-04-24T08:17:59.970308Z"},"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-24T08:17:59.972592Z","iopub.execute_input":"2023-04-24T08:17:59.973335Z","iopub.status.idle":"2023-04-24T08:17:59.984099Z","shell.execute_reply.started":"2023-04-24T08:17:59.973296Z","shell.execute_reply":"2023-04-24T08:17:59.983126Z"},"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-24T08:17:59.987937Z","iopub.execute_input":"2023-04-24T08:17:59.988301Z","iopub.status.idle":"2023-04-24T08:17:59.994046Z","shell.execute_reply.started":"2023-04-24T08:17:59.988265Z","shell.execute_reply":"2023-04-24T08:17:59.992985Z"},"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-24T08:17:59.996024Z","iopub.execute_input":"2023-04-24T08:17:59.996744Z","iopub.status.idle":"2023-04-24T08:18:00.0029Z","shell.execute_reply.started":"2023-04-24T08:17:59.996704Z","shell.execute_reply":"2023-04-24T08:18:00.001694Z"},"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-24T08:18:00.00476Z","iopub.execute_input":"2023-04-24T08:18:00.00525Z","iopub.status.idle":"2023-04-24T08:18:00.013092Z","shell.execute_reply.started":"2023-04-24T08:18:00.005212Z","shell.execute_reply":"2023-04-24T08:18:00.012082Z"},"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-24T08:18:00.014656Z","iopub.execute_input":"2023-04-24T08:18:00.015141Z","iopub.status.idle":"2023-04-24T08:18:00.02475Z","shell.execute_reply.started":"2023-04-24T08:18:00.015102Z","shell.execute_reply":"2023-04-24T08:18:00.02369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Settings","metadata":{}},{"cell_type":"code","source":"# model setting\nmodel = SegmentationModel()\nmodel.to(device)\nmodel.load_state_dict(torch.load('/kaggle/input/stage1-window-multilabel-best/stage1_sagittal_best.ckpt'))\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-24T08:18:00.026499Z","iopub.execute_input":"2023-04-24T08:18:00.026886Z","iopub.status.idle":"2023-04-24T08:18:09.365034Z","shell.execute_reply.started":"2023-04-24T08:18:00.026848Z","shell.execute_reply":"2023-04-24T08:18:09.363983Z"},"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-24T08:18:09.366732Z","iopub.execute_input":"2023-04-24T08:18:09.367107Z","iopub.status.idle":"2023-04-24T08:18:09.377278Z","shell.execute_reply.started":"2023-04-24T08:18:09.367069Z","shell.execute_reply":"2023-04-24T08:18:09.375941Z"},"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-24T08:18:09.378663Z","iopub.execute_input":"2023-04-24T08:18:09.37978Z","iopub.status.idle":"2023-04-24T08:18:09.395846Z","shell.execute_reply.started":"2023-04-24T08:18:09.379738Z","shell.execute_reply":"2023-04-24T08:18:09.394777Z"},"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-24T08:18:09.399295Z","iopub.execute_input":"2023-04-24T08:18:09.400745Z","iopub.status.idle":"2023-04-24T08:18:23.419278Z","shell.execute_reply.started":"2023-04-24T08:18:09.400704Z","shell.execute_reply":"2023-04-24T08:18:23.418254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2023-04-24T08:18:23.420974Z","iopub.execute_input":"2023-04-24T08:18:23.421346Z","iopub.status.idle":"2023-04-24T08:18:23.436227Z","shell.execute_reply.started":"2023-04-24T08:18:23.421305Z","shell.execute_reply":"2023-04-24T08:18:23.435122Z"},"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-24T08:18:23.437931Z","iopub.execute_input":"2023-04-24T08:18:23.438559Z","iopub.status.idle":"2023-04-24T08:18:23.460074Z","shell.execute_reply.started":"2023-04-24T08:18:23.438519Z","shell.execute_reply":"2023-04-24T08:18:23.459093Z"},"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-24T08:18:23.46169Z","iopub.execute_input":"2023-04-24T08:18:23.462084Z"},"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-24T18:51:44.605397Z","iopub.execute_input":"2023-04-24T18:51:44.605874Z","iopub.status.idle":"2023-04-24T18:51:46.03794Z","shell.execute_reply.started":"2023-04-24T18:51:44.60583Z","shell.execute_reply":"2023-04-24T18:51:46.036635Z"},"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-24T18:51:49.528833Z","iopub.execute_input":"2023-04-24T18:51:49.529264Z","iopub.status.idle":"2023-04-24T18:51:49.547486Z","shell.execute_reply.started":"2023-04-24T18:51:49.529224Z","shell.execute_reply":"2023-04-24T18:51:49.541586Z"},"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-24T18:51:51.706557Z","iopub.execute_input":"2023-04-24T18:51:51.707571Z","iopub.status.idle":"2023-04-24T18:51:51.718534Z","shell.execute_reply.started":"2023-04-24T18:51:51.707531Z","shell.execute_reply":"2023-04-24T18:51:51.713805Z"},"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-24T18:51:53.84972Z","iopub.execute_input":"2023-04-24T18:51:53.850082Z","iopub.status.idle":"2023-04-24T18:51:56.196524Z","shell.execute_reply.started":"2023-04-24T18:51:53.85005Z","shell.execute_reply":"2023-04-24T18:51:56.195567Z"},"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-24T18:51:56.198928Z","iopub.execute_input":"2023-04-24T18:51:56.199568Z","iopub.status.idle":"2023-04-24T18:51:57.395864Z","shell.execute_reply.started":"2023-04-24T18:51:56.19953Z","shell.execute_reply":"2023-04-24T18:51:57.394468Z"},"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-24T18:52:06.056279Z","iopub.execute_input":"2023-04-24T18:52:06.056646Z","iopub.status.idle":"2023-04-24T18:52:06.817249Z","shell.execute_reply.started":"2023-04-24T18:52:06.056595Z","shell.execute_reply":"2023-04-24T18:52:06.815882Z"},"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-24T18:52:17.716045Z","iopub.execute_input":"2023-04-24T18:52:17.716441Z","iopub.status.idle":"2023-04-24T18:52:20.419516Z","shell.execute_reply.started":"2023-04-24T18:52:17.716405Z","shell.execute_reply":"2023-04-24T18:52:20.417948Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_batch(imgs, real_mask, size=5)","metadata":{"execution":{"iopub.status.busy":"2023-04-24T18:52:23.001224Z","iopub.execute_input":"2023-04-24T18:52:23.001837Z","iopub.status.idle":"2023-04-24T18:52:24.154899Z","shell.execute_reply.started":"2023-04-24T18:52:23.001796Z","shell.execute_reply":"2023-04-24T18:52:24.153397Z"},"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-24T18:52:29.143235Z","iopub.execute_input":"2023-04-24T18:52:29.143832Z","iopub.status.idle":"2023-04-24T18:52:29.904663Z","shell.execute_reply.started":"2023-04-24T18:52:29.143793Z","shell.execute_reply":"2023-04-24T18:52:29.903498Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}