{"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":"Reference:\n\nhttps://www.kaggle.com/code/abebe9849/swin-v2-unet-upernet","metadata":{}},{"cell_type":"code","source":"import sys\nsys.path.append('../input/pytorch-modules/segmentation_models.pytorch-master/segmentation_models.pytorch-master/')\nsys.path.append('../input/pytorch-modules/pytorch-image-models-master/pytorch-image-models-master/')\nsys.path.append('../input/pytorch-modules/pretrained-models.pytorch-master/pretrained-models.pytorch-master/')\nsys.path.append('../input/pytorch-modules/EfficientNet-PyTorch-master/EfficientNet-PyTorch-master/')\nsys.path.append('../input/pytorch-modules/lightning-master/lightning-master/')\nsys.path.append('../input/swin-transformer-v2/')\nimport timm\nimport torch\nfrom torch.utils import data as torch_data\nfrom torch.utils.data import DataLoader\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom  timm.models.layers.cbam import *\nfrom swin_v2 import unet_swin\nimport segmentation_models_pytorch as smp\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport pandas as pd \nimport numpy as np\nimport os\nimport cv2\nimport gc\nimport random\nfrom glob import glob\nfrom PIL import Image\nfrom tqdm.notebook import tqdm\nimport matplotlib.gridspec as gridspec\nimport matplotlib.patches as mpatches\nimport matplotlib.pyplot as plt \nfrom sklearn.model_selection import StratifiedGroupKFold\nimport os\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n\ndef set_seed(seed):\n    random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    if torch.cuda.is_available():\n        torch.cuda.manual_seed_all(seed)\n        torch.backends.cudnn.deterministic = True\n\nset_seed(1)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-09-14T22:14:33.720261Z","iopub.execute_input":"2022-09-14T22:14:33.721233Z","iopub.status.idle":"2022-09-14T22:14:39.333755Z","shell.execute_reply.started":"2022-09-14T22:14:33.721142Z","shell.execute_reply":"2022-09-14T22:14:39.332409Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_splits = 5\nfold_selected = [1,2,4]\nim_width = 256\nim_height = 256\npr = 0.50","metadata":{"execution":{"iopub.status.busy":"2022-09-14T22:14:39.335319Z","iopub.execute_input":"2022-09-14T22:14:39.336030Z","iopub.status.idle":"2022-09-14T22:14:39.344468Z","shell.execute_reply.started":"2022-09-14T22:14:39.335985Z","shell.execute_reply":"2022-09-14T22:14:39.343065Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train = pd.read_csv('../input/hubmap-organ-segmentation/train.csv')\ntrain['path'] = ['../input/hubmap-organ-segmentation/train_images/'+str(x)+'.tiff' for x in train.id]\ntrain['count'] = train.groupby('organ').agg('count')['rle']\ntrain['count'].fillna(0, inplace=True)\ntrain.reset_index(inplace=True)\n\nskf = StratifiedGroupKFold(n_splits=n_splits, shuffle=True, random_state=42)\nfor fold, (_, val_idx) in enumerate(skf.split(X=train, y=train['count'], groups=train['age']), 1):\n    train.loc[val_idx, 'fold'] = fold\n\ntrain_ids = train[train[\"fold\"]!=1].index\nvalid_ids = train[train[\"fold\"]==1].index\n\nX_train = train[train.index.isin(train_ids)]\nX_valid = train[train.index.isin(valid_ids)]\n\ndisplay(train.groupby('fold').size())\ndisplay(train.head(1))\n\nfig, axes = plt.subplots(nrows=1, ncols=2, figsize=(10, 5))\nfor i, col in enumerate([\"organ\", \"sex\"]):\n    _= train[[col]].value_counts().plot.pie(ax=axes[i], autopct='%1.1f%%', ylabel=col)\n\ntrain[[\"age\"]].hist(bins=40, figsize=(10, 5))","metadata":{"execution":{"iopub.status.busy":"2022-09-14T22:14:39.346320Z","iopub.execute_input":"2022-09-14T22:14:39.347450Z","iopub.status.idle":"2022-09-14T22:14:40.078182Z","shell.execute_reply.started":"2022-09-14T22:14:39.347418Z","shell.execute_reply":"2022-09-14T22:14:40.077084Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def rle_decode(mask_rle: str, shape: tuple = None) -> np.ndarray:\n    seq = mask_rle.split()\n    starts = np.array(list(map(int, seq[0::2])))\n    lengths = np.array(list(map(int, seq[1::2])))\n    assert len(starts) == len(lengths)\n    ends = starts + lengths\n    img = np.zeros((np.product(shape),), dtype=np.uint8)\n    for begin, end in zip(starts, ends):\n        img[begin:end] = 1\n    img.shape = shape\n    return img    \n\ndef rle_encode(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.T.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)    \n\ndef round_clip_0_1(x, **kwargs):\n    return x.round().clip(0, 1)\n\ndata_transforms = A.Compose([\n        A.OneOf([\n        A.HorizontalFlip(p=pr),\n        A.ShiftScaleRotate(scale_limit=0.5, rotate_limit=0, shift_limit=0.1, p=pr, border_mode=0),\n        A.IAAAdditiveGaussianNoise(p=pr),\n        A.IAAPerspective(p=pr),\n        A.RandomBrightness(p=pr),\n        ], p=1.0),\n        ], p=1.0)    \n","metadata":{"execution":{"iopub.status.busy":"2022-09-14T22:14:40.081240Z","iopub.execute_input":"2022-09-14T22:14:40.081576Z","iopub.status.idle":"2022-09-14T22:14:40.093405Z","shell.execute_reply.started":"2022-09-14T22:14:40.081548Z","shell.execute_reply":"2022-09-14T22:14:40.091787Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class DataGenerator(torch_data.Dataset):\n    def __init__(self, df, subset=\"train\", shuffle=False, transforms=None):\n        super().__init__()\n        self.df = df\n        self.shuffle = shuffle\n        self.subset = subset\n        self.indexes = np.arange(len(df))\n        self.on_epoch_end()\n        self.transforms = transforms\n        \n    def __len__(self):\n        return int(np.floor(len(self.df)))\n    \n    def on_epoch_end(self):\n        if self.shuffle == True:\n            np.random.shuffle(self.indexes)\n    \n    def __getitem__(self, index):\n        \n        w=self.df['img_width'].iloc[index]\n        h=self.df['img_height'].iloc[index]\n        img = self.__load_image(self.df['path'].iloc[index],h,w) \n        \n        if self.subset=='train':\n            mask = rle_decode(self.df.iloc[index].rle, shape=(h, w))\n            mask = cv2.resize(mask, (im_height,im_width))\n            mask = mask.astype(np.float32) \n            mask = np.expand_dims(mask.T, axis=-1) \n\n            if self.transforms:\n                sample = self.transforms(image=img, mask=mask)\n                img, mask = sample['image'], sample['mask']\n        \n        if self.subset=='train':\n            return img.T,mask.T\n        else:\n            return img.T\n        \n    def __load_image(self, img_path,h,w):\n        img = cv2.imread(img_path, cv2.IMREAD_COLOR)\n        img = cv2.resize(img, (im_height,im_width))\n        img = img.astype(np.float32) / 255.\n\n        return img","metadata":{"execution":{"iopub.status.busy":"2022-09-14T22:14:40.094886Z","iopub.execute_input":"2022-09-14T22:14:40.095430Z","iopub.status.idle":"2022-09-14T22:14:40.110614Z","shell.execute_reply.started":"2022-09-14T22:14:40.095393Z","shell.execute_reply":"2022-09-14T22:14:40.109512Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"n_cpu = os.cpu_count()\ntrain_dataset = DataGenerator(X_train)\nvalid_dataset = DataGenerator(X_valid)  \ntrain_dataloader = DataLoader(train_dataset, batch_size=4, shuffle=True, num_workers=n_cpu)\nvalid_dataloader = DataLoader(valid_dataset, batch_size=4, shuffle=False, num_workers=n_cpu)\n\n# lets look at some samples\n\nsample = train_dataset[0]\nplt.subplot(1,2,1)\nplt.imshow(sample[0].T) # for visualization we have to transpose back to HWC\nplt.subplot(1,2,2)\nplt.imshow(sample[1].squeeze())  # for visualization we have to remove 3rd dimension of mask\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-09-14T22:14:40.112264Z","iopub.execute_input":"2022-09-14T22:14:40.112879Z","iopub.status.idle":"2022-09-14T22:14:40.495571Z","shell.execute_reply.started":"2022-09-14T22:14:40.112834Z","shell.execute_reply":"2022-09-14T22:14:40.494620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class HubModel(pl.LightningModule):\n\n    def __init__(self, model, preprocessor, opt):\n        super().__init__()\n        self.model = model\n        self.opt = opt\n        # preprocessing parameteres for image\n        params = smp.encoders.get_preprocessing_params(preprocessor)\n        self.register_buffer(\"std\", torch.tensor(params[\"std\"]).view(1, 3, 1, 1))\n        self.register_buffer(\"mean\", torch.tensor(params[\"mean\"]).view(1, 3, 1, 1))\n\n        # for image segmentation dice loss could be the best first choice\n        self.loss_fn = smp.losses.DiceLoss(smp.losses.BINARY_MODE, from_logits=True)\n\n    def forward(self, image):\n        # normalize image here\n        image = (image - self.mean) / self.std\n        mask = self.model(image)\n        return mask\n\n    def shared_step(self, batch, stage):\n        \n        image = batch[0]\n\n        # Shape of the image should be (batch_size, num_channels, height, width)\n        # if you work with grayscale images, expand channels dim to have [batch_size, 1, height, width]\n        assert image.ndim == 4\n\n        # Check that image dimensions are divisible by 32, \n        # encoder and decoder connected by `skip connections` and usually encoder have 5 stages of \n        # downsampling by factor 2 (2 ^ 5 = 32); e.g. if we have image with shape 65x65 we will have \n        # following shapes of features in encoder and decoder: 84, 42, 21, 10, 5 -> 5, 10, 20, 40, 80\n        # and we will get an error trying to concat these features\n        h, w = image.shape[2:]\n        assert h % 32 == 0 and w % 32 == 0\n\n        mask = batch[1]\n\n        # Shape of the mask should be [batch_size, num_classes, height, width]\n        # for binary segmentation num_classes = 1\n        assert mask.ndim == 4\n\n        # Check that mask values in between 0 and 1, NOT 0 and 255 for binary segmentation\n        assert mask.max() <= 1.0 and mask.min() >= 0\n\n        logits_mask = self.forward(image)\n        \n        # Predicted mask contains logits, and loss_fn param `from_logits` is set to True\n        loss = self.loss_fn(logits_mask, mask)\n\n        # Lets compute metrics for some threshold\n        # first convert mask values to probabilities, then \n        # apply thresholding\n        prob_mask = logits_mask.sigmoid()\n        pred_mask = (prob_mask > 0.5).float()\n\n        # We will compute IoU metric by two ways\n        #   1. dataset-wise\n        #   2. image-wise\n        # but for now we just compute true positive, false positive, false negative and\n        # true negative 'pixels' for each image and class\n        # these values will be aggregated in the end of an epoch\n        tp, fp, fn, tn = smp.metrics.get_stats(pred_mask.long(), mask.long(), mode=\"binary\")\n\n        return {\n            \"loss\": loss,\n            \"tp\": tp,\n            \"fp\": fp,\n            \"fn\": fn,\n            \"tn\": tn,\n        }\n\n    def shared_epoch_end(self, outputs, stage):\n        # aggregate step metics\n        tp = torch.cat([x[\"tp\"] for x in outputs])\n        fp = torch.cat([x[\"fp\"] for x in outputs])\n        fn = torch.cat([x[\"fn\"] for x in outputs])\n        tn = torch.cat([x[\"tn\"] for x in outputs])\n\n        # per image IoU means that we first calculate IoU score for each image \n        # and then compute mean over these scores\n        per_image_f1 = smp.metrics.f1_score(tp, fp, fn, tn, reduction=\"micro-imagewise\")\n        \n        # dataset IoU means that we aggregate intersection and union over whole dataset\n        # and then compute IoU score. The difference between dataset_iou and per_image_iou scores\n        # in this particular case will not be much, however for dataset \n        # with \"empty\" images (images without target class) a large gap could be observed. \n        # Empty images influence a lot on per_image_iou and much less on dataset_iou.\n        dataset_f1 = smp.metrics.f1_score(tp, fp, fn, tn, reduction=\"micro\")\n\n        metrics = {\n            # f\"{stage}_per_image_f1\": per_image_f1,\n            f\"{stage}_dataset_f1\": dataset_f1,\n        }\n        \n        self.log_dict(metrics, prog_bar=True)\n\n    def training_step(self, batch, batch_idx):\n        return self.shared_step(batch, \"train\")            \n\n    def training_epoch_end(self, outputs):\n        return self.shared_epoch_end(outputs, \"train\")\n\n    def validation_step(self, batch, batch_idx):\n        return self.shared_step(batch, \"valid\")\n\n    def validation_epoch_end(self, outputs):\n        return self.shared_epoch_end(outputs, \"valid\")\n\n    def test_step(self, batch, batch_idx):\n        return self.shared_step(batch, \"test\")  \n\n    def test_epoch_end(self, outputs):\n        return self.shared_epoch_end(outputs, \"test\")\n\n    def configure_optimizers(self):\n\n        if self.opt=='Adam':\n            return torch.optim.Adam(self.parameters(), lr=lr)\n        elif self.opt=='Adamax':\n            return torch.optim.Adamax(self.parameters(), lr=lr)","metadata":{"execution":{"iopub.status.busy":"2022-09-14T22:14:40.497436Z","iopub.execute_input":"2022-09-14T22:14:40.497830Z","iopub.status.idle":"2022-09-14T22:14:40.690227Z","shell.execute_reply.started":"2022-09-14T22:14:40.497792Z","shell.execute_reply":"2022-09-14T22:14:40.689026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"BATCH_SIZE = 4\nEPOCHS = 5\nlr = 1e-4\npreprocessor = \"mit_b5\"\noptimizer = 'Adam'\n\nfor fold in fold_selected:\n    print('\\nFold: ',fold)\n    \n    train_ids = train[train[\"fold\"]!=fold].index\n    valid_ids = train[train[\"fold\"]==fold].index\n\n    X_train = train[train.index.isin(train_ids)]\n    X_valid = train[train.index.isin(valid_ids)]\n\n    train_dataset = DataGenerator(X_train)\n    valid_dataset = DataGenerator(X_valid)  \n    train_dataloader = DataLoader(train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=n_cpu)\n    valid_dataloader = DataLoader(valid_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=n_cpu)\n\n    model = unet_swin(img_size=im_width,size=\"swinv2_base_window16_256\")\n    model = HubModel(model, preprocessor, optimizer)\n    \n    trainer = pl.Trainer(\n    gpus=1, \n    max_epochs=EPOCHS\n    )\n\n    trainer.fit(\n    model, \n    train_dataloaders=train_dataloader, \n    val_dataloaders=valid_dataloader,\n    )\n\n    torch.save(model.state_dict(), f'swin_{preprocessor}_{fold}.pt')\n    \n    del X_train, X_valid, train_dataset, valid_dataloader, model\n    gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-09-14T22:14:40.692234Z","iopub.execute_input":"2022-09-14T22:14:40.693026Z","iopub.status.idle":"2022-09-14T22:19:26.191475Z","shell.execute_reply.started":"2022-09-14T22:14:40.692981Z","shell.execute_reply":"2022-09-14T22:19:26.190349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = pd.read_csv('../input/hubmap-organ-segmentation/test.csv')\ntest_df['path'] = ['../input/hubmap-organ-segmentation/test_images/'+str(x)+'.tiff' for x in test_df.id]\ntest_df['image_path'] = [x.replace('../input/hubmap-organ-segmentation/test_images/','./test_images/').replace('.tiff','.png') for x in test_df.path]\nsubmission = pd.read_csv('../input/hubmap-organ-segmentation/sample_submission.csv')\ntest_df.head()","metadata":{"execution":{"iopub.status.busy":"2022-09-14T22:19:26.198223Z","iopub.execute_input":"2022-09-14T22:19:26.198847Z","iopub.status.idle":"2022-09-14T22:19:26.254828Z","shell.execute_reply.started":"2022-09-14T22:19:26.198813Z","shell.execute_reply":"2022-09-14T22:19:26.253690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = unet_swin(img_size=im_width,size=\"swinv2_base_window16_256\")\nmodel = HubModel(model, preprocessor, optimizer)\ndataset = DataGenerator(test_df, subset='test')\ndataloader = DataLoader(dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=n_cpu)\ntrainer = pl.Trainer()\npreds = []\nrle = []\n\nfor m in fold_selected:\n    model.load_state_dict(torch.load(f'swin_{preprocessor}_{m}.pt'))\n    pred = trainer.predict(model, dataloaders=dataloader)\n    preds.append(np.vstack(pred))\n    \ndel model, pred, dataset, dataloader\ngc.collect()    \n    \npreds = np.mean(preds, axis=0)\n\nfor i in tqdm(test_df.index):\n    pred_img = cv2.resize(preds[i].T, (test_df.loc[i,\"img_width\"], test_df.loc[i,\"img_height\"]), interpolation=cv2.INTER_NEAREST) \n    pred_img = (pred_img>0.5).astype(dtype='uint8')   \n    rle.append(rle_encode(pred_img))","metadata":{"execution":{"iopub.status.busy":"2022-09-14T22:19:26.256339Z","iopub.execute_input":"2022-09-14T22:19:26.257144Z","iopub.status.idle":"2022-09-14T22:19:47.524494Z","shell.execute_reply.started":"2022-09-14T22:19:26.257089Z","shell.execute_reply":"2022-09-14T22:19:47.523410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df['rle'] = rle\nsubmission = submission[['id']].merge(test_df[['id','rle']], on='id', how='left')\nsubmission.to_csv('submission.csv',index=False)\nsubmission.head()","metadata":{"execution":{"iopub.status.busy":"2022-09-14T22:19:47.526582Z","iopub.execute_input":"2022-09-14T22:19:47.527324Z","iopub.status.idle":"2022-09-14T22:19:47.577024Z","shell.execute_reply.started":"2022-09-14T22:19:47.527283Z","shell.execute_reply":"2022-09-14T22:19:47.575690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}