{"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-04T11:51:44.087697Z","iopub.execute_input":"2023-04-04T11:51:44.087973Z","iopub.status.idle":"2023-04-04T11:52:23.487551Z","shell.execute_reply.started":"2023-04-04T11:51:44.087947Z","shell.execute_reply":"2023-04-04T11:52:23.486306Z"},"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-04T11:53:55.543186Z","iopub.execute_input":"2023-04-04T11:53:55.543878Z","iopub.status.idle":"2023-04-04T11:53:55.552869Z","shell.execute_reply.started":"2023-04-04T11:53:55.543842Z","shell.execute_reply":"2023-04-04T11:53:55.551474Z"},"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 = 4\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 = []","metadata":{"id":"u7dgRAJ7ihtl","execution":{"iopub.status.busy":"2023-04-04T11:52:38.164194Z","iopub.execute_input":"2023-04-04T11:52:38.164576Z","iopub.status.idle":"2023-04-04T11:52:38.227816Z","shell.execute_reply.started":"2023-04-04T11:52:38.164538Z","shell.execute_reply":"2023-04-04T11:52:38.226685Z"},"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-04T11:52:38.231238Z","iopub.execute_input":"2023-04-04T11:52:38.231688Z","iopub.status.idle":"2023-04-04T11:52:54.39436Z","shell.execute_reply.started":"2023-04-04T11:52:38.231648Z","shell.execute_reply":"2023-04-04T11:52:54.393314Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"vert_df","metadata":{"execution":{"iopub.status.busy":"2023-04-04T11:52:54.396016Z","iopub.execute_input":"2023-04-04T11:52:54.396403Z","iopub.status.idle":"2023-04-04T11:52:54.420331Z","shell.execute_reply.started":"2023-04-04T11:52:54.396365Z","shell.execute_reply":"2023-04-04T11:52:54.419308Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from matplotlib import pyplot as plt\nimport nibabel as nib\nimport numpy as np\n\n# testing for axial viewed image with labels\nbase_path = '/kaggle/input'\nuid_id = '1.2.826.0.1.3680043.10633'\n\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)[180]\nprint(mask)\nprint(np.unique(mask))\nprint(mask.shape)\n\nplt.imshow(mask, interpolation='nearest')\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-04-04T11:52:54.421707Z","iopub.execute_input":"2023-04-04T11:52:54.422137Z","iopub.status.idle":"2023-04-04T11:52:56.23329Z","shell.execute_reply.started":"2023-04-04T11:52:54.422102Z","shell.execute_reply":"2023-04-04T11:52:56.23229Z"},"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-04T11:52:56.234819Z","iopub.execute_input":"2023-04-04T11:52:56.235187Z","iopub.status.idle":"2023-04-04T11:52:56.429215Z","shell.execute_reply.started":"2023-04-04T11:52:56.23515Z","shell.execute_reply":"2023-04-04T11:52:56.428285Z"},"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-04T11:52:56.430691Z","iopub.execute_input":"2023-04-04T11:52:56.431031Z","iopub.status.idle":"2023-04-04T11:52:56.460033Z","shell.execute_reply.started":"2023-04-04T11:52:56.430994Z","shell.execute_reply":"2023-04-04T11:52:56.458841Z"},"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-04T11:52:56.461836Z","iopub.execute_input":"2023-04-04T11:52:56.462198Z","iopub.status.idle":"2023-04-04T11:52:56.46823Z","shell.execute_reply.started":"2023-04-04T11:52:56.462163Z","shell.execute_reply":"2023-04-04T11:52:56.467048Z"},"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-04T11:52:56.473911Z","iopub.execute_input":"2023-04-04T11:52:56.47418Z","iopub.status.idle":"2023-04-04T11:52:56.481079Z","shell.execute_reply.started":"2023-04-04T11:52:56.474155Z","shell.execute_reply":"2023-04-04T11:52:56.479855Z"},"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-04T11:52:56.482587Z","iopub.execute_input":"2023-04-04T11:52:56.483039Z","iopub.status.idle":"2023-04-04T11:52:56.492759Z","shell.execute_reply.started":"2023-04-04T11:52:56.483003Z","shell.execute_reply":"2023-04-04T11:52:56.491493Z"},"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-04T11:52:56.494441Z","iopub.execute_input":"2023-04-04T11:52:56.494832Z","iopub.status.idle":"2023-04-04T11:52:56.505841Z","shell.execute_reply.started":"2023-04-04T11:52:56.494792Z","shell.execute_reply":"2023-04-04T11:52:56.504906Z"},"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        image = np.load(f'{base_path}/3-channel-preprocessed-dataset/prep_train/{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-1):\n            if channel+1 in list(np.unique(mask)): real_mask.append(np.where(mask==channel+1, 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-04T11:52:56.507665Z","iopub.execute_input":"2023-04-04T11:52:56.508056Z","iopub.status.idle":"2023-04-04T11:52:56.521917Z","shell.execute_reply.started":"2023-04-04T11:52:56.507988Z","shell.execute_reply":"2023-04-04T11:52:56.52084Z"},"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-04T11:52:56.523436Z","iopub.execute_input":"2023-04-04T11:52:56.523774Z","iopub.status.idle":"2023-04-04T11:52:56.533525Z","shell.execute_reply.started":"2023-04-04T11:52:56.52374Z","shell.execute_reply":"2023-04-04T11:52:56.532547Z"},"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-04T11:52:56.535277Z","iopub.execute_input":"2023-04-04T11:52:56.535753Z","iopub.status.idle":"2023-04-04T11:52:56.545002Z","shell.execute_reply.started":"2023-04-04T11:52:56.535719Z","shell.execute_reply":"2023-04-04T11:52:56.544027Z"},"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-04T11:52:56.546614Z","iopub.execute_input":"2023-04-04T11:52:56.547021Z","iopub.status.idle":"2023-04-04T11:52:56.553856Z","shell.execute_reply.started":"2023-04-04T11:52:56.546987Z","shell.execute_reply":"2023-04-04T11:52:56.552281Z"},"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-04T11:52:56.55558Z","iopub.execute_input":"2023-04-04T11:52:56.556406Z","iopub.status.idle":"2023-04-04T11:52:56.564429Z","shell.execute_reply.started":"2023-04-04T11:52:56.556373Z","shell.execute_reply":"2023-04-04T11:52:56.563455Z"},"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-04T11:52:56.565941Z","iopub.execute_input":"2023-04-04T11:52:56.566418Z","iopub.status.idle":"2023-04-04T11:52:56.573686Z","shell.execute_reply.started":"2023-04-04T11:52:56.566283Z","shell.execute_reply":"2023-04-04T11:52:56.572685Z"},"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-04T11:52:56.574996Z","iopub.execute_input":"2023-04-04T11:52:56.575709Z","iopub.status.idle":"2023-04-04T11:53:07.795592Z","shell.execute_reply.started":"2023-04-04T11:52:56.575673Z","shell.execute_reply":"2023-04-04T11:53:07.794517Z"},"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-04T11:53:07.79708Z","iopub.execute_input":"2023-04-04T11:53:07.797447Z","iopub.status.idle":"2023-04-04T11:53:07.809337Z","shell.execute_reply.started":"2023-04-04T11:53:07.79741Z","shell.execute_reply":"2023-04-04T11:53:07.806285Z"},"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-04T11:53:07.810948Z","iopub.execute_input":"2023-04-04T11:53:07.811651Z","iopub.status.idle":"2023-04-04T11:53:07.826178Z","shell.execute_reply.started":"2023-04-04T11:53:07.811614Z","shell.execute_reply":"2023-04-04T11:53:07.82516Z"},"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-04T11:53:07.828073Z","iopub.execute_input":"2023-04-04T11:53:07.828481Z","iopub.status.idle":"2023-04-04T11:53:20.625448Z","shell.execute_reply.started":"2023-04-04T11:53:07.828405Z","shell.execute_reply":"2023-04-04T11:53:20.624322Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2023-04-04T11:53:20.626834Z","iopub.execute_input":"2023-04-04T11:53:20.627722Z","iopub.status.idle":"2023-04-04T11:53:20.644825Z","shell.execute_reply.started":"2023-04-04T11:53:20.627683Z","shell.execute_reply":"2023-04-04T11:53:20.643966Z"},"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-04T11:53:20.64646Z","iopub.execute_input":"2023-04-04T11:53:20.646857Z","iopub.status.idle":"2023-04-04T11:53:20.664288Z","shell.execute_reply.started":"2023-04-04T11:53:20.646821Z","shell.execute_reply":"2023-04-04T11:53:20.663391Z"},"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-04T11:54:01.910332Z","iopub.execute_input":"2023-04-04T11:54:01.910704Z","iopub.status.idle":"2023-04-04T20:10:15.49898Z","shell.execute_reply.started":"2023-04-04T11:54:01.910672Z","shell.execute_reply":"2023-04-04T20:10:15.497995Z"},"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-04T20:10:15.801936Z","iopub.execute_input":"2023-04-04T20:10:15.802533Z","iopub.status.idle":"2023-04-04T20:10:15.811151Z","shell.execute_reply.started":"2023-04-04T20:10:15.802496Z","shell.execute_reply":"2023-04-04T20:10:15.809574Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}