{"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":"# **Inference Notebook**","metadata":{}},{"cell_type":"markdown","source":"https://www.kaggle.com/code/vexxingbanana/hubmap-unet-semantic-approach-infer","metadata":{}},{"cell_type":"markdown","source":"# **Install segmentation_models_pytorch**","metadata":{}},{"cell_type":"code","source":"!pip install segmentation_models_pytorch","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-06-25T01:28:38.456432Z","iopub.execute_input":"2022-06-25T01:28:38.457028Z","iopub.status.idle":"2022-06-25T01:28:55.157443Z","shell.execute_reply.started":"2022-06-25T01:28:38.456935Z","shell.execute_reply":"2022-06-25T01:28:55.156192Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Import Libraries**","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport time\nimport matplotlib.pyplot as plt\nimport cv2\nimport glob\nimport os\nimport shutil\nimport timm\nimport random\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.cuda import amp\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport transformers\nfrom sklearn.model_selection import StratifiedKFold, KFold, StratifiedGroupKFold\nimport multiprocessing as mp\nimport segmentation_models_pytorch as smp\nimport copy\nfrom collections import defaultdict\nimport gc\nfrom tqdm import tqdm\nimport tifffile\nfrom colorama import Fore, Back, Style\nc_  = Fore.GREEN\nsr_ = Style.RESET_ALL","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-06-25T01:28:55.16051Z","iopub.execute_input":"2022-06-25T01:28:55.160952Z","iopub.status.idle":"2022-06-25T01:29:06.299507Z","shell.execute_reply.started":"2022-06-25T01:28:55.160909Z","shell.execute_reply":"2022-06-25T01:29:06.298479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Config**","metadata":{}},{"cell_type":"code","source":"class CFG:\n    seed = 0\n    batch_size = 16\n    head = \"UNet\"\n    backbone = \"efficientnet-b0\"\n    img_size = [512, 512]\n    lr = 1e-3\n    scheduler = 'CosineAnnealingLR' #['CosineAnnealingLR']\n    epochs = 20\n    warmup_epochs = 2\n    n_folds = 5\n    folds_to_run = [0]\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    base_path = '../input/hubmap-organ-segmentation'\n    num_workers = mp.cpu_count()\n    num_classes = 1\n    n_accumulate = max(1, 16//batch_size)\n    loss = 'Dice'\n    optimizer = 'Adam'\n    weight_decay = 1e-6","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.301016Z","iopub.execute_input":"2022-06-25T01:29:06.301801Z","iopub.status.idle":"2022-06-25T01:29:06.379128Z","shell.execute_reply.started":"2022-06-25T01:29:06.301762Z","shell.execute_reply":"2022-06-25T01:29:06.378106Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Helper Functions**","metadata":{}},{"cell_type":"code","source":"# ref: https://www.kaggle.com/paulorzp/run-length-encode-and-decode\ndef rle_decode(mask_rle, shape):\n    '''\n    mask_rle: run-length as string formated (start length)\n    shape: (height,width) of array to return \n    Returns numpy array, 1 - mask, 0 - background\n\n    '''\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape)  # Needed to align to RLE direction\n\n\n# ref.: https://www.kaggle.com/stainsby/fast-tested-rle\ndef rle_encode(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.383029Z","iopub.execute_input":"2022-06-25T01:29:06.383996Z","iopub.status.idle":"2022-06-25T01:29:06.394179Z","shell.execute_reply.started":"2022-06-25T01:29:06.383956Z","shell.execute_reply":"2022-06-25T01:29:06.392978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def read_tiff(path, scale=None, verbose=0): #Modified from https://www.kaggle.com/code/abhinand05/hubmap-extensive-eda-what-are-we-hacking\n    image = tifffile.imread(path)\n    if len(image.shape) == 5:\n        image = image.squeeze().transpose(1, 2, 0)\n    \n    if verbose:\n        print(f\"[{path}] Image shape: {image.shape}\")\n    \n    if scale:\n        new_size = (image.shape[1] // scale, image.shape[0] // scale)\n        image = cv2.resize(image, new_size)\n        \n        if verbose:\n            print(f\"[{path}] Resized Image shape: {image.shape}\")\n        \n    mx = np.max(image)\n    image = image.astype(np.float32)\n    if mx:\n        image /= mx # scale image to [0, 1]\n    return image","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.395994Z","iopub.execute_input":"2022-06-25T01:29:06.39649Z","iopub.status.idle":"2022-06-25T01:29:06.407966Z","shell.execute_reply.started":"2022-06-25T01:29:06.396452Z","shell.execute_reply":"2022-06-25T01:29:06.40697Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def prepare_loaders(fold):\n    train_df = df.query(\"fold!=@fold\").reset_index(drop=True)\n    valid_df = df.query(\"fold==@fold\").reset_index(drop=True)\n\n    train_dataset = HuBMAP_Dataset(train_df, transforms=data_transforms['train'])\n    valid_dataset = HuBMAP_Dataset(valid_df, transforms=data_transforms['valid'])\n\n    train_loader = torch.utils.data.DataLoader(train_dataset, batch_size=CFG.batch_size,\n                              num_workers=CFG.num_workers, shuffle=True, pin_memory=True, drop_last=False)\n    valid_loader = torch.utils.data.DataLoader(valid_dataset, batch_size=CFG.batch_size,\n                              num_workers=CFG.num_workers, shuffle=False, pin_memory=True)\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.409568Z","iopub.execute_input":"2022-06-25T01:29:06.409994Z","iopub.status.idle":"2022-06-25T01:29:06.419469Z","shell.execute_reply.started":"2022-06-25T01:29:06.409959Z","shell.execute_reply":"2022-06-25T01:29:06.41843Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Set Seed For Reproducibility**","metadata":{}},{"cell_type":"code","source":"def set_seed(seed = 42): #From https://www.kaggle.com/code/awsaf49/uwmgi-unet-train-pytorch/\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    # When running on the CuDNN backend, two further options must be set\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    # Set a fixed value for the hash seed\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    print('> SEEDING DONE')\n    \nset_seed(CFG.seed)","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.421306Z","iopub.execute_input":"2022-06-25T01:29:06.421777Z","iopub.status.idle":"2022-06-25T01:29:06.435734Z","shell.execute_reply.started":"2022-06-25T01:29:06.42172Z","shell.execute_reply":"2022-06-25T01:29:06.4345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Grab Metadata**","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(os.path.join(CFG.base_path, \"train.csv\"))\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.437179Z","iopub.execute_input":"2022-06-25T01:29:06.438099Z","iopub.status.idle":"2022-06-25T01:29:06.768424Z","shell.execute_reply.started":"2022-06-25T01:29:06.438063Z","shell.execute_reply":"2022-06-25T01:29:06.767503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Data Processing**","metadata":{}},{"cell_type":"code","source":"kf = KFold(n_splits=CFG.n_folds, shuffle=True, random_state=CFG.seed)\nfor fold, (train_idx, val_idx) in enumerate(kf.split(df)):\n    df.loc[val_idx, 'fold'] = fold\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.770006Z","iopub.execute_input":"2022-06-25T01:29:06.770354Z","iopub.status.idle":"2022-06-25T01:29:06.800893Z","shell.execute_reply.started":"2022-06-25T01:29:06.770319Z","shell.execute_reply":"2022-06-25T01:29:06.799858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['image_path'] = df['id'].apply(lambda x: os.path.join(CFG.base_path, 'train_images', str(x) + '.tiff'))","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.805395Z","iopub.execute_input":"2022-06-25T01:29:06.805695Z","iopub.status.idle":"2022-06-25T01:29:06.812767Z","shell.execute_reply.started":"2022-06-25T01:29:06.805669Z","shell.execute_reply":"2022-06-25T01:29:06.811572Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.81462Z","iopub.execute_input":"2022-06-25T01:29:06.815265Z","iopub.status.idle":"2022-06-25T01:29:06.852505Z","shell.execute_reply.started":"2022-06-25T01:29:06.815228Z","shell.execute_reply":"2022-06-25T01:29:06.851492Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Dataset**","metadata":{}},{"cell_type":"code","source":"class HuBMAP_Dataset(torch.utils.data.Dataset):\n    def __init__(self, df, labeled=True, transforms=None):\n        self.df = df\n        self.labeled = labeled\n        self.transforms = transforms\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, index):\n        img_path = self.df.loc[index, 'image_path']\n        img_height = self.df.loc[index, 'img_height']\n        img_width = self.df.loc[index, 'img_width']\n        img = read_tiff(img_path)\n        \n        if self.labeled:\n            rle_mask = self.df.loc[index, 'rle']\n            mask = rle_decode(rle_mask, (img_height, img_width))\n            \n            if self.transforms:\n                data = self.transforms(image=img, mask=mask)\n                img  = data['image']\n                mask  = data['mask']\n            \n            mask = np.expand_dims(mask, axis=0)\n            img = np.transpose(img, (2, 0, 1))\n#             mask = np.transpose(mask, (2, 0, 1))\n            \n            return torch.tensor(img), torch.tensor(mask)\n        \n        else:\n            if self.transforms:\n                data = self.transforms(image=img)\n                img  = data['image']\n                \n            img = np.transpose(img, (2, 0, 1))\n            \n            return torch.tensor(img)","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.85588Z","iopub.execute_input":"2022-06-25T01:29:06.856147Z","iopub.status.idle":"2022-06-25T01:29:06.867747Z","shell.execute_reply.started":"2022-06-25T01:29:06.856116Z","shell.execute_reply":"2022-06-25T01:29:06.86651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Augmentations**","metadata":{}},{"cell_type":"code","source":"data_transforms = {\n    \"train\": A.Compose([\n        A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n        A.HorizontalFlip(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.05, rotate_limit=10, p=0.5),\n    ]),\n    \n    \"valid\": A.Compose([\n        A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n        ], p=1.0)\n}","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.869673Z","iopub.execute_input":"2022-06-25T01:29:06.870461Z","iopub.status.idle":"2022-06-25T01:29:06.880865Z","shell.execute_reply.started":"2022-06-25T01:29:06.870412Z","shell.execute_reply":"2022-06-25T01:29:06.879892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Losses and Metrics**","metadata":{}},{"cell_type":"code","source":"JaccardLoss = smp.losses.JaccardLoss(mode='binary')\nDiceLoss    = smp.losses.DiceLoss(mode='binary')\nBCELoss     = smp.losses.SoftBCEWithLogitsLoss()\nLovaszLoss  = smp.losses.LovaszLoss(mode='binary', per_image=False)\nTverskyLoss = smp.losses.TverskyLoss(mode='binary', log_loss=False)\n\ndef dice_coef(y_true, y_pred, thr=0.5, dim=(2,3), epsilon=0.001):\n    y_true = y_true.to(torch.float32)\n    y_pred = (y_pred>thr).to(torch.float32)\n    inter = (y_true*y_pred).sum(dim=dim)\n    den = y_true.sum(dim=dim) + y_pred.sum(dim=dim)\n    dice = ((2*inter+epsilon)/(den+epsilon)).mean(dim=(1,0))\n    return dice\n\ndef iou_coef(y_true, y_pred, thr=0.5, dim=(2,3), epsilon=0.001):\n    y_true = y_true.to(torch.float32)\n    y_pred = (y_pred>thr).to(torch.float32)\n    inter = (y_true*y_pred).sum(dim=dim)\n    union = (y_true + y_pred - y_true*y_pred).sum(dim=dim)\n    iou = ((inter+epsilon)/(union+epsilon)).mean(dim=(1,0))\n    return iou\n\nlosses = {\n    \"Dice\": DiceLoss,\n    \"Jaccard\": JaccardLoss,\n    \"BCE\": BCELoss,\n    \"Lovasz\": LovaszLoss,\n    \"Tversky\": TverskyLoss,\n}","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.882505Z","iopub.execute_input":"2022-06-25T01:29:06.883113Z","iopub.status.idle":"2022-06-25T01:29:06.895584Z","shell.execute_reply.started":"2022-06-25T01:29:06.883078Z","shell.execute_reply":"2022-06-25T01:29:06.894544Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Schedulers**","metadata":{}},{"cell_type":"code","source":"def get_scheduler(optimizer):\n    \n    if len(df[df['fold'] == CFG.folds_to_run[0]]) % CFG.batch_size != 0:\n        num_steps = len(df[df['fold'] != CFG.folds_to_run[0]]) // CFG.batch_size + 1\n    \n    else:\n        len(df[df['fold'] != CFG.folds_to_run[0]]) // CFG.batch_size\n    \n    if CFG.scheduler == 'CosineAnnealingLR':\n        scheduler = transformers.get_cosine_schedule_with_warmup(optimizer, CFG.warmup_epochs * num_steps, CFG.epochs * num_steps)\n        \n    elif CFG.scheduler == 'ReduceLROnPlateau':\n        scheduler = lr_scheduler.ReduceLROnPlateau(optimizer, mode='min', factor=0.1, patience=7, threshold=0.0001, min_lr=1e-6)\n    elif CFG.scheduer == 'ExponentialLR':\n        scheduler = lr_scheduler.ExponentialLR(optimizer, gamma=0.85)\n        \n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.897079Z","iopub.execute_input":"2022-06-25T01:29:06.897499Z","iopub.status.idle":"2022-06-25T01:29:06.908029Z","shell.execute_reply.started":"2022-06-25T01:29:06.897463Z","shell.execute_reply":"2022-06-25T01:29:06.907138Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Optimizers**","metadata":{}},{"cell_type":"code","source":"def get_optimizer(optimizer_name=CFG.optimizer):\n    if CFG.optimizer == 'Adam':\n        optimizer = optim.Adam(model.parameters(), lr=CFG.lr)\n    \n    elif CFG.optimizer == 'AdamW':\n        optimizer = optim.AdamW(model.parameters(), lr=CFG.lr, weight_decay=CFG.weight_decay)\n        \n    return optimizer","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.909933Z","iopub.execute_input":"2022-06-25T01:29:06.910501Z","iopub.status.idle":"2022-06-25T01:29:06.918744Z","shell.execute_reply.started":"2022-06-25T01:29:06.910464Z","shell.execute_reply":"2022-06-25T01:29:06.917753Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Models**","metadata":{}},{"cell_type":"code","source":"def build_model():\n    model = smp.Unet(\n        encoder_name=CFG.backbone,      \n        encoder_weights=\"imagenet\",     \n        in_channels=3,                  \n        classes=CFG.num_classes,\n        activation=None,\n    )\n    model.to(CFG.device)\n    return model\n\ndef load_model(path):\n    model = build_model()\n    model.load_state_dict(torch.load(path))\n    model.eval()\n    return model","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.920318Z","iopub.execute_input":"2022-06-25T01:29:06.920748Z","iopub.status.idle":"2022-06-25T01:29:06.929352Z","shell.execute_reply.started":"2022-06-25T01:29:06.920711Z","shell.execute_reply":"2022-06-25T01:29:06.928296Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Training Functions**","metadata":{}},{"cell_type":"markdown","source":"Modified from https://www.kaggle.com/code/awsaf49/uwmgi-unet-train-pytorch/","metadata":{}},{"cell_type":"code","source":"def train_one_epoch(model, optimizer, scheduler, dataloader, device, epoch):\n    model.train()\n    scaler = amp.GradScaler()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    criterion = losses[CFG.loss]\n    \n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    for step, (images, masks) in pbar:         \n        images = images.to(device, dtype=torch.float)\n        masks  = masks.to(device, dtype=torch.float)\n        \n        batch_size = images.size(0)\n        \n        with amp.autocast(enabled=True):\n            y_pred = model(images)\n            loss   = criterion(y_pred, masks)\n            \n        scaler.scale(loss).backward()\n    \n        if (step + 1) % CFG.n_accumulate == 0:\n            scaler.step(optimizer)\n            scaler.update()\n\n            # zero the parameter gradients\n            optimizer.zero_grad()\n\n            if scheduler is not None:\n                scheduler.step()\n                \n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss / dataset_size\n        \n        mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n        current_lr = optimizer.param_groups[0]['lr']\n        pbar.set_postfix(train_loss=f'{epoch_loss:0.4f}',\n                        lr=f'{current_lr:0.5f}',\n                        gpu_mem=f'{mem:0.2f} GB')\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return epoch_loss","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.93086Z","iopub.execute_input":"2022-06-25T01:29:06.931332Z","iopub.status.idle":"2022-06-25T01:29:06.945075Z","shell.execute_reply.started":"2022-06-25T01:29:06.931264Z","shell.execute_reply":"2022-06-25T01:29:06.944202Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"@torch.no_grad()\ndef valid_one_epoch(model, dataloader, device, epoch):\n    model.eval()\n    \n    dataset_size = 0\n    running_loss = 0.0\n    criterion = losses[CFG.loss]\n    \n    val_scores = []\n    \n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Valid ')\n    for step, (images, masks) in pbar:        \n        images  = images.to(device, dtype=torch.float)\n        masks   = masks.to(device, dtype=torch.float)\n        \n        batch_size = images.size(0)\n        \n        y_pred  = model(images)\n        loss    = criterion(y_pred, masks)\n        \n        running_loss += (loss.item() * batch_size)\n        dataset_size += batch_size\n        \n        epoch_loss = running_loss / dataset_size\n        \n        y_pred = nn.Sigmoid()(y_pred)\n        val_dice = dice_coef(masks, y_pred).cpu().detach().numpy()\n        val_jaccard = iou_coef(masks, y_pred).cpu().detach().numpy()\n        val_scores.append([val_dice, val_jaccard])\n        \n        mem = torch.cuda.memory_reserved() / 1E9 if torch.cuda.is_available() else 0\n        current_lr = optimizer.param_groups[0]['lr']\n        pbar.set_postfix(valid_loss=f'{epoch_loss:0.4f}',\n                        lr=f'{current_lr:0.5f}',\n                        gpu_memory=f'{mem:0.2f} GB')\n    val_scores  = np.mean(val_scores, axis=0)\n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    return epoch_loss, val_scores","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.946604Z","iopub.execute_input":"2022-06-25T01:29:06.947049Z","iopub.status.idle":"2022-06-25T01:29:06.960468Z","shell.execute_reply.started":"2022-06-25T01:29:06.946971Z","shell.execute_reply":"2022-06-25T01:29:06.959546Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def run_training(model, optimizer, scheduler, device, num_epochs):\n    # To automatically log gradients\n    \n    if torch.cuda.is_available():\n        print(\"cuda: {}\\n\".format(torch.cuda.get_device_name()))\n    \n    start = time.time()\n    best_model_wts = copy.deepcopy(model.state_dict())\n    best_dice      = -np.inf\n    best_epoch     = -1\n    history = defaultdict(list)\n    \n    for epoch in range(1, num_epochs + 1): \n        gc.collect()\n        print(f'Epoch {epoch}/{num_epochs}', end='')\n        train_loss = train_one_epoch(model, optimizer, scheduler, \n                                           dataloader=train_loader, \n                                           device=CFG.device, epoch=epoch)\n        \n        val_loss, val_scores = valid_one_epoch(model, valid_loader, \n                                                 device=CFG.device, \n                                                 epoch=epoch)\n        val_dice, val_jaccard = val_scores\n    \n        history['Train Loss'].append(train_loss)\n        history['Valid Loss'].append(val_loss)\n        history['Valid Dice'].append(val_dice)\n        history['Valid Jaccard'].append(val_jaccard)\n        \n        print(f'Valid Dice: {val_dice:0.4f} | Valid Jaccard: {val_jaccard:0.4f}')\n        \n        # deep copy the model\n        if val_dice >= best_dice:\n            print(f\"{c_}Valid Score Improved ({best_dice:0.4f} ---> {val_dice:0.4f})\")\n            best_dice    = val_dice\n            best_jaccard = val_jaccard\n            best_epoch   = epoch\n            best_model_wts = copy.deepcopy(model.state_dict())\n            PATH = f\"best_epoch-{fold:02d}.bin\"\n            torch.save(model.state_dict(), PATH)\n            # Save a model file from the current directory\n            print(f\"Model Saved{sr_}\")\n            \n        last_model_wts = copy.deepcopy(model.state_dict())\n        PATH = f\"last_epoch-{fold:02d}.bin\"\n        torch.save(model.state_dict(), PATH)\n            \n        print(); print()\n    \n    end = time.time()\n    time_elapsed = end - start\n    print('Training complete in {:.0f}h {:.0f}m {:.0f}s'.format(\n        time_elapsed // 3600, (time_elapsed % 3600) // 60, (time_elapsed % 3600) % 60))\n    print(\"Best Score: {:.4f}\".format(best_dice))\n    \n    # load best model weights\n    model.load_state_dict(best_model_wts)\n    \n    return model, history","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.964078Z","iopub.execute_input":"2022-06-25T01:29:06.964486Z","iopub.status.idle":"2022-06-25T01:29:06.978589Z","shell.execute_reply.started":"2022-06-25T01:29:06.964438Z","shell.execute_reply":"2022-06-25T01:29:06.977508Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Training**","metadata":{}},{"cell_type":"code","source":"for fold in CFG.folds_to_run:\n    print(f'#'*15)\n    print(f'### Fold: {fold}')\n    print(f'#'*15)\n    train_loader, valid_loader = prepare_loaders(fold=fold)\n    model = build_model()\n    optimizer = get_optimizer()\n    scheduler = get_scheduler(optimizer)\n    model, history = run_training(model, optimizer, scheduler,\n                                  device=CFG.device,\n                                  num_epochs=CFG.epochs)","metadata":{"execution":{"iopub.status.busy":"2022-06-25T01:29:06.980023Z","iopub.execute_input":"2022-06-25T01:29:06.980706Z","iopub.status.idle":"2022-06-25T01:31:00.034375Z","shell.execute_reply.started":"2022-06-25T01:29:06.98067Z","shell.execute_reply":"2022-06-25T01:31:00.032493Z"},"trusted":true},"execution_count":null,"outputs":[]}]}