{"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":"### Model Architecture\n\n<center>\n<img src='https://i.postimg.cc/ZYV9RwXt/2-D-model-architecture.png' width=800>\n</center>\n\n<br>\n\n### Links\n\n1. [RSNA Fracture Detection - in-depth EDA](https://www.kaggle.com/code/samuelcortinhas/rsna-fracture-detection-in-depth-eda)\n2. [RSNA 2022 Spine Fracture Detection - Metadata](https://www.kaggle.com/datasets/samuelcortinhas/rsna-2022-spine-fracture-detection-metadata)\n3. [RNSA - 2D model [Train] [PyTorch]](https://www.kaggle.com/code/samuelcortinhas/rnsa-2d-model-train-pytorch)\n4. [RNSA - 2D model [Re-Train] [PyTorch]](https://www.kaggle.com/code/samuelcortinhas/rnsa-2d-model-re-train-pytorch)\n5. [RSNA - Trained 2D models [PyTorch]](https://www.kaggle.com/datasets/samuelcortinhas/rsna-trained-2d-models-pytorch)\n6. [RNSA - 2D model [Validate] [PyTorch]](https://www.kaggle.com/code/samuelcortinhas/rnsa-2d-model-validate-pytorch)\n7. [RNSA - 2D model [Infer] [PyTorch]](https://www.kaggle.com/code/samuelcortinhas/rnsa-2d-model-infer-pytorch)\n\n<hr>\n\nMuch of my work is inspired by https://www.kaggle.com/code/vslaykovsky/train-pytorch-effnetv2-baseline-cv-0-49","metadata":{}},{"cell_type":"markdown","source":"The purpose of this notebook is to load in a trained model and train it further.","metadata":{}},{"cell_type":"markdown","source":"# Libraries","metadata":{}},{"cell_type":"code","source":"!pip install -qU ../input/for-pydicom/python_gdcm-3.0.14-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl ../input/for-pydicom/pylibjpeg-1.4.0-py3-none-any.whl --find-links frozen_packages --no-index","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-18T11:41:12.860359Z","iopub.execute_input":"2022-10-18T11:41:12.860981Z","iopub.status.idle":"2022-10-18T11:41:26.486739Z","shell.execute_reply.started":"2022-10-18T11:41:12.860869Z","shell.execute_reply":"2022-10-18T11:41:26.485284Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!pip install 'git+https://github.com/katsura-jp/pytorch-cosine-annealing-with-warmup'","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-18T11:41:26.489634Z","iopub.execute_input":"2022-10-18T11:41:26.490125Z","iopub.status.idle":"2022-10-18T11:41:40.267367Z","shell.execute_reply.started":"2022-10-18T11:41:26.490074Z","shell.execute_reply":"2022-10-18T11:41:40.266162Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"!mkdir -p /root/.cache/torch/hub/checkpoints/\n!cp ../input/rsna-2022-whl/efficientnet_v2_s-dd5fe13b.pth  /root/.cache/torch/hub/checkpoints/\n!pip install /kaggle/input/rsna-2022-whl/{pydicom-2.3.0-py3-none-any.whl,pylibjpeg-1.4.0-py3-none-any.whl,python_gdcm-3.0.15-cp37-cp37m-manylinux_2_17_x86_64.manylinux2014_x86_64.whl}\n!pip install /kaggle/input/rsna-2022-whl/{torch-1.12.1-cp37-cp37m-manylinux1_x86_64.whl,torchvision-0.13.1-cp37-cp37m-manylinux1_x86_64.whl}","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-18T11:41:40.269077Z","iopub.execute_input":"2022-10-18T11:41:40.269555Z","iopub.status.idle":"2022-10-18T11:43:03.696423Z","shell.execute_reply.started":"2022-10-18T11:41:40.269515Z","shell.execute_reply":"2022-10-18T11:43:03.694809Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n%matplotlib inline\nimport matplotlib.patches as patches\nimport seaborn as sns\nsns.set(style='darkgrid', font_scale=1.6)\nimport cv2\nimport os\nfrom os import listdir\nimport re\nimport gc\nimport random\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nfrom tqdm.auto import tqdm\nfrom pprint import pprint\nfrom time import time\nimport itertools\nfrom skimage import measure\nfrom mpl_toolkits.mplot3d.art3d import Poly3DCollection\nimport nibabel as nib\nfrom glob import glob\nimport warnings\n#warnings.filterwarnings(\"ignore\", category=DeprecationWarning)\nwarnings.filterwarnings(\"ignore\", category=UserWarning)\n#warnings.filterwarnings(\"ignore\", category=FutureWarning)\nimport zipfile\nfrom scipy import ndimage\nfrom sklearn.model_selection import train_test_split, GroupKFold\nfrom joblib import Parallel, delayed\nfrom PIL import Image\nfrom dipy.denoise.nlmeans import nlmeans\nfrom dipy.denoise.noise_estimate import estimate_sigma\n#from kaggle_volclassif.utils import interpolate_volume\nfrom skimage import exposure\n\n# Pytorch\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision as tv\nimport torchvision.transforms as transforms\nimport torch.nn.functional as F\nimport kornia\nimport kornia.augmentation as augmentation\nimport albumentations as A\nfrom cosine_annealing_warmup import CosineAnnealingWarmupRestarts\n\nfrom torch.cuda.amp import GradScaler, autocast\nfrom torchvision.models.feature_extraction import create_feature_extractor\n#from tqdm.notebook import tqdm\n#import wandb","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-18T11:43:03.700306Z","iopub.execute_input":"2022-10-18T11:43:03.700821Z","iopub.status.idle":"2022-10-18T11:43:06.961103Z","shell.execute_reply.started":"2022-10-18T11:43:03.700772Z","shell.execute_reply":"2022-10-18T11:43:06.959902Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Reproducibility","metadata":{}},{"cell_type":"code","source":"# Set random seeds\ndef set_seed(seed=0):\n    np.random.seed(seed)\n    random.seed(seed)\n    torch.manual_seed(seed)\n\nset_seed(11)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:06.962662Z","iopub.execute_input":"2022-10-18T11:43:06.963602Z","iopub.status.idle":"2022-10-18T11:43:06.973291Z","shell.execute_reply.started":"2022-10-18T11:43:06.963558Z","shell.execute_reply":"2022-10-18T11:43:06.971908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"# Hyperparameters\nBATCH_SIZE = 4\n#LEARNING_RATE = 1e-4\n#ONE_CYCLE_MAX_LR = 0.0005\nN_EPOCHS = 1\n#PATIENCE = 3\nEXPERIMENTAL = False\nAUGMENTATIONS = True\nN_FOLDS = 5\nFOLD = 0     # 0,1,2,3,4\nIMG_SIZE = (256,256)\n\n# Scheduler\nMAX_LR = 0.00001\nMIN_LR = 0.000005\nGAMMA = 0.5\n\n# Config device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndevice","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:06.975303Z","iopub.execute_input":"2022-10-18T11:43:06.975792Z","iopub.status.idle":"2022-10-18T11:43:06.991989Z","shell.execute_reply.started":"2022-10-18T11:43:06.975748Z","shell.execute_reply":"2022-10-18T11:43:06.990543Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data","metadata":{}},{"cell_type":"markdown","source":"### Load tables","metadata":{}},{"cell_type":"code","source":"# Load tables\ntrain_df = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train.csv\")\ntrain_bbox = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/train_bounding_boxes.csv\")\ntest_df = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/test.csv\")\nss = pd.read_csv(\"../input/rsna-2022-cervical-spine-fracture-detection/sample_submission.csv\")\n\n# Print dataframe shapes\nprint('train shape:', train_df.shape)\nprint('train bbox shape:', train_bbox.shape)\nprint('test shape:', test_df.shape)\nprint('ss shape:', ss.shape)\nprint('')\n\n# Show first few entries\ntrain_df.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:06.994227Z","iopub.execute_input":"2022-10-18T11:43:06.994735Z","iopub.status.idle":"2022-10-18T11:43:07.066099Z","shell.execute_reply.started":"2022-10-18T11:43:06.994693Z","shell.execute_reply":"2022-10-18T11:43:07.064901Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Slice dataframe with vertebrae predictions","metadata":{}},{"cell_type":"code","source":"# rsna-2022-spine-fracture-detection-metadata contains inference of C1-C7 vertebrae for all training sample\ntrain_df_slices = pd.read_csv('../input/rsna-2022-spine-fracture-detection-metadata/train_segmented.csv')\nc1c7 = [f'C{i}' for i in range(1,8)]\ntrain_df_slices[c1c7] = (train_df_slices[c1c7] > 0.5).astype(int)\n\n# Merge dfs\ntrain_df_vert = train_df_slices.set_index('StudyInstanceUID').join(train_df.set_index('StudyInstanceUID'), rsuffix='_fracture').reset_index().copy()\ntrain_df_vert.head(3)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:10.076451Z","iopub.execute_input":"2022-10-18T11:43:10.076876Z","iopub.status.idle":"2022-10-18T11:43:13.257747Z","shell.execute_reply.started":"2022-10-18T11:43:10.07684Z","shell.execute_reply":"2022-10-18T11:43:13.256723Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Drop bad scans","metadata":{}},{"cell_type":"code","source":"# https://www.kaggle.com/competitions/rsna-2022-cervical-spine-fracture-detection/discussion/344862\nbad_scans = ['1.2.826.0.1.3680043.20574','1.2.826.0.1.3680043.29952']\n\nfor uid in bad_scans:\n    train_df.drop(train_df[train_df['StudyInstanceUID']==uid].index, axis=0, inplace=True)\n    train_df_vert.drop(train_df_vert[train_df_vert['StudyInstanceUID']==uid].index, axis=0, inplace=True)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:13.259409Z","iopub.execute_input":"2022-10-18T11:43:13.259815Z","iopub.status.idle":"2022-10-18T11:43:13.593022Z","shell.execute_reply.started":"2022-10-18T11:43:13.259786Z","shell.execute_reply":"2022-10-18T11:43:13.591615Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Feature engineering","metadata":{}},{"cell_type":"code","source":"# Calculate slice ratio\nslice_max = train_df_vert.groupby('StudyInstanceUID')['Slice'].max().to_dict()\ntrain_df_vert['SliceRatio'] = 0\ntrain_df_vert['SliceRatio'] = train_df_vert['Slice']/train_df_vert['StudyInstanceUID'].map(slice_max)\n\n# Number of slices\ntrain_df_vert['SliceTotal'] = 0\ntrain_df_vert['SliceTotal'] = train_df_vert['StudyInstanceUID'].map(slice_max)\n\n# Reversed indicator\nz_reversed = ((train_df_vert.groupby('StudyInstanceUID')['ImagePositionPatient_z'].first()-train_df_vert.groupby('StudyInstanceUID')['ImagePositionPatient_z'].last())<0).astype('int')\ntrain_df_vert['Reversed'] = 0\ntrain_df_vert['Reversed'] = train_df_vert['StudyInstanceUID'].map(z_reversed)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:13.594592Z","iopub.execute_input":"2022-10-18T11:43:13.594949Z","iopub.status.idle":"2022-10-18T11:43:13.934251Z","shell.execute_reply.started":"2022-10-18T11:43:13.594917Z","shell.execute_reply":"2022-10-18T11:43:13.932913Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Debug","metadata":{}},{"cell_type":"code","source":"if len(ss)==3:\n    # Fix mismatch with test_images folder\n    test_df = pd.DataFrame(columns = ['row_id','StudyInstanceUID','prediction_type'])\n    for i in ['1.2.826.0.1.3680043.22327','1.2.826.0.1.3680043.25399','1.2.826.0.1.3680043.5876']:\n        for j in ['C1','C2','C3','C4','C5','C6','C7','patient_overall']:\n            test_df = test_df.append({'row_id':i+'_'+j,'StudyInstanceUID':i,'prediction_type':j},ignore_index=True)\n    \n    # Sample submission\n    ss = pd.DataFrame(test_df['row_id'])\n    ss['fractured'] = 0.5\n    \n    display(test_df.head(3))","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:15.387202Z","iopub.execute_input":"2022-10-18T11:43:15.38759Z","iopub.status.idle":"2022-10-18T11:43:15.460487Z","shell.execute_reply.started":"2022-10-18T11:43:15.387556Z","shell.execute_reply":"2022-10-18T11:43:15.459381Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Helper functions","metadata":{}},{"cell_type":"code","source":"def load_dicom(path):\n    \"\"\"\n    This supports loading both regular and compressed JPEG images. \n    See the first sell with `pip install` commands for the necessary dependencies\n    \"\"\"\n    img = pydicom.dcmread(path)\n    img.PhotometricInterpretation = 'YBR_FULL'\n    data = img.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return cv2.cvtColor(data, cv2.COLOR_GRAY2RGB), img","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:17.033113Z","iopub.execute_input":"2022-10-18T11:43:17.034441Z","iopub.status.idle":"2022-10-18T11:43:17.040669Z","shell.execute_reply.started":"2022-10-18T11:43:17.034397Z","shell.execute_reply":"2022-10-18T11:43:17.039763Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Augmentations","metadata":{}},{"cell_type":"code","source":"# Data augmentations (albumentations)\nif AUGMENTATIONS:\n    augs = A.Compose([\n        A.Resize(*IMG_SIZE, interpolation=cv2.INTER_LINEAR),\n        A.HorizontalFlip(p=0.35),\n        #A.VerticalFlip(p=0.2),\n        A.ShiftScaleRotate(shift_limit=0.07, scale_limit=0.125, rotate_limit=35, border_mode=cv2.BORDER_CONSTANT, value=0., p=0.45),\n        A.RandomBrightnessContrast(brightness_limit=0.175, contrast_limit=0.175, p=0.5),\n        A.OneOf([\n            A.GridDistortion(num_steps=5, distort_limit=0.05, p=1.0),\n           #A.OpticalDistortion(distort_limit=0.05, shift_limit=0.05, p=1.0),\n            A.ElasticTransform(alpha=1, sigma=50, alpha_affine=25, border_mode=cv2.BORDER_CONSTANT, value=0., p=1.0)\n        ], p=0.25),\n        #A.CoarseDropout(max_holes=8, max_height=512//20, max_width=512//20, min_holes=5, fill_value=0, mask_fill_value=0, p=0.5),\n        ], p=1.0)\nelse:\n    augs = A.Resize(*IMG_SIZE, interpolation=cv2.INTER_LINEAR)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:18.471257Z","iopub.execute_input":"2022-10-18T11:43:18.471933Z","iopub.status.idle":"2022-10-18T11:43:18.481212Z","shell.execute_reply.started":"2022-10-18T11:43:18.471895Z","shell.execute_reply":"2022-10-18T11:43:18.479981Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualise augmentations","metadata":{}},{"cell_type":"code","source":"if AUGMENTATIONS:\n    path = '../input/rsna-2022-cervical-spine-fracture-detection/train_images/1.2.826.0.1.3680043.10515/115.dcm'\n    img, _ = load_dicom(path)\n    plt.imshow(img, cmap='bone')\n    plt.axis('off')\n    plt.title('Original')\n    plt.show()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-18T11:43:20.097069Z","iopub.execute_input":"2022-10-18T11:43:20.097821Z","iopub.status.idle":"2022-10-18T11:43:20.339175Z","shell.execute_reply.started":"2022-10-18T11:43:20.097782Z","shell.execute_reply":"2022-10-18T11:43:20.337934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if AUGMENTATIONS:\n    # Plot images\n    fig, axes = plt.subplots(nrows=3, ncols=6, figsize=(24,12))\n    plt.suptitle('Augmentation examples', y=0.92)\n    for i in range(18):\n        aug_img = augs(image=img)['image']\n\n        # Plot the image\n        x = i // 6\n        y = i % 6\n        axes[x, y].imshow(aug_img, cmap=\"bone\")\n        axes[x, y].axis('off')","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-18T11:43:21.377679Z","iopub.execute_input":"2022-10-18T11:43:21.378082Z","iopub.status.idle":"2022-10-18T11:43:23.663042Z","shell.execute_reply.started":"2022-10-18T11:43:21.378051Z","shell.execute_reply":"2022-10-18T11:43:23.661907Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Torch dataset","metadata":{}},{"cell_type":"code","source":"# Dataset for train/valid sets only\nclass RSNADataset(Dataset):\n    # Initialise\n    def __init__(self, df_table = train_df_vert, transform=None):\n        super().__init__()\n        self.df_table = df_table.reset_index(drop=True)\n        self.transform = transform\n        self.train_path = '../input/rsna-2022-cervical-spine-fracture-detection/train_images'\n        self.meta_cols = ['SliceRatio','SliceTotal','SliceThickness','ImagePositionPatient_x','ImagePositionPatient_y','ImagePositionPatient_z','Reversed']\n        \n    # Get item in position given by index\n    def __getitem__(self, index):\n        # Load image\n        path = os.path.join(self.train_path, self.df_table.iloc[index].StudyInstanceUID, f'{self.df_table.iloc[index].Slice}.dcm')\n        img = load_dicom(path)[0]\n        \n        # Data augmentations\n        if self.transform is not None:\n            img = self.transform(image=img)['image']\n        \n        # Pytorch uses (batch, channel, height, width) order. Converting (height, width, channel) -> (channel, height, width)\n        img = np.transpose(img, (2, 0, 1))\n            \n        # Convert to tensor\n        img = torch.from_numpy(img.astype('float32'))\n        \n        # Targets\n        frac_targets = torch.as_tensor(self.df_table.iloc[index][[f'C{i}_fracture' for i in range(1,8)]].astype('float32').values)\n        vert_targets = torch.as_tensor(self.df_table.iloc[index][[f'C{i}' for i in range(1,8)]].astype('float32').values)\n        frac_targets = frac_targets * vert_targets  # we only enable targets that are visible on the current slice\n        \n        # Metadata\n        meta = torch.as_tensor(self.df_table.iloc[index][self.meta_cols].astype('float32').values)\n        \n        return img, meta, frac_targets, vert_targets\n\n    # Length of dataset\n    def __len__(self):\n        return len(self.df_table)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:23.664866Z","iopub.execute_input":"2022-10-18T11:43:23.665275Z","iopub.status.idle":"2022-10-18T11:43:23.682951Z","shell.execute_reply.started":"2022-10-18T11:43:23.665238Z","shell.execute_reply":"2022-10-18T11:43:23.681474Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Train/valid split (grouped by patients)\nif EXPERIMENTAL:\n    train_patients, valid_patients = train_test_split(train_df['StudyInstanceUID'].values, train_size=0.08, test_size=0.02, random_state=0)\n    \n    # Select train/valid tables\n    train_table = train_df_vert[train_df_vert['StudyInstanceUID'].isin(train_patients)]\n    valid_table = train_df_vert[train_df_vert['StudyInstanceUID'].isin(valid_patients)]\n    \n    # Define train/valid dataset\n    train_dataset = RSNADataset(df_table = train_table, transform=augs)\n    valid_dataset = RSNADataset(df_table = valid_table, transform=A.Resize(*IMG_SIZE, interpolation=cv2.INTER_LINEAR))\nelse:\n    gkf = GroupKFold(N_FOLDS)\n    for k, (train_idx, valid_idx) in enumerate(gkf.split(train_df_vert, groups=train_df_vert.StudyInstanceUID)):\n        if k==FOLD:\n            # Select train/valid tables\n            train_table = train_df_vert.iloc[train_idx]\n            valid_table = train_df_vert.iloc[valid_idx]\n\n    # Define train/valid dataset\n    train_dataset = RSNADataset(df_table = train_table, transform=augs)\n    valid_dataset = RSNADataset(df_table = valid_table, transform=A.Resize(*IMG_SIZE, interpolation=cv2.INTER_LINEAR))","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:23.684468Z","iopub.execute_input":"2022-10-18T11:43:23.684873Z","iopub.status.idle":"2022-10-18T11:43:24.397479Z","shell.execute_reply.started":"2022-10-18T11:43:23.684836Z","shell.execute_reply":"2022-10-18T11:43:24.396352Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Torch dataloaders","metadata":{}},{"cell_type":"code","source":"# Dataloaders\ntrain_loader = DataLoader(dataset=train_dataset, batch_size=BATCH_SIZE, shuffle=True)\nvalid_loader = DataLoader(dataset=valid_dataset, batch_size=BATCH_SIZE, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:24.400204Z","iopub.execute_input":"2022-10-18T11:43:24.401848Z","iopub.status.idle":"2022-10-18T11:43:24.408594Z","shell.execute_reply.started":"2022-10-18T11:43:24.401805Z","shell.execute_reply":"2022-10-18T11:43:24.407349Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class ImageMetaModel(torch.nn.Module):\n    def __init__(self):\n        super().__init__()\n        # Image model layers\n        effnet = tv.models.efficientnet_v2_s()\n        self.effnet = create_feature_extractor(effnet, ['flatten'])\n        self.img_linear = nn.Linear(1280, 64)\n        self.swish = nn.SiLU()\n        \n        # Metadata model layers\n        self.tab_linear1 = nn.Linear(7, 128)\n        self.tab_linear2 = nn.Linear(128, 64)\n        self.relu = nn.ReLU()\n        \n        # Combined model layers\n        self.drop = nn.Dropout(p=0.3)\n        self.linear1 = nn.Linear(128, 256)\n        self.linear2 = nn.Linear(256, 7)\n        self.linear3 = nn.Linear(256, 7)\n        \n        \n    # Forward pass\n    def forward(self, img, meta):\n        # Image model\n        x_img = self.effnet(img)['flatten']\n        x_img = self.img_linear(x_img)\n        x_img = self.swish(x_img)\n        \n        # Metadata model\n        x_meta = self.tab_linear1(meta)\n        x_meta = self.relu(x_meta)\n        x_meta = self.tab_linear2(x_meta)\n        x_meta = self.relu(x_meta)\n        \n        # Concatenate\n        x = torch.cat([x_img, x_meta], dim=1)\n        \n        # Combined model\n        x = self.drop(x)\n        x = self.linear1(x)\n        #x = self.relu(x)\n        \n        # Split\n        x_frac = self.linear2(x)\n        x_vert = self.linear3(x)\n        \n        return x_frac, x_vert\n\n    def predict(self, img, meta):\n        frac, vert = self.forward(img, meta)\n        return torch.sigmoid(frac), torch.sigmoid(vert)\n\nmodel = ImageMetaModel().to(device)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:25.871103Z","iopub.execute_input":"2022-10-18T11:43:25.87152Z","iopub.status.idle":"2022-10-18T11:43:26.59108Z","shell.execute_reply.started":"2022-10-18T11:43:25.871484Z","shell.execute_reply":"2022-10-18T11:43:26.589938Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Load trained model","metadata":{}},{"cell_type":"code","source":"# Load checkpoint\nPATH='../input/rsna-trained-2d-models-pytorch/EffNet_model_fold0_epoch5.pt'\ncheckpoint = torch.load(PATH, map_location=torch.device(device))\n\n# Load states\nmodel.load_state_dict(checkpoint['model_state_dict'])\nprev_epoch = checkpoint['epoch']\nprev_loss = checkpoint['loss']","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:27.535507Z","iopub.execute_input":"2022-10-18T11:43:27.536219Z","iopub.status.idle":"2022-10-18T11:43:30.300368Z","shell.execute_reply.started":"2022-10-18T11:43:27.536167Z","shell.execute_reply":"2022-10-18T11:43:30.298785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Print  previous final loss and epoch\nprint('Current epoch:', prev_epoch)\nprint('Current loss:', prev_loss)","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:30.305044Z","iopub.execute_input":"2022-10-18T11:43:30.305467Z","iopub.status.idle":"2022-10-18T11:43:30.311492Z","shell.execute_reply.started":"2022-10-18T11:43:30.305433Z","shell.execute_reply":"2022-10-18T11:43:30.310246Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Loss & optimiser","metadata":{}},{"cell_type":"code","source":"loss_fn = nn.BCEWithLogitsLoss(reduction='none')\n#loss_fn = nn.BCELoss(reduction=\"none\")\n\ncompetition_weights = {\n    '-' : torch.tensor([1, 1, 1, 1, 1, 1, 1], dtype=torch.float, device=device),\n    '+' : torch.tensor([2, 2, 2, 2, 2, 2, 2], dtype=torch.float, device=device),\n}\n\n# y_hat.shape = (batch_size, num_classes)\n# y.shape = (batch_size, num_classes)\n\ndef FracLoss(y_hat, y):\n    '''Weighted Multilabel Log Loss'''\n    loss = loss_fn(y_hat, y.to(y_hat.dtype))\n    weights = y * competition_weights['+'] + (1 - y) * competition_weights['-']\n    return (loss * weights).sum(axis=1)\n    \ndef VertLoss(y_hat, y):\n    '''BCE Loss'''\n    loss = loss_fn(y_hat, y.to(y_hat.dtype))\n    return loss.sum(axis=1)\n\ndef CombinedLoss(y_hat_frac, y_frac, y_hat_vert, y_vert):\n    # Loss for fracture detection\n    L1 = FracLoss(y_hat_frac, y_frac)\n    \n    # Loss for vertebrae detection\n    L2 = VertLoss(y_hat_vert, y_vert)\n    \n    return (L1+L2).mean()","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:32.70568Z","iopub.execute_input":"2022-10-18T11:43:32.706397Z","iopub.status.idle":"2022-10-18T11:43:32.716954Z","shell.execute_reply.started":"2022-10-18T11:43:32.706356Z","shell.execute_reply":"2022-10-18T11:43:32.715558Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# AdamW optimiser\noptimiser = optim.AdamW(params=model.parameters(), lr=MIN_LR)\n\n# Load state dict\noptimiser.load_state_dict(checkpoint['optimiser_state_dict'])\n\n# Learning rate scheduler\n#scheduler = lr_scheduler.CosineAnnealingLR(optimiser, T_max=N_EPOCHS)\n#scheduler = lr_scheduler.OneCycleLR(optimiser, max_lr=ONE_CYCLE_MAX_LR, total_steps=N_EPOCHS*len(train_loader))\n\n# https://github.com/katsura-jp/pytorch-cosine-annealing-with-warmup\nscheduler = CosineAnnealingWarmupRestarts(optimiser,\n                                          first_cycle_steps=len(train_loader),\n                                          cycle_mult=1.0,\n                                          max_lr=MAX_LR,\n                                          min_lr=MIN_LR,\n                                          warmup_steps=len(train_loader)/5,\n                                          gamma=GAMMA)\n\n# Load state dict\n#scheduler.load_state_dict(checkpoint['scheduler_state_dict'])","metadata":{"execution":{"iopub.status.busy":"2022-10-18T11:43:35.727598Z","iopub.execute_input":"2022-10-18T11:43:35.727979Z","iopub.status.idle":"2022-10-18T11:43:35.946726Z","shell.execute_reply.started":"2022-10-18T11:43:35.727949Z","shell.execute_reply":"2022-10-18T11:43:35.945747Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Visualise learning rate schedule","metadata":{}},{"cell_type":"code","source":"def plot_schedule():\n    opt = optim.AdamW(params=model.parameters(), lr=MIN_LR)\n    #sch = lr_scheduler.OneCycleLR(opt, max_lr=ONE_CYCLE_MAX_LR, total_steps=N_EPOCHS*len(train_loader))\n    #sch = lr_scheduler.CosineAnnealingLR(opt, T_max=N_EPOCHS*len(train_loader), eta_min=0.001)\n    sch = CosineAnnealingWarmupRestarts(opt,\n                                          first_cycle_steps=len(train_loader),\n                                          cycle_mult=1.0,\n                                          max_lr=0.0001,\n                                          min_lr=0.000005,\n                                          warmup_steps=len(train_loader)/5,\n                                          gamma=0.5)\n    \n    lr_hist = []\n    for i in range(3*len(train_loader)):\n        sch.step()\n        lr_hist.append(opt.param_groups[0]['lr'])\n    \n    plt.figure(figsize=(10,4))\n    plt.plot(lr_hist)\n    plt.title('LR schedule')\n    plt.xlabel('Step')\n    plt.ylabel('Learning rate')\n    plt.show()\n    \nplot_schedule()","metadata":{"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-18T11:43:48.031237Z","iopub.execute_input":"2022-10-18T11:43:48.031626Z","iopub.status.idle":"2022-10-18T11:43:49.615273Z","shell.execute_reply.started":"2022-10-18T11:43:48.031595Z","shell.execute_reply":"2022-10-18T11:43:49.613893Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Train model\n\nDue to time constraints, we train the model for an epoch and then do validation in another notebook.","metadata":{}},{"cell_type":"code","source":"loss_hist = []\nloss_step_hist = []\n\n# Loop over epochs\nfor epoch in tqdm(range(N_EPOCHS)):\n    loss_acc = 0\n    loss_step = 0\n    \n    # Loop over batches\n    for idx, (img, meta, frac_targets, vert_targets) in enumerate(train_loader):\n        # Send to device\n        img = img.to(device)\n        meta = meta.to(device)\n        frac_targets = frac_targets.to(device)\n        vert_targets = vert_targets.to(device)\n        \n        # Forward pass\n        frac_preds, vert_preds = model(img, meta)\n        L = CombinedLoss(frac_preds, frac_targets, vert_preds, vert_targets)\n\n        # Backprop\n        L.backward()\n\n        # Update parameters\n        optimiser.step()\n\n        # Zero gradients\n        optimiser.zero_grad()\n        \n        # Track loss\n        loss_step += L.detach().item()\n    \n        # Update learning rate\n        scheduler.step()\n        \n        # Print progress\n        if (idx%1000)==0:\n            loss_acc += loss_step\n            print(f'Minibatch iteration: {idx}/{len(train_loader)}, Loss: {loss_step/1000:.8f}')\n            \n            # Save loss history\n            loss_step_hist.append(loss_step/1000)\n            \n            # Reset\n            loss_step = 0\n    \n    \n    # Save loss history\n    loss_hist.append(loss_acc/len(train_loader))\n    \n    # Print loss\n    print(f'Epoch {epoch+1}/{N_EPOCHS}, loss {loss_acc/len(train_loader):.8f}')\n    \n    # Save model\n    print('Saving model')\n    torch.save({\n        'epoch': epoch+1+prev_epoch,\n        'model_state_dict': model.state_dict(),\n        'optimiser_state_dict': optimiser.state_dict(),\n        'scheduler_state_dict': scheduler.state_dict(),\n        'loss': loss_acc/len(train_loader),\n        }, f\"EffNet_model_fold{FOLD}_epoch{epoch+1+prev_epoch}.pt\")\n    \nprint('')\nprint('Training complete!')","metadata":{"_kg_hide-output":true,"execution":{"iopub.status.busy":"2022-10-12T09:48:41.035845Z","iopub.execute_input":"2022-10-12T09:48:41.036262Z","iopub.status.idle":"2022-10-12T12:30:14.584377Z","shell.execute_reply.started":"2022-10-12T09:48:41.036227Z","shell.execute_reply":"2022-10-12T12:30:14.580529Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# This version includes validation calculations (it wasn't used because of time constraints)\n'''\nloss_hist = []\nloss_step_hist = []\nval_loss_hist = []\npatience_counter = 0\nbest_val_loss = np.inf\n\n# Loop over epochs\nfor epoch in tqdm(range(N_EPOCHS)):\n    loss_acc = 0\n    loss_step = 0\n    val_loss_acc = 0\n    \n    # Loop over batches\n    for idx, (img, meta, frac_targets, vert_targets) in enumerate(train_loader):\n        # Send to device\n        img = img.to(device)\n        meta = meta.to(device)\n        frac_targets = frac_targets.to(device)\n        vert_targets = vert_targets.to(device)\n        \n        # Forward pass\n        frac_preds, vert_preds = model(img, meta)\n        L = CombinedLoss(frac_preds, frac_targets, vert_preds, vert_targets)\n\n        # Backprop\n        L.backward()\n\n        # Update parameters\n        optimiser.step()\n\n        # Zero gradients\n        optimiser.zero_grad()\n        \n        # Track loss\n        loss_step += L.detach().item()\n    \n        # Update learning rate\n        scheduler.step()\n        \n        # Print progress\n        if (idx%1000)==0:\n            loss_acc += loss_step\n            print(f'Minibatch iteration: {idx}/{len(train_loader)}, Loss: {loss_step/1000:.8f}')\n            \n            # Save loss history\n            loss_step_hist.append(loss_step/1000)\n            \n            # Reset\n            loss_step = 0\n    \n    # Don't update weights\n    with torch.no_grad():\n        # Validate\n        for img, meta, frac_targets, vert_targets in valid_loader:\n            # Send to device\n            img = img.to(device)\n            meta = meta.to(device)\n            frac_targets = frac_targets.to(device)\n            vert_targets = vert_targets.to(device)\n\n            # Forward pass\n            frac_preds, vert_preds = model(img, meta)\n            L = CombinedLoss(frac_preds, frac_targets, vert_preds, vert_targets)\n\n            # Track loss\n            val_loss_acc += L.item()\n    \n    # Save loss history\n    loss_hist.append(loss_acc/len(train_loader))\n    val_loss_hist.append(val_loss_acc/len(valid_loader))\n    \n    # Print loss\n    if (epoch+1)%1==0:\n        print(f'Epoch {epoch+1}/{N_EPOCHS}, loss {loss_acc/len(train_loader):.8f}, val_loss {val_loss_acc/len(valid_loader):.8f}')\n    \n    # Save model (& early stopping)\n    if (val_loss_acc/len(valid_loader)) < best_val_loss:\n        best_val_loss = val_loss_acc/len(valid_loader)\n        patience_counter=0\n        print('Valid loss improved --> saving model')\n        torch.save({\n            'epoch': epoch+1+prev_epoch,\n            'model_state_dict': model.state_dict(),\n            'optimiser_state_dict': optimiser.state_dict(),\n            'scheduler_state_dict': scheduler.state_dict(),\n            'loss': loss_acc/len(train_loader),\n            'val_loss': val_loss_acc/len(valid_loader),\n            }, f\"EffNet_model_fold{FOLD}_epoch{epoch+1+prev_epoch}.pt\")\n    else:\n        patience_counter+=1\n        \n        if patience_counter==PATIENCE:\n            break\n    \nprint('')\nprint('Training complete!')\n'''","metadata":{"_kg_hide-output":true,"_kg_hide-input":true,"execution":{"iopub.status.busy":"2022-10-10T09:43:41.22607Z","iopub.status.idle":"2022-10-10T09:43:41.226759Z","shell.execute_reply.started":"2022-10-10T09:43:41.22643Z","shell.execute_reply":"2022-10-10T09:43:41.226482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Learning curves","metadata":{}},{"cell_type":"code","source":"# Plot loss every 1000 minibatches\nplt.figure(figsize=(10,5))\nplt.plot(loss_step_hist[1:], c='C0', label='loss')\nplt.title('Combined loss')\nplt.xlabel('Step (x1000)')\nplt.ylabel('Loss')\nplt.legend()\nplt.show()","metadata":{"_kg_hide-input":false,"execution":{"iopub.status.busy":"2022-10-10T09:43:41.229969Z","iopub.status.idle":"2022-10-10T09:43:41.231197Z","shell.execute_reply.started":"2022-10-10T09:43:41.230929Z","shell.execute_reply":"2022-10-10T09:43:41.230955Z"},"trusted":true},"execution_count":null,"outputs":[]}]}