{"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-19T06:05:32.996769Z","iopub.execute_input":"2023-04-19T06:05:32.997265Z","iopub.status.idle":"2023-04-19T06:06:15.377105Z","shell.execute_reply.started":"2023-04-19T06:05:32.997198Z","shell.execute_reply":"2023-04-19T06:06:15.375751Z"},"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-19T06:06:15.382025Z","iopub.execute_input":"2023-04-19T06:06:15.382428Z","iopub.status.idle":"2023-04-19T06:06:32.604609Z","shell.execute_reply.started":"2023-04-19T06:06:15.382386Z","shell.execute_reply":"2023-04-19T06:06:32.603306Z"},"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-19T06:06:32.606304Z","iopub.execute_input":"2023-04-19T06:06:32.606921Z","iopub.status.idle":"2023-04-19T06:06:32.689393Z","shell.execute_reply.started":"2023-04-19T06:06:32.60688Z","shell.execute_reply":"2023-04-19T06:06:32.688269Z"},"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-19T06:06:32.692599Z","iopub.execute_input":"2023-04-19T06:06:32.692879Z","iopub.status.idle":"2023-04-19T06:06:49.715652Z","shell.execute_reply.started":"2023-04-19T06:06:32.692852Z","shell.execute_reply":"2023-04-19T06:06:49.714447Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vert_df","metadata":{"execution":{"iopub.status.busy":"2023-04-19T06:06:49.717288Z","iopub.execute_input":"2023-04-19T06:06:49.718415Z","iopub.status.idle":"2023-04-19T06:06:49.741656Z","shell.execute_reply.started":"2023-04-19T06:06:49.718371Z","shell.execute_reply":"2023-04-19T06:06:49.740064Z"},"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-19T06:06:49.743977Z","iopub.execute_input":"2023-04-19T06:06:49.744349Z","iopub.status.idle":"2023-04-19T06:06:51.168974Z","shell.execute_reply.started":"2023-04-19T06:06:49.744313Z","shell.execute_reply":"2023-04-19T06:06:51.168024Z"},"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-19T06:06:51.170049Z","iopub.execute_input":"2023-04-19T06:06:51.171147Z","iopub.status.idle":"2023-04-19T06:06:51.384057Z","shell.execute_reply.started":"2023-04-19T06:06:51.171101Z","shell.execute_reply":"2023-04-19T06:06:51.382934Z"},"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.\n3. (additional) 1.2.826.0.1.3680043.23904: when converted to jpg file, cannot see anything from the image.","metadata":{"id":"nx4jxfKXSJAi"}},{"cell_type":"code","source":"bad_scans = ['1.2.826.0.1.3680043.20574','1.2.826.0.1.3680043.29952', '1.2.826.0.1.3680043.23904']\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-19T06:06:51.386546Z","iopub.execute_input":"2023-04-19T06:06:51.387635Z","iopub.status.idle":"2023-04-19T06:06:51.423007Z","shell.execute_reply.started":"2023-04-19T06:06:51.387594Z","shell.execute_reply":"2023-04-19T06:06:51.422062Z"},"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-19T06:06:51.424895Z","iopub.execute_input":"2023-04-19T06:06:51.425389Z","iopub.status.idle":"2023-04-19T06:06:51.432166Z","shell.execute_reply.started":"2023-04-19T06:06:51.42534Z","shell.execute_reply":"2023-04-19T06:06:51.430787Z"},"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-19T06:06:51.438151Z","iopub.execute_input":"2023-04-19T06:06:51.438655Z","iopub.status.idle":"2023-04-19T06:06:51.445762Z","shell.execute_reply.started":"2023-04-19T06:06:51.43862Z","shell.execute_reply":"2023-04-19T06:06:51.444457Z"},"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-19T06:06:51.447644Z","iopub.execute_input":"2023-04-19T06:06:51.448035Z","iopub.status.idle":"2023-04-19T06:06:51.460494Z","shell.execute_reply.started":"2023-04-19T06:06:51.447998Z","shell.execute_reply":"2023-04-19T06:06:51.459597Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-19T06:06:51.461778Z","iopub.execute_input":"2023-04-19T06:06:51.462814Z","iopub.status.idle":"2023-04-19T06:06:51.473084Z","shell.execute_reply.started":"2023-04-19T06:06:51.462772Z","shell.execute_reply":"2023-04-19T06:06:51.472312Z"},"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-19T06:06:51.47438Z","iopub.execute_input":"2023-04-19T06:06:51.475019Z","iopub.status.idle":"2023-04-19T06:06:51.486107Z","shell.execute_reply.started":"2023-04-19T06:06:51.474981Z","shell.execute_reply":"2023-04-19T06:06:51.48512Z"},"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.Resize(512, 512),\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([alb.Resize(512, 512)])\n}","metadata":{"id":"xoFZZMAgSJAk","tags":[],"execution":{"iopub.status.busy":"2023-04-19T06:06:51.487317Z","iopub.execute_input":"2023-04-19T06:06:51.487997Z","iopub.status.idle":"2023-04-19T06:06:51.498638Z","shell.execute_reply.started":"2023-04-19T06:06:51.487958Z","shell.execute_reply":"2023-04-19T06:06:51.497753Z"},"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        image = cv2.imread(f'{base_path}/preprocess-yolo-cropped-image-window/yolo_image/{new_uid_id}.jpg')\n        image = np.asarray(image) # img_size x img_size 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_path = f'{base_path}/preprocess-yolo-cropped-image-window/yolo_seg/{uid_id}.npz'\n        try: mask = np.load(mask_path)['arr_0'] # already rotated to sagittal view\n        except: mask = np.zeros((512, 512)) # if there is no mask exist\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-19T06:06:51.500118Z","iopub.execute_input":"2023-04-19T06:06:51.500949Z","iopub.status.idle":"2023-04-19T06:06:51.517503Z","shell.execute_reply.started":"2023-04-19T06:06:51.500907Z","shell.execute_reply":"2023-04-19T06:06:51.516472Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test dataset","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-19T06:10:56.250039Z","iopub.execute_input":"2023-04-19T06:10:56.250607Z","iopub.status.idle":"2023-04-19T06:11:09.466876Z","shell.execute_reply.started":"2023-04-19T06:10:56.250564Z","shell.execute_reply":"2023-04-19T06:11:09.465771Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dataset = CustomDataset(df=df, transform=data_transforms['valid'], test=False)\ntest_loader = DataLoader(test_dataset, batch_size=5, shuffle=True)\n\nimgs, real_mask = next(iter(test_loader))\nplot_batch(imgs, real_mask, size=5)","metadata":{"execution":{"iopub.status.busy":"2023-04-19T06:13:49.969326Z","iopub.execute_input":"2023-04-19T06:13:49.970456Z","iopub.status.idle":"2023-04-19T06:13:51.67935Z","shell.execute_reply.started":"2023-04-19T06:13:49.970406Z","shell.execute_reply":"2023-04-19T06:13:51.678435Z"},"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-18T08:36:34.041006Z","iopub.execute_input":"2023-04-18T08:36:34.041513Z","iopub.status.idle":"2023-04-18T08:36:34.052109Z","shell.execute_reply.started":"2023-04-18T08:36:34.041429Z","shell.execute_reply":"2023-04-18T08:36:34.051163Z"},"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-18T08:36:34.053738Z","iopub.execute_input":"2023-04-18T08:36:34.054138Z","iopub.status.idle":"2023-04-18T08:36:34.064517Z","shell.execute_reply.started":"2023-04-18T08:36:34.054102Z","shell.execute_reply":"2023-04-18T08:36:34.063473Z"},"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-18T08:36:34.066096Z","iopub.execute_input":"2023-04-18T08:36:34.06712Z","iopub.status.idle":"2023-04-18T08:36:34.074165Z","shell.execute_reply.started":"2023-04-18T08:36:34.067076Z","shell.execute_reply":"2023-04-18T08:36:34.073146Z"},"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-18T08:36:34.075738Z","iopub.execute_input":"2023-04-18T08:36:34.076153Z","iopub.status.idle":"2023-04-18T08:36:34.083621Z","shell.execute_reply.started":"2023-04-18T08:36:34.076119Z","shell.execute_reply":"2023-04-18T08:36:34.082584Z"},"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-18T08:36:34.085065Z","iopub.execute_input":"2023-04-18T08:36:34.085539Z","iopub.status.idle":"2023-04-18T08:36:34.094603Z","shell.execute_reply.started":"2023-04-18T08:36:34.085505Z","shell.execute_reply":"2023-04-18T08:36:34.093609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Settings","metadata":{}},{"cell_type":"code","source":"# model setting\nmodel = SegmentationModel()\nmodel.to(device)\n#model.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-18T08:36:34.096212Z","iopub.execute_input":"2023-04-18T08:36:34.097057Z","iopub.status.idle":"2023-04-18T08:36:44.576065Z","shell.execute_reply.started":"2023-04-18T08:36:34.097023Z","shell.execute_reply":"2023-04-18T08:36:44.574848Z"},"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-18T08:36:44.57829Z","iopub.execute_input":"2023-04-18T08:36:44.578757Z","iopub.status.idle":"2023-04-18T08:36:44.592317Z","shell.execute_reply.started":"2023-04-18T08:36:44.578709Z","shell.execute_reply":"2023-04-18T08:36:44.591075Z"},"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-18T08:36:44.59385Z","iopub.execute_input":"2023-04-18T08:36:44.594945Z","iopub.status.idle":"2023-04-18T08:36:44.610562Z","shell.execute_reply.started":"2023-04-18T08:36:44.594899Z","shell.execute_reply":"2023-04-18T08:36:44.609478Z"},"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-18T08:36:44.612434Z","iopub.execute_input":"2023-04-18T08:36:44.612847Z","iopub.status.idle":"2023-04-18T08:36:58.954697Z","shell.execute_reply.started":"2023-04-18T08:36:44.612804Z","shell.execute_reply":"2023-04-18T08:36:58.953667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2023-04-18T08:36:58.95624Z","iopub.execute_input":"2023-04-18T08:36:58.956631Z","iopub.status.idle":"2023-04-18T08:36:58.970956Z","shell.execute_reply.started":"2023-04-18T08:36:58.956586Z","shell.execute_reply":"2023-04-18T08:36:58.969935Z"},"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-18T08:36:58.972622Z","iopub.execute_input":"2023-04-18T08:36:58.973299Z","iopub.status.idle":"2023-04-18T08:36:58.994288Z","shell.execute_reply.started":"2023-04-18T08:36:58.973247Z","shell.execute_reply":"2023-04-18T08:36:58.9934Z"},"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-18T08:36:58.995665Z","iopub.execute_input":"2023-04-18T08:36:58.996186Z","iopub.status.idle":"2023-04-18T18:07:10.394823Z","shell.execute_reply.started":"2023-04-18T08:36:58.996151Z","shell.execute_reply":"2023-04-18T18:07:10.39374Z"},"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":{"execution":{"iopub.status.busy":"2023-04-18T18:07:10.396226Z","iopub.execute_input":"2023-04-18T18:07:10.397245Z","iopub.status.idle":"2023-04-18T18:07:10.40847Z","shell.execute_reply.started":"2023-04-18T18:07:10.397204Z","shell.execute_reply":"2023-04-18T18:07:10.407327Z"},"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-18T18:07:10.41009Z","iopub.execute_input":"2023-04-18T18:07:10.410963Z","iopub.status.idle":"2023-04-18T18:07:11.658332Z","shell.execute_reply.started":"2023-04-18T18:07:10.41092Z","shell.execute_reply":"2023-04-18T18:07:11.65571Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualization","metadata":{}},{"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-18T18:07:11.66417Z","iopub.status.idle":"2023-04-18T18:07:11.665044Z","shell.execute_reply.started":"2023-04-18T18:07:11.664778Z","shell.execute_reply":"2023-04-18T18:07:11.664804Z"},"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-18T18:07:11.666667Z","iopub.status.idle":"2023-04-18T18:07:11.667532Z","shell.execute_reply.started":"2023-04-18T18:07:11.667258Z","shell.execute_reply":"2023-04-18T18:07:11.667297Z"},"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-18T18:07:11.668959Z","iopub.status.idle":"2023-04-18T18:07:11.670948Z","shell.execute_reply.started":"2023-04-18T18:07:11.67065Z","shell.execute_reply":"2023-04-18T18:07:11.670685Z"},"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-18T18:07:11.67254Z","iopub.status.idle":"2023-04-18T18:07:11.673077Z","shell.execute_reply.started":"2023-04-18T18:07:11.672811Z","shell.execute_reply":"2023-04-18T18:07:11.672845Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_batch(imgs, real_mask, size=5)","metadata":{"execution":{"iopub.status.busy":"2023-04-18T18:07:11.674464Z","iopub.status.idle":"2023-04-18T18:07:11.675543Z","shell.execute_reply.started":"2023-04-18T18:07:11.675258Z","shell.execute_reply":"2023-04-18T18:07:11.675305Z"},"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-18T18:07:11.677157Z","iopub.status.idle":"2023-04-18T18:07:11.677665Z","shell.execute_reply.started":"2023-04-18T18:07:11.677411Z","shell.execute_reply":"2023-04-18T18:07:11.677436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}