{"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-28T04:14:09.371293Z","iopub.execute_input":"2022-08-28T04:14:09.372078Z","iopub.status.idle":"2022-08-28T04:14:10.840711Z","shell.execute_reply.started":"2022-08-28T04:14:09.372043Z","shell.execute_reply":"2022-08-28T04:14:10.829337Z"},"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-28T04:14:10.843272Z","iopub.status.idle":"2022-08-28T04:14:10.844319Z","shell.execute_reply.started":"2022-08-28T04:14:10.843849Z","shell.execute_reply":"2022-08-28T04:14:10.84389Z"},"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']\n    epochs = 20\n    warmup_epochs = 2\n    n_folds = 3  # changed to 4\n    folds_to_run = [0, 1, 2]  # changed to 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_tiled = '/kaggle/input/hubmap-2022-512x512' # Or '/kaggle/input/hubmap-2022-512x512'，这个是[TILED的小图]\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  # 这是什么？\n    # ====================添加部分\n    crops_path_notile_resize = '../input/d/zwzwsun/hubmap-2022-512x512/'  # 这个是原图进行resize\n    crops_path_notile_randcrop = '../input/hubmap-2022-origin'  # 这个是原图进行resize\n    notile_resize_count = 2  # 原图resize所占的权重（重复几遍）\n    notile_randcrop_count = 5  # 原图resize所占的权重（重复几遍）\n    crop_size = [1024, 1024]\n    crop_resize = [512, 512]\n    ","metadata":{"execution":{"iopub.status.busy":"2022-08-28T04:14:38.559298Z","iopub.execute_input":"2022-08-28T04:14:38.559778Z","iopub.status.idle":"2022-08-28T04:14:38.572598Z","shell.execute_reply.started":"2022-08-28T04:14:38.559709Z","shell.execute_reply":"2022-08-28T04:14:38.570118Z"},"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-28T04:14:46.700237Z","iopub.execute_input":"2022-08-28T04:14:46.700671Z","iopub.status.idle":"2022-08-28T04:14:46.714424Z","shell.execute_reply.started":"2022-08-28T04:14:46.700637Z","shell.execute_reply":"2022-08-28T04:14:46.712681Z"},"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-28T04:14:10.869286Z","iopub.status.idle":"2022-08-28T04:14:10.870105Z","shell.execute_reply.started":"2022-08-28T04:14:10.869673Z","shell.execute_reply":"2022-08-28T04:14:10.869707Z"},"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-28T04:14:10.872511Z","iopub.status.idle":"2022-08-28T04:14:10.874233Z","shell.execute_reply.started":"2022-08-28T04:14:10.873893Z","shell.execute_reply":"2022-08-28T04:14:10.873925Z"},"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-08-28T04:14:10.877842Z","iopub.status.idle":"2022-08-28T04:14:10.886259Z","shell.execute_reply.started":"2022-08-28T04:14:10.885926Z","shell.execute_reply":"2022-08-28T04:14:10.885957Z"},"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()\n# 存train所有信息","metadata":{"execution":{"iopub.status.busy":"2022-08-28T04:14:52.702991Z","iopub.execute_input":"2022-08-28T04:14:52.703409Z","iopub.status.idle":"2022-08-28T04:14:52.879404Z","shell.execute_reply.started":"2022-08-28T04:14:52.703375Z","shell.execute_reply":"2022-08-28T04:14:52.877983Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 将[n+1]维的矩阵打散为[n]维\ndef flatten_l_o_l(nested_list):\n    \"\"\" Flatten a list of lists \"\"\"\n    return [item for sublist in nested_list for item in sublist]","metadata":{"execution":{"iopub.status.busy":"2022-08-28T04:14:54.923432Z","iopub.execute_input":"2022-08-28T04:14:54.9239Z","iopub.status.idle":"2022-08-28T04:14:54.932108Z","shell.execute_reply.started":"2022-08-28T04:14:54.923864Z","shell.execute_reply":"2022-08-28T04:14:54.930166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images_list_tiled = []\nimages_list_notile = []\nimages_status_tiled = []  # 记录数据来源，以后可以考虑对于不同数据来源图像进行不同处理\nimages_status_notile = []  # 记录数据来源，以后可以考虑对于不同数据来源图像进行不同处理\n# status  |    explanation\n# --------------------------\n#   1     |       tiled\n#   2     |    notile_orig\n#   3     |  orig_random_crop\n\n\nimages_list_tiled = list(glob.glob(f'{CFG.crops_path_tiled}/train/*.png'))  # 分割'_'单独分类\nfor i in images_list_tiled:\n    images_status_tiled.append(1)\n    \nimages_list1 = list(glob.glob(f'{CFG.crops_path_notile_resize}/train/*.png'))\nfor i in range(CFG.notile_resize_count):\n    images_list_notile.append(images_list1)  # 多尺度训练，数据从多个不同数据集中加载\n    for i in images_list1:\n        images_status_notile.append(2)\n        \nimages_list2 = list(glob.glob(f'{CFG.crops_path_notile_randcrop}/train/*.png'))\nfor i in range(CFG.notile_randcrop_count):\n    images_list_notile.append(images_list2)  # 多尺度训练，数据从多个不同数据集中加载\n    for i in images_list2:\n        images_status_notile.append(3)\n\nimages_list_notile = flatten_l_o_l(images_list_notile)\n\nprint(len(images_list_tiled), len(images_list_notile), len(images_status_tiled), len(images_status_notile), 'total:', len(images_list_tiled + images_list_notile))\n# print(images_list1[:5])","metadata":{"execution":{"iopub.status.busy":"2022-08-28T04:15:02.004959Z","iopub.execute_input":"2022-08-28T04:15:02.00563Z","iopub.status.idle":"2022-08-28T04:15:02.036771Z","shell.execute_reply.started":"2022-08-28T04:15:02.005594Z","shell.execute_reply":"2022-08-28T04:15:02.035234Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# tiled\ndf = pd.DataFrame({\n    \"id\": [x.split('/')[-1].split('_')[0] for x in images_list_tiled],\n    \"image_paths\": images_list_tiled,\n    \"mask_paths\": [x.replace('train', 'masks') for x in images_list_tiled],\n    \"organ\": [comp_train_df[comp_train_df['id'] == int(x.split('/')[-1].split('_')[0])]['organ'].item() for x in images_list_tiled],\n    \"status\": images_status_tiled\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})\n\n# notile\ndf2 = pd.DataFrame({\n    \"id\": [x.split('/')[-1].split('.')[0] for x in images_list_notile],\n    \"image_paths\": images_list_notile,\n    \"mask_paths\": [x.replace('train', 'masks') for x in images_list_notile],\n    \"organ\": [comp_train_df[comp_train_df['id'] == int(x.split('/')[-1].split('.')[0])]['organ'].item() for x in images_list_notile],\n    \"status\": images_status_notile\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})\n\ndf = df.append(df2, ignore_index=True)\ndf","metadata":{"execution":{"iopub.status.busy":"2022-08-28T04:15:03.951495Z","iopub.execute_input":"2022-08-28T04:15:03.951944Z","iopub.status.idle":"2022-08-28T04:15:07.625915Z","shell.execute_reply.started":"2022-08-28T04:15:03.951909Z","shell.execute_reply":"2022-08-28T04:15:07.624267Z"},"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-28T04:15:10.629019Z","iopub.execute_input":"2022-08-28T04:15:10.633151Z","iopub.status.idle":"2022-08-28T04:15:10.703778Z","shell.execute_reply.started":"2022-08-28T04:15:10.633078Z","shell.execute_reply":"2022-08-28T04:15:10.702313Z"},"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\ndata_transforms_randcrop = A.Compose([\n        A.RandomCrop(*CFG.crop_size),\n        A.Resize(*CFG.crop_resize, interpolation=cv2.INTER_NEAREST),\n        ], p=1.0)\n","metadata":{"execution":{"iopub.status.busy":"2022-08-28T04:15:16.89459Z","iopub.execute_input":"2022-08-28T04:15:16.89547Z","iopub.status.idle":"2022-08-28T04:15:16.906961Z","shell.execute_reply.started":"2022-08-28T04:15:16.895434Z","shell.execute_reply":"2022-08-28T04:15:16.905182Z"},"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                # =====================需要进行部分修改\n                if(self.df.loc[index, 'status'] == 3):  # 是需要进行rand_crop的原图\n                    data = data_transforms_randcrop(image=img, mask=mask)\n                    img  = data['image']\n                    mask  = data['mask']\n                    # print(img.shape, mask.shape)  # (512, 512, 3)   (512, 512)\n                # =====================\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-28T04:15:19.294008Z","iopub.execute_input":"2022-08-28T04:15:19.294439Z","iopub.status.idle":"2022-08-28T04:15:19.308597Z","shell.execute_reply.started":"2022-08-28T04:15:19.294393Z","shell.execute_reply":"2022-08-28T04:15:19.307088Z"},"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}\n# 本文使用的是Dice","metadata":{"execution":{"iopub.status.busy":"2022-08-28T04:15:24.566441Z","iopub.execute_input":"2022-08-28T04:15:24.566908Z","iopub.status.idle":"2022-08-28T04:15:24.582514Z","shell.execute_reply.started":"2022-08-28T04:15:24.566874Z","shell.execute_reply":"2022-08-28T04:15:24.580997Z"},"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-08-28T04:15:29.259166Z","iopub.execute_input":"2022-08-28T04:15:29.259554Z","iopub.status.idle":"2022-08-28T04:15:29.272939Z","shell.execute_reply.started":"2022-08-28T04:15:29.259522Z","shell.execute_reply":"2022-08-28T04:15:29.271469Z"},"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-28T04:15:31.894843Z","iopub.execute_input":"2022-08-28T04:15:31.895298Z","iopub.status.idle":"2022-08-28T04:15:31.904843Z","shell.execute_reply.started":"2022-08-28T04:15:31.895265Z","shell.execute_reply":"2022-08-28T04:15:31.902806Z"},"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-28T04:15:34.358432Z","iopub.execute_input":"2022-08-28T04:15:34.35887Z","iopub.status.idle":"2022-08-28T04:15:34.369536Z","shell.execute_reply.started":"2022-08-28T04:15:34.358838Z","shell.execute_reply":"2022-08-28T04:15:34.367791Z"},"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-28T04:15:36.816059Z","iopub.execute_input":"2022-08-28T04:15:36.81648Z","iopub.status.idle":"2022-08-28T04:15:36.831836Z","shell.execute_reply.started":"2022-08-28T04:15:36.816431Z","shell.execute_reply":"2022-08-28T04:15:36.830104Z"},"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-28T04:15:40.194201Z","iopub.execute_input":"2022-08-28T04:15:40.194615Z","iopub.status.idle":"2022-08-28T04:15:40.209166Z","shell.execute_reply.started":"2022-08-28T04:15:40.194583Z","shell.execute_reply":"2022-08-28T04:15:40.207681Z"},"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-28T04:15:44.108858Z","iopub.execute_input":"2022-08-28T04:15:44.109245Z","iopub.status.idle":"2022-08-28T04:15:44.126164Z","shell.execute_reply.started":"2022-08-28T04:15:44.109215Z","shell.execute_reply":"2022-08-28T04:15:44.124667Z"},"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-28T04:15:48.297104Z","iopub.execute_input":"2022-08-28T04:15:48.297505Z"},"trusted":true},"execution_count":null,"outputs":[]}]}