{"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":"code","source":"! pip install segmentation-models-pytorch\n\nfrom skimage import io, filters, transform \nimport tifffile as tiff\nimport albumentations as A\n \nimport os\nimport numpy as np\nimport pandas as pd\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nimport segmentation_models_pytorch as smp","metadata":{"papermill":{"duration":20.508921,"end_time":"2022-07-24T10:28:24.508462","exception":false,"start_time":"2022-07-24T10:28:03.999541","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-26T12:18:16.583669Z","iopub.execute_input":"2022-07-26T12:18:16.584255Z","iopub.status.idle":"2022-07-26T12:18:38.876315Z","shell.execute_reply.started":"2022-07-26T12:18:16.584122Z","shell.execute_reply":"2022-07-26T12:18:38.875141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class config:\n    IMAGE_SIZE = 1024\n    TRAIN_BATCH_SIZE = 4\n    VALID_BATCH_SIZE = 2*TRAIN_BATCH_SIZE\n    EPOCHS = 37\n    DEVICE = \"cuda\"\n    OUTPUT = \"/hubmap/\"\n    MODEL_PATH = \"nvidia/segformer-b0-finetuned-cityscapes-1024-1024\"\n    FOLDS = 5\n    LR = 8e-5\n    \n#     MEAN =  [0.78036435, 0.75635034, 0.77327976]\n#     STD = [0.24925208, 0.26279064, 0.258655 ] ","metadata":{"papermill":{"duration":0.012983,"end_time":"2022-07-24T10:28:24.526294","exception":false,"start_time":"2022-07-24T10:28:24.513311","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-26T12:18:38.878562Z","iopub.execute_input":"2022-07-26T12:18:38.879824Z","iopub.status.idle":"2022-07-26T12:18:38.886288Z","shell.execute_reply.started":"2022-07-26T12:18:38.879782Z","shell.execute_reply":"2022-07-26T12:18:38.885191Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HubDataset(torch.utils.data.Dataset):\n    def __init__(self,image_path,mask_path, augmentations=None):\n        self.image_path = image_path\n        self.mask_path = mask_path\n        self.augmentations = augmentations\n        \n    def __len__(self):\n        return len(self.image_path)\n    \n    def __getitem__(self,item):\n        \n        image = tiff.imread(self.image_path[item])\n        mask = io.imread(self.mask_path[item])\n        mask = mask.reshape(mask.shape[0],mask.shape[1],1)\n        \n        \n        ## resize\n        image = transform.resize(image, (config.IMAGE_SIZE, config.IMAGE_SIZE, 3))\n        mask = transform.resize(mask, (config.IMAGE_SIZE, config.IMAGE_SIZE, 1))\n        \n        image = image - np.min(image)\n        image = image / np.max(image)\n        \n        if self.augmentations is not None:\n            augmented = self.augmentations(image=image, mask=mask)\n            image = augmented[\"image\"]\n            mask = augmented[\"mask\"]    \n        image = np.transpose(image, (2, 0, 1)).astype(np.float32)\n        mask = np.transpose(mask, (2, 0, 1)).astype(np.float32)\n        \n        return {\n            \"image\": torch.tensor(image, dtype=torch.float),\n            \"mask\" : torch.tensor(mask, dtype=torch.float),\n        }","metadata":{"papermill":{"duration":0.017721,"end_time":"2022-07-24T10:28:24.548370","exception":false,"start_time":"2022-07-24T10:28:24.530649","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-26T12:18:38.887330Z","iopub.execute_input":"2022-07-26T12:18:38.888184Z","iopub.status.idle":"2022-07-26T12:18:38.902615Z","shell.execute_reply.started":"2022-07-26T12:18:38.888147Z","shell.execute_reply":"2022-07-26T12:18:38.901426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def dice_metric(_mask1,_mask2):\n    \n    batch_size = _mask1.shape[0]\n    dice_total = 0.0\n    \n    for idx in range(batch_size):\n        mask1 = _mask1[idx].reshape(_mask1[idx].shape[1],_mask1[idx].shape[2]).cpu().numpy()\n        mask2 = _mask2[idx].reshape(_mask2[idx].shape[1],_mask2[idx].shape[2]).cpu().numpy()\n        intersect = np.sum(mask1*mask2)\n        fsum = np.sum(mask1)\n        ssum = np.sum(mask2)\n        eps = 1e-7 ##  for empty masks the numerator should not be divivded by zero\n        dice = (2 * intersect + eps) / (fsum + ssum + eps)\n        dice = np.mean(dice)\n        \n        dice_total += dice\n        \n    final_dice = dice_total/batch_size\n    final_dice = round(final_dice, 4)  # for easy reading till 4 decimal places\n    return final_dice  ","metadata":{"papermill":{"duration":0.016353,"end_time":"2022-07-24T10:28:24.569252","exception":false,"start_time":"2022-07-24T10:28:24.552899","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-26T12:18:38.905919Z","iopub.execute_input":"2022-07-26T12:18:38.906331Z","iopub.status.idle":"2022-07-26T12:18:38.915273Z","shell.execute_reply.started":"2022-07-26T12:18:38.906306Z","shell.execute_reply":"2022-07-26T12:18:38.914078Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\n\nJaccardLoss = smp.losses.JaccardLoss(mode='binary',smooth=1.0)\nDiceLoss = smp.losses.DiceLoss(mode='binary', smooth=1.0)\nFocalLoss = smp.losses.FocalLoss(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 criterion(y_pred, y_true):\n    return DiceLoss(y_pred, y_true)\n\nscaler = torch.cuda.amp.GradScaler()\n\ndef train(model,train_loader,device,optimizer):\n    model.train()\n    running_train_loss = 0.0\n    for data in train_loader:\n        inputs = data['image']\n        masks = data['mask']\n\n        ## CUDA\n        inputs = inputs.to(device, dtype=torch.float)\n        masks = masks.to(device, dtype=torch.float)\n\n        ## forward \n        with torch.cuda.amp.autocast():\n            outputs = model(inputs,)\n            loss = criterion(outputs, masks)\n\n        optimizer.zero_grad()\n        scaler.scale(loss).backward()\n        scaler.step(optimizer)\n        scaler.update()\n\n        running_train_loss +=loss.item()\n        \n    train_loss_value = running_train_loss/len(train_loader)\n    print(f'train custom Loss is {train_loss_value}')\n\n@torch.no_grad()    \ndef evaluation(model,valid_loader,device,optimizer):\n    model.eval()\n    running_dice_score = 0.0\n    running_val_loss = 0.0\n    for data in valid_loader:\n        inputs = data['image']\n        masks = data['mask']\n        \n        inputs = inputs.to(device, dtype=torch.float)\n        masks = masks.to(device, dtype=torch.float)\n\n        output = model(inputs,)\n        running_val_loss +=  criterion(output, masks)\n        \n        output = torch.sigmoid(output)\n        running_dice_score += dice_metric(masks, output)\n\n    val_loss = running_val_loss/len(valid_loader) \n    dice_score = running_dice_score/len(valid_loader)\n        \n    print(f'valid custom Loss is {val_loss}')\n        \n    return dice_score","metadata":{"papermill":{"duration":0.080248,"end_time":"2022-07-24T10:28:24.653806","exception":false,"start_time":"2022-07-24T10:28:24.573558","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-26T12:18:38.916944Z","iopub.execute_input":"2022-07-26T12:18:38.917427Z","iopub.status.idle":"2022-07-26T12:18:38.994083Z","shell.execute_reply.started":"2022-07-26T12:18:38.917391Z","shell.execute_reply":"2022-07-26T12:18:38.993000Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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 *F.interpolate(x, scale_factor=self.scale_factor, mode='bilinear', align_corners=False) \\\n            + (1-self.mixing )*F.interpolate(x, scale_factor=self.scale_factor, mode='nearest')\n        return x","metadata":{"execution":{"iopub.status.busy":"2022-07-26T12:18:38.995571Z","iopub.execute_input":"2022-07-26T12:18:38.995995Z","iopub.status.idle":"2022-07-26T12:18:39.008390Z","shell.execute_reply.started":"2022-07-26T12:18:38.995938Z","shell.execute_reply":"2022-07-26T12:18:39.007415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport segmentation_models_pytorch as smp\n\nfrom transformers import SegformerForSemanticSegmentation\n\nclass HubmapModel(nn.Module):\n    def __init__(self):\n        super(HubmapModel, self).__init__()\n\n        self.model = SegformerForSemanticSegmentation.from_pretrained(config.MODEL_PATH,\n                                                         num_labels=1,ignore_mismatched_sizes=True)\n        self.mixup = MixUpSample()\n    def forward(self, image):\n        img_segs = self.model(image)\n\n        upsampled_logits = self.mixup(img_segs.logits)\n        return upsampled_logits\n\n\n","metadata":{"papermill":{"duration":0.560097,"end_time":"2022-07-24T10:28:25.218820","exception":false,"start_time":"2022-07-24T10:28:24.658723","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-26T12:18:39.010522Z","iopub.execute_input":"2022-07-26T12:18:39.011099Z","iopub.status.idle":"2022-07-26T12:18:39.637382Z","shell.execute_reply.started":"2022-07-26T12:18:39.011043Z","shell.execute_reply":"2022-07-26T12:18:39.636251Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## MAIN","metadata":{"papermill":{"duration":0.004752,"end_time":"2022-07-24T10:28:25.228675","exception":false,"start_time":"2022-07-24T10:28:25.223923","status":"completed"},"tags":[]}},{"cell_type":"code","source":"df = pd.read_csv(\"../input/hubmap-folds/train_unsplit_data.csv\")","metadata":{"papermill":{"duration":0.296782,"end_time":"2022-07-24T10:28:25.530010","exception":false,"start_time":"2022-07-24T10:28:25.233228","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-26T12:18:39.638713Z","iopub.execute_input":"2022-07-26T12:18:39.639262Z","iopub.status.idle":"2022-07-26T12:18:40.013487Z","shell.execute_reply.started":"2022-07-26T12:18:39.639227Z","shell.execute_reply":"2022-07-26T12:18:40.012364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# from sklearn import model_selection\n# df_train, df_valid = model_selection.train_test_split(df, test_size=0.3,)","metadata":{"papermill":{"duration":0.011347,"end_time":"2022-07-24T10:28:25.545889","exception":false,"start_time":"2022-07-24T10:28:25.534542","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-26T12:18:40.014900Z","iopub.execute_input":"2022-07-26T12:18:40.015309Z","iopub.status.idle":"2022-07-26T12:18:40.020889Z","shell.execute_reply.started":"2022-07-26T12:18:40.015273Z","shell.execute_reply":"2022-07-26T12:18:40.019913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for fold in {3}:\n    best_score = 0.0\n\n    model = HubmapModel()\n    model.to(\"cuda\")\n    \n    df_train = df[df.fold != fold].reset_index(drop=True)\n    df_valid = df[df.fold == fold].reset_index(drop=True)\n\n    df_train = df_train.drop(columns = 'fold')\n    df_valid = df_valid.drop(columns = 'fold')\n\n    train_ids = df_train.id.values.tolist()\n    valid_ids = df_valid.id.values.tolist()\n\n    train_images = [os.path.join(\"../input/hubmap-organ-segmentation/train_images\",str(i) + \".tiff\") for i in train_ids]\n    train_masks =  [os.path.join(\"../input/hubmap-hpa-2022-maskdataset/hubmap_2022_MaskDataset\",str(i) + \".png\") for i in train_ids]\n\n    valid_images = [os.path.join(\"../input/hubmap-organ-segmentation/train_images\",str(i) + \".tiff\") for i in valid_ids]\n    valid_masks =  [os.path.join(\"../input/hubmap-hpa-2022-maskdataset/hubmap_2022_MaskDataset\",str(i) + \".png\") for i in valid_ids]\n\n    train_dataset = HubDataset(image_path=train_images,mask_path=train_masks)\n    train_loader = torch.utils.data.DataLoader(train_dataset,batch_size=config.TRAIN_BATCH_SIZE,shuffle=True,pin_memory=True) \n    valid_dataset = HubDataset(image_path=valid_images, mask_path=valid_masks)\n    valid_loader = torch.utils.data.DataLoader(valid_dataset,batch_size=config.VALID_BATCH_SIZE,shuffle=False,pin_memory=True) \n\n    optimizer = torch.optim.Adam(model.parameters(),lr=config.LR,)\n    scheduler = torch.optim.lr_scheduler.ExponentialLR(optimizer, gamma=0.8, verbose=True)\n\n    print(f'============================== FOLD -- {fold} ==============================')\n\n    for epoch in range(config.EPOCHS):\n        print(f'==================== Epoch -- {epoch} ====================')\n        train(model=model,train_loader=train_loader,device=config.DEVICE,optimizer=optimizer)\n        scheduler.step()\n        dice_score = evaluation(model=model,valid_loader=valid_loader,device=config.DEVICE,optimizer=optimizer)\n\n        print(f'validation DICE Metric={dice_score}')\n        \n        if dice_score > best_score:\n            best_score = dice_score\n            torch.save(model.state_dict(),'model-'+str(fold) +'.pth')\n","metadata":{"papermill":{"duration":9165.724466,"end_time":"2022-07-24T13:01:11.274681","exception":false,"start_time":"2022-07-24T10:28:25.550215","status":"completed"},"tags":[],"execution":{"iopub.status.busy":"2022-07-26T12:18:40.024606Z","iopub.execute_input":"2022-07-26T12:18:40.025439Z"},"trusted":true},"execution_count":null,"outputs":[]}]}