{"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":"## References:\n1. Fine-Tune a Semantic Segmentation Model with a Custom Dataset: https://huggingface.co/blog/fine-tune-segformer\n\n2. Notebook: https://github.com/NielsRogge/Transformers-Tutorials/blob/master/SegFormer/Fine_tune_SegFormer_on_custom_dataset.ipynb\n\n","metadata":{}},{"cell_type":"code","source":"!pip install -q transformers datasets","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-07T14:06:59.287729Z","iopub.execute_input":"2022-09-07T14:06:59.288158Z","iopub.status.idle":"2022-09-07T14:07:09.160720Z","shell.execute_reply.started":"2022-09-07T14:06:59.288123Z","shell.execute_reply":"2022-09-07T14:07:09.159158Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !pip install /kaggle/input/staintools-offline/spams-2.6.5.4-cp37-cp37m-linux_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.164606Z","iopub.execute_input":"2022-09-07T14:07:09.165323Z","iopub.status.idle":"2022-09-07T14:07:09.171184Z","shell.execute_reply.started":"2022-09-07T14:07:09.165287Z","shell.execute_reply":"2022-09-07T14:07:09.169882Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Doubts (Work in progress)\n\n**This Version**\n1. Training with unstained data\n\n**To DO**\n\n1. Incorporating Stocastic Weight Averaging for better generalisation\n2. Resize using scale factor and cv2.INTER_CUBIC (Multi zoom?)\n3. Change segformer config -> increase attention heads, depth and add dropout. \n4. Try with 1024 x 1024 ( though as per discussion 768 - > 1024 0.002LB+ )\n\n5. Lung predicitons extremely poor (Top 100 get 0.2, our output is blank) -> External data?\n6. Insted of mixUpsampling try only nearest.","metadata":{}},{"cell_type":"markdown","source":"## Importing Libraries","metadata":{}},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os\nimport matplotlib.pyplot as plt\nfrom skimage import io,img_as_float\nimport cv2\nimport torch\nfrom torch.utils.data import Dataset,DataLoader\nfrom sklearn.model_selection import StratifiedKFold, KFold, StratifiedGroupKFold\n# import rasterio\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport random\nfrom torchvision.utils import make_grid\nfrom torchvision.io import read_image\nfrom pathlib import Path\nimport torchvision.transforms.functional as F\nimport time\nimport copy\nfrom collections import defaultdict\nimport gc\n# import segmentation_models_pytorch as smp\nfrom torch.optim import lr_scheduler\nimport transformers \n\nimport torch.nn as nn\nfrom tqdm import tqdm\nfrom torchmetrics.functional import dice\nimport torch.optim as optim\nimport argparse\n\nfrom yaml import parse\n\nfrom torch.cuda import amp\nimport albumentations as A\nfrom torch.nn import functional as nnF\n\nfrom colorama import Fore, Back, Style\nc_  = Fore.GREEN\nlo_  = Fore.RED\nsr_ = Style.RESET_ALL\n\nplt.rcParams[\"savefig.bbox\"] = 'tight'","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.172934Z","iopub.execute_input":"2022-09-07T14:07:09.173948Z","iopub.status.idle":"2022-09-07T14:07:09.187411Z","shell.execute_reply.started":"2022-09-07T14:07:09.173908Z","shell.execute_reply":"2022-09-07T14:07:09.186286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config","metadata":{}},{"cell_type":"code","source":"class CFG:\n    def __init__(self,n_fold = 5,seed = 42,batch_size = 4,debug=False,kaggle = True):\n        \n        # Directories\n        self.root_dir = \"../input/hubmap-organ-segmentation\" if kaggle else '/media/harish/4tb_hdd/mustaffa/saharsh_hubmap/hubmap-organ-segmentation'\n#         self.IMAGES   = '../input/hubmap-2022-256x256/train' if kaggle else '/media/harish/4tb_hdd/mustaffa/saharsh_hubmap/hubmap-2022-256x256/train' \n#         self.MASKS    = \"../input/hubmap-2022-256x256/masks\" if kaggle else '/media/harish/4tb_hdd/mustaffa/saharsh_hubmap/hubmap-2022-256x256/masks'\n        \n        self.IMAGES   = '../input/hubmap-stain-norm-original/train_images_stain_normalised' #'../input/hubmap-hpa-2022-png-dataset/train_images_png'#'../input/hubmap-organ-segmentation/train_images'\n        self.MASKS    = '../input/hubmap-hpa-2022-png-dataset/train_masks_png'\n        \n        self.fold_no       = 0\n        self.n_fold        = n_fold\n        self.folds_to_run  = [0,1,2,3,4]\n        self.seed          = seed\n        self.batch_size    = batch_size\n        self.train_bs      = batch_size\n        self.valid_bs      = batch_size*2\n        self.debug         = debug\n        self.img_size      = [1024, 1024] #..[768,768] # [1617,1617] -> try with batch size of 1\n        self.exp_name      = 'Hubmap256-training'\n        self.epochs        = 50\n        self.lr            = 5e-5 #as per LRRT 4.64e-4 for DICEBCE #0.00085\n        \n        self.optimizer     = 'Adam'\n        self.weight_decay  = 1e-6\n        # For scheduler\n        self.scheduler     = \"ReduceLROnPlateau\"\n        self.min_lr        = 1e-6 # 0.00005\n        self.T_max         = int(280/self.batch_size*self.epochs)+50\n        self.T_0           = 25\n        self.warmup_epochs = 0\n        self.wd            = 1e-6\n        self.n_accumulate  = 1 #max(1, 32//self.batch_size)\n        \n        self.num_classes   = 1\n        self.device        = torch.device(\"cuda:0\" if torch.cuda.is_available() else \"cpu\")\n        \n        self.hub_model_id  = 'nvidia/mit-b2'\n        self.model_name    = \"segformer\"\n        \n    def display(self):\n        print(f\"{self.exp_name}\")\n        print(f\"debug is {self.debug}\")        \n        print(f\"Batch size is  {self.batch_size}\")\n        print(f\"img_size is    {self.img_size}\")\n        print(f\"fold_no is     {self.fold_no}\")\n        print(f\"Model is       {self.model_name}_{self.hub_model_id}\")\n        print(f\"epochs is      {self.epochs}\")\n        print(f\"Optimiser is   {self.optimizer}\")\n        print(f\"Scheduler is   {self.scheduler}\")\n        print(f\"LR is          {self.lr}\")\n        \n    \n\ndef set_seed(seed = 42):\n    '''Sets the seed of the entire notebook so results are the same every time we run.\n    This is for REPRODUCIBILITY.'''\n    print(f\"Setting seed as {seed}\")\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\n\ndef initialize_config(debug=False,batch_size=4):\n    cfg = CFG(batch_size = batch_size,debug = debug)    \n    set_seed(cfg.seed)\n    cfg.display()\n    return cfg\n","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.190997Z","iopub.execute_input":"2022-09-07T14:07:09.191947Z","iopub.status.idle":"2022-09-07T14:07:09.208400Z","shell.execute_reply.started":"2022-09-07T14:07:09.191912Z","shell.execute_reply":"2022-09-07T14:07:09.207449Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Utility Functions\n* create folds -> Splits data using stratified kfolds\n* TO DO: Implement Albumentation transformation with huggingface dataset class","metadata":{}},{"cell_type":"code","source":"def create_folds(cfg = None):\n    \n    train_csv = pd.read_csv(os.path.join(cfg.root_dir,\"train.csv\"))\n    train_csv.set_index('id',inplace = True)\n    \n    images_path = os.listdir(cfg.IMAGES)\n    images_path = sorted(images_path)\n    masks_path =  os.listdir(cfg.MASKS)\n    masks_path = sorted(masks_path)\n    \n    ids  = [filename[:-4] for filename in masks_path]\n    \n#     files = [filename[:-4] for filename in masks_path]\n#     ids = [f[:-5] for f in files]\n\n    organs = [train_csv['organ'][int(idx)] for idx in ids]\n    images_path = [os.path.join(cfg.IMAGES,f) for f in images_path]\n    masks_path = [os.path.join(cfg.MASKS,f) for f in masks_path]\n   \n    maping = {\n        'id': ids,\n#         'id_grid':files,\n        'organ':organs,\n        'image_path':images_path,\n        'mask_path': masks_path\n\n    }\n    df = pd.DataFrame.from_dict(maping)    \n    skf = KFold(n_splits=cfg.n_fold, shuffle=True,random_state=cfg.seed)\n\n    df.loc[:,'fold']=-1\n    for f,(t_idx, v_idx) in enumerate(skf.split(X=df['id'], y=df['organ'])):\n        df.iloc[v_idx,-1]=f\n\n    return df\n\ndef get_optimizer(cfg,optimizer_name= 'Adam'):\n    if optimizer_name == 'Adam':\n        optimizer = optim.Adam(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\n    \n    elif optimizer_name == 'AdamW':\n        optimizer = optim.AdamW(model.parameters(), lr=cfg.lr, weight_decay=cfg.weight_decay)\n        \n    return optimizer\n\n\ndef get_scheduler(cfg,optimizer,df):\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\n    if cfg.scheduler == 'CosineAnnealingLR':\n        scheduler = lr_scheduler.CosineAnnealingLR(optimizer,T_max=cfg.T_max,eta_min=cfg.min_lr)\n        \n    elif cfg.scheduler == 'ReduceLROnPlateau':\n        scheduler = lr_scheduler.ReduceLROnPlateau( optimizer  = optimizer, verbose=True, factor=0.7,mode=\"max\",patience=5, threshold=0.001,threshold_mode = 'abs',min_lr =1e-6)\n        \n    elif cfg.scheduer == 'ExponentialLR':\n        scheduler = lr_scheduler.ExponentialLR(optimizer, gamma=0.85)\n        \n    return scheduler\n","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.209870Z","iopub.execute_input":"2022-09-07T14:07:09.210490Z","iopub.status.idle":"2022-09-07T14:07:09.227191Z","shell.execute_reply.started":"2022-09-07T14:07:09.210446Z","shell.execute_reply":"2022-09-07T14:07:09.226139Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Augmentations**","metadata":{}},{"cell_type":"code","source":"# --------- Augmnetation functions LB = 0.6-------\ndef do_random_flip(image, mask):\n    if np.random.rand()>0.5:\n        image = cv2.flip(image,0)\n        mask = cv2.flip(mask,0)\n    if np.random.rand()>0.5:\n        image = cv2.flip(image,1)\n        mask = cv2.flip(mask,1)\n    if np.random.rand()>0.5:\n        image = image.transpose(1,0,2)\n        mask = mask.transpose(1,0)\n    \n    image = np.ascontiguousarray(image)\n    mask = np.ascontiguousarray(mask)\n    return image, mask\n\ndef do_random_rot90(image, mask):\n    r = np.random.choice([\n        0,\n        cv2.ROTATE_90_CLOCKWISE,\n        cv2.ROTATE_90_COUNTERCLOCKWISE,\n        cv2.ROTATE_180,\n    ])\n    if r==0:\n        return image, mask\n    else:\n        image = cv2.rotate(image, r)\n        mask = cv2.rotate(mask, r)\n        return image, mask\n    \ndef do_random_contast(image, mask, mag=0.3):\n    alpha = 1 + random.uniform(-1,1)*mag\n    image = image * alpha\n    image = np.clip(image,0,1)\n    return image, mask\n\ndef do_random_hsv(image, mask, mag=[0.15,0.25,0.25]):\n    image = (image*255).astype(np.uint8)\n    hsv = cv2.cvtColor(image, cv2.COLOR_BGR2HSV)\n\n    h = hsv[:, :, 0].astype(np.float32)  # hue\n    s = hsv[:, :, 1].astype(np.float32)  # saturation\n    v = hsv[:, :, 2].astype(np.float32)  # value\n    h = (h*(1 + random.uniform(-1,1)*mag[0]))%180\n    s =  s*(1 + random.uniform(-1,1)*mag[1])\n    v =  v*(1 + random.uniform(-1,1)*mag[2])\n\n    hsv[:, :, 0] = np.clip(h,0,180).astype(np.uint8)\n    hsv[:, :, 1] = np.clip(s,0,255).astype(np.uint8)\n    hsv[:, :, 2] = np.clip(v,0,255).astype(np.uint8)\n    image = cv2.cvtColor(hsv, cv2.COLOR_HSV2BGR)\n    image = image.astype(np.float32)/255\n    return image, mask\n\ndef do_random_noise(image, mask, mag=0.1):\n    height, width = image.shape[:2]\n    noise = np.random.uniform(-1,1, (height, width,1))*mag\n    image = image + noise\n    image = np.clip(image,0,1)\n    return image, mask\n\ndef do_random_rotate_scale(image, mask, angle=30, scale=[0.8,1.2] ):\n    angle = np.random.uniform(-angle, angle)\n    scale = np.random.uniform(*scale) if scale is not None else 1\n    \n    height, width = image.shape[:2]\n    center = (height // 2, width // 2)\n    \n    transform = cv2.getRotationMatrix2D(center, angle, scale)\n    image = cv2.warpAffine( image, transform, (width, height), flags=cv2.INTER_LINEAR,\n                            borderMode=cv2.BORDER_CONSTANT, borderValue=(0,0,0))\n    mask  = cv2.warpAffine( mask, transform, (width, height), flags=cv2.INTER_LINEAR,\n                            borderMode=cv2.BORDER_CONSTANT, borderValue=0)\n    return image, mask\n\ndef elastic_transform(image,alpha, sigma, alpha_affine, random_state=None):\n    \"\"\"Elastic deformation of images as described in [Simard2003]_ (with modifications).\n    .. [Simard2003] Simard, Steinkraus and Platt, \"Best Practices for\n         Convolutional Neural Networks applied to Visual Document Analysis\", in\n         Proc. of the International Conference on Document Analysis and\n         Recognition, 2003.\n\n     Based on https://gist.github.com/erniejunior/601cdf56d2b424757de5\n    \"\"\"\n    if random_state is None:\n        random_state = np.random.RandomState(None)\n\n    shape = image.shape\n    shape_size = shape[:2]\n    \n    # Random affine\n    center_square = np.float32(shape_size) // 2\n    square_size = min(shape_size) // 3\n    pts1 = np.float32([center_square + square_size, [center_square[0]+square_size, center_square[1]-square_size], center_square - square_size])\n    pts2 = pts1 + random_state.uniform(-alpha_affine, alpha_affine, size=pts1.shape).astype(np.float32)\n    M = cv2.getAffineTransform(pts1, pts2)\n    image = cv2.warpAffine(image, M, shape_size[::-1], borderMode=cv2.BORDER_REFLECT_101)\n\n    dx = gaussian_filter((random_state.rand(*shape) * 2 - 1), sigma) * alpha\n    dy = gaussian_filter((random_state.rand(*shape) * 2 - 1), sigma) * alpha\n    dz = np.zeros_like(dx)\n\n    x, y, z = np.meshgrid(np.arange(shape[1]), np.arange(shape[0]), np.arange(shape[2]))\n    indices = np.reshape(y+dy, (-1, 1)), np.reshape(x+dx, (-1, 1)), np.reshape(z, (-1, 1))\n    print('transforming elastic. . . ')\n    im_merge_t = map_coordinates(image, indices, order=1, mode='reflect').reshape(shape)\n    \n    img = im_merge_t[:,:,:3]\n    mask = im_merge_t[:,:,3]\n    return img, mask\n\n#-----Previous Augmentation version LB = 0.5 --------------\n\n# def get_transforms(cfg=None):\n#     A.Compose([\n          \n#             #A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.05, rotate_limit=10, p=0.5),\n#             A.OneOf([\n#                 A.GridDistortion(num_steps=5, distort_limit=0.05, p=1.0),\n#                 A.OpticalDistortion(distort_limit=0.05, shift_limit=0.05, p=1.0),\n#                 A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=1.0)\n#             ], p=0.25),\n#             A.CoarseDropout(max_holes=8, max_height=cfg.img_size[0]//20, max_width=cfg.img_size[1]//20,\n#                             min_holes=5, fill_value=0, mask_fill_value=0, p=0.5),\n#             ], p=1.0)\n\n\n# def get_transforms(train = True,cfg=None):\n#     data_transforms = {\n#         \"train\":  A.Compose([\n#                     A.augmentations.transforms.ColorJitter(p=0.5),\n#                     A.OneOf([\n#                         A.OpticalDistortion(p=0.5),\n#                         A.GridDistortion(p=.5),\n#                         A.PiecewiseAffine(p=0.5),\n#                     ], p=0.5),\n#                     A.OneOf([\n#                         A.HueSaturationValue(10, 15, 10),\n#                         A.CLAHE(clip_limit=4),\n#                         A.RandomBrightnessContrast(),            \n#                     ], p=0.5),\n#                     ])              \n               \n        \n#         \"valid\": A.Compose([\n        \n#             ])\n#     }\n#     if train==True:\n#         return data_transforms[\"train\"] \n#     else:\n#         return data_transforms['valid']\n","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.230538Z","iopub.execute_input":"2022-09-07T14:07:09.230809Z","iopub.status.idle":"2022-09-07T14:07:09.259221Z","shell.execute_reply.started":"2022-09-07T14:07:09.230780Z","shell.execute_reply":"2022-09-07T14:07:09.258294Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def valid_augment5(image, mask):\n    #image, mask  = do_crop(image, mask, image_size, xy=(None,None))    \n    return image, mask\n\ndef train_augment5b(image, mask):\n    more_transform = A.Compose([\n        \n         A.OneOf([ \n             A.ElasticTransform(p=1, alpha=image.shape[1]*3, sigma=image.shape[1] * 0.07, alpha_affine=image.shape[1] * 0.09),\n             A.OpticalDistortion(p=0.5),\n             A.GridDistortion(p=.5),\n              ], p=0.5),\n         A.OneOf([             \n             A.RandomBrightnessContrast(p =0.5), \n             A.RandomGamma(p=0.5),\n                ], p=0.5), \n#          A.CoarseDropout(max_holes=6, max_height=64, max_width=64,p=0.2), # what should be the dropout tile size = 5% of image size?\n         A.CLAHE(clip_limit=4, p =0.5),\n            \n    ])\n    image, mask = do_random_flip(image, mask)\n    image, mask = do_random_rot90(image, mask)\n\n\n      \n    for fn in np.random.choice([\n        lambda image, mask: (image, mask),\n        lambda image, mask: do_random_noise(image, mask, mag=0.1),\n        lambda image, mask: do_random_contast(image, mask, mag=0.40),\n        lambda image, mask: do_random_hsv(image, mask, mag=[0.40, 0.40, 0])\n    ], 2): image, mask = fn(image, mask)\n\n    for fn in np.random.choice([\n        lambda image, mask: (image, mask),\n        lambda image, mask: do_random_rotate_scale(image, mask, angle=45, scale=[0.50, 2.0]),\n    ], 1): image, mask = fn(image, mask)\n    \n    augmented = more_transform(image= (image*255).astype(np.uint8), mask=mask)\n    image = augmented['image'].astype(np.float32)/255\n    mask = augmented['mask']\n        \n    return image, mask","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.260913Z","iopub.execute_input":"2022-09-07T14:07:09.261630Z","iopub.status.idle":"2022-09-07T14:07:09.274268Z","shell.execute_reply.started":"2022-09-07T14:07:09.261595Z","shell.execute_reply":"2022-09-07T14:07:09.273219Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset\n* prepare_loaders -> Prepares the dataset and dataloader, Trainer API needs dataset","metadata":{}},{"cell_type":"code","source":"class HuBMAPDataset(Dataset):\n    def __init__(self, train_csv= None, cfg = None, transforms = None, mode='train'):\n        self.train_csv = train_csv\n        print(len(self.train_csv))\n        # remove that one faulty image from train_csv\n        self.mode = mode\n        self.cfg = cfg\n        self.transforms = transforms       \n        self.image_paths = self.train_csv['image_path'].tolist()  \n        self.mask_paths = self.train_csv['mask_path'].tolist()\n        \n    def __len__(self):\n        return len(self.image_paths)\n    \n    def __getitem__(self,idx):\n        \n        try:\n            # Read image\n            image_size = cfg.img_size[0]\n            image = cv2.cvtColor(cv2.imread(self.image_paths[idx]), cv2.COLOR_BGR2RGB)\n\n            if self.mode == 'train':\n                \n                # Read mask\n                mask = cv2.imread(self.mask_paths[idx],cv2.IMREAD_GRAYSCALE)\n                mask  = (mask/255).astype(np.uint8) \n#                 print('Read image and mask')\n\n                # Resize \n    \n#                 s = self.pixel_size/0.4 * (image_size/image.shape[0])\n#                 print('s is:',s)\n#                 image = cv2.resize(image,dsize=None,fx=s,fy=s,interpolation = cv2.INTER_CUBIC)\n#                 mask = cv2.resize(mask,dsize=None,fx=s,fy=s,interpolation = cv2.INTER_CUBIC)\n               \n#                 image = (image - mean)/std\n                image = cv2.resize(image,dsize=(image_size,image_size),interpolation=cv2.INTER_LINEAR)\n                mask  = cv2.resize(mask, dsize=(image_size,image_size),interpolation=cv2.INTER_LINEAR)\n                              \n                image = image.astype(np.float32)/255                                   \n                \n#                 image = image - np.min(image)\n#                 image = image /(np.max(image)+0.00001)\n \n                # -------------- Sanity check------------------------\n#                 plt.imshow(image)\n#                 print('4 - Shapes: img,mask',image.shape,mask.shape)\n                # ---------------------------------------------------                \n                # -------------- Transformations --------------------\n                if self.transforms:\n                   \n                    image, mask = train_augment5b(image, mask)\n                else:\n                    image, mask = valid_augment5(image, mask)\n               \n                    \n                # ---------------------------------------------------\n                mask  = mask.astype(np.float32)\n                \n                image = np.transpose(image, (2, 0, 1))\n#                 print('5 - Shapes: img,mask, Unique',image.shape,mask.shape,np.unique(mask))\n              \n                return torch.tensor(image), torch.tensor(mask)\n           \n        except:\n            print(idx,\"error occured in dataloader/dataset\")\n            return None\n\n","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.275773Z","iopub.execute_input":"2022-09-07T14:07:09.276241Z","iopub.status.idle":"2022-09-07T14:07:09.290721Z","shell.execute_reply.started":"2022-09-07T14:07:09.276205Z","shell.execute_reply":"2022-09-07T14:07:09.289770Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"  \ndef prepare_loaders(fold,df,cfg, debug=False):\n    \n    train_df = df.query(\"fold!=@fold\").reset_index(drop=True)\n    valid_df = df.query(\"fold==@fold\").reset_index(drop=True)\n    \n    if debug:\n        train_df = train_df.head(4*5)\n        valid_df = valid_df.head(4*5)\n\n    train_dataset = HuBMAPDataset(train_df, transforms=True,cfg=cfg,mode = 'train') #get_transforms(train = True,cfg=cfg)\n    valid_dataset = HuBMAPDataset(valid_df, transforms=False,cfg=cfg, mode = 'train') #get_transforms(train = False,cfg=None)\n    \n    train_loader = DataLoader(train_dataset, batch_size=cfg.train_bs if not cfg.debug else 20, \n                            num_workers=1, shuffle=True, pin_memory=True, drop_last=False)\n    valid_loader = DataLoader(valid_dataset, batch_size=cfg.valid_bs if not cfg.debug else 20, \n                            num_workers=1, shuffle=False, pin_memory=True)\n    \n    return train_loader, valid_loader","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.292232Z","iopub.execute_input":"2022-09-07T14:07:09.292791Z","iopub.status.idle":"2022-09-07T14:07:09.303838Z","shell.execute_reply.started":"2022-09-07T14:07:09.292757Z","shell.execute_reply":"2022-09-07T14:07:09.302685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Metrics**","metadata":{}},{"cell_type":"code","source":"def 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","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.308221Z","iopub.execute_input":"2022-09-07T14:07:09.308583Z","iopub.status.idle":"2022-09-07T14:07:09.317347Z","shell.execute_reply.started":"2022-09-07T14:07:09.308544Z","shell.execute_reply":"2022-09-07T14:07:09.316145Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Dice BCE loss**","metadata":{}},{"cell_type":"code","source":"'''\nEp 1:\nDice value:  tensor(0.6109, device='cuda:0')\nBCE value:  tensor(0.7966, device='cuda:0')\n'''\n\nclass DiceBCELoss(nn.Module):\n\n    def __init__(self, weight=None, size_average=True):\n        super(DiceBCELoss, self).__init__()\n\n    def forward(self, inputs, targets, smooth=1):\n        \n        #comment out if your model contains a sigmoid or equivalent activation layer\n#         inputs = nnF.sigmoid(inputs)       \n        \n        #flatten label and prediction tensors\n        inputs = inputs.view(-1)\n        targets = targets.view(-1)        \n        \n        BCE = nnF.binary_cross_entropy_with_logits(inputs, targets, reduction='mean')\n#         return BCE\n    \n        intersection = (inputs * targets).sum()                            \n        dice_loss = 1 - (2.*intersection + smooth)/(inputs.sum() + targets.sum() + smooth)  \n#         print('Dice value: ',dice_loss)\n#         print('BCE value: ',BCE)\n        Dice_BCE = BCE + dice_loss\n        \n        return Dice_BCE\n","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.319008Z","iopub.execute_input":"2022-09-07T14:07:09.319415Z","iopub.status.idle":"2022-09-07T14:07:09.331447Z","shell.execute_reply.started":"2022-09-07T14:07:09.319342Z","shell.execute_reply":"2022-09-07T14:07:09.330341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"from transformers import SegformerForSemanticSegmentation\nimport json\nfrom huggingface_hub import cached_download, hf_hub_url\nimport torch\nfrom torch import nn\nfrom sklearn.metrics import accuracy_score\nfrom tqdm.notebook import tqdm\nfrom datasets import load_metric\n\ndef build_model():\n    \n    # Tried:\n    # 1. nvidia/b0 (768,768)\n    # 2. nvidia/b2 (768,768)\n    \n    #To do:\n    # 1. 'nvidia/segformer-b2-finetuned-cityscapes-1024-1024' and on (1024,1024) after improving LB on (768,768)\n\n    model = SegformerForSemanticSegmentation.from_pretrained('nvidia/mit-b2',num_labels=1,ignore_mismatched_sizes=True)    \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\n","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.333082Z","iopub.execute_input":"2022-09-07T14:07:09.333744Z","iopub.status.idle":"2022-09-07T14:07:09.342843Z","shell.execute_reply.started":"2022-09-07T14:07:09.333709Z","shell.execute_reply":"2022-09-07T14:07:09.341775Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch.nn.functional as Fx\n'''\nUpsampling function that combines both neareast and bilinear\n'''\nclass MixUpSample(nn.Module):\n    def __init__( self, scale_factor=4):\n        super().__init__()\n        self.mixing = nn.Parameter(torch.tensor(0.5))\n        self.scale_factor = scale_factor\n\n    def forward(self, x):\n        x = self.mixing *Fx.interpolate(x, scale_factor=self.scale_factor, mode='bilinear', align_corners=False) \\\n            + (1-self.mixing )*Fx.interpolate(x, scale_factor=self.scale_factor, mode='nearest')\n        \n        return x\n","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.344471Z","iopub.execute_input":"2022-09-07T14:07:09.345121Z","iopub.status.idle":"2022-09-07T14:07:09.355977Z","shell.execute_reply.started":"2022-09-07T14:07:09.345067Z","shell.execute_reply":"2022-09-07T14:07:09.354956Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Train epoch**","metadata":{}},{"cell_type":"code","source":"'''\nFLOW OF TRAINING:\n\n1. Load batch of data (add to gpu if accessible)\n2. forward pass\n   |_ output = model (input)\n   |_ loss   = criterion (output,target)\n   \n3. backward pass + optimize (only in training)\n   |_ loss.backward()\n   |_ optimizer.step()\n   \n4. Calculate statistics\n   |_ running_loss += loss.item()*batch_size\n\nreturn epoch_loss = running_loss/len(dataloader)\n'''\n\ndef train_one_epoch(cfg,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 = DiceBCELoss() \n    \n    pbar = tqdm(enumerate(dataloader), total=len(dataloader), desc='Train ')\n    for step, (images, masks) in pbar:  #Shapes: img: (8,3,512,512)\n\n        # get a batch of inputs\n        images = images.to(device, dtype=torch.float)\n        masks  = masks.to(device, dtype=torch.float)\n        masks =masks.unsqueeze(1)\n        batch_size = images.size(0)\n        \n        # ------------- forward ---------------\n        with amp.autocast(enabled=True):\n           \n            outputs = model(images)\n            logits  = outputs.logits # shape (batch_size, num_labels, height/4, width/4) \n            logits  = nn.Sigmoid()(logits)\n            \n#             logits = upsample_obj(logits)\n            logits = nn.functional.interpolate(logits, size=masks.shape[-2:], mode=\"nearest\")\n\n            y_pred  = logits\n            loss    = criterion(y_pred, masks)           \n            \n        # -------- backward + optimize --------\n        scaler.scale(loss).backward() #/cfg.n_accumulate\n    \n        if ((step+1)%cfg.n_accumulate==0 or (step+1)==len(dataloader)):\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        # statistics\n        running_loss += (loss.item() * batch_size) #pytorch loss.item gives average loss of the batch\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.6f}',\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-09-07T14:07:09.357438Z","iopub.execute_input":"2022-09-07T14:07:09.358394Z","iopub.status.idle":"2022-09-07T14:07:09.372183Z","shell.execute_reply.started":"2022-09-07T14:07:09.358358Z","shell.execute_reply":"2022-09-07T14:07:09.371282Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Val epoch**","metadata":{}},{"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 = DiceBCELoss()\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        masks   = masks.unsqueeze(1)\n#         print(\"Masks: \",masks.shape)\n        batch_size = images.size(0)\n        \n        outputs = model(images)\n        logits  = outputs.logits \n        logits = nn.Sigmoid()(logits)\n        logits = nn.functional.interpolate(logits, size=masks.shape[-2:], mode=\"nearest\")\n#         logits = upsample_obj(logits)\n#         print(\"Logits: \",logits.shape)\n#         logits  = nn.functional.interpolate(logits, size=masks.shape[-2:], mode=\"bilinear\", align_corners=False) #upsample to masks size\n        y_pred  = logits\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        \n#         print(\"Ypred: \",y_pred.shape)\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.6f}',\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-09-07T14:07:09.373528Z","iopub.execute_input":"2022-09-07T14:07:09.374522Z","iopub.status.idle":"2022-09-07T14:07:09.387994Z","shell.execute_reply.started":"2022-09-07T14:07:09.374486Z","shell.execute_reply":"2022-09-07T14:07:09.387028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Run training**","metadata":{}},{"cell_type":"code","source":"import copy\ndef run_training(cfg,fold, 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(cfg,model, optimizer, scheduler,dataloader=train_loader,device=cfg.device, epoch=epoch)\n        \n        val_loss, val_scores = valid_one_epoch(model, valid_loader,device=cfg.device, epoch=epoch)\n        \n        val_dice, val_jaccard = val_scores\n        \n        #ReduceLR scheduler\n        if scheduler :\n#             print('not none scheduler...')\n            scheduler.step(val_dice) #Monitor metric\n            \n        \n        history['Fold'].append(fold)\n        history['Epoch'].append(epoch)\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        pd.DataFrame(history).to_csv('./segformer_mit-b2_fold_0_1024_lr-1e-4_200ep.csv', index=False)\n        \n        print(f'Valid Dice: {val_dice:0.4f} | Valid Jaccard: {val_jaccard:0.4f}')\n        print(f'{lo_}Valid Loss: {val_loss:0.4f} | Train Loss: {train_loss:0.4f}{sr_}')\n        # deep copy the model\n        if val_dice >= best_dice:\n            \n            print(f\"{c_}Valid Dice Score Improved ({best_dice:0.4f} ---> {val_dice:0.4f}){sr_}\")\n            best_dice    = val_dice\n            best_jaccard = val_jaccard\n            best_epoch   = epoch\n            \n            \n            best_model_wts = copy.deepcopy(model.state_dict())\n            PATH = f\"segformer_mit-b2_1024_lr-1e-4_200ep-{best_epoch}.bin\"            \n            \n#         if epoch > 7:\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\"segformer_mit-b2_1024_last_epoch.bin\"\n        torch.save(model.state_dict(), PATH)\n            \n        print(); print(f'{sr_}')\n        \n        # ---------Add patience factor ------------\n        \n        # -----------------------------------------\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-09-07T14:07:09.391072Z","iopub.execute_input":"2022-09-07T14:07:09.391371Z","iopub.status.idle":"2022-09-07T14:07:09.407186Z","shell.execute_reply.started":"2022-09-07T14:07:09.391347Z","shell.execute_reply":"2022-09-07T14:07:09.406130Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **Check DataLoader**","metadata":{}},{"cell_type":"code","source":"# # Sanity check Dataloader\n# mean = np.array([0.485,0.456,0.406])\n# std = np.array([0.229,0.224,0.225])\n\n\nupsample_obj = MixUpSample()\n\n# from torch_lr_finder import LRFinder\nfold = 0\ncfg  = initialize_config(debug=False,batch_size=2)\ndf   = create_folds(cfg=cfg)\npd.set_option('display.max_colwidth', None)\ndisplay(df.head(5))\n\n","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.410481Z","iopub.execute_input":"2022-09-07T14:07:09.410931Z","iopub.status.idle":"2022-09-07T14:07:09.605126Z","shell.execute_reply.started":"2022-09-07T14:07:09.410904Z","shell.execute_reply":"2022-09-07T14:07:09.604144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check Tensor shapes\ni = 0\nfor i in range(5):\n    train_loader, valid_loader = prepare_loaders(fold=fold,df = df,cfg=cfg)\n    batch = next(iter(train_loader))\n    images, labels = batch\n#     print(images.shape, labels.shape, type(images), type(labels), images.dtype, labels.dtype)\n\n    # Sanity check sample n=1 \n    testImg = images[0]\n    testMsk = labels[0]\n    print(\"Img: \",testImg.shape, testImg.dtype, type(testImg))\n    print(\"Mask: \",testMsk.shape,testMsk.dtype, type(testMsk))\n\n    # Plot exmaple mask \n    x=testMsk #.permute(1,2,0)\n    x=x[:,:]\n\n\n    plt.figure(figsize=(8,8))\n    plt.subplot(1, 2, 1)\n    plt.imshow((testImg.permute(1,2,0)), interpolation='none')\n    plt.subplot(1, 2, 2)\n    plt.imshow(testImg.permute(1,2,0), 'gray', interpolation='none')\n    plt.imshow(x, 'jet', interpolation='none', alpha=0.7)\n    plt.show()\n    del testImg,x,testMsk","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:09.606639Z","iopub.execute_input":"2022-09-07T14:07:09.607512Z","iopub.status.idle":"2022-09-07T14:07:31.747660Z","shell.execute_reply.started":"2022-09-07T14:07:09.607473Z","shell.execute_reply":"2022-09-07T14:07:31.744160Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## **MAIN**","metadata":{}},{"cell_type":"code","source":"# Initlaisations\n\n# mean = np.array([0.485,0.456,0.406])\n# std = np.array([0.229,0.224,0.225])\nupsample_obj = MixUpSample()\n\nfor fold in range(1):\n    print(f'#'*15*2)\n    print(f'### Fold: {fold}')\n    print(f'#'*15*2)\n    \n    cfg                        = initialize_config(debug=False,batch_size=2)\n    df                         = create_folds(cfg=cfg)\n    train_loader, valid_loader = prepare_loaders(fold=fold,df = df,cfg=cfg)\n    model                      = build_model().to(cfg.device) \n    optimizer                  = get_optimizer(cfg,optimizer_name=cfg.optimizer)\n    scheduler                  = get_scheduler(cfg,optimizer,df)\n    model, history             = run_training(cfg,fold, model, optimizer, scheduler,device=cfg.device,num_epochs=cfg.epochs)\n    \n    \n    #Plot 1\n    #plt.subplot(3, 1, 1)\n    plt.plot(history['Epoch'], history['Train Loss'], 'r--')\n    plt.plot(history['Epoch'], history['Valid Loss'], 'b-')\n    plt.legend(['Training Loss', 'Valid Loss'])\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.savefig('./Epoch_vs_Loss.png')\n    plt.show()\n\n    #Plot 2\n    #plt.subplot(3,1, 2)\n    plt.plot(history['Epoch'], history['Valid Dice'])\n    plt.xlabel('Epoch')\n    plt.ylabel('Valid Dice')\n    plt.savefig('./Epoch_vs_Valid-Dice.png')\n    plt.show()\n\n    #Plot 3\n    #plt.subplot(3,1, 3)\n    plt.plot(history['Valid Dice'], history['Train Loss'] )\n    plt.xlabel('Valid Dice') \n    plt.ylabel('Train Loss')\n    plt.savefig('./Valid-Dice_vs_Loss.png')\n    plt.show()\n    \n    \n    \n   ","metadata":{"execution":{"iopub.status.busy":"2022-09-07T14:07:31.749854Z","iopub.execute_input":"2022-09-07T14:07:31.750280Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}