{"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":"<img src=\"https://i.imgur.com/rrdabd5.png\">\n\n<center><h1> - TRAIN: PyTorch Model - </h1></center>\n\n> 🦴 **Goal**: Detect and localize cervical spine fractures within CT scans.\n\n### 🆕 What is MONAI\n\n📌 [MONAI](https://monai.io/) is a freely available (open-source), community-supported, PyTorch-based framework for deep learning in *healthcare imaging*. It provides domain-optimized foundational capabilities for developing healthcare imaging training workflows in a native PyTorch paradigm.\n\n<div class=\"alert alert-block alert-info\">\n  <p>💡<b> Special Thanks</b>: A huge thank you to <a href=\"https://www.kaggle.com/boliu0\">Bo</a> and his notebook on <a href=\"https://www.kaggle.com/code/boliu0/monai-3d-cnn-training/notebook\">MONAI 3D CNN - Training</a>. This one helped me tremendously, as it put an order into the steps I had to follow to tackle 2 areas I haven't encountered before: multi-target approach and multi-scans for same instance scenario.</p>\n</div>\n\n### ⬇ Libraries","metadata":{}},{"cell_type":"code","source":"# libjpeg & gdcm without internet access\n# src: https://www.kaggle.com/code/awsaf49/pydicom-conda-helper\n\n# MONAI 3D model\n!pip install -q monai\n!pip install -q git+https://github.com/ildoonet/pytorch-gradual-warmup-lr.git","metadata":{"_kg_hide-input":false,"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-18T10:42:49.425811Z","iopub.execute_input":"2022-10-18T10:42:49.426204Z","iopub.status.idle":"2022-10-18T10:43:12.892923Z","shell.execute_reply.started":"2022-10-18T10:42:49.426171Z","shell.execute_reply":"2022-10-18T10:43:12.891564Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Libraries\nimport os\nimport re\nimport gc\nimport cv2\nimport wandb\nfrom PIL import Image\nimport random\nimport math\nimport shutil\nfrom glob import glob\nfrom tqdm import tqdm\nfrom pprint import pprint\nfrom time import time\nimport warnings\nimport pandas as pd\nimport numpy as np\nimport seaborn as sns\nimport matplotlib as mpl\nfrom matplotlib import cm\nimport matplotlib.patches as patches\nimport matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\nfrom matplotlib.offsetbox import AnnotationBbox, OffsetImage\nfrom matplotlib.colors import ListedColormap, LinearSegmentedColormap\nfrom matplotlib.patches import Rectangle\nfrom IPython.display import display_html\nplt.rcParams.update({'font.size': 16})\n\n# Environment check\nwarnings.filterwarnings(\"ignore\")\nos.environ[\"WANDB_SILENT\"] = \"true\"\nCONFIG = {'competition': 'RSNA_SpineFructure', '_wandb_kernel': 'aot'}\n\n# Custom colors\nclass clr:\n    S = '\\033[1m' + '\\033[94m'\n    E = '\\033[0m'\n    \nmy_colors = [\"#5EAFD9\", \"#449DD1\", \"#3977BB\", \n             \"#2D51A5\", \"#5C4C8F\", \"#8B4679\",\n             \"#C53D4C\", \"#E23836\", \"#FF4633\", \"#FF5746\"]\nCMAP1 = ListedColormap(my_colors)\n\nprint(clr.S+\"Notebook Color Schemes:\"+clr.E)\nsns.palplot(sns.color_palette(my_colors))\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:12.895512Z","iopub.execute_input":"2022-10-18T10:43:12.896279Z","iopub.status.idle":"2022-10-18T10:43:12.99733Z","shell.execute_reply.started":"2022-10-18T10:43:12.896228Z","shell.execute_reply":"2022-10-18T10:43:12.996165Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# PyTorch\nimport torch\nfrom torch.utils.data import TensorDataset, DataLoader, Dataset\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.optim import lr_scheduler\nfrom torch.utils.data.sampler import SubsetRandomSampler, RandomSampler, SequentialSampler\nfrom torch.optim.lr_scheduler import StepLR, ReduceLROnPlateau, CosineAnnealingLR\nimport torchvision\nimport torchvision.transforms as transforms\nfrom warmup_scheduler import GradualWarmupScheduler\nimport albumentations\n\nfrom sklearn.model_selection import GroupKFold, train_test_split, StratifiedKFold\nfrom sklearn.metrics import roc_auc_score, cohen_kappa_score, confusion_matrix\n\n# MONAI 3D\nfrom monai.transforms import Randomizable, apply_transform\nfrom monai.transforms import Compose, Resize, ScaleIntensity, ToTensor, RandAffine\nfrom monai.networks.nets import densenet","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.0035Z","iopub.execute_input":"2022-10-18T10:43:13.004153Z","iopub.status.idle":"2022-10-18T10:43:13.021708Z","shell.execute_reply.started":"2022-10-18T10:43:13.004079Z","shell.execute_reply":"2022-10-18T10:43:13.020216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 🐝 W&B Fork & Run\n\nIn order to run this notebook you will need to input your own **secret API key** within the `! wandb login $secret_value_0` line. \n\n🐝**How do you get your own API key?**\n\nSuper simple! Go to **https://wandb.ai/site** -> Login -> Click on your profile in the top right corner -> Settings -> Scroll down to API keys -> copy your very own key (for more info check [this amazing notebook for ML Experiment Tracking on Kaggle](https://www.kaggle.com/ayuraj/experiment-tracking-with-weights-and-biases)).\n\n<center><img src=\"https://i.imgur.com/fFccmoS.png\" width=500></center>","metadata":{}},{"cell_type":"code","source":"import wandb\n\nwandb.login()","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.029373Z","iopub.execute_input":"2022-10-18T10:43:13.029862Z","iopub.status.idle":"2022-10-18T10:43:13.056582Z","shell.execute_reply.started":"2022-10-18T10:43:13.029789Z","shell.execute_reply":"2022-10-18T10:43:13.055402Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### ⬇ Helper Functions","metadata":{}},{"cell_type":"code","source":"def set_seed(seed=0):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)  \n    torch.cuda.manual_seed(seed)  \n    torch.cuda.manual_seed_all(seed)  \n    torch.backends.cudnn.deterministic = True\n    \n    \ndef show_values_on_bars(axs, h_v=\"v\", space=0.4):\n    '''Plots the value at the end of the a seaborn barplot.\n    axs: the ax of the plot\n    h_v: weather or not the barplot is vertical/ horizontal'''\n    \n    def _show_on_single_plot(ax):\n        if h_v == \"v\":\n            for p in ax.patches:\n                _x = p.get_x() + p.get_width() / 2\n                _y = p.get_y() + p.get_height()\n                value = int(p.get_height())\n                ax.text(_x, _y, format(value, ','), ha=\"center\") \n        elif h_v == \"h\":\n            for p in ax.patches:\n                _x = p.get_x() + p.get_width() + float(space)\n                _y = p.get_y() + p.get_height()\n                value = int(p.get_width())\n                ax.text(_x, _y, format(value, ','), ha=\"left\")\n\n    if isinstance(axs, np.ndarray):\n        for idx, ax in np.ndenumerate(axs):\n            _show_on_single_plot(ax)\n    else:\n        _show_on_single_plot(axs)\n        \n        \ndef atoi(text):\n    return int(text) if text.isdigit() else text\n\ndef natural_keys(text):\n    '''\n    alist.sort(key=natural_keys) sorts in human order\n    http://nedbatchelder.com/blog/200712/human_sorting.html\n    (See Toothy's implementation in the comments)\n    '''\n    return [ atoi(c) for c in re.split(r'(\\d+)', text) ]\n        \n        \n# === 🐝 W&B ===\ndef save_dataset_artifact(run_name, artifact_name, path, data_type=\"dataset\"):\n    '''Saves dataset to W&B Artifactory.\n    run_name: name of the experiment\n    artifact_name: under what name should the dataset be stored\n    path: path to the dataset'''\n    \n    run = wandb.init(project='RSNA_SpineFructure', \n                     name=run_name, \n                     config=CONFIG)\n    artifact = wandb.Artifact(name=artifact_name, \n                              type=data_type)\n    artifact.add_file(path)\n\n    wandb.log_artifact(artifact)\n    wandb.finish()\n    print(\"Artifact has been saved successfully.\")\n    \n    \ndef create_wandb_plot(x_data=None, y_data=None, x_name=None, y_name=None, title=None, log=None, plot=\"line\"):\n    '''Create and save lineplot/barplot in W&B Environment.\n    x_data & y_data: Pandas Series containing x & y data\n    x_name & y_name: strings containing axis names\n    title: title of the graph\n    log: string containing name of log'''\n    \n    data = [[label, val] for (label, val) in zip(x_data, y_data)]\n    table = wandb.Table(data=data, columns = [x_name, y_name])\n    \n    if plot == \"line\":\n        wandb.log({log : wandb.plot.line(table, x_name, y_name, title=title)})\n    elif plot == \"bar\":\n        wandb.log({log : wandb.plot.bar(table, x_name, y_name, title=title)})\n    elif plot == \"scatter\":\n        wandb.log({log : wandb.plot.scatter(table, x_name, y_name, title=title)})\n        \n        \ndef create_wandb_hist(x_data=None, x_name=None, title=None, log=None):\n    '''Create and save histogram in W&B Environment.\n    x_data: Pandas Series containing x values\n    x_name: strings containing axis name\n    title: title of the graph\n    log: string containing name of log'''\n    \n    data = [[x] for x in x_data]\n    table = wandb.Table(data=data, columns=[x_name])\n    wandb.log({log : wandb.plot.histogram(table, x_name, title=title)})","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-18T10:43:13.062391Z","iopub.execute_input":"2022-10-18T10:43:13.065758Z","iopub.status.idle":"2022-10-18T10:43:13.103389Z","shell.execute_reply.started":"2022-10-18T10:43:13.065707Z","shell.execute_reply":"2022-10-18T10:43:13.102372Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# ⬇ Global Params\n\n📌 **Note**: For the moment I run out of memory in the Kaggle environment, no matter how low I set the `epoch_size`, `img_size`, `learning_rate` etc. So I will come back with the results and models I have locally & compare.\n\n<center><img src=\"https://i.imgur.com/SBhl7OB.jpg\" width=800></center>","metadata":{}},{"cell_type":"code","source":"# 🌱 Seed  \ntorch.manual_seed(0)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.104558Z","iopub.execute_input":"2022-10-18T10:43:13.104922Z","iopub.status.idle":"2022-10-18T10:43:13.119054Z","shell.execute_reply.started":"2022-10-18T10:43:13.104892Z","shell.execute_reply":"2022-10-18T10:43:13.118127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DEVICE = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(clr.S+\"Device:\"+clr.E, DEVICE)\n\n# Kaggle Notebook Setup\nDF_SIZE = 0.03\nN_SPLITS = 5\nKERNEL_TYPE = 'densenet121_baseline'\nIMG_RESIZE = 100\nSTACK_RESIZE = 50\nuse_amp = False\nNUM_WORKERS = 1\nBATCH_SIZE = 2\nLR = 0.05\nOUT_DIM = 8\nEPOCHS = 2\n\n# Local Setup\n# DF_SIZE = 1\n# N_SPLITS = 5\n# KERNEL_TYPE = 'densenet121_baseline'\n# IMG_RESIZE = 150\n# STACK_RESIZE = 50\n# use_amp = False\n# NUM_WORKERS = 4\n# BATCH_SIZE = 16\n# LR = 0.0005\n# OUT_DIM = 8\n# EPOCHS = 5","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.122419Z","iopub.execute_input":"2022-10-18T10:43:13.1227Z","iopub.status.idle":"2022-10-18T10:43:13.132406Z","shell.execute_reply.started":"2022-10-18T10:43:13.122675Z","shell.execute_reply":"2022-10-18T10:43:13.13144Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"📌 The difference in this competition is that **we don't have 1 target column, but 8**, as the fracture could be present in multiple locations throughout the cervical area of the spine.","metadata":{}},{"cell_type":"code","source":"target_cols = ['C1', 'C2', 'C3', \n               'C4', 'C5', 'C6', 'C7',\n               'patient_overall']","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.133826Z","iopub.execute_input":"2022-10-18T10:43:13.134075Z","iopub.status.idle":"2022-10-18T10:43:13.140615Z","shell.execute_reply.started":"2022-10-18T10:43:13.134052Z","shell.execute_reply":"2022-10-18T10:43:13.139642Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1. Evaluation Metric Understanding\n\nBelow is an example of how (I think) the formula works.\n\nAs seen in [this discussion post](https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/341854) and [this weighted log loss post](https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/340392) we now know that the `w` (weight) within the given Log Loss formula can have 4 values:\n* 1 - if the *label* is a vertebrae and it's NOT present\n* 2 - if the *label* is a vertebrae and it IS present\n* 7 - if the *label* is from *overall patient* and it's NOT present\n* 14 - if the *label* is from *overall patient* and it IS present","metadata":{}},{"cell_type":"code","source":"# src: https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/341854\ncompetition_weights = {\n    '-' : torch.tensor([1, 1, 1, 1, 1, 1, 1, 7], dtype=torch.float, device=DEVICE),\n    '+' : torch.tensor([2, 2, 2, 2, 2, 2, 2, 14], dtype=torch.float, device=DEVICE),\n}","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.142049Z","iopub.execute_input":"2022-10-18T10:43:13.142837Z","iopub.status.idle":"2022-10-18T10:43:13.151762Z","shell.execute_reply.started":"2022-10-18T10:43:13.142802Z","shell.execute_reply":"2022-10-18T10:43:13.150968Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Example\n\n# # Prediction (very bad)\n# logits = torch.tensor([[0.2221, 0.1037, 0.0739, 0.1112, 0.1026, 0.0902, 0.1597, 0.1365],\n#                        [0.1702, 0.0952, 0.0815, 0.1262, 0.1185, 0.1097, 0.1675, 0.1312]],\n#                       device=DEVICE)\n# print(clr.S+\"Prediction:\"+clr.E, \"\\n\", logits)\n\n# # Actual\n# targets = torch.tensor([[0., 0., 0., 0., 0., 0., 0., 0.],\n#                         [1., 0., 0., 0., 0., 0., 0., 1.]], device=DEVICE)\n# print(clr.S+\"Target:\"+clr.E, \"\\n\", targets)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.156197Z","iopub.execute_input":"2022-10-18T10:43:13.156746Z","iopub.status.idle":"2022-10-18T10:43:13.163762Z","shell.execute_reply.started":"2022-10-18T10:43:13.15672Z","shell.execute_reply":"2022-10-18T10:43:13.162702Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Compute the weights\n# weights = targets * competition_weights['+'] + (1 - targets) * competition_weights['-']\n# print(clr.S+\"Weights:\"+clr.E, \"\\n\", weights)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.165879Z","iopub.execute_input":"2022-10-18T10:43:13.166165Z","iopub.status.idle":"2022-10-18T10:43:13.173213Z","shell.execute_reply.started":"2022-10-18T10:43:13.166131Z","shell.execute_reply":"2022-10-18T10:43:13.17227Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Compute losses on label and exam level\n# L = torch.zeros(targets.shape, device=DEVICE)\n\n# w = weights\n# y = targets\n# p = logits\n\n# for i in range(L.shape[0]):\n#     for j in range(L.shape[1]):\n#         L[i, j] = -w[i, j] * (\n#             y[i, j] * math.log(p[i, j]) +\n#             (1 - y[i, j]) * math.log(1 - p[i, j]))\n        \n# print(clr.S+\"LOSSES:\"+clr.E, \"\\n\", L)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.175774Z","iopub.execute_input":"2022-10-18T10:43:13.176612Z","iopub.status.idle":"2022-10-18T10:43:13.181914Z","shell.execute_reply.started":"2022-10-18T10:43:13.176578Z","shell.execute_reply":"2022-10-18T10:43:13.180909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Average Loss on Exam (or patient)\n# Exams_Loss = torch.div(torch.sum(L, dim=1), torch.sum(w, dim=1))\n\n# print(clr.S+\"Exam Losses:\"+clr.E, \"\\n\", Exams_Loss)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.185143Z","iopub.execute_input":"2022-10-18T10:43:13.185582Z","iopub.status.idle":"2022-10-18T10:43:13.191523Z","shell.execute_reply.started":"2022-10-18T10:43:13.185556Z","shell.execute_reply":"2022-10-18T10:43:13.190869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"> Below is a breakdown on how the formula works (this breakdown helped me A LOT - curtosy of [🦴 RSNA Fracture Detection - in-depth EDA](https://www.kaggle.com/code/samuelcortinhas/rsna-fracture-detection-in-depth-eda) notebook by [Samuel Cortinhas](https://www.kaggle.com/samuelcortinhas).\n\n<center><img src=\"https://i.imgur.com/GXxNezD.jpg\" width=800></center>\n\n## ⤵ Custom Loss Function\n\n🦴 **Side Note**: the `eps` parameter is so that we never have a `log(0)` - which is *undefined*. More on this in [this discussion post](https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/349669).","metadata":{}},{"cell_type":"code","source":"def get_custom_loss(logits, targets):\n    \n    # Compute the weights\n    weights = targets * competition_weights['+'] + (1 - targets) * competition_weights['-']\n    \n    # Losses on label and exam level\n    L = torch.zeros(targets.shape, device=DEVICE)\n\n    w = weights\n    y = targets\n    p = logits\n    eps=1e-8\n\n    for i in range(L.shape[0]):\n        for j in range(L.shape[1]):\n            L[i, j] = -w[i, j] * (\n                y[i, j] * math.log(p[i, j] + eps) +\n                (1 - y[i, j]) * math.log(1 - p[i, j] + eps))\n            \n    # Average Loss on Exam (or patient)\n    Exams_Loss = torch.div(torch.sum(L, dim=1), torch.sum(w, dim=1))\n    \n    return Exams_Loss","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.192626Z","iopub.execute_input":"2022-10-18T10:43:13.193181Z","iopub.status.idle":"2022-10-18T10:43:13.201884Z","shell.execute_reply.started":"2022-10-18T10:43:13.193139Z","shell.execute_reply":"2022-10-18T10:43:13.20116Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 2. Data Split\n\n📌 We will split the data on folds, **bearing in mind to group based on the `StudyInstanceUID`** - meaning that we want to have all CT scans from one study grouped together, not scattered throughout the training and validation data.","metadata":{}},{"cell_type":"code","source":"np.random.seed(0)\n\ndf = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\n\n# Sample down df\ninstances = df.StudyInstanceUID.unique().tolist()\ninstances = random.sample(instances, k=int(len(instances)*DF_SIZE))\ndf = df[df[\"StudyInstanceUID\"].isin(instances)].reset_index(drop=True)\nprint(clr.S+\"Dataframe size:\"+clr.E, df.shape)\n\n# Create folds\nkfold = GroupKFold(n_splits=N_SPLITS)\ndf['fold'] = -1\n\n# Append fold\nfor k, (_, valid_i) in enumerate(kfold.split(df,\n                                             groups=df.StudyInstanceUID)):\n    df.loc[valid_i, 'fold'] = k\n    \nprint(clr.S+\"K Folds Count:\"+clr.E)\ndf[\"fold\"].value_counts()","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.203175Z","iopub.execute_input":"2022-10-18T10:43:13.203768Z","iopub.status.idle":"2022-10-18T10:43:13.232722Z","shell.execute_reply.started":"2022-10-18T10:43:13.203733Z","shell.execute_reply":"2022-10-18T10:43:13.231803Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3. PyTorch Dataset\n\nThe dataset!\n\nAnother **difference** that we have in this competition is that we have **multiple CT scans** for one case. Meaning that we can use multiple images to predict if there is one or multiple fractures involved in the particular study.\n\nThis is an area I got stuck again (and this is why Bo's [notebook](https://www.kaggle.com/code/boliu0/monai-3d-cnn-training/notebook) helped clear things out for me), cuz how was I supposed to use all images at once?\n\nI gather from multiple solutions that there are *many* ways this could be approached. In here we will use `np.stack()` to stack all the scans together.\n\nBelow is a detailed explanation on how the code below works.\n\n<center><img src=\"https://i.imgur.com/sVqtvWe.jpg\" width=800></center>","metadata":{}},{"cell_type":"code","source":"class RSNADataset(Dataset):\n    \n    def __init__(self, csv, mode, transform=None):\n        self.csv = csv\n        self.mode = mode\n        self.transform = transform\n        \n    def __len__(self):\n        return self.csv.shape[0]\n        \n    def __getitem__(self, index):\n        # Set Random Seed\n        dt = self.csv.iloc[index, :]\n        study_paths = glob(f\"../input/rsna-fracture-detection/zip_png_images/{dt.StudyInstanceUID}/*\")\n        \n        \n        # Load images\n        study_images = [cv2.imread(path)[:,:,::-1] for path in study_paths]\n        # Stack all scans into 1\n        stacked_image = np.stack([img.astype(np.float32) for img in study_images], \n                                 axis=2).transpose(3,0,1,2)\n        \n       \n        stacked_image = apply_transform(self.transform, stacked_image)\n            \n        if self.mode==\"test\":\n            return {\"image\": stacked_image}\n        else:\n            targets = torch.tensor(dt[target_cols]).float()\n            return {\"image\": stacked_image,\n                    \"targets\": targets}","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.233993Z","iopub.execute_input":"2022-10-18T10:43:13.234572Z","iopub.status.idle":"2022-10-18T10:43:13.243757Z","shell.execute_reply.started":"2022-10-18T10:43:13.234538Z","shell.execute_reply":"2022-10-18T10:43:13.2426Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.1 Custom to_device() function\n\nThis function sends the data to **GPU**.","metadata":{}},{"cell_type":"code","source":"def data_to_device(data):\n    \n    image, targets = data.values()\n    return image.to(DEVICE), targets.to(DEVICE)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.24701Z","iopub.execute_input":"2022-10-18T10:43:13.247349Z","iopub.status.idle":"2022-10-18T10:43:13.256231Z","shell.execute_reply.started":"2022-10-18T10:43:13.247324Z","shell.execute_reply":"2022-10-18T10:43:13.255146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 3.2 Custom transform\n\n*TODO: Add more transforms to the training part.*","metadata":{}},{"cell_type":"code","source":"train_transforms = Compose([ScaleIntensity(), \n                            Resize((IMG_RESIZE, IMG_RESIZE, STACK_RESIZE)), \n                            # TODO - add more here\n                            ToTensor()])\nvalid_transforms = Compose([ScaleIntensity(), \n                          Resize((IMG_RESIZE, IMG_RESIZE, STACK_RESIZE)), \n                          ToTensor()])","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.258378Z","iopub.execute_input":"2022-10-18T10:43:13.258702Z","iopub.status.idle":"2022-10-18T10:43:13.269035Z","shell.execute_reply.started":"2022-10-18T10:43:13.258668Z","shell.execute_reply":"2022-10-18T10:43:13.268146Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 🩺 3.3 Sanity Check\n\nLet's test the function with an example and see if everything works as it should.\n\n<center><img src=\"https://i.imgur.com/NRn4HDz.jpg\" width=800></center>","metadata":{}},{"cell_type":"code","source":"# # Sample data\n# sample_df = df.head(6)\n\n# # Instantiate Dataset object\n# dataset = RSNADataset(csv=sample_df, mode=\"train\", transform=train_transforms)\n# # The Dataloader\n# dataloader = DataLoader(dataset, batch_size=2, shuffle=True)\n\n# # Output of the Dataloader\n# for k, data in enumerate(dataloader):\n#     image, targets = data_to_device(data)\n#     print(clr.S + f\"Batch: {k}\" + clr.E, \"\\n\" +\n#           clr.S + \"Image:\" + clr.E, image.shape, \"\\n\" +\n#           clr.S + \"Targets:\" + clr.E, targets, \"\\n\" +\n#           \"=\"*50)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.272691Z","iopub.execute_input":"2022-10-18T10:43:13.273625Z","iopub.status.idle":"2022-10-18T10:43:13.278387Z","shell.execute_reply.started":"2022-10-18T10:43:13.27359Z","shell.execute_reply":"2022-10-18T10:43:13.277337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# del dataset, dataloader, image, targets\n# gc.collect()","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.279714Z","iopub.execute_input":"2022-10-18T10:43:13.280846Z","iopub.status.idle":"2022-10-18T10:43:13.290807Z","shell.execute_reply.started":"2022-10-18T10:43:13.280814Z","shell.execute_reply":"2022-10-18T10:43:13.289382Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4. Loss & Gradual Warmup","metadata":{}},{"cell_type":"code","source":"CRITERION = nn.BCEWithLogitsLoss(reduction='none')\n\ndef get_criterion(logits, target): \n    loss = CRITERION(logits.view(-1), target.view(-1))\n    return loss","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.292357Z","iopub.execute_input":"2022-10-18T10:43:13.2929Z","iopub.status.idle":"2022-10-18T10:43:13.300055Z","shell.execute_reply.started":"2022-10-18T10:43:13.292866Z","shell.execute_reply":"2022-10-18T10:43:13.298871Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### 🦴 What is Gradual Warmup\n\n[Warm-up](https://hasty.ai/docs/mp-wiki/scheduler/warm-up) is a way to **reduce the primacy effect for adaptive schedulers** like *Adam* or *AdamW* of the early training examples. It allows them to compute the correct gradients from the beginning on. Without it, you *may need to run a few extra epochs* to get the convergence desired.\n\n**Learning Rate**: Using a too large learning rate may result in numerical instability *especially at the very beginning of the training*, where parameters are randomly initialized. The **warmup strategy increases the learning rate from 0 to the initial learning rate linearly** during the initial N epochs or m batches.\n\n📍 **Example**:\n* if we have a total of 5000 epochs\n* if `Last Epoch` = 1000\n* and we use Gradual Warmup THEN:\n    * the first 1000 iterations the model will learn the corpus with **minimal learning rate** than the rate which you've specified in the model\n    * from the 1001th iteration, **model will use the previously defined base learning rate**","metadata":{}},{"cell_type":"code","source":"class GradualWarmupSchedulerV2(GradualWarmupScheduler):\n    '''\n    src: https://www.kaggle.com/code/boliu0/monai-3d-cnn-training/notebook\n    '''\n    \n    def __init__(self, optimizer, multiplier, total_epoch, after_scheduler=None):\n        super(GradualWarmupSchedulerV2, self).__init__(optimizer, multiplier, \n                                                       total_epoch, after_scheduler)\n    \n    def get_lr(self):\n        if self.last_epoch > self.total_epoch:\n            if self.after_scheduler:\n                if not self.finished:\n                    self.after_scheduler.base_lrs = [base_lr * self.multiplier \n                                                     for base_lr in self.base_lrs]\n                    self.finished = True\n                return self.after_scheduler.get_lr()\n            return [base_lr * self.multiplier \n                    for base_lr in self.base_lrs]\n        \n        if self.multiplier == 1.0:\n            return [base_lr * (float(self.last_epoch) / self.total_epoch) \n                    for base_lr in self.base_lrs]\n        else:\n            return [base_lr * ((self.multiplier - 1.) * self.last_epoch / self.total_epoch + 1.) \n                    for base_lr in self.base_lrs]","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.302096Z","iopub.execute_input":"2022-10-18T10:43:13.302488Z","iopub.status.idle":"2022-10-18T10:43:13.313914Z","shell.execute_reply.started":"2022-10-18T10:43:13.302453Z","shell.execute_reply":"2022-10-18T10:43:13.31298Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 5. Training Preparation\n\n### Custom File\n\nAs I won't be doing the actual training of the model in this notebook (due to **out of memory issues**), I am also implementing within the functions an `add_in_file()` custom that will allow me to store all the logs during training and then share them with you here :).","metadata":{}},{"cell_type":"code","source":"def add_in_file(text, f):\n    \n    with open(f'log_{KERNEL_TYPE}.txt', 'a+') as f:\n        print(text, file=f)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.315681Z","iopub.execute_input":"2022-10-18T10:43:13.31604Z","iopub.status.idle":"2022-10-18T10:43:13.32374Z","shell.execute_reply.started":"2022-10-18T10:43:13.316007Z","shell.execute_reply":"2022-10-18T10:43:13.322664Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.1 Train epoch\n\nThe function below does all the steps for **training** the model for only 1 epoch.","metadata":{}},{"cell_type":"code","source":"def train_epoch(model, dataloader, optimizer, epoch, f):\n    \n    # Add info to file\n    print(\"Training...\")\n    add_in_file('Training...', f)\n    \n    # Track training time for 1 epoch\n    start_time = time()\n    \n    # === TRAIN ===\n    model.train()\n    train_losses, train_comp_losses = [], []\n    \n    # Loop through the data\n    bar = tqdm(dataloader)\n    for data in bar:\n        image, targets = data_to_device(data)\n        \n        # Train & Optimize\n        optimizer.zero_grad()\n        logits = model(image)\n        loss = get_criterion(logits, targets)\n        loss.sum().backward()\n        optimizer.step()\n        \n        # === COMP LOSS ===\n        comp_loss = get_custom_loss(logits, targets)\n\n        # Save losses\n        train_losses.append(loss.detach().cpu().numpy())\n        train_comp_losses.append(comp_loss.detach().cpu().numpy().mean())\n        \n        gc.collect()\n\n    # Compute Overall Loss\n    mean_train_loss = np.mean(train_losses)\n    mean_comp_loss = np.mean(train_comp_losses)\n    \n    # Save info\n    total_time = round((time() - start_time)/60, 3)\n    add_in_file('Train Mean Loss: {}'.format(mean_train_loss), f)\n    add_in_file('Train Mean Comp Loss: {}'.format(mean_comp_loss), f)\n    add_in_file('~~~ Train Time: {} mins ~~~'.format(total_time), f)\n    \n    # 🐝 Log to W&B\n    wandb.log({\"train_loss\": mean_train_loss,\n               \"train_comp_loss\": mean_comp_loss,}, step=epoch)\n                \n    # Print info\n    print(clr.S+\"Train Mean Loss:\"+clr.E, mean_train_loss)\n    print(clr.S+\"Train Mean Comp Loss:\"+clr.E, mean_comp_loss)\n    print(clr.S+f\"~~~ Train Time: {total_time} mins ~~~\"+clr.E)\n    \n    return mean_train_loss","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.325548Z","iopub.execute_input":"2022-10-18T10:43:13.325934Z","iopub.status.idle":"2022-10-18T10:43:13.338078Z","shell.execute_reply.started":"2022-10-18T10:43:13.325898Z","shell.execute_reply":"2022-10-18T10:43:13.337094Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.2 Validation epoch\n\nThis function does all the steps to **validate** the model for 1 epoch.","metadata":{}},{"cell_type":"code","source":"def valid_epoch(model, dataloader, epoch, f):\n    \n    # Add info to file\n    print(\"Validation...\")\n    add_in_file('Validation...', f)\n    \n    # Track validation time for 1 epoch\n    start_time = time()\n    \n    # === EVAL ===\n    model.eval()\n    valid_preds, valid_targets, valid_comp_loss = [], [], []\n    \n    with torch.no_grad():\n        for data in dataloader:\n            \n            image, targets = data_to_device(data)\n            logits = model(image)\n            \n            # === COMP LOSS ===\n            comp_loss = get_custom_loss(logits, targets)\n            # Save actuals, preds and losses\n            valid_targets.append(targets.detach().cpu())\n            valid_preds.append(logits.detach().cpu())\n            valid_comp_loss.append(comp_loss.detach().cpu().numpy().mean())\n            \n            gc.collect()\n\n    # Overall Valid Loss\n    valid_losses = get_criterion(torch.cat(valid_preds), torch.cat(valid_targets)).numpy()\n    mean_valid_loss = np.mean(valid_losses)\n    \n    # Overall Competition Loss\n    mean_comp_valid_loss = np.mean(valid_comp_loss)\n    \n    # Compute Area Under Curve\n    PREDS = np.concatenate(torch.cat(valid_preds).numpy())\n    TARGETS = np.concatenate(torch.cat(valid_targets).numpy())\n    auc = roc_auc_score(TARGETS, PREDS)\n    \n    # Save info\n    total_time = round((time() - start_time)/60, 3)\n    add_in_file('Valid Mean Loss: {}'.format(mean_valid_loss), f)\n    add_in_file('Valid Mean Comp Loss: {}'.format(mean_comp_valid_loss), f)\n    add_in_file('Valid AUC: {}'.format(auc), f)\n    add_in_file('~~~ Valid Time: {} mins ~~~'.format(total_time), f)\n    \n    # 🐝 Log to W&B\n    wandb.log({\"valid_loss\": mean_valid_loss,\n               \"valid_comp_loss\": mean_comp_valid_loss,\n               \"valid_auc\": auc}, step=epoch)\n        \n    # Print info\n    print(clr.S+\"Valid Mean Loss:\"+clr.E, mean_valid_loss)\n    print(clr.S+\"Valid Mean Comp Loss:\"+clr.E, mean_comp_valid_loss)\n    print(clr.S+\"Valid AUC:\"+clr.E, auc)\n    print(clr.S+f\"~~~ Validation Time: {total_time} mins ~~~\"+clr.E)\n    \n    return mean_valid_loss","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.339603Z","iopub.execute_input":"2022-10-18T10:43:13.340414Z","iopub.status.idle":"2022-10-18T10:43:13.353576Z","shell.execute_reply.started":"2022-10-18T10:43:13.340371Z","shell.execute_reply":"2022-10-18T10:43:13.352707Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.3 Experiment running function\n\nThis is the main function, where we initiate the Dataset and Dataloaders, the Model, Optimizer and Scheduler and where we train and validate the model, after which we save it.\n\n<center><img src=\"https://i.imgur.com/9AoM0Wt.jpg\" width=800></center>","metadata":{}},{"cell_type":"code","source":"def run_train(fold):\n    \n    # 🐝 W&B Tracking\n    RUN_CONFIG = CONFIG.copy()\n    params = dict(model=\"densenet121\", \n                  epochs=EPOCHS, \n                  split=N_SPLITS, \n                  batch=BATCH_SIZE, lr=LR,\n                  img_size=IMG_RESIZE, stack_size=STACK_RESIZE,\n                  data_size=DF_SIZE)\n    RUN_CONFIG.update(params)\n    run = wandb.init(project='RSNA_SpineFructure', config=CONFIG)\n    \n    # Get the train and valid data\n    train = df[df[\"fold\"] != fold].reset_index(drop=True)\n    valid = df[df[\"fold\"] == fold].reset_index(drop=True)\n    \n    # Create the Dataset & Dataloader\n    train_dataset = RSNADataset(csv=train, mode=\"train\", \n                                transform=train_transforms)\n    valid_dataset = RSNADataset(csv=valid, mode=\"train\", \n                                transform=valid_transforms)\n    \n    trainloader = DataLoader(train_dataset, batch_size=BATCH_SIZE,\n                             sampler=RandomSampler(train_dataset))\n    validloader = DataLoader(valid_dataset, batch_size=BATCH_SIZE)\n    \n    # Model\n    model = densenet.densenet121(spatial_dims=3, in_channels=3,\n                                 out_channels=OUT_DIM)\n    model.class_layers.out = nn.Sequential(nn.Linear(in_features=1024, out_features=OUT_DIM), \n                                           nn.Softmax(dim=1))\n    model.to(DEVICE)\n    wandb.watch(model, log_freq=100) # 🐝\n    \n    # Optimizer & Scheduler\n    optimizer = optim.Adam(model.parameters(), lr=LR)\n    scheduler_cosine = lr_scheduler.CosineAnnealingWarmRestarts(optimizer, 2)\n    scheduler_warmup = GradualWarmupSchedulerV2(optimizer, multiplier=10, \n                                                total_epoch=1, \n                                                after_scheduler=scheduler_cosine)\n    \n    # Initiate initial loss\n    valid_loss_BEST = 1000\n    # Create model name\n    model_file = f'{KERNEL_TYPE}_best_fold{fold}.pth'\n    # Create file to save outputs\n    f = open(f'log_{KERNEL_TYPE}.txt', 'a')\n    \n    \n    for epoch in range(EPOCHS):\n        \n        add_in_file('======== Epoch: {}/{} ========'.format(epoch+1, EPOCHS), f)\n        print(\"=\"*8, clr.S+f\"Epoch {epoch}\"+clr.E, \"=\"*8)\n        \n        scheduler_warmup.step(epoch-1)\n        \n        # Train & Validate\n        mean_train_loss = train_epoch(model, trainloader, optimizer, epoch, f)\n        mean_valid_loss = valid_epoch(model, validloader, epoch, f)\n        \n        # Save model\n        if mean_valid_loss < valid_loss_BEST:\n            print('Saving model ...')\n            add_in_file('Saving model => {}'.format(model_file), f)\n            torch.save(model.state_dict(), model_file)\n            valid_loss_BEST = mean_valid_loss\n            \n    torch.cuda.empty_cache()\n    gc.collect()\n    \n    # 🐝 Experiment End\n    wandb.finish()","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.355347Z","iopub.execute_input":"2022-10-18T10:43:13.356211Z","iopub.status.idle":"2022-10-18T10:43:13.371251Z","shell.execute_reply.started":"2022-10-18T10:43:13.356171Z","shell.execute_reply":"2022-10-18T10:43:13.369944Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 6. Training ... 🏃‍♂️","metadata":{}},{"cell_type":"code","source":"# Line commented due to out of memory issue\nfor i in range(4):\n    run_train(fold=i)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T10:43:13.374267Z","iopub.execute_input":"2022-10-18T10:43:13.374918Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"🐝 Below are some of my experiments and logs (trained on just a couple of epochs - as tests - vs 5 epochs)\n<center><img src=\"https://i.imgur.com/f1OevsB.jpg\" width=900></center>","metadata":{}},{"cell_type":"code","source":"# Print the output of the run (done on local machine)\nf = open('../input/rsna-fracture-detection/log_densenet121_baseline.txt', 'r')\nprint(f.read())\nf.close()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 🐝 Save Artifacts\nsave_dataset_artifact(run_name=\"save_logs\", artifact_name=\"logs\",\n                      path=\"../input/rsna-fracture-detection/log_densenet121_baseline.txt\", data_type=\"dataset\")\nsave_dataset_artifact(run_name=\"save_model\", artifact_name=\"model\",\n                      path=\"../input/rsna-fracture-detection/densenet121_baseline_best_fold0.pth\", data_type=\"model\")","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<center><img src=\"https://i.imgur.com/0cx4xXI.png\"></center>\n\n### 🐝 W&B Dashboard\n\n> My [W&B Dashboard](https://wandb.ai/andrada/RSNA_SpineFructure?workspace=user-andrada) - updated with training logs.\n\n<center><video src=\"https://i.imgur.com/5WOgE7W.mp4\" width=800 controls></center>\n\n<center><img src=\"https://i.imgur.com/knxTRkO.png\"></center>\n\n### My Specs\n\n* 🖥 Z8 G4 Workstation\n* 💾 2 CPUs & 96GB Memory\n* 🎮 2x NVIDIA A6000\n* 💻 Zbook Studio G7 on the go","metadata":{}}]}