{"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 ⚡ SMP Train & Infer\n\n## Combine powers of [PyTorch Lightning](https://www.pytorchlightning.ai/) and [Segmentation Models PyTorch](https://github.com/qubvel/segmentation_models.pytorch)","metadata":{}},{"cell_type":"markdown","source":"# Installs","metadata":{}},{"cell_type":"code","source":"!cd ../input/hubmap-downloads && \\\npip install efficientnet_pytorch-0.6.3.tar.gz pretrainedmodels-0.7.4.tar.gz timm-0.4.12-py3-none-any.whl  segmentation_models_pytorch-0.2.1-py3-none-any.whl && \\\npip install monai-0.9.0-202206131636-py3-none-any.whl","metadata":{"execution":{"iopub.status.busy":"2022-07-21T10:37:37.237668Z","iopub.execute_input":"2022-07-21T10:37:37.238103Z","iopub.status.idle":"2022-07-21T10:38:40.918654Z","shell.execute_reply.started":"2022-07-21T10:37:37.238016Z","shell.execute_reply":"2022-07-21T10:38:40.917491Z"},"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 albumentations as A\nimport cv2\nimport monai\nimport numpy as np\nimport pandas as pd\nimport pytorch_lightning as pl\nimport seaborn as sns\nimport segmentation_models_pytorch as smp\nimport tifffile\nimport torch\nimport torch.nn as nn\nfrom albumentations.pytorch import ToTensorV2\nfrom matplotlib import pyplot as plt\nfrom sklearn.model_selection import StratifiedKFold\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import Dataset\nfrom torchmetrics import Dice\nfrom torchmetrics import MetricCollection\nfrom tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-07-21T10:38:40.921227Z","iopub.execute_input":"2022-07-21T10:38:40.921964Z","iopub.status.idle":"2022-07-21T10:38:48.982300Z","shell.execute_reply.started":"2022-07-21T10:38:40.921921Z","shell.execute_reply":"2022-07-21T10:38:48.981282Z"},"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-21T10:38:48.985152Z","iopub.execute_input":"2022-07-21T10:38:48.986537Z","iopub.status.idle":"2022-07-21T10:38:48.994356Z","shell.execute_reply.started":"2022-07-21T10:38:48.986492Z","shell.execute_reply":"2022-07-21T10:38:48.993094Z"},"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-21T10:38:48.997233Z","iopub.execute_input":"2022-07-21T10:38:48.997692Z","iopub.status.idle":"2022-07-21T10:38:49.011231Z","shell.execute_reply.started":"2022-07-21T10:38:48.997657Z","shell.execute_reply":"2022-07-21T10:38:49.010283Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = prepare_data(data_dir=COMPETITION_DATA_DIR, stage=\"train\", n_splits=N_SPLITS, random_seed=RANDOM_SEED)\ntest_df = prepare_data(data_dir=COMPETITION_DATA_DIR, stage=\"test\", n_splits=N_SPLITS, random_seed=RANDOM_SEED)","metadata":{"execution":{"iopub.status.busy":"2022-07-21T10:38:49.012645Z","iopub.execute_input":"2022-07-21T10:38:49.013255Z","iopub.status.idle":"2022-07-21T10:38:49.890702Z","shell.execute_reply.started":"2022-07-21T10:38:49.013218Z","shell.execute_reply":"2022-07-21T10:38:49.889652Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df","metadata":{"execution":{"iopub.status.busy":"2022-07-21T10:38:49.892226Z","iopub.execute_input":"2022-07-21T10:38:49.892597Z","iopub.status.idle":"2022-07-21T10:38:49.929389Z","shell.execute_reply.started":"2022-07-21T10:38:49.892560Z","shell.execute_reply":"2022-07-21T10:38:49.928392Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df","metadata":{"execution":{"iopub.status.busy":"2022-07-21T10:38:49.930838Z","iopub.execute_input":"2022-07-21T10:38:49.931243Z","iopub.status.idle":"2022-07-21T10:38:49.943804Z","shell.execute_reply.started":"2022-07-21T10:38:49.931210Z","shell.execute_reply":"2022-07-21T10:38:49.942677Z"},"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-21T10:38:49.945477Z","iopub.execute_input":"2022-07-21T10:38:49.946123Z","iopub.status.idle":"2022-07-21T10:38:49.957886Z","shell.execute_reply.started":"2022-07-21T10:38:49.946086Z","shell.execute_reply":"2022-07-21T10:38:49.956835Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_masks(train_df)","metadata":{"execution":{"iopub.status.busy":"2022-07-21T10:38:49.960677Z","iopub.execute_input":"2022-07-21T10:38:49.961103Z","iopub.status.idle":"2022-07-21T10:38:55.712565Z","shell.execute_reply.started":"2022-07-21T10:38:49.961067Z","shell.execute_reply":"2022-07-21T10:38:55.711326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning DataModule","metadata":{}},{"cell_type":"code","source":"class TIFFDataset(Dataset):\n    def __init__(self, src: pd.DataFrame, transform: Callable, load_mask: bool = True):\n        self.src = src\n        self.transform = transform\n        self.load_mask = load_mask\n\n    def __len__(self) -> int:\n        return len(self.src)\n\n    def __getitem__(self, idx: int) -> Dict[str, Any]:\n        row = self.src.iloc[idx]\n\n        image = tifffile.imread(row[\"image\"])\n\n        if self.load_mask:\n            mask = np.load(row[\"mask\"])\n\n            data = self.transform(image=image, mask=mask)\n            image, mask = data[\"image\"], data[\"mask\"]\n\n            return {\"id\": row[\"id\"], \"image\": image, \"mask\": mask, \"img_height\": row[\"img_height\"], \"img_width\": row[\"img_width\"]}\n        else:\n            data = self.transform(image=image)\n            image = data[\"image\"]\n\n            return {\"id\": row[\"id\"], \"image\": image, \"img_height\": row[\"img_height\"], \"img_width\": row[\"img_width\"]}","metadata":{"execution":{"iopub.status.busy":"2022-07-21T10:44:26.262299Z","iopub.execute_input":"2022-07-21T10:44:26.263090Z","iopub.status.idle":"2022-07-21T10:44:26.273569Z","shell.execute_reply.started":"2022-07-21T10:44:26.263051Z","shell.execute_reply":"2022-07-21T10:44:26.272119Z"},"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\n        # TODO: add more augmentations\n        train_transform = A.Compose(\n            [\n                A.HorizontalFlip(),\n                A.VerticalFlip(),\n                A.RandomRotate90(),\n                A.ShiftScaleRotate(\n                    shift_limit=0.0625, scale_limit=0.2, rotate_limit=15, p=0.9, border_mode=cv2.BORDER_REFLECT\n                ),\n                A.OneOf(\n                    [\n                        A.OpticalDistortion(p=0.3),\n                        A.GridDistortion(p=0.1),\n                        A.PiecewiseAffine(p=0.3),\n                    ],\n                    p=0.3,\n                ),\n                A.OneOf(\n                    [\n                        A.HueSaturationValue(10, 15, 10),\n                        A.CLAHE(clip_limit=2),\n                        A.RandomBrightnessContrast(),\n                    ],\n                    p=0.3,\n                ),\n                A.ToFloat(),\n                A.Resize(height=spatial_size[0], width=spatial_size[1]),\n                ToTensorV2(),\n            ]\n        )\n\n        val_transform = A.Compose(\n            [\n                A.ToFloat(),\n                A.Resize(height=spatial_size[0], width=spatial_size[1]),\n                ToTensorV2(),\n            ]\n        )\n\n        test_transform = A.Compose(\n            [\n                A.ToFloat(),\n                A.Resize(height=spatial_size[0], width=spatial_size[1]),\n                ToTensorV2(),\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, load_mask=False)\n\n    def _dataset(self, df: pd.DataFrame, transform: Callable, load_mask: bool = True) -> TIFFDataset:\n        return TIFFDataset(src=df, transform=transform, load_mask=load_mask)\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: TIFFDataset, 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-21T10:45:29.165741Z","iopub.execute_input":"2022-07-21T10:45:29.166311Z","iopub.status.idle":"2022-07-21T10:45:29.189654Z","shell.execute_reply.started":"2022-07-21T10:45:29.166271Z","shell.execute_reply":"2022-07-21T10:45:29.187942Z"},"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    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 = batch[\"mask\"][idx].numpy() if show_mask else None\n\n        show_image(title, image, mask)","metadata":{"execution":{"iopub.status.busy":"2022-07-21T10:39:02.966385Z","iopub.execute_input":"2022-07-21T10:39:02.967609Z","iopub.status.idle":"2022-07-21T10:39:02.982972Z","shell.execute_reply.started":"2022-07-21T10:39:02.967563Z","shell.execute_reply":"2022-07-21T10:39:02.981982Z"},"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-21T10:39:03.001540Z","iopub.execute_input":"2022-07-21T10:39:03.001918Z","iopub.status.idle":"2022-07-21T10:39:03.152577Z","shell.execute_reply.started":"2022-07-21T10:39:03.001883Z","shell.execute_reply":"2022-07-21T10:39:03.150998Z"},"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-21T10:39:03.154306Z","iopub.execute_input":"2022-07-21T10:39:03.154775Z","iopub.status.idle":"2022-07-21T10:39:47.314876Z","shell.execute_reply.started":"2022-07-21T10:39:03.154735Z","shell.execute_reply":"2022-07-21T10:39:47.313750Z"},"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-21T10:39:47.316063Z","iopub.execute_input":"2022-07-21T10:39:47.316406Z","iopub.status.idle":"2022-07-21T10:39:48.953289Z","shell.execute_reply.started":"2022-07-21T10:39:47.316375Z","shell.execute_reply":"2022-07-21T10:39:48.952272Z"},"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        threshold: float = 0.5,\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        self.metrics = self._init_metrics()\n\n    def _init_model(self) -> nn.Module:\n        # TODO: try other networks\n        return smp.create_model(\n            arch=\"Unet\",\n            encoder_name=\"resnet18\",\n            encoder_weights=None,\n            in_channels=3,\n            classes=1,\n            activation=None,\n        )\n\n    def _init_loss_fn(self):\n        # TODO: try other losses\n        return smp.losses.DiceLoss(mode=\"binary\")\n\n    def _init_metrics(self):\n        return torch.nn.ModuleDict(\n            {\n                \"train_metrics\": MetricCollection({\"train_dice\": Dice()}),\n                \"val_metrics\": MetricCollection({\"val_dice\": Dice()}),\n            }\n        )\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        metrics = self.metrics[\"train_metrics\"](outputs, masks)\n\n        self.log(\"train_loss\", loss, batch_size=images.shape[0])\n        self.log_dict(metrics, 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        metrics = self.metrics[\"val_metrics\"](outputs, masks)\n\n        self.log(\"val_loss\", loss, prog_bar=True, batch_size=images.shape[0])\n        self.log_dict(metrics, 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-21T10:39:48.955394Z","iopub.execute_input":"2022-07-21T10:39:48.956180Z","iopub.status.idle":"2022-07-21T10:39:48.974655Z","shell.execute_reply.started":"2022-07-21T10:39:48.956132Z","shell.execute_reply":"2022-07-21T10:39:48.973737Z"},"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-21T10:39:48.976354Z","iopub.execute_input":"2022-07-21T10:39:48.976761Z","iopub.status.idle":"2022-07-21T10:39:48.989309Z","shell.execute_reply.started":"2022-07-21T10:39:48.976725Z","shell.execute_reply":"2022-07-21T10:39:48.988342Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = train()","metadata":{"execution":{"iopub.status.busy":"2022-07-21T10:39:48.990729Z","iopub.execute_input":"2022-07-21T10:39:48.991090Z","iopub.status.idle":"2022-07-21T10:41:50.203834Z","shell.execute_reply.started":"2022-07-21T10:39:48.991053Z","shell.execute_reply":"2022-07-21T10:41:50.202675Z"},"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-21T10:41:50.205537Z","iopub.execute_input":"2022-07-21T10:41:50.205926Z","iopub.status.idle":"2022-07-21T10:41:50.626014Z","shell.execute_reply.started":"2022-07-21T10:41:50.205887Z","shell.execute_reply":"2022-07-21T10:41:50.624192Z"},"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-21T10:44:33.966473Z","iopub.execute_input":"2022-07-21T10:44:33.966831Z","iopub.status.idle":"2022-07-21T10:44:33.982314Z","shell.execute_reply.started":"2022-07-21T10:44:33.966803Z","shell.execute_reply":"2022-07-21T10:44:33.981356Z"},"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-21T10:45:37.881388Z","iopub.execute_input":"2022-07-21T10:45:37.882084Z","iopub.status.idle":"2022-07-21T10:45:53.585547Z","shell.execute_reply.started":"2022-07-21T10:45:37.882045Z","shell.execute_reply":"2022-07-21T10:45:53.584461Z"},"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-21T10:45:53.588091Z","iopub.execute_input":"2022-07-21T10:45:53.588495Z","iopub.status.idle":"2022-07-21T10:45:53.606221Z","shell.execute_reply.started":"2022-07-21T10:45:53.588454Z","shell.execute_reply":"2022-07-21T10:45:53.605371Z"},"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-21T10:45:53.608087Z","iopub.execute_input":"2022-07-21T10:45:53.609115Z","iopub.status.idle":"2022-07-21T10:45:53.623836Z","shell.execute_reply.started":"2022-07-21T10:45:53.609079Z","shell.execute_reply":"2022-07-21T10:45:53.622851Z"},"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-21T10:45:53.627139Z","iopub.execute_input":"2022-07-21T10:45:53.628644Z","iopub.status.idle":"2022-07-21T10:45:53.662449Z","shell.execute_reply.started":"2022-07-21T10:45:53.628608Z","shell.execute_reply":"2022-07-21T10:45:53.661563Z"},"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-21T10:45:53.664474Z","iopub.execute_input":"2022-07-21T10:45:53.665090Z","iopub.status.idle":"2022-07-21T10:45:53.702139Z","shell.execute_reply.started":"2022-07-21T10:45:53.665053Z","shell.execute_reply":"2022-07-21T10:45:53.701254Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"save_masks(val_pred_df)","metadata":{"execution":{"iopub.status.busy":"2022-07-21T10:45:53.703573Z","iopub.execute_input":"2022-07-21T10:45:53.703908Z","iopub.status.idle":"2022-07-21T10:45:54.529421Z","shell.execute_reply.started":"2022-07-21T10:45:53.703876Z","shell.execute_reply":"2022-07-21T10:45:54.528398Z"},"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-21T10:45:54.530893Z","iopub.execute_input":"2022-07-21T10:45:54.531918Z","iopub.status.idle":"2022-07-21T10:45:54.702401Z","shell.execute_reply.started":"2022-07-21T10:45:54.531880Z","shell.execute_reply":"2022-07-21T10:45:54.701385Z"},"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-21T10:45:54.704026Z","iopub.execute_input":"2022-07-21T10:45:54.704417Z","iopub.status.idle":"2022-07-21T10:45:59.597271Z","shell.execute_reply.started":"2022-07-21T10:45:54.704379Z","shell.execute_reply":"2022-07-21T10:45:59.596263Z"},"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-21T10:45:59.598735Z","iopub.execute_input":"2022-07-21T10:45:59.599669Z","iopub.status.idle":"2022-07-21T10:46:04.314897Z","shell.execute_reply.started":"2022-07-21T10:45:59.599622Z","shell.execute_reply":"2022-07-21T10:46:04.314018Z"},"jupyter":{"source_hidden":true},"trusted":true},"execution_count":null,"outputs":[]}]}