{"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"},"kaggle":{"accelerator":"tpu1vmV38","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"sourceType":"competition"},{"sourceId":4696088,"sourceType":"datasetVersion","datasetId":2687741}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# PyTorch Lightning GPU & TPU Trainer with KFolds and F1 Loss + W&B logging ⚡️\n\nTo run the model on TPU, un-comment and run the below cell and change the `gpus=1` argument to `tpu_cores=1` or `tpu_cores=8` in the `Trainer` object at the bottom of the notebook.\n\nYou can also extend the notebook to use Multi-GPU setting by changing the current `gpus=1` to `gpus=2` when using the 2x T4 GPUs.\n    \nI am training a `swin_base_patch4_window7_224` model with `img_size=224` on 1x P100 GPU. As indicated above, you can always perform minimal changes in this training code to extend that to single or mulitple GPUs or TPU cores.\n\n**Feel free to fork and change the models and do some preprocessing, but if you do please leave an upvote :)**","metadata":{}},{"cell_type":"markdown","source":"<center>\n<img src=\"https://img.shields.io/badge/Upvote-If%20you%20like%20my%20work-07b3c8?style=for-the-badge&logo=kaggle\">\n</center>","metadata":{}},{"cell_type":"markdown","source":"Uncomment this cell to run the notebook with TPUs 👇🏻","metadata":{}},{"cell_type":"code","source":"# %%capture\n# ! curl https://raw.githubusercontent.com/pytorch/xla/master/contrib/scripts/env-setup.py -o pytorch-xla-env-setup.py\n# ! python pytorch-xla-env-setup.py --version 1.7 --apt-packages libomp5 libopenblas-dev\n# ! python pytorch-xla-env-setup.py --apt-packages libomp5 libopenblas-dev\n# !pip install cloud-tpu-client==0.10 https://storage.googleapis.com/tpu-pytorch/wheels/torch_xla-1.7-cp36-cp36m-linux_x86_64.whl\n# !pip install cloud-tpu-client==0.10 https://storage.googleapis.com/tpu-pytorch/wheels/torch_xla-1.9-cp37-cp37m-linux_x86_64.whl","metadata":{"execution":{"iopub.status.busy":"2022-12-06T07:31:21.367921Z","iopub.execute_input":"2022-12-06T07:31:21.368806Z","iopub.status.idle":"2022-12-06T07:31:49.516417Z","shell.execute_reply.started":"2022-12-06T07:31:21.368717Z","shell.execute_reply":"2022-12-06T07:31:49.515395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Installation and Imports","metadata":{}},{"cell_type":"code","source":"%%capture\n!pip install timm\n!pip install einops\n!apt-get update && apt-get install -y python3-opencv\n!pip install -q opencv-python\n!pip install albumentations\n!pip install --upgrade wandb\n!pip install --upgrade torchmetrics\n!pip install --upgrade pytorch-lightning","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-08T05:49:35.161222Z","iopub.execute_input":"2022-12-08T05:49:35.16217Z","iopub.status.idle":"2022-12-08T05:50:33.729482Z","shell.execute_reply.started":"2022-12-08T05:49:35.162067Z","shell.execute_reply":"2022-12-08T05:50:33.72815Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch_xla","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport sys\nimport cv2\nimport glob\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n\nimport wandb\nfrom einops import rearrange\n\nimport timm\nimport torch\nimport transformers\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.autograd import Variable\nfrom torch.utils.data import DataLoader, Dataset\n\nimport torchmetrics\nimport pytorch_lightning as pl\nfrom pytorch_lightning.callbacks import ModelCheckpoint\n\nfrom sklearn.model_selection import StratifiedGroupKFold\n\nfrom albumentations import (\n    HorizontalFlip, VerticalFlip, Flip, OneOf, \n    Compose, Normalize,Resize\n)\nfrom albumentations.pytorch import ToTensorV2\n\nimport warnings\nwarnings.simplefilter('ignore')","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2022-12-08T06:14:18.272352Z","iopub.execute_input":"2022-12-08T06:14:18.273072Z","iopub.status.idle":"2022-12-08T06:14:18.281828Z","shell.execute_reply.started":"2022-12-08T06:14:18.273032Z","shell.execute_reply":"2022-12-08T06:14:18.280712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nos.environ.pop('TPU_PROCESS_ADDRESSES')\nos.environ.pop('CLOUD_TPU_TASK_ID')\nos.environ[\"NUMBA_NUM_THREADS\"] = \"1\"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Config File and Wandb","metadata":{}},{"cell_type":"code","source":"Config = {\n    'TRAIN_BS': 16,\n    'VALID_BS': 16,\n    'IMG_SIZE': (224, 224),\n    'MODEL_NAME': 'swin_base_patch4_window7_224',\n    'NUM_WORKERS': 8,\n    'PARENT_PATH': '/kaggle/input/rsna-mammography-images-as-pngs/images_as_pngs_cv2_384/train_images_processed_cv2_384/',\n    'FILE_PATH': '/kaggle/input/rsna-breast-cancer-detection/train.csv',\n    'LOSS': 'BCEWithLogitsLoss',\n    'EVAL_METRIC': 'F1',\n    'NB_EPOCHS': 2,\n    'SPLITS': 5,\n    'min_lr': 1e-6,\n    'T_max': 20,\n    'T_0': 25,\n    'fc_dropout': 0.2,\n    'betas': (0.9, 0.999),\n    'NUM_LABELS': 1,\n    'LR': 2e-4,\n    'competition': 'rsna_mammography',\n    '_wandb_kernel': 'tanaym',\n}","metadata":{"execution":{"iopub.status.busy":"2022-12-08T06:14:18.67538Z","iopub.execute_input":"2022-12-08T06:14:18.676354Z","iopub.status.idle":"2022-12-08T06:14:18.683495Z","shell.execute_reply.started":"2022-12-08T06:14:18.676304Z","shell.execute_reply":"2022-12-08T06:14:18.682289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### About W&B:\n<center><img src=\"https://i.imgur.com/gb6B4ig.png\" width=\"400\" alt=\"Weights & Biases\"/></center><br>\n<p style=\"text-align:center\">WandB is a developer tool for companies turn deep learning research projects into deployed software by helping teams track their models, visualize model performance and easily automate training and improving models.\nWe will use their tools to log hyperparameters and output metrics from your runs, then visualize and compare results and quickly share findings with your colleagues.<br><br></p>","metadata":{}},{"cell_type":"markdown","source":"To login to W&B, you can use below snippet.\n\n```python\nfrom kaggle_secrets import UserSecretsClient\nuser_secrets = UserSecretsClient()\nwb_key = user_secrets.get_secret(\"WANDB_API_KEY\")\n\nwandb.login(key=wb_key)\n```\nMake sure you have your W&B key stored as `WANDB_API_KEY` under Add-ons -> Secrets\n\nYou can view [this](https://www.kaggle.com/ayuraj/experiment-tracking-with-weights-and-biases) notebook to learn more about W&B tracking.\n\nIf you don't want to login to W&B, the kernel will still work and log everything to W&B in anonymous mode.","metadata":{}},{"cell_type":"code","source":"# Start W&B logging\n# W&B Login\nfrom kaggle_secrets import UserSecretsClient\n# user_secrets = UserSecretsClient()\n# wb_key = user_secrets.get_secret(\"WANDB_API_KEY\")\n\n# wandb.login(key=wb_key)\n\n# run = wandb.init(\n#     project='pytorch_lightning',\n#     config=Config,\n#     group='vision',\n#     job_type='train',\n# )","metadata":{"execution":{"iopub.status.busy":"2022-12-08T06:14:19.243341Z","iopub.execute_input":"2022-12-08T06:14:19.244475Z","iopub.status.idle":"2022-12-08T06:14:27.060202Z","shell.execute_reply.started":"2022-12-08T06:14:19.244401Z","shell.execute_reply":"2022-12-08T06:14:27.059246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Focal Loss\nclass FocalLoss(nn.Module):\n    def __init__(self, gamma=0, alpha=None, size_average=True):\n        \"\"\"\n        Focal Loss function class taken from:\n        https://github.com/clcarwin/focal_loss_pytorch\n        \"\"\"\n        super(FocalLoss, self).__init__()\n        self.gamma = gamma\n        self.alpha = alpha\n        if isinstance(alpha,(float,int)): self.alpha = torch.Tensor([alpha,1-alpha])\n        if isinstance(alpha,list): self.alpha = torch.Tensor(alpha)\n        self.size_average = size_average\n\n    def forward(self, input, target):\n        if input.dim()>2:\n            input = input.view(input.size(0),input.size(1),-1)  # N,C,H,W => N,C,H*W\n            input = input.transpose(1,2)    # N,C,H*W => N,H*W,C\n            input = input.contiguous().view(-1,input.size(2))   # N,H*W,C => N*H*W,C\n        target = target.view(-1,1)\n\n        logpt = F.log_softmax(input)\n        logpt = logpt.gather(1,target)\n        logpt = logpt.view(-1)\n        pt = Variable(logpt.data.exp())\n\n        if self.alpha is not None:\n            if self.alpha.type()!=input.data.type():\n                self.alpha = self.alpha.type_as(input.data)\n            at = self.alpha.gather(0,target.data.view(-1))\n            logpt = logpt * Variable(at)\n\n        loss = -1 * (1-pt)**self.gamma * logpt\n        if self.size_average: return loss.mean()\n        else: return loss.sum()\n        \ndef probabilistic_f1(labels, preds, beta=1):\n    \"\"\"\n    Function taken from Awsaf's notebook:\n    https://www.kaggle.com/code/awsaf49/rsna-bcd-efficientnet-tf-tpu-1vm-train\n    \"\"\"\n    eps = 1e-5\n    preds = preds.clip(0, 1)\n    y_true_count = labels.sum()\n    ctp = preds[labels==1].sum()\n    cfp = preds[labels==0].sum()\n    beta_squared = beta * beta\n    c_precision = ctp / (ctp + cfp + eps)\n    c_recall = ctp / (y_true_count + eps)\n    if (c_precision > 0 and c_recall > 0):\n        result = (1 + beta_squared) * (c_precision * c_recall) / (beta_squared * c_precision + c_recall + eps)\n        return result\n    else:\n        return 0.0\n    \ndef wandb_log(**kwargs):\n    for k, v in kwargs.items():\n        wandb.log({k: v})","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-12-08T06:14:27.062764Z","iopub.execute_input":"2022-12-08T06:14:27.063459Z","iopub.status.idle":"2022-12-08T06:14:27.083656Z","shell.execute_reply.started":"2022-12-08T06:14:27.063402Z","shell.execute_reply":"2022-12-08T06:14:27.082407Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Dataset Class - for loading in data","metadata":{}},{"cell_type":"code","source":"class RSNAData(Dataset):\n    def __init__(self, df, img_folder, augments=None, is_test=False):\n        self.df = df\n        self.is_test = is_test\n        self.augments = augments\n        self.img_folder = img_folder\n        \n    def __getitem__(self, idx):\n        img_path = os.path.join(self.img_folder, self.df['img_name'][idx])\n        img = cv2.imread(img_path)\n        img = cv2.resize(img, Config['IMG_SIZE'])\n        if self.augments:\n            img = self.augments(image=img)['image']\n        img = torch.tensor(img, dtype=torch.float)\n        # Rearrange the image dimensions so that channels are first in format\n        # This is because Swin Transformer Model requires Channels (c) to come first\n        img = rearrange(img, 'h w c -> c h w')\n        \n        if not self.is_test:\n            target = self.df['cancer'][idx]\n            target = torch.tensor(target, dtype=torch.float)\n            return (img, target)\n        return (img)\n    \n    def __len__(self):\n        return len(self.df)","metadata":{"execution":{"iopub.status.busy":"2022-12-08T06:14:27.085521Z","iopub.execute_input":"2022-12-08T06:14:27.086337Z","iopub.status.idle":"2022-12-08T06:14:27.09927Z","shell.execute_reply.started":"2022-12-08T06:14:27.0863Z","shell.execute_reply":"2022-12-08T06:14:27.098217Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Basic Image Augmentations","metadata":{}},{"cell_type":"code","source":"class Augments:\n    \"\"\"\n    Contains Train, Validation Augments\n    \"\"\"\n    train_augments = Compose([\n        Resize(*Config['IMG_SIZE'], p=1.0),\n        HorizontalFlip(p=0.5),\n        VerticalFlip(p=0.5),\n        Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225],\n            max_pixel_value=255.0,\n            p=1.0\n        ),\n        ToTensorV2(p=1.0),\n    ],p=1.)\n    \n    valid_augments = Compose([\n        Resize(*Config['IMG_SIZE'], p=1.0),\n        Normalize(\n            mean=[0.485, 0.456, 0.406],\n            std=[0.229, 0.224, 0.225],\n            max_pixel_value=255.0,\n            p=1.0\n        ),\n        ToTensorV2(p=1.0),\n    ], p=1.)","metadata":{"execution":{"iopub.status.busy":"2022-12-08T06:14:27.104285Z","iopub.execute_input":"2022-12-08T06:14:27.10723Z","iopub.status.idle":"2022-12-08T06:14:27.118701Z","shell.execute_reply.started":"2022-12-08T06:14:27.107193Z","shell.execute_reply":"2022-12-08T06:14:27.117786Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Class using `pl.LightningModule` ⚡️","metadata":{}},{"cell_type":"code","source":"class MyNet(pl.LightningModule):\n    def __init__(self, model_name, num_labels, pretrained):\n        super(MyNet, self).__init__()\n        self.model = timm.create_model(model_name, pretrained=pretrained)\n        self.n_features = self.model.head.in_features\n        self.model.reset_classifier(0)\n        self.fc = nn.Linear(self.n_features, num_labels)\n    def forward(self, x):\n        features = self.model(x)\n        out = self.fc(features)\n        return out","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class RSNAModel(pl.LightningModule):\n    def __init__(self, pretrained=True):\n        super(RSNAModel, self).__init__()\n        # Model Architecture\n#         self.model = timm.create_model(Config['MODEL_NAME'], pretrained=pretrained)\n#         self.n_features = self.model.head.in_features\n#         self.model.reset_classifier(0)\n#         self.fc = nn.Linear(self.n_features, Config['NUM_LABELS'])\n        self.model = MyNet(Config['MODEL_NAME'], Config['NUM_LABELS'], pretrained)\n        # Loss functions\n        self.train_loss = nn.BCEWithLogitsLoss()\n        self.valid_loss = nn.BCEWithLogitsLoss()\n\n        # Metric\n        self.f1 = torchmetrics.F1Score(task='binary')\n        \n    def forward(self, x):\n#         features = self.model(x)\n#         output = self.fc(features)\n        output = self.model(x)\n        return output\n    \n    def training_step(self, batch, batch_idx):\n        imgs = batch[0]\n        target = batch[1]\n        \n        out = self(imgs).view(-1)\n        train_loss = self.train_loss(out, target)\n        \n        logs = {'train_loss': train_loss}\n#         wandb_log(train_step_loss=train_loss.item())\n        return {'loss': train_loss, 'log': logs}\n    \n    def validation_step(self, batch, batch_idx):\n        imgs = batch[0]\n        target = batch[1]\n        \n        out = self(imgs).view(-1)\n        valid_loss = self.valid_loss(out, target)\n        \n        self.f1(out, target)\n        f1_current = self.f1(out, target)\n        self.log('f1_valid_epoch', self.f1, on_epoch=True, on_step=True)\n        \n#         wandb_log(val_step_loss=valid_loss.item(), f1_step_valid=f1_current)\n        \n        return {'val_loss': valid_loss, 'f1_score': f1_current}\n    \n    def validation_end(self, outputs):\n        avg_loss = torch.stack([x['val_loss'] for x in outputs]).mean()\n        \n        logs = {'val_loss': avg_loss}\n#         wandb_log(val_loss=avg_loss)\n        \n        print(f\"val_loss: {avg_loss}\")\n        return {'avg_val_loss': avg_loss, 'log': logs}\n    \n    def configure_optimizers(self):\n        opt = torch.optim.Adam(self.parameters(), lr=Config['LR'])\n        sch = torch.optim.lr_scheduler.CosineAnnealingLR(\n            opt, \n            T_max=Config['T_max'],\n            eta_min=Config['min_lr']\n        )\n        \n        return [opt], [sch]","metadata":{"execution":{"iopub.status.busy":"2022-12-08T06:14:27.123645Z","iopub.execute_input":"2022-12-08T06:14:27.126517Z","iopub.status.idle":"2022-12-08T06:14:27.144929Z","shell.execute_reply.started":"2022-12-08T06:14:27.126481Z","shell.execute_reply":"2022-12-08T06:14:27.143953Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## CSV File loading and Model Training","metadata":{}},{"cell_type":"code","source":"# Load the data and pass it onto the training function\ndf = pd.read_csv(\"/kaggle/input/rsna-breast-cancer-detection/train.csv\")\ndf['img_name'] = df['patient_id'].astype(str) + \"/\" + df['image_id'].astype(str) + \".png\"\ndf = df.sample(frac=1).reset_index(drop=True)\ndf.head()\ndf=df[:1000]","metadata":{"execution":{"iopub.status.busy":"2022-12-08T06:14:27.146494Z","iopub.execute_input":"2022-12-08T06:14:27.14731Z","iopub.status.idle":"2022-12-08T06:14:27.405596Z","shell.execute_reply.started":"2022-12-08T06:14:27.147275Z","shell.execute_reply":"2022-12-08T06:14:27.404505Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    kfold = StratifiedGroupKFold(n_splits=Config['SPLITS'])\n    for fold_, (train_idx, valid_idx) in enumerate(kfold.split(df, df['cancer'].values, df['patient_id'].values)):\n        print(f\"{'='*40} Fold: {fold_} / 5 {'='*40}\")\n        \n        train_df = df.loc[train_idx].reset_index(drop=True)\n        valid_df = df.loc[valid_idx].reset_index(drop=True)\n        \n        train_dataset = RSNAData(\n            df = train_df,\n            img_folder = Config['PARENT_PATH']\n        )\n        valid_dataset = RSNAData(\n            df = valid_df,\n            img_folder = Config['PARENT_PATH'],\n        )\n        train_loader = DataLoader(\n            train_dataset,\n            batch_size=Config['TRAIN_BS'],\n            shuffle=True,\n            num_workers=Config['NUM_WORKERS'],\n        )\n        valid_loader = DataLoader(\n            valid_dataset,\n            batch_size=Config['VALID_BS'],\n            shuffle=False,\n            num_workers=Config['NUM_WORKERS'],\n        )\n\n        model = RSNAModel()\n        trainer = pl.Trainer(\n            max_epochs=Config['NB_EPOCHS'],\n            accelerator='auto',\n            devices='auto',\n            deterministic=True,\n        )\n        trainer.fit(model, train_loader, valid_loader)","metadata":{"execution":{"iopub.status.busy":"2022-12-08T06:14:27.407144Z","iopub.execute_input":"2022-12-08T06:14:27.407826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Finish the logging run\n# run.finish()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<center>\n<img src=\"https://img.shields.io/badge/Upvote-If%20you%20like%20my%20work-07b3c8?style=for-the-badge&logo=kaggle\">\n</center>","metadata":{}}]}