{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":10676978,"sourceType":"datasetVersion","datasetId":6613756}],"dockerImageVersionId":30840,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# Import lib","metadata":{}},{"cell_type":"code","source":"import os\nimport gc\nimport re\nimport sys\nfrom PIL import Image\nimport cv2\nimport math, random\nimport numpy as np\nimport pandas as pd\nfrom glob import glob\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nfrom sklearn.model_selection import KFold\n\nfrom collections import OrderedDict\n\nimport torch\nimport torch.nn.functional as F\nfrom torch import nn\nfrom torch.utils.data import DataLoader, Dataset\nfrom torch.optim import AdamW\n\nimport timm\nfrom transformers import get_cosine_schedule_with_warmup\n\nimport albumentations as A\n\nfrom sklearn.model_selection import KFold","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:29:29.750614Z","iopub.execute_input":"2025-02-07T16:29:29.751026Z","iopub.status.idle":"2025-02-07T16:29:29.757687Z","shell.execute_reply.started":"2025-02-07T16:29:29.750985Z","shell.execute_reply":"2025-02-07T16:29:29.756743Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Config","metadata":{}},{"cell_type":"code","source":"path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:34.941428Z","iopub.execute_input":"2025-02-07T14:10:34.941932Z","iopub.status.idle":"2025-02-07T14:10:34.945448Z","shell.execute_reply.started":"2025-02-07T14:10:34.941907Z","shell.execute_reply":"2025-02-07T14:10:34.944517Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"NOT_DEBUG = True # True -> run naormally, False -> debug mode, with lesser computing cost\n\nOUTPUT_DIR = f'rsna24-results'\ndevice = 'cuda:0' if torch.cuda.is_available() else 'cpu'\nN_WORKERS = os.cpu_count() \nUSE_AMP = True # can change True if using T4 or newer than Ampere\nSEED = 8620\n\nIMG_SIZE = [512, 512]\nIN_CHANS = 30\nN_LABELS = 25\nN_CLASSES = 3 * N_LABELS\n\nAUG_PROB = 0.75\n\nN_FOLDS = 5 if NOT_DEBUG else 2\nEPOCHS = 20 if NOT_DEBUG else 2\n#MODEL_NAME = \"tf_efficientnet_b3.ns_jft_in1k\" if NOT_DEBUG else \"tf_efficientnet_b0.ns_jft_in1k\"\nMODEL_NAME = \"convnext_base_in22k\" if NOT_DEBUG else \"convnext_tiny_in22k\"\nGRAD_ACC = 2\nTGT_BATCH_SIZE = 32\nBATCH_SIZE = TGT_BATCH_SIZE // GRAD_ACC\nMAX_GRAD_NORM = None\nEARLY_STOPPING_EPOCH = 3\n\nLR = 2e-4 * TGT_BATCH_SIZE / 32\nWD = 1e-2\nAUG = True","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:34.946981Z","iopub.execute_input":"2025-02-07T14:10:34.947310Z","iopub.status.idle":"2025-02-07T14:10:35.034947Z","shell.execute_reply.started":"2025-02-07T14:10:34.947280Z","shell.execute_reply":"2025-02-07T14:10:35.034187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"os.makedirs(OUTPUT_DIR, exist_ok=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:35.036263Z","iopub.execute_input":"2025-02-07T14:10:35.036598Z","iopub.status.idle":"2025-02-07T14:10:35.047197Z","shell.execute_reply.started":"2025-02-07T14:10:35.036566Z","shell.execute_reply":"2025-02-07T14:10:35.046566Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def set_random_seed(seed: int = 8620, deterministic: bool = False):\n    \"\"\"Set seeds\"\"\"\n    random.seed(seed)\n    np.random.seed(seed)\n    os.environ[\"PYTHONHASHSEED\"] = str(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)  # type: ignore\n    torch.backends.cudnn.benchmark = True\n    torch.backends.cudnn.deterministic = deterministic  # type: ignore\n\nset_random_seed(SEED)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:35.047997Z","iopub.execute_input":"2025-02-07T14:10:35.048258Z","iopub.status.idle":"2025-02-07T14:10:35.065741Z","shell.execute_reply.started":"2025-02-07T14:10:35.048215Z","shell.execute_reply":"2025-02-07T14:10:35.064978Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataframes","metadata":{}},{"cell_type":"code","source":"df = pd.read_csv(f'{path}/train.csv')\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:35.066569Z","iopub.execute_input":"2025-02-07T14:10:35.066899Z","iopub.status.idle":"2025-02-07T14:10:35.123019Z","shell.execute_reply.started":"2025-02-07T14:10:35.066868Z","shell.execute_reply":"2025-02-07T14:10:35.122299Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df = df.fillna(-100)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:35.123795Z","iopub.execute_input":"2025-02-07T14:10:35.124039Z","iopub.status.idle":"2025-02-07T14:10:35.132970Z","shell.execute_reply.started":"2025-02-07T14:10:35.124017Z","shell.execute_reply":"2025-02-07T14:10:35.132306Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label2id = {'Normal/Mild': 0, 'Moderate':1, 'Severe':2}\ndf = df.replace(label2id)\ndf.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:35.135554Z","iopub.execute_input":"2025-02-07T14:10:35.135823Z","iopub.status.idle":"2025-02-07T14:10:35.171307Z","shell.execute_reply.started":"2025-02-07T14:10:35.135802Z","shell.execute_reply":"2025-02-07T14:10:35.170561Z"}},"outputs":[],"execution_count":null},{"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]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:35.172379Z","iopub.execute_input":"2025-02-07T14:10:35.172594Z","iopub.status.idle":"2025-02-07T14:10:35.175988Z","shell.execute_reply.started":"2025-02-07T14:10:35.172576Z","shell.execute_reply":"2025-02-07T14:10:35.175152Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def atoi(text):\n    return int(text) if text.isdigit() else text\n\ndef natural_keys(text):\n    return [ atoi(c) for c in re.split(r'(\\d+)', text) ]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:27:57.974452Z","iopub.execute_input":"2025-02-07T16:27:57.974842Z","iopub.status.idle":"2025-02-07T16:27:57.980186Z","shell.execute_reply.started":"2025-02-07T16:27:57.974801Z","shell.execute_reply":"2025-02-07T16:27:57.979343Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataset","metadata":{}},{"cell_type":"code","source":"class RSNA24Dataset(Dataset):\n    def __init__(self, df, phase='train', transform=None):\n        self.df = df\n        self.transform = transform\n        self.phase = phase\n    \n    def __len__(self):\n        return len(self.df)\n\n    def __getitem__(self, idx):\n        x = np.zeros((512, 512, IN_CHANS), dtype=np.uint8)\n        t = self.df.iloc[idx]\n        st_id = int(t['study_id'])\n        label = t[1:].values.astype(np.int64)\n        \n        # Sagittal T1\n        for i in range(0, 10, 1):\n            try:\n                p = f'/kaggle/input/lsdc-redefine-dataset/cvt_png/{st_id}/Sagittal T1/{i:03d}.png'\n                img = Image.open(p).convert('L')\n                img = np.array(img)\n                x[..., i] = img.astype(np.uint8)\n            except:\n                #print(f'failed to load on {st_id}, Sagittal T1')\n                pass\n            \n        # Sagittal T2/STIR\n        for i in range(0, 10, 1):\n            try:\n                p = f'/kaggle/input/lsdc-redefine-dataset/cvt_png/{st_id}/Sagittal T2_STIR/{i:03d}.png'\n                img = Image.open(p).convert('L')\n                img = np.array(img)\n                x[..., i+10] = img.astype(np.uint8)\n            except:\n                #print(f'failed to load on {st_id}, Sagittal T2/STIR')\n                pass\n            \n        # Axial T2\n        axt2 = glob(f'/kaggle/input/lsdc-redefine-dataset/cvt_png/{st_id}/Axial T2/*.png')\n        axt2 = sorted(axt2)\n    \n        step = len(axt2) / 10.0\n        st = len(axt2)/2.0 - 4.0*step\n        end = len(axt2)+0.0001\n                \n        for i, j in enumerate(np.arange(st, end, step)):\n            try:\n                p = axt2[max(0, int((j-0.5001).round()))]\n                img = Image.open(p).convert('L')\n                img = np.array(img)\n                x[..., i+20] = img.astype(np.uint8)\n            except:\n                #print(f'failed to load on {st_id}, Sagittal T2/STIR')\n                pass  \n            \n        assert np.sum(x)>0\n            \n        if self.transform is not None:\n            x = self.transform(image=x)['image']\n\n        x = x.transpose(2, 0, 1)\n                \n        return x, label","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:35.176788Z","iopub.execute_input":"2025-02-07T14:10:35.176982Z","iopub.status.idle":"2025-02-07T14:10:35.190389Z","shell.execute_reply.started":"2025-02-07T14:10:35.176964Z","shell.execute_reply":"2025-02-07T14:10:35.189620Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data augmentation","metadata":{}},{"cell_type":"code","source":"transforms_train = A.Compose([\n    #A.RandomBrightnessContrast(brightness_limit=(-0.2, 0.2), contrast_limit=(-0.2, 0.2), p=AUG_PROB),\n    # A.OneOf([\n    #     A.MotionBlur(blur_limit=5),\n    #     A.MedianBlur(blur_limit=5),\n    #     A.GaussianBlur(blur_limit=5),\n    #     A.GaussNoise(var_limit=(5.0, 30.0)),\n    # ], p=AUG_PROB),\n\n    # A.OneOf([\n    #     A.OpticalDistortion(distort_limit=1.0),\n    #     A.GridDistortion(num_steps=5, distort_limit=1.),\n    #     A.ElasticTransform(alpha=3),\n    # ], p=AUG_PROB),\n\n    A.ShiftScaleRotate(shift_limit=0.1, scale_limit=0.1, rotate_limit=15, border_mode=0, p=AUG_PROB),\n    A.Resize(IMG_SIZE[0], IMG_SIZE[1]),\n    #A.CoarseDropout(max_holes=16, max_height=64, max_width=64, min_holes=1, min_height=8, min_width=8, p=AUG_PROB),    \n    A.Normalize(mean=0.5, std=0.5)\n])\n\ntransforms_val = A.Compose([\n    A.Resize(IMG_SIZE[0], IMG_SIZE[1]),\n    A.Normalize(mean=0.5, std=0.5)\n])\n\nif not NOT_DEBUG or not AUG:\n    transforms_train = transforms_val","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:35.191077Z","iopub.execute_input":"2025-02-07T14:10:35.191350Z","iopub.status.idle":"2025-02-07T14:10:35.206958Z","shell.execute_reply.started":"2025-02-07T14:10:35.191328Z","shell.execute_reply":"2025-02-07T14:10:35.206225Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Dataloader","metadata":{}},{"cell_type":"code","source":"tmp_ds = RSNA24Dataset(df, phase='train', transform=transforms_train)\ntmp_dl = DataLoader(\n            tmp_ds,\n            batch_size=1,\n            shuffle=False,\n            pin_memory=True,\n            drop_last=False,\n            num_workers=0\n            )\n\nfor i, (x, t) in enumerate(tmp_dl):\n    if i==5:break\n    print('x stat:', x.shape, x.min(), x.max(),x.mean(), x.std())\n    print(t, t.shape)\n    y = x.numpy().transpose(0,2,3,1)[0,...,:3]\n    y = (y + 1) / 2\n    plt.imshow(y)\n    plt.show()\n    print('y stat:', y.shape, y.min(), y.max(),y.mean(), y.std())\n    print()\nplt.close()\ndel tmp_ds, tmp_dl","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:35.207712Z","iopub.execute_input":"2025-02-07T14:10:35.207976Z","iopub.status.idle":"2025-02-07T14:10:39.344826Z","shell.execute_reply.started":"2025-02-07T14:10:35.207950Z","shell.execute_reply":"2025-02-07T14:10:39.344065Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model","metadata":{}},{"cell_type":"code","source":"class RSNA24Model(nn.Module):\n    def __init__(self, model_name, in_c=30, n_classes=75, pretrained=True, features_only=False):\n        super().__init__()\n        self.model = timm.create_model(\n                                    model_name,\n                                    pretrained=pretrained, \n                                    features_only=features_only,\n                                    in_chans=in_c,\n                                    num_classes=n_classes,\n                                    global_pool='avg'\n                                    )\n    \n    def forward(self, x):\n        y = self.model(x)\n        return y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:39.345733Z","iopub.execute_input":"2025-02-07T14:10:39.346045Z","iopub.status.idle":"2025-02-07T14:10:39.351004Z","shell.execute_reply.started":"2025-02-07T14:10:39.346011Z","shell.execute_reply":"2025-02-07T14:10:39.350122Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"m = RSNA24Model(MODEL_NAME, in_c=IN_CHANS, n_classes=N_CLASSES, pretrained=False)\ni = torch.randn(2, IN_CHANS, 512, 512)\nout = m(i)\nfor o in out:\n    print(o.shape, o.min(), o.max())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:39.351810Z","iopub.execute_input":"2025-02-07T14:10:39.352021Z","iopub.status.idle":"2025-02-07T14:10:43.617670Z","shell.execute_reply.started":"2025-02-07T14:10:39.352002Z","shell.execute_reply":"2025-02-07T14:10:43.616651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"del m, i, out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:43.618484Z","iopub.execute_input":"2025-02-07T14:10:43.618777Z","iopub.status.idle":"2025-02-07T14:10:43.623054Z","shell.execute_reply.started":"2025-02-07T14:10:43.618755Z","shell.execute_reply":"2025-02-07T14:10:43.622042Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Training","metadata":{}},{"cell_type":"code","source":"#autocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.bfloat16) # if your gpu is newer Ampere, you can use this, lesser appearance of nan than half\nautocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.half) # you can use with T4 gpu. or newer\nscaler = torch.cuda.amp.GradScaler(enabled=USE_AMP, init_scale=4096)\n\nskf = KFold(n_splits=N_FOLDS, shuffle=True, random_state=SEED)\nfor fold, (trn_idx, val_idx) in enumerate(skf.split(range(len(df)))):\n    print('#'*30)\n    print(f'start fold{fold}')\n    print('#'*30)\n    print(len(trn_idx), len(val_idx))\n    df_train = df.iloc[trn_idx]\n    df_valid = df.iloc[val_idx]\n\n    train_ds = RSNA24Dataset(df_train, phase='train', transform=transforms_train)\n    train_dl = DataLoader(\n                train_ds,\n                batch_size=BATCH_SIZE,\n                shuffle=True,\n                pin_memory=True,\n                drop_last=True,\n                num_workers=N_WORKERS\n                )\n\n    valid_ds = RSNA24Dataset(df_valid, phase='valid', transform=transforms_val)\n    valid_dl = DataLoader(\n                valid_ds,\n                batch_size=BATCH_SIZE*2,\n                shuffle=False,\n                pin_memory=True,\n                drop_last=False,\n                num_workers=N_WORKERS\n                )\n\n    model = RSNA24Model(MODEL_NAME, IN_CHANS, N_CLASSES, pretrained=True)\n    model.to(device)\n    \n    optimizer = AdamW(model.parameters(), lr=LR, weight_decay=WD)\n\n    warmup_steps = EPOCHS/10 * len(train_dl) // GRAD_ACC\n    num_total_steps = EPOCHS * len(train_dl) // GRAD_ACC\n    num_cycles = 0.475\n    scheduler = get_cosine_schedule_with_warmup(optimizer,\n                                                num_warmup_steps=warmup_steps,\n                                                num_training_steps=num_total_steps,\n                                                num_cycles=num_cycles)\n\n    weights = torch.tensor([1.0, 2.0, 4.0])\n    criterion = nn.CrossEntropyLoss(weight=weights.to(device))\n    criterion2 = nn.CrossEntropyLoss(weight=weights)\n\n    best_loss = 1.2\n    best_wll = 1.2\n    es_step = 0\n\n    for epoch in range(1, EPOCHS+1):\n        print(f'start epoch {epoch}')\n        model.train()\n        total_loss = 0\n        with tqdm(train_dl, leave=True) as pbar:\n            optimizer.zero_grad()\n            for idx, (x, t) in enumerate(pbar):  \n                x = x.to(device)\n                t = t.to(device)\n                \n                with autocast:\n                    loss = 0\n                    y = model(x)\n                    for col in range(N_LABELS):\n                        pred = y[:,col*3:col*3+3]\n                        gt = t[:,col]\n                        loss = loss + criterion(pred, gt) / N_LABELS\n                        \n                    total_loss += loss.item()\n                    if GRAD_ACC > 1:\n                        loss = loss / GRAD_ACC\n    \n                if not math.isfinite(loss):\n                    print(f\"Loss is {loss}, stopping training\")\n                    sys.exit(1)\n    \n                pbar.set_postfix(\n                    OrderedDict(\n                        loss=f'{loss.item()*GRAD_ACC:.6f}',\n                        lr=f'{optimizer.param_groups[0][\"lr\"]:.3e}'\n                    )\n                )\n                scaler.scale(loss).backward()\n\n                torch.nn.utils.clip_grad_norm_(model.parameters(), MAX_GRAD_NORM or 1e9)\n                \n                if (idx + 1) % GRAD_ACC == 0:\n                    scaler.step(optimizer)\n                    scaler.update()\n                    optimizer.zero_grad()\n                    if scheduler is not None:\n                        scheduler.step()                    \n    \n        train_loss = total_loss/len(train_dl)\n        print(f'train_loss:{train_loss:.6f}')\n\n        total_loss = 0\n        y_preds = []\n        labels = []\n        \n        model.eval()\n        with tqdm(valid_dl, leave=True) as pbar:\n            with torch.no_grad():\n                for idx, (x, t) in enumerate(pbar):\n                    \n                    x = x.to(device)\n                    t = t.to(device)\n                        \n                    with autocast:\n                        loss = 0\n                        loss_ema = 0\n                        y = model(x)\n                        for col in range(N_LABELS):\n                            pred = y[:,col*3:col*3+3]\n                            gt = t[:,col]\n \n                            loss = loss + criterion(pred, gt) / N_LABELS\n                            y_pred = pred.float()\n                            y_preds.append(y_pred.cpu())\n                            labels.append(gt.cpu())\n                        \n                        total_loss += loss.item()   \n    \n        val_loss = total_loss/len(valid_dl)\n        \n        y_preds = torch.cat(y_preds, dim=0)\n        labels = torch.cat(labels)\n        val_wll = criterion2(y_preds, labels)\n        \n        print(f'val_loss:{val_loss:.6f}, val_wll:{val_wll:.6f}')\n\n        if val_loss < best_loss or val_wll < best_wll:\n            \n            es_step = 0\n\n            if device!='cuda:0':\n                model.to('cuda:0')                \n                \n            if val_loss < best_loss:\n                print(f'epoch:{epoch}, best loss updated from {best_loss:.6f} to {val_loss:.6f}')\n                best_loss = val_loss\n                \n            if val_wll < best_wll:\n                print(f'epoch:{epoch}, best wll_metric updated from {best_wll:.6f} to {val_wll:.6f}')\n                best_wll = val_wll\n                fname = f'{OUTPUT_DIR}/best_wll_model_fold-{fold}.pt'\n                torch.save(model.state_dict(), fname)\n            \n            if device!='cuda:0':\n                model.to(device)\n            \n        else:\n            es_step += 1\n            if es_step >= EARLY_STOPPING_EPOCH:\n                print('early stopping')\n                break  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T14:10:43.623988Z","iopub.execute_input":"2025-02-07T14:10:43.624214Z","iopub.status.idle":"2025-02-07T16:19:27.087667Z","shell.execute_reply.started":"2025-02-07T14:10:43.624195Z","shell.execute_reply":"2025-02-07T16:19:27.084040Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CV score","metadata":{}},{"cell_type":"code","source":"cv = 0\ny_preds = []\nlabels = []\nweights = torch.tensor([1.0, 2.0, 4.0])\ncriterion2 = nn.CrossEntropyLoss(weight=weights)\n\nfor fold, (trn_idx, val_idx) in enumerate(skf.split(range(len(df)))):\n    print('#'*30)\n    print(f'start fold{fold}')\n    print('#'*30)\n    df_valid = df.iloc[val_idx]\n    valid_ds = RSNA24Dataset(df_valid, phase='valid', transform=transforms_val)\n    valid_dl = DataLoader(\n                valid_ds,\n                batch_size=1,\n                shuffle=False,\n                pin_memory=True,\n                drop_last=False,\n                num_workers=N_WORKERS\n                )\n\n    model = RSNA24Model(MODEL_NAME, IN_CHANS, N_CLASSES, pretrained=False)\n    fname = f'{OUTPUT_DIR}/best_wll_model_fold-{fold}.pt'\n    model.load_state_dict(torch.load(fname))\n    model.to(device)   \n    \n    model.eval()\n    with tqdm(valid_dl, leave=True) as pbar:\n        with torch.no_grad():\n            for idx, (x, t) in enumerate(pbar):\n                \n                x = x.to(device)\n                t = t.to(device)\n                    \n                with autocast:\n                    y = model(x)\n                    for col in range(N_LABELS):\n                        pred = y[:,col*3:col*3+3]\n                        gt = t[:,col] \n                        y_pred = pred.float()\n                        y_preds.append(y_pred.cpu())\n                        labels.append(gt.cpu())\n\ny_preds = torch.cat(y_preds)\nlabels = torch.cat(labels)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:19:34.190158Z","iopub.execute_input":"2025-02-07T16:19:34.190628Z","iopub.status.idle":"2025-02-07T16:20:27.614743Z","shell.execute_reply.started":"2025-02-07T16:19:34.190585Z","shell.execute_reply":"2025-02-07T16:20:27.613128Z"},"collapsed":true,"jupyter":{"outputs_hidden":true}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"cv = criterion2(y_preds, labels)\nprint('cv score:', cv.item())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:20:27.615481Z","iopub.status.idle":"2025-02-07T16:20:27.615772Z","shell.execute_reply":"2025-02-07T16:20:27.615665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import log_loss\ny_pred_np = y_preds.softmax(1).numpy()\nlabels_np = labels.numpy()\ny_pred_nan = np.zeros((y_preds.shape[0], 1))\ny_pred2 = np.concatenate([y_pred_nan, y_pred_np],axis=1)\nweights = []\nfor l in labels:\n    if l==0: weights.append(1)\n    elif l==1: weights.append(2)\n    elif l==2: weights.append(4)\n    else: weights.append(0)\ncv2 = log_loss(labels, y_pred2, normalize=True, sample_weight=weights)\nprint('cv score:', cv2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:20:27.616511Z","iopub.status.idle":"2025-02-07T16:20:27.616749Z","shell.execute_reply":"2025-02-07T16:20:27.616651Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"np.save(f'{OUTPUT_DIR}/labels.npy', labels_np)\nnp.save(f'{OUTPUT_DIR}/final_oof.npy', y_pred2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:19:27.092002Z","iopub.status.idle":"2025-02-07T16:19:27.092365Z","shell.execute_reply":"2025-02-07T16:19:27.092195Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Predict","metadata":{}},{"cell_type":"code","source":"random_pred = np.ones((y_preds.shape[0], 3)) / 3.0\ny_pred3 = np.concatenate([y_pred_nan, random_pred],axis=1)\ncv3 = log_loss(labels, y_pred3, normalize=True, sample_weight=weights)\nprint('random score:', cv3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:19:27.093209Z","iopub.status.idle":"2025-02-07T16:19:27.093502Z","shell.execute_reply":"2025-02-07T16:19:27.093398Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Submission","metadata":{}},{"cell_type":"code","source":"dft = pd.read_csv(f'{path}/test_series_descriptions.csv')\ndft.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:26:13.852851Z","iopub.execute_input":"2025-02-07T16:26:13.853153Z","iopub.status.idle":"2025-02-07T16:26:13.885865Z","shell.execute_reply.started":"2025-02-07T16:26:13.853132Z","shell.execute_reply":"2025-02-07T16:26:13.885201Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class RSNA24Test(Dataset):\n    def __init__(self, df, study_ids, phase='test', transform=None):\n        self.df = df\n        self.study_ids = study_ids\n        self.transform = transform\n        self.phase = phase\n    \n    def __len__(self):\n        return len(self.study_ids)\n    \n    def get_img_paths(self, study_id, series_desc):\n        pdf = self.df[self.df['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'{path}/test_images/{study_id}/{row[\"series_id\"]}/*.dcm')\n            pimgs = sorted(pimgs, key=natural_keys)\n            allimgs.extend(pimgs)\n            \n        return allimgs\n    \n    def read_dcm_ret_arr(self, src_path):\n        dicom_data = pydicom.dcmread(src_path)\n        image = dicom_data.pixel_array\n        image = (image - image.min()) / (image.max() - image.min() + 1e-6) * 255\n        img = cv2.resize(image, (IMG_SIZE[0], IMG_SIZE[1]),interpolation=cv2.INTER_CUBIC)\n        assert img.shape==(IMG_SIZE[0], IMG_SIZE[1])\n        return img\n\n    def __getitem__(self, idx):\n        x = np.zeros((IMG_SIZE[0], IMG_SIZE[1], IN_CHANS), dtype=np.uint8)\n        st_id = self.study_ids[idx]        \n        \n        # Sagittal T1\n        allimgs_st1 = self.get_img_paths(st_id, 'Sagittal T1')\n        if len(allimgs_st1)==0:\n            print(st_id, ': Sagittal T1, has no images')\n        \n        else:\n            step = len(allimgs_st1) / 10.0\n            st = len(allimgs_st1)/2.0 - 4.0*step\n            end = len(allimgs_st1)+0.0001\n            for j, i in enumerate(np.arange(st, end, step)):\n                try:\n                    ind2 = max(0, int((i-0.5001).round()))\n                    img = self.read_dcm_ret_arr(allimgs_st1[ind2])\n                    x[..., j] = img.astype(np.uint8)\n                except:\n                    print(f'failed to load on {st_id}, Sagittal T1')\n                    pass\n            \n        # Sagittal T2/STIR\n        allimgs_st2 = self.get_img_paths(st_id, 'Sagittal T2/STIR')\n        if len(allimgs_st2)==0:\n            print(st_id, ': Sagittal T2/STIR, has no images')\n            \n        else:\n            step = len(allimgs_st2) / 10.0\n            st = len(allimgs_st2)/2.0 - 4.0*step\n            end = len(allimgs_st2)+0.0001\n            for j, i in enumerate(np.arange(st, end, step)):\n                try:\n                    ind2 = max(0, int((i-0.5001).round()))\n                    img = self.read_dcm_ret_arr(allimgs_st2[ind2])\n                    x[..., j+10] = img.astype(np.uint8)\n                except:\n                    print(f'failed to load on {st_id}, Sagittal T2/STIR')\n                    pass\n            \n        # Axial T2\n        allimgs_at2 = self.get_img_paths(st_id, 'Axial T2')\n        if len(allimgs_at2)==0:\n            print(st_id, ': Axial T2, has no images')\n            \n        else:\n            step = len(allimgs_at2) / 10.0\n            st = len(allimgs_at2)/2.0 - 4.0*step\n            end = len(allimgs_at2)+0.0001\n\n            for j, i in enumerate(np.arange(st, end, step)):\n                try:\n                    ind2 = max(0, int((i-0.5001).round()))\n                    img = self.read_dcm_ret_arr(allimgs_at2[ind2])\n                    x[..., j+20] = img.astype(np.uint8)\n                except:\n                    print(f'failed to load on {st_id}, Axial T2')\n                    pass  \n            \n            \n        if self.transform is not None:\n            x = self.transform(image=x)['image']\n\n        x = x.transpose(2, 0, 1)\n                \n        return x, str(st_id)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:26:15.665157Z","iopub.execute_input":"2025-02-07T16:26:15.665516Z","iopub.status.idle":"2025-02-07T16:26:15.678527Z","shell.execute_reply.started":"2025-02-07T16:26:15.665488Z","shell.execute_reply":"2025-02-07T16:26:15.677648Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transforms_test = A.Compose([\n    A.Resize(IMG_SIZE[0], IMG_SIZE[1]),\n    A.Normalize(mean=0.5, std=0.5)\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:20:36.374152Z","iopub.execute_input":"2025-02-07T16:20:36.374519Z","iopub.status.idle":"2025-02-07T16:20:36.396385Z","shell.execute_reply.started":"2025-02-07T16:20:36.374487Z","shell.execute_reply":"2025-02-07T16:20:36.395478Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"study_ids = list(dft['study_id'].unique())\nsample_sub = pd.read_csv(f'{path}/sample_submission.csv')\nLABELS = list(sample_sub.columns[1:])\nLABELS","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:38:35.655117Z","iopub.execute_input":"2025-02-07T16:38:35.655483Z","iopub.status.idle":"2025-02-07T16:38:35.672134Z","shell.execute_reply.started":"2025-02-07T16:38:35.655457Z","shell.execute_reply":"2025-02-07T16:38:35.671483Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_ds = RSNA24Test(dft, study_ids, transform=transforms_test)\ntest_dl = DataLoader(\n    test_ds, \n    batch_size=1, \n    shuffle=False,\n    num_workers=N_WORKERS,\n    pin_memory=True,\n    drop_last=False\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:26:19.496469Z","iopub.execute_input":"2025-02-07T16:26:19.496795Z","iopub.status.idle":"2025-02-07T16:26:19.502445Z","shell.execute_reply.started":"2025-02-07T16:26:19.496765Z","shell.execute_reply":"2025-02-07T16:26:19.501455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"models = []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:26:21.585011Z","iopub.execute_input":"2025-02-07T16:26:21.585359Z","iopub.status.idle":"2025-02-07T16:26:21.610417Z","shell.execute_reply.started":"2025-02-07T16:26:21.585328Z","shell.execute_reply":"2025-02-07T16:26:21.609729Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import glob\nCKPT_PATHS = glob.glob('/kaggle/working/rsna24-results/best_wll_model_fold-*.pt')\nCKPT_PATHS = sorted(CKPT_PATHS)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:37:18.820378Z","iopub.execute_input":"2025-02-07T16:37:18.820760Z","iopub.status.idle":"2025-02-07T16:37:18.826144Z","shell.execute_reply.started":"2025-02-07T16:37:18.820716Z","shell.execute_reply":"2025-02-07T16:37:18.825325Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for i, cp in enumerate(CKPT_PATHS):\n    print(f'loading {cp}...')\n    model = RSNA24Model(MODEL_NAME, IN_CHANS, N_CLASSES, pretrained=False)\n    model.load_state_dict(torch.load(cp))\n    model.eval()\n    model.half()\n    model.to(device)\n    models.append(model)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:26:25.844781Z","iopub.execute_input":"2025-02-07T16:26:25.845072Z","iopub.status.idle":"2025-02-07T16:27:13.130544Z","shell.execute_reply.started":"2025-02-07T16:26:25.845048Z","shell.execute_reply":"2025-02-07T16:27:13.129841Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"autocast = torch.cuda.amp.autocast(enabled=USE_AMP, dtype=torch.half)\ny_preds = []\nrow_names = []\n\nwith tqdm(test_dl, leave=True) as pbar:\n    with torch.no_grad():\n        for idx, (x, si) in enumerate(pbar):\n            x = x.to(device)\n            pred_per_study = np.zeros((25, 3))\n            \n            for cond in CONDITIONS:\n                for level in LEVELS:\n                    row_names.append(si[0] + '_' + cond + '_' + level)\n            \n            with autocast:\n                for m in models:\n                    y = m(x)[0]\n                    for col in range(N_LABELS):\n                        pred = y[col*3:col*3+3]\n                        y_pred = pred.float().softmax(0).cpu().numpy()\n                        pred_per_study[col] += y_pred / len(models)\n                y_preds.append(pred_per_study)\n\ny_preds = np.concatenate(y_preds, axis=0)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:38:41.221749Z","iopub.execute_input":"2025-02-07T16:38:41.222037Z","iopub.status.idle":"2025-02-07T16:38:41.969214Z","shell.execute_reply.started":"2025-02-07T16:38:41.222014Z","shell.execute_reply":"2025-02-07T16:38:41.968209Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub = pd.DataFrame()\nsub['row_id'] = row_names\nsub[LABELS] = y_preds\nsub.head(25)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:39:03.242073Z","iopub.execute_input":"2025-02-07T16:39:03.242469Z","iopub.status.idle":"2025-02-07T16:39:03.259936Z","shell.execute_reply.started":"2025-02-07T16:39:03.242416Z","shell.execute_reply":"2025-02-07T16:39:03.259097Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"sub.to_csv('submission.csv', index=False)\npd.read_csv('submission.csv').head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-02-07T16:39:05.181799Z","iopub.execute_input":"2025-02-07T16:39:05.182092Z","iopub.status.idle":"2025-02-07T16:39:05.209705Z","shell.execute_reply.started":"2025-02-07T16:39:05.182070Z","shell.execute_reply":"2025-02-07T16:39:05.208792Z"}},"outputs":[],"execution_count":null}]}