{"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":"# HuBMAP PyTorch ⚡ MONAI Train & Infer\n\n## Combine powers of [PyTorch Lightning](https://www.pytorchlightning.ai/) and [MONAI](https://monai.io/)","metadata":{}},{"cell_type":"markdown","source":"# Installs","metadata":{}},{"cell_type":"code","source":"!cd ../input/hubmap-downloads && \\\npip install monai-0.9.0-202206131636-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:20:56.529342Z","iopub.execute_input":"2022-07-22T12:20:56.529791Z","iopub.status.idle":"2022-07-22T12:21:27.771723Z","shell.execute_reply.started":"2022-07-22T12:20:56.529685Z","shell.execute_reply":"2022-07-22T12:21:27.770594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nfrom typing import Any\nfrom typing import Callable\nfrom typing import Dict\nfrom typing import Optional\nfrom typing import Tuple\n\nimport monai\nimport numpy as np\nimport pandas as pd\nimport pytorch_lightning as pl\nimport seaborn as sns\nimport tifffile\nimport torch\nimport torch.nn as nn\nfrom matplotlib import pyplot as plt\nfrom monai.data import CSVDataset\nfrom monai.data import DataLoader\nfrom monai.data import ImageReader\nfrom sklearn.model_selection import StratifiedKFold\nfrom tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:27.773955Z","iopub.execute_input":"2022-07-22T12:21:27.774754Z","iopub.status.idle":"2022-07-22T12:21:35.126000Z","shell.execute_reply.started":"2022-07-22T12:21:27.774691Z","shell.execute_reply":"2022-07-22T12:21:35.124917Z"},"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\"\nOUTPUT_DIR = KAGGLE_DIR / \"working\"\n\nCOMPETITION_DATA_DIR = INPUT_DIR / \"hubmap-organ-segmentation\"\n\nTRAIN_PREPARED_CSV_PATH = \"train_prepared.csv\"\nVAL_PRED_PREPARED_CSV_PATH = \"val_pred_prepared.csv\"\nTEST_PREPARED_CSV_PATH = \"test_prepared.csv\"\n\nN_SPLITS = 4\nRANDOM_SEED = 2022\nSPATIAL_SIZE = 1024\nVAL_FOLD = 0\nNUM_WORKERS = 2\nBATCH_SIZE = 16\nLEARNING_RATE = 1e-3\nWEIGHT_DECAY = 0.0\nFAST_DEV_RUN = False\nGPUS = 1\nMAX_EPOCHS = 20\nPRECISION = 16\nDEBUG = False\n\nDEVICE = \"cuda\"\nTHRESHOLD = 0.5","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:35.127604Z","iopub.execute_input":"2022-07-22T12:21:35.128621Z","iopub.status.idle":"2022-07-22T12:21:35.137750Z","shell.execute_reply.started":"2022-07-22T12:21:35.128582Z","shell.execute_reply":"2022-07-22T12:21:35.136759Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare DataFrames (Add paths and create folds)","metadata":{}},{"cell_type":"code","source":"def add_path_to_df(df: pd.DataFrame, data_dir: Path, type_: str, stage: str) -> pd.DataFrame:\n    ending = \".tiff\" if type_ == \"image\" else \".npy\"\n    \n    dir_ = str(data_dir / f\"{stage}_{type_}s\") if type_ == \"image\" else f\"{stage}_{type_}s\"\n    df[type_] = dir_ + \"/\" + df[\"id\"].astype(str) + ending\n    return df\n\n\ndef add_paths_to_df(df: pd.DataFrame, data_dir: Path, stage: str) -> pd.DataFrame:\n    df = add_path_to_df(df, data_dir, \"image\", stage)\n    df = add_path_to_df(df, data_dir, \"mask\", stage)\n    return df\n\n\ndef create_folds(df: pd.DataFrame, n_splits: int, random_seed: int) -> pd.DataFrame:\n    skf = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=random_seed)\n    for fold, (_, val_idx) in enumerate(skf.split(X=df, y=df[\"organ\"])):\n        df.loc[val_idx, \"fold\"] = fold\n\n    return df\n\n\ndef prepare_data(data_dir: Path, stage: str, n_splits: int, random_seed: int) -> None:\n    df = pd.read_csv(data_dir / f\"{stage}.csv\")\n    df = add_paths_to_df(df, data_dir, stage)\n\n    if stage == \"train\":\n        df = create_folds(df, n_splits, random_seed)\n\n    filename = f\"{stage}_prepared.csv\"\n    df.to_csv(filename, index=False)\n\n    print(f\"Created {filename} with shape {df.shape}\")\n\n    return df","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:35.140743Z","iopub.execute_input":"2022-07-22T12:21:35.141102Z","iopub.status.idle":"2022-07-22T12:21:35.154511Z","shell.execute_reply.started":"2022-07-22T12:21:35.141067Z","shell.execute_reply":"2022-07-22T12:21:35.153386Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = prepare_data(COMPETITION_DATA_DIR, \"train\", N_SPLITS, RANDOM_SEED)\ntest_df = prepare_data(COMPETITION_DATA_DIR, \"test\", N_SPLITS, RANDOM_SEED)","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:35.156208Z","iopub.execute_input":"2022-07-22T12:21:35.156858Z","iopub.status.idle":"2022-07-22T12:21:36.022738Z","shell.execute_reply.started":"2022-07-22T12:21:35.156822Z","shell.execute_reply":"2022-07-22T12:21:36.021616Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:36.024532Z","iopub.execute_input":"2022-07-22T12:21:36.024949Z","iopub.status.idle":"2022-07-22T12:21:36.064919Z","shell.execute_reply.started":"2022-07-22T12:21:36.024910Z","shell.execute_reply":"2022-07-22T12:21:36.063774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:36.066610Z","iopub.execute_input":"2022-07-22T12:21:36.066997Z","iopub.status.idle":"2022-07-22T12:21:36.079657Z","shell.execute_reply.started":"2022-07-22T12:21:36.066960Z","shell.execute_reply":"2022-07-22T12:21:36.078238Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Save Train Masks as NumPy Arrays","metadata":{}},{"cell_type":"code","source":"def rle2mask(mask_rle: str, shape: Tuple[int, int]) -> np.ndarray:\n    \"\"\"\n    mask_rle: run-length as string formated (start length)\n    shape: (width,height) of array to return\n    Returns numpy array, 1 - mask, 0 - background\n    Source: https://www.kaggle.com/paulorzp/rle-functions-run-lenght-encode-decode\n    \"\"\"\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0] * shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape(shape).T\n\n\ndef save_array(file_path: str, array: np.ndarray) -> None:\n    file_path = Path(file_path)\n    file_path.parent.mkdir(parents=True, exist_ok=True)\n    np.save(file_path, array)\n\n\ndef save_masks(df: pd.DataFrame) -> None:\n    for row in tqdm(df.itertuples(), total=len(df)):\n        mask = rle2mask(row.rle, shape=(row.img_width, row.img_height))\n        save_array(row.mask, mask)","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:36.081300Z","iopub.execute_input":"2022-07-22T12:21:36.082531Z","iopub.status.idle":"2022-07-22T12:21:36.094730Z","shell.execute_reply.started":"2022-07-22T12:21:36.082486Z","shell.execute_reply":"2022-07-22T12:21:36.093739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_masks(train_df)","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:36.097716Z","iopub.execute_input":"2022-07-22T12:21:36.097984Z","iopub.status.idle":"2022-07-22T12:21:42.087885Z","shell.execute_reply.started":"2022-07-22T12:21:36.097961Z","shell.execute_reply":"2022-07-22T12:21:42.086690Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning DataModule","metadata":{}},{"cell_type":"code","source":"class TIFFImageReader(ImageReader):\n    def read(self, data: str) -> np.ndarray:\n        return tifffile.imread(data)\n\n    def get_data(self, img: np.ndarray) -> Tuple[np.ndarray, Dict[str, Any]]:\n        return img, {\"spatial_shape\": np.asarray(img.shape), \"original_channel_dim\": -1}\n\n    def verify_suffix(self, filename: str) -> bool:\n        return \".tiff\" in filename","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:42.092946Z","iopub.execute_input":"2022-07-22T12:21:42.093823Z","iopub.status.idle":"2022-07-22T12:21:42.103389Z","shell.execute_reply.started":"2022-07-22T12:21:42.093774Z","shell.execute_reply":"2022-07-22T12:21:42.101907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class LitDataModule(pl.LightningDataModule):\n    def __init__(\n        self,\n        train_csv_path: str,\n        test_csv_path: str,\n        spatial_size: int,\n        val_fold: int,\n        batch_size: int,\n        num_workers: int,\n    ):\n        super().__init__()\n\n        self.save_hyperparameters()\n\n        self.train_df = pd.read_csv(train_csv_path)\n        self.test_df = pd.read_csv(test_csv_path)\n\n        self.train_transform, self.val_transform, self.test_transform = self._init_transforms()\n\n    def _init_transforms(self) -> Tuple[Callable, Callable, Callable]:\n        spatial_size = (self.hparams.spatial_size, self.hparams.spatial_size)\n        train_transform = monai.transforms.Compose(\n            [\n                monai.transforms.LoadImaged(keys=[\"image\"], reader=TIFFImageReader),\n                monai.transforms.EnsureChannelFirstd(keys=[\"image\"]),\n                monai.transforms.ScaleIntensityd(keys=[\"image\"]),\n                monai.transforms.LoadImaged(keys=[\"mask\"]),\n                monai.transforms.AddChanneld(keys=[\"mask\"]),\n                monai.transforms.Resized(keys=[\"image\", \"mask\"], spatial_size=spatial_size),\n                monai.transforms.RandAxisFlipd(keys=[\"image\", \"mask\"], prob=0.5),\n                monai.transforms.RandRotate90d(keys=[\"image\", \"mask\"], prob=0.5),\n                monai.transforms.RandGridDistortiond(keys=[\"image\", \"mask\"], prob=0.5, distort_limit=0.2),\n                monai.transforms.OneOf(\n                    [\n                        monai.transforms.RandAdjustContrastd(keys=[\"image\"], prob=0.5, gamma=(1.5, 2.5)),\n                        monai.transforms.RandHistogramShiftd(keys=[\"image\"], prob=0.5),\n                    ]\n                ),\n            ]\n        )\n\n        val_transform = monai.transforms.Compose(\n            [\n                monai.transforms.LoadImaged(keys=[\"image\"], reader=TIFFImageReader),\n                monai.transforms.EnsureChannelFirstd(keys=[\"image\"]),\n                monai.transforms.ScaleIntensityd(keys=[\"image\"]),\n                monai.transforms.LoadImaged(keys=[\"mask\"]),\n                monai.transforms.AddChanneld(keys=[\"mask\"]),\n                monai.transforms.Resized(keys=[\"image\", \"mask\"], spatial_size=spatial_size),\n            ]\n        )\n\n        test_transform = monai.transforms.Compose(\n            [\n                monai.transforms.LoadImaged(keys=[\"image\"], reader=TIFFImageReader),\n                monai.transforms.EnsureChannelFirstd(keys=[\"image\"]),\n                monai.transforms.ScaleIntensityd(keys=[\"image\"]),\n                monai.transforms.Resized(keys=[\"image\"], spatial_size=spatial_size),\n            ]\n        )\n\n        return train_transform, val_transform, test_transform\n\n    def setup(self, stage: str = None):\n        if stage == \"fit\" or stage is None:\n            train_df = self.train_df[self.train_df.fold != self.hparams.val_fold].reset_index(drop=True)\n            val_df = self.train_df[self.train_df.fold == self.hparams.val_fold].reset_index(drop=True)\n\n            self.train_dataset = self._dataset(train_df, transform=self.train_transform)\n            self.val_dataset = self._dataset(val_df, transform=self.val_transform)\n\n        if stage == \"test\" or stage is None:\n            self.test_dataset = self._dataset(self.test_df, transform=self.test_transform)\n\n    def _dataset(self, df: pd.DataFrame, transform: Callable) -> CSVDataset:\n        return CSVDataset(src=df, transform=transform)\n\n    def train_dataloader(self) -> DataLoader:\n        return self._dataloader(self.train_dataset, train=True)\n\n    def val_dataloader(self) -> DataLoader:\n        return self._dataloader(self.val_dataset)\n\n    def test_dataloader(self) -> DataLoader:\n        return self._dataloader(self.test_dataset)\n\n    def _dataloader(self, dataset: CSVDataset, train: bool = False) -> DataLoader:\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        )","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:42.104999Z","iopub.execute_input":"2022-07-22T12:21:42.105352Z","iopub.status.idle":"2022-07-22T12:21:48.800560Z","shell.execute_reply.started":"2022-07-22T12:21:42.105317Z","shell.execute_reply":"2022-07-22T12:21:48.799480Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize Images and Masks","metadata":{}},{"cell_type":"code","source":"def show_image(title: str, image: np.ndarray, mask: Optional[np.ndarray] = None):\n    plt.title(title)\n    print(\"Image Shape: {}\".format(image.shape))\n    plt.imshow(image)\n\n    if mask is not None:\n        plt.imshow(mask, alpha=0.2)\n\n    plt.tight_layout()\n    plt.axis(\"off\")\n\n\ndef show_batch(batch: Dict, nrows: int, show_mask: bool = True):\n    fig, _ = plt.subplots(figsize=(3 * nrows, 3 * nrows))\n\n    for idx, _ in enumerate(batch[\"image\"]):\n        plt.subplot(nrows, nrows, idx + 1)\n\n        title = batch[\"id\"][idx].numpy()\n        image = np.transpose(batch[\"image\"][idx].numpy(), axes=(1, 2, 0))\n        mask = np.transpose(batch[\"mask\"][idx].numpy(), axes=(1, 2, 0)) if show_mask else None\n\n        show_image(title, image, mask)","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:48.806858Z","iopub.execute_input":"2022-07-22T12:21:48.807710Z","iopub.status.idle":"2022-07-22T12:21:48.831900Z","shell.execute_reply.started":"2022-07-22T12:21:48.807655Z","shell.execute_reply":"2022-07-22T12:21:48.830650Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Setup DataModule","metadata":{}},{"cell_type":"code","source":"nrows = 3\n\ndata_module = LitDataModule(\n    train_csv_path=TRAIN_PREPARED_CSV_PATH,\n    test_csv_path=TEST_PREPARED_CSV_PATH,\n    spatial_size=SPATIAL_SIZE,\n    val_fold=VAL_FOLD,\n    batch_size=nrows ** 2,\n    num_workers=NUM_WORKERS,\n)\ndata_module.setup()","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:48.838345Z","iopub.execute_input":"2022-07-22T12:21:48.841766Z","iopub.status.idle":"2022-07-22T12:21:49.050890Z","shell.execute_reply.started":"2022-07-22T12:21:48.841718Z","shell.execute_reply":"2022-07-22T12:21:49.049827Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Train Images","metadata":{}},{"cell_type":"code","source":"train_batch = next(iter(data_module.train_dataloader()))\nshow_batch(train_batch, nrows)","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:21:49.054878Z","iopub.execute_input":"2022-07-22T12:21:49.055231Z","iopub.status.idle":"2022-07-22T12:22:13.274362Z","shell.execute_reply.started":"2022-07-22T12:21:49.055199Z","shell.execute_reply":"2022-07-22T12:22:13.271782Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Test Images","metadata":{}},{"cell_type":"code","source":"test_batch = next(iter(data_module.test_dataloader()))\nshow_batch(test_batch, nrows, show_mask=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:22:13.275609Z","iopub.execute_input":"2022-07-22T12:22:13.276790Z","iopub.status.idle":"2022-07-22T12:22:14.232681Z","shell.execute_reply.started":"2022-07-22T12:22:13.276750Z","shell.execute_reply":"2022-07-22T12:22:14.231640Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning Module","metadata":{}},{"cell_type":"code","source":"class LitModule(pl.LightningModule):\n    def __init__(\n        self,\n        learning_rate: float,\n        weight_decay: float,\n    ):\n        super().__init__()\n\n        self.save_hyperparameters()\n\n        self.model = self._init_model()\n\n        self.loss_fn = self._init_loss_fn()\n\n        # TODO: add metric\n\n    def _init_model(self) -> nn.Module:\n        # TODO: try other networks\n        model = monai.networks.nets.AttentionUnet(\n            spatial_dims=2,\n            in_channels=3,\n            out_channels=1,\n            channels=(16, 32, 64, 128, 256),\n            strides=(2, 2, 2, 2),\n        )\n        \n#         model = monai.networks.nets.EfficientNetBN(\n#             model_name = 'efficientnet-b0', \n#             pretrained=True, \n#             progress=True,\n#             spatial_dims=2,\n#             in_channels=3,\n#         )\n#         strides = [2, 2, 2, 2, 2]\n#         model = monai.networks.nets.DynUNet(\n#             spatial_dims = 2,\n#             in_channels = 3,\n#             out_channels = 1,\n#             kernel_size = [3,3,3,3,3],\n#             upsample_kernel_size = [3,3,3,3,3],\n#             filters = [96, 128, 192, 256, 384, 512, 768, 1024][: len(strides)],\n#             strides = strides,\n# #             dropout = [0, 0.4, 0.5, 0.3, 0],\n#         )\n        return model\n\n    def _init_loss_fn(self):\n        # TODO: try other losses\n        return monai.losses.DiceLoss(sigmoid=True)\n\n    def configure_optimizers(self):\n        # TODO: try other optimizers and schedulers\n        return torch.optim.Adam(\n            params=self.parameters(), lr=self.hparams.learning_rate, weight_decay=self.hparams.weight_decay\n        )\n\n    def forward(self, images: torch.Tensor) -> torch.Tensor:\n        return self.model(images)\n\n    def training_step(self, batch: Dict, batch_idx: int) -> torch.Tensor:\n        images, masks = batch[\"image\"], batch[\"mask\"]\n        outputs = self(images)\n\n        loss = self.loss_fn(outputs, masks)\n\n        self.log(\"train_loss\", loss, batch_size=images.shape[0])\n\n        return loss\n\n    def validation_step(self, batch: Dict, batch_idx: int) -> None:\n        images, masks = batch[\"image\"], batch[\"mask\"]\n        outputs = self(images)\n\n        loss = self.loss_fn(outputs, masks)\n\n        self.log(\"val_loss\", loss, prog_bar=True, batch_size=images.shape[0])\n\n    @classmethod\n    def load_eval_checkpoint(cls, checkpoint_path: str, device: str) -> nn.Module:\n        module = cls.load_from_checkpoint(checkpoint_path=checkpoint_path).to(device)\n        module.eval()\n\n        return module","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:22:14.234637Z","iopub.execute_input":"2022-07-22T12:22:14.235342Z","iopub.status.idle":"2022-07-22T12:22:14.250007Z","shell.execute_reply.started":"2022-07-22T12:22:14.235298Z","shell.execute_reply":"2022-07-22T12:22:14.249135Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"def train(\n    random_seed: int = RANDOM_SEED,\n    train_csv_path: str = str(TRAIN_PREPARED_CSV_PATH),\n    test_csv_path: str = str(TEST_PREPARED_CSV_PATH),\n    spatial_size: Tuple[int, int] = SPATIAL_SIZE,\n    val_fold: str = VAL_FOLD,\n    batch_size: int = BATCH_SIZE,\n    num_workers: int = NUM_WORKERS,\n    learning_rate: float = LEARNING_RATE,\n    weight_decay: float = WEIGHT_DECAY,\n    fast_dev_run: bool = FAST_DEV_RUN,\n    gpus: int = GPUS,\n    max_epochs: int = MAX_EPOCHS,\n    precision: int = PRECISION,\n    debug: bool = DEBUG,\n) -> None:\n    pl.seed_everything(random_seed)\n\n    data_module = LitDataModule(\n        train_csv_path=train_csv_path,\n        test_csv_path=test_csv_path,\n        spatial_size=spatial_size,\n        val_fold=val_fold,\n        batch_size=2 if debug else batch_size,\n        num_workers=num_workers,\n    )\n\n    module = LitModule(\n        learning_rate=learning_rate,\n        weight_decay=weight_decay,\n    )\n\n    trainer = pl.Trainer(\n        fast_dev_run=fast_dev_run,\n        gpus=gpus,\n        limit_train_batches=0.1 if debug else 1.0,\n        limit_val_batches=0.1 if debug else 1.0,\n        log_every_n_steps=5,\n        logger=pl.loggers.CSVLogger(save_dir='logs/'),\n        max_epochs=2 if debug else max_epochs,\n        precision=precision,\n    )\n\n    trainer.fit(module, datamodule=data_module)\n    \n    return trainer","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:22:14.250987Z","iopub.execute_input":"2022-07-22T12:22:14.252048Z","iopub.status.idle":"2022-07-22T12:22:14.264648Z","shell.execute_reply.started":"2022-07-22T12:22:14.252011Z","shell.execute_reply":"2022-07-22T12:22:14.263576Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = train()","metadata":{"execution":{"iopub.status.busy":"2022-07-22T12:22:31.669279Z","iopub.execute_input":"2022-07-22T12:22:31.669891Z","iopub.status.idle":"2022-07-22T12:23:37.810917Z","shell.execute_reply.started":"2022-07-22T12:22:31.669853Z","shell.execute_reply":"2022-07-22T12:23:37.808642Z"},"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\")[[\"epoch\", \"train_loss\", \"val_loss\"]]\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":"2022-07-17T11:56:26.696324Z","iopub.execute_input":"2022-07-17T11:56:26.697049Z","iopub.status.idle":"2022-07-17T11:56:27.115355Z","shell.execute_reply.started":"2022-07-17T11:56:26.697009Z","shell.execute_reply":"2022-07-17T11:56:27.114445Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Infer","metadata":{}},{"cell_type":"code","source":"def mask2rle(img):\n    '''\n    Efficient implementation of mask2rle, from @paulorzp\n    --\n    img: numpy array, 1 - mask, 0 - background\n    Returns run length as string formated\n    Source: https://www.kaggle.com/xhlulu/efficient-mask2rle\n    '''\n    pixels = img.T.flatten()\n    pixels = np.pad(pixels, ((1, 1), ))\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\n\n@torch.no_grad()\ndef create_pred_df(module, dataloader, threshold):\n    ids = []\n    rles = []\n    for batch in tqdm(dataloader):\n        id_ = batch[\"id\"].numpy()[0]\n        height = batch[\"img_height\"].numpy()[0]\n        width = batch[\"img_width\"].numpy()[0]\n        \n        images = batch[\"image\"].to(module.device)\n        outputs = module(images)[0]\n        \n        post_pred_transform = monai.transforms.Compose(\n            [\n                monai.transforms.Resize(spatial_size=(height, width), mode=\"nearest\"),\n                monai.transforms.Activations(sigmoid=True),\n                monai.transforms.AsDiscrete(threshold=threshold),\n            ]\n        )\n        \n        mask = post_pred_transform(outputs).to(torch.uint8).cpu().detach().numpy()[0]\n        \n        rle = mask2rle(mask)\n        \n        ids.append(id_)\n        rles.append(rle)\n        \n    return pd.DataFrame({\"id\": ids, \"rle\": rles})\n\n\ndef infer(\n    checkpoint_path: str,\n    device: str = DEVICE,\n    train_csv_path: str = TRAIN_PREPARED_CSV_PATH,\n    test_csv_path: str = TEST_PREPARED_CSV_PATH,\n    spatial_size: int = SPATIAL_SIZE,\n    num_workers: int = NUM_WORKERS,\n    threshold: float = THRESHOLD,\n):\n    module = LitModule.load_eval_checkpoint(checkpoint_path, device)\n\n    data_module = LitDataModule(\n        train_csv_path=train_csv_path,\n        test_csv_path=test_csv_path,\n        spatial_size=spatial_size,\n        val_fold=0,\n        batch_size=1,\n        num_workers=num_workers,\n    )\n    data_module.setup()\n    \n    val_dataloader = data_module.val_dataloader()\n    test_dataloader = data_module.test_dataloader()\n    \n    val_pred_df = create_pred_df(module, val_dataloader, threshold)\n    test_pred_df = create_pred_df(module, test_dataloader, threshold)\n    \n    return val_pred_df, test_pred_df\n","metadata":{"execution":{"iopub.status.busy":"2022-07-17T11:56:27.116864Z","iopub.execute_input":"2022-07-17T11:56:27.119798Z","iopub.status.idle":"2022-07-17T11:56:27.134327Z","shell.execute_reply.started":"2022-07-17T11:56:27.119755Z","shell.execute_reply":"2022-07-17T11:56:27.133362Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoint_path = list((Path(trainer.logger.log_dir) / \"checkpoints\").glob(\"*.ckpt\"))[0]\nval_pred_df, test_pred_df = infer(checkpoint_path)","metadata":{"execution":{"iopub.status.busy":"2022-07-17T11:56:27.135884Z","iopub.execute_input":"2022-07-17T11:56:27.13625Z","iopub.status.idle":"2022-07-17T11:57:43.667063Z","shell.execute_reply.started":"2022-07-17T11:56:27.136211Z","shell.execute_reply":"2022-07-17T11:57:43.665919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submit","metadata":{}},{"cell_type":"code","source":"test_pred_df.to_csv(\"submission.csv\", index=False)\ntest_pred_df","metadata":{"execution":{"iopub.status.busy":"2022-07-17T11:57:43.668947Z","iopub.execute_input":"2022-07-17T11:57:43.669687Z","iopub.status.idle":"2022-07-17T11:57:43.803867Z","shell.execute_reply.started":"2022-07-17T11:57:43.669645Z","shell.execute_reply":"2022-07-17T11:57:43.80274Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualize Val Predictions","metadata":{}},{"cell_type":"code","source":"val_pred_df = add_path_to_df(val_pred_df, COMPETITION_DATA_DIR, \"mask\", \"pred\")\nval_pred_df","metadata":{"execution":{"iopub.status.busy":"2022-07-17T11:57:43.805659Z","iopub.execute_input":"2022-07-17T11:57:43.806041Z","iopub.status.idle":"2022-07-17T11:57:44.011595Z","shell.execute_reply.started":"2022-07-17T11:57:43.806003Z","shell.execute_reply":"2022-07-17T11:57:44.010624Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_df = train_df[train_df.fold == VAL_FOLD].reset_index(drop=True)\nval_df","metadata":{"execution":{"iopub.status.busy":"2022-07-17T11:57:44.013014Z","iopub.execute_input":"2022-07-17T11:57:44.013494Z","iopub.status.idle":"2022-07-17T11:57:44.046441Z","shell.execute_reply.started":"2022-07-17T11:57:44.013455Z","shell.execute_reply":"2022-07-17T11:57:44.045416Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_pred_df = val_pred_df.merge(val_df, on=\"id\", suffixes=(\"\", \"_gt\"))\nval_pred_df","metadata":{"execution":{"iopub.status.busy":"2022-07-17T11:57:44.048093Z","iopub.execute_input":"2022-07-17T11:57:44.048464Z","iopub.status.idle":"2022-07-17T11:57:44.274179Z","shell.execute_reply.started":"2022-07-17T11:57:44.048427Z","shell.execute_reply":"2022-07-17T11:57:44.272921Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_masks(val_pred_df)","metadata":{"execution":{"iopub.status.busy":"2022-07-17T11:57:44.276379Z","iopub.execute_input":"2022-07-17T11:57:44.276861Z","iopub.status.idle":"2022-07-17T11:58:18.376206Z","shell.execute_reply.started":"2022-07-17T11:57:44.27682Z","shell.execute_reply":"2022-07-17T11:58:18.375203Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"val_pred_df.to_csv(VAL_PRED_PREPARED_CSV_PATH, index=False)","metadata":{"execution":{"iopub.status.busy":"2022-07-17T11:58:18.377933Z","iopub.execute_input":"2022-07-17T11:58:18.37864Z","iopub.status.idle":"2022-07-17T11:58:27.933858Z","shell.execute_reply.started":"2022-07-17T11:58:18.378602Z","shell.execute_reply":"2022-07-17T11:58:27.932837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## GT Val Images","metadata":{}},{"cell_type":"code","source":"nrows = 3\n\ndata_module = LitDataModule(\n    train_csv_path=TRAIN_PREPARED_CSV_PATH,\n    test_csv_path=TEST_PREPARED_CSV_PATH,\n    spatial_size=SPATIAL_SIZE,\n    val_fold=VAL_FOLD,\n    batch_size=nrows ** 2,\n    num_workers=0,\n)\ndata_module.setup()\n\nval_batch = next(iter(data_module.val_dataloader()))\nshow_batch(val_batch, nrows)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-17T11:58:27.939422Z","iopub.execute_input":"2022-07-17T11:58:27.940006Z","iopub.status.idle":"2022-07-17T11:58:44.644274Z","shell.execute_reply.started":"2022-07-17T11:58:27.939974Z","shell.execute_reply":"2022-07-17T11:58:44.639939Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Pred Val Images","metadata":{}},{"cell_type":"code","source":"data_module = LitDataModule(\n    train_csv_path=VAL_PRED_PREPARED_CSV_PATH,\n    test_csv_path=TEST_PREPARED_CSV_PATH,\n    spatial_size=SPATIAL_SIZE,\n    val_fold=VAL_FOLD,\n    batch_size=nrows ** 2,\n    num_workers=0,\n)\ndata_module.setup()\n\nval_batch = next(iter(data_module.val_dataloader()))\nshow_batch(val_batch, nrows)","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-07-17T11:58:44.646097Z","iopub.execute_input":"2022-07-17T11:58:44.646739Z","iopub.status.idle":"2022-07-17T11:59:04.149753Z","shell.execute_reply.started":"2022-07-17T11:58:44.6467Z","shell.execute_reply":"2022-07-17T11:59:04.148843Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]}]}