{"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":"# RBCD PyTorch⚡MONAI Train & Infer\n\n## Combine powers of [PyTorch Lightning](https://pytorch-lightning.readthedocs.io/en/stable/) and [MONAI](https://docs.monai.io/en/stable/)\n\n## Sources\n- [Theo Viel](https://www.kaggle.com/theoviel)'s [RSNA Breast Cancer Detection - 512x512 pngs](https://www.kaggle.com/datasets/theoviel/rsna-breast-cancer-512-pngs) dataset\n- [Theo Viel](https://www.kaggle.com/theoviel)'s [Dicom conversion](https://www.kaggle.com/code/theoviel/rsna-breast-baseline-inference)\n- My [RBCD Downloads](https://www.kaggle.com/code/clemchris/rbcd-downloads) notebook for package installations","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"markdown","source":"# Installs","metadata":{}},{"cell_type":"code","source":"!pip install monai dicomsdl pytorch_lightning --no-index --find-links=../input/rbcd-downloads","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-12-19T10:10:01.349237Z","iopub.execute_input":"2022-12-19T10:10:01.349673Z","iopub.status.idle":"2022-12-19T10:10:13.215806Z","shell.execute_reply.started":"2022-12-19T10:10:01.34958Z","shell.execute_reply":"2022-12-19T10:10:13.214609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints\n!cp ../input/rbcd-downloads/efficientnet-b0-355c32eb.pth /root/.cache/torch/hub/checkpoints\n!cp ../input/rbcd-downloads/efficientnet-b4-6ed6700e.pth /root/.cache/torch/hub/checkpoints","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:13.21938Z","iopub.execute_input":"2022-12-19T10:10:13.220143Z","iopub.status.idle":"2022-12-19T10:10:19.15721Z","shell.execute_reply.started":"2022-12-19T10:10:13.220082Z","shell.execute_reply":"2022-12-19T10:10:19.155805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!ls -al /root/.cache/torch/hub/checkpoints","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:19.15936Z","iopub.execute_input":"2022-12-19T10:10:19.15979Z","iopub.status.idle":"2022-12-19T10:10:20.200731Z","shell.execute_reply.started":"2022-12-19T10:10:19.159743Z","shell.execute_reply":"2022-12-19T10:10:20.199468Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import multiprocessing as mp\nfrom pathlib import Path\n\nimport cv2\nfrom matplotlib import pyplot as plt\nimport monai\nimport numpy as np\nimport pandas as pd\nimport dicomsdl\nimport pytorch_lightning as pl\nimport seaborn as sns\nimport torch\nimport torchvision\nfrom joblib import delayed\nfrom joblib import Parallel\nfrom monai.data import CSVDataset\nfrom monai.data import DataLoader\nfrom sklearn.model_selection import StratifiedGroupKFold\nimport torch.nn.functional as F\nfrom torchmetrics import MetricCollection\nfrom torchmetrics.classification import BinaryF1Score\nfrom tqdm.notebook import tqdm","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:20.204074Z","iopub.execute_input":"2022-12-19T10:10:20.205275Z","iopub.status.idle":"2022-12-19T10:10:27.637139Z","shell.execute_reply.started":"2022-12-19T10:10:20.205228Z","shell.execute_reply":"2022-12-19T10:10:27.635977Z"},"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\nDATA_ROOT_DIR = INPUT_DIR / \"rsna-breast-cancer-detection\"\nCHECKPOINTS_DIR = INPUT_DIR / \"rbcd-checkpoints\"\n\nTRAIN_IMAGES_DIR = INPUT_DIR / \"rsna-breast-cancer-512-pngs\"\nTEST_IMAGES_DIR = DATA_ROOT_DIR / \"test_images\"\n\nTRAIN_CSV_PATH = DATA_ROOT_DIR / \"train.csv\"\nTEST_CSV_PATH = DATA_ROOT_DIR / \"test.csv\"\n\nOUTPUT_TEST_IMAGES_DIR = OUTPUT_DIR / \"test_images\"\nOUTPUT_TEST_IMAGES_DIR.mkdir(exist_ok=True)\n\nACCELERATOR = \"gpu\"\nBATCH_SIZE = 64\nDEVICES = 1\nETA_MIN = 1e-6\nFAST_DEV_RUN = False\nINTENSITY_TRANSFORM = \"normalize\"\nLEARNING_RATE = 1e-4\nLOSS = \"sigmoid_focal_loss\"\nMAX_EPOCHS = 5\nMODEL_NAME = \"efficientnet-b0\"\nNUM_SPLITS = 4\nNUM_WORKERS = mp.cpu_count()\nPRECISION = 16\nSEED = 42\nSPATIAL_SIZE = 512\nUPSAMPLE = 10\nVAL_FOLD = 0.0\nWEIGHT_DECAY = 1e-6\n\nprint(f\"NUM_WORKERS={NUM_WORKERS}\")","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:27.638669Z","iopub.execute_input":"2022-12-19T10:10:27.641839Z","iopub.status.idle":"2022-12-19T10:10:27.653186Z","shell.execute_reply.started":"2022-12-19T10:10:27.641797Z","shell.execute_reply":"2022-12-19T10:10:27.651424Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Convert DCM to PNG","metadata":{}},{"cell_type":"code","source":"def convert_dcm_to_png(image_path, size, output_image_dir):\n    patient_id = image_path.parent.name\n    image_id = image_path.stem\n\n    dicom = dicomsdl.open(str(image_path))\n    img = dicom.pixelData()\n\n    img = (img - img.min()) / (img.max() - img.min())\n\n    if dicom.getPixelDataInfo()[\"PhotometricInterpretation\"] == \"MONOCHROME1\":\n        img = 1 - img\n\n    img = cv2.resize(img, (size, size))\n\n    output_image_path = output_image_dir / f\"{patient_id}_{image_id}.png\"\n    cv2.imwrite(str(output_image_path), (img * 255).astype(np.uint8))","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:27.654575Z","iopub.execute_input":"2022-12-19T10:10:27.655155Z","iopub.status.idle":"2022-12-19T10:10:27.665734Z","shell.execute_reply.started":"2022-12-19T10:10:27.655111Z","shell.execute_reply":"2022-12-19T10:10:27.66461Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_image_paths = sorted(TEST_IMAGES_DIR.glob(\"*/*.dcm\"))\n\n_ = Parallel(n_jobs=NUM_WORKERS)(\n    delayed(convert_dcm_to_png)(test_image_path, size=SPATIAL_SIZE, output_image_dir=OUTPUT_TEST_IMAGES_DIR)\n    for test_image_path in tqdm(test_image_paths)\n)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:27.667415Z","iopub.execute_input":"2022-12-19T10:10:27.667815Z","iopub.status.idle":"2022-12-19T10:10:30.143207Z","shell.execute_reply.started":"2022-12-19T10:10:27.667776Z","shell.execute_reply":"2022-12-19T10:10:30.14193Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare DataFrames","metadata":{}},{"cell_type":"code","source":"def prepare_data(csv_path, images_dir, create_splits: bool = False):\n    df = pd.read_csv(csv_path)\n\n    df[\"image\"] = (\n        str(images_dir)\n        + \"/\"\n        + df[\"patient_id\"].astype(str)\n        + \"_\"\n        + df[\"image_id\"].astype(str)\n        + \".png\"\n    )\n    \n    if create_splits:\n        skf = StratifiedGroupKFold(n_splits=NUM_SPLITS)\n        for fold, (_, val_) in enumerate(\n            skf.split(X=df, y=df.cancer, groups=df.patient_id)\n        ):\n            df.loc[val_, \"fold\"] = fold\n            \n    # Save\n    file_path = csv_path.name\n    df.to_csv(file_path, index=False)\n    \n    print(f\"Created {file_path} with {len(df)} rows\")\n    \n    return df","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:30.145016Z","iopub.execute_input":"2022-12-19T10:10:30.145717Z","iopub.status.idle":"2022-12-19T10:10:30.154379Z","shell.execute_reply.started":"2022-12-19T10:10:30.145679Z","shell.execute_reply":"2022-12-19T10:10:30.153292Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_df = prepare_data(TRAIN_CSV_PATH, TRAIN_IMAGES_DIR, create_splits=True)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:30.156371Z","iopub.execute_input":"2022-12-19T10:10:30.157262Z","iopub.status.idle":"2022-12-19T10:10:34.234044Z","shell.execute_reply.started":"2022-12-19T10:10:30.157175Z","shell.execute_reply":"2022-12-19T10:10:34.232993Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_df = prepare_data(TEST_CSV_PATH, OUTPUT_TEST_IMAGES_DIR)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:34.241514Z","iopub.execute_input":"2022-12-19T10:10:34.243663Z","iopub.status.idle":"2022-12-19T10:10:34.262062Z","shell.execute_reply.started":"2022-12-19T10:10:34.243624Z","shell.execute_reply":"2022-12-19T10:10:34.260968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# LightningDataModule","metadata":{}},{"cell_type":"code","source":"class RBCDDataModule(pl.LightningDataModule):\n    INTENSITY_TRANSFORMS = {\n        \"normalize\": monai.transforms.NormalizeIntensityd,\n        \"scale\": monai.transforms.ScaleIntensityd,\n    }\n    \n    def __init__(\n        self,\n        batch_size: int,\n        intensity_transform: str,\n        data_csv_path: str,\n        num_workers: int,\n        spatial_size: int,\n        upsample: int,\n        val_fold: float,\n    ):\n        super().__init__()\n\n        self.save_hyperparameters()\n\n        self.df = pd.read_csv(data_csv_path)\n\n        self.keys = (\"image\",)\n        self.spatial_size = (spatial_size, spatial_size)\n        self.train_transform = self._init_train_transform()\n        self.val_transform = self._init_val_transform()\n\n    def _init_train_transform(self):\n        mode = \"bilinear\"\n\n        return monai.transforms.Compose(\n            [\n                monai.transforms.LoadImaged(\n                    keys=self.keys,\n                    ensure_channel_first=True,\n                ),\n                self.INTENSITY_TRANSFORMS[self.hparams.intensity_transform](keys=self.keys),\n                monai.transforms.ResizeWithPadOrCropd(\n                    keys=self.keys,\n                    spatial_size=self.spatial_size,\n                ),\n                monai.transforms.RandRotated(\n                    keys=self.keys,\n                    range_x=0.3,\n                    range_y=0.3,\n                    mode=mode,\n                    prob=0.2,\n                ),\n                monai.transforms.RandZoomd(\n                    keys=self.keys,\n                    min_zoom=[0.8, 0.8],\n                    max_zoom=[1.2, 1.2],\n                    mode=mode,\n                    prob=0.16,\n                ),\n                monai.transforms.RandGaussianSmoothd(\n                    keys=self.keys,\n                    sigma_x=[0.5, 1.15],\n                    sigma_y=[0.5, 1.15],\n                    prob=0.2,\n                ),\n                monai.transforms.RandScaleIntensityd(\n                    keys=self.keys,\n                    factors=0.3,\n                    prob=0.5,\n                ),\n                monai.transforms.RandShiftIntensityd(\n                    keys=self.keys,\n                    offsets=0.1,\n                    prob=0.5,\n                ),\n                monai.transforms.RandGaussianNoised(\n                    keys=self.keys,\n                    mean=0.0,\n                    std=0.1,\n                    prob=0.2,\n                ),\n                monai.transforms.RandFlipd(\n                    keys=self.keys,\n                    prob=0.5,\n                    spatial_axis=0,\n                ),\n                monai.transforms.RandFlipd(\n                    keys=self.keys,\n                    prob=0.5,\n                    spatial_axis=1,\n                ),\n            ]\n        )\n\n    def _init_val_transform(self):\n        return monai.transforms.Compose(\n            [\n                monai.transforms.LoadImaged(\n                    keys=self.keys,\n                    ensure_channel_first=True,\n                ),\n                self.INTENSITY_TRANSFORMS[self.hparams.intensity_transform](keys=self.keys),\n                monai.transforms.ResizeWithPadOrCropd(\n                    keys=self.keys,\n                    spatial_size=self.spatial_size,\n                ),\n            ]\n        )\n\n    def setup(self, stage=None):\n        if self.hparams.data_csv_path == \"train.csv\":\n            train_df = self.df[self.df.fold != self.hparams.val_fold].reset_index(drop=True)\n            val_df = self.df[self.df.fold == self.hparams.val_fold].reset_index(drop=True)\n        \n        if stage == \"fit\" or stage is None:\n            # Upsample cancer data (from https://www.kaggle.com/code/awsaf49/rsna-bcd-efficientnet-tf-tpu-1vm-train?scriptVersionId=113444994&cellId=67)\n            pos_df = train_df[train_df.cancer == 1].sample(frac=self.hparams.upsample, replace=True)\n            neg_df = train_df[train_df.cancer == 0]\n            train_df = pd.concat([pos_df, neg_df], axis=0, ignore_index=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        if stage == \"predict\" or stage is None:\n            if self.hparams.data_csv_path == \"train.csv\":\n                # Use val_df to optimize threshold\n                self.predict_dataset = self._dataset(val_df, self.val_transform)\n            else:\n                self.predict_dataset = self._dataset(self.df, self.val_transform, col_names=[\"image\"])\n\n\n    def _dataset(self, df, transform, col_names=[\"image\", \"fold\", \"cancer\"]):\n        return CSVDataset(src=df, transform=transform, col_names=col_names)\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        )","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:34.266552Z","iopub.execute_input":"2022-12-19T10:10:34.268754Z","iopub.status.idle":"2022-12-19T10:10:34.301033Z","shell.execute_reply.started":"2022-12-19T10:10:34.268715Z","shell.execute_reply":"2022-12-19T10:10:34.300141Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Visualize Data","metadata":{}},{"cell_type":"code","source":"num_patiens = 1\nnum_images_per_patient = 4\n    \ndata_module = RBCDDataModule(\n    batch_size=num_patiens * num_images_per_patient,\n    data_csv_path=\"train.csv\",\n    intensity_transform=INTENSITY_TRANSFORM,\n    num_workers=0,\n    spatial_size=SPATIAL_SIZE,\n    upsample=UPSAMPLE,\n    val_fold=0.0,\n)\ndata_module.setup()\n\ndataloaders = {\n    \"train\": data_module.train_dataloader(),\n    \"val\": data_module.val_dataloader(),\n    \"predict\": data_module.predict_dataloader(),\n}","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:34.305521Z","iopub.execute_input":"2022-12-19T10:10:34.307761Z","iopub.status.idle":"2022-12-19T10:10:34.814929Z","shell.execute_reply.started":"2022-12-19T10:10:34.307724Z","shell.execute_reply":"2022-12-19T10:10:34.813937Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"nrows = num_patiens\nncols = num_images_per_patient\nfor stage, dataloader in dataloaders.items():\n    batch = next(iter(dataloader))\n\n    fig, _ = plt.subplots(figsize=(ncols * 4, nrows * 4))\n    plt.suptitle(f\"{stage.title()} data\")\n    for idx, image in enumerate(batch[\"image\"]):\n        plt.subplot(nrows, ncols, idx + 1)\n\n        title = f\"{image.shape}, {image.min().item():.2f}, {image.max().item():.2f}\"\n\n        plt.title(title)\n        plt.imshow(image[0], cmap=\"gray\")\n\n        plt.axis(\"off\")","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:34.81629Z","iopub.execute_input":"2022-12-19T10:10:34.817441Z","iopub.status.idle":"2022-12-19T10:10:36.880034Z","shell.execute_reply.started":"2022-12-19T10:10:34.817386Z","shell.execute_reply":"2022-12-19T10:10:36.879152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# LightningModule","metadata":{}},{"cell_type":"code","source":"class RBCDModule(pl.LightningModule):\n    LOSS_FNS = {\n        \"binary_cross_entropy_with_logits\": F.binary_cross_entropy_with_logits,\n        \"sigmoid_focal_loss\": torchvision.ops.sigmoid_focal_loss,\n    }\n    \n    def __init__(\n        self,\n        eta_min: float,\n        learning_rate: float,\n        loss: str, \n        model_name: str,\n        T_max: 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_fn = self._init_loss_fn()\n\n        self.metrics = self._init_metrics()\n\n    def _init_model(self):\n        model_kwargs = {\n            \"model_name\": self.hparams.model_name,\n            \"in_channels\": 1,\n            \"num_classes\": 1,\n            \"spatial_dims\": 2,\n        }\n\n        return monai.networks.nets.EfficientNetBN(**model_kwargs)\n\n    def _init_loss_fn(self):\n        return self.LOSS_FNS[self.hparams.loss]\n\n    def _init_metrics(self):\n        metrics = {\n            \"f1\": BinaryF1Score(),\n        }\n        metric_collection = MetricCollection(metrics)\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 = torch.optim.Adam(\n            params=self.parameters(), lr=self.hparams.learning_rate, weight_decay=self.hparams.weight_decay\n        )\n\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n            optimizer, T_max=self.hparams.T_max, eta_min=self.hparams.eta_min\n        )\n\n        return [optimizer], [scheduler]\n\n    def forward(self, images):\n        return self.model(images)\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        try:\n            images, labels, logits = self._forward_pass(batch)\n            preds = logits.sigmoid()\n            return preds, labels\n        except:\n            images = batch[\"image\"].as_tensor()\n            logits = self(images).view(-1)\n            preds = logits.sigmoid()\n            return preds\n\n    def _shared_step(self, batch, stage):\n        images, labels, logits = self._forward_pass(batch)\n        \n        if self.hparams.loss == \"sigmoid_focal_loss\":\n            loss = self.loss_fn(logits, labels, alpha=0.80, gamma=2.0, reduction=\"mean\")\n        else:\n            loss = self.loss_fn(logits, labels)\n\n        self.metrics[f\"{stage}_metrics\"](logits, labels)\n\n        self._log(loss, stage, batch_size=len(images))\n\n        return loss\n\n    def _forward_pass(self, batch):\n        images, labels = batch[\"image\"].as_tensor(), batch[\"cancer\"].float()\n        logits = self(images).view(-1)\n        return images, labels, logits\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":"2022-12-19T10:10:36.881854Z","iopub.execute_input":"2022-12-19T10:10:36.882581Z","iopub.status.idle":"2022-12-19T10:10:36.902833Z","shell.execute_reply.started":"2022-12-19T10:10:36.88254Z","shell.execute_reply":"2022-12-19T10:10:36.901762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"monai.utils.set_determinism(SEED)\npl.seed_everything(SEED, workers=True)\n\ndata_module = RBCDDataModule(\n    batch_size=BATCH_SIZE,\n    data_csv_path=\"train.csv\",\n    intensity_transform=INTENSITY_TRANSFORM,\n    num_workers=NUM_WORKERS,\n    spatial_size=SPATIAL_SIZE,\n    upsample=UPSAMPLE,\n    val_fold=VAL_FOLD,\n)\n\nmodule = RBCDModule(\n    eta_min=ETA_MIN,\n    learning_rate=LEARNING_RATE,\n    loss=LOSS,\n    model_name=MODEL_NAME,\n    T_max=MAX_EPOCHS,\n    weight_decay=WEIGHT_DECAY,\n)\n\ntrainer = pl.Trainer(\n    accelerator=ACCELERATOR,\n    benchmark=True,\n    devices=1,\n    fast_dev_run=FAST_DEV_RUN,\n    logger=pl.loggers.CSVLogger(save_dir='logs/'),\n    log_every_n_steps=5,\n    max_epochs=MAX_EPOCHS,\n    precision=PRECISION,\n)\n\ntrainer.fit(module, datamodule=data_module)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:10:36.908682Z","iopub.execute_input":"2022-12-19T10:10:36.909071Z","iopub.status.idle":"2022-12-19T10:11:52.098509Z","shell.execute_reply.started":"2022-12-19T10:10:36.909029Z","shell.execute_reply":"2022-12-19T10:11:52.097097Z"},"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\", \"train_f1\", \"val_f1\"]]\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-12-19T10:12:05.646734Z","iopub.execute_input":"2022-12-19T10:12:05.647652Z","iopub.status.idle":"2022-12-19T10:12:06.078976Z","shell.execute_reply.started":"2022-12-19T10:12:05.647603Z","shell.execute_reply":"2022-12-19T10:12:06.078043Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"checkpoints_dir = Path(trainer.logger.log_dir) / \"checkpoints\"\ncheckpoint_path = next(checkpoints_dir.iterdir())\ncheckpoint_path","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:12:18.569871Z","iopub.execute_input":"2022-12-19T10:12:18.57027Z","iopub.status.idle":"2022-12-19T10:12:18.578151Z","shell.execute_reply.started":"2022-12-19T10:12:18.570234Z","shell.execute_reply":"2022-12-19T10:12:18.577206Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimize Threshold","metadata":{}},{"cell_type":"markdown","source":"## Helpers","metadata":{}},{"cell_type":"code","source":"# https://www.kaggle.com/code/awsaf49/metric-probabilistic-fscore-tf-torch-numpy?scriptVersionId=113097239&cellId=19 # noqa: E501\ndef pfbeta_torch(preds, labels, beta=1):\n    preds = preds.clip(0, 1)\n\n    y_true_count = labels.sum()\n    ctp = preds[labels == 1].sum()\n    cfp = preds[labels == 0].sum()\n\n    beta_squared = beta * beta\n\n    c_precision = ctp / (ctp + cfp)\n    c_recall = ctp / y_true_count\n\n    if c_precision > 0 and c_recall > 0:\n        return ((1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall)).item()\n    else:\n        return 0.0\n\n\ndef pfbeta_thresh(preds, labels):\n    optimized_preds = optimize_preds(preds, labels)\n    return pfbeta_torch(optimized_preds, labels)\n\n\ndef optimize_preds(preds, labels, return_thresh=False, print_results=False):\n    preds = preds.clone()\n\n    without_thresh = pfbeta_torch(preds, labels)\n\n    threshs = np.linspace(0, 1, 101)\n    f1s = [pfbeta_torch((preds > thr).float(), labels) for thr in threshs]\n    idx = np.argmax(f1s)\n    thresh, best_pfbeta = threshs[idx], f1s[idx]\n\n    preds = (preds > thresh).float()\n\n    if print_results:\n        print(f\"without optimization: {without_thresh:.3f}\")\n        pfbeta = pfbeta_torch(preds, labels)\n        print(f\"with optimization: {pfbeta:.3f}\")\n        print(f\"best_thresh: {thresh}\")\n\n    if return_thresh:\n        return thresh\n\n    return preds","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:12:28.335909Z","iopub.execute_input":"2022-12-19T10:12:28.336322Z","iopub.status.idle":"2022-12-19T10:12:28.347111Z","shell.execute_reply.started":"2022-12-19T10:12:28.336288Z","shell.execute_reply":"2022-12-19T10:12:28.346136Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Optimize","metadata":{}},{"cell_type":"code","source":"data_module = RBCDDataModule(\n    batch_size=BATCH_SIZE*2,\n    data_csv_path=\"train.csv\",\n    intensity_transform=INTENSITY_TRANSFORM,\n    num_workers=NUM_WORKERS,\n    spatial_size=SPATIAL_SIZE,\n    upsample=UPSAMPLE,\n    val_fold=VAL_FOLD,\n)\n\nmodule = RBCDModule.load_from_checkpoint(checkpoint_path)\n\ntrainer = pl.Trainer(\n    accelerator=ACCELERATOR,\n    devices=DEVICES,\n    logger=None,\n    precision=16 if ACCELERATOR == \"gpu\" else 32,\n)\n\npredictions = trainer.predict(module, datamodule=data_module)\n\npreds, labels = [], []\nfor pred, label in predictions:\n    preds.append(pred)\n    labels.append(label)\n\npreds = torch.cat(preds)\nlabels = torch.cat(labels)\n\nthreshold = optimize_preds(preds.float(), labels.float(), return_thresh=True, print_results=True)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:12:42.326992Z","iopub.execute_input":"2022-12-19T10:12:42.328471Z","iopub.status.idle":"2022-12-19T10:17:49.382013Z","shell.execute_reply.started":"2022-12-19T10:12:42.328423Z","shell.execute_reply":"2022-12-19T10:17:49.380707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Infer","metadata":{}},{"cell_type":"code","source":"data_module = RBCDDataModule(\n    batch_size=BATCH_SIZE*2,\n    data_csv_path=\"test.csv\",\n    intensity_transform=INTENSITY_TRANSFORM,\n    num_workers=NUM_WORKERS,\n    spatial_size=SPATIAL_SIZE,\n    upsample=UPSAMPLE,\n    val_fold=VAL_FOLD,\n)\n\nmodule = RBCDModule.load_from_checkpoint(checkpoint_path)\n\ntrainer = pl.Trainer(\n    accelerator=ACCELERATOR,\n    devices=DEVICES,\n    logger=None,\n    precision=16 if ACCELERATOR == \"gpu\" else 32,\n)\n\npredictions = trainer.predict(module, datamodule=data_module)\n\npredictions = torch.cat(predictions).numpy()\n    \nprint(predictions.shape)","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:22:09.073812Z","iopub.execute_input":"2022-12-19T10:22:09.074215Z","iopub.status.idle":"2022-12-19T10:22:11.04195Z","shell.execute_reply.started":"2022-12-19T10:22:09.074181Z","shell.execute_reply":"2022-12-19T10:22:11.040681Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Submit","metadata":{}},{"cell_type":"code","source":"predictions = (predictions > threshold).astype(int)\n\ntest_df[\"cancer\"] = predictions\n\nsub_df = test_df[[\"prediction_id\", \"cancer\"]].groupby(\"prediction_id\").mean().reset_index()\n\nsub_df.to_csv(\"submission.csv\",index=False)\n\nprint(sub_df.head())","metadata":{"execution":{"iopub.status.busy":"2022-12-19T10:22:11.044733Z","iopub.execute_input":"2022-12-19T10:22:11.045183Z","iopub.status.idle":"2022-12-19T10:22:11.06135Z","shell.execute_reply.started":"2022-12-19T10:22:11.045141Z","shell.execute_reply":"2022-12-19T10:22:11.060062Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}