{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":61446,"databundleVersionId":6962461,"sourceType":"competition"},{"sourceId":6929841,"sourceType":"datasetVersion","datasetId":3978555},{"sourceId":7455059,"sourceType":"datasetVersion","datasetId":4339413}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Baseline for image segmentation with Lightning⚡Torch","metadata":{}},{"cell_type":"code","source":"import os, glob\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport sys\n\nsys.path.append(\"/kaggle/input/rle-run-length-encoding-py-module\")  # this will be path to this notebook\nDATASET_FOLDER = \"/kaggle/input/blood-vessel-segmentation\"","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-01-22T12:10:27.741400Z","iopub.execute_input":"2024-01-22T12:10:27.741665Z","iopub.status.idle":"2024-01-22T12:10:28.693063Z","shell.execute_reply.started":"2024-01-22T12:10:27.741641Z","shell.execute_reply":"2024-01-22T12:10:28.692083Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(os.path.join(DATASET_FOLDER, \"train_rles.csv\"))\ndf_train[[\"dataset\", \"slice\"]] = df_train['id'].str.rsplit(pat='_', n=1, expand=True)\ndisplay(df_train.head())","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:10:28.694816Z","iopub.execute_input":"2024-01-22T12:10:28.695233Z","iopub.status.idle":"2024-01-22T12:10:29.964420Z","shell.execute_reply.started":"2024-01-22T12:10:28.695200Z","shell.execute_reply":"2024-01-22T12:10:29.963358Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"uq_datasets = df_train[\"dataset\"].unique()\nprint(f\"parsed {len(uq_datasets)} unique datasets: {uq_datasets}\")","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:10:29.966455Z","iopub.execute_input":"2024-01-22T12:10:29.966813Z","iopub.status.idle":"2024-01-22T12:10:29.975508Z","shell.execute_reply.started":"2024-01-22T12:10:29.966780Z","shell.execute_reply":"2024-01-22T12:10:29.974533Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## RLE decoding & encoding\n\n> In order to reduce the submission file size, teams must submit segmentation results using run-length encoding on the pixel values. That is, instead of submitting an exhaustive list of indices for your segmentation, you will submit pairs of values that contain a start position and a run length. E.g. '0 3' implies starting at pixel 0 and running a total of 3 pixels (0,1,2). The competition format requires a space delimited list of pairs. For example, '0 3 10 5' implies pixels 0,1,2, and 10,11,12,13,14 are to be included in the mask. The metric checks that the pairs are sorted, positive, and the decoded pixel values are not duplicated. The pixels are numbered from top to bottom, then left to right: 0 is pixel (0,0), 1 is pixel (1,0), and 2 is pixel (2,0) etc. [source](https://www.kaggle.com/code/leahscherschel/run-length-encoding)","metadata":{}},{"cell_type":"code","source":"from rle import rle_decode, rle_encode","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2024-01-22T12:10:29.977962Z","iopub.execute_input":"2024-01-22T12:10:29.978266Z","iopub.status.idle":"2024-01-22T12:10:29.995561Z","shell.execute_reply.started":"2024-01-22T12:10:29.978236Z","shell.execute_reply":"2024-01-22T12:10:29.994567Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset & DataModule\n\nCreating standard PyTorch dataset to define how the data shall be loaded and set representations. We define the sample pair as:\n\nA DataModule standardizes the training, val, test splits, data preparation and transforms. The main advantage is consistent data splits, data preparation and transforms across models.","metadata":{}},{"cell_type":"code","source":"import os\nimport torch\nfrom PIL import Image\nfrom torch.utils.data import Dataset\n\nclass SenNetHOADataset(Dataset):\n    split: float = 0.90\n\n    def __init__(\n        self,\n        df_data: pd.DataFrame,\n        path_img_dir: str =  os.path.join(DATASET_FOLDER, \"train\"),\n        transforms = None,\n        mode: str = 'train',\n        labels_lut = None\n    ):\n        self.path_img_dir = path_img_dir\n        self.transforms = transforms\n        self.mode = mode\n\n        # shuffle data\n        self.data = df_data.sample(frac=1, random_state=42).reset_index(drop=True)\n        # split dataset\n        assert 0.0 <= self.split <= 1.0\n        frac = int(self.split * len(self.data))\n        self.data = self.data[:frac] if mode == 'train' else self.data[frac:]\n        \n        self.img_paths = []\n        self.rles = []\n        for _, row in self.data.iterrows():\n            p_img = os.path.join(self.path_img_dir, row[\"dataset\"], \"images\", f'{row[\"slice\"]}.tif')\n            if not os.path.isfile(p_img):\n                continue\n            self.img_paths.append(p_img)\n            self.rles.append(row[\"rle\"])\n\n    def __getitem__(self, idx: int) -> tuple:\n        img = plt.imread(self.img_paths[idx])\n        if img.ndim == 3:\n            img = np.mean(img, axis=2)\n        if np.max(img) > 255:\n            img = np.clip(img / 255, 0, 255).astype(np.uint8)\n        img_shape = img.shape[:2]\n        mask = np.zeros((img_shape[0], img_shape[1], 2), dtype=np.uint8)\n        mask[..., 1] = rle_decode(self.rles[idx], img_shape=img_shape)\n\n        # augmentation / TODO\n        if self.transforms:\n            img = self.transforms(Image.fromarray(img))\n            mask = self.transforms(Image.fromarray(mask))\n        #print(f\"img dim: {img.shape}\")\n        return img, mask\n\n    def __len__(self) -> int:\n        assert len(self.img_paths) == len(self.rles)\n        return len(self.img_paths)\n\n# ==============================\n# ==============================\n\ndataset = SenNetHOADataset(df_train)\n\nfrom skimage import color\n\nfor i in range(6):\n    fig, axarr = plt.subplots(ncols=3, figsize=(12, 6))\n    img, mask = dataset[i]\n    mask = np.argmax(mask, axis=-1)\n    print(f\"image {img.shape} with range({np.min(img)}, {np.max(img)})\")\n    axarr[0].imshow(img, cmap=\"gray\")\n    axarr[1].imshow(color.label2rgb(mask, img, bg_label=0, bg_color=(1.,1.,1.), alpha=0.25))\n    axarr[2].imshow(mask, vmin=0, interpolation='antialiased', interpolation_stage='rgba')\n\n    # for i in range(3):\n    #     axarr[i].set_axis_off()","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:18:34.499072Z","iopub.execute_input":"2024-01-22T12:18:34.500017Z","iopub.status.idle":"2024-01-22T12:18:51.142926Z","shell.execute_reply.started":"2024-01-22T12:18:34.499977Z","shell.execute_reply":"2024-01-22T12:18:51.141996Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchvision import transforms as T\nfrom torchvision.transforms import InterpolationMode\n\nTRAIN_TRANSFORM = T.Compose([\n    # TODO add random\n    T.Resize(512), # resize to have the smaller size 512\n    T.CenterCrop(512), # srop remaining dimension\n    T.ToTensor(),\n])\n\nVALID_TRANSFORM = T.Compose([\n    T.Resize(512), # resize to have the smaller size 512\n    T.CenterCrop(512), # srop remaining dimension\n    T.ToTensor(),\n])","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:20:03.648357Z","iopub.execute_input":"2024-01-22T12:20:03.648714Z","iopub.status.idle":"2024-01-22T12:20:04.135505Z","shell.execute_reply.started":"2024-01-22T12:20:03.648686Z","shell.execute_reply":"2024-01-22T12:20:04.134620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"The DataModule include creating training and validation dataset with given split and feading it to particular data loaders...","metadata":{}},{"cell_type":"code","source":"import multiprocessing as mproc\nimport pytorch_lightning as pl\nfrom torch.utils.data import DataLoader\n\nclass SenNetDM(pl.LightningDataModule):\n\n    def __init__(\n        self,\n        df_data,\n        path_img_dir: str = os.path.join(DATASET_FOLDER, \"train\"),\n        batch_size: int = 8,\n        num_workers: int = None,\n        train_transforms = TRAIN_TRANSFORM,\n        valid_transforms = VALID_TRANSFORM\n    ):\n        super().__init__()\n        self.df_data = df_data\n        self.path_img_dir = path_img_dir\n        self.batch_size = batch_size\n        self.num_workers = num_workers or mproc.cpu_count()\n        self.train_dataset = None\n        self.valid_dataset = None\n        self.train_transforms = train_transforms\n        self.valid_transforms = valid_transforms\n\n    def prepare_data(self):\n        pass\n\n    def setup(self, stage=None):\n        self.train_dataset = SenNetHOADataset(\n            self.df_data, self.path_img_dir, mode='train', transforms=self.train_transforms)\n        print(f\"training dataset: {len(self.train_dataset)}\")\n        self.valid_dataset = SenNetHOADataset(\n            self.df_data, self.path_img_dir, mode='valid', transforms=self.valid_transforms)\n        print(f\"validation dataset: {len(self.valid_dataset)}\")\n\n    def train_dataloader(self):\n        return DataLoader(\n            self.train_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=True,\n        )\n\n    def val_dataloader(self):\n        return DataLoader(\n            self.valid_dataset,\n            batch_size=self.batch_size,\n            num_workers=self.num_workers,\n            shuffle=False,\n        )\n\n    def test_dataloader(self):\n        pass\n\n# ==============================\n# ==============================\n\ndm = SenNetDM(df_train, batch_size=14)\ndm.setup()\n\n# quick view\nfig = plt.figure(figsize=(3, 7))\nfor imgs, masks in dm.train_dataloader():\n    print(f'image size: {imgs[0].shape}')\n    print(f'mask size: {masks[0].shape}')\n    for i in range(3):\n        fig, axarr = plt.subplots(ncols=3, figsize=(10, 4))\n        img, mask = imgs[i][0].numpy(), masks[i][1].numpy()\n        axarr[0].imshow(img, cmap=\"gray\")\n        axarr[1].imshow(color.label2rgb(mask, img, bg_label=0, bg_color=(1.,1.,1.), alpha=0.25))\n        axarr[2].imshow(mask, vmin=0, interpolation='antialiased', interpolation_stage='rgba')\n    break","metadata":{"execution":{"iopub.status.busy":"2024-01-22T12:20:04.137171Z","iopub.execute_input":"2024-01-22T12:20:04.137514Z","iopub.status.idle":"2024-01-22T12:20:18.487170Z","shell.execute_reply.started":"2024-01-22T12:20:04.137485Z","shell.execute_reply":"2024-01-22T12:20:18.485919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# CNN Model","metadata":{}},{"cell_type":"code","source":"!pip install segmentation-models-pytorch -f /kaggle/input/segmentation-models-pytorch-python-package/ --no-index","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-22T12:20:18.488819Z","iopub.execute_input":"2024-01-22T12:20:18.489449Z","iopub.status.idle":"2024-01-22T12:20:32.914385Z","shell.execute_reply.started":"2024-01-22T12:20:18.489407Z","shell.execute_reply":"2024-01-22T12:20:32.913143Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torchvision\nimport segmentation_models_pytorch as smp\n# from adan_pytorch import Adan\n# from lion_pytorch import Lion\nfrom torch import nn\nfrom torch.nn import functional as F\n\n\nclass LitSenNet(pl.LightningModule):\n\n    def __init__(self, net, lr: float = 1e-4):\n        super().__init__()\n        self.net = net\n        #self.arch = net.pretrained_cfg.get('architecture')\n        # self.train_accuracy = MulticlassAccuracy(num_classes=self.num_classes)\n        # self.val_accuracy = MulticlassAccuracy(num_classes=self.num_classes)\n        # self.val_f1_score = MulticlassF1Score(num_classes=self.num_classes)\n        self.learn_rate = lr\n        self.loss = smp.losses.DiceLoss(mode='binary')\n\n    def forward(self, x):\n        return self.net(x)\n\n    def compute_loss(self, y_hat, y):\n        return self.loss(y_hat, y)\n\n    def training_step(self, batch, batch_idx):\n        x, y = batch\n        y_hat = self(x)\n        loss = self.compute_loss(y_hat, y)\n        self.log(\"train_loss\", loss, logger=True, prog_bar=True)\n        # self.log(\"train_acc\", self.train_accuracy(y_hat, lbs), logger=True, prog_bar=True)\n        return loss\n\n    def validation_step(self, batch, batch_idx):\n        x, y = batch\n        y_hat = self(x)\n        loss = self.compute_loss(y_hat, y)\n        self.log(\"valid_loss\", loss, logger=True, prog_bar=False)\n        # self.log(\"valid_acc\", self.val_accuracy(y_hat, lbs), logger=True, prog_bar=False)\n        # self.log(\"valid_f1\", self.val_f1_score(y_hat, lbs), logger=True, prog_bar=True)\n\n    def configure_optimizers(self):\n        optimizer = torch.optim.AdamW(self.parameters(), lr=self.learn_rate)\n        #optimizer = Lion(self.parameters(), lr=self.learn_rate, weight_decay=1e-2)\n        #optimizer = Adan(self.parameters(), lr=self.learn_rate, betas=(0.02, 0.08, 0.01), weight_decay=0.02)\n        #scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n        #    optimizer, T_max=self.trainer.max_epochs, eta_min=1e-6, verbose=True)\n        scheduler = torch.optim.lr_scheduler.CyclicLR(\n           optimizer, base_lr=self.learn_rate, max_lr=self.learn_rate * 10,\n           step_size_up=10, cycle_momentum=False, mode=\"triangular2\", verbose=True)\n        #scheduler = torch.optim.lr_scheduler.OneCycleLR(\n        #    optimizer, max_lr=1e-4, steps_per_epoch=1, epochs=self.trainer.max_epochs)\n        return [optimizer], [scheduler]\n\n# ==============================\n# ==============================\n\nnet = smp.Unet(\n    encoder_name=\"efficientnet-b4\", # choose encoder, e.g. mobilenet_v2 or resnet34\n    #encoder_weights=\"imagenet\", # use `imagenet` pre-trained weights for encoder initialization\n    in_channels=1, # model input channels (1 for gray-scale images, 3 for RGB, etc.)\n    classes=2, # model output channels (number of classes in your dataset)\n)\nmodel = LitSenNet(net=net, lr=1e-4)\nprint(model)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-22T12:20:32.917147Z","iopub.execute_input":"2024-01-22T12:20:32.917523Z","iopub.status.idle":"2024-01-22T12:20:36.406126Z","shell.execute_reply.started":"2024-01-22T12:20:32.917487Z","shell.execute_reply":"2024-01-22T12:20:36.405247Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training\n\nWe use Pytorch Lightning which allow us to drop all the boilet plate code and simplify all training just to use/call Trainer...","metadata":{}},{"cell_type":"code","source":"logger = pl.loggers.CSVLogger(save_dir='logs/', name=\"LitSenNet\")\nnb_epochs = 10 if torch.cuda.is_available() else 2\n\n# ==============================\n\ntrainer = pl.Trainer(\n    # fast_dev_run=True,\n    # callbacks=[swa],\n    logger=logger,\n    max_epochs=nb_epochs,\n    precision=16,\n    accumulate_grad_batches=4,\n    #val_check_interval=0.5,\n)\n\n# ==============================\n\n# trainer.tune(model, datamodule=dm)\ntrainer.fit(model=model, datamodule=dm)","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2024-01-22T12:20:36.407555Z","iopub.execute_input":"2024-01-22T12:20:36.408211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"Quick visualization of the training process...","metadata":{}},{"cell_type":"code","source":"import seaborn as sn\n\nmetrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')\ndel metrics[\"step\"]\nmetrics.set_index(\"epoch\", inplace=True)\n# display(metrics.dropna(axis=1, how=\"all\").head())\ng = sn.relplot(data=metrics, kind=\"line\")\nplt.gcf().set_size_inches(12, 4)\n# plt.gca().set_yscale('log')\nplt.grid()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.save_checkpoint(\"image_segmentation_model.pt\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Inference ?","metadata":{}},{"cell_type":"code","source":"!head /kaggle/input/blood-vessel-segmentation/sample_submission.csv","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"ls_images = glob.glob(os.path.join(DATASET_FOLDER, \"test\", \"*\", \"*\", \"*.tif\"))\nprint(f\"found images: {len(ls_images)}\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.auto import tqdm\n\nsubmission = []\nfor p_img in tqdm(ls_images):\n    path_ = p_img.split(os.path.sep)\n    # parse the submission ID\n    dataset = path_[-3]\n    slice_id, _ = os.path.splitext(path_[-1])\n#     # load image to get dimension\n#     img = plt.imread(p_img)\n#     # sample mask with rectangle\n#     mask = np.zeros(img.shape[:2])\n    # submission entry\n    submission.append({\n        \"id\": f\"{dataset}_{slice_id}\",\n        #\"rle\": rle_encode(mask),\n        \"rle\": \"1 2\",  # dummy\n    })","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_sub = pd.DataFrame(submission)\ndisplay(df_sub.head())\ndf_sub.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!head submission.csv","metadata":{"_kg_hide-output":true,"trusted":true},"execution_count":null,"outputs":[]}]}