{"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":30699,"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-11T11:57:10.219548Z","iopub.execute_input":"2024-08-11T11:57:10.220117Z","iopub.status.idle":"2024-08-11T12:01:03.467328Z","shell.execute_reply.started":"2024-08-11T11:57:10.220086Z","shell.execute_reply":"2024-08-11T12:01:03.466113Z"},"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.25075Z","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-11T12:01:03.469682Z","iopub.execute_input":"2024-08-11T12:01:03.470489Z","iopub.status.idle":"2024-08-11T12:01:13.533918Z","shell.execute_reply.started":"2024-08-11T12:01:03.470446Z","shell.execute_reply":"2024-08-11T12:01:13.532916Z"},"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-11T12:01:13.535336Z","iopub.execute_input":"2024-08-11T12:01:13.535802Z","iopub.status.idle":"2024-08-11T12:01:13.540055Z","shell.execute_reply.started":"2024-08-11T12:01:13.535745Z","shell.execute_reply":"2024-08-11T12:01:13.539145Z"},"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-11T12:01:13.542522Z","iopub.execute_input":"2024-08-11T12:01:13.542960Z","iopub.status.idle":"2024-08-11T12:01:13.553164Z","shell.execute_reply.started":"2024-08-11T12:01:13.542926Z","shell.execute_reply":"2024-08-11T12:01:13.552351Z"},"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 = False,\n    with_clip = False,\n)\n\nseeding(CONFIG['seed'])","metadata":{"execution":{"iopub.status.busy":"2024-08-11T12:02:11.465970Z","iopub.execute_input":"2024-08-11T12:02:11.466354Z","iopub.status.idle":"2024-08-11T12:02:11.478271Z","shell.execute_reply.started":"2024-08-11T12:02:11.466325Z","shell.execute_reply":"2024-08-11T12:02:11.477449Z"},"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\")\nlist_ids=train_main['study_id'].unique()\nselected_ids=np.random.choice(list_ids,20)\ntrain_main = train_main[~train_main['study_id'].isin(selected_ids)]\ntrain_desc = train_desc[~train_desc['study_id'].isin(selected_ids)]\ntrain_labels = train_labels[~train_labels['study_id'].isin(selected_ids)]\n\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-11T12:04:22.469293Z","iopub.execute_input":"2024-08-11T12:04:22.470209Z","iopub.status.idle":"2024-08-11T12:04:22.619428Z","shell.execute_reply.started":"2024-08-11T12:04:22.470174Z","shell.execute_reply":"2024-08-11T12:04:22.618364Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_labels_true = pd.read_csv(DATA_PATH/\"train_label_coordinates.csv\")\ntest_labels_true = test_labels_true[test_labels_true['study_id'].isin(selected_ids)]\n","metadata":{"execution":{"iopub.status.busy":"2024-08-11T14:52:34.499399Z","iopub.execute_input":"2024-08-11T14:52:34.500384Z","iopub.status.idle":"2024-08-11T14:52:34.583480Z","shell.execute_reply.started":"2024-08-11T14:52:34.500333Z","shell.execute_reply":"2024-08-11T14:52:34.582511Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_labels_true","metadata":{"execution":{"iopub.status.busy":"2024-08-11T14:52:40.190530Z","iopub.execute_input":"2024-08-11T14:52:40.191248Z","iopub.status.idle":"2024-08-11T14:52:40.206793Z","shell.execute_reply.started":"2024-08-11T14:52:40.191212Z","shell.execute_reply":"2024-08-11T14:52:40.205620Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"list_ids=train_main['study_id'].unique()\nselected_ids=np.random.choice(list_ids,20)","metadata":{"execution":{"iopub.status.busy":"2024-08-11T12:03:28.390908Z","iopub.execute_input":"2024-08-11T12:03:28.391781Z","iopub.status.idle":"2024-08-11T12:03:28.396545Z","shell.execute_reply.started":"2024-08-11T12:03:28.391746Z","shell.execute_reply":"2024-08-11T12:03:28.395614Z"},"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-11T12:07:28.398270Z","iopub.execute_input":"2024-08-11T12:07:28.398663Z","iopub.status.idle":"2024-08-11T12:07:28.405724Z","shell.execute_reply.started":"2024-08-11T12:07:28.398631Z","shell.execute_reply":"2024-08-11T12:07:28.404649Z"},"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-11T12:07:35.680537Z","iopub.execute_input":"2024-08-11T12:07:35.680955Z","iopub.status.idle":"2024-08-11T12:07:35.701303Z","shell.execute_reply.started":"2024-08-11T12:07:35.680925Z","shell.execute_reply":"2024-08-11T12:07:35.700369Z"},"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-11T12:07:36.926496Z","iopub.execute_input":"2024-08-11T12:07:36.926877Z","iopub.status.idle":"2024-08-11T12:07:36.939486Z","shell.execute_reply.started":"2024-08-11T12:07:36.926848Z","shell.execute_reply":"2024-08-11T12:07:36.938549Z"},"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-11T12:07:37.988387Z","iopub.execute_input":"2024-08-11T12:07:37.989071Z","iopub.status.idle":"2024-08-11T12:07:37.994791Z","shell.execute_reply.started":"2024-08-11T12:07:37.989037Z","shell.execute_reply":"2024-08-11T12:07:37.993756Z"},"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-11T12:07:40.838061Z","iopub.execute_input":"2024-08-11T12:07:40.838781Z","iopub.status.idle":"2024-08-11T12:07:54.557427Z","shell.execute_reply.started":"2024-08-11T12:07:40.838747Z","shell.execute_reply":"2024-08-11T12:07:54.556450Z"},"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-11T12:08:02.382440Z","iopub.execute_input":"2024-08-11T12:08:02.382870Z","iopub.status.idle":"2024-08-11T12:08:02.404171Z","shell.execute_reply.started":"2024-08-11T12:08:02.382833Z","shell.execute_reply":"2024-08-11T12:08:02.403259Z"},"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-11T12:08:07.477062Z","iopub.execute_input":"2024-08-11T12:08:07.477698Z","iopub.status.idle":"2024-08-11T12:08:07.485061Z","shell.execute_reply.started":"2024-08-11T12:08:07.477663Z","shell.execute_reply":"2024-08-11T12:08:07.484052Z"},"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-11T12:08:26.258612Z","iopub.execute_input":"2024-08-11T12:08:26.258962Z","iopub.status.idle":"2024-08-11T12:08:26.281551Z","shell.execute_reply.started":"2024-08-11T12:08:26.258936Z","shell.execute_reply":"2024-08-11T12:08:26.280544Z"},"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-11T12:08:29.501027Z","iopub.execute_input":"2024-08-11T12:08:29.502044Z","iopub.status.idle":"2024-08-11T12:08:29.509945Z","shell.execute_reply.started":"2024-08-11T12:08:29.502008Z","shell.execute_reply":"2024-08-11T12:08:29.508830Z"},"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-11T12:09:14.388106Z","iopub.execute_input":"2024-08-11T12:09:14.388604Z","iopub.status.idle":"2024-08-11T12:09:14.967959Z","shell.execute_reply.started":"2024-08-11T12:09:14.388541Z","shell.execute_reply":"2024-08-11T12:09:14.966801Z"},"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-11T12:10:15.840197Z","iopub.execute_input":"2024-08-11T12:10:15.840565Z","iopub.status.idle":"2024-08-11T12:10:15.855899Z","shell.execute_reply.started":"2024-08-11T12:10:15.840534Z","shell.execute_reply":"2024-08-11T12:10:15.854901Z"},"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-11T12:10:19.763623Z","iopub.execute_input":"2024-08-11T12:10:19.764332Z","iopub.status.idle":"2024-08-11T14:49:20.124398Z","shell.execute_reply.started":"2024-08-11T12:10:19.764296Z","shell.execute_reply":"2024-08-11T14:49:20.122076Z"},"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":"markdown","source":"# infrence","metadata":{}},{"cell_type":"code","source":"test_labels_true = pd.read_csv(DATA_PATH/\"train_label_coordinates.csv\")\ntest_labels_true = test_labels_true[test_labels_true['study_id'].isin(selected_ids)]\n","metadata":{"execution":{"iopub.status.busy":"2024-08-11T14:49:29.945651Z","iopub.execute_input":"2024-08-11T14:49:29.946102Z","iopub.status.idle":"2024-08-11T14:49:29.952784Z","shell.execute_reply.started":"2024-08-11T14:49:29.946065Z","shell.execute_reply":"2024-08-11T14:49:29.951902Z"},"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\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-11T14:53:54.764426Z","iopub.execute_input":"2024-08-11T14:53:54.765064Z","iopub.status.idle":"2024-08-11T14:53:54.776565Z","shell.execute_reply.started":"2024-08-11T14:53:54.765030Z","shell.execute_reply":"2024-08-11T14:53:54.775627Z"},"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-11T14:54:01.286023Z","iopub.execute_input":"2024-08-11T14:54:01.286767Z","iopub.status.idle":"2024-08-11T14:54:01.293496Z","shell.execute_reply.started":"2024-08-11T14:54:01.286717Z","shell.execute_reply":"2024-08-11T14:54:01.292585Z"},"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\n    img_size = 224, # 224, 384\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 = 1,\n    warmup = 1,\n    num_cycles = 0.375,\n    device = torch.device(\"cuda:0\") if torch.cuda.is_available() else \"cpu\",\n    seed = 2024,\n    log_wandb = False,\n    with_clip = False,\n)\n\n\n\nseeding(CONFIG['seed'])","metadata":{"execution":{"iopub.status.busy":"2024-08-11T14:54:11.902568Z","iopub.execute_input":"2024-08-11T14:54:11.903291Z","iopub.status.idle":"2024-08-11T14:54:11.913771Z","shell.execute_reply.started":"2024-08-11T14:54:11.903260Z","shell.execute_reply":"2024-08-11T14:54:11.912837Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_labels_true = pd.read_csv(DATA_PATH/\"train_label_coordinates.csv\")\ntest_labels_true = test_labels_true[test_labels_true['study_id'].isin(selected_ids)]\n","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_desc = test_labels_true","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_desc=pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\")\ntrain_desc = train_desc[train_desc['study_id'].isin(selected_ids)]\n","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:02:21.043087Z","iopub.execute_input":"2024-08-11T15:02:21.043800Z","iopub.status.idle":"2024-08-11T15:02:21.056004Z","shell.execute_reply.started":"2024-08-11T15:02:21.043768Z","shell.execute_reply":"2024-08-11T15:02:21.054934Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_desc","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:02:23.360127Z","iopub.execute_input":"2024-08-11T15:02:23.360505Z","iopub.status.idle":"2024-08-11T15:02:23.375426Z","shell.execute_reply.started":"2024-08-11T15:02:23.360473Z","shell.execute_reply":"2024-08-11T15:02:23.374493Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_desc[train_desc['study_id']]","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:00:35.152891Z","iopub.execute_input":"2024-08-11T15:00:35.153663Z","iopub.status.idle":"2024-08-11T15:00:35.166311Z","shell.execute_reply.started":"2024-08-11T15:00:35.153630Z","shell.execute_reply":"2024-08-11T15:00:35.165336Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_labels_true","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:00:01.441649Z","iopub.execute_input":"2024-08-11T15:00:01.442548Z","iopub.status.idle":"2024-08-11T15:00:01.456906Z","shell.execute_reply.started":"2024-08-11T15:00:01.442513Z","shell.execute_reply":"2024-08-11T15:00:01.455872Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DATA_PATH = Path(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification\")\ntrain_main = pd.read_csv(DATA_PATH/\"train.csv\")\ntest_desc =train_desc #pd.read_csv(DATA_PATH/\"test_series_descriptions.csv\")\nsample_df = pd.read_csv(DATA_PATH/\"sample_submission.csv\")\nstudy_ids = test_desc['study_id'].unique().tolist()","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:02:53.866951Z","iopub.execute_input":"2024-08-11T15:02:53.867614Z","iopub.status.idle":"2024-08-11T15:02:53.890545Z","shell.execute_reply.started":"2024-08-11T15:02:53.867583Z","shell.execute_reply":"2024-08-11T15:02:53.889629Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_desc","metadata":{"execution":{"iopub.status.busy":"2024-08-11T14:58:18.844455Z","iopub.execute_input":"2024-08-11T14:58:18.845170Z","iopub.status.idle":"2024-08-11T14:58:18.856065Z","shell.execute_reply.started":"2024-08-11T14:58:18.845136Z","shell.execute_reply":"2024-08-11T14:58:18.855054Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"AXIAL_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}","metadata":{"execution":{"iopub.status.busy":"2024-08-11T14:54:49.701209Z","iopub.execute_input":"2024-08-11T14:54:49.701905Z","iopub.status.idle":"2024-08-11T14:54:49.707679Z","shell.execute_reply.started":"2024-08-11T14:54:49.701871Z","shell.execute_reply":"2024-08-11T14:54:49.706742Z"},"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-11T14:54:55.869064Z","iopub.execute_input":"2024-08-11T14:54:55.869769Z","iopub.status.idle":"2024-08-11T14:54:55.875488Z","shell.execute_reply.started":"2024-08-11T14:54:55.869731Z","shell.execute_reply":"2024-08-11T14:54:55.874477Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class Spine25DDataset(Dataset):\n    def __init__(self, data, st_ids, transform=None):\n        self.data = data\n        self.st_ids = st_ids\n        self.transform = transform\n        \n    \n    def __len__(self):\n        return len(self.st_ids)\n    \n    def get_img_paths(self, study_id, series_desc):\n        pdf = self.data[self.data['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)}/test_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 __getitem__(self, idx):\n        row = self.data.iloc[idx]\n        study_id = self.st_ids[idx]\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        \n        # Sagittal\n        allimgs_sag = self.get_img_paths(study_id, 'Sagittal T2/STIR')\n        if len(allimgs_sag)==0:\n            pass\n        \n        else:\n            for i in range(10):\n                try:\n                    img = self.read_dcm(allimgs_sag[i])\n                    img = PIL.Image.fromarray(img).convert(\"RGB\")\n                    img = np.array(img).astype(np.uint8)\n                    img = self.transform(image=img)['image']\n                    sagittal[i, ...] = img\n                    \n                except:\n                    pass\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            for i in range(10):\n                try:\n                    img = self.read_dcm(allimgs_cor[i])\n                    img = PIL.Image.fromarray(img).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        allimgs_ax = self.get_img_paths(study_id, 'Axial T2')\n        if len(allimgs_ax)==0:\n            pass\n        \n        else:\n            for i in range(10):\n                try:\n                    img = self.read_dcm(allimgs_ax[i])\n                    img = PIL.Image.fromarray(img).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        return {\"axial\": axial, \"coronal\": coronal, \"sagittal\": sagittal, \"study_id\": str(study_id)}","metadata":{"execution":{"iopub.status.busy":"2024-08-11T14:55:00.400969Z","iopub.execute_input":"2024-08-11T14:55:00.401336Z","iopub.status.idle":"2024-08-11T14:55:00.422894Z","shell.execute_reply.started":"2024-08-11T14:55:00.401307Z","shell.execute_reply":"2024-08-11T14:55:00.421992Z"},"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.3, scale_limit=0.3, rotate_limit=45, border_mode=4, 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        ], p=0.5),\n        \n        A.CoarseDropout(max_holes=2, max_height=int(height * 0.25), max_width=int(width * 0.25), 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, ids, 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 = Spine25DDataset(data=data, st_ids=ids, 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 = Spine25DDataset(data=data, st_ids=ids, transform=eval_tsfm)\n        dls = DataLoader(ds, \n                         batch_size=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\n","metadata":{"execution":{"iopub.status.busy":"2024-08-11T14:55:06.152065Z","iopub.execute_input":"2024-08-11T14:55:06.152470Z","iopub.status.idle":"2024-08-11T14:55:06.165640Z","shell.execute_reply.started":"2024-08-11T14:55:06.152439Z","shell.execute_reply":"2024-08-11T14:55:06.164698Z"},"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        \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    \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    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        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","metadata":{"execution":{"iopub.status.busy":"2024-08-11T14:55:13.321153Z","iopub.execute_input":"2024-08-11T14:55:13.321526Z","iopub.status.idle":"2024-08-11T14:55:13.343882Z","shell.execute_reply.started":"2024-08-11T14:55:13.321495Z","shell.execute_reply":"2024-08-11T14:55:13.342740Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"CONDITIONS = [\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]\n\ndls = get_dataloaders(test_desc, study_ids, CONFIG, split='test')","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:03:02.745964Z","iopub.execute_input":"2024-08-11T15:03:02.747152Z","iopub.status.idle":"2024-08-11T15:03:02.753611Z","shell.execute_reply.started":"2024-08-11T15:03:02.747107Z","shell.execute_reply":"2024-08-11T15:03:02.752609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_ids","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:27:19.387685Z","iopub.execute_input":"2024-08-11T15:27:19.388086Z","iopub.status.idle":"2024-08-11T15:27:19.394924Z","shell.execute_reply.started":"2024-08-11T15:27:19.388053Z","shell.execute_reply":"2024-08-11T15:27:19.393947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def inference(model, dataloader):\n    model.to(CONFIG[\"device\"])\n    model.eval()\n    y_preds = []\n    row_names = []\n\n    axial_indices = list(AXIAL_COLS.values())\n    coronal_indices = list(SAGT1_COLS.values())\n    sagittal_indices = list(SAGT2_COLS.values())\n    \n    pbar = tqdm(dls, leave=True)\n    \n    with torch.no_grad():\n        for idx, batch in enumerate(pbar):\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            si = batch['study_id']\n            pred_per_study = np.ones((25, 3)) * (1/3)\n            for cond in CONDITIONS:\n                for level in LEVELS:\n                    row_names.append(si[0] + '_' + cond + '_' + level)\n                \n            with torch.autocast(device_type=\"cuda\", dtype=torch.float16):\n                y_axial, y_coronal, y_sagittal = model(axial, coronal, sagittal)\n#                 y_axial = y_axial.squeeze().mean(dim=0) \n#                 y_coronal = y_coronal.squeeze().mean(dim=0) \n#                 y_sagittal = y_sagittal.squeeze().mean(dim=0)\n                \n                y_axial = y_axial.squeeze().amax(dim=0) \n                y_coronal = y_coronal.squeeze().amax(dim=0) \n                y_sagittal = y_sagittal.squeeze().amax(dim=0)\n                \n                # axial\n                for col in range(CONFIG['axial_labels']):\n                    pred = y_axial[col*3:col*3+3]\n                    y_pred = pred.float().softmax(dim=-1).cpu().numpy()\n                    pred_per_study[axial_indices[col]] = y_pred\n                    \n                # coronal\n                for col in range(CONFIG['sagT1_labels']):\n                    pred = y_coronal[col*3:col*3+3]\n                    y_pred = pred.float().softmax(dim=-1).cpu().numpy()\n                    pred_per_study[coronal_indices[col]] = y_pred\n                    \n                # sagittal\n                for col in range(CONFIG['sagT2_labels']):\n                    pred = y_sagittal[col*3:col*3+3]\n                    y_pred = pred.float().softmax(dim=-1).cpu().numpy()\n                    pred_per_study[sagittal_indices[col]] = y_pred\n            y_preds.append(pred_per_study)\n            \n    y_preds = np.concatenate(y_preds, axis=0)\n    return y_preds, row_names","metadata":{"execution":{"iopub.status.busy":"2024-08-11T14:55:42.345368Z","iopub.execute_input":"2024-08-11T14:55:42.345770Z","iopub.status.idle":"2024-08-11T14:55:42.360541Z","shell.execute_reply.started":"2024-08-11T14:55:42.345737Z","shell.execute_reply":"2024-08-11T14:55:42.359637Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"models = []\nnet = Clf(backbone=CONFIG['backbone'], increase_stride=False)\nweights_path = \"/kaggle/working/rsna_2024_lumbar_spine_fold_0_epoch_16.pth\"\nweights = torch.load(weights_path, map_location=torch.device(\"cpu\"))\nnet.load_state_dict(weights)\nnet.eval()\nmodels.append(net.to(CONFIG['device']))","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:03:06.779241Z","iopub.execute_input":"2024-08-11T15:03:06.779633Z","iopub.status.idle":"2024-08-11T15:03:07.679417Z","shell.execute_reply.started":"2024-08-11T15:03:06.779595Z","shell.execute_reply":"2024-08-11T15:03:07.678628Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_id=[df['row_id'][i].split(\"_\")[0] for i in range(len(row_ids))]","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:18:54.899493Z","iopub.execute_input":"2024-08-11T15:18:54.900456Z","iopub.status.idle":"2024-08-11T15:18:54.912661Z","shell.execute_reply.started":"2024-08-11T15:18:54.900409Z","shell.execute_reply":"2024-08-11T15:18:54.911553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds, row_ids = inference(net, dls)","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:03:08.645141Z","iopub.execute_input":"2024-08-11T15:03:08.645545Z","iopub.status.idle":"2024-08-11T15:03:12.748003Z","shell.execute_reply.started":"2024-08-11T15:03:08.645495Z","shell.execute_reply":"2024-08-11T15:03:12.746897Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['study_id']=study_id","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:29:48.802621Z","iopub.execute_input":"2024-08-11T15:29:48.803294Z","iopub.status.idle":"2024-08-11T15:29:48.808145Z","shell.execute_reply.started":"2024-08-11T15:29:48.803260Z","shell.execute_reply":"2024-08-11T15:29:48.807207Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"TARGET_COLS = sample_df.columns.tolist()\ndf = pd.DataFrame()\ndf['row_id'] = row_ids\ndf[['normal', 'mild', 'severe']] = preds\ndf.columns = TARGET_COLS\ndf = df.sort_values(\"row_id\").reset_index(drop=True)\ndf","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:29:47.476849Z","iopub.execute_input":"2024-08-11T15:29:47.477531Z","iopub.status.idle":"2024-08-11T15:29:47.495573Z","shell.execute_reply.started":"2024-08-11T15:29:47.477493Z","shell.execute_reply":"2024-08-11T15:29:47.494634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df['max_condition'] = df[['normal_mild', 'moderate', 'severe']].idxmax(axis=1)\n","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:30:04.193699Z","iopub.execute_input":"2024-08-11T15:30:04.194075Z","iopub.status.idle":"2024-08-11T15:30:04.201283Z","shell.execute_reply.started":"2024-08-11T15:30:04.194045Z","shell.execute_reply":"2024-08-11T15:30:04.200281Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"code","source":"df = df.sort_values(\"row_id\").reset_index(drop=True)\n","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:11:23.254037Z","iopub.execute_input":"2024-08-11T15:11:23.254752Z","iopub.status.idle":"2024-08-11T15:11:23.260836Z","shell.execute_reply.started":"2024-08-11T15:11:23.254719Z","shell.execute_reply":"2024-08-11T15:11:23.259762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_labels=df[\"max_condition\"]","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:42:31.528499Z","iopub.execute_input":"2024-08-11T15:42:31.529426Z","iopub.status.idle":"2024-08-11T15:42:31.534733Z","shell.execute_reply.started":"2024-08-11T15:42:31.529390Z","shell.execute_reply":"2024-08-11T15:42:31.533645Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_labels,true_labels\ntrue_labels=new_df[\"label\"]","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:34:21.162623Z","iopub.execute_input":"2024-08-11T15:34:21.163012Z","iopub.status.idle":"2024-08-11T15:34:21.171897Z","shell.execute_reply.started":"2024-08-11T15:34:21.162977Z","shell.execute_reply":"2024-08-11T15:34:21.170760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"true_labels=new_df[\"label\"]","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:42:20.709270Z","iopub.execute_input":"2024-08-11T15:42:20.709967Z","iopub.status.idle":"2024-08-11T15:42:20.714498Z","shell.execute_reply.started":"2024-08-11T15:42:20.709931Z","shell.execute_reply":"2024-08-11T15:42:20.713491Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"true_labels","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:42:25.188887Z","iopub.execute_input":"2024-08-11T15:42:25.189656Z","iopub.status.idle":"2024-08-11T15:42:25.196801Z","shell.execute_reply.started":"2024-08-11T15:42:25.189621Z","shell.execute_reply":"2024-08-11T15:42:25.195916Z"},"collapsed":true,"jupyter":{"outputs_hidden":true},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_df","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:33:07.087023Z","iopub.execute_input":"2024-08-11T15:33:07.087710Z","iopub.status.idle":"2024-08-11T15:33:07.099415Z","shell.execute_reply.started":"2024-08-11T15:33:07.087673Z","shell.execute_reply":"2024-08-11T15:33:07.098500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Reshaping the DataFrame\nnew_df = train_data_csv.melt(id_vars=[\"study_id\"], var_name=\"condition\", value_name=\"label\")\n\n# Creating the row_id column\nnew_df[\"row_id\"] = new_df[\"study_id\"].astype(str) + \"_\" + new_df[\"condition\"]","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:42:11.445795Z","iopub.execute_input":"2024-08-11T15:42:11.446532Z","iopub.status.idle":"2024-08-11T15:42:11.457519Z","shell.execute_reply.started":"2024-08-11T15:42:11.446499Z","shell.execute_reply":"2024-08-11T15:42:11.456651Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_df = new_df.sort_values(\"row_id\").reset_index(drop=True)\n","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:42:13.227607Z","iopub.execute_input":"2024-08-11T15:42:13.227983Z","iopub.status.idle":"2024-08-11T15:42:13.234299Z","shell.execute_reply.started":"2024-08-11T15:42:13.227951Z","shell.execute_reply":"2024-08-11T15:42:13.233321Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_csv","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:41:45.542770Z","iopub.execute_input":"2024-08-11T15:41:45.543495Z","iopub.status.idle":"2024-08-11T15:41:45.572668Z","shell.execute_reply.started":"2024-08-11T15:41:45.543461Z","shell.execute_reply":"2024-08-11T15:41:45.571738Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_csv = train_data_csv.sort_values(\"study_id\").reset_index(drop=True)\n","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:14:00.051085Z","iopub.execute_input":"2024-08-11T15:14:00.051472Z","iopub.status.idle":"2024-08-11T15:14:00.058522Z","shell.execute_reply.started":"2024-08-11T15:14:00.051444Z","shell.execute_reply":"2024-08-11T15:14:00.057391Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_csv","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:14:02.315989Z","iopub.execute_input":"2024-08-11T15:14:02.316373Z","iopub.status.idle":"2024-08-11T15:14:02.346690Z","shell.execute_reply.started":"2024-08-11T15:14:02.316341Z","shell.execute_reply":"2024-08-11T15:14:02.345755Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_labels,true_labels","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:42:42.741968Z","iopub.execute_input":"2024-08-11T15:42:42.742704Z","iopub.status.idle":"2024-08-11T15:42:42.750268Z","shell.execute_reply.started":"2024-08-11T15:42:42.742672Z","shell.execute_reply":"2024-08-11T15:42:42.749385Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_labels=predicted_labels.map(test_dict)\ntrue_labels=true_labels.map(train_dict)","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:42:45.436201Z","iopub.execute_input":"2024-08-11T15:42:45.436802Z","iopub.status.idle":"2024-08-11T15:42:45.442563Z","shell.execute_reply.started":"2024-08-11T15:42:45.436772Z","shell.execute_reply":"2024-08-11T15:42:45.441641Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"predicted_labels","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:42:48.148500Z","iopub.execute_input":"2024-08-11T15:42:48.148889Z","iopub.status.idle":"2024-08-11T15:42:48.156535Z","shell.execute_reply.started":"2024-08-11T15:42:48.148855Z","shell.execute_reply":"2024-08-11T15:42:48.155619Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"true_labels=np.array(true_labels,dtype='int64')\npredicted_labels=np.array(predicted_labels,dtype='int64')","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:42:51.684670Z","iopub.execute_input":"2024-08-11T15:42:51.685531Z","iopub.status.idle":"2024-08-11T15:42:51.690242Z","shell.execute_reply.started":"2024-08-11T15:42:51.685497Z","shell.execute_reply":"2024-08-11T15:42:51.689363Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nfrom sklearn.metrics import accuracy_score, confusion_matrix\n\n# Example true and predicted labels\n# y_true = [0, 1, 2, 2, 1, 0, 1, 2, 0, 1]  # Replace with your true labels\n# y_pred = [0, 2, 1, 2, 1, 0, 0, 2, 0, 1]  # Replace with your predicted labels\n\n# Calculate accuracy\naccuracy = accuracy_score(true_labels, predicted_labels)\nprint(f'Accuracy: {accuracy:.2f}')\n\n# Calculate confusion matrix\nconf_matrix = confusion_matrix(true_labels, predicted_labels)\nprint('Confusion Matrix:')\nprint(conf_matrix)\n","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:42:53.661092Z","iopub.execute_input":"2024-08-11T15:42:53.661963Z","iopub.status.idle":"2024-08-11T15:42:53.670852Z","shell.execute_reply.started":"2024-08-11T15:42:53.661925Z","shell.execute_reply":"2024-08-11T15:42:53.669943Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_dict={\n    \"normal_mild\":0,\n    \"moderate\":1,\n    \"severe\": 2 \n}\ntrain_dict={\n    \"Normal/Mild\":0,\n    \"Moderate\":1,\n    \"Severe\": 2\n}","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:41:24.724327Z","iopub.execute_input":"2024-08-11T15:41:24.724725Z","iopub.status.idle":"2024-08-11T15:41:24.729756Z","shell.execute_reply.started":"2024-08-11T15:41:24.724693Z","shell.execute_reply":"2024-08-11T15:41:24.728712Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df[df['study_id']==\"3234424112\"]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_csv","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:23:28.277250Z","iopub.execute_input":"2024-08-11T15:23:28.277744Z","iopub.status.idle":"2024-08-11T15:23:28.308233Z","shell.execute_reply.started":"2024-08-11T15:23:28.277709Z","shell.execute_reply":"2024-08-11T15:23:28.307305Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_csv.loc[14]","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:22:46.559297Z","iopub.execute_input":"2024-08-11T15:22:46.559993Z","iopub.status.idle":"2024-08-11T15:22:46.567195Z","shell.execute_reply.started":"2024-08-11T15:22:46.559957Z","shell.execute_reply":"2024-08-11T15:22:46.566335Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_csv[train_data_csv[\"study_id\"]==3234424112]","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:21:20.192930Z","iopub.execute_input":"2024-08-11T15:21:20.193597Z","iopub.status.idle":"2024-08-11T15:21:20.212542Z","shell.execute_reply.started":"2024-08-11T15:21:20.193552Z","shell.execute_reply":"2024-08-11T15:21:20.211678Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data_csv=pd.read_csv('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv')\ntrain_data_csv = train_data_csv[train_data_csv['study_id'].isin(selected_ids)]\n","metadata":{"execution":{"iopub.status.busy":"2024-08-11T15:05:21.838252Z","iopub.execute_input":"2024-08-11T15:05:21.838946Z","iopub.status.idle":"2024-08-11T15:05:21.869346Z","shell.execute_reply.started":"2024-08-11T15:05:21.838914Z","shell.execute_reply":"2024-08-11T15:05:21.868590Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_labels_true","metadata":{},"execution_count":null,"outputs":[]}]}