{"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⚡Timm Train","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"markdown","source":"# Installs","metadata":{}},{"cell_type":"code","source":"!pip install pytorch_lightning timm --no-index --find-links=../input/rbcd-downloads","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2023-01-17T16:07:23.106753Z","iopub.execute_input":"2023-01-17T16:07:23.107058Z","iopub.status.idle":"2023-01-17T16:07:35.772263Z","shell.execute_reply.started":"2023-01-17T16:07:23.106995Z","shell.execute_reply":"2023-01-17T16:07:35.771091Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints\n!cp ../input/rbcd-downloads/tf_efficientnetv2_s-eb54923e.pth /root/.cache/torch/hub/checkpoints","metadata":{"execution":{"iopub.status.busy":"2023-01-17T16:07:35.774854Z","iopub.execute_input":"2023-01-17T16:07:35.775291Z","iopub.status.idle":"2023-01-17T16:07:39.649698Z","shell.execute_reply.started":"2023-01-17T16:07:35.775229Z","shell.execute_reply":"2023-01-17T16:07:39.648353Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Imports","metadata":{}},{"cell_type":"code","source":"import multiprocessing as mp\nfrom pathlib import Path\nfrom typing import Tuple\n\nfrom matplotlib import pyplot as plt\nimport numpy as np\nimport pandas as pd\nimport pytorch_lightning as pl\nimport seaborn as sns\nimport timm\nimport torch\nimport torch.nn.functional as F\nfrom PIL import Image\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom timm.data.transforms_factory import create_transform\nfrom timm.loss import BinaryCrossEntropy\nfrom timm.optim import create_optimizer_v2\nfrom torch.utils.data import DataLoader\nfrom torch.utils.data import Dataset\nfrom torchmetrics import MetricCollection\nfrom torchmetrics.classification import BinaryF1Score","metadata":{"execution":{"iopub.status.busy":"2023-01-17T16:07:39.651384Z","iopub.execute_input":"2023-01-17T16:07:39.651885Z","iopub.status.idle":"2023-01-17T16:07:44.570998Z","shell.execute_reply.started":"2023-01-17T16:07:39.651831Z","shell.execute_reply":"2023-01-17T16:07:44.569919Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Paths & Settings","metadata":{}},{"cell_type":"code","source":"KAGGLE_DIR = Path(\"/\") / \"kaggle\"\n\nINPUT_DIR = KAGGLE_DIR / \"input\"\n\nDATA_ROOT_DIR = INPUT_DIR / \"rsna-breast-cancer-detection\"\n\nTRAIN_IMAGES_DIR = INPUT_DIR / \"rsna-mammo-dicomsdl-1024\" / \"train_images_processed_cv2_dicomsdl_1024\"\nTRAIN_CSV_PATH = DATA_ROOT_DIR / \"train.csv\"\n\nACCELERATOR = \"gpu\"\nBATCH_SIZE = 8\nDEVICES = 1\nDROP_RATE = 0.5\nDROP_PATH_RATE = 0.4\nETA_MIN = 1e-6\nFAST_DEV_RUN = False\nIMAGE_SIZE = 1024\nLEARNING_RATE = 3e-4\nMAX_EPOCHS = 5\nMODEL_NAME = \"tf_efficientnetv2_s\"\nNUM_SPLITS = 4\nNUM_WORKERS = mp.cpu_count()\nOVERFIT_BATCHES = 0\nOPTIMIZER = \"AdamW\"\nPRECISION = 16\nSEED = 42\nUPSAMPLE = 10\nVAL_FOLD = 0.0\nWEIGHT_DECAY = 1e-6","metadata":{"execution":{"iopub.status.busy":"2023-01-17T16:07:44.574263Z","iopub.execute_input":"2023-01-17T16:07:44.575091Z","iopub.status.idle":"2023-01-17T16:07:44.586197Z","shell.execute_reply.started":"2023-01-17T16:07:44.575047Z","shell.execute_reply":"2023-01-17T16:07:44.582164Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Prepare Data","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":"2023-01-17T16:07:44.587396Z","iopub.execute_input":"2023-01-17T16:07:44.587669Z","iopub.status.idle":"2023-01-17T16:07:44.604312Z","shell.execute_reply.started":"2023-01-17T16:07:44.587644Z","shell.execute_reply":"2023-01-17T16:07:44.603324Z"},"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":"2023-01-17T16:07:44.605853Z","iopub.execute_input":"2023-01-17T16:07:44.60625Z","iopub.status.idle":"2023-01-17T16:07:48.376464Z","shell.execute_reply.started":"2023-01-17T16:07:44.606212Z","shell.execute_reply":"2023-01-17T16:07:48.375211Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class RBCDDataset(Dataset):\n    def __init__(self, df, transform):\n        self.df = df\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n\n        image = Image.open(row.image).convert(\"RGB\")\n\n        if self.transform is not None:\n            image = self.transform(image)\n\n        try:\n            return image, row.cancer\n        except:\n            return image","metadata":{"execution":{"iopub.status.busy":"2023-01-17T16:07:48.378315Z","iopub.execute_input":"2023-01-17T16:07:48.378721Z","iopub.status.idle":"2023-01-17T16:07:48.388464Z","shell.execute_reply.started":"2023-01-17T16:07:48.378667Z","shell.execute_reply":"2023-01-17T16:07:48.387348Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# LightningDataModule","metadata":{}},{"cell_type":"code","source":"class TimmDataModule(pl.LightningDataModule):\n    def __init__(\n        self,\n        batch_size: int,\n        data_csv_path: str,\n        num_workers: 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.spatial_size = (IMAGE_SIZE, IMAGE_SIZE)\n\n        self.train_transform = self._init_train_transform()\n        self.val_transform = self._init_val_transform()\n\n    def _init_train_transform(self):\n        return create_transform(\n            input_size=self.spatial_size,\n            is_training=True,\n            scale=(0.75, 1.33),\n            ratio=(0.08, 1.0),\n            hflip=0.5,\n            vflip=0.5,\n            color_jitter=0.4,\n            interpolation=\"random\",\n        )\n    def _init_val_transform(self):\n        return create_transform(\n            input_size=self.spatial_size,\n            is_training=False,\n            interpolation=\"bilinear\",\n        )\n        \n    def setup(self, stage=None):\n        if self.hparams.data_csv_path == \"train.csv\":\n            val_fold = self.hparams.val_fold\n            train_df = self.df[self.df.fold != val_fold].reset_index(drop=True)\n            val_df = self.df[self.df.fold == val_fold].reset_index(drop=True)\n\n            # Upsample cancer data (from https://www.kaggle.com/code/awsaf49/rsna-bcd-efficientnet-tf-tpu-1vm-train?scriptVersionId=113444994&cellId=67) # noqa: E501\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            self.predict_dataset = self._dataset(val_df, self.val_transform)\n        else:\n            self.predict_dataset = self._dataset(self.df, self.val_transform)\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 _dataset(self, df, transform):\n        return RBCDDataset(df=df, transform=transform)\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":"2023-01-17T16:07:48.390298Z","iopub.execute_input":"2023-01-17T16:07:48.390752Z","iopub.status.idle":"2023-01-17T16:07:48.477813Z","shell.execute_reply.started":"2023-01-17T16:07:48.390645Z","shell.execute_reply":"2023-01-17T16:07:48.476746Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# LightningModule","metadata":{}},{"cell_type":"code","source":"class TimmModule(pl.LightningModule):\n    def __init__(\n        self,\n        drop_rate: float,\n        drop_path_rate: float,\n        eta_min: float,\n        learning_rate: float,\n        max_epochs: int,\n        model_name: str,\n        optimizer: str,\n        steps_per_epoch: int,\n        weight_decay: float,\n    ):\n        super().__init__()\n        \n        self.save_hyperparameters()\n\n        self.model = self._init_model()\n\n        self.metrics = self._init_metrics()\n\n        self._loss_fn = self._init_loss_fn()\n        \n    def _init_model(self):\n        return timm.create_model(\n            self.hparams.model_name,\n            pretrained=True,\n            num_classes=1,\n            drop_rate=self.hparams.drop_rate,\n            drop_path_rate=self.hparams.drop_path_rate,\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 _init_loss_fn(self):\n        return F.binary_cross_entropy_with_logits\n\n    def configure_optimizers(self):\n        optimizer = create_optimizer_v2(\n            self.parameters(),\n            opt=self.hparams.optimizer,\n            lr=self.hparams.learning_rate,\n            weight_decay=self.hparams.weight_decay,\n        )\n\n        scheduler = torch.optim.lr_scheduler.OneCycleLR(\n            optimizer,\n            max_lr=self.hparams.learning_rate,\n            epochs=self.hparams.max_epochs,\n            steps_per_epoch=self.hparams.steps_per_epoch,\n        )\n\n        return {\n            \"optimizer\": optimizer,\n            \"lr_scheduler\": {\"scheduler\": scheduler, \"interval\": \"step\"},\n        }\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        images = batch[0]\n        logits = self(images).view(-1)\n        preds = logits.sigmoid()\n        return preds\n    \n    def predict_step(self, batch, batch_idx):\n        try:\n            images, labels = batch[0], batch[1].float()\n            logits = self(images).view(-1)\n            preds = logits.sigmoid()\n            return preds, labels\n        except:\n            images = batch\n            logits = self(images).view(-1)\n            preds = logits.sigmoid()\n            return preds\n\n    def _shared_step(self, batch, stage):\n        images, labels = batch[0], batch[1].float()\n        logits = self(images).view(-1)\n        \n        loss = self._loss_fn(logits, labels)\n\n        self.metrics[f\"{stage}_metrics\"](logits, labels)\n\n        self._log(stage, loss, batch_size=len(images))\n\n        return loss\n    \n    def _log(self, stage, loss, batch_size):\n        self.log(f\"{stage}_loss\", loss, batch_size=batch_size)\n        self.log_dict(self.metrics[f\"{stage}_metrics\"], batch_size=batch_size)","metadata":{"execution":{"iopub.status.busy":"2023-01-17T16:07:48.479509Z","iopub.execute_input":"2023-01-17T16:07:48.479916Z","iopub.status.idle":"2023-01-17T16:07:48.499643Z","shell.execute_reply.started":"2023-01-17T16:07:48.479882Z","shell.execute_reply":"2023-01-17T16:07:48.498622Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train","metadata":{}},{"cell_type":"code","source":"pl.seed_everything(SEED, workers=True)\n\ndata_module = TimmDataModule(\n    batch_size=BATCH_SIZE,\n    data_csv_path=\"train.csv\",\n    num_workers=NUM_WORKERS,\n    upsample=UPSAMPLE,\n    val_fold=VAL_FOLD,\n)\ndata_module.setup()\n\nmodule = TimmModule(\n    drop_rate=DROP_RATE,\n    drop_path_rate=DROP_PATH_RATE,\n    eta_min=ETA_MIN,\n    learning_rate=LEARNING_RATE,\n    max_epochs=MAX_EPOCHS,\n    model_name=MODEL_NAME,\n    optimizer=OPTIMIZER,\n    steps_per_epoch=len(data_module.train_dataloader()),\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    overfit_batches=OVERFIT_BATCHES,\n    precision=PRECISION,\n)\n\ntrainer.fit(module, datamodule=data_module)","metadata":{"execution":{"iopub.status.busy":"2023-01-17T16:07:48.501041Z","iopub.execute_input":"2023-01-17T16:07:48.501475Z"},"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":{"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\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":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Optimize","metadata":{}},{"cell_type":"code","source":"checkpoints_dir = Path(trainer.logger.log_dir) / \"checkpoints\"\ncheckpoint_path = next(checkpoints_dir.iterdir())\ncheckpoint_path","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_module = TimmDataModule(\n    batch_size=BATCH_SIZE,\n    data_csv_path=\"train.csv\",\n    num_workers=NUM_WORKERS,\n    upsample=UPSAMPLE,\n    val_fold=VAL_FOLD,\n)\n\nmodule = TimmModule.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":{"trusted":true},"execution_count":null,"outputs":[]}]}