{"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":192246722,"sourceType":"kernelVersion"}],"dockerImageVersionId":30716,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Inference\n\nTraining code: [rsna-pytorch-train-woptimizemetric](https://www.kaggle.com/code/samu2505/rsna-pytorch-train-woptimizemetric?scriptVersionId=192246722)\n\nLoading Model weights and Cross-validation: [wandb&cross-validation](https://www.kaggle.com/code/samu2505/rsna-wandbmodelweights-crossvalidation/notebook)","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19"}},{"cell_type":"markdown","source":"# Import libraries","metadata":{}},{"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\nimport timm\n\nimport cv2\ncv2.setNumThreads(0)\nimport PIL\nimport pydicom\nimport warnings\nwarnings.filterwarnings(\"ignore\")","metadata":{"execution":{"iopub.status.busy":"2024-08-21T01:14:41.008866Z","iopub.execute_input":"2024-08-21T01:14:41.009246Z","iopub.status.idle":"2024-08-21T01:14:48.448865Z","shell.execute_reply.started":"2024-08-21T01:14:41.009205Z","shell.execute_reply":"2024-08-21T01:14:48.447860Z"},"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 = True\n        torch.backends.cudnn.benchmark = False\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-21T01:14:48.450613Z","iopub.execute_input":"2024-08-21T01:14:48.450898Z","iopub.status.idle":"2024-08-21T01:14:48.457823Z","shell.execute_reply.started":"2024-08-21T01:14:48.450873Z","shell.execute_reply":"2024-08-21T01:14:48.456369Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Configuration","metadata":{}},{"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 = 15,\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.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\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-21T01:14:48.459117Z","iopub.execute_input":"2024-08-21T01:14:48.459435Z","iopub.status.idle":"2024-08-21T01:14:48.474437Z","shell.execute_reply.started":"2024-08-21T01:14:48.459409Z","shell.execute_reply":"2024-08-21T01:14:48.473382Z"},"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 = 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-21T01:14:48.475668Z","iopub.execute_input":"2024-08-21T01:14:48.476001Z","iopub.status.idle":"2024-08-21T01:14:48.529809Z","shell.execute_reply.started":"2024-08-21T01:14:48.475963Z","shell.execute_reply":"2024-08-21T01:14:48.528584Z"},"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) + 1e-4)\n    data = (data * 255).astype(np.uint8)\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-08-21T01:15:11.676222Z","iopub.execute_input":"2024-08-21T01:15:11.676644Z","iopub.status.idle":"2024-08-21T01:15:11.682749Z","shell.execute_reply.started":"2024-08-21T01:15:11.676609Z","shell.execute_reply":"2024-08-21T01:15:11.681581Z"},"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 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 = self.st_ids[idx]\n        H, W = CONFIG['img_size'], CONFIG['img_size']\n        H, W = CONFIG['img_size'], CONFIG['img_size']\n        nchans = CONFIG['in_chans']\n        sagittal = np.zeros((H, W, nchans), dtype=np.uint8)\n        coronal = np.zeros((H, W, nchans), dtype=np.uint8)\n        axial = np.zeros((H, W, nchans), dtype=np.uint8)\n        \n        # Sagittal\n        allimgs_sag = self.get_img_paths(study_id, 'Sagittal T2/STIR')\n        sagT2_scans = len(allimgs_sag)\n        sagT2_indices = np.quantile(list(range(sagT2_scans)), np.linspace(0., 1., nchans)).round().astype(int)\n        allimgs_sag = [allimgs_sag[i] for i in sagT2_indices]\n        \n        if len(allimgs_sag)==0:\n            pass\n        \n        else:\n            sagittal = self.get_images(nslides=nchans, image_paths=allimgs_sag)\n            \n        # coronal\n        allimgs_cor = self.get_img_paths(study_id, 'Sagittal T1')\n        sagT1_scans = len(allimgs_cor)\n        sagT1_indices = np.quantile(list(range(sagT1_scans)), np.linspace(0., 1., nchans)).round().astype(int)\n        allimgs_cor = [allimgs_cor[i] for i in sagT1_indices]\n        \n        if len(allimgs_cor)==0:\n            pass\n        \n        else:\n            coronal = self.get_images(nslides=nchans, image_paths=allimgs_cor)\n                \n        # Axial\n        allimgs_ax = self.get_img_paths(study_id, 'Axial T2')\n        ax_scans = len(allimgs_ax)\n        ax_indices = np.quantile(list(range(ax_scans)), np.linspace(0., 1., nchans)).round().astype(int)\n        allimgs_ax = [allimgs_ax[i] for i in ax_indices]\n        \n        if len(allimgs_ax)==0:\n            pass\n        \n        else:\n            axial = self.get_images(nslides=nchans, 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        return {\"axial\": axial, \"coronal\": coronal, \"sagittal\": sagittal, \"study_id\": str(study_id)}","metadata":{"execution":{"iopub.status.busy":"2024-08-21T01:15:13.188716Z","iopub.execute_input":"2024-08-21T01:15:13.189122Z","iopub.status.idle":"2024-08-21T01:15:18.466058Z","shell.execute_reply.started":"2024-08-21T01:15:13.189092Z","shell.execute_reply":"2024-08-21T01:15:18.464983Z"},"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","metadata":{"execution":{"iopub.status.busy":"2024-08-21T01:15:18.467965Z","iopub.execute_input":"2024-08-21T01:15:18.468306Z","iopub.status.idle":"2024-08-21T01:15:18.481603Z","shell.execute_reply.started":"2024-08-21T01:15:18.468277Z","shell.execute_reply":"2024-08-21T01:15:18.480714Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dls = get_dataloaders(test_desc, study_ids, CONFIG, split='test')\nb = next(iter(dls))","metadata":{"execution":{"iopub.status.busy":"2024-08-21T01:15:18.482789Z","iopub.execute_input":"2024-08-21T01:15:18.483097Z","iopub.status.idle":"2024-08-21T01:15:19.612137Z","shell.execute_reply.started":"2024-08-21T01:15:18.483072Z","shell.execute_reply":"2024-08-21T01:15:19.610853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model","metadata":{}},{"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    \nnet = Clf(backbone=CONFIG['backbone'])\nnet.eval()\nnet(b['axial'], b['coronal'], b['sagittal']).shape","metadata":{"execution":{"iopub.status.busy":"2024-08-21T01:15:19.616053Z","iopub.execute_input":"2024-08-21T01:15:19.616541Z","iopub.status.idle":"2024-08-21T01:15:20.331265Z","shell.execute_reply.started":"2024-08-21T01:15:19.616475Z","shell.execute_reply":"2024-08-21T01:15:20.330185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Predictions","metadata":{}},{"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-21T01:15:20.332498Z","iopub.execute_input":"2024-08-21T01:15:20.332861Z","iopub.status.idle":"2024-08-21T01:15:20.339049Z","shell.execute_reply.started":"2024-08-21T01:15:20.332834Z","shell.execute_reply":"2024-08-21T01:15:20.337937Z"},"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    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                logits = model(axial, coronal, sagittal)[0]\n#                 logits = logits.reshape(3*25)\n                for col in range(25):\n#                     pred = logits[col*3:col*3+3]\n                    pred = logits[:, col]\n                    y_pred = pred.float().softmax(dim=-1).cpu().numpy()\n                    pred_per_study[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-21T01:15:20.340747Z","iopub.execute_input":"2024-08-21T01:15:20.341064Z","iopub.status.idle":"2024-08-21T01:15:20.351225Z","shell.execute_reply.started":"2024-08-21T01:15:20.341036Z","shell.execute_reply":"2024-08-21T01:15:20.350255Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Load model weights","metadata":{}},{"cell_type":"code","source":"weights_path = \"/kaggle/input/rsna-pytorch-train-woptimizemetric/rsna_2024_lumbar_spine_fold_0_epoch_5.pth\"\nweights = torch.load(weights_path, map_location=torch.device(\"cpu\"))\nmodel = Clf(backbone=CONFIG['backbone'], pretrained=False)\nmodel.load_state_dict(weights)","metadata":{"execution":{"iopub.status.busy":"2024-08-21T01:15:20.352449Z","iopub.execute_input":"2024-08-21T01:15:20.352883Z","iopub.status.idle":"2024-08-21T01:15:20.909701Z","shell.execute_reply.started":"2024-08-21T01:15:20.352853Z","shell.execute_reply":"2024-08-21T01:15:20.908467Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"preds, row_ids = inference(model, dls)","metadata":{"execution":{"iopub.status.busy":"2024-08-21T01:15:20.910967Z","iopub.execute_input":"2024-08-21T01:15:20.911276Z","iopub.status.idle":"2024-08-21T01:15:22.449073Z","shell.execute_reply.started":"2024-08-21T01:15:20.911246Z","shell.execute_reply":"2024-08-21T01:15:22.447636Z"},"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.to_csv('submission.csv', index=False)","metadata":{"execution":{"iopub.status.busy":"2024-08-21T01:15:22.451353Z","iopub.execute_input":"2024-08-21T01:15:22.451773Z","iopub.status.idle":"2024-08-21T01:15:22.471578Z","shell.execute_reply.started":"2024-08-21T01:15:22.451733Z","shell.execute_reply":"2024-08-21T01:15:22.470563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pd.read_csv(\"submission.csv\")","metadata":{"execution":{"iopub.status.busy":"2024-08-21T01:15:22.474577Z","iopub.execute_input":"2024-08-21T01:15:22.474941Z","iopub.status.idle":"2024-08-21T01:15:22.494230Z","shell.execute_reply.started":"2024-08-21T01:15:22.474911Z","shell.execute_reply":"2024-08-21T01:15:22.493181Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}