{"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":"# **Acknowledgements**","metadata":{}},{"cell_type":"markdown","source":"This is a notebook combining crops of the train tiffs and the UNet notebooks I have created. This allows you to use segmentation_models_pytorch with cropped images. **Make sure to upvote the cropping notebooks!**\n\nCropping notebooks:\nhttps://www.kaggle.com/code/thedevastator/converting-to-256x256\n\nhttps://www.kaggle.com/iafoss/256x256-images\n\n\nOriginal UNet notebooks:\nhttps://www.kaggle.com/code/vexxingbanana/hubmap-unet-semantic-approach-train/notebook\n\nhttps://www.kaggle.com/code/vexxingbanana/hubmap-unet-semantic-approach-infer","metadata":{}},{"cell_type":"markdown","source":"# **Inference Notebook**","metadata":{}},{"cell_type":"markdown","source":"https://www.kaggle.com/code/vexxingbanana/hubmap-unet-w-256x256-crops-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-08-26T03:18:42.448816Z","iopub.execute_input":"2022-08-26T03:18:42.449305Z","iopub.status.idle":"2022-08-26T03:18:51.948180Z","shell.execute_reply.started":"2022-08-26T03:18:42.449262Z","shell.execute_reply":"2022-08-26T03:18:51.946744Z"},"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 re\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, GroupKFold\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-08-26T03:18:51.954447Z","iopub.execute_input":"2022-08-26T03:18:51.956975Z","iopub.status.idle":"2022-08-26T03:18:51.969947Z","shell.execute_reply.started":"2022-08-26T03:18:51.956933Z","shell.execute_reply":"2022-08-26T03:18:51.968024Z"},"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-b3\" #['efficientnet-b0', efficientnet-b1', ... , efficientnet-b7'] and many other backbone architectures\n    img_size = [512, 512] # Or [512, 512]\n    lr = 1e-3\n    scheduler = 'CosineAnnealingLR' #['CosineAnnealingLR', 'ReduceLROnPlateau', 'ExponentialLR', 'OneCycleLR']  # 调整学习率\n    epochs = 20\n    warmup_epochs = 2\n    n_folds = 5\n    folds_to_run = [0, 1, 2, 3, 4]\n    device = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n    base_path = '/kaggle/input/hubmap-organ-segmentation'\n    crops_path = '/kaggle/input/hubmap-2022-512x512' # Or '/kaggle/input/hubmap-2022-512x512'\n    num_workers = mp.cpu_count()\n    num_classes = 1\n    n_accumulate = max(1, 16//batch_size)\n    loss = 'Dice'\n    optimizer = 'AdamW'\n    weight_decay = 1e-6","metadata":{"execution":{"iopub.status.busy":"2022-08-26T03:18:51.975715Z","iopub.execute_input":"2022-08-26T03:18:51.978312Z","iopub.status.idle":"2022-08-26T03:18:51.988689Z","shell.execute_reply.started":"2022-08-26T03:18:51.978274Z","shell.execute_reply":"2022-08-26T03:18:51.987781Z"},"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-08-26T03:18:51.995037Z","iopub.execute_input":"2022-08-26T03:18:51.997867Z","iopub.status.idle":"2022-08-26T03:18:52.010211Z","shell.execute_reply.started":"2022-08-26T03:18:51.997830Z","shell.execute_reply":"2022-08-26T03:18:52.008999Z"},"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\n\ndef read_img(path, verbose=0):\n    image = cv2.cvtColor(cv2.imread(path), cv2.COLOR_BGR2RGB)\n    \n    if verbose:\n        print(f\"[{path}] Image shape: {image.shape}\")\n        \n#     mx = np.max(image)\n#     image = image.astype(np.float32)\n#     if mx:\n#         image /= mx\n        \n    return image","metadata":{"execution":{"iopub.status.busy":"2022-08-26T03:18:52.015062Z","iopub.execute_input":"2022-08-26T03:18:52.017688Z","iopub.status.idle":"2022-08-26T03:18:52.028641Z","shell.execute_reply.started":"2022-08-26T03:18:52.017630Z","shell.execute_reply":"2022-08-26T03:18:52.027684Z"},"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-08-26T03:18:52.034059Z","iopub.execute_input":"2022-08-26T03:18:52.036612Z","iopub.status.idle":"2022-08-26T03:18:52.047431Z","shell.execute_reply.started":"2022-08-26T03:18:52.036576Z","shell.execute_reply":"2022-08-26T03:18:52.045791Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Set Seed For Reproducibility**","metadata":{}},{"cell_type":"code","source":"def set_seed(seed = 46): #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-08-26T03:18:52.051207Z","iopub.execute_input":"2022-08-26T03:18:52.051586Z","iopub.status.idle":"2022-08-26T03:18:52.065392Z","shell.execute_reply.started":"2022-08-26T03:18:52.051551Z","shell.execute_reply":"2022-08-26T03:18:52.064353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Grab Metadata**","metadata":{}},{"cell_type":"code","source":"comp_train_df = pd.read_csv(os.path.join(CFG.base_path, 'train.csv'))\ncomp_train_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-26T03:18:52.066636Z","iopub.execute_input":"2022-08-26T03:18:52.067317Z","iopub.status.idle":"2022-08-26T03:18:52.297627Z","shell.execute_reply.started":"2022-08-26T03:18:52.067281Z","shell.execute_reply":"2022-08-26T03:18:52.296244Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images_list = list(glob.glob(f'{CFG.crops_path}/train/*.png'))\nprint(images_list[:5])","metadata":{"execution":{"iopub.status.busy":"2022-08-26T03:18:52.302585Z","iopub.execute_input":"2022-08-26T03:18:52.302990Z","iopub.status.idle":"2022-08-26T03:18:52.317592Z","shell.execute_reply.started":"2022-08-26T03:18:52.302953Z","shell.execute_reply":"2022-08-26T03:18:52.316388Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = pd.DataFrame({\n    \"id\": [x.split('/')[-1].split('.')[0] for x in images_list],\n    \"image_paths\": images_list,\n    \"mask_paths\": [x.replace('train', 'masks') for x in images_list],\n    \"organ\": [comp_train_df[comp_train_df['id'] == int(x.split('/')[-1].split('.')[0])]['organ'].item() for x in images_list]\n    #You can grab any metadata from the original train.csv for each cropped image by modifying ['organ'] to the column you want above\n})\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-26T03:18:52.324379Z","iopub.execute_input":"2022-08-26T03:18:52.326590Z","iopub.status.idle":"2022-08-26T03:18:52.542547Z","shell.execute_reply.started":"2022-08-26T03:18:52.326555Z","shell.execute_reply":"2022-08-26T03:18:52.541669Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# **Data Processing**","metadata":{}},{"cell_type":"code","source":"kf = GroupKFold(n_splits=CFG.n_folds)\nfor fold, (train_idx, val_idx) in enumerate(kf.split(df, groups=df['id'])):\n    df.loc[val_idx, 'fold'] = fold\ndf.head()","metadata":{"execution":{"iopub.status.busy":"2022-08-26T03:18:52.546592Z","iopub.execute_input":"2022-08-26T03:18:52.548726Z","iopub.status.idle":"2022-08-26T03:18:52.574042Z","shell.execute_reply.started":"2022-08-26T03:18:52.548689Z","shell.execute_reply":"2022-08-26T03:18:52.573165Z"},"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_paths']\n        img = read_img(img_path)\n        \n        if self.labeled:\n            mask_path = self.df.loc[index, 'mask_paths']\n            mask = cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE)\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-08-26T03:18:52.577910Z","iopub.execute_input":"2022-08-26T03:18:52.580544Z","iopub.status.idle":"2022-08-26T03:18:52.593908Z","shell.execute_reply.started":"2022-08-26T03:18:52.580509Z","shell.execute_reply":"2022-08-26T03:18:52.592792Z"},"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        A.RandomRotate90(p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.5, border_mode=cv2.BORDER_REFLECT),\n        A.OneOf([\n            A.HueSaturationValue(10,15,10),\n            A.CLAHE(clip_limit=2),\n            A.RandomBrightnessContrast(),            \n        ], p=0.4),\n        A.Normalize(),\n    ]),\n    \n    \"valid\": A.Compose([\n        A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n        A.Normalize(),\n        ], p=1.0),\n}\n# data_transforms = {\n#     \"train\": A.Compose([\n#         A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n#         A.HorizontalFlip(p=0.5),\n#         A.RandomRotate90(),  # 随即旋转\n#         A.VerticalFlip(p=0.5),  # 垂直翻转\n#         A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.6, \n#                          border_mode=cv2.BORDER_REFLECT),  # 随机平移，缩放和旋转输入\n#         A.OneOf([\n#             A.OpticalDistortion(p=0.3),\n#             A.GridDistortion(p=0.1),\n#             A.PiecewiseAffine(p=0.3),\n#         ], p=0.3),\n#         A.OneOf([\n#             A.HueSaturationValue(10,15,10),\n#             A.CLAHE(clip_limit=2),\n#             A.RandomBrightnessContrast(),            \n#         ], p=0.3),\n#         A.Normalize(),\n#     ]),\n    \n#     \"valid\": A.Compose([\n#         A.Resize(*CFG.img_size, interpolation=cv2.INTER_NEAREST),\n#         A.Normalize(),\n#         ], p=1.0),\n# }","metadata":{"execution":{"iopub.status.busy":"2022-08-26T03:18:52.598516Z","iopub.execute_input":"2022-08-26T03:18:52.601542Z","iopub.status.idle":"2022-08-26T03:18:52.612685Z","shell.execute_reply.started":"2022-08-26T03:18:52.601471Z","shell.execute_reply":"2022-08-26T03:18:52.611739Z"},"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-08-26T03:18:52.617592Z","iopub.execute_input":"2022-08-26T03:18:52.620243Z","iopub.status.idle":"2022-08-26T03:18:52.636096Z","shell.execute_reply.started":"2022-08-26T03:18:52.620207Z","shell.execute_reply":"2022-08-26T03:18:52.635148Z"},"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        num_steps = len(df[df['fold'] != CFG.folds_to_run[0]]) // CFG.batch_size\n    # kFolds相关，确定训练除了CFG.folds_to_run[i]步骤的次数\n    \n    # scheduler = 'OneCycleLR' #['CosineAnnealingLR', 'ReduceLROnPlateau', 'ExponentialLR', 'OneCycleLR']  # 调整学习率\n    if CFG.scheduler == 'CosineAnnealingLR':\n        scheduler = transformers.get_cosine_schedule_with_warmup(optimizer, CFG.warmup_epochs * num_steps, CFG.epochs * num_steps)\n        # 先逐渐增加，再逐渐减小; CFG.warmup_epochs = 2; CFG.epochs = 20\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    # 学习率调整为：当某一指标不再提升时，降低学习率。\n    elif CFG.scheduler == 'ExponentialLR':\n        scheduler = lr_scheduler.ExponentialLR(optimizer, gamma=0.85)\n    # 学习率调整为：指数衰减初始学习率。（lr = lr * gamma**epoch）\n    elif CFG.scheduler == 'OneCycleLR':\n        scheduler = lr_scheduler.OneCycleLR(optimizer, pct_start=(CFG.warmup_epochs * 1.0 / CFG.epochs), div_factor=1e3, \n                                              max_lr=CFG.lr, epochs=CFG.epochs, steps_per_epoch=num_steps)\n    # 学习率调整为：根据“1cycle”策略，设置各参数组的学习率。1cycle策略将学习率从初始学习率退火到最大学习率，然后从最大学习率退火到远低于初始学习率的最小学习率。\n    return scheduler","metadata":{"execution":{"iopub.status.busy":"2022-08-26T03:18:52.640802Z","iopub.execute_input":"2022-08-26T03:18:52.643573Z","iopub.status.idle":"2022-08-26T03:18:52.654597Z","shell.execute_reply.started":"2022-08-26T03:18:52.643539Z","shell.execute_reply":"2022-08-26T03:18:52.653569Z"},"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-08-26T03:18:52.659645Z","iopub.execute_input":"2022-08-26T03:18:52.662260Z","iopub.status.idle":"2022-08-26T03:18:52.670208Z","shell.execute_reply.started":"2022-08-26T03:18:52.662224Z","shell.execute_reply":"2022-08-26T03:18:52.669229Z"},"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-08-26T03:18:52.675273Z","iopub.execute_input":"2022-08-26T03:18:52.677904Z","iopub.status.idle":"2022-08-26T03:18:52.687013Z","shell.execute_reply.started":"2022-08-26T03:18:52.677868Z","shell.execute_reply":"2022-08-26T03:18:52.686026Z"},"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-08-26T03:18:52.692017Z","iopub.execute_input":"2022-08-26T03:18:52.694551Z","iopub.status.idle":"2022-08-26T03:18:52.709392Z","shell.execute_reply.started":"2022-08-26T03:18:52.694516Z","shell.execute_reply":"2022-08-26T03:18:52.708320Z"},"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-08-26T03:18:52.714634Z","iopub.execute_input":"2022-08-26T03:18:52.717295Z","iopub.status.idle":"2022-08-26T03:18:52.731839Z","shell.execute_reply.started":"2022-08-26T03:18:52.717259Z","shell.execute_reply":"2022-08-26T03:18:52.730940Z"},"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-08-26T03:18:52.736436Z","iopub.execute_input":"2022-08-26T03:18:52.738452Z","iopub.status.idle":"2022-08-26T03:18:52.758588Z","shell.execute_reply.started":"2022-08-26T03:18:52.738418Z","shell.execute_reply":"2022-08-26T03:18:52.757577Z"},"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-08-26T03:18:52.763788Z","iopub.execute_input":"2022-08-26T03:18:52.766638Z","iopub.status.idle":"2022-08-26T03:32:27.747383Z","shell.execute_reply.started":"2022-08-26T03:18:52.766591Z","shell.execute_reply":"2022-08-26T03:32:27.746212Z"},"trusted":true},"execution_count":null,"outputs":[]}]}