{"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":"# Vesuvis PyTorch ⚡ MONAI \n\n## Nice [MONAI](https://docs.monai.io/en/stable/index.html) Features:\n- [CSVDataset](https://docs.monai.io/en/stable/data.html#csvdataset) to easily create dataset from DataFrame containing paths to volumes, masks, and labels\n- [RandWeightedCropd](https://docs.monai.io/en/stable/transforms.html#randweightedcropd) to create multiple random crops weighted with the mask\n- [matshow3d()](https://docs.monai.io/en/stable/visualize.html#monai.visualize.utils.matshow3d) function to quickly visualize volumes, masks, and labels\n- [UNet](https://docs.monai.io/en/stable/networks.html#unet) implementation\n- [DiceLoss](https://docs.monai.io/en/stable/losses.html#diceloss) implementation (Jaccard & DiceCELoss available as well)\n- [sliding_window_inference](https://docs.monai.io/en/stable/inferers.html#monai.inferers.sliding_window_inference) to run prediction on whole volume using patches\n\n## Nice [PyTorch Lightning](https://lightning.ai/docs/pytorch/stable/) Features:\n- [LightningDataModule](https://lightning.ai/docs/pytorch/stable/api/lightning.pytorch.core.LightningDataModule.html) to set up the train and val datasets, transforms, and dataloaders\n- [LightningModule](https://lightning.ai/docs/pytorch/stable/api/lightning.pytorch.core.LightningModule.html) to set up the model, loss, metrics, optimizer, scheduler, logging, callbacks, training and validation steps\n- [Trainer](https://lightning.ai/docs/pytorch/stable/api/lightning.pytorch.trainer.trainer.Trainer.html) to run training on (multiple) GPUs with mixed precision","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"markdown","source":"# Installs","metadata":{}},{"cell_type":"code","source":"!pip install monai lovely-numpy -q --no-index --find-links=../input/vesuvis-downloads","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:28:23.956670Z","iopub.execute_input":"2023-05-02T06:28:23.957235Z","iopub.status.idle":"2023-05-02T06:28:37.085768Z","shell.execute_reply.started":"2023-05-02T06:28:23.957206Z","shell.execute_reply":"2023-05-02T06:28:37.084552Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints\n!cp ../input/vesuvis-downloads/efficientnet-b0-355c32eb.pth /root/.cache/torch/hub/checkpoints","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:28:37.089193Z","iopub.execute_input":"2023-05-02T06:28:37.089592Z","iopub.status.idle":"2023-05-02T06:28:39.715334Z","shell.execute_reply.started":"2023-05-02T06:28:37.089548Z","shell.execute_reply":"2023-05-02T06:28:39.713972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"from collections import defaultdict\nfrom io import StringIO\nfrom pathlib import Path\nfrom typing import Tuple\n\nimport lovely_numpy as ln\nimport monai\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport PIL.Image as Image\nimport pytorch_lightning as pl\nimport seaborn as sns\nimport torch\nfrom monai.data import CSVDataset\nfrom monai.data import DataLoader\nfrom monai.inferers import sliding_window_inference\nfrom monai.visualize import matshow3d\nfrom torchmetrics import Dice\nfrom torchmetrics import MetricCollection\nfrom torchmetrics.classification import BinaryFBetaScore\nfrom tqdm.auto import tqdm","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:28:39.718232Z","iopub.execute_input":"2023-05-02T06:28:39.718637Z","iopub.status.idle":"2023-05-02T06:28:54.417577Z","shell.execute_reply.started":"2023-05-02T06:28:39.718593Z","shell.execute_reply":"2023-05-02T06:28:54.416541Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Paths & Settings","metadata":{}},{"cell_type":"code","source":"KAGGLE_DIR = Path(\"/\") / \"kaggle\"\n\nINPUT_DIR = KAGGLE_DIR / \"input\"\n\nCOMPETITION_DATA_DIR = INPUT_DIR / \"vesuvius-challenge-ink-detection\"\nPREPARED_DATA_DIR = INPUT_DIR / \"vesuvis-data-preparation\"\n\nTRAIN_DATA_CSV_PATH = PREPARED_DATA_DIR / \"data_1.0.csv\"\nTEST_DATA_CSV_PATH = \"test.csv\"\n\nDOWNSAMPLING = float(TRAIN_DATA_CSV_PATH.name.split(\"_\")[-1].replace(\".csv\", \"\"))\nNUM_Z_SLICES = 4\n\nACCELERATOR = \"gpu\"\nBATCH_SIZE = 1\nDEVICES = 1\nDROPOUT = 0.0\nETA_MIN = 1e-6\nFAST_DEV_RUN = False\nINTENSITY_TRANSFORM = \"NormalizeIntensity\"\nLEARNING_RATE = 0.01\nLOSS = \"BCE\"\nMODEL_NAME = \"FlexibleUNet_efficientnet-b0\"\nMAX_EPOCHS = 100\nNUM_WORKERS = 2\nNUM_SAMPLES = 12\nOPTIMIZER = \"SGD\"\nOVERFIT_BATCHES = 0\nPATCH_SIZE = (512, 512)\nPRECISION = 16\nRAND_TRANSFORMS = \"RandZoom-RandGaussianNoise-RandGaussianSmooth-RandScaleIntensity-RandFlip0-RandFlip1\"\nSCHEDULER = \"CosineAnnealingLR\"\nSEED = 2023\nSW_BATCH_SIZE = 4\nVAL_FRAGMENT_ID = \"1\"\nWEIGHT_DECAY = 1e-6\n\nTHRESHOLD = 0.5","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:28:54.420791Z","iopub.execute_input":"2023-05-02T06:28:54.421247Z","iopub.status.idle":"2023-05-02T06:28:54.430142Z","shell.execute_reply.started":"2023-05-02T06:28:54.421203Z","shell.execute_reply":"2023-05-02T06:28:54.427999Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning Datamodule","metadata":{}},{"cell_type":"code","source":"class VesuvisDataModule(pl.LightningDataModule):\n    def __init__(\n        self,\n        batch_size: int,\n        data_csv_path: str,\n        intensity_transform: str,\n        num_workers: int,\n        num_samples: int,\n        patch_size: Tuple[int, int],\n        rand_transforms: str,\n        val_fragment_id: str,\n    ):\n        super().__init__()\n\n        self.save_hyperparameters()\n\n        self.df = pd.read_csv(data_csv_path)\n\n        self.keys = (\"volume_npy\", \"mask_npy\", \"label_npy\")\n        self.train_transform = self._init_train_transform()\n        self.val_transform = self._init_val_transform()\n        self.predict_transform = self._init_predict_transform()\n        \n    def _load_transforms(self, predict: bool = False):\n        return [\n            monai.transforms.LoadImaged(\n                keys=\"volume_npy\",\n            ),\n            monai.transforms.LoadImaged(\n                keys=(\"mask_npy\", \"label_npy\") if not predict else \"mask_npy\",\n                ensure_channel_first=True,\n            ),\n        ]\n\n    @property\n    def _intensity_transforms(self):\n        if self.hparams.intensity_transform == \"NormalizeIntensity\":\n            return [\n                monai.transforms.NormalizeIntensityd(\n                    keys=\"volume_npy\",\n                    nonzero=True,\n                    channel_wise=True,\n                ),\n            ]\n        elif self.hparams.intensity_transform == \"ScaleIntensity\":\n            return [\n                monai.transforms.ScaleIntensityd(\n                    keys=\"volume_npy\",\n                ),\n            ]\n        else:\n            raise ValueError(f\"{self.hparams.intensity_transform} is not implemented\")\n            \n    @property\n    def _rand_transforms(self):\n        all_rand_transforms = {\n            \"RandAffine\": monai.transforms.RandAffined(\n                keys=self.keys,\n                prob=0.75,\n                rotate_range=(np.pi / 4, np.pi / 4),\n                translate_range=(0.0625, 0.0625),\n                scale_range=(0.1, 0.1),\n            ),\n            \"RandFlip0\": monai.transforms.RandFlipd(\n                keys=self.keys,\n                spatial_axis=0,\n                prob=0.5,\n            ),\n            \"RandFlip1\": monai.transforms.RandFlipd(\n                keys=self.keys,\n                spatial_axis=1,\n                prob=0.5,\n            ),\n            \"RandGaussianNoise\": monai.transforms.RandGaussianNoised(\n                keys=\"volume_npy\",\n                prob=0.15,\n                mean=0.0,\n                std=0.01,\n            ),\n            \"RandGaussianSmooth\": monai.transforms.RandGaussianSmoothd(\n                keys=\"volume_npy\",\n                prob=0.15,\n                sigma_x=(0.5, 1.15),\n                sigma_y=(0.5, 1.15),\n            ),\n            \"RandScaleIntensity\": monai.transforms.RandScaleIntensityd(\n                keys=\"volume_npy\",\n                factors=0.3,\n                prob=0.15,\n            ),\n            \"RandZoom\": monai.transforms.RandZoomd(\n                keys=self.keys,\n                min_zoom=0.9,\n                max_zoom=1.2,\n                mode=(\"bilinear\", \"nearest\", \"nearest\"),\n                align_corners=(True, None, None),\n                prob=0.15,\n            ),\n        }\n\n        rand_transforms = [\n            monai.transforms.RandCropByPosNegLabeld(\n                keys=self.keys,\n                label_key=\"label_npy\",\n                spatial_size=self.hparams.patch_size,\n                num_samples=self.hparams.num_samples,\n                image_key=\"volume_npy\",\n                image_threshold=0,\n            ),\n        ]\n\n        if self.hparams.rand_transforms is not None:\n            for rand_transform in self.hparams.rand_transforms.split(\"-\"):\n                rand_transforms.append(all_rand_transforms[rand_transform])\n\n        return rand_transforms\n    \n    def _init_train_transform(self):\n        return monai.transforms.Compose(self._load_transforms() + self._intensity_transforms + self._rand_transforms)\n\n    def _init_val_transform(self):\n        return monai.transforms.Compose(self._load_transforms() + self._intensity_transforms)\n\n    def _init_predict_transform(self):\n        return monai.transforms.Compose(self._load_transforms(predict=True) + self._intensity_transforms)\n\n    def setup(self, stage=None):\n        if stage == \"fit\" or stage is None:\n            train_val_df = self.df[self.df.stage == \"train\"].reset_index(drop=True)\n\n            train_df = train_val_df[train_val_df.fragment_id != int(self.hparams.val_fragment_id)].reset_index(\n                drop=True\n            )\n\n            val_df = train_val_df[train_val_df.fragment_id == int(self.hparams.val_fragment_id)].reset_index(drop=True)\n\n            self.train_dataset = self._dataset(train_df, self.train_transform)\n            self.val_dataset = self._dataset(val_df, self.val_transform)\n\n            print(f\"# train: {len(self.train_dataset)}\")\n            print(f\"# val: {len(self.val_dataset)}\")\n\n        if stage == \"predict\" or stage is None:\n            predict_df = self.df[self.df.stage == \"test\"].reset_index(drop=True)\n            self.predict_dataset = self._dataset(predict_df, self.predict_transform)\n\n            print(f\"# predict: {len(self.predict_dataset)}\")\n\n    def _dataset(self, df, transform):\n        return CSVDataset(\n            src=df,\n            transform=transform,\n        )\n\n    def train_dataloader(self):\n        return self._dataloader(self.train_dataset, train=True)\n\n    def val_dataloader(self):\n        return self._dataloader(self.val_dataset)\n\n    def predict_dataloader(self):\n        return self._dataloader(self.predict_dataset)\n\n    def _dataloader(self, dataset, train=False):\n        return DataLoader(\n            dataset,\n            batch_size=self.hparams.batch_size,\n            shuffle=train,\n            num_workers=self.hparams.num_workers,\n            pin_memory=True,\n            drop_last=train,\n        )","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:28:54.432323Z","iopub.execute_input":"2023-05-02T06:28:54.433210Z","iopub.status.idle":"2023-05-02T06:28:54.459158Z","shell.execute_reply.started":"2023-05-02T06:28:54.433165Z","shell.execute_reply":"2023-05-02T06:28:54.458109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize Data","metadata":{}},{"cell_type":"code","source":"def visualize_dataloaders(dataloaders, train=True):\n    for stage, dataloader in dataloaders.items():\n        for batch_idx, batch in enumerate(dataloader):\n            volumes = batch[\"volume_npy\"]\n            masks = batch[\"mask_npy\"]\n            \n            if train:\n                labels = batch[\"label_npy\"]\n            else: \n                labels = masks\n                \n            for volume, mask, label in zip(volumes, masks, labels):\n                fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n                plt.suptitle(f\"stage: {stage}, fragment: {batch_idx}\")\n\n                for idx, image in enumerate((volume, mask, label)):\n                    matshow3d(\n                        volume=image,\n                        fig=axes[idx],\n                        title=f\"{list(image.shape)}, {image.min().item():.2f}, {image.max().item():.2f}\",\n                        vmin=0.0,\n                        vmax=1.0,\n                        every_n=1,\n                        fill_value=1.0,\n                        margin=4,\n                        cmap=\"gray\",\n                    )","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:28:54.460733Z","iopub.execute_input":"2023-05-02T06:28:54.461179Z","iopub.status.idle":"2023-05-02T06:28:54.475015Z","shell.execute_reply.started":"2023-05-02T06:28:54.461137Z","shell.execute_reply":"2023-05-02T06:28:54.474033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_module = VesuvisDataModule(\n    batch_size=BATCH_SIZE,\n    data_csv_path=TRAIN_DATA_CSV_PATH,\n    intensity_transform=\"ScaleIntensity\",\n    num_workers=NUM_WORKERS,\n    num_samples=2,\n    rand_transforms=RAND_TRANSFORMS,\n    patch_size=PATCH_SIZE,\n    val_fragment_id=VAL_FRAGMENT_ID,\n)\ndata_module.setup()\n\ndataloaders = {\n    \"train\": data_module.train_dataloader(),\n    \"val\": data_module.val_dataloader(),\n}\n\nvisualize_dataloaders(dataloaders)","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:28:54.476457Z","iopub.execute_input":"2023-05-02T06:28:54.476833Z","iopub.status.idle":"2023-05-02T06:29:47.775135Z","shell.execute_reply.started":"2023-05-02T06:28:54.476787Z","shell.execute_reply":"2023-05-02T06:29:47.773919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning Module","metadata":{}},{"cell_type":"code","source":"class VesuvisModule(pl.LightningModule):\n    def __init__(\n        self,\n        dropout: float,\n        eta_min: float,\n        learning_rate: float,\n        loss: str,\n        model_name: str,\n        max_epochs: int,\n        num_z_slices: int,\n        optimizer: str,\n        patch_size: Tuple[int, int],\n        scheduler: str,\n        sw_batch_size: int,\n        weight_decay: float,\n    ):\n        super().__init__()\n\n        self.save_hyperparameters()\n\n        self.model = self._init_model()\n\n        self.loss = self._init_loss()\n\n        self.metrics = self._init_metrics()\n\n    def _init_model(self):\n        # TODO: add more models\n        if self.hparams.model_name == \"UNet\":\n            return monai.networks.nets.UNet(\n                spatial_dims=2,\n                in_channels=self.hparams.num_z_slices,\n                out_channels=1,\n                channels=(16, 32, 64, 128, 256),\n                strides=(2, 2, 2, 2),\n                num_res_units=2,\n                dropout=self.hparams.dropout,\n            )\n        elif \"FlexibleUNet\" in self.hparams.model_name:\n            return monai.networks.nets.FlexibleUNet(\n                in_channels=self.hparams.num_z_slices,\n                out_channels=1,\n                backbone=self.hparams.model_name.split(\"_\")[1],\n                pretrained=True,\n                spatial_dims=2,\n                dropout=self.hparams.dropout,\n            )\n        else:\n            raise ValueError(f\"{self.hparams.model_name} is not implemented\")\n\n    def _init_loss(self):\n        if self.hparams.loss == \"BCE\":\n            loss = torch.nn.BCEWithLogitsLoss()\n        elif self.hparams.loss == \"Dice\":\n            loss = monai.losses.DiceLoss(sigmoid=True)\n        elif self.hparams.loss == \"Jaccard\":\n            loss = monai.losses.DiceLoss(\n                sigmoid=True,\n                jaccard=True,\n            )\n        elif self.hparams.loss == \"DiceCE\":\n            loss = monai.losses.DiceCELoss(sigmoid=True)\n        else:\n            raise ValueError(f\"{self.hparams.loss} is not implemented\")\n\n        return monai.losses.MaskedLoss(loss)\n\n    def _init_metrics(self):\n        metric_collection = MetricCollection(\n            {\n                \"dice\": Dice(),\n                \"fbeta\": BinaryFBetaScore(beta=0.5),\n            }\n        )\n\n        return torch.nn.ModuleDict(\n            {\n                \"train_metrics\": metric_collection.clone(prefix=\"train_\"),\n                \"val_metrics\": metric_collection.clone(prefix=\"val_\"),\n            }\n        )\n\n    def configure_optimizers(self):\n        optimizer = self._init_optimizer()\n        scheduler = self._init_scheduler(optimizer)\n\n        return {\n            \"optimizer\": optimizer,\n            \"lr_scheduler\": {\n                \"scheduler\": scheduler,\n                \"interval\": \"epoch\",\n            },\n        }\n\n    def _init_optimizer(self):\n        if self.hparams.optimizer == \"Adam\":\n            return torch.optim.Adam(\n                params=self.parameters(),\n                lr=self.hparams.learning_rate,\n                weight_decay=self.hparams.weight_decay,\n            )\n        elif self.hparams.optimizer == \"AdamW\":\n            return torch.optim.AdamW(\n                params=self.parameters(),\n                lr=self.hparams.learning_rate,\n                weight_decay=self.hparams.weight_decay,\n            )\n        elif self.hparams.optimizer == \"SGD\":\n            return torch.optim.SGD(\n                params=self.parameters(),\n                lr=self.hparams.learning_rate,\n                momentum=0.99,\n                nesterov=True,\n            )\n        else:\n            raise ValueError(f\"{self.hparams.optimizer} is not implemented\")\n\n    def _init_scheduler(self, optimizer):\n        if self.hparams.scheduler == \"CosineAnnealingLR\":\n            return torch.optim.lr_scheduler.CosineAnnealingLR(\n                optimizer,\n                T_max=self.hparams.max_epochs,\n                eta_min=self.hparams.eta_min,\n            )\n        elif self.hparams.scheduler == \"StepLR\":\n            return torch.optim.lr_scheduler.StepLR(\n                optimizer,\n                step_size=self.hparams.max_epochs // 5,\n                gamma=0.95,\n            )\n        elif self.hparams.scheduler == \"PolyLR\":\n            return torch.optim.lr_scheduler.LambdaLR(\n                optimizer, lr_lambda=lambda epoch: (1 - epoch / self.hparams.max_epochs) ** 0.9\n            )\n        else:\n            raise ValueError(f\"{self.hparams.scheduler} is not implemented\")\n\n    def forward(self, x):\n        return self.model(x)\n\n    def training_step(self, batch):\n        return self._shared_step(batch, \"train\")\n\n    def validation_step(self, batch, batch_idx):\n        self._shared_step(batch, \"val\")\n\n    def predict_step(self, batch, batch_idx):\n        outputs = self._forward_pass(batch, \"predict\")\n        return outputs.sigmoid().squeeze()\n\n    def _shared_step(self, batch, stage):\n        outputs, labels, masks = self._forward_pass(batch, stage)\n\n        loss = self.loss(outputs, labels, masks)\n\n        self.metrics[f\"{stage}_metrics\"](outputs, labels)\n\n        self._log(loss, stage, batch_size=len(outputs))\n\n        return loss\n\n    def _forward_pass(self, batch, stage):\n        volumes = batch[\"volume_npy\"].as_tensor()\n        masks = batch[\"mask_npy\"].as_tensor()\n\n        if stage == \"train\":\n            outputs = self(volumes)\n        elif stage == \"val\":\n            outputs = sliding_window_inference(\n                inputs=volumes,\n                roi_size=self.hparams.patch_size,\n                sw_batch_size=self.hparams.sw_batch_size,\n                predictor=self,\n                overlap=0.5,\n                mode=\"gaussian\",\n            )\n        elif stage == \"predict\":\n            outputs = sliding_window_inference(\n                inputs=volumes,\n                roi_size=self.hparams.patch_size,\n                sw_batch_size=self.hparams.sw_batch_size,\n                predictor=self,\n                overlap=0.5,\n                mode=\"gaussian\",\n            )\n\n            ct = 1.0\n            for dims in [[2], [3], [2, 3]]:\n                flip_inputs = torch.flip(volumes, dims)\n                flip_outputs = torch.flip(\n                    sliding_window_inference(\n                        inputs=flip_inputs,\n                        roi_size=self.hparams.patch_size,\n                        sw_batch_size=self.hparams.sw_batch_size,\n                        predictor=self,\n                        overlap=0.5,\n                        mode=\"gaussian\",\n                    ),\n                    dims,\n                )\n                del flip_inputs\n                outputs += flip_outputs\n                del flip_outputs\n                ct += 1.0\n\n            outputs /= ct\n            \n            return outputs\n\n        try:\n            labels = batch[\"label_npy\"].as_tensor().long()\n            return outputs, labels, masks\n        except KeyError:\n            return outputs, masks\n\n    def _log(self, loss, stage, batch_size):\n        self.log(f\"{stage}_loss\", loss, batch_size=batch_size)\n        self.log_dict(self.metrics[f\"{stage}_metrics\"], batch_size=batch_size)","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:29:47.777354Z","iopub.execute_input":"2023-05-02T06:29:47.778181Z","iopub.status.idle":"2023-05-02T06:29:47.809885Z","shell.execute_reply.started":"2023-05-02T06:29:47.778130Z","shell.execute_reply":"2023-05-02T06:29:47.808765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"def train(\n    accelerator=ACCELERATOR,\n    batch_size=BATCH_SIZE,\n    data_csv_path=TRAIN_DATA_CSV_PATH,\n    devices=DEVICES,\n    dropout=DROPOUT,\n    eta_min=ETA_MIN,\n    fast_dev_run=FAST_DEV_RUN,\n    intensity_transform=INTENSITY_TRANSFORM,\n    learning_rate=LEARNING_RATE,\n    loss=LOSS,\n    model_name=MODEL_NAME,\n    max_epochs=MAX_EPOCHS,\n    num_workers=NUM_WORKERS,\n    num_samples=NUM_SAMPLES,\n    num_z_slices=NUM_Z_SLICES,\n    optimizer=OPTIMIZER,\n    overfit_batches=OVERFIT_BATCHES,\n    patch_size=PATCH_SIZE,\n    precision=PRECISION,\n    rand_transforms=RAND_TRANSFORMS,\n    scheduler=SCHEDULER,\n    seed=SEED,\n    sw_batch_size=SW_BATCH_SIZE,\n    val_fragment_id=VAL_FRAGMENT_ID,\n    weight_decay=WEIGHT_DECAY,\n):\n    monai.utils.set_determinism(seed)\n    pl.seed_everything(seed, workers=True)\n\n    data_module = VesuvisDataModule(\n        batch_size=batch_size,\n        data_csv_path=data_csv_path,\n        intensity_transform=intensity_transform,\n        num_workers=num_workers,\n        num_samples=num_samples,\n        patch_size=patch_size,\n        rand_transforms=rand_transforms,\n        val_fragment_id=val_fragment_id,\n    )\n\n    module = VesuvisModule(\n        dropout=dropout,\n        eta_min=eta_min,\n        learning_rate=learning_rate,\n        loss=loss,\n        model_name=model_name,\n        max_epochs=max_epochs,\n        num_z_slices=num_z_slices,\n        optimizer=optimizer,\n        patch_size=patch_size,\n        scheduler=scheduler,\n        sw_batch_size=sw_batch_size,\n        weight_decay=weight_decay,\n    )\n\n    trainer = pl.Trainer(\n        accelerator=accelerator,\n        benchmark=True,\n        check_val_every_n_epoch=1,\n        devices=devices,\n        fast_dev_run=fast_dev_run,\n        logger=pl.loggers.CSVLogger(save_dir='logs/'),\n        log_every_n_steps=1,\n        max_epochs=max_epochs,\n        overfit_batches=overfit_batches,\n        precision=precision,\n        strategy=\"ddp\" if devices > 1 else None,\n    )\n\n    trainer.fit(module, datamodule=data_module)\n\n    return module, trainer","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:29:47.812801Z","iopub.execute_input":"2023-05-02T06:29:47.813673Z","iopub.status.idle":"2023-05-02T06:29:47.827763Z","shell.execute_reply.started":"2023-05-02T06:29:47.813634Z","shell.execute_reply":"2023-05-02T06:29:47.826704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"module, trainer = train()","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:29:47.832501Z","iopub.execute_input":"2023-05-02T06:29:47.832777Z","iopub.status.idle":"2023-05-02T06:44:17.133740Z","shell.execute_reply.started":"2023-05-02T06:29:47.832752Z","shell.execute_reply":"2023-05-02T06:44:17.132562Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# From https://www.kaggle.com/code/jirkaborovec?scriptVersionId=93358967&cellId=22\nmetrics = pd.read_csv(f\"{trainer.logger.log_dir}/metrics.csv\")\nmetrics = metrics[[\"epoch\", \"train_loss\", \"val_loss\", \"val_dice\", \"val_fbeta\"]]\nmetrics.set_index(\"epoch\", inplace=True)\n\nsns.relplot(data=metrics, kind=\"line\", height=5, aspect=1.5)\nplt.grid()","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:44:17.139620Z","iopub.execute_input":"2023-05-02T06:44:17.139959Z","iopub.status.idle":"2023-05-02T06:44:17.915008Z","shell.execute_reply.started":"2023-05-02T06:44:17.139915Z","shell.execute_reply":"2023-05-02T06:44:17.914014Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare Test Data\n\n### Follows Training Data Preparation from https://www.kaggle.com/code/clemchris/vesuvis-data-preparation","metadata":{}},{"cell_type":"markdown","source":"## Create Test DataFrame","metadata":{}},{"cell_type":"code","source":"def create_df_from_mask_paths(stage, downsampling):\n    mask_paths = sorted(COMPETITION_DATA_DIR.glob(f\"{stage}/*/mask.png\"))\n\n    df = pd.DataFrame({\"mask_png\": mask_paths})\n\n    df[\"mask_png\"] = df[\"mask_png\"].astype(str)\n\n    df[\"stage\"] = df[\"mask_png\"].str.split(\"/\").str[-3]\n    df[\"fragment_id\"] = df[\"mask_png\"].str.split(\"/\").str[-2]\n\n    df[\"mask_npy\"] = df[\"mask_png\"].str.replace(\n        stage, f\"{stage}_{downsampling}\", regex=False\n    )\n    df[\"mask_npy\"] = df[\"mask_npy\"].str.replace(\"input\", \"working\", regex=False)\n    df[\"mask_npy\"] = df[\"mask_npy\"].str.replace(\"png\", \"npy\", regex=False)\n\n    if stage == \"train\":\n        df[\"label_png\"] = df[\"mask_png\"].str.replace(\"mask\", \"inklabels\", regex=False)\n        df[\"label_npy\"] = df[\"mask_npy\"].str.replace(\"mask\", \"inklabels\", regex=False)\n\n    df[\"volumes_dir\"] = df[\"mask_png\"].str.replace(\n        \"mask.png\", \"surface_volume\", regex=False\n    )\n    df[\"volume_npy\"] = df[\"mask_npy\"].str.replace(\"mask\", \"volume\", regex=False)\n\n    return df","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:44:17.916461Z","iopub.execute_input":"2023-05-02T06:44:17.917452Z","iopub.status.idle":"2023-05-02T06:44:17.926982Z","shell.execute_reply.started":"2023-05-02T06:44:17.917410Z","shell.execute_reply":"2023-05-02T06:44:17.925886Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = create_df_from_mask_paths(\"test\", DOWNSAMPLING)\n\ntest_df.to_csv(TEST_DATA_CSV_PATH)","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:44:17.928558Z","iopub.execute_input":"2023-05-02T06:44:17.929260Z","iopub.status.idle":"2023-05-02T06:44:17.963799Z","shell.execute_reply.started":"2023-05-02T06:44:17.929222Z","shell.execute_reply":"2023-05-02T06:44:17.962892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:44:17.965998Z","iopub.execute_input":"2023-05-02T06:44:17.966716Z","iopub.status.idle":"2023-05-02T06:44:17.983234Z","shell.execute_reply.started":"2023-05-02T06:44:17.966678Z","shell.execute_reply":"2023-05-02T06:44:17.982334Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Convert Data To NumPy","metadata":{}},{"cell_type":"code","source":"def load_image(path):\n    return Image.open(path)\n\n\ndef resize_image(image, downsampling):\n    size = int(image.size[0] * downsampling), int(image.size[1] * downsampling)\n    return image.resize(size)\n\n\ndef load_and_resize_image(path, downsampling):\n    image = load_image(path)\n    return resize_image(image, downsampling)\n\n\ndef load_label_npy(path, downsampling):\n    label = load_and_resize_image(path, downsampling)\n    return np.array(label) > 0\n\n\ndef load_mask_npy(path, downsampling):\n    mask = load_and_resize_image(path, downsampling).convert(\"1\")\n    return np.array(mask)\n\n\ndef load_z_slice_npy(path, downsampling):\n    z_slice = load_and_resize_image(path, downsampling)\n    return np.array(z_slice, dtype=np.float32) / 65535.0\n\n\ndef load_volume_npy(volumes_dir, num_z_slices, downsampling):\n    mid = 65 // 2\n    start = mid - num_z_slices // 2\n    end = mid + num_z_slices // 2\n\n    z_slices_paths = sorted(Path(volumes_dir).glob(\"*.tif\"))[start:end]\n\n    batch_size = num_z_slices // 4\n    paths_batches = [\n        z_slices_paths[i : i + batch_size]\n        for i in range(0, len(z_slices_paths), batch_size)\n    ]\n\n    volumes = []\n    for paths_batch in tqdm(\n        paths_batches, leave=False, desc=\"Processing batches\", position=1\n    ):\n        z_slices = [\n            load_z_slice_npy(path, downsampling)\n            for path in tqdm(\n                paths_batch, leave=False, desc=\"Processing paths\", position=2\n            )\n        ]\n        volumes.append(np.stack(z_slices, axis=0))\n        del z_slices\n\n        # break\n\n    volume = np.concatenate(volumes, axis=0)\n\n    return volume","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:44:17.984581Z","iopub.execute_input":"2023-05-02T06:44:17.984924Z","iopub.status.idle":"2023-05-02T06:44:17.997526Z","shell.execute_reply.started":"2023-05-02T06:44:17.984889Z","shell.execute_reply":"2023-05-02T06:44:17.996205Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_data_as_npy(df, train=True):\n    for row in tqdm(\n        df.itertuples(), total=len(df), desc=\"Processing fragments\", position=0\n    ):\n        mask_npy = load_mask_npy(row.mask_png, DOWNSAMPLING)\n        volume_npy = load_volume_npy(row.volumes_dir, NUM_Z_SLICES, DOWNSAMPLING)\n\n        Path(row.mask_npy).parent.mkdir(exist_ok=True, parents=True)\n        np.save(row.mask_npy, mask_npy)\n        np.save(row.volume_npy, volume_npy)\n\n        if train:\n            label_npy = load_label_npy(row.label_png, DOWNSAMPLING)\n            np.save(row.label_npy, label_npy)\n\n        tqdm.write(f\"Created {row.volume_npy} with shape {volume_npy.shape}\")","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:44:17.999415Z","iopub.execute_input":"2023-05-02T06:44:17.999772Z","iopub.status.idle":"2023-05-02T06:44:18.011539Z","shell.execute_reply.started":"2023-05-02T06:44:17.999737Z","shell.execute_reply":"2023-05-02T06:44:18.010415Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_data_as_npy(test_df, train=False)","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:44:18.013223Z","iopub.execute_input":"2023-05-02T06:44:18.013777Z","iopub.status.idle":"2023-05-02T06:44:25.984247Z","shell.execute_reply.started":"2023-05-02T06:44:18.013739Z","shell.execute_reply":"2023-05-02T06:44:25.982851Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize","metadata":{}},{"cell_type":"code","source":"data_module = VesuvisDataModule(\n    batch_size=BATCH_SIZE,\n    data_csv_path=TEST_DATA_CSV_PATH,\n    intensity_transform=INTENSITY_TRANSFORM,\n    num_workers=NUM_WORKERS,\n    num_samples=NUM_SAMPLES,\n    rand_transforms=RAND_TRANSFORMS,\n    patch_size=PATCH_SIZE,\n    val_fragment_id=VAL_FRAGMENT_ID,\n)\ndata_module.setup(stage=\"predict\")\n\ndataloaders = {\n    \"predict\": data_module.predict_dataloader(),\n}\n\nvisualize_dataloaders(dataloaders, train=False)","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:44:25.986353Z","iopub.execute_input":"2023-05-02T06:44:25.987863Z","iopub.status.idle":"2023-05-02T06:44:50.505712Z","shell.execute_reply.started":"2023-05-02T06:44:25.987806Z","shell.execute_reply":"2023-05-02T06:44:50.504136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predict","metadata":{}},{"cell_type":"code","source":"def predict(\n    module,\n    accelerator=ACCELERATOR,\n    batch_size=BATCH_SIZE,\n    data_csv_path=TEST_DATA_CSV_PATH,\n    devices=DEVICES,\n    intensity_transform=INTENSITY_TRANSFORM,\n    num_workers=NUM_WORKERS,\n    num_samples=NUM_SAMPLES,\n    patch_size=PATCH_SIZE,\n    precision=PRECISION,\n    rand_transforms=RAND_TRANSFORMS,\n    seed=SEED,\n    val_fragment_id=VAL_FRAGMENT_ID,\n):\n    monai.utils.set_determinism(seed)\n    pl.seed_everything(seed, workers=True)\n\n    data_module = VesuvisDataModule(\n        batch_size=batch_size,\n        data_csv_path=data_csv_path,\n        intensity_transform=intensity_transform,\n        num_workers=num_workers,\n        num_samples=num_samples,\n        patch_size=patch_size,\n        rand_transforms=rand_transforms,\n        val_fragment_id=val_fragment_id,\n    )\n\n    trainer = pl.Trainer(\n        accelerator=accelerator,\n        devices=devices,\n        precision=precision,\n    )\n\n    predictions = trainer.predict(module, datamodule=data_module)\n\n    return predictions","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:44:50.508343Z","iopub.execute_input":"2023-05-02T06:44:50.509299Z","iopub.status.idle":"2023-05-02T06:44:50.520566Z","shell.execute_reply.started":"2023-05-02T06:44:50.509241Z","shell.execute_reply":"2023-05-02T06:44:50.519190Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predictions = predict(module)","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:44:50.522709Z","iopub.execute_input":"2023-05-02T06:44:50.523220Z","iopub.status.idle":"2023-05-02T06:45:57.973686Z","shell.execute_reply.started":"2023-05-02T06:44:50.523172Z","shell.execute_reply":"2023-05-02T06:45:57.972451Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"def plot_image(image, title):\n    fig = plt.figure()\n    plt.title(title)\n    plt.imshow(image, cmap=\"gray\")\n    \n# ref.: https://www.kaggle.com/stainsby/fast-tested-rle\ndef rle(img):\n    '''\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    '''\n    pixels = img.flatten()\n    # pixels = (pixels >= thr).astype(int)\n    \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","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:45:57.976243Z","iopub.execute_input":"2023-05-02T06:45:57.976677Z","iopub.status.idle":"2023-05-02T06:45:57.984866Z","shell.execute_reply.started":"2023-05-02T06:45:57.976630Z","shell.execute_reply":"2023-05-02T06:45:57.983724Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission_df = pd.read_csv(COMPETITION_DATA_DIR / \"sample_submission.csv\")\n\npredictions_rle = []\nfor mask_png_path, prediction in zip(test_df[\"mask_png\"].values, predictions):\n    prediction = prediction.numpy()\n    plot_image(prediction, f\"{ln.lovely(prediction)}\")\n\n    mask = load_image(mask_png_path)\n\n    prediction = prediction * mask\n    plot_image(prediction, f\"{ln.lovely(prediction)}\")\n\n    prediction = np.where(prediction > THRESHOLD, 1, 0).astype(np.uint8)\n    plot_image(prediction, f\"{ln.lovely(prediction)}\")\n\n    prediction_rle = rle(prediction)\n    predictions_rle.append(prediction_rle)\n\n    plot_image(prediction, f\"{ln.lovely(prediction)}\")\n        \n    del prediction\n    \nsubmission_df[\"Predicted\"] = predictions_rle\nsubmission_df.to_csv(\"submission.csv\", index=False)\nsubmission_df","metadata":{"execution":{"iopub.status.busy":"2023-05-02T06:52:30.066352Z","iopub.execute_input":"2023-05-02T06:52:30.067357Z","iopub.status.idle":"2023-05-02T06:52:55.174041Z","shell.execute_reply.started":"2023-05-02T06:52:30.067315Z","shell.execute_reply":"2023-05-02T06:52:55.172938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}