{"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":45867,"databundleVersionId":6924515,"sourceType":"competition"},{"sourceId":151170812,"sourceType":"kernelVersion"}],"dockerImageVersionId":30588,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport os\nfrom sklearn.preprocessing import LabelEncoder\ntrain = pd.read_csv(\"/kaggle/input/UBC-OCEAN/train.csv\")\ntrain.head()\nimage_ids = [int(i.split('.')[0]) for i in os.listdir(\"/kaggle/input/mask-ocean-eda-numpy-labels\") if i.endswith('npy')]\ntrain['is_masked'] = train['image_id'].isin(image_ids)\nle = LabelEncoder()\ntrain['label'] = le.fit_transform(train['label'])\ntrain['image_path'] = '/kaggle/input/UBC-OCEAN/train_thumbnails/'+train['image_id'].astype(str) + '_thumbnail.png'\ntrain['mask_path'] = '/kaggle/input/mask-ocean-eda-numpy-labels/' + train['image_id'].astype(str) + '.npy'\ntrain = train[train['is_masked']].reset_index(drop = True)\ntrain","metadata":{"execution":{"iopub.status.busy":"2023-11-18T17:54:25.265426Z","iopub.execute_input":"2023-11-18T17:54:25.265788Z","iopub.status.idle":"2023-11-18T17:54:26.038561Z","shell.execute_reply.started":"2023-11-18T17:54:25.265759Z","shell.execute_reply":"2023-11-18T17:54:26.037591Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torch.utils.data import Dataset, DataLoader\nimport torch\nimport pandas as pd\nfrom sklearn.model_selection import train_test_split\nimport numpy as np\nimport os\nfrom PIL import Image\nimport torchvision.transforms as transforms\nimport matplotlib.pyplot as plt\nimport cv2","metadata":{"execution":{"iopub.status.busy":"2023-11-18T17:54:26.044263Z","iopub.execute_input":"2023-11-18T17:54:26.045021Z","iopub.status.idle":"2023-11-18T17:54:27.881587Z","shell.execute_reply.started":"2023-11-18T17:54:26.044984Z","shell.execute_reply":"2023-11-18T17:54:27.880685Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, df, size, mode):\n        self.df = df\n        self.size = size\n        self.mode = mode\n        # self.df = self.df.sample(frac=1).reset_index(drop=True)\n        self.df = self.df.sample(frac=1, random_state=42).reset_index(drop=True)\n        self.df_train, self.df_val = train_test_split(self.df, test_size=0.2, random_state=42, stratify=self.df['label'])\n        if self.mode == 'train':\n            self.df = self.df_train\n        else:\n            self.df = self.df_val\n        self.df = self.df.reset_index(drop=True)\n        self.transform = transforms.Compose([\n            transforms.Resize((size, size)),\n            transforms.ToTensor(),\n        ])\n\n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        image_path = self.df['image_path'][idx]\n        mask_path = self.df['mask_path'][idx]\n        image = Image.open(image_path)\n        image = self.transform(image)\n        mask = np.load(mask_path)\n        mask = cv2.resize(mask, (self.size, self.size), interpolation=cv2.INTER_NEAREST)\n#         mask = np.eye(4)[mask]\n        mask = torch.from_numpy(mask)\n#         mask = mask.permute(2, 0, 1)\n        return {'image': image, 'mask': mask}\n    \ntrain_dataset = CustomDataset(train, 512, 'train')\nval_dataset = CustomDataset(train, 512, 'val')\n\n# make a dataloader for train and val\ntrain_loader = DataLoader(train_dataset, batch_size=4, shuffle=False, num_workers= 3)\nval_loader = DataLoader(val_dataset, batch_size=8, shuffle=False, num_workers= 3)","metadata":{"execution":{"iopub.status.busy":"2023-11-18T18:42:07.222407Z","iopub.execute_input":"2023-11-18T18:42:07.222781Z","iopub.status.idle":"2023-11-18T18:42:07.243607Z","shell.execute_reply.started":"2023-11-18T18:42:07.222751Z","shell.execute_reply":"2023-11-18T18:42:07.242653Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # batch = next(iter(train_loader))\n# batch['mask'][0].long()","metadata":{"execution":{"iopub.status.busy":"2023-11-18T18:41:39.848155Z","iopub.execute_input":"2023-11-18T18:41:39.848804Z","iopub.status.idle":"2023-11-18T18:41:39.857393Z","shell.execute_reply.started":"2023-11-18T18:41:39.848768Z","shell.execute_reply":"2023-11-18T18:41:39.856479Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"%%capture\n!pip install segmentation-models-pytorch\n!pip install --upgrade pytorch-lightning","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import pytorch_lightning as pl\nimport segmentation_models_pytorch as smp","metadata":{"execution":{"iopub.status.busy":"2023-11-18T17:54:32.526583Z","iopub.execute_input":"2023-11-18T17:54:32.527214Z","iopub.status.idle":"2023-11-18T17:54:35.965068Z","shell.execute_reply.started":"2023-11-18T17:54:32.527181Z","shell.execute_reply":"2023-11-18T17:54:35.964218Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport pytorch_lightning as pl\nimport segmentation_models_pytorch as smp\n\nclass CancerSegModel(pl.LightningModule):\n\n    def __init__(self, arch, encoder_name, in_channels, out_classes, **kwargs):\n        super().__init__()\n        self.model = smp.create_model(\n            arch, encoder_name=encoder_name, in_channels=in_channels, classes=out_classes, **kwargs\n        )\n        params = smp.encoders.get_preprocessing_params(encoder_name)\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        self.loss_fn = smp.losses.DiceLoss(smp.losses.MULTICLASS_MODE, from_logits=True)\n        \n        self.train_outputs = []\n        self.valid_outputs = []\n\n    def forward(self, image):\n        image = (image - self.mean) / self.std\n        mask = self.model(image)\n        return mask\n\n    def shared_step(self, batch, stage):\n        image = batch[\"image\"]\n        assert image.ndim == 4\n        h, w = image.shape[2:]\n        assert h % 32 == 0 and w % 32 == 0\n\n        mask = batch[\"mask\"]\n#         assert mask.ndim == 4\n#         assert mask.max() <= 1.0 and mask.min() >= 0\n\n        logits_mask = self.forward(image)\n#         print(logits_mask.shape)\n        loss = self.loss_fn(logits_mask, mask.long())\n        prob_mask = logits_mask.softmax(dim = 1)\n#         pred_mask = (prob_mask > 0.5).float()\n        pred_mask = torch.argmax(prob_mask, dim = 1)\n        tp, fp, fn, tn = smp.metrics.get_stats(pred_mask.long(), mask.long(), mode=\"multiclass\", num_classes=4)\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, mode):\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        per_image_iou = smp.metrics.iou_score(tp, fp, fn, tn, reduction=\"micro-imagewise\")\n        dataset_iou = smp.metrics.iou_score(tp, fp, fn, tn, reduction=\"micro\")\n\n        metrics = {\n            f\"{mode}_per_image_iou\": per_image_iou,\n            f\"{mode}_dataset_iou\": dataset_iou,\n        }\n#         self.log_dict(metrics, prog_bar=True)\n        return metrics\n\n    def training_step(self, batch, batch_idx):\n        step = self.shared_step(batch, \"train\")\n        self.log(\"train_loss\", step['loss'], prog_bar=True)\n        self.train_outputs.append(step)\n        return step['loss']\n\n    def validation_step(self, batch, batch_idx):\n        step = self.shared_step(batch, \"valid\")\n        self.log(\"valid_loss\", step['loss'], prog_bar=True)\n        self.valid_outputs.append(step)\n        return step['loss']\n    def on_train_epoch_end(self):\n        metrics = self.shared_epoch_end(self.train_outputs, 'train')\n        self.train_outputs.clear()\n        self.log_dict(metrics, prog_bar = True)\n        return metrics\n    def on_validation_epoch_end(self):\n        metrics = self.shared_epoch_end(self.valid_outputs, 'valid')\n        self.valid_outputs.clear()\n        self.log_dict(metrics, prog_bar = True)\n        return metrics\n\n    def configure_optimizers(self):\n        return torch.optim.Adam(self.parameters(), lr=0.0001)","metadata":{"execution":{"iopub.status.busy":"2023-11-18T18:59:48.746566Z","iopub.execute_input":"2023-11-18T18:59:48.747743Z","iopub.status.idle":"2023-11-18T18:59:48.770376Z","shell.execute_reply.started":"2023-11-18T18:59:48.7477Z","shell.execute_reply":"2023-11-18T18:59:48.769238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model = CancerSegModel(\"unetplusplus\", \"efficientnet-b3\", in_channels=3, out_classes=4)","metadata":{"execution":{"iopub.status.busy":"2023-11-18T18:59:51.031929Z","iopub.execute_input":"2023-11-18T18:59:51.032646Z","iopub.status.idle":"2023-11-18T18:59:51.290906Z","shell.execute_reply.started":"2023-11-18T18:59:51.032611Z","shell.execute_reply":"2023-11-18T18:59:51.289829Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"logger = pl.loggers.CSVLogger(save_dir='logs/', name=\"unetpp_effnetb3_20\")\ntrainer = pl.Trainer(\n    logger=logger,\n    accelerator = 'gpu', \n    max_epochs=40,\n    log_every_n_steps=2,\n    accumulate_grad_batches=4\n)\n\ntrainer.fit(\n    model, \n    train_dataloaders=train_loader, \n    val_dataloaders=val_loader,\n)","metadata":{"execution":{"iopub.status.busy":"2023-11-18T18:59:52.943246Z","iopub.execute_input":"2023-11-18T18:59:52.944167Z","iopub.status.idle":"2023-11-18T19:14:11.742319Z","shell.execute_reply.started":"2023-11-18T18:59:52.944128Z","shell.execute_reply":"2023-11-18T19:14:11.741222Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics = pd.read_csv(f'{trainer.logger.log_dir}/metrics.csv')","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:14:14.024465Z","iopub.execute_input":"2023-11-18T19:14:14.024828Z","iopub.status.idle":"2023-11-18T19:14:14.034523Z","shell.execute_reply.started":"2023-11-18T19:14:14.024799Z","shell.execute_reply":"2023-11-18T19:14:14.033564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics['train_loss'].dropna().plot()","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:14:15.722929Z","iopub.execute_input":"2023-11-18T19:14:15.723677Z","iopub.status.idle":"2023-11-18T19:14:16.023237Z","shell.execute_reply.started":"2023-11-18T19:14:15.723643Z","shell.execute_reply":"2023-11-18T19:14:16.022333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics['valid_loss'].dropna().plot()","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:14:18.448243Z","iopub.execute_input":"2023-11-18T19:14:18.448618Z","iopub.status.idle":"2023-11-18T19:14:18.70161Z","shell.execute_reply.started":"2023-11-18T19:14:18.448592Z","shell.execute_reply":"2023-11-18T19:14:18.700525Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics['train_per_image_iou'].dropna().plot()","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:14:25.858181Z","iopub.execute_input":"2023-11-18T19:14:25.858614Z","iopub.status.idle":"2023-11-18T19:14:26.141205Z","shell.execute_reply.started":"2023-11-18T19:14:25.85858Z","shell.execute_reply":"2023-11-18T19:14:26.140097Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics['train_dataset_iou'].dropna().plot()","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:14:30.741266Z","iopub.execute_input":"2023-11-18T19:14:30.74199Z","iopub.status.idle":"2023-11-18T19:14:30.972773Z","shell.execute_reply.started":"2023-11-18T19:14:30.74195Z","shell.execute_reply":"2023-11-18T19:14:30.971816Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics['valid_dataset_iou'].dropna().plot()","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:14:38.624836Z","iopub.execute_input":"2023-11-18T19:14:38.625634Z","iopub.status.idle":"2023-11-18T19:14:38.901946Z","shell.execute_reply.started":"2023-11-18T19:14:38.625595Z","shell.execute_reply":"2023-11-18T19:14:38.900737Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"metrics['valid_per_image_iou'].dropna().plot()","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:14:56.427399Z","iopub.execute_input":"2023-11-18T19:14:56.427765Z","iopub.status.idle":"2023-11-18T19:14:56.647192Z","shell.execute_reply.started":"2023-11-18T19:14:56.427733Z","shell.execute_reply":"2023-11-18T19:14:56.646127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch = next(iter(val_loader))\nwith torch.no_grad():\n    model.eval()\n    logits = model(batch[\"image\"])\n# pr_masks = logits.softmax(dim = 1)\npr_masks = torch.argmax(logits.softmax(dim = 1), dim = 1)\n\nfor image, gt_mask, pr_mask in zip(batch[\"image\"], batch[\"mask\"], pr_masks):\n    plt.figure(figsize=(10, 5))\n\n    plt.subplot(1, 3, 1)\n    plt.imshow(image.numpy().transpose(1, 2, 0))  # convert CHW -> HWC\n    plt.title(\"Image\")\n    plt.axis(\"off\")\n\n    plt.subplot(1, 3, 2)\n    plt.imshow(gt_mask.numpy()) # just squeeze classes dim, because we have only one class\n    plt.title(\"Ground truth\")\n    plt.axis(\"off\")\n\n    plt.subplot(1, 3, 3)\n    plt.imshow(pr_mask.numpy()) # just squeeze classes dim, because we have only one class\n    plt.title(\"Prediction\")\n    plt.axis(\"off\")\n\n    plt.show()","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:15:01.069979Z","iopub.execute_input":"2023-11-18T19:15:01.070958Z","iopub.status.idle":"2023-11-18T19:15:21.369733Z","shell.execute_reply.started":"2023-11-18T19:15:01.070924Z","shell.execute_reply":"2023-11-18T19:15:21.368507Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"batch = next(iter(train_loader))\nwith torch.no_grad():\n    model.eval()\n    logits = model(batch[\"image\"])\n# pr_masks = logits.softmax(dim = 1)\npr_masks = torch.argmax(logits.softmax(dim = 1), dim = 1)\n\nfor image, gt_mask, pr_mask in zip(batch[\"image\"], batch[\"mask\"], pr_masks):\n    plt.figure(figsize=(10, 5))\n\n    plt.subplot(1, 3, 1)\n    plt.imshow(image.numpy().transpose(1, 2, 0))  # convert CHW -> HWC\n    plt.title(\"Image\")\n    plt.axis(\"off\")\n\n    plt.subplot(1, 3, 2)\n    plt.imshow(gt_mask.numpy()) # just squeeze classes dim, because we have only one class\n    plt.title(\"Ground truth\")\n    plt.axis(\"off\")\n\n    plt.subplot(1, 3, 3)\n    plt.imshow(pr_mask.numpy()) # just squeeze classes dim, because we have only one class\n    plt.title(\"Prediction\")\n    plt.axis(\"off\")\n\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:15:58.032428Z","iopub.execute_input":"2023-11-18T19:15:58.03284Z","iopub.status.idle":"2023-11-18T19:16:07.689136Z","shell.execute_reply.started":"2023-11-18T19:15:58.032806Z","shell.execute_reply":"2023-11-18T19:16:07.687797Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer.save_checkpoint(\"image_segmentation_model.pt\")","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:18:17.514212Z","iopub.execute_input":"2023-11-18T19:18:17.51539Z","iopub.status.idle":"2023-11-18T19:18:17.986747Z","shell.execute_reply.started":"2023-11-18T19:18:17.515351Z","shell.execute_reply":"2023-11-18T19:18:17.985903Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ThumbnailDataset(Dataset):\n    def __init__(self, size, image_paths):\n        self.size = size\n        self.image_paths = image_paths\n        self.transform = transforms.Compose([\n            transforms.Resize((size, size)),\n            transforms.ToTensor(),\n        ])\n\n    def __len__(self):\n        return len(self.image_paths)\n    \n    def __getitem__(self, idx):\n        image_path = self.image_paths[idx]\n        image = Image.open(image_path)\n        image = self.transform(image)\n#         mask = mask.permute(2, 0, 1)\n        return {\n            'image':image,\n            'id':int(image_path.split('/')[-1].split('.')[0].split('_')[0])\n        }\n    \nall_train_thumbnails = ThumbnailDataset(512, [\"/kaggle/input/UBC-OCEAN/train_thumbnails/\"+i for i in os.listdir('/kaggle/input/UBC-OCEAN/train_thumbnails')])\n\nthumbnail_loader = DataLoader(all_train_thumbnails, batch_size=64, shuffle=False, num_workers= 3)","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:33:12.097392Z","iopub.execute_input":"2023-11-18T19:33:12.098244Z","iopub.status.idle":"2023-11-18T19:33:12.109396Z","shell.execute_reply.started":"2023-11-18T19:33:12.098211Z","shell.execute_reply":"2023-11-18T19:33:12.108427Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir preds","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:33:17.24157Z","iopub.execute_input":"2023-11-18T19:33:17.242452Z","iopub.status.idle":"2023-11-18T19:33:18.346183Z","shell.execute_reply.started":"2023-11-18T19:33:17.242416Z","shell.execute_reply":"2023-11-18T19:33:18.344646Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from tqdm.auto import tqdm\nfor batch in tqdm(thumbnail_loader):\n    with torch.no_grad():\n        model.eval()\n        logits = model(batch['image'])\n    pr_masks = torch.argmax(logits.softmax(dim = 1), dim = 1)\n    for image, pr_mask, id_ in zip(batch[\"image\"], pr_masks, batch['id']):\n        np.save('preds/'+str(id_.item())+'.npy',pr_mask.numpy())","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:33:23.712463Z","iopub.execute_input":"2023-11-18T19:33:23.712902Z","iopub.status.idle":"2023-11-18T19:50:22.864541Z","shell.execute_reply.started":"2023-11-18T19:33:23.712851Z","shell.execute_reply":"2023-11-18T19:50:22.862881Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import random\nfor i in random.sample(os.listdir('/kaggle/working/preds'), 10):\n    img = np.load('/kaggle/working/preds/'+i)\n    thumb = Image.open(\"/kaggle/input/UBC-OCEAN/train_thumbnails/\"+i.split('.')[0]+\"_thumbnail.png\")\n    plt.figure()\n    plt.subplot(1,2,1)\n    plt.imshow(thumb)\n    plt.subplot(1,2,2)\n    plt.imshow(img)\nplt.show","metadata":{"execution":{"iopub.status.busy":"2023-11-18T19:54:19.655586Z","iopub.execute_input":"2023-11-18T19:54:19.656281Z","iopub.status.idle":"2023-11-18T19:54:35.956721Z","shell.execute_reply.started":"2023-11-18T19:54:19.656247Z","shell.execute_reply":"2023-11-18T19:54:35.955691Z"},"trusted":true},"execution_count":null,"outputs":[]}]}