{"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 ⚡ Train","metadata":{}},{"cell_type":"markdown","source":"# Installs","metadata":{}},{"cell_type":"code","source":"!pip install segmentation_models_pytorch","metadata":{"execution":{"iopub.status.busy":"2022-07-12T14:45:05.193923Z","iopub.execute_input":"2022-07-12T14:45:05.194555Z","iopub.status.idle":"2022-07-12T14:45:14.598625Z","shell.execute_reply.started":"2022-07-12T14:45:05.194519Z","shell.execute_reply":"2022-07-12T14:45:14.597326Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"from pathlib import Path\nfrom typing import Any\nfrom typing import Callable\nfrom typing import Dict\nfrom typing import Tuple\n\n\nimport albumentations as A\nimport cv2\nimport numpy as np\nimport pandas as pd\nimport pytorch_lightning as pl\nimport torch\nimport torch.nn as nn\nimport segmentation_models_pytorch as smp\nfrom albumentations.pytorch import ToTensorV2\nfrom matplotlib import pyplot as plt\nfrom sklearn.model_selection import StratifiedKFold\nfrom torch.utils.data import Dataset\nfrom torch.utils.data import DataLoader\nfrom torchmetrics import Dice\nfrom torchmetrics import MetricCollection","metadata":{"execution":{"iopub.status.busy":"2022-07-12T14:48:32.030694Z","iopub.execute_input":"2022-07-12T14:48:32.031155Z","iopub.status.idle":"2022-07-12T14:48:32.045675Z","shell.execute_reply.started":"2022-07-12T14:48:32.031112Z","shell.execute_reply":"2022-07-12T14:48:32.044474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Paths & Settings","metadata":{}},{"cell_type":"code","source":"KAGGLE_DIR = Path(\"/\") / \"kaggle\"\nINPUT_DIR = KAGGLE_DIR / \"input\"\n\nCOMPETITION_DATA_DIR = INPUT_DIR / \"hubmap-organ-segmentation\"\nPREPROCESSED_DATA_DIR = INPUT_DIR / \"hubmap-2022-256x256\"\n\nTRAIN_CSV_PATH = COMPETITION_DATA_DIR / \"train.csv\"\nTRAIN_FOLDS_CSV_PATH = \"train_folds.csv\"\nN_SPLITS = 4\nRANDOM_SEED = 2022\n\nVAL_FOLD = 0\nBATCH_SIZE = 64\nNUM_WORKERS = 2\nARCH = \"Unet\"\nENCODER_NAME = \"efficientnet-b1\"\nENCODER_WEIGHTS = \"imagenet\"\nLOSS = \"dice\"\nOPTIMIZER = \"Adam\"\nLEARNING_RATE = 3e-4\nWEIGHT_DECAY = 1e-6\nSCHEDULER = None\nMIN_LR = 1e-6\n\nFAST_DEV_RUN = False # Debug training\nGPUS = 1\nMAX_EPOCHS = 20\nPRECISION = 16","metadata":{"execution":{"iopub.status.busy":"2022-07-12T14:48:32.763120Z","iopub.execute_input":"2022-07-12T14:48:32.763829Z","iopub.status.idle":"2022-07-12T14:48:32.771401Z","shell.execute_reply.started":"2022-07-12T14:48:32.763790Z","shell.execute_reply":"2022-07-12T14:48:32.770210Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Create Folds","metadata":{}},{"cell_type":"code","source":"def 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","metadata":{"execution":{"iopub.status.busy":"2022-07-12T14:48:33.574988Z","iopub.execute_input":"2022-07-12T14:48:33.576175Z","iopub.status.idle":"2022-07-12T14:48:33.582523Z","shell.execute_reply.started":"2022-07-12T14:48:33.576131Z","shell.execute_reply":"2022-07-12T14:48:33.581623Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = pd.read_csv(TRAIN_CSV_PATH)\ntrain_df = create_folds(train_df, N_SPLITS, RANDOM_SEED)\ntrain_df","metadata":{"execution":{"iopub.status.busy":"2022-07-12T14:48:34.115443Z","iopub.execute_input":"2022-07-12T14:48:34.116710Z","iopub.status.idle":"2022-07-12T14:48:34.282618Z","shell.execute_reply.started":"2022-07-12T14:48:34.116661Z","shell.execute_reply":"2022-07-12T14:48:34.281440Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class HuBMAPDataset(Dataset):\n    def __init__(self, df: pd.DataFrame, preprocessed_data_dir: str, transforms: Callable):\n        ids = df.id.astype(str).values\n        \n        preprocessed_data_dir = Path(preprocessed_data_dir)\n        \n        all_image_paths = sorted((preprocessed_data_dir / \"train\").iterdir())\n        self.image_paths = [path for path in all_image_paths if path.stem.split(\"_\")[0] in ids]\n\n        all_mask_paths = sorted((preprocessed_data_dir / \"masks\").iterdir())\n        self.mask_paths = [path for path in all_mask_paths if path.stem.split(\"_\")[0] in ids]\n\n        self.transforms = transforms\n\n    def __len__(self) -> int:\n        return len(self.image_paths)\n\n    def __getitem__(self, idx: int) -> Tuple[torch.Tensor, torch.Tensor]:\n        image_path = self.image_paths[idx]\n        mask_path = self.mask_paths[idx]\n\n        assert image_path.stem == mask_path.stem\n\n        image = self._load_image(image_path)\n        mask = self._load_mask(mask_path)\n\n        data = self.transforms(image=image, mask=mask)\n        image, mask = data[\"image\"], data[\"mask\"]\n\n        return image, mask\n\n    @staticmethod\n    def _load_image(image_path: Path) -> np.ndarray:\n        return cv2.cvtColor(cv2.imread(str(image_path)), cv2.COLOR_BGR2RGB)\n\n    @staticmethod\n    def _load_mask(mask_path: Path) -> np.ndarray:\n        return cv2.imread(str(mask_path), cv2.IMREAD_GRAYSCALE)","metadata":{"execution":{"iopub.status.busy":"2022-07-12T15:01:02.821344Z","iopub.execute_input":"2022-07-12T15:01:02.822339Z","iopub.status.idle":"2022-07-12T15:01:02.835387Z","shell.execute_reply.started":"2022-07-12T15:01:02.822285Z","shell.execute_reply":"2022-07-12T15:01:02.834417Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning Data Module","metadata":{"execution":{"iopub.status.busy":"2022-07-04T13:33:09.510698Z","iopub.execute_input":"2022-07-04T13:33:09.511071Z","iopub.status.idle":"2022-07-04T13:33:09.540288Z","shell.execute_reply.started":"2022-07-04T13:33:09.511040Z","shell.execute_reply":"2022-07-04T13:33:09.538479Z"}}},{"cell_type":"code","source":"class LitDataModule(pl.LightningDataModule):\n    def __init__(\n        self,\n        train_df: pd.DataFrame,\n        preprocessed_data_dir: str,\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_transforms, self.val_transforms = self._init_transforms()\n\n    def _init_transforms(self) -> Tuple[Callable, Callable]:\n        train_transforms = [\n            A.Normalize(\n                mean=(0.7720342, 0.74582646, 0.76392896),\n                std=(0.24745085, 0.26182273, 0.25782376),\n            ),\n            ToTensorV2(),\n        ]\n\n        val_transforms = [\n            A.Normalize(\n                mean=(0.7720342, 0.74582646, 0.76392896),\n                std=(0.24745085, 0.26182273, 0.25782376),\n            ),\n            ToTensorV2(),\n        ]\n\n        return A.Compose(train_transforms), A.Compose(val_transforms)\n\n    def setup(self, stage: str = None):\n        train_df = self.hparams.train_df[self.hparams.train_df.fold != self.hparams.val_fold].reset_index(drop=True)\n        val_df = self.hparams.train_df[self.hparams.train_df.fold == self.hparams.val_fold].reset_index(drop=True)\n\n        self.train_dataset = self._dataset(train_df, transforms=self.train_transforms)\n        self.val_dataset = self._dataset(val_df, transforms=self.val_transforms)\n\n    def _dataset(self, df: pd.DataFrame, transforms: Callable) -> HuBMAPDataset:\n        return HuBMAPDataset(df, self.hparams.preprocessed_data_dir, transforms)\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 _dataloader(self, dataset: HuBMAPDataset, 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-12T15:01:03.634449Z","iopub.execute_input":"2022-07-12T15:01:03.634784Z","iopub.status.idle":"2022-07-12T15:01:03.649355Z","shell.execute_reply.started":"2022-07-12T15:01:03.634755Z","shell.execute_reply":"2022-07-12T15:01:03.648212Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize Train Batch","metadata":{}},{"cell_type":"code","source":"nrows = 2\nncols = 2\nbatch_size = nrows * ncols\n\ndata_module = LitDataModule(train_df, PREPROCESSED_DATA_DIR, VAL_FOLD, batch_size, NUM_WORKERS)\ndata_module.setup()\ndata_loader = data_module.train_dataloader()\n\nimages, masks = next(iter(data_loader))\n\nfig, _ = plt.subplots(figsize=(10, 10))\nfor i, (image, mask) in enumerate(zip(images, masks)):\n    plt.subplot(nrows, ncols, i + 1)\n    plt.tight_layout()\n    plt.axis('off')\n    \n    image = image.permute(1, 2, 0).numpy()\n    mask = mask.numpy()\n    \n    print(image.shape, image.min(), image.max(), image.mean(), image.std())\n    print(mask.shape, mask.min(), mask.max(), mask.mean(), mask.std())\n    \n    plt.imshow(image)\n    plt.imshow(mask, alpha=0.2)","metadata":{"execution":{"iopub.status.busy":"2022-07-12T15:01:04.388085Z","iopub.execute_input":"2022-07-12T15:01:04.389319Z","iopub.status.idle":"2022-07-12T15:01:05.614134Z","shell.execute_reply.started":"2022-07-12T15:01:04.389247Z","shell.execute_reply":"2022-07-12T15:01:05.612895Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Lightning Module","metadata":{}},{"cell_type":"code","source":"class LitModule(pl.LightningModule):\n    LOSS_FNS = {\n        \"bce\": smp.losses.SoftBCEWithLogitsLoss(),\n        \"dice\": smp.losses.DiceLoss(mode=\"binary\"),\n        \"focal\": smp.losses.FocalLoss(mode=\"binary\"),\n        \"jaccard\": smp.losses.JaccardLoss(mode=\"binary\"),\n        \"lovasz\": smp.losses.LovaszLoss(mode=\"binary\"),\n        \"tversky\": smp.losses.TverskyLoss(mode=\"binary\"),\n    }\n\n    def __init__(\n        self,\n        arch: str,\n        encoder_name: str,\n        encoder_weights: str,\n        loss: str,\n        optimizer: str,\n        learning_rate: float,\n        weight_decay: float,\n        scheduler: str,\n        T_max: int,\n        T_0: int,\n        min_lr: int,\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        return smp.create_model(\n            self.hparams.arch,\n            encoder_name=self.hparams.encoder_name,\n            encoder_weights=self.hparams.encoder_weights,\n            classes=1,\n            activation=\"sigmoid\",\n        )\n\n    def _init_loss_fn(self) -> Callable:\n        losses = self.hparams.loss.split(\"_\")\n        loss_fns = [self.LOSS_FNS[loss] for loss in losses]\n\n        def criterion(y_pred, y_true):\n            return sum(loss_fn(y_pred, y_true) for loss_fn in loss_fns) / len(loss_fns)\n\n        return criterion\n\n    def _init_metrics(self) -> nn.ModuleDict:\n        train_metrics = MetricCollection({\"train_dice\": Dice()})\n        val_metrics = MetricCollection({\"val_dice\": Dice()})\n\n        return nn.ModuleDict(\n            {\n                \"train_metrics\": train_metrics,\n                \"val_metrics\": val_metrics,\n            }\n        )\n\n    def configure_optimizers(self) -> Dict[str, Any]:\n        optimizer_kwargs = dict(\n            params=self.parameters(), lr=self.hparams.learning_rate, weight_decay=self.hparams.weight_decay\n        )\n        if self.hparams.optimizer == \"Adam\":\n            optimizer = torch.optim.Adam(**optimizer_kwargs)\n        elif self.hparams.optimizer == \"AdamW\":\n            optimizer = torch.optim.AdamW(**optimizer_kwargs)\n        elif self.hparams.optimizer == \"SGD\":\n            optimizer = torch.optim.SGD(**optimizer_kwargs)\n        else:\n            raise ValueError(f\"Unknown optimizer: {self.hparams.optimizer}\")\n\n        if self.hparams.scheduler is not None:\n            if self.hparams.scheduler == \"CosineAnnealingLR\":\n                scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n                    optimizer, T_max=self.hparams.T_max, eta_min=self.hparams.min_lr\n                )\n            elif self.hparams.scheduler == \"CosineAnnealingWarmRestarts\":\n                scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(\n                    optimizer, T_0=self.hparams.T_0, eta_min=self.hparams.min_lr\n                )\n            else:\n                raise ValueError(f\"Unknown scheduler: {self.hparams.scheduler}\")\n\n            return {\"optimizer\": optimizer, \"lr_scheduler\": {\"scheduler\": scheduler, \"interval\": \"step\"}}\n        else:\n            return {\"optimizer\": optimizer}\n\n    def forward(self, images: torch.Tensor) -> torch.Tensor:\n        return self.model(images)\n\n    def training_step(self, batch: Tuple[torch.Tensor, torch.Tensor], batch_idx: int) -> torch.Tensor:\n        return self.shared_step(batch, \"train\")\n\n    def validation_step(self, batch: Tuple[torch.Tensor, torch.Tensor], batch_idx: int):\n        self.shared_step(batch, \"val\")\n\n    def shared_step(self, batch: Tuple[torch.Tensor, torch.Tensor], stage: str) -> torch.Tensor:\n        images, masks = batch\n        y_pred = self(images)\n\n        loss = self.loss_fn(y_pred, masks)\n        metrics = self.metrics[f\"{stage}_metrics\"](y_pred, masks)\n\n        self._log(loss, metrics, stage)\n\n        return loss\n\n    def _log(self, loss: torch.Tensor, metrics: dict, stage: str):\n        on_step = True if stage == \"train\" else False\n        self.log(f\"{stage}_loss\", loss, on_step=on_step, on_epoch=True, prog_bar=not on_step)\n        self.log_dict(metrics, on_step=False, on_epoch=True)\n\n    @classmethod\n    def load_eval_checkpoint(cls, checkpoint_path: Path, 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-12T15:01:05.616645Z","iopub.execute_input":"2022-07-12T15:01:05.617001Z","iopub.status.idle":"2022-07-12T15:01:05.653556Z","shell.execute_reply.started":"2022-07-12T15:01:05.616966Z","shell.execute_reply":"2022-07-12T15:01:05.652291Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"def train(\n    train_df,\n    random_seed: int = RANDOM_SEED,\n    preprocessed_data_dir: str = str(PREPROCESSED_DATA_DIR),\n    val_fold: int = VAL_FOLD,\n    batch_size: int = BATCH_SIZE,\n    num_workers: int = NUM_WORKERS,\n    arch: str = ARCH,\n    encoder_name: str = ENCODER_NAME,\n    encoder_weights: str = ENCODER_WEIGHTS,\n    loss: str = LOSS,\n    optimizer: str = OPTIMIZER,\n    learning_rate: float = LEARNING_RATE,\n    weight_decay: float = WEIGHT_DECAY,\n    scheduler: str = SCHEDULER,\n    min_lr: float = MIN_LR,\n    gpus: int = GPUS,\n    fast_dev_run: bool = FAST_DEV_RUN,\n    max_epochs: int = MAX_EPOCHS,\n    precision: int = PRECISION,\n):\n    pl.seed_everything(random_seed)\n\n    data_module = LitDataModule(\n        train_df=train_df,\n        preprocessed_data_dir=preprocessed_data_dir,\n        val_fold=val_fold,\n        batch_size=batch_size,\n        num_workers=num_workers,\n    )\n\n    module = LitModule(\n        arch=arch,\n        encoder_name=encoder_name,\n        encoder_weights=encoder_weights,\n        loss=loss,\n        optimizer=optimizer,\n        learning_rate=learning_rate,\n        weight_decay=weight_decay,\n        scheduler=scheduler,\n        T_max=int(30_000 / batch_size * max_epochs) + 50,\n        T_0=25,\n        min_lr=min_lr,\n    )\n    \n    trainer = pl.Trainer(\n        fast_dev_run=fast_dev_run,\n        gpus=gpus,\n        #logger=pl.loggers.CSVLogger(save_dir='logs/'),\n        logger=None,\n        log_every_n_steps=10,\n        max_epochs=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-12T15:01:54.771695Z","iopub.execute_input":"2022-07-12T15:01:54.772114Z","iopub.status.idle":"2022-07-12T15:01:54.784996Z","shell.execute_reply.started":"2022-07-12T15:01:54.772068Z","shell.execute_reply":"2022-07-12T15:01:54.783862Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainer = train(train_df)","metadata":{"execution":{"iopub.status.busy":"2022-07-12T15:01:55.302970Z","iopub.execute_input":"2022-07-12T15:01:55.303562Z","iopub.status.idle":"2022-07-12T15:02:08.006751Z","shell.execute_reply.started":"2022-07-12T15:01:55.303525Z","shell.execute_reply":"2022-07-12T15:02:08.005441Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}