{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":183439689,"sourceType":"kernelVersion"}],"dockerImageVersionId":30698,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 2.5D Model Training\n\nInference notebook [rsna-2-5d-inference](https://www.kaggle.com/code/samu2505/rsna-2-5d-inference?scriptVersionId=191084741) is here\n\nModel weights and cross validation [rsna-2-5dmodelcheckpoints-cross-validation](https://www.kaggle.com/code/samu2505/rsna-2-5dmodelcheckpoints-cross-validation) is here","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"code","source":"!unzip -q /kaggle/input/rsna2024-lsdc-making-dataset/_output_.zip","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:42:26.975098Z","iopub.execute_input":"2024-08-10T08:42:26.976190Z","iopub.status.idle":"2024-08-10T08:48:28.245816Z","shell.execute_reply.started":"2024-08-10T08:42:26.976141Z","shell.execute_reply":"2024-08-10T08:48:28.242496Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# !python -m pip install -q lightning\n# !pip install -q git+https://github.com/ildoonet/pytorch-gradual-warmup-lr.git","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:28.250750Z","iopub.execute_input":"2024-08-10T08:48:28.252162Z","iopub.status.idle":"2024-08-10T08:48:28.259716Z","shell.execute_reply.started":"2024-08-10T08:48:28.252112Z","shell.execute_reply":"2024-08-10T08:48:28.258303Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os, gc, sys, copy, pickle\nfrom pathlib import Path\nimport glob\nfrom tqdm.auto import tqdm\ntqdm.pandas()\n\nimport math\nimport random\nimport pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\n\nfrom joblib import Parallel, delayed\nimport multiprocessing as mp\n\nimport albumentations as A\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport torch.cuda.amp as amp\nimport torch.nn.functional as F\nfrom torch.utils.data import DataLoader, Dataset\nimport torchvision.transforms as transforms\n\nfrom torch.utils.data import WeightedRandomSampler\nfrom sklearn.utils.class_weight import compute_class_weight\nfrom sklearn import model_selection\n\nfrom transformers import get_cosine_schedule_with_warmup, get_cosine_with_hard_restarts_schedule_with_warmup\n\nimport timm\n\nimport cv2\ncv2.setNumThreads(0)\nimport PIL\nimport pydicom\nfrom IPython import display\nimport warnings\nwarnings.filterwarnings(\"ignore\")\nos.environ[\"CUDA_LAUNCH_BLOCKING\"] = \"1\"","metadata":{"execution":{"iopub.status.busy":"2024-08-10T09:05:36.155246Z","iopub.execute_input":"2024-08-10T09:05:36.156132Z","iopub.status.idle":"2024-08-10T09:05:45.417514Z","shell.execute_reply.started":"2024-08-10T09:05:36.156084Z","shell.execute_reply":"2024-08-10T09:05:45.416551Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# timm.create_model('efficientnetv2_rw_m.ra2_in1k', num_classes=0, pretrained=True)\n# timm.list_pretrained(\"efficientnet*\")","metadata":{"execution":{"iopub.status.busy":"2024-08-10T09:05:45.419615Z","iopub.execute_input":"2024-08-10T09:05:45.420064Z","iopub.status.idle":"2024-08-10T09:05:45.425000Z","shell.execute_reply.started":"2024-08-10T09:05:45.420025Z","shell.execute_reply":"2024-08-10T09:05:45.423818Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seeding(SEED):\n    np.random.seed(SEED)\n    random.seed(SEED)\n    os.environ['PYTHONHASHSEED'] = str(SEED)\n    torch.manual_seed(SEED)\n    if torch.cuda.is_available(): \n        torch.cuda.manual_seed(SEED)\n        torch.cuda.manual_seed_all(SEED)\n        torch.backends.cudnn.deterministic = False\n        torch.backends.cudnn.benchmark = True\n    print('seeding done!!!')\n\ndef flush():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        torch.cuda.reset_peak_memory_stats()","metadata":{"execution":{"iopub.status.busy":"2024-08-10T09:05:45.426379Z","iopub.execute_input":"2024-08-10T09:05:45.426746Z","iopub.status.idle":"2024-08-10T09:05:45.439105Z","shell.execute_reply.started":"2024-08-10T09:05:45.426719Z","shell.execute_reply":"2024-08-10T09:05:45.438033Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONFIG = dict(\n    project_name = \"RSNA-2024-25DModel\",\n    artifact_name = \"rsnaEffNetModel\",\n    load_kernel = None,\n    load_last = True,\n    n_folds = 5,\n    backbone = \"convnext_pico.d1_in1k\", # tf_efficientnetv2_s_in21ft1k, tf_efficientnet_b0.ns_jft_in1k, convnext_pico.d1_in1k\n    img_size = 224,\n    n_slice_per_c = 10,\n    in_chans = 3,\n    axial_chans = 10,\n    axial_labels = 10,\n    axial_classes = 3 * 10,\n    \n    sagT1_chans = 10,\n    sagT1_labels = 10,\n    sagT1_classes = 3 * 10,\n    \n    sagT2_chans = 10,\n    sagT2_labels = 5,\n    sagT2_classes = 3 * 5,\n    \n    n_classes = 3 * 25,\n\n    drop_rate = 0.,\n    drop_rate_last = 0.3,\n    drop_path_rate = 0.,\n    p_mixup = 0.5,\n    p_rand_order_v1 = 0.2,\n    lr = 1e-4,\n    wd = 1e-6,\n\n    epochs = 20,\n    batch_size = 8,\n    warmup = 1,\n    num_cycles = 0.475,\n    device = torch.device(\"cuda:0\") if torch.cuda.is_available() else \"cpu\",\n    seed = 2024,\n    log_wandb = True,\n    with_clip = False,\n)\n\nif CONFIG['log_wandb']:\n    import wandb\n    from kaggle_secrets import UserSecretsClient\n    user_secrets = UserSecretsClient()\n    secret_value_0 = user_secrets.get_secret(\"WANDB_API_KEY\")\n    wandb.login(key=secret_value_0)\n\nseeding(CONFIG['seed'])","metadata":{"execution":{"iopub.status.busy":"2024-08-10T09:05:45.441220Z","iopub.execute_input":"2024-08-10T09:05:45.441587Z","iopub.status.idle":"2024-08-10T09:05:48.547263Z","shell.execute_reply.started":"2024-08-10T09:05:45.441559Z","shell.execute_reply":"2024-08-10T09:05:48.545906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_PATH = Path(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification\")\n\ntrain_main = pd.read_csv(DATA_PATH/\"train.csv\")\ntrain_desc = pd.read_csv(DATA_PATH/\"train_series_descriptions.csv\")\ntrain_labels = pd.read_csv(DATA_PATH/\"train_label_coordinates.csv\")\n\ntrain_main = train_main.fillna(-100)\nlabel2id = {'Normal/Mild': 0, 'Moderate':1, 'Severe':2}\ntrain_main = train_main.replace(label2id)\n\n# df_train = train_main.merge(train_desc, on='study_id', how='inner')","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:39.330642Z","iopub.execute_input":"2024-08-10T08:48:39.331085Z","iopub.status.idle":"2024-08-10T08:48:39.569033Z","shell.execute_reply.started":"2024-08-10T08:48:39.331055Z","shell.execute_reply":"2024-08-10T08:48:39.567844Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TARGET_COLS = train_main.columns.tolist()[1:]\nAXIAL_COLS = {col:i for i, col in enumerate(train_main.columns[1:]) if 'subarticular_stenosis' in col}\nSAGT1_COLS = {col:i for i, col in enumerate(train_main.columns[1:]) if 'neural_foraminal_narrowing' in col}\nSAGT2_COLS = {col:i for i, col in enumerate(train_main.columns[1:]) if 'spinal_canal_stenosis' in col}\n\nCONDITIONS = [\n    'Spinal Canal Stenosis', \n    'Left Neural Foraminal Narrowing', \n    'Right Neural Foraminal Narrowing',\n    'Left Subarticular Stenosis',\n    'Right Subarticular Stenosis'\n]\n\nLEVELS = [\n    'L1/L2',\n    'L2/L3',\n    'L3/L4',\n    'L4/L5',\n    'L5/S1',\n]","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:39.570695Z","iopub.execute_input":"2024-08-10T08:48:39.571155Z","iopub.status.idle":"2024-08-10T08:48:39.580175Z","shell.execute_reply.started":"2024-08-10T08:48:39.571115Z","shell.execute_reply":"2024-08-10T08:48:39.579008Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Setting up Dataset and model\n\nCredit to [haqishen](https://www.kaggle.com/code/haqishen/rsna-2022-1st-place-solution-train-stage2-type1), most of the code is based on his winning solution","metadata":{}},{"cell_type":"markdown","source":"# Create dataset","metadata":{}},{"cell_type":"code","source":"class SpineDataset25(Dataset):\n    def __init__(self, df, mode='train', transform=None):\n        self.df = df\n        self.mode = mode\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        study_id = row.study_id\n        axial_label = self.df.loc[idx, AXIAL_COLS.keys()].values.astype(int)\n        coronal_label = self.df.loc[idx, SAGT1_COLS.keys()].values.astype(int)\n        sagittal_label = self.df.loc[idx, SAGT2_COLS.keys()].values.astype(int)\n        H, W = CONFIG['img_size'], CONFIG['img_size']\n        sagittal = np.zeros((10, H, W, 3), dtype=np.uint8)\n        coronal = np.zeros((10, H, W, 3), dtype=np.uint8)\n        axial = np.zeros((10, H, W, 3), dtype=np.uint8)\n#         SLIDES = np.zeros((CONFIG[\"n_slice_per_c\"], H, W, 3), dtype=np.uint8)\n        \n        # Sagittal\n        for i in range(10):\n            try:\n                img_path = f\"/kaggle/working/cvt_png/{study_id}/Sagittal T2_STIR/{i:03d}.png\"\n                img = PIL.Image.open(img_path).convert(\"RGB\")\n                img = np.array(img).astype(np.uint8)\n                img = self.transform(image=img)['image']\n                sagittal[i, ...] = img\n            except:\n                pass\n            \n        # coronal\n        for i in range(10):\n            try:\n                img_path = f\"/kaggle/working/cvt_png/{study_id}/Sagittal T1/{i:03d}.png\"\n                img = PIL.Image.open(img_path).convert(\"RGB\")\n                img = np.array(img).astype(np.uint8)\n                img = self.transform(image=img)['image']\n                coronal[i, ...] = img\n            except:\n                pass\n            \n        # Axial\n        for i in range(10):\n            try:\n                img_path = f\"/kaggle/working/cvt_png/{study_id}/Axial T2/{i:03d}.png\"\n                img = PIL.Image.open(img_path).convert(\"RGB\")\n                img = np.array(img).astype(np.uint8)\n                img = self.transform(image=img)['image']\n                axial[i, ...] = img\n            except:\n                pass\n        \n        axial = axial.transpose(0, 3, 1, 2).astype(np.float32) / 255.0 \n        coronal = coronal.transpose(0, 3, 1, 2).astype(np.float32) / 255.0 \n        sagittal = sagittal.transpose(0, 3, 1, 2).astype(np.float32) / 255.0 \n        \n        if self.mode != 'test':\n            axial = torch.tensor(axial).float()\n            coronal = torch.tensor(coronal).float()\n            sagittal = torch.tensor(sagittal).float()\n            axial_label = torch.tensor([axial_label] * CONFIG[\"n_slice_per_c\"]).float()\n            coronal_label = torch.tensor([coronal_label] * CONFIG[\"n_slice_per_c\"]).float()\n            sagittal_label = torch.tensor([sagittal_label] * CONFIG[\"n_slice_per_c\"]).float()\n            \n            if self.mode == 'train' and random.random() < CONFIG['p_rand_order_v1']:\n                axial_indices = torch.randperm(axial.size(0))\n                coronal_indices = torch.randperm(coronal.size(0))\n                sagittal_indices = torch.randperm(sagittal.size(0))\n                axial = axial[axial_indices]\n                coronal = coronal[coronal_indices]\n                sagittal = sagittal[sagittal_indices]\n            return {\"axial\": axial, \"coronal\": coronal, \"sagittal\": sagittal, \n                    \"axial_target\": axial_label, \"coronal_target\": coronal_label, \"sagittal_target\": sagittal_label}\n        \n        else:\n            return {\"axial\": torch.tensor(axial).float(), \n                    \"coronal\": torch.tensor(coronal).float(), \n                    \"sagittal\": torch.tensor(sagittal).float()}","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:39.582138Z","iopub.execute_input":"2024-08-10T08:48:39.582533Z","iopub.status.idle":"2024-08-10T08:48:39.608029Z","shell.execute_reply.started":"2024-08-10T08:48:39.582503Z","shell.execute_reply":"2024-08-10T08:48:39.606854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def get_transforms(height, width):\n    train_tsfm = A.Compose([\n        A.Resize(height=height, width=height),\n        A.Perspective(p=0.5),\n        A.HorizontalFlip(p=0.5),\n        A.VerticalFlip(p=0.5),\n        A.Rotate(-25, 25, p=0.5),\n        A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, border_mode=0, p=0.7),\n        \n        A.OneOf([\n            A.MotionBlur(blur_limit=3),\n            A.MedianBlur(blur_limit=3),\n            A.GaussianBlur(blur_limit=3),\n            A.GaussNoise(var_limit=(3.0, 9.0)),\n        ], p=0.5),\n        \n        A.OneOf([\n            A.OpticalDistortion(distort_limit=1.),\n            A.GridDistortion(num_steps=5, distort_limit=1.),\n            A.ElasticTransform(alpha=3),\n        ], p=0.5),\n        \n        A.CoarseDropout(max_holes=1, max_height=int(height * 0.275), max_width=int(width * 0.275), p=0.3),\n    ])\n    \n    valid_tsfm = A.Compose([\n        A.Resize(height=height, width=width),\n#         A.CenterCrop(height=height, width=width, p=1.0),\n    ])\n    return {\"train\": train_tsfm, \"eval\": valid_tsfm}\n\ndef get_dataloaders(data, cfg, split=\"train\"):\n    img_size = cfg['img_size']\n    height, width = img_size, img_size\n    tsfm = get_transforms(height=height, width=width)\n    if split == 'train':\n        tr_tsfm = tsfm['train']\n        ds = SpineDataset25(data, transform=tr_tsfm, mode=split)\n        dls = DataLoader(ds, \n                         batch_size=cfg['batch_size'], \n                         shuffle=True,\n                         num_workers=os.cpu_count(), \n                         drop_last=True, \n                         pin_memory=True)\n        \n    elif split == 'valid' or split == 'test':\n        eval_tsfm = tsfm['eval']\n        ds = SpineDataset25(data, transform=eval_tsfm, mode=split)\n        dls = DataLoader(ds, \n                         batch_size=2*cfg['batch_size'], \n                         shuffle=False, \n                         num_workers=os.cpu_count(), \n                         drop_last=False, \n                         pin_memory=True)\n    else:\n        raise Exception(\"Split should be 'train' or 'valid' or 'test'!!!\")\n    return dls","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:39.609839Z","iopub.execute_input":"2024-08-10T08:48:39.610506Z","iopub.status.idle":"2024-08-10T08:48:39.628242Z","shell.execute_reply.started":"2024-08-10T08:48:39.610465Z","shell.execute_reply":"2024-08-10T08:48:39.627175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls = get_dataloaders(train_main, CONFIG, split='train')","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:39.629819Z","iopub.execute_input":"2024-08-10T08:48:39.630737Z","iopub.status.idle":"2024-08-10T08:48:39.647047Z","shell.execute_reply.started":"2024-08-10T08:48:39.630694Z","shell.execute_reply":"2024-08-10T08:48:39.646036Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"b = next(iter(dls))\n# b['axial'].mean()\n\nk = 3\nfig, axes = plt.subplots(2, 5, figsize=(12, 12))\naxes = axes.flatten()\nsag = b['axial'][k].detach().cpu().numpy().transpose(0,2,3,1)\nfor i in range(10):\n    axes[i].imshow(sag[i, ...])\n    axes[i].axis(False)\nplt.tight_layout()\nplt.show()\n# sag[0].squeeze()","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:39.648478Z","iopub.execute_input":"2024-08-10T08:48:39.649474Z","iopub.status.idle":"2024-08-10T08:48:55.163398Z","shell.execute_reply.started":"2024-08-10T08:48:39.649434Z","shell.execute_reply":"2024-08-10T08:48:55.161879Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data splitting","metadata":{}},{"cell_type":"code","source":"from sklearn import model_selection\n\nkfold = model_selection.KFold(n_splits=5, shuffle=True, random_state=2024)\ndf = train_main.sample(frac=1.0, random_state=2024).reset_index(drop=True)\nx = df.index.values\ndf['fold'] = -1\nfor fold, (tr_idx, val_idx) in enumerate(kfold.split(x)):\n    df.loc[val_idx, 'fold'] = fold\n    \ndf.fold.value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:55.165143Z","iopub.execute_input":"2024-08-10T08:48:55.165506Z","iopub.status.idle":"2024-08-10T08:48:55.194905Z","shell.execute_reply.started":"2024-08-10T08:48:55.165474Z","shell.execute_reply":"2024-08-10T08:48:55.193704Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"def gem(x, p=3, eps=1e-4):\n    return F.avg_pool2d(x.clamp(min=eps), (x.size(-2), x.size(-1))).pow(1.0 / p)\n\nclass GeM(nn.Module):\n    def __init__(self, p=3, eps=1e-6, p_trainable=False):\n        super(GeM, self).__init__()\n        if p_trainable:\n            self.p = Parameter(torch.ones(1) * p)\n        else:\n            self.p = p\n        self.eps = eps\n\n\n    def forward(self, x):\n        ret = gem(x, p=self.p, eps=self.eps)\n        return ret","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:55.200051Z","iopub.execute_input":"2024-08-10T08:48:55.200423Z","iopub.status.idle":"2024-08-10T08:48:55.208926Z","shell.execute_reply.started":"2024-08-10T08:48:55.200391Z","shell.execute_reply":"2024-08-10T08:48:55.207757Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BaseModel(nn.Module):\n    def __init__(self, backbone, in_chans=3, pretrained=False, increase_stride=False):\n        super(BaseModel, self).__init__()\n\n        self.encoder = timm.create_model(\n            backbone,\n            in_chans=in_chans,\n            num_classes=0,\n            features_only=False,\n            drop_rate=CONFIG[\"drop_rate\"],\n            drop_path_rate=CONFIG[\"drop_path_rate\"],\n            pretrained=pretrained\n        )\n        self.encoder.name = backbone\n        self.nb_fts = self.encoder.num_features\n        self.gap = nn.AdaptiveAvgPool2d(1)\n#         self.gap = GeM(p_trainable=False)\n        \n        if increase_stride:\n            self.increase_stride()\n        \n    def increase_stride(self):\n        \"\"\"\n        Increase the stride of the first layer of the encoder\n        \"\"\"\n        if \"efficientnet\" in self.encoder.name:\n            self.encoder.conv_stem.stride = (4, 4)\n        elif \"nfnet\" in self.encoder.name:\n            self.encoder.stem.conv1.stride = (4, 4)\n        else:\n            raise NotImplementedError\n            \n    def forward(self, x):\n        x = self.encoder.forward_features(x)\n        x = self.gap(x)[:,:,0,0]\n        return x\n    \n    \nclass Clf(nn.Module):\n    def __init__(self, backbone, pretrained=False, increase_stride=False):\n        super(Clf, self).__init__()\n        self.axial_encoder = BaseModel(backbone=backbone, in_chans=3, \n                                       pretrained=pretrained, increase_stride=increase_stride)\n        self.coronal_encoder = BaseModel(backbone=backbone, in_chans=3, \n                                         pretrained=pretrained, increase_stride=increase_stride)\n        self.sagittal_encoder = BaseModel(backbone=backbone, in_chans=3, \n                                          pretrained=pretrained, increase_stride=increase_stride)\n        \n        self.in_chans = 3\n        self.nb_fts = 3*self.axial_encoder.nb_fts\n        self.lstm = nn.LSTM(self.nb_fts, 256, num_layers=2, dropout=CONFIG[\"drop_rate\"], bidirectional=True, batch_first=True)\n        self.axial_head = self.get_head(CONFIG['axial_classes'])\n        self.coronal_head = self.get_head(CONFIG['sagT1_classes'])\n        self.sagittal_head = self.get_head(CONFIG['sagT2_classes'])\n    \n    \n    def get_head(self, n_classes):\n        head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.Dropout(CONFIG[\"drop_rate_last\"]),\n            nn.LeakyReLU(0.1),\n            nn.Linear(256, n_classes),\n        )\n        return head\n    \n    \n    def extract_features(self, x, view='axial'):\n        bs = x.shape[0]\n        x = x.view(bs * CONFIG[\"n_slice_per_c\"], self.in_chans, CONFIG[\"img_size\"], CONFIG[\"img_size\"])\n        if view == 'axial':\n            feat = self.axial_encoder(x)\n        elif view == 'coronal':\n            feat = self.coronal_encoder(x)\n        elif view == 'sagittal':\n            feat = self.sagittal_encoder(x)\n        else:\n            raise NotImplementedError\n            \n        feat = feat.view(bs, CONFIG[\"n_slice_per_c\"], -1)\n        return feat\n        \n        \n    def forward(self, ax, cor, sag):\n        bs = ax.shape[0]\n        ax_fts = self.extract_features(ax, view='axial')\n        cor_fts = self.extract_features(cor, view='coronal')\n        sag_fts = self.extract_features(sag, view='sagittal')\n        fts = torch.concatenate([ax_fts, cor_fts, sag_fts], dim=-1)\n        fts, _ = self.lstm(fts)\n        \n        fts = fts.contiguous().view(bs * CONFIG[\"n_slice_per_c\"], -1)\n        ax_fts = self.axial_head(fts)\n        y_ax = ax_fts.view(bs, CONFIG[\"n_slice_per_c\"], CONFIG[\"axial_classes\"]).contiguous()\n        cor_fts = self.coronal_head(fts)\n        y_cor = cor_fts.view(bs, CONFIG[\"n_slice_per_c\"], CONFIG[\"sagT1_classes\"]).contiguous()\n        sag_fts = self.sagittal_head(fts)\n        y_sag = sag_fts.view(bs, CONFIG[\"n_slice_per_c\"], CONFIG[\"sagT2_classes\"]).contiguous()\n        return y_ax, y_cor, y_sag\n\n# gc.collect()\n# net = Clf(backbone=CONFIG['backbone'], pretrained=False)\n# net.eval()\n# o1, o2, o3 = net(b['axial'], b['coronal'], b['sagittal'])","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:55.210885Z","iopub.execute_input":"2024-08-10T08:48:55.211301Z","iopub.status.idle":"2024-08-10T08:48:55.238699Z","shell.execute_reply.started":"2024-08-10T08:48:55.211262Z","shell.execute_reply":"2024-08-10T08:48:55.237647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# o1.shape, o2.shape, o3.shape","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:55.240288Z","iopub.execute_input":"2024-08-10T08:48:55.241091Z","iopub.status.idle":"2024-08-10T08:48:55.254518Z","shell.execute_reply.started":"2024-08-10T08:48:55.241051Z","shell.execute_reply":"2024-08-10T08:48:55.253373Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from collections import Counter, defaultdict\n\nclass MetricMonitor:\n    def __init__(self, float_precision=4):\n        self.float_precision = float_precision\n        self.reset()\n\n    def reset(self):\n        self.metrics = defaultdict(lambda: {\"val\": 0, \"count\": 0, \"avg\": 0})\n\n    def update(self, metric_name, val):\n        metric = self.metrics[metric_name]\n\n        metric[\"val\"] += val\n        metric[\"count\"] += 1\n        metric[\"avg\"] = metric[\"val\"] / metric[\"count\"]\n\n    def __str__(self):\n        return \" | \".join(\n            [\n                \"{metric_name}: {avg:.{float_precision}f}\".format(\n                    metric_name=metric_name, avg=metric[\"avg\"], float_precision=self.float_precision\n                )\n                for (metric_name, metric) in self.metrics.items()\n            ]\n        )","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:55.255908Z","iopub.execute_input":"2024-08-10T08:48:55.256219Z","iopub.status.idle":"2024-08-10T08:48:55.267959Z","shell.execute_reply.started":"2024-08-10T08:48:55.256192Z","shell.execute_reply":"2024-08-10T08:48:55.266727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gc.collect()\n\ndef get_loss(logits, labels, criterion):\n    loss = criterion(logits.view(-1, 3), labels.view(-1).to(torch.int64))\n    return loss\n\n# get_loss(o3, b['sagittal_target'], criterion)\ndef shared_step(model, batch, criterion):\n    axial = batch['axial'].to(CONFIG[\"device\"], non_blocking=True)\n    coronal = batch['coronal'].to(CONFIG[\"device\"], non_blocking=True)\n    sagittal = batch['sagittal'].to(CONFIG[\"device\"], non_blocking=True)\n    axial_target = batch['axial_target'].to(CONFIG[\"device\"], non_blocking=True)\n    coronal_target = batch['coronal_target'].to(CONFIG[\"device\"], non_blocking=True)\n    sagittal_target = batch['sagittal_target'].to(CONFIG[\"device\"], non_blocking=True)\n    \n    axial_logits, coronal_logits, sagittal_logits = model(axial, coronal, sagittal)\n    axial_loss = get_loss(axial_logits, axial_target, criterion=criterion)\n    coronal_loss = get_loss(coronal_logits, coronal_target, criterion)\n    sagittal_loss = get_loss(sagittal_logits, sagittal_target, criterion)\n    loss = axial_loss + coronal_loss + sagittal_loss\n    return {\n        \"loss\": loss / 3\n    }\n\n# shared_step(net, b, criterion)\ndef train(train_loader, model, criterion, optimizer, epoch, scaler, scheduler=None):\n    metric_monitor = MetricMonitor()\n    model.train()\n    stream = tqdm(train_loader)\n    train_loss = 0\n    for i, batch in enumerate(stream, start=1):\n        optimizer.zero_grad(set_to_none=True)\n        \n        with torch.autocast(device_type='cuda', dtype=torch.float16):\n            outputs = shared_step(model, batch, criterion)\n            loss =  outputs['loss']\n        \n        metric_monitor.update(\"Loss\", loss)\n        train_loss += loss.detach().float()\n        CONFIG['example_ct'] += len(batch[\"axial\"])\n        scaler.scale(loss).backward()\n        \n        # clip the gradient\n        if CONFIG['with_clip']:\n            scaler.unscale_(optimizer)\n            nn.utils.clip_grad_value_(model.parameters(), clip_value=1.0)\n        \n        lr = optimizer.param_groups[0]['lr']\n        scaler.step(optimizer)\n        scaler.update()\n        \n        _train_metrics = {\n            \"train/step_loss\": loss,\n            \"train/epoch\": (i + 1 + CONFIG['n_steps_per_epoch'] * CONFIG['epochs']),\n            \"train/example_ct\": CONFIG['example_ct'],\n            \"lr\": lr,\n        }\n        \n        if CONFIG['log_wandb'] and (i+1 < CONFIG['n_steps_per_epoch']):\n            wandb.log(_train_metrics)\n        \n        CONFIG['step_ct'] += 1\n        if scheduler is not None:\n            scheduler.step()\n        \n        stream.set_description(\n            \"Epoch: {epoch}. Train.      {metric_monitor}\".format(epoch=epoch, metric_monitor=metric_monitor)\n        )\n        \n    total_train_loss = train_loss / len(train_loader)\n    _train_metrics['train/epoch_loss'] = total_train_loss\n    \n    flush()\n    return _train_metrics\n\n\ndef validate(val_loader, model, criterion, epoch):\n    metric_monitor = MetricMonitor()\n    model.eval()\n    stream = tqdm(val_loader)\n    valid_loss = 0\n    \n    with torch.no_grad():\n        for i, batch in enumerate(stream, start=1):\n            with torch.autocast(device_type='cuda', dtype=torch.float16):\n                outputs = shared_step(model, batch, criterion)\n                loss =  outputs['loss']\n\n            metric_monitor.update(\"Loss\", loss)\n            valid_loss += loss.detach().float()\n            _valid_metrics = {\n                    \"valid/step_loss\": loss,\n                }\n            \n#             if CONFIG['log_wandb']:\n#                 wandb.log(_valid_metrics)\n            \n            stream.set_description(\n                \"Epoch: {epoch}. Validation. {metric_monitor}\".format(epoch=epoch, metric_monitor=metric_monitor)\n            )\n            \n    total_valid_loss = valid_loss / len(val_loader)\n    _valid_metrics['valid/epoch_loss'] = total_valid_loss\n    flush()\n    return _valid_metrics","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:55.269722Z","iopub.execute_input":"2024-08-10T08:48:55.270132Z","iopub.status.idle":"2024-08-10T08:48:55.566850Z","shell.execute_reply.started":"2024-08-10T08:48:55.270102Z","shell.execute_reply":"2024-08-10T08:48:55.565613Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_and_validate(model, train_dataset, val_dataset, fold=0):\n    if CONFIG['log_wandb']:\n        run = wandb.init(\n            project=CONFIG[\"project_name\"],\n            resume=\"allow\",\n        )\n        artifact = wandb.Artifact(f\"{CONFIG['artifact_name']}_{fold}\", type=\"model\")\n    \n    if torch.cuda.is_available():\n        if torch.cuda.device_count() > 1:\n            DEVICE_IDS = list(range(torch.cuda.device_count()))\n            print(f\"\\nUsing {len(DEVICE_IDS)} GPUs to train ...\\n\")\n            model = nn.DataParallel(model, device_ids=DEVICE_IDS)\n            \n    model = model.to(CONFIG[\"device\"])\n#     model.apply(init_weights)\n    train_loader = get_dataloaders(train_dataset, CONFIG, split=\"train\")\n    valid_loader = get_dataloaders(val_dataset, CONFIG, split=\"valid\")\n    \n    n_steps_per_epoch = math.ceil(len(train_loader.dataset) / CONFIG['batch_size'])\n    CONFIG['n_steps_per_epoch'] = n_steps_per_epoch\n    CONFIG['example_ct'] = 0\n    CONFIG['step_ct'] = 0\n    \n    # weighted cross entropy loss\n    class_weights = torch.tensor([1, 2, 4], dtype=torch.float32)\n    criterion = nn.CrossEntropyLoss(weight=class_weights).to(CONFIG[\"device\"])\n \n    optimizer = torch.optim.AdamW(model.parameters(), lr=CONFIG[\"lr\"], weight_decay=CONFIG['wd'])\n    scaler = torch.cuda.amp.GradScaler()\n\n    scheduler = get_cosine_schedule_with_warmup(\n            optimizer,\n            num_warmup_steps=CONFIG[\"warmup\"] * CONFIG['n_steps_per_epoch'],\n            num_training_steps=CONFIG[\"epochs\"]* CONFIG['n_steps_per_epoch'],\n            num_cycles = CONFIG[\"num_cycles\"],\n        )\n    \n#     scheduler = get_cosine_with_hard_restarts_schedule_with_warmup(\n#         optimizer,\n#         num_warmup_steps=CONFIG[\"warmup\"] * CONFIG['n_steps_per_epoch'],\n#         num_training_steps=CONFIG[\"epochs\"]* CONFIG['n_steps_per_epoch'],\n#         num_cycles = 2,\n#     )\n    \n    best_metric = np.inf\n    loss_min = np.inf\n    es = 0\n    ES_RATIO = 0.3 if CONFIG[\"epochs\"] < 30 else 0.20\n    weights_file = \"rsna_2024_lumbar_spine_fold_{fold}_epoch_{epoch}.pth\"\n    for epoch in range(1, CONFIG[\"epochs\"] + 1):\n        _train_metrics = train(train_loader, model, criterion, optimizer, epoch, scaler, scheduler=scheduler)\n        _valid_metrics = validate(valid_loader, model, criterion, epoch)\n        \n        val_loss = _valid_metrics['valid/epoch_loss']\n        if CONFIG['log_wandb']:\n            wandb.log({**_train_metrics, **_valid_metrics})\n        \n        if val_loss < best_metric:\n            print(f\"Best metric: ({best_metric:.6f} --> {val_loss:.6f}). Saving model ...\")\n            if torch.cuda.device_count() > 2:\n                torch.save(model.module.state_dict(), weights_file.format(fold=fold, epoch=epoch))\n            else:\n                torch.save(model.state_dict(), weights_file.format(fold=fold, epoch=epoch))\n            best_metric = val_loss\n            if CONFIG['log_wandb']:\n                if epoch == 1:\n                    artifact.add_file(weights_file.format(fold=fold, epoch=epoch))\n                    run.log_artifact(artifact)\n                else:\n                    draft_artifact = wandb.Artifact(f\"{CONFIG['artifact_name']}_{fold}\", type=\"model\")\n                    draft_artifact.add_file(weights_file.format(fold=fold, epoch=epoch))\n                    run.log_artifact(draft_artifact)\n                \n            es = 0\n            \n        else:\n            es += 1\n            \n        if es > math.ceil(ES_RATIO*CONFIG[\"epochs\"]):\n            print(f\"Early stopping on epoch {epoch} ...\")\n            break\n    \n    if CONFIG['log_wandb']:\n        wandb.config = CONFIG\n        wandb.finish()\n        \n    del model, train_loader, valid_loader\n    flush()","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:55.568337Z","iopub.execute_input":"2024-08-10T08:48:55.568722Z","iopub.status.idle":"2024-08-10T08:48:55.591979Z","shell.execute_reply.started":"2024-08-10T08:48:55.568691Z","shell.execute_reply":"2024-08-10T08:48:55.590799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for fold in range(5):\n    model = Clf(backbone=CONFIG['backbone'], pretrained=True, increase_stride=False)\n    train_ds = df[df['fold'] != fold].reset_index(drop=True)\n    valid_ds = df[df['fold'] == fold].reset_index(drop=True)\n    train_and_validate(model, train_ds, valid_ds, fold=fold)\n    \n    break\ngc.collect()\nflush()","metadata":{"execution":{"iopub.status.busy":"2024-08-10T08:48:55.593416Z","iopub.execute_input":"2024-08-10T08:48:55.593839Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"```python\nclass TimmModel(nn.Module):\n    def __init__(self, backbone, pretrained=False):\n        super(TimmModel, self).__init__()\n\n        self.encoder = timm.create_model(\n            backbone,\n            in_chans=CONFIG[\"in_chans\"],\n            num_classes=CONFIG[\"axial_classes\"],\n            features_only=False,\n            drop_rate=CONFIG[\"drop_rate\"],\n            drop_path_rate=CONFIG[\"drop_path_rate\"],\n            pretrained=pretrained\n        )\n\n        if 'efficient' in backbone:\n            hdim = self.encoder.conv_head.out_channels\n            self.encoder.classifier = nn.Identity()\n        elif 'convnext' in backbone:\n            hdim = self.encoder.head.fc.in_features\n            self.encoder.head.fc = nn.Identity()\n\n\n        self.lstm = nn.LSTM(hdim, 256, num_layers=2, dropout=CONFIG[\"drop_rate\"], bidirectional=True, batch_first=True)\n        self.head = nn.Sequential(\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.Dropout(CONFIG[\"drop_rate_last\"]),\n            nn.LeakyReLU(0.1),\n            nn.Linear(256, CONFIG[\"n_classes\"]),\n        )\n\n    def forward(self, x):  # (bs, nslice, ch, sz, sz)\n        bs = x.shape[0]\n        x = x.view(bs * CONFIG[\"n_slice_per_c\"], CONFIG[\"in_chans\"], CONFIG[\"img_size\"], CONFIG[\"img_size\"])\n        feat = self.encoder(x)\n        feat = feat.view(bs, CONFIG[\"n_slice_per_c\"], -1)\n        feat, _ = self.lstm(feat)\n        feat = feat.contiguous().view(bs * CONFIG[\"n_slice_per_c\"], -1)\n        feat = self.head(feat)\n        feat = feat.view(bs, CONFIG[\"n_slice_per_c\"], CONFIG[\"n_classes\"]).contiguous()\n\n        return feat\n\nnet = TimmModel(backbone=CONFIG['backbone'], pretrained=False)\nnet.eval()\n\noutputs = net(b['axial'])\n```","metadata":{}},{"cell_type":"markdown","source":"```python\nclass SpineLightningModel(pl.LightningModule):\n    def __init__(self, backbone, pretrained=False):\n        super().__init__()\n        self.model = TimmModel(backbone, pretrained=pretrained)\n        weights = torch.tensor([1.0, 2.0, 4.0])\n        self.loss_fn = nn.CrossEntropyLoss(weight=weights)\n    \n    def forward(self, images):\n        return self.model(images)\n    \n    \n    def shared_step(self, batch, stage='train'):\n        images, labels = batch['slides'], batch['target']\n        logits = self.forward(images)\n        logits = logits.amax(dim=1)\n        loss = 0\n        n_labels = labels.shape[-1]\n        for col in range(n_labels):\n            pred = logits[:, col*3:col*3+3]\n            target = labels[:,col]\n            loss = loss + self.loss_fn(pred, target.to(torch.int64)) / n_labels\n            \n        self.log(f\"{stage}_loss\", loss, on_step=True, on_epoch=True, prog_bar=True)\n        \n        outputs = {\n            \"loss\": loss\n        }\n        return outputs\n    \n    \n    def training_step(self, batch, batch_idx):\n        loss = self.shared_step(batch, \"train\")\n        return loss\n    \n    def validation_step(self, batch, batch_idx):\n        loss = self.shared_step(batch, \"valid\")\n        return loss\n    \n    def configure_optimizers(self):\n        optimizer = torch.optim.Adam(self.parameters(), lr=CONFIG['lr'])\n        after_scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=CONFIG[\"epochs\"], eta_min=0)\n        scheduler = GradualWarmupScheduler(optimizer, multiplier=1, \n                                           total_epoch=math.ceil(CONFIG[\"warmup_ratio\"]*CONFIG['epochs']), \n                                           after_scheduler=after_scheduler)\n        return [optimizer], [scheduler]\n```","metadata":{}},{"cell_type":"markdown","source":"```python\nfor fold in range(5):\n    train_ds = df[df['fold'] != fold].reset_index(drop=True)\n    valid_ds = df[df['fold'] == fold].reset_index(drop=True)\n    \n    train_loader = get_dataloaders(train_ds, CONFIG, split=\"train\")\n    valid_loader = get_dataloaders(valid_ds, CONFIG, split=\"valid\")\n    \n    if CONFIG['log_wandb']:\n        wandb_logger = WandbLogger(\n            project=CONFIG[\"project_name\"],\n            checkpoint_name=f'{CONFIG[\"artifact_name\"]}_{fold}',\n            log_model=\"all\",\n        )\n        \n    logger = wandb_logger if CONFIG['log_wandb'] else None\n\n    callbacks = [\n        ModelCheckpoint(save_weights_only=True, \n                        mode=\"min\", \n                        monitor=\"valid_loss\"),\n        LearningRateMonitor(\"epoch\"),\n    ]\n    \n    net = SpineLightningModel(backbone=CONFIG['backbone'], pretrained=True)\n    \n    trainer = pl.Trainer(accelerator=\"gpu\", devices=1, \n                         precision=\"16-mixed\",\n#                          gradient_clip_val=0.5, gradient_clip_algorithm=\"value\", \n                         max_epochs=CONFIG['epochs'], \n                         logger=logger, callbacks=callbacks, default_root_dir=os.getcwd())\n    \n    trainer.fit(net, train_dataloaders=train_loader, val_dataloaders=valid_loader)\n    break\n    \n    \nif CONFIG['log_wandb']:\n    wandb.finish()\n```","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}