{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","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,"isSourceIdPinned":false,"sourceType":"competition"},{"sourceId":14535717,"sourceType":"datasetVersion","datasetId":9009659},{"sourceId":14697941,"sourceType":"datasetVersion","datasetId":9388508}],"dockerImageVersionId":31260,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport gc\nimport cv2\nimport math\nimport random\nimport numpy as np\nimport pandas as pd\nimport pydicom\nimport torch\nimport torch.nn as nn\nimport torchvision\nimport timm\nimport albumentations as A\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.metrics import accuracy_score, f1_score, roc_auc_score, classification_report\n\n# ==========================================\n# 1. INFERENCE CONFIGURATION\n# ==========================================\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nSEED = 101\n\n# --- HYPERPARAMETERS (Original Model) ---\nPATCH_H, PATCH_W = 512, 512\nAXIAL_SIZE = 384\npatch_size = 64\nBS = 48\nLmax = 15  # Fixed: Original model uses 15 slices\n\n# --- MODEL PATHS (Update these for the 15-slice model) ---\n# Path to the fully trained ORIGINAL model (15 slices)\nMAIN_MODEL_PATH = \"/kaggle/input/ablation-models-for-btp/DualView_ViT_Fold1_final.pth\" \n\n# Path to the pretrained Sagittal Discriminator (Specific fold used during training)\nSAG_ENCODER_PATH = \"/kaggle/input/lumbar-spine-keypoint-detection-models/Sagittal_T2_spine_discriminator_1\"\n\n# Path to the pretrained Axial weights\nAXIAL_WEIGHTS_PATH = \"/kaggle/input/lumbar-spine-keypoint-detection-models/Axial_PreTrain_EffNetV2.pth\"\n\n# Data Paths\nIMAGES_DIR = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images'\nLABELS_PATH = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\"\nCOORDS_PATH = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv\"\nDESC_PATH = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\"\n\nLEVELS = ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']\nLABELS_MAP = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\nINV_LABELS_MAP = {0: 'Normal/Mild', 1: 'Moderate', 2: 'Severe'}\n\ndef seed_everything(seed):\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)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n\nseed_everything(SEED)\n\n# ==========================================\n# 2. DATA PREPARATION\n# ==========================================\nprint(\"--> Parsing Dataset (Original 15-slice config)...\")\ndf_coors = pd.read_csv(COORDS_PATH)\nlabels_df = pd.read_csv(LABELS_PATH)\ntrain_desc = pd.read_csv(DESC_PATH)\n\ndf = df_coors.merge(train_desc, on=['study_id', 'series_id'], how='left')\ndf = df[df['series_description'] != 'Sagittal T1']\n\ndata = dict()   \nfor i in range(df.shape[0]):\n    study_id = str(df.iloc[i]['study_id'])\n    series_id = str(df.iloc[i]['series_id'])\n    instance_number = str(df.iloc[i]['instance_number'])\n    level = str(df.iloc[i]['level'])\n    x, y = df.iloc[i]['x'], df.iloc[i]['y']\n    \n    val_desc = str(df.iloc[i]['series_description'])\n    desc = val_desc.split()[0].lower() if len(val_desc.split()) > 0 else 'unknown'\n    \n    if study_id not in data:\n        data[study_id] = dict()\n\n    if level not in data[study_id]:\n        data[study_id][level] = {'sagittal': {}, 'axial': []}\n        \n    full_path = os.path.join(IMAGES_DIR, str(study_id), str(series_id), str(instance_number) + '.dcm')\n    \n    if desc == 'axial':\n        data[study_id][level]['axial'].append({'x': x, 'y': y, 'path': full_path})\n    else:\n        data[study_id][level]['sagittal'] = {'x': x, 'y': y, 'path': full_path}\n\nrows = []\nfor study_id, levels_dict in data.items():\n    for level, views in levels_dict.items():\n        row = {'study_id': study_id, 'level': level}\n        if 'sagittal' in views and views['sagittal']:\n            row['sag_x'] = views['sagittal']['x']\n            row['sag_y'] = views['sagittal']['y']\n            row['sag_path'] = views['sagittal']['path']\n        if 'axial' in views:\n            for idx, ax_slice in enumerate(views['axial']):\n                if idx > 1: break \n                row[f'ax_{idx}_x'] = ax_slice['x']\n                row[f'ax_{idx}_y'] = ax_slice['y']\n                row[f'ax_{idx}_path'] = ax_slice['path']\n        rows.append(row)\n\nsot_df = pd.DataFrame(rows)\nlabels_df['study_id'] = labels_df['study_id'].astype(str)\nprint(f\"--> Dataset parsed. Found {len(sot_df)} study-level rows.\")\n\n# ==========================================\n# 3. HELPER FUNCTIONS & DATASET\n# ==========================================\ndef load_axial_slice(path, x, y):\n    if pd.isna(path) or not isinstance(path, str) or not os.path.exists(path):\n        return torch.zeros(1, AXIAL_SIZE, AXIAL_SIZE)\n    try:\n        dcm = pydicom.dcmread(path)\n        img = dcm.pixel_array.astype(np.float32)\n        img = img - img.min()\n        v_max = np.percentile(img, 99)\n        img = (img / v_max if v_max > 0 else img / (img.max() + 1e-6))\n        img = np.clip(img, 0, 1)\n        img = (img * 255).astype(np.uint8)\n\n        h, w = img.shape\n        cx, cy = (int(x), int(y)) if (pd.notna(x) and pd.notna(y)) else (w//2, h//2)\n        pad_h, pad_w = int(0.15 * h), int(0.15 * w)\n        \n        crop = img[max(0, cy-pad_h):min(h, cy+pad_h), max(0, cx-pad_w):min(w, cx+pad_w)]\n        if crop.size == 0: return torch.zeros(1, AXIAL_SIZE, AXIAL_SIZE)\n        \n        img_resized = cv2.resize(crop, (AXIAL_SIZE, AXIAL_SIZE), interpolation=cv2.INTER_LINEAR)\n        return torch.from_numpy(img_resized).unsqueeze(0).float() / 255.0\n    except Exception:\n        return torch.zeros(1, AXIAL_SIZE, AXIAL_SIZE)\n\nclass DualView_Spinal_Dataset(Dataset):\n    def __init__(self, sot_df, label_df=None, VALID=False, augment=False, P=patch_size):\n        self.sot_df = sot_df.copy()\n        self.sot_df['study_id'] = self.sot_df['study_id'].astype(str)\n        self.VALID = VALID\n        self.augment = augment \n        self.P = P\n        self.sag_resize = torchvision.transforms.Resize((PATCH_H, PATCH_W), antialias=True)\n        \n        self.label_lookup = {}\n        if label_df is not None:\n            self.label_lookup = label_df.copy().set_index('study_id').to_dict('index')\n\n        self.study_groups = self.sot_df.groupby('study_id')\n        self.study_ids = list(self.study_groups.groups.keys())\n\n    def __len__(self): return len(self.study_ids)\n\n    def __getitem__(self, index):\n        study_id = self.study_ids[index]\n        group = self.study_groups.get_group(study_id).set_index('level')\n        \n        sag_imgs = torch.zeros(5, Lmax, 2*self.P, 2*self.P)\n        sag_masks = torch.ones(5, Lmax).bool()\n        ax_imgs = torch.zeros(5, 2, 1, AXIAL_SIZE, AXIAL_SIZE)\n        targets = []\n\n        lbl_row = self.label_lookup.get(study_id, {})\n\n        for idx, level in enumerate(LEVELS):\n            val = lbl_row.get(f\"spinal_canal_stenosis_{level.replace('/', '_').lower()}\", 'UNK')\n            targets.append(LABELS_MAP.get(val, -100))\n\n            if level not in group.index: continue\n            row = group.loc[level]\n\n            ax_imgs[idx, 0] = load_axial_slice(row['ax_0_path'], row['ax_0_x'], row['ax_0_y'])\n            ax_imgs[idx, 1] = load_axial_slice(row['ax_1_path'], row['ax_1_x'], row['ax_1_y'])\n\n            if pd.notna(row['sag_path']) and os.path.exists(row['sag_path']):\n                try:\n                    self._load_sagittal_volume(row, sag_imgs, sag_masks, idx)\n                except: pass\n\n        c = self.P // 2\n        sag_imgs = sag_imgs[:, :, c:c+self.P, c:c+self.P]\n\n        return [sag_imgs, sag_masks, ax_imgs], torch.tensor(targets, dtype=torch.long)\n\n    def _load_sagittal_volume(self, row, sag_imgs, sag_masks, idx):\n        dcm_ref = pydicom.dcmread(row['sag_path'])\n        img_ref = dcm_ref.pixel_array.astype(np.float32)\n        H_orig, W_orig = img_ref.shape\n        \n        cx, cy = (row['sag_x'], row['sag_y']) if pd.notna(row['sag_x']) else (W_orig/2, H_orig/2)\n        \n        transpose = False\n        if H_orig > W_orig:\n            cy -= (H_orig - W_orig) // 2\n            H_orig, transpose = W_orig, True\n        elif H_orig < W_orig:\n            cx -= (W_orig - H_orig) // 2\n            W_orig = H_orig\n\n        sc_y = int(cy * PATCH_H / H_orig + self.P)\n        sc_x = int(cx * PATCH_W / W_orig + self.P)\n\n        folder = os.path.dirname(row['sag_path'])\n        try:\n            center_inst = int(os.path.basename(row['sag_path']).split('.')[0])\n        except: return\n\n        for k in range(Lmax):\n            inst = center_inst - Lmax // 2 + k\n            fpath = os.path.join(folder, f\"{inst}.dcm\")\n            \n            if os.path.exists(fpath):\n                px = pydicom.dcmread(fpath).pixel_array.astype(np.float32)\n                v_max = np.quantile(px, 0.99)\n                px = (px - px.min()) / (v_max + 1e-6)\n                px = torch.tensor(np.clip(px, 0, 1))\n                \n                if transpose: \n                    diff = px.shape[0] - px.shape[1]\n                    start = diff // 2\n                    px = px[start : start + px.shape[1], :]\n                elif px.shape[0] < px.shape[1]: \n                    diff = px.shape[1] - px.shape[0]\n                    start = diff // 2\n                    px = px[:, start : start + px.shape[0]]\n\n                px = self.sag_resize(px.unsqueeze(0))\n                px = nn.functional.pad(px, [self.P]*4, 'reflect').squeeze(0)\n                \n                if sc_y-self.P >= 0 and sc_x-self.P >= 0:\n                     if sc_y+self.P <= px.shape[0] and sc_x+self.P <= px.shape[1]:\n                        sag_imgs[idx, k] = px[sc_y-self.P:sc_y+self.P, sc_x-self.P:sc_x+self.P]\n                        sag_masks[idx, k] = False\n\n# ==========================================\n# 4. MODEL ARCHITECTURE\n# ==========================================\nclass Sagittal_T2_spine_Discriminator(nn.Module):\n    def __init__(self, dim=512):\n        super().__init__()\n        self.emb = torchvision.models.resnet18(weights=None)\n        self.emb.conv1 = nn.Conv2d(1, patch_size, kernel_size=(7, 7), stride=(2, 2), padding=(3, 3), bias=False)\n        self.emb.fc = nn.Identity()\n        self.proj_out = nn.Linear(dim, 2)\n\nclass SpatialDropout(nn.Module):\n    def __init__(self, drop=0.1):\n        super(SpatialDropout, self).__init__()\n        self.drop = drop\n    def forward(self, inputs, noise_shape=None):\n        outputs = inputs.clone()\n        if noise_shape is None:\n            noise_shape = (inputs.shape[0], inputs.shape[1], *([1] * (inputs.dim() - 2)))\n        if not self.training or self.drop == 0:\n            return inputs\n        else:\n            noises = inputs.new().resize_(noise_shape)\n            if self.drop == 1: noises.fill_(0.0)\n            else: noises.bernoulli_(1 - self.drop).div_(1 - self.drop)\n            noises = noises.expand_as(inputs)    \n            outputs.mul_(noises)\n            return outputs\n\nclass Axial_Feature_Extractor(nn.Module):\n    def __init__(self, weight_path=None):\n        super().__init__()\n        self.backbone = timm.create_model(\"efficientnetv2_rw_t.ra2_in1k\", pretrained=False, in_chans=1, num_classes=0)\n        for p in self.backbone.parameters(): p.requires_grad = False\n        self.proj = nn.Linear(self.backbone.num_features, 512)\n        \n        if weight_path and os.path.exists(weight_path):\n            state = torch.load(weight_path, map_location=\"cpu\", weights_only=True)\n            if \"state_dict\" in state: state = state[\"state_dict\"]\n            state = {k.replace(\"module.\", \"\"): v for k, v in state.items()}\n            self.load_state_dict(state, strict=False)\n\n    def forward(self, x):\n        return self.proj(self.backbone(x))\n\nclass DualView_ViT(nn.Module):\n    def __init__(self, axial_weight_path=None, dim=512, depth=6, head_size=64):\n        super().__init__()\n        self.sag_encoder = torchvision.models.resnet18(weights=None)\n        self.sag_encoder.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        self.sag_encoder.fc = nn.Identity()\n        self.sag_proj = nn.Linear(512, dim)\n        self.ax_encoder = Axial_Feature_Extractor(axial_weight_path)\n        self.ax_adapt = nn.Linear(512, dim)\n\n        self.register_buffer('slices_enc', self._get_sinusoid(Lmax, dim))\n        self.register_buffer('pos_enc', self._get_sinusoid(5, dim))\n        \n        encoder_layer = nn.TransformerEncoderLayer(d_model=dim, nhead=dim//head_size, dim_feedforward=4*dim, dropout=0.1, activation=nn.GELU(), batch_first=True, norm_first=True)\n        self.sag_slice_transformer = nn.TransformerEncoder(encoder_layer, depth)\n        self.level_transformer = nn.TransformerEncoder(encoder_layer, depth)\n        \n        self.cross_attn = nn.MultiheadAttention(dim, dim//head_size, batch_first=True)\n        self.norm_sag = nn.LayerNorm(dim)\n        self.norm_ax = nn.LayerNorm(dim)\n        self.s_dropout = SpatialDropout(0.1)\n        self.proj_out = nn.Linear(dim, 3)\n\n    def _get_sinusoid(self, n_pos, d_model):\n        position = torch.arange(0, n_pos, dtype=torch.float).unsqueeze(1)\n        div_term = torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model))\n        pe = torch.zeros(n_pos, d_model)\n        pe[:, 0::2] = torch.sin(position * div_term)\n        pe[:, 1::2] = torch.cos(position * div_term)\n        return pe\n\n    def forward(self, x):\n        sag_imgs, sag_masks, ax_imgs = x\n        B = sag_imgs.shape[0]\n\n        s = self.sag_encoder(sag_imgs.view(-1, 1, patch_size, patch_size))\n        s = self.sag_proj(s).view(B*5, Lmax, -1) + self.slices_enc\n        flat_mask = sag_masks.view(B*5, Lmax)\n        s = self.sag_slice_transformer(s, src_key_padding_mask=flat_mask)\n        s = s.masked_fill(flat_mask.unsqueeze(-1), 0)\n        s_emb = s.sum(1) / (~flat_mask).sum(1).unsqueeze(-1).clamp(min=1)\n\n        a = self.ax_encoder(ax_imgs.view(-1, 1, AXIAL_SIZE, AXIAL_SIZE))\n        a = self.ax_adapt(a).view(B*5, 2, -1)\n\n        q = self.norm_sag(s_emb.unsqueeze(1))\n        kv = self.norm_ax(a)\n        attn_out, _ = self.cross_attn(q, kv, kv)\n        fused = s_emb + attn_out.squeeze(1)\n\n        lvl_mask = (sag_masks.sum(2) == Lmax)\n        out = self.level_transformer(fused.view(B, 5, -1) + self.pos_enc, src_key_padding_mask=lvl_mask)\n        out = out.permute(0, 2, 1) \n        out = self.s_dropout(out)\n        out = out.permute(0, 2, 1) \n        return self.proj_out(out)\n\n# ==========================================\n# 5. INFERENCE LOOP & METRICS\n# ==========================================\ndef run_full_inference():\n    print(f\"--> Initializing Original Inference (Lmax={Lmax}) on Device: {device}\")\n    \n    # 1. Setup Data\n    ds = DualView_Spinal_Dataset(sot_df, labels_df, VALID=True, augment=False)\n    dl = DataLoader(ds, batch_size=BS, shuffle=False, num_workers=4, pin_memory=True)\n    \n    # 2. Setup Model\n    model = DualView_ViT(axial_weight_path=AXIAL_WEIGHTS_PATH).to(device)\n    \n    # 3. Load Weights\n    if os.path.exists(SAG_ENCODER_PATH):\n        try:\n            print(f\"--> Loading Sagittal Encoder from {SAG_ENCODER_PATH}\")\n            old = torch.load(SAG_ENCODER_PATH, map_location=device, weights_only=False)\n            model.sag_encoder.load_state_dict(old.emb.state_dict(), strict=False)\n        except Exception as e:\n            print(f\"Warning: Failed to load Sagittal Encoder: {e}\")\n\n    if os.path.exists(MAIN_MODEL_PATH):\n        print(f\"--> Loading Main Model weights from {MAIN_MODEL_PATH}\")\n        try:\n            state = torch.load(MAIN_MODEL_PATH, map_location=device)\n            model.load_state_dict(state, strict=True)\n        except Exception as e:\n             raise RuntimeError(f\"Failed to load main model weights: {e}\")\n    else:\n        raise FileNotFoundError(f\"Main model weights not found at {MAIN_MODEL_PATH}\")\n\n    model.eval()\n    \n    all_preds = []\n    all_targets = []\n    \n    print(\"--> Starting Prediction Loop...\")\n    with torch.no_grad():\n        for inputs, targets in dl:\n            sag_imgs, sag_masks, ax_imgs = inputs\n            inputs_device = [\n                sag_imgs.to(device),\n                sag_masks.to(device),\n                ax_imgs.to(device)\n            ]\n            \n            logits = model(inputs_device) # Shape: (B, 5, 3)\n            logits_flat = logits.view(-1, 3)\n            targets_flat = targets.view(-1)\n            \n            all_preds.append(logits_flat.cpu())\n            all_targets.append(targets_flat.cpu())\n            \n    all_preds = torch.cat(all_preds, dim=0)\n    all_targets = torch.cat(all_targets, dim=0)\n    \n    # --- SANITIZATION STEP 1: Fix Raw Logits ---\n    if torch.isnan(all_preds).any() or torch.isinf(all_preds).any():\n        print(\"Warning: NaNs/Infs in raw logits. Sanitizing...\")\n        all_preds = torch.nan_to_num(all_preds, nan=0.0, posinf=10.0, neginf=-10.0)\n\n    # --- METRIC CALCULATION START ---\n    print(\"\\n\" + \"=\"*40)\n    print(\"       INFERENCE RESULTS (ORIGINAL)       \")\n    print(\"=\"*40)\n\n    # 1. CALCULATE WEIGHTED LOSS\n    # Logits and targets are on CPU, so we use CPU for loss to avoid mismatch\n    loss_weights = torch.tensor([1.0, 2.0, 4.0])\n    criterion = nn.CrossEntropyLoss(weight=loss_weights, ignore_index=-100)\n    \n    try:\n        final_loss = criterion(all_preds, all_targets).item()\n        print(f\"Weighted CE Loss: {final_loss:.4f}\")\n    except Exception as e:\n        print(f\"Weighted CE Loss: N/A ({e})\")\n\n    # Filter targets for classification metrics\n    mask = all_targets != -100\n    preds_valid = all_preds[mask]\n    targs_valid = all_targets[mask]\n    \n    # Probabilities\n    probs_valid = torch.softmax(preds_valid, dim=1).numpy()\n    \n    # --- SANITIZATION STEP 2: Fix Probabilities ---\n    if np.isnan(probs_valid).any():\n        print(\"Warning: NaNs in probabilities. Replacing with uniform.\")\n        probs_valid = np.nan_to_num(probs_valid, nan=1.0/3.0)\n\n    preds_cls = preds_valid.argmax(dim=1).numpy()\n    targs_np = targs_valid.numpy()\n    \n    # 2. Accuracy\n    acc = accuracy_score(targs_np, preds_cls)\n    print(f\"Overall Accuracy: {acc:.4f}\")\n    \n    # 3. F1\n    f1_macro = f1_score(targs_np, preds_cls, average='macro')\n    print(f\"Macro F1 Score:   {f1_macro:.4f}\")\n    \n    # 4. AUROC\n    try:\n        auroc_ovr = roc_auc_score(targs_np, probs_valid, multi_class='ovr', average='macro')\n        print(f\"AUROC (OvR):      {auroc_ovr:.4f}\")\n    except Exception as e:\n        print(f\"AUROC (OvR):      N/A ({e})\")\n        \n    print(\"\\n--- Binary AUROC (Per Class) ---\")\n    for i, cls_name in INV_LABELS_MAP.items():\n        try:\n            binary_target = (targs_np == i).astype(int)\n            binary_prob = probs_valid[:, i]\n            auc_cls = roc_auc_score(binary_target, binary_prob)\n            print(f\"  {cls_name:<12}: {auc_cls:.4f}\")\n        except Exception as e:\n            print(f\"  {cls_name:<12}: N/A\")\n\n    print(\"\\n--- Detailed Classification Report ---\")\n    print(classification_report(targs_np, preds_cls, target_names=[INV_LABELS_MAP[0], INV_LABELS_MAP[1], INV_LABELS_MAP[2]], digits=4))\n    \n    print(\"=\"*40)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-01T14:06:08.662953Z","iopub.execute_input":"2026-02-01T14:06:08.6636Z","iopub.status.idle":"2026-02-01T14:06:21.571479Z","shell.execute_reply.started":"2026-02-01T14:06:08.663567Z","shell.execute_reply":"2026-02-01T14:06:21.570764Z"}},"outputs":[{"name":"stderr","text":"/usr/local/lib/python3.12/dist-packages/pydantic/_internal/_generate_schema.py:2249: UnsupportedFieldAttributeWarning: The 'repr' attribute with value False was provided to the `Field()` function, which has no effect in the context it was used. 'repr' is field-specific metadata, and can only be attached to a model field using `Annotated` metadata or by assignment. This may have happened because an `Annotated` type alias using the `type` statement was used, or if the `Field()` function was attached to a single member of a union type.\n  warnings.warn(\n/usr/local/lib/python3.12/dist-packages/pydantic/_internal/_generate_schema.py:2249: UnsupportedFieldAttributeWarning: The 'frozen' attribute with value True was provided to the `Field()` function, which has no effect in the context it was used. 'frozen' is field-specific metadata, and can only be attached to a model field using `Annotated` metadata or by assignment. This may have happened because an `Annotated` type alias using the `type` statement was used, or if the `Field()` function was attached to a single member of a union type.\n  warnings.warn(\n","output_type":"stream"},{"name":"stdout","text":"--> Parsing Dataset (Original 15-slice config)...\n--> Dataset parsed. Found 9765 study-level rows.\n","output_type":"stream"}],"execution_count":1},{"cell_type":"code","source":"run_full_inference()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-02-01T14:06:21.572533Z","iopub.execute_input":"2026-02-01T14:06:21.572841Z","iopub.status.idle":"2026-02-01T14:21:48.873058Z","shell.execute_reply.started":"2026-02-01T14:06:21.572813Z","shell.execute_reply":"2026-02-01T14:21:48.872043Z"}},"outputs":[{"name":"stdout","text":"--> Initializing Original Inference (Lmax=15) on Device: cuda\n","output_type":"stream"},{"name":"stderr","text":"/usr/local/lib/python3.12/dist-packages/torch/nn/modules/transformer.py:392: UserWarning: enable_nested_tensor is True, but self.use_nested_tensor is False because encoder_layer.norm_first was True\n  warnings.warn(\n","output_type":"stream"},{"name":"stdout","text":"--> Loading Sagittal Encoder from /kaggle/input/lumbar-spine-keypoint-detection-models/Sagittal_T2_spine_discriminator_1\n--> Loading Main Model weights from /kaggle/input/ablation-models-for-btp/DualView_ViT_Fold1_final.pth\n--> Starting Prediction Loop...\nWarning: NaNs/Infs in raw logits. Sanitizing...\n\n========================================\n       INFERENCE RESULTS (ORIGINAL)       \n========================================\nWeighted CE Loss: 0.2692\nOverall Accuracy: 0.9222\nMacro F1 Score:   0.7373\nAUROC (OvR):      0.9700\n\n--- Binary AUROC (Per Class) ---\n  Normal/Mild : 0.9799\n  Moderate    : 0.9442\n  Severe      : 0.9858\n\n--- Detailed Classification Report ---\n              precision    recall  f1-score   support\n\n Normal/Mild     0.9789    0.9603    0.9695      8664\n    Moderate     0.5216    0.5258    0.5237       736\n      Severe     0.6280    0.8404    0.7188       470\n\n    accuracy                         0.9222      9870\n   macro avg     0.7095    0.7755    0.7373      9870\nweighted avg     0.9281    0.9222    0.9243      9870\n\n========================================\n","output_type":"stream"}],"execution_count":2},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}