{"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"}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Training\n\nInference notebook is here: [rsna-lumbar-inference](https://www.kaggle.com/code/samu2505/rsna-lumbar-inference-lb-0-84-cv-0-54)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"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\"] = \"0\"","metadata":{"execution":{"iopub.status.busy":"2024-08-12T04:52:39.286359Z","iopub.execute_input":"2024-08-12T04:52:39.287207Z","iopub.status.idle":"2024-08-12T04:52:49.133214Z","shell.execute_reply.started":"2024-08-12T04:52:39.287168Z","shell.execute_reply":"2024-08-12T04:52:49.132181Z"},"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-12T04:52:49.135346Z","iopub.execute_input":"2024-08-12T04:52:49.135704Z","iopub.status.idle":"2024-08-12T04:52:49.143319Z","shell.execute_reply.started":"2024-08-12T04:52:49.135672Z","shell.execute_reply":"2024-08-12T04:52:49.142287Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONFIG = dict(\n    project_name = \"RSNA-2024-Baseline\",\n    artifact_name = \"rsnaEffNetModel\",\n    load_kernel = None,\n    load_last = True,\n    n_folds = 5,\n    backbone = \"tf_efficientnet_b0.ns_jft_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-3,\n    wd = 1e-4,\n\n    epochs = 5,\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-12T04:53:37.427596Z","iopub.execute_input":"2024-08-12T04:53:37.428183Z","iopub.status.idle":"2024-08-12T04:53:43.190912Z","shell.execute_reply.started":"2024-08-12T04:53:37.428109Z","shell.execute_reply":"2024-08-12T04:53:43.189674Z"},"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)\n\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-12T04:52:49.169507Z","iopub.execute_input":"2024-08-12T04:52:49.169938Z","iopub.status.idle":"2024-08-12T04:52:49.410686Z","shell.execute_reply.started":"2024-08-12T04:52:49.169901Z","shell.execute_reply":"2024-08-12T04:52:49.409667Z"},"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-12T04:52:49.413515Z","iopub.execute_input":"2024-08-12T04:52:49.413892Z","iopub.status.idle":"2024-08-12T04:52:49.422037Z","shell.execute_reply.started":"2024-08-12T04:52:49.413862Z","shell.execute_reply":"2024-08-12T04:52:49.420748Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom(path):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    if dicom.PhotometricInterpretation == \"MONOCHROME1\":\n        data = data - np.min(data)\n        \n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-08-12T04:52:49.423673Z","iopub.execute_input":"2024-08-12T04:52:49.424555Z","iopub.status.idle":"2024-08-12T04:52:49.440579Z","shell.execute_reply.started":"2024-08-12T04:52:49.424496Z","shell.execute_reply":"2024-08-12T04:52:49.439602Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class SpineDataset(Dataset):\n    def __init__(self, data, desc, mode='train', transform=None):\n        self.data = data\n        self.desc = desc\n        self.mode = mode\n        self.transform = transform\n        \n    \n    def __len__(self):\n        return len(self.data)\n    \n    def get_img_paths(self, study_id, series_desc):\n        pdf = self.desc[self.desc['study_id'] == study_id]\n        pdf_ = pdf[pdf['series_description'] == series_desc]\n        allimgs = []\n        for i, row in pdf_.iterrows():\n            pimgs = glob.glob(f\"{str(DATA_PATH)}/train_images/{study_id}/{row['series_id']}/*.dcm\")\n            pimgs = sorted(pimgs, key=lambda p: int(os.path.basename(p).split('.')[0]))\n            allimgs.extend(pimgs)\n        return allimgs\n    \n    def read_dcm(self, src_path):\n        img = load_dicom(src_path)\n        return img\n    \n    def get_images(self, nslides, image_paths):\n        H, W = CONFIG['img_size'], CONFIG['img_size']\n        IMAGES = np.zeros((H, W, nslides), dtype=np.uint8)\n        for i in range(nslides):\n            try:\n                img = self.read_dcm(image_paths[i])\n                img = cv2.resize(img, (H,W)).astype(np.uint8)\n                IMAGES[..., i] = img\n            except:\n                pass\n            \n        return IMAGES\n    \n    def __getitem__(self, idx):\n        row = self.data.iloc[idx]\n        study_id = row.study_id\n        labels = self.data.loc[idx, TARGET_COLS].values.astype(int)\n        \n        H, W = CONFIG['img_size'], CONFIG['img_size']\n        sagittal = np.zeros((H, W, 15), dtype=np.uint8)\n        coronal = np.zeros((H, W, 15), dtype=np.uint8)\n        axial = np.zeros((H, W, 15), dtype=np.uint8)\n        \n        # Sagittal\n        allimgs_sag = self.get_img_paths(study_id, 'Sagittal T2/STIR')\n        \n        if len(allimgs_sag)==0:\n            pass\n        \n        else:\n            sagittal = self.get_images(nslides=15, image_paths=allimgs_sag)\n            \n        # coronal\n        allimgs_cor = self.get_img_paths(study_id, 'Sagittal T1')\n        if len(allimgs_cor)==0:\n            pass\n        \n        else:\n            coronal = self.get_images(nslides=15, image_paths=allimgs_cor)\n                \n        # Axial\n        allimgs_ax = self.get_img_paths(study_id, 'Axial T2')\n        if len(allimgs_ax)==0:\n            pass\n        \n        else:\n            axial = self.get_images(nslides=15, image_paths=allimgs_ax)\n        \n        axial = self.transform(image=axial)['image']\n        axial = axial.transpose(2,0,1).astype(np.float32) / 255.0 \n        axial = torch.tensor(axial).float()\n        \n        coronal = self.transform(image=coronal)['image']\n        coronal = coronal.transpose(2,0,1).astype(np.float32) / 255.0 \n        coronal = torch.tensor(coronal).float()\n        \n        sagittal = self.transform(image=sagittal)['image']\n        sagittal = sagittal.transpose(2,0,1).astype(np.float32) / 255.0 \n        sagittal = torch.tensor(sagittal).float()\n        \n        if self.mode != 'test':\n            return {\"axial\": axial, \"coronal\": coronal, \"sagittal\": sagittal, \"target\": torch.tensor(labels)}\n        \n        else:\n            return {\"axial\": axial, \"coronal\": coronal, \"sagittal\": sagittal}\n    \n    \n# ts = A.Compose([\n#     A.Resize(height=224, width=224),\n# ])\n\n# ds = SpineDataset(df_train, transform=ts)\n# dls = DataLoader(ds, batch_size=4, shuffle=True, num_workers=os.cpu_count(), drop_last=True)\n# b = next(iter(dls))","metadata":{"execution":{"iopub.status.busy":"2024-08-12T04:52:49.442084Z","iopub.execute_input":"2024-08-12T04:52:49.442484Z","iopub.status.idle":"2024-08-12T04:52:49.466827Z","shell.execute_reply.started":"2024-08-12T04:52:49.442454Z","shell.execute_reply":"2024-08-12T04:52:49.465727Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# b['axial'].shape","metadata":{"execution":{"iopub.status.busy":"2024-08-12T04:52:49.468021Z","iopub.execute_input":"2024-08-12T04:52:49.468370Z","iopub.status.idle":"2024-08-12T04:52:49.489377Z","shell.execute_reply.started":"2024-08-12T04:52:49.468342Z","shell.execute_reply":"2024-08-12T04:52:49.488295Z"},"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        \n#         A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, border_mode=4, p=0.7),\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=2, 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, desc, 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    \n    if split == 'train':\n        tr_tsfm = tsfm['train']\n        ds = SpineDataset(data=data, desc=desc, mode='train', transform=tr_tsfm)\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 = SpineDataset(data=data, desc=desc, mode='valid', transform=eval_tsfm)\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-12T04:52:49.490624Z","iopub.execute_input":"2024-08-12T04:52:49.490992Z","iopub.status.idle":"2024-08-12T04:52:49.512137Z","shell.execute_reply.started":"2024-08-12T04:52:49.490963Z","shell.execute_reply":"2024-08-12T04:52:49.510972Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# dls = get_dataloaders(data=train_main, desc=train_desc, cfg=CONFIG, split='train')\n# b = next(iter(dls))","metadata":{"execution":{"iopub.status.busy":"2024-08-12T04:52:49.513488Z","iopub.execute_input":"2024-08-12T04:52:49.513861Z","iopub.status.idle":"2024-08-12T04:52:49.536051Z","shell.execute_reply.started":"2024-08-12T04:52:49.513832Z","shell.execute_reply":"2024-08-12T04:52:49.534880Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# for b in tqdm(dls, total=len(dls)):\n#     pass","metadata":{"execution":{"iopub.status.busy":"2024-08-12T04:52:49.537387Z","iopub.execute_input":"2024-08-12T04:52:49.537733Z","iopub.status.idle":"2024-08-12T04:52:49.554483Z","shell.execute_reply.started":"2024-08-12T04:52:49.537704Z","shell.execute_reply":"2024-08-12T04:52:49.553378Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# k = 2\n# fig, axes = plt.subplots(1, 5, figsize=(12, 12))\n# axes = axes.flatten()\n# sag = b['coronal'][k].detach().cpu().numpy().transpose(1,2,0)\n# IMAGE = np.zeros((5, 224, 224, 3)).astype(np.float32)\n\n# for i in range(5):\n#     IMAGE[i, ...] = sag[..., i*3:i*3+3]\n#     axes[i].imshow(IMAGE[i, ...])\n#     axes[i].axis(False)\n# plt.tight_layout()\n# plt.show()\n# # sag[0].squeeze()","metadata":{"execution":{"iopub.status.busy":"2024-08-12T04:52:49.556130Z","iopub.execute_input":"2024-08-12T04:52:49.556472Z","iopub.status.idle":"2024-08-12T04:52:49.570128Z","shell.execute_reply.started":"2024-08-12T04:52:49.556443Z","shell.execute_reply":"2024-08-12T04:52:49.569075Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Split data","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-12T04:52:49.571622Z","iopub.execute_input":"2024-08-12T04:52:49.572008Z","iopub.status.idle":"2024-08-12T04:52:49.619299Z","shell.execute_reply.started":"2024-08-12T04:52:49.571978Z","shell.execute_reply":"2024-08-12T04:52:49.618125Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class BaseModel(nn.Module):\n    def __init__(self, backbone, in_chans=15, pretrained=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            \n    def forward(self, x):\n        x = self.encoder.forward_features(x)\n        x = self.gap(x)[:,:,0,0]\n        return x\n    \nclass Clf(nn.Module):\n    def __init__(self, backbone, pretrained=False, in_chans=15):\n        super(Clf, self).__init__()\n        self.axial_encoder = BaseModel(backbone=backbone, in_chans=in_chans, pretrained=pretrained)\n        self.coronal_encoder = BaseModel(backbone=backbone, in_chans=in_chans, pretrained=pretrained)\n        self.sagittal_encoder = BaseModel(backbone=backbone, in_chans=in_chans, pretrained=pretrained)\n        \n        self.in_chans = in_chans\n        self.out_chans = 3 * 25 \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.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, self.out_chans),\n        )\n    \n    def forward(self, axial, coronal, sagittal):\n        bs = axial.shape[0]\n        ax_fts = self.axial_encoder(axial)\n        cor_fts = self.coronal_encoder(coronal)\n        sag_fts = self.sagittal_encoder(sagittal)\n        fts = torch.concatenate([ax_fts, cor_fts, sag_fts], dim=-1)\n        fts, _ = self.lstm(fts)\n        fts = self.head(fts)\n        fts = fts.reshape(bs, 3, 25)\n        return fts\n    \n# net = Clf(backbone=CONFIG['backbone'])\n# net.eval()\n# out = net(b['axial'], b['coronal'], b['sagittal'])","metadata":{"execution":{"iopub.status.busy":"2024-08-12T04:52:49.623059Z","iopub.execute_input":"2024-08-12T04:52:49.623422Z","iopub.status.idle":"2024-08-12T04:52:49.641448Z","shell.execute_reply.started":"2024-08-12T04:52:49.623393Z","shell.execute_reply":"2024-08-12T04:52:49.640286Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Severe loss","metadata":{}},{"cell_type":"code","source":"import torch.nn.functional as F\nfrom torch.nn.modules.loss import _Loss\n\nclass SevereLoss(_Loss):\n    \"\"\"\n    For Kaggle RSNA 2024\n    criterion = SevereLoss()\n    loss = criterion(y_pred, y)\n    \"\"\"\n    def __init__(self, temperature=1.0):\n        \"\"\"\n        Use max if temperature = 0\n        \"\"\"\n        super().__init__()\n        self.t = temperature\n        assert self.t >= 0\n\n    def __repr__(self):\n        return 'SevereLoss(t=%.1f)' % self.t\n\n    def forward(self, y_pred: torch.Tensor, y: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Args:\n          y_pred (Tensor[float]): logit             (batch_size, 3, 25)\n          y      (Tensor[int]):   true label index  (batch_size, 25)\n        \"\"\"\n        assert y_pred.size(0) == y.size(0)\n        assert y_pred.size(1) == 3 and y_pred.size(2) == 25\n        assert y.size(1) == 25\n        assert y.size(0) > 0\n\n        slices = [slice(0, 5), slice(5, 15), slice(15, 25)]\n        w = 2 ** y  # sample_weight w = (1, 2, 4) for y = 0, 1, 2 (batch_size, 25)\n\n        loss = F.cross_entropy(y_pred, y, reduction='none')  # (batch_size, 25)\n\n        # Weighted sum of losses for spinal (:5), foraminal (5:15), and subarticular (15:25)\n        wloss_sums = []\n        for k, idx in enumerate(slices):\n            wloss_sums.append((w[:, idx] * loss[:, idx]).sum())\n\n        # Spinal max\n        y_spinal_prob = y_pred[:, :, :5].softmax(dim=1)             # (batch_size, 3,  5)\n        w_max = torch.amax(w[:, :5], dim=1)                         # batch_size\n        #y_max = torch.amax(y[:, :5] == 2, dim=1).to(torch.float32)  # 0 or 1\n        y_max = torch.amax(y[:, :5] == 2, dim=1).to(y_pred.dtype)\n\n        if self.t > 0:\n            # Attention for the maximum value\n            attn = F.softmax(y_spinal_prob[:, 2, :] / self.t, dim=1)         # (batch_size, 5)\n\n            # Approximately the max among 5 y_spinal_probs\n            y_pred_max = (attn * y_spinal_prob[:, 2, :]).sum(dim=1)     # weighted average among 5 spinal columns,\n        else:\n            # Exact max; this works too\n            y_pred_max = y_spinal_prob[:, 2, :].amax(dim=1)\n\n        loss_max = F.binary_cross_entropy(y_pred_max, y_max, reduction='none')\n        wloss_sums.append((w_max * loss_max).sum())\n\n        # See `compute_global_normalization` for the numbers\n        loss = (wloss_sums[0] / 6.084050632911392 +\n                wloss_sums[1] / 12.962531645569621 +\n                wloss_sums[2] / 14.38632911392405 +\n                wloss_sums[3] / 1.729113924050633) / (4 * y.size(0))\n\n        return loss\n\n\ndef compute_global_normalization(train):\n    # Compute the weight global average\n    weight_map = {'Normal/Mild': 1,\n                'Moderate': 2,\n                'Severe': 4,\n                None: 0}\n\n    w_sum = [0, ] * 4\n\n    for r in train.iter_rows():\n        w = np.array([weight_map[x] for x in r[1:]])  # array[int] (25, )\n        assert len(w) == 25\n\n        w_sum[0] += w[:5].sum()    # spinal\n        w_sum[1] += w[5:15].sum()  # foraminal\n        w_sum[2] += w[15:25].sum() # subarticular\n        w_sum[3] += w[:5].max()    # any_severe_spinal\n\n    for k in range(4):\n        w_sum[k] /= len(train)\n\n    # (6.084050632911392, 12.962531645569621, 14.38632911392405, 1.729113924050633)\n    return w_sum","metadata":{"execution":{"iopub.status.busy":"2024-08-12T04:52:49.643296Z","iopub.execute_input":"2024-08-12T04:52:49.643761Z","iopub.status.idle":"2024-08-12T04:52:49.667689Z","shell.execute_reply.started":"2024-08-12T04:52:49.643724Z","shell.execute_reply":"2024-08-12T04:52:49.666503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training with Pytorch","metadata":{}},{"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-12T04:52:49.669353Z","iopub.execute_input":"2024-08-12T04:52:49.669713Z","iopub.status.idle":"2024-08-12T04:52:49.692679Z","shell.execute_reply.started":"2024-08-12T04:52:49.669682Z","shell.execute_reply":"2024-08-12T04:52:49.691482Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def init_weights(m):\n    if type(m) == nn.Linear:\n        torch.nn.init.xavier_uniform(m.weight)\n        m.bias.data.fill_(0.01)\n        \n\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    target = batch['target'].to(CONFIG[\"device\"], non_blocking=True)\n    \n    logits = model(axial, coronal, sagittal)\n#     loss = criterion(logits.view(-1, CONFIG[\"out_dim\"]), target.view(-1).to(torch.int64))\n    loss = criterion(logits, target)\n\n    return {\n        \"loss\": loss,\n        \"logits\": logits,\n        \"target\": target\n    }\n\n\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['target'])\n        # backward pass, with gradient scaling\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    n_sum = 0\n    loss0_sum = 0.0  # loss for the criterion\n    \n    # 4 losses for the evaluation metric\n    loss4_sum = torch.zeros(4, device=CONFIG['device'])\n    w_sum = torch.zeros(4, device=CONFIG['device'])\n    slices = [slice(0, 5), slice(5, 15), slice(15, 25)]  # spinal, foraminal, subarticular\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            loss0 =  outputs['loss']\n            y_pred = outputs['logits']\n            y = outputs['target']\n\n            w = 2 ** y  # sample_weight w = (1, 2, 4) for y = 0, 1, 2 (batch_size, 25)\n            bs = len(batch['target'])\n            n_sum += bs\n            loss0_sum += loss0.item() * bs\n            \n            # Compute score\n            # - weighted loss for spinal, foraminal, subarticular\n            # - binary cross entropy for maximum spinal severe\n            ce_loss = F.cross_entropy(y_pred, y, reduction='none')  # (batch_size, 25)\n            for k, idx in enumerate(slices):\n                w_sum[k] += w[:, idx].sum()\n                loss4_sum[k] += (w[:, idx] * ce_loss[:, idx]).sum()\n                \n            # Spinal max\n            y_spinal_prob = y_pred[:, :, :5].softmax(dim=1)            # (batch_size, 3,  5)\n            w_max = torch.amax(w[:, :5], dim=1)                        # (batch_size, )\n            y_max = torch.amax(y[:, :5] == 2, dim=1).to(torch.float)   # 0 or 1\n            y_pred_max = y_spinal_prob[:, 2, :].amax(dim=1)            # max in severe (class=2)\n\n            loss_max = F.binary_cross_entropy(y_pred_max, y_max, reduction='none')\n            loss4_sum[3] += (w_max * loss_max).sum()\n            w_sum[3] += w_max.sum()\n            \n            metric_monitor.update(\"Loss\", loss0)\n#             valid_loss += loss.detach().float()\n            \n            _valid_metrics = {\n#                     \"valid/step_loss\": loss,\n                }\n            \n            stream.set_description(\n                \"Epoch: {epoch}. Validation. {metric_monitor}\".format(epoch=epoch, metric_monitor=metric_monitor)\n            )\n    # Average over spinal, foraminal, subarticular, and any_severe_spinal\n    score = (loss4_sum / w_sum).sum().item() / 4\n    \n#     total_valid_loss = valid_loss / len(val_loader)\n    _valid_metrics['valid/epoch_loss'] = loss0_sum / n_sum\n    _valid_metrics['valid/score'] = score\n    flush()\n    return _valid_metrics","metadata":{"execution":{"iopub.status.busy":"2024-08-12T04:52:49.694272Z","iopub.execute_input":"2024-08-12T04:52:49.694663Z","iopub.status.idle":"2024-08-12T04:52:49.725571Z","shell.execute_reply.started":"2024-08-12T04:52:49.694632Z","shell.execute_reply":"2024-08-12T04:52:49.724603Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_and_validate(model, train_dataset, val_dataset, desc, 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(data=train_dataset, desc=desc, cfg=CONFIG, split=\"train\")\n    valid_loader = get_dataloaders(data=val_dataset, desc=desc, cfg=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    criterion = SevereLoss(temperature=0).to(CONFIG['device'])\n \n    optimizer = torch.optim.AdamW(model.parameters(), lr=CONFIG[\"lr\"])\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    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        score = _valid_metrics['valid/score']\n#         val_loss = score\n        print(f\"score: {score:.6f}\")\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-12T04:52:49.726777Z","iopub.execute_input":"2024-08-12T04:52:49.727100Z","iopub.status.idle":"2024-08-12T04:52:49.752738Z","shell.execute_reply.started":"2024-08-12T04:52:49.727069Z","shell.execute_reply":"2024-08-12T04:52:49.751563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for fold in range(5):\n    model = Clf(backbone=CONFIG[\"backbone\"], pretrained=True)\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, train_desc, fold=fold)\n    \n    break\ngc.collect()\nflush()","metadata":{"execution":{"iopub.status.busy":"2024-08-12T04:53:20.005025Z","iopub.execute_input":"2024-08-12T04:53:20.005506Z","iopub.status.idle":"2024-08-12T04:53:26.390406Z","shell.execute_reply.started":"2024-08-12T04:53:20.005464Z","shell.execute_reply":"2024-08-12T04:53:26.373928Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}