{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","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":11125637,"sourceType":"datasetVersion","datasetId":6938318}],"dockerImageVersionId":31012,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"!pip install -q iterative-stratification ","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:03:05.308623Z","iopub.execute_input":"2025-04-28T17:03:05.309288Z","iopub.status.idle":"2025-04-28T17:03:09.739696Z","shell.execute_reply.started":"2025-04-28T17:03:05.309263Z","shell.execute_reply":"2025-04-28T17:03:09.738796Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport timm\nimport torch\nimport random\nimport shutil\nimport pydicom\nimport numpy as np\nimport pandas as pd\nimport transformers\nfrom tqdm import tqdm\nimport torch.nn as nn\nfrom typing import List\nfrom torch import Tensor\nimport matplotlib.pyplot as plt\nimport torch.nn.functional as F\nimport torchvision.transforms.v2 as v2\nfrom torch.optim import AdamW, lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.utils.tensorboard import SummaryWriter","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:03:36.306696Z","iopub.execute_input":"2025-04-28T17:03:36.306942Z","iopub.status.idle":"2025-04-28T17:03:36.312388Z","shell.execute_reply.started":"2025-04-28T17:03:36.306926Z","shell.execute_reply":"2025-04-28T17:03:36.311735Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"seed = 210\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\ndef seed_everything(seed):\n    random.seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n    print('Finish seeding with seed {}'.format(seed))\n\nseed_everything(seed)\nprint('Training on device {}'.format(device))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:03:36.313282Z","iopub.execute_input":"2025-04-28T17:03:36.313565Z","iopub.status.idle":"2025-04-28T17:03:36.426835Z","shell.execute_reply.started":"2025-04-28T17:03:36.313537Z","shell.execute_reply":"2025-04-28T17:03:36.426086Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Load Files","metadata":{}},{"cell_type":"code","source":"train = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\")\ntrain_coor = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv\")\ntrain_series = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\")\ntrain_dummy = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\")\ntrain_meta = pd.read_csv(\"/kaggle/input/meta-csv/meta.csv\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:03:36.427731Z","iopub.execute_input":"2025-04-28T17:03:36.428063Z","iopub.status.idle":"2025-04-28T17:03:37.110124Z","shell.execute_reply.started":"2025-04-28T17:03:36.428043Z","shell.execute_reply":"2025-04-28T17:03:37.109356Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dummy = train_dummy.fillna(\"Normal/Mild\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:03:37.110872Z","iopub.execute_input":"2025-04-28T17:03:37.111143Z","iopub.status.idle":"2025-04-28T17:03:37.120216Z","shell.execute_reply.started":"2025-04-28T17:03:37.111120Z","shell.execute_reply":"2025-04-28T17:03:37.119457Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_coor = train_coor.merge(train_series[['study_id', 'series_id', 'series_description']], on=['study_id', 'series_id'])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:03:37.120945Z","iopub.execute_input":"2025-04-28T17:03:37.121143Z","iopub.status.idle":"2025-04-28T17:03:37.164146Z","shell.execute_reply.started":"2025-04-28T17:03:37.121128Z","shell.execute_reply":"2025-04-28T17:03:37.163480Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class SpineCoorDataset(Dataset):\n    def __init__(self, coor, meta, condition, mode):\n        if condition == 'scs':\n            self.coor = coor.loc[coor.condition == \"Spinal Canal Stenosis\"]\n        elif condition == 'nfn':\n            self.coor = coor.loc[coor.condition.isin([\n                'Left Neural Foraminal Narrowing',\n                'Right Neural Foraminal Narrowing'\n            ])]\n        elif condition == 'ss':\n            self.coor = coor.loc[coor.condition.isin([\n                'Left Subarticular Stenosis',\n                'Right Subarticular Stenosis'\n            ])]\n\n        g_coor = self.coor.groupby(['study_id']).count()\n        if condition == 'scs':\n            self.id = g_coor[g_coor.series_id == 5].reset_index().study_id.unique()\n        else:\n            self.id = g_coor[g_coor.series_id == 10].reset_index().study_id.unique()\n\n        if condition == 'ss':\n            self.resize = v2.Resize((256, 256))\n        else:\n            self.resize = v2.Resize((384, 384)) # it can be adjusted accordingly to system resources\n\n        self.id = list(set(self.id) - set([3637444890]))\n\n        self.condition = condition\n        self.meta = meta\n        self.mode = mode\n\n    def __len__(self):\n        return len(self.id)\n\n    def __getitem__(self, idx):\n        study_id = self.id[idx]\n        if self.condition == 'scs':\n            volume, label = self.volume_scs(study_id)\n        elif self.condition == 'nfn':\n            volume, label = self.volume_nfn(study_id)\n        elif self.condition == 'ss':\n            volume, label = self.volume_ss(study_id)\n\n        return volume, label\n\n    def volume_scs(self, study_id):\n        all_levels = ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']\n        meta = self.meta.loc[(self.meta.study_id == study_id) & (self.meta.series_description == 'Sagittal T2/STIR')]\n        meta = meta.sort_values('ipp_x', ascending=True).reset_index(drop=True)\n        coor = self.coor.loc[(self.coor.study_id == study_id) & (self.coor.series_description == 'Sagittal T2/STIR')]\n        coor_dict = {}\n        meta_list = []\n        for _, row in coor.iterrows():\n            series_id, instance_number = row.series_id, row.instance_number\n            meta_list.append(meta.loc[(meta.series_id == series_id) & (meta.instance_number == instance_number)])\n        sub_meta = pd.concat(meta_list)\n        idx = meta.loc[meta.ipp_x == sub_meta.ipp_x.median()].index[0]\n        # idx = meta.loc[meta.instance_number == sub_meta.instance_number.median()].index[0]\n        img_row = meta.loc[idx]\n        before_img_row = meta.loc[idx -1]\n        after_img_row = meta.loc[idx + 1]\n        img = self.normalize(self.load_dicom(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{img_row.series_id}/{img_row.instance_number}.dcm\"))\n        bimg = self.normalize(self.load_dicom(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{before_img_row.series_id}/{before_img_row.instance_number}.dcm\"))\n        aimg = self.normalize(self.load_dicom(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{after_img_row.series_id}/{after_img_row.instance_number}.dcm\"))\n        height, width = img.shape\n        for _, row in coor.iterrows():\n            series_id, instance_number = row.series_id, row.instance_number\n            # print(row.x)\n            x = row.x/width\n            # print(x)\n            y = row.y/height\n            coor_dict[row.level] = torch.tensor([x, y]).to(torch.float32)\n\n        updated_dict = {}\n        for level in all_levels:\n            if level in coor_dict:\n                updated_dict[level] = coor_dict[level]\n            else:\n                print(f\"Missing level '{level}' in study ID: {study_id}\")\n                raise ValueError(f\"Missing coordinate for level '{level}' in study ID: {study_id}\")\n        coor_dict = updated_dict\n\n        img = self.resize(torch.tensor(img[None, ...]))\n        bimg = self.resize(torch.tensor(bimg[None, ...]))\n        aimg = self.resize(torch.tensor(aimg[None, ...]))\n        img = torch.cat([bimg, img, aimg]).to(torch.float32)\n\n        return img, coor_dict\n\n\n    def volume_nfn(self, study_id):\n        all_levels = ['left_L1/L2', 'left_L2/L3', 'left_L3/L4', 'left_L4/L5', 'left_L5/S1', 'right_L1/L2', 'right_L2/L3', 'right_L3/L4', 'right_L4/L5', 'right_L5/S1']\n        meta = self.meta.loc[(self.meta.study_id == study_id) & (self.meta.series_description == 'Sagittal T1')]\n        meta = meta.sort_values('ipp_x', ascending=True).reset_index(drop=True)\n        coor = self.coor.loc[(self.coor.study_id == study_id) & (self.coor.series_description == 'Sagittal T1')]\n        coor_dict = {}\n        right_meta_list = []\n        left_meta_list = []\n        for _, row in coor.iterrows():\n            series_id, instance_number = row.series_id, row.instance_number\n            if row.condition == \"Right Neural Foraminal Narrowing\":\n                right_meta_list.append(meta.loc[(meta.series_id == series_id) & (meta.instance_number == instance_number)])\n            else:\n                left_meta_list.append(meta.loc[(meta.series_id == series_id) & (meta.instance_number == instance_number)])\n        right_sub_meta = pd.concat(right_meta_list)\n        left_sub_meta = pd.concat(left_meta_list)\n        right_idx = meta.loc[meta.ipp_x == right_sub_meta.ipp_x.median()].index[0]\n        left_idx = meta.loc[meta.ipp_x == left_sub_meta.ipp_x.median()].index[0]\n        # idx = meta.loc[meta.instance_number == sub_meta.instance_number.median()].index[0]\n        right_img_row = meta.loc[right_idx]\n        right_before_img_row = meta.loc[right_idx -1]\n        right_after_img_row = meta.loc[right_idx + 1]\n\n        left_img_row = meta.loc[left_idx]\n        left_before_img_row = meta.loc[left_idx -1]\n        left_after_img_row = meta.loc[left_idx + 1]\n\n        rimg = self.normalize(self.load_dicom(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{right_img_row.series_id}/{right_img_row.instance_number}.dcm\"))\n        rbimg = self.normalize(self.load_dicom(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{right_before_img_row.series_id}/{right_before_img_row.instance_number}.dcm\"))\n        ramig = self.normalize(self.load_dicom(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{right_after_img_row.series_id}/{right_after_img_row.instance_number}.dcm\"))\n\n        limg = self.normalize(self.load_dicom(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{left_img_row.series_id}/{left_img_row.instance_number}.dcm\"))\n        lbimg = self.normalize(self.load_dicom(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{left_before_img_row.series_id}/{left_before_img_row.instance_number}.dcm\"))\n        laimg = self.normalize(self.load_dicom(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{left_after_img_row.series_id}/{left_after_img_row.instance_number}.dcm\"))\n\n        rheight, rwidth = rimg.shape\n        lheight, lwidth = limg.shape\n\n        for _, row in coor.iterrows():\n            series_id, instance_number = row.series_id, row.instance_number\n            if row.condition == 'Right Neural Foraminal Narrowing':\n                x = row.x/rwidth\n                y = row.y/rheight\n                coor_dict['right_' + row.level] = torch.tensor([x, y]).to(torch.float32)\n            else:\n                x = row.x/lwidth\n                y = row.y/lheight\n                coor_dict['left_' + row.level] = torch.tensor([x, y]).to(torch.float32)\n\n        updated_dict = {}\n        for level in all_levels:\n            if level in coor_dict:\n                updated_dict[level] = coor_dict[level]\n            else:\n                print(f\"Missing level '{level}' in study ID: {study_id}\")\n                raise ValueError(f\"Missing coordinate for level '{level}' in study ID: {study_id}\")\n        coor_dict = updated_dict\n\n        rimg = self.resize(torch.tensor(rimg[None, ...]))\n        rbimg = self.resize(torch.tensor(rbimg[None, ...]))\n        ramig = self.resize(torch.tensor(ramig[None, ...]))\n\n        rimg = torch.cat([rbimg, rimg, ramig]).to(torch.float32)\n\n        limg = self.resize(torch.tensor(limg[None, ...]))\n        lbimg = self.resize(torch.tensor(lbimg[None, ...]))\n        laimg = self.resize(torch.tensor(laimg[None, ...]))\n\n        limg = torch.cat([lbimg, limg, laimg]).to(torch.float32)\n\n        img = torch.stack([limg, rimg]).to(torch.float32).contiguous()\n\n        return img, coor_dict\n\n    def volume_ss(self, study_id):\n        all_levels = ['left_L1/L2', 'left_L2/L3', 'left_L3/L4', 'left_L4/L5', 'left_L5/S1', 'right_L1/L2', 'right_L2/L3', 'right_L3/L4', 'right_L4/L5', 'right_L5/S1']\n        meta = self.meta.loc[(self.meta.study_id == study_id) & (self.meta.series_description == 'Axial T2')]\n        meta = meta.sort_values('ipp_z', ascending=True).reset_index(drop=True)\n        coor = self.coor.loc[(self.coor.study_id == study_id) & (self.coor.series_description == 'Axial T2')]\n        coor_dict = {}\n        img_dict = {}\n        for _, row in coor.iterrows():\n            series_id, instance_number = row.series_id, row.instance_number\n            img = self.load_dicom(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{series_id}/{instance_number}.dcm\")\n            height, width = img.shape\n            img = self.resize(torch.tensor(img[None, ...]))\n            img = self.normalize(img.to(torch.float32))\n            x = row.x/width\n            y = row.y/height\n            if row.condition == 'Left Subarticular Stenosis':\n                coor_dict['left_' + row.level] = torch.tensor([x, y]).to(torch.float32)\n                img_dict['left_' + row.level] = img\n            else:\n                coor_dict['right_' + row.level] = torch.tensor([x, y]).to(torch.float32)\n                img_dict['right_' + row.level] = img\n            \n\n        updated_dict = {}\n        img_list = []\n        for level in all_levels:\n            if level in coor_dict:\n                updated_dict[level] = coor_dict[level]\n                img_list.append(img_dict[level])\n            else:\n                print(f\"Missing level '{level}' in study ID: {study_id}\")\n                raise ValueError(f\"Missing coordinate for level '{level}' in study ID: {study_id}\")\n        coor_dict = updated_dict\n\n        volume = torch.stack(img_list).contiguous()\n\n        return volume, coor_dict\n\n    def normalize(self, x):\n        if self.condition == 'ss':\n            lower, upper = torch.quantile(x, torch.tensor(0.01)), torch.quantile(x, torch.tensor(0.99))\n            x = torch.clamp(x, lower, upper)\n            x = x - torch.min(x)\n            x = x / torch.max(x)\n        else:\n            lower, upper = np.percentile(x, (1, 99))\n            x = np.clip(x, lower, upper)\n            x = x - np.min(x)\n            x = x / np.max(x)\n        return x\n\n    def load_dicom(self, path):\n        return pydicom.dcmread(path).pixel_array","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:03:37.166057Z","iopub.execute_input":"2025-04-28T17:03:37.166259Z","iopub.status.idle":"2025-04-28T17:03:37.192810Z","shell.execute_reply.started":"2025-04-28T17:03:37.166238Z","shell.execute_reply":"2025-04-28T17:03:37.192169Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\n\n# ----------------- Gated Attention Block (with Residual) -----------------\nclass GatedAttention(nn.Module):\n    def __init__(self, in_features):\n        super().__init__()\n        self.gate = nn.Sequential(\n            nn.Linear(in_features, in_features),\n            nn.Sigmoid()\n        )\n    def forward(self, x):\n        gate = self.gate(x)\n        return x * gate + x  # Residual connection\n\n# ----------------- MobileNetV3Small Spine Detection Model -----------------\nclass MobileNetV3GatedSpine(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        # Load MobileNetV3 Small backbone\n        self.encoder = timm.create_model(\n            'convnext_base.fb_in22k_ft_in1k_384',\n            in_chans=3,\n            pretrained=True,\n            features_only=False,\n            num_classes=0\n        )\n\n\n        # Feature dimension\n        self.in_features = self.encoder.num_features\n        print(self.in_features)\n\n        # Adaptive Pooling + Flatten\n        self.flatten = nn.Sequential(\n            nn.AdaptiveAvgPool2d((1,1)),\n            nn.Flatten(1)\n        )\n\n        # Gated Attention after encoder output\n        self.attention = GatedAttention(self.in_features)\n\n        # Project feature\n        self.projector = nn.Sequential(\n            nn.Linear(self.in_features, 1024),\n            nn.SiLU(inplace=True)\n        )\n\n        # Dropout for regularization\n        self.dropout = nn.Dropout(p=0.2)\n\n        # Heads for 5 vertebral levels\n        self.heads = nn.ModuleList([\n            nn.Linear(1024, 2) for _ in range(5)\n        ])\n\n    def forward(self, x):\n        # Feature extraction\n        x = self.encoder.forward_features(x)\n        x = self.flatten(x)\n\n        # Gated Attention\n        x = self.attention(x)\n\n        # Project to feature dimension\n        x = self.projector(x)\n\n        # Small dropout\n        x = self.dropout(x)\n\n        # Predict each level\n        output = {}\n        levels = ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']\n\n        for i, level in enumerate(levels):\n            output[level] = self.heads[i](x).sigmoid()\n        \n        return output\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T16:08:17.784070Z","iopub.execute_input":"2025-04-28T16:08:17.784797Z","iopub.status.idle":"2025-04-28T16:08:17.792606Z","shell.execute_reply.started":"2025-04-28T16:08:17.784771Z","shell.execute_reply":"2025-04-28T16:08:17.791824Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Validation","metadata":{}},{"cell_type":"code","source":"from iterstrat.ml_stratifiers import MultilabelStratifiedKFold\n\nlabel = train.columns[1:]\ntrain_dummy['fold'] = -1  # Initialize before assigning\nkfold = MultilabelStratifiedKFold(n_splits=5, shuffle=True, random_state=42)\nfor i, (train_idx, valid_idx) in enumerate(kfold.split(train_dummy, train_dummy[label])):\n    train_dummy.loc[valid_idx, 'fold'] = i\ntrain_series = train_series.merge(train_dummy[['study_id', 'fold']], on='study_id')\ntrain_coor = train_coor.merge(train_dummy[['study_id', 'fold']], on='study_id')\ntrain = train.merge(train_dummy[['study_id', 'fold']], on='study_id')\ntrain_meta = train_meta.merge(train_dummy[['study_id', 'fold']], on='study_id')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:03:37.193381Z","iopub.execute_input":"2025-04-28T17:03:37.193564Z","iopub.status.idle":"2025-04-28T17:03:37.737218Z","shell.execute_reply.started":"2025-04-28T17:03:37.193550Z","shell.execute_reply":"2025-04-28T17:03:37.736670Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loss","metadata":{}},{"cell_type":"code","source":"class SCSLoss(nn.Module):\n    def __init__(self, condition=\"scs\"):\n        super(SCSLoss, self).__init__()\n        self.condition = condition\n\n    def forward(self, outputs, targets):\n        loss = 0\n        count = 0\n        if self.condition == 'scs':\n            expected_level = ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']\n        else:\n            expected_level = ['left_L1/L2', 'left_L2/L3', 'left_L3/L4', 'left_L4/L5', 'left_L5/S1', 'right_L1/L2', 'right_L2/L3', 'right_L3/L4', 'right_L4/L5', 'right_L5/S1']\n        for level in expected_level:\n            if level in targets and level in outputs:\n                _loss = nn.functional.l1_loss(outputs[level], targets[level])\n                loss += _loss\n                count += 1\n\n        # Return average loss, avoiding division by zero\n        return loss / max(count, 1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T18:18:03.810515Z","iopub.execute_input":"2025-04-28T18:18:03.811448Z","iopub.status.idle":"2025-04-28T18:18:03.818099Z","shell.execute_reply.started":"2025-04-28T18:18:03.811412Z","shell.execute_reply":"2025-04-28T18:18:03.817294Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Score","metadata":{}},{"cell_type":"code","source":"def calculate_mse_score(outputs, targets):\n    \"\"\"Computes Mean Squared Error (MSE) for model predictions.\"\"\"\n    total_mse = 0\n    total_samples = 0\n    \n    # Only calculate for levels present in both outputs and targets\n    for level in set(outputs.keys()).intersection(targets.keys()):\n        # Ensure shapes match\n        if outputs[level].shape == targets[level].shape:\n            mse = nn.functional.mse_loss(outputs[level], targets[level], reduction='sum')\n            total_mse += mse.item()\n            total_samples += targets[level].numel()\n    \n    # Return average, avoiding division by zero\n    return total_mse / max(total_samples, 1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T18:18:04.428902Z","iopub.execute_input":"2025-04-28T18:18:04.429327Z","iopub.status.idle":"2025-04-28T18:18:04.433962Z","shell.execute_reply.started":"2025-04-28T18:18:04.429306Z","shell.execute_reply":"2025-04-28T18:18:04.433253Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Tolerance Matrix","metadata":{}},{"cell_type":"code","source":"def calculate_regression_tolerance(preds, targets, tolerances=[0, 1, 2]):\n    \"\"\"\n    Calculate regression tolerance matrix showing what percentage of predictions\n    fall within various tolerance thresholds.\n    \n    Args:\n        preds: Dictionary of model predictions\n        targets: Dictionary of ground truth values\n        tolerances: List of tolerance thresholds in mm\n        \n    Returns:\n        Dictionary with tolerance percentages for each threshold\n    \"\"\"\n    tolerance_counts = {f\"±{tol}\": 0 for tol in tolerances}\n    tolerance_counts[\">±2\"] = 0\n    total = 0\n\n    for level in preds:\n        if level not in targets:\n            continue\n\n        pred_vals = preds[level].detach().cpu().numpy().round().astype(int)\n        true_vals = targets[level].detach().cpu().numpy().astype(int)\n\n        for i in range(len(pred_vals)):\n            for j in range(len(pred_vals[i])):\n                diff = abs(pred_vals[i][j] - true_vals[i][j])\n                matched = False\n                for tol in tolerances:\n                    if diff <= tol:\n                        tolerance_counts[f\"±{tol}\"] += 1\n                        matched = True\n                        break\n                if not matched:\n                    tolerance_counts[\">±2\"] += 1\n                total += 1\n\n    return tolerance_counts, total","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T18:18:06.242020Z","iopub.execute_input":"2025-04-28T18:18:06.242458Z","iopub.status.idle":"2025-04-28T18:18:06.248940Z","shell.execute_reply.started":"2025-04-28T18:18:06.242436Z","shell.execute_reply":"2025-04-28T18:18:06.248232Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def train_spine_model(model, train_coor, train_meta, n_folds=5, epochs=18, batch_size=4,\n                      learning_rate=0.001, weight_decay=0.0001, patience=5,\n                      mixed_precision=True, experiment_name=\"spine_model\",\n                      tolerances=[0, 1, 2], condition='scs'):\n    \"\"\"\n    Complete training function for spine model with improved practices\n\n    Args:\n        model: Your defined model\n        train_coor: DataFrame containing coordinate annotations\n        train_meta: DataFrame containing metadata\n        n_folds: Number of folds for cross-validation\n        epochs: Number of training epochs\n        batch_size: Batch size for training\n        learning_rate: Initial learning rate\n        weight_decay: Weight decay for optimizer\n        patience: Patience for early stopping\n        mixed_precision: Whether to use mixed precision training\n        experiment_name: Name for experiment logs\n        tolerances: List of tolerance thresholds for regression tolerance calculation\n    \"\"\"\n\n    # Create output directory for models\n    os.makedirs(\"models\", exist_ok=True)\n\n    # Initialize tensorboard writer\n    writer = SummaryWriter(f'runs/{experiment_name}')\n\n    # Print training configuration\n    print(f\"\\n===== TRAINING CONFIGURATION =====\")\n    print(f\"Number of folds: {n_folds}\")\n    print(f\"Epochs: {epochs}\")\n    print(f\"Batch size: {batch_size}\")\n    print(f\"Learning rate: {learning_rate}\")\n    print(f\"Weight decay: {weight_decay}\")\n    print(f\"Mixed precision: {mixed_precision}\")\n    print(f\"Tolerance thresholds: {tolerances}\")\n    print(f\"==============================\\n\")\n\n    # Initialize scaler for mixed precision\n    scaler = torch.amp.GradScaler() if mixed_precision else None\n\n    # Cross-validation loop\n    all_val_mses = []\n\n    for fold in range(n_folds):\n        print(f\"\\n{'='*20} FOLD {fold+1}/{n_folds} {'='*20}\")\n\n        # Initialize fold-specific writer\n        fold_writer = SummaryWriter(f'runs/{experiment_name}/fold_{fold}')\n\n        # Initialize datasets and dataloaders\n        train_dataset = SpineCoorDataset(train_coor.loc[train_coor.fold!=fold],\n                                    train_meta.loc[train_meta.fold!=fold],\n                                    condition, 'train')\n\n        valid_dataset = SpineCoorDataset(train_coor.loc[train_coor.fold==fold],\n                                    train_meta.loc[train_meta.fold==fold],\n                                    condition, 'valid')\n\n        train_loader = DataLoader(train_dataset, batch_size=batch_size,\n                                 shuffle=True, num_workers=2, pin_memory=True)\n\n        valid_loader = DataLoader(valid_dataset, batch_size=batch_size,\n                                 shuffle=False, num_workers=2, pin_memory=True)\n\n        optimizer = AdamW(model.parameters(), lr=learning_rate, weight_decay=weight_decay)\n\n        # Learning rate scheduler with warmup\n        total_steps = epochs * len(train_loader)\n        warmup_steps = int(0.1 * total_steps)  # 10% warmup\n\n        scheduler = transformers.get_cosine_schedule_with_warmup(\n            optimizer=optimizer,\n            num_warmup_steps=warmup_steps,\n            num_training_steps=total_steps,\n            num_cycles=0.5\n        )\n\n        criterion = SCSLoss(condition=condition)\n\n        # Initialize tracking variables\n        train_losses, val_losses = [], []\n        train_mses, val_mses = [], []\n        best_val_mse = float('inf')\n        counter = 0  # For early stopping\n\n        # Ensure all expected levels are present\n        if train_dataset.condition == 'scs':\n            expected_levels = [\"L1/L2\", \"L2/L3\", \"L3/L4\", \"L4/L5\", \"L5/S1\"]\n        else:\n            expected_levels = ['left_L1/L2', 'left_L2/L3', 'left_L3/L4', 'left_L4/L5', 'left_L5/S1',\n                              'right_L1/L2', 'right_L2/L3', 'right_L3/L4', 'right_L4/L5', 'right_L5/S1']\n\n        # Training loop\n        for epoch in range(epochs):\n            print(f\"\\n🚀 Epoch {epoch+1}/{epochs} - Fold {fold+1}/{n_folds}\")\n\n            # === Training phase ===\n            model.train()\n            running_loss = 0.0\n            total_mse = 0.0\n            epoch_train_tolerances = {f\"±{tol}\": 0 for tol in tolerances}\n            epoch_train_tolerances[\">±2\"] = 0\n            train_total_predictions = 0\n\n            progress_bar = tqdm(enumerate(train_loader), total=len(train_loader),\n                               desc=\"Training Progress\", leave=False)\n\n            for batch_idx, (volume, batch) in progress_bar:\n                volume = volume.to(device)\n                # Ensure all expected levels are present\n                batch = {key: value.to(device) for key, value in batch['coor'].items() if key in expected_levels}\n\n                # Reset gradients\n                optimizer.zero_grad()\n\n                # Mixed precision training\n                if mixed_precision:\n                    with torch.amp.autocast('cuda'):\n                        outputs_raw = model(volume)\n                        outputs = {}\n                        for level in outputs_raw:\n                            if isinstance(outputs_raw[level], dict):\n                                outputs[level] = outputs_raw[level]['coor']  # Take only coordinate tensor\n                            else:\n                                outputs[level] = outputs_raw[level]  # In case it's not dict (safety)\n                        loss = criterion(outputs, batch)\n\n                    # Scale loss and backward pass\n                    scaler.scale(loss).backward()\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                    scaler.step(optimizer)\n                    scaler.update()\n                else:\n                    # Standard precision training\n                    outputs = model(volume)\n                    loss = criterion(outputs, batch)\n                    loss.backward()\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n                    optimizer.step()\n\n                # Step the scheduler\n                scheduler.step()\n\n                # Track metrics\n                running_loss += loss.item()\n                batch_mse = calculate_mse_score(outputs, batch)\n                total_mse += batch_mse\n\n                # Calculate tolerance metrics for this batch\n                batch_tolerance_counts, batch_total = calculate_regression_tolerance(outputs, batch, tolerances)\n                for k in batch_tolerance_counts:\n                    epoch_train_tolerances[k] += batch_tolerance_counts[k]\n                train_total_predictions += batch_total\n\n                # Update progress bar\n                current_lr = optimizer.param_groups[0]['lr']\n                progress_bar.set_postfix(\n                    loss=f\"{loss.item():.4f}\",\n                    mse=f\"{batch_mse:.4f}\",\n                    lr=f\"{current_lr:.6f}\",\n                    t_mse=f\"{total_mse:.4f}\"\n                )\n\n                # Free memory\n                del volume, batch, outputs, loss\n                torch.cuda.empty_cache()\n\n            # Calculate epoch metrics\n            epoch_train_loss = running_loss / len(train_loader)\n            epoch_train_mse = total_mse / len(train_loader)\n            train_losses.append(epoch_train_loss)\n            train_mses.append(epoch_train_mse)\n\n            # Convert tolerance counts to percentages\n            for k in epoch_train_tolerances:\n                if train_total_predictions > 0:\n                    epoch_train_tolerances[k] = 100 * epoch_train_tolerances[k] / train_total_predictions\n                else:\n                    epoch_train_tolerances[k] = 0.0\n\n            # Log metrics\n            fold_writer.add_scalar('Loss/train', epoch_train_loss, epoch)\n            fold_writer.add_scalar('MSE/train', epoch_train_mse, epoch)\n            writer.add_scalar(f'Loss/train/fold_{fold}', epoch_train_loss, epoch)\n            writer.add_scalar(f'MSE/train/fold_{fold}', epoch_train_mse, epoch)\n\n            # Log tolerance metrics\n            for tolerance_key, tolerance_value in epoch_train_tolerances.items():\n                fold_writer.add_scalar(f'Tolerance/train/{tolerance_key}', tolerance_value, epoch)\n                writer.add_scalar(f'Tolerance/train/{tolerance_key}/fold_{fold}', tolerance_value, epoch)\n\n            print(f\"🔥 Training Loss: {epoch_train_loss:.4f} | MSE Score: {epoch_train_mse:.4f}\")\n            print(f\"📏 Training Tolerances: \" + \" | \".join([f\"{k}: {v:.1f}%\" for k, v in epoch_train_tolerances.items()]))\n\n            # === Validation phase ===\n            model.eval()\n            val_running_loss = 0.0\n            val_total_mse = 0.0\n            epoch_val_tolerances = {f\"±{tol}\": 0 for tol in tolerances}\n            epoch_val_tolerances[\">±2\"] = 0\n            val_total_predictions = 0\n\n            with torch.no_grad():\n                for volume, batch in tqdm(valid_loader, desc=\"Validation Progress\", leave=False):\n\n                    volume = volume.to(device)\n                    # Ensure all expected levels are present\n                    batch = {key: value.to(device) for key, value in batch['coor'].items() if key in expected_levels}\n\n                    outputs_raw = model(volume)\n                    outputs = {}\n                    for level in outputs_raw:\n                        outputs[level] = outputs_raw[level]['coor']\n\n                    loss = criterion(outputs, batch)\n                    \n                    val_running_loss += loss.item()\n                    batch_mse = calculate_mse_score(outputs, batch)\n                    val_total_mse += batch_mse\n\n                    # Calculate tolerance metrics for validation batch\n                    batch_tolerance_counts, batch_total = calculate_regression_tolerance(outputs, batch, tolerances)\n                    for k in batch_tolerance_counts:\n                        epoch_val_tolerances[k] += batch_tolerance_counts[k]\n                    val_total_predictions += batch_total\n\n                    del volume, batch, outputs, loss\n                    torch.cuda.empty_cache()\n\n            # Calculate validation metrics\n            epoch_val_loss = val_running_loss / len(valid_loader)\n            epoch_val_mse = val_total_mse / len(valid_loader)\n            val_losses.append(epoch_val_loss)\n            val_mses.append(epoch_val_mse)\n\n            # Convert tolerance counts to percentages\n            for k in epoch_val_tolerances:\n                if val_total_predictions > 0:\n                    epoch_val_tolerances[k] = 100 * epoch_val_tolerances[k] / val_total_predictions\n                else:\n                    epoch_val_tolerances[k] = 0.0\n\n            # Log validation metrics\n            fold_writer.add_scalar('Loss/val', epoch_val_loss, epoch)\n            fold_writer.add_scalar('MSE/val', epoch_val_mse, epoch)\n            writer.add_scalar(f'Loss/val/fold_{fold}', epoch_val_loss, epoch)\n            writer.add_scalar(f'MSE/val/fold_{fold}', epoch_val_mse, epoch)\n\n            # Log validation tolerance metrics\n            for tolerance_key, tolerance_value in epoch_val_tolerances.items():\n                fold_writer.add_scalar(f'Tolerance/val/{tolerance_key}', tolerance_value, epoch)\n                writer.add_scalar(f'Tolerance/val/{tolerance_key}/fold_{fold}', tolerance_value, epoch)\n\n            print(f\"✅ Validation Loss: {epoch_val_loss:.4f} | MSE Score: {epoch_val_mse:.4f}\")\n            print(f\"📏 Validation Tolerances: \" + \" | \".join([f\"{k}: {v:.1f}%\" for k, v in epoch_val_tolerances.items()]))\n\n            # Model checkpointing - save best model based on validation MSE\n            if epoch_val_mse < best_val_mse:\n                best_val_mse = epoch_val_mse\n                torch.save({\n                    'epoch': epoch,\n                    'model_state_dict': model.state_dict(),\n                    'optimizer_state_dict': optimizer.state_dict(),\n                    'val_loss': epoch_val_loss,\n                    'val_mse': epoch_val_mse,\n                    'val_tolerances': epoch_val_tolerances,\n                }, f'models/{experiment_name}_fold_{fold}_best.pt')\n\n                print(f\"📌 New best model saved with MSE: {best_val_mse:.4f}\")\n                counter = 0  # Reset early stopping counter\n            else:\n                counter += 1\n                print(f\"⚠️ No improvement for {counter}/{patience} epochs\")\n\n            # Early stopping\n            if counter >= patience:\n                print(f\"⛔ Early stopping triggered after {epoch+1} epochs\")\n                break\n\n        # End of fold - record best validation MSE\n        all_val_mses.append(best_val_mse)\n        print(f\"Fold {fold+1} completed. Best validation MSE: {best_val_mse:.4f}\")\n\n        # Save final model for this fold\n        torch.save({\n            'epoch': epoch,\n            'model_state_dict': model.state_dict(),\n            'optimizer_state_dict': optimizer.state_dict(),\n            'val_loss': val_losses[-1],\n            'val_mse': val_mses[-1],\n            'val_tolerances': epoch_val_tolerances,\n        }, f'models/{experiment_name}_fold_{fold}_final.pt')\n\n        # Close fold writer\n        fold_writer.close()\n\n    # End of cross-validation\n    avg_val_mse = sum(all_val_mses) / len(all_val_mses)\n    print(f\"\\n===== TRAINING COMPLETE =====\")\n    print(f\"Cross-validation results:\")\n    for fold, mse in enumerate(all_val_mses):\n        print(f\"Fold {fold+1}: MSE = {mse:.4f}\")\n    print(f\"Average validation MSE: {avg_val_mse:.4f}\")\n\n    # Save experiment summary\n    with open(f'models/{experiment_name}_summary.txt', 'w') as f:\n        f.write(f\"Experiment: {experiment_name}\\n\")\n        f.write(f\"Folds: {n_folds}\\n\")\n        f.write(f\"Epochs: {epochs}\\n\")\n        f.write(f\"Batch size: {batch_size}\\n\")\n        f.write(f\"Learning rate: {learning_rate}\\n\")\n        f.write(f\"Weight decay: {weight_decay}\\n\")\n        f.write(f\"Mixed precision: {mixed_precision}\\n\")\n        f.write(f\"Tolerance thresholds: {tolerances}\\n\")\n        f.write(\"\\nResults:\\n\")\n        for fold, mse in enumerate(all_val_mses):\n            f.write(f\"Fold {fold+1}: MSE = {mse:.4f}\\n\")\n        f.write(f\"Average validation MSE: {avg_val_mse:.4f}\\n\")\n\n    # Close main writer\n    writer.close()\n\n    plt.figure(figsize=(10, 6))\n    plt.plot(all_val_mses, label='Train MSE', marker='o')\n    plt.plot(avg_val_mse, label='Validation MSE', marker='x')\n    plt.title('Train vs Validation MSE')\n    plt.xlabel('Epoch')\n    plt.ylabel('MSE')\n    plt.legend()\n    plt.grid(True)\n    plt.tight_layout()\n    plt.savefig(f'models/{experiment_name}_mse_epoch_plot.png')\n    plt.show()\n\n\n    return all_val_mses, avg_val_mse\n\n\nmodel = MobileNetV3GatedSpine().to(device)\nmses, avg_mse = train_spine_model(model, train_coor, train_meta, n_folds=5, condition='scs')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T19:10:11.161893Z","iopub.execute_input":"2025-04-28T19:10:11.162568Z","iopub.status.idle":"2025-04-28T19:10:11.210613Z","shell.execute_reply.started":"2025-04-28T19:10:11.162543Z","shell.execute_reply":"2025-04-28T19:10:11.209724Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Maxvit ","metadata":{}},{"cell_type":"markdown","source":"> it work only on 224 size image","metadata":{}},{"cell_type":"code","source":"class ConvNextSCSDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        #self.size = 384\n        self.encoder = timm.create_model(\n            'maxvit_tiny_rw_224',   # <-- model name from timm\n            pretrained=True,\n            in_chans=3,\n            features_only=False,\n            num_classes=0\n        )\n        self.in_features = self.encoder.num_features\n        print(self.in_features)\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1),\n                                    #nn.LayerNorm(self.in_features)\n                                    )\n\n        self.l1 = nn.Linear(self.in_features, 2)\n        self.l2 = nn.Linear(self.in_features, 2)\n        self.l3 = nn.Linear(self.in_features, 2)\n        self.l4 = nn.Linear(self.in_features, 2)\n        self.l5 = nn.Linear(self.in_features, 2)\n\n    def forward(self, x, label=None):\n        #for loc, img in x.items():\n            #print(img.shape)\n        #    img = self.encoder.forward_features(img)\n        #    img = self.flatten(img)\n        #    x[loc] = img\n        x = self.encoder.forward_features(x)\n        x = self.flatten(x)\n        l1 = self.l1(x)\n        l2 = self.l2(x)\n        l3 = self.l3(x)\n        l4 = self.l4(x)\n        l5 = self.l5(x)\n        return {'L1/L2': l1.sigmoid(), 'L2/L3': l2.sigmoid(), 'L3/L4': l3.sigmoid(), 'L4/L5': l4.sigmoid(), 'L5/S1': l5.sigmoid()}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T16:50:43.997225Z","iopub.execute_input":"2025-04-28T16:50:43.997561Z","iopub.status.idle":"2025-04-28T16:50:44.005214Z","shell.execute_reply.started":"2025-04-28T16:50:43.997535Z","shell.execute_reply":"2025-04-28T16:50:44.004447Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = ConvNextSCSDetect().to(device)\nmses, avg_mse = train_spine_model(model, train_coor, train_meta, n_folds=5, condition='scs')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T16:53:56.390551Z","iopub.execute_input":"2025-04-28T16:53:56.390923Z","iopub.status.idle":"2025-04-28T16:53:56.411639Z","shell.execute_reply.started":"2025-04-28T16:53:56.390902Z","shell.execute_reply":"2025-04-28T16:53:56.410654Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class ConvNextSCSDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        #self.size = 384\n        self.encoder = timm.create_model(\n            'efficientnetv2_s',   # <-- model name from timm\n            pretrained=False,\n            in_chans=3,\n            features_only=False,\n            num_classes=0\n        )\n        self.in_features = self.encoder.num_features\n        print(self.in_features)\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1),\n                                    #nn.LayerNorm(self.in_features)\n                                    )\n\n        self.l1 = nn.Linear(self.in_features, 2)\n        self.l2 = nn.Linear(self.in_features, 2)\n        self.l3 = nn.Linear(self.in_features, 2)\n        self.l4 = nn.Linear(self.in_features, 2)\n        self.l5 = nn.Linear(self.in_features, 2)\n\n    def forward(self, x, label=None):\n        #for loc, img in x.items():\n            #print(img.shape)\n        #    img = self.encoder.forward_features(img)\n        #    img = self.flatten(img)\n        #    x[loc] = img\n        x = self.encoder.forward_features(x)\n        x = self.flatten(x)\n        l1 = self.l1(x)\n        l2 = self.l2(x)\n        l3 = self.l3(x)\n        l4 = self.l4(x)\n        l5 = self.l5(x)\n        return {'L1/L2': l1.sigmoid(), 'L2/L3': l2.sigmoid(), 'L3/L4': l3.sigmoid(), 'L4/L5': l4.sigmoid(), 'L5/S1': l5.sigmoid()}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T16:54:34.369566Z","iopub.execute_input":"2025-04-28T16:54:34.369870Z","iopub.status.idle":"2025-04-28T16:54:34.376501Z","shell.execute_reply.started":"2025-04-28T16:54:34.369849Z","shell.execute_reply":"2025-04-28T16:54:34.375878Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = ConvNextSCSDetect().to(device)\nmses, avg_mse = train_spine_model(model, train_coor, train_meta, n_folds=5, condition='scs')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:03:54.768907Z","iopub.execute_input":"2025-04-28T17:03:54.769375Z","iopub.status.idle":"2025-04-28T17:03:54.788693Z","shell.execute_reply.started":"2025-04-28T17:03:54.769354Z","shell.execute_reply":"2025-04-28T17:03:54.787981Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import gc\n\n# ✅ After each epoch\ntorch.cuda.empty_cache()\ntorch.cuda.ipc_collect()\ngc.collect()  # Optional, forces Python garbage collection\n\ntorch.cuda.empty_cache()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:01:41.516063Z","iopub.execute_input":"2025-04-28T17:01:41.516773Z","iopub.status.idle":"2025-04-28T17:01:42.026969Z","shell.execute_reply.started":"2025-04-28T17:01:41.516724Z","shell.execute_reply":"2025-04-28T17:01:42.026414Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\n\n# ----------------- Gated Attention Block (with Residual) -----------------\nclass GatedAttention(nn.Module):\n    def __init__(self, in_features):\n        super().__init__()\n        self.gate = nn.Sequential(\n            nn.Linear(in_features, in_features),\n            nn.Sigmoid()\n        )\n    def forward(self, x):\n        gate = self.gate(x)\n        return x * gate + x  # Residual connection\n\n# ----------------- EfficientNetV2-S Spine Detection Model -----------------\nclass EfficientNetV2GatedSpine(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        # Load EfficientNetV2-S backbone\n        self.encoder = timm.create_model(\n            'efficientnetv2_s',\n            pretrained=False,\n            in_chans=3,\n            features_only=False,\n            num_classes=0\n        )\n\n        # Feature dimension\n        self.in_features = self.encoder.num_features  # 1280 for EfficientNetV2-S\n\n        # Adaptive Pooling + Flatten\n        self.flatten = nn.Sequential(\n            nn.AdaptiveAvgPool2d((1, 1)),\n            nn.Flatten(1)\n        )\n\n        # Gated Attention after encoder output\n        self.attention = GatedAttention(self.in_features)\n\n        # Project feature\n        self.projector = nn.Identity()\n\n        # Dropout for regularization\n        self.dropout = nn.Dropout(p=0.2)\n\n        # Heads for 5 vertebral levels\n        self.heads = nn.ModuleList([\n            nn.Linear(self.in_features, 2) for _ in range(5)\n        ])\n\n    def forward(self, x):\n        # Feature extraction\n        x = self.encoder.forward_features(x)\n        x = self.flatten(x)\n\n        # Gated Attention\n        x = self.attention(x)\n\n        # Project to feature dimension\n        x = self.projector(x)\n\n        # Small dropout\n        x = self.dropout(x)\n\n        # Predict each level\n        output = {}\n        levels = ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']\n\n        for i, level in enumerate(levels):\n            output[level] = self.heads[i](x).sigmoid()\n\n        return output\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T18:07:42.379263Z","iopub.execute_input":"2025-04-28T18:07:42.379806Z","iopub.status.idle":"2025-04-28T18:07:42.387501Z","shell.execute_reply.started":"2025-04-28T18:07:42.379784Z","shell.execute_reply":"2025-04-28T18:07:42.386804Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# EfficientNetV2-S\nmodel = EfficientNetV2GatedSpine().to(device)\nmses, avg_mse = train_spine_model(model, train_coor, train_meta, n_folds=5, condition='scs')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T17:05:25.393015Z","iopub.execute_input":"2025-04-28T17:05:25.393244Z","iopub.status.idle":"2025-04-28T17:36:38.445817Z","shell.execute_reply.started":"2025-04-28T17:05:25.393226Z","shell.execute_reply":"2025-04-28T17:36:38.444889Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\n\n# ----------------- Gated Attention Block (with Residual) -----------------\nclass GatedAttention(nn.Module):\n    def __init__(self, in_features):\n        super().__init__()\n        self.gate = nn.Sequential(\n            nn.Linear(in_features, in_features),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        gate = self.gate(x)\n        return x * gate + x  # Residual connection\n\n# ----------------- EfficientNetV2-S NFN Detection Model -----------------\nclass EfficientNetV2SNFN(nn.Module):\n    def __init__(self):\n        super().__init__()\n\n        # Load EfficientNetV2-S backbone\n        self.encoder = timm.create_model(\n            'efficientnetv2_s',\n            in_chans=3,\n            pretrained=False,\n            features_only=False,\n            num_classes=0\n        )\n\n        # Feature dimension\n        self.in_features = self.encoder.num_features\n\n        # Adaptive Pooling + Flatten\n        self.flatten = nn.Sequential(\n            nn.AdaptiveAvgPool2d((1,1)),\n            nn.Flatten(1)\n        )\n\n        # Gated Attention after encoder output\n        self.attention = GatedAttention(self.in_features)\n\n        # Project feature (increased from 1024 to 1280 for improved representational capacity)\n        self.projector = nn.Sequential(\n            nn.Linear(self.in_features, 1280),\n            nn.SiLU(inplace=True)\n        )\n\n        # Dropout for regularization\n        self.dropout = nn.Dropout(p=0.2)\n\n        # Heads for 5 vertebral levels (left and right)\n        self.ll1 = nn.Linear(1280, 2)\n        self.ll2 = nn.Linear(1280, 2)\n        self.ll3 = nn.Linear(1280, 2)\n        self.ll4 = nn.Linear(1280, 2)\n        self.ll5 = nn.Linear(1280, 2)\n        \n        self.rl1 = nn.Linear(1280, 2)\n        self.rl2 = nn.Linear(1280, 2)\n        self.rl3 = nn.Linear(1280, 2)\n        self.rl4 = nn.Linear(1280, 2)\n        self.rl5 = nn.Linear(1280, 2)\n\n    def forward(self, x, label=None):\n        # Reshape input to handle both left and right side images\n        shape = x.shape\n        x = x.reshape(shape[0] * shape[1], 3, shape[-2], shape[-1])  # Flatten the batch and side\n\n        # Feature extraction through EfficientNetV2-S\n        x = self.encoder.forward_features(x)\n        x = self.flatten(x)\n        x = x.reshape(shape[0], shape[1], -1)  # Split back into left and right\n\n        # Split features for left and right\n        x_left = x[:, 0, :]\n        x_right = x[:, 1, :]\n\n        # Apply Gated Attention\n        x_left = self.attention(x_left)\n        x_right = self.attention(x_right)\n\n        # Project features into desired size\n        x_left = self.projector(x_left)\n        x_right = self.projector(x_right)\n\n        # Apply Dropout\n        x_left = self.dropout(x_left)\n        x_right = self.dropout(x_right)\n\n        # Predict each level for left and right sides\n        output = {\n            'left_L1/L2': self.ll1(x_left).sigmoid(),\n            'left_L2/L3': self.ll2(x_left).sigmoid(),\n            'left_L3/L4': self.ll3(x_left).sigmoid(),\n            'left_L4/L5': self.ll4(x_left).sigmoid(),\n            'left_L5/S1': self.ll5(x_left).sigmoid(),\n            'right_L1/L2': self.rl1(x_right).sigmoid(),\n            'right_L2/L3': self.rl2(x_right).sigmoid(),\n            'right_L3/L4': self.rl3(x_right).sigmoid(),\n            'right_L4/L5': self.rl4(x_right).sigmoid(),\n            'right_L5/S1': self.rl5(x_right).sigmoid()\n        }\n\n        return output\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Extrea","metadata":{}},{"cell_type":"code","source":"# -------------------------------\n# Dataset\n# -------------------------------\nclass SpineCoorDataset(Dataset):\n    def __init__(self, coor, meta, condition, mode):\n        if condition == 'scs':\n            self.coor = coor[coor.condition == \"Spinal Canal Stenosis\"]\n        elif condition == 'nfn':\n            self.coor = coor[coor.condition.isin([\n                'Left Neural Foraminal Narrowing',\n                'Right Neural Foraminal Narrowing'\n            ])]\n        elif condition == 'ss':\n            self.coor = coor[coor.condition.isin([\n                'Left Subarticular Stenosis',\n                'Right Subarticular Stenosis'\n            ])]\n\n        g_coor = self.coor.groupby('study_id').count()\n        if condition == 'scs':\n            self.id = g_coor[g_coor.series_id == 5].reset_index().study_id.unique()\n        else:\n            self.id = g_coor[g_coor.series_id == 10].reset_index().study_id.unique()\n\n        if condition == 'ss':\n            self.resize = v2.Resize((256, 256))\n        else:\n            self.resize = v2.Resize((384, 384))\n\n        # remove problematic ID\n        self.id = list(set(self.id) - {3637444890})\n\n        self.condition = condition\n        self.meta = meta\n        self.mode = mode\n\n    def __len__(self):\n        return len(self.id)\n\n    def __getitem__(self, idx):\n        study_id = self.id[idx]\n        volume, label = self.volume_scs(study_id)\n        return volume, label\n\n    def volume_scs(self, study_id):\n        all_levels = ['L1/L2','L2/L3','L3/L4','L4/L5','L5/S1']\n        meta = self.meta[(self.meta.study_id==study_id)&\n                         (self.meta.series_description=='Sagittal T2/STIR')]\n        meta = meta.sort_values('ipp_x').reset_index(drop=True)\n        coor = self.coor[(self.coor.study_id==study_id)&\n                         (self.coor.series_description=='Sagittal T2/STIR')]\n        coor_dict, conf_dict = {}, {}\n\n        # find central slice\n        x_positions = torch.tensor([row.ipp_x for _,row in coor.iterrows()])\n        median_x = x_positions.median().item()\n        idx_row = meta[meta.ipp_x==median_x].index[0]\n        slices = [idx_row-1, idx_row, idx_row+1]\n        imgs = []\n        for i in slices:\n            r = meta.loc[i]\n            arr = self.load_dicom(\n                f\"/kaggle/input/.../train_images/{study_id}/{r.series_id}/{r.instance_number}.dcm\"\n            )\n            imgs.append(self.normalize(arr))\n        h,w = imgs[1].shape\n\n        for lvl in all_levels:\n            lvl_coor = coor[coor.level==lvl]\n            if not lvl_coor.empty:\n                x = lvl_coor.iloc[0].x / w\n                y = lvl_coor.iloc[0].y / h\n                coor_dict[lvl] = torch.tensor([x,y])\n                conf_dict[lvl] = torch.tensor(1.0)\n            else:\n                coor_dict[lvl] = torch.tensor([0.0,0.0])\n                conf_dict[lvl] = torch.tensor(0.0)\n\n        imgs = [self.resize(torch.tensor(im[None,...])) for im in imgs]\n        volume = torch.cat(imgs,0).float()\n        return volume, {'coor':coor_dict,'conf':conf_dict}\n\n    def normalize(self, x):\n        if self.condition=='ss':\n            low,high = torch.quantile(torch.tensor(x),0.01).item(), torch.quantile(torch.tensor(x),0.99).item()\n            x = np.clip(x,low,high)\n        else:\n            low,high = np.percentile(x,(1,99))\n            x = np.clip(x,low,high)\n        x = (x - x.min())/(x.max()-x.min())\n        return x\n\n    def load_dicom(self,path): return pydicom.dcmread(path).pixel_array","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T19:36:22.918730Z","iopub.execute_input":"2025-04-28T19:36:22.919009Z","iopub.status.idle":"2025-04-28T19:36:22.932408Z","shell.execute_reply.started":"2025-04-28T19:36:22.918989Z","shell.execute_reply":"2025-04-28T19:36:22.931677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MobileNetV3SmallSCSDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = timm.create_model('mobilenetv3_small_100',pretrained=True,in_chans=3,num_classes=0)\n        self.in_features = self.encoder.num_features\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),nn.Flatten(1))\n        self.gated_attention = GatedAttention(self.in_features)\n        self.heads = nn.ModuleDict({lvl:nn.Linear(self.in_features,3)for lvl in ['L1/L2','L2/L3','L3/L4','L4/L5','L5/S1']})\n\n    def forward(self,x):\n        f = self.encoder.forward_features(x)\n        f = self.flatten(f)\n        f = self.gated_attention(f)\n        out={}\n        for lvl,head in self.heads.items():\n            p = head(f)\n            out[lvl] = {\n                'coor': torch.sigmoid(p[:,:2]),\n                'conf': p[:,2].unsqueeze(-1)  # raw logits\n            }\n        return out","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T19:36:24.716739Z","iopub.execute_input":"2025-04-28T19:36:24.717006Z","iopub.status.idle":"2025-04-28T19:36:24.723291Z","shell.execute_reply.started":"2025-04-28T19:36:24.716988Z","shell.execute_reply":"2025-04-28T19:36:24.722561Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------\n# Loss\n# -------------------------------\nclass SCSLoss(nn.Module):\n    def __init__(self,lambda_conf=1.0):\n        super().__init__()\n        self.lambda_conf=lambda_conf\n\n    def forward(self,outputs,targets):\n        coord_loss=0.0; conf_loss=0.0; num_present=0.0\n        B = next(iter(targets.values()))['conf'].shape[0]\n        L = len(outputs)\n        for lvl, pred in outputs.items():\n            p_xy = pred['coor']; p_logit = pred['conf'].squeeze(-1)\n            t_xy = targets[lvl]['coor']; t_conf = targets[lvl]['conf']\n            mask = t_conf.unsqueeze(-1).expand_as(t_xy)\n            coord_loss += F.smooth_l1_loss(p_xy*mask, t_xy*mask, reduction='sum')\n            conf_loss  += F.binary_cross_entropy_with_logits(p_logit, t_conf, reduction='sum')\n            num_present += t_conf.sum()\n        coord_loss /= (num_present+1e-6)\n        conf_loss  /= (B*L)\n        return coord_loss + self.lambda_conf*conf_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T19:36:44.364726Z","iopub.execute_input":"2025-04-28T19:36:44.365441Z","iopub.status.idle":"2025-04-28T19:36:44.371312Z","shell.execute_reply.started":"2025-04-28T19:36:44.365416Z","shell.execute_reply":"2025-04-28T19:36:44.370621Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------\n# Metrics\n# -------------------------------\ndef calculate_mse_score(outputs, targets):\n    total_mse = 0.0\n    count = 0\n    for lvl in outputs:\n        p_xy = outputs[lvl]['coor']\n        t_xy = targets[lvl]['coor']\n        mse = F.mse_loss(p_xy, t_xy, reduction='sum').item()\n        total_mse += mse\n        count += p_xy.numel()\n    return total_mse / (count + 1e-6)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T19:36:48.434021Z","iopub.execute_input":"2025-04-28T19:36:48.434715Z","iopub.status.idle":"2025-04-28T19:36:48.438936Z","shell.execute_reply.started":"2025-04-28T19:36:48.434685Z","shell.execute_reply":"2025-04-28T19:36:48.438265Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def calculate_regression_tolerance(outputs, targets, tolerances=[0,1,2]):\n    counts = {f\"±{tol}\": 0 for tol in tolerances}\n    counts[\">±2\"] = 0\n    total = 0\n    for lvl in outputs:\n        p = outputs[lvl]['coor'].detach().cpu().numpy()\n        t = targets[lvl]['coor'].detach().cpu().numpy()\n        diffs = np.abs(p - t).astype(int)\n        for drow in diffs:\n            for d in drow:\n                matched = False\n                for tol in tolerances:\n                    if d <= tol:\n                        counts[f\"±{tol}\"] += 1\n                        matched = True\n                        break\n                if not matched:\n                    counts[\">±2\"] += 1\n                total += 1\n    return counts, total","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T19:36:48.743921Z","iopub.execute_input":"2025-04-28T19:36:48.744144Z","iopub.status.idle":"2025-04-28T19:36:48.749855Z","shell.execute_reply.started":"2025-04-28T19:36:48.744128Z","shell.execute_reply":"2025-04-28T19:36:48.749138Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# -------------------------------\n# Training\n# -------------------------------\ndef train_spine_model(model, train_coor, train_meta,\n                      n_folds=5, epochs=18, batch_size=4,\n                      learning_rate=1e-3, weight_decay=1e-4,\n                      patience=5, mixed_precision=True,\n                      experiment_name=\"spine_model\"):\n\n    os.makedirs(\"models\", exist_ok=True)\n    writer = SummaryWriter(f\"runs/{experiment_name}\")\n    scaler = torch.amp.GradScaler() if mixed_precision else None\n\n    all_val_mses = []\n    for fold in range(n_folds):\n        print(f\"-- Fold {fold+1}/{n_folds} --\")\n        # Datasets & Loaders\n        train_ds = SpineCoorDataset(\n            train_coor[train_coor.fold != fold],\n            train_meta[train_meta.fold != fold],\n            'scs','train'\n        )\n        val_ds = SpineCoorDataset(\n            train_coor[train_coor.fold == fold],\n            train_meta[train_meta.fold == fold],\n            'scs','valid'\n        )\n        train_loader = DataLoader(train_ds, batch_size=batch_size,\n                                  shuffle=True, num_workers=2, pin_memory=True)\n        val_loader   = DataLoader(val_ds,   batch_size=batch_size,\n                                  shuffle=False, num_workers=2, pin_memory=True)\n\n        optimizer = AdamW(model.parameters(), lr=learning_rate,\n                          weight_decay=weight_decay)\n        total_steps = epochs * len(train_loader)\n        warmup = int(0.1 * total_steps)\n        scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(\n            optimizer, T_max=total_steps)\n\n        criterion = SCSLoss(lambda_conf=1.0)\n        expected_levels = ['L1/L2','L2/L3','L3/L4','L4/L5','L5/S1']\n\n        best_mse = float('inf'); counter=0\n        for epoch in range(epochs):\n            print(f\" Epoch {epoch+1}/{epochs}\")\n            # Training\n            model.train()\n            train_loss = train_mse = 0.0\n            tol_counts = {f\"±{t}\":0 for t in [0,1,2]}; tol_counts[\">±2\"]=0\n            total_preds=0\n            for vol, batch in tqdm(train_loader):\n                vol = vol.to(device)\n                # prepare targets per level\n                targets = {lvl: {'coor': batch['coor'][lvl].to(device),\n                                 'conf': batch['conf'][lvl].to(device)}\n                           for lvl in expected_levels}\n\n                optimizer.zero_grad()\n                if mixed_precision:\n                    with torch.amp.autocast(device_type='cuda'):\n                        out = model(vol)\n                        loss = criterion(out, targets)\n                    scaler.scale(loss).backward()\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(),1.0)\n                    scaler.step(optimizer)\n                    scaler.update()\n                else:\n                    out = model(vol)\n                    loss = criterion(out, targets)\n                    loss.backward()\n                    torch.nn.utils.clip_grad_norm_(model.parameters(),1.0)\n                    optimizer.step()\n\n                scheduler.step()\n                train_loss += loss.item()\n                mse = calculate_mse_score(out, targets)\n                train_mse += mse\n                bc, bt = calculate_regression_tolerance(out, targets)\n                for k in tol_counts: tol_counts[k] += bc[k]\n                total_preds += bt\n                del vol, batch, out, loss\n\n            # compute averages\n            train_loss /= len(train_loader)\n            train_mse /= len(train_loader)\n            for k in tol_counts: tol_counts[k] = 100 * tol_counts[k]/(total_preds+1e-6)\n            print(f\" Train Loss: {train_loss:.4f}, MSE: {train_mse:.4f}\")\n            print(\" Tolerances: \", tol_counts)\n\n            # Validation\n            model.eval()\n            val_loss= val_mse=0.0\n            val_tcounts = {f\"±{t}\":0 for t in [0,1,2]}; val_tcounts[\">±2\"]=0\n            vpreds=0\n            with torch.no_grad():\n                for vol, batch in tqdm(val_loader):\n                    vol=vol.to(device)\n                    targets = {lvl:{'coor': batch['coor'][lvl].to(device),\n                                    'conf': batch['conf'][lvl].to(device)}\n                               for lvl in expected_levels}\n                    out = model(vol)\n                    loss = criterion(out, targets)\n                    val_loss += loss.item()\n                    mse = calculate_mse_score(out, targets)\n                    val_mse += mse\n                    bc,bt = calculate_regression_tolerance(out, targets)\n                    for k in val_tcounts: val_tcounts[k]+=bc[k]\n                    vpreds += bt\n            val_loss /= len(val_loader)\n            val_mse  /= len(val_loader)\n            for k in val_tcounts: val_tcounts[k]=100*val_tcounts[k]/(vpreds+1e-6)\n            print(f\" Val Loss: {val_loss:.4f}, MSE: {val_mse:.4f}\")\n            print(\" Val Tolerances: \", val_tcounts)\n\n            # Early stopping & checkpoints\n            if val_mse < best_mse:\n                best_mse=val_mse; counter=0\n                torch.save(model.state_dict(), f\"models/best_fold{fold}.pth\")\n                print(\" Saved best model.\")\n            else:\n                counter+=1\n                if counter>=patience:\n                    print(\"Early stopping.\")\n                    break\n\n        all_val_mses.append(best_mse)\n\n    avg_mse = sum(all_val_mses)/len(all_val_mses)\n    print(\"CV MSEs:\", all_val_mses)\n    print(\"Avg CV MSE:\", avg_mse)\n    return all_val_mses, avg_mse\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T19:36:50.124646Z","iopub.execute_input":"2025-04-28T19:36:50.124934Z","iopub.status.idle":"2025-04-28T19:36:50.141584Z","shell.execute_reply.started":"2025-04-28T19:36:50.124903Z","shell.execute_reply":"2025-04-28T19:36:50.141039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = MobileNetV3SmallSCSDetect().to(device)\nmses, avg_mse = train_spine_model(model, train_coor, train_meta, n_folds=5, )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T19:36:52.791607Z","iopub.execute_input":"2025-04-28T19:36:52.791918Z","iopub.status.idle":"2025-04-28T19:36:53.230372Z","shell.execute_reply.started":"2025-04-28T19:36:52.791900Z","shell.execute_reply":"2025-04-28T19:36:53.229043Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load model weights\ndef load_model_weights(model, weight_path, device):\n    pretrained_dict = torch.load(weight_path, map_location=device)\n    model_dict = model.state_dict()\n    pretrained_dict = {k: v for k, v in pretrained_dict.items() if k in model_dict}\n    model_dict.update(pretrained_dict)\n    model.load_state_dict(model_dict)\n\n# Image preprocessing\ndef preprocess_image(img_path):\n    dicom_data = pydicom.dcmread(img_path)\n    img = dicom_data.pixel_array\n\n    # Normalize image (scaling intensity to [0, 1])\n    lower, upper = np.percentile(img, (1, 99))\n    img = np.clip(img, lower, upper)\n    img = img - np.min(img)\n    img = img / np.max(img)\n\n    img = torch.tensor(img).unsqueeze(0).float()  # (1, H, W)\n    img = v2.Resize((384, 384))(img)\n    img = img.repeat(3, 1, 1)                     # (3, H, W)\n    img = img.unsqueeze(0).to(device)              # (B=1, C=3, H, W)\n\n    return img\n\n# Predict coordinates and confidences\n@torch.no_grad()\ndef predict_coordinates(img_path, model, device):\n    img = preprocess_image(img_path)\n    outputs = model(img)\n\n    coordinates = {}\n    confidences = {}\n\n    for level, out in outputs.items():\n        # out is a Tensor directly, (x, y, confidence)\n        coor = out['coor'].squeeze(0).cpu().numpy()  # (x, y)\n        conf = out['conf'].squeeze(0).cpu().numpy()  # confidence\n        \n        coordinates[level] = (float(coor[0]), float(coor[1]))\n        confidences[level] = float(conf[0])\n\n    return coordinates, confidences\n\n# Example usage\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = MobileNetV3SmallSCSDetect()\nload_model_weights(model, '/kaggle/working/models/spine_model_fold_0_final.pt', device)\nmodel = model.to(device)\nmodel.eval()\n\n# Example image path (change to actual path in your dataset)\nimg_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/4003253/702807833/1.dcm'\n\n# Get coordinates and confidences\ncoordinates, confidences = predict_coordinates(img_path, model, device)\n\nthreshold = 0.5\nfor level in coordinates:\n    x, y = coordinates[level]\n    conf = confidences[level]\n    if conf > threshold:\n        print(f\"{level}: x = {x:.2f}, y = {y:.2f}, confidence = {conf:.2f}\")\n    else:\n        print(f\"{level}: Prediction is below threshold\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T18:47:10.751738Z","iopub.execute_input":"2025-04-28T18:47:10.752001Z","iopub.status.idle":"2025-04-28T18:47:11.340039Z","shell.execute_reply.started":"2025-04-28T18:47:10.751982Z","shell.execute_reply":"2025-04-28T18:47:11.339432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ----------------- Inference -----------------\ndef load_model_weights(model, weight_path, device='cuda'):\n    \"\"\"Load model weights from checkpoint\"\"\"\n    checkpoint = torch.load(weight_path, map_location=device)\n    \n    # Check if checkpoint has 'model_state_dict' key\n    if 'model_state_dict' in checkpoint:\n        state_dict = checkpoint['model_state_dict']\n    else:\n        state_dict = checkpoint  # Assume it's just the state dict\n    \n    # Load weights\n    model.load_state_dict(state_dict, strict=False)\n    return model\n\n\ndef preprocess_image(img_path, device='cuda'):\n    \"\"\"Preprocess image for inference\"\"\"\n    dicom_data = pydicom.dcmread(img_path)\n    img = dicom_data.pixel_array\n    \n    # Handle different pixel value ranges\n    if img.max() > 0:  # Avoid division by zero\n        lower, upper = np.percentile(img, (1, 99))\n        img = np.clip(img, lower, upper)\n        img = (img - lower) / max(upper - lower, 1e-8)  # Normalize to [0,1]\n    else:\n        # Handle empty or invalid images\n        img = np.zeros_like(img)\n    \n    # Convert to tensor\n    img = torch.tensor(img).unsqueeze(0).float()  # (1, H, W)\n    img = v2.Resize((384, 384))(img)\n    img = img.repeat(3, 1, 1)                     # (3, H, W)\n    img = img.unsqueeze(0).to(device)             # (B=1, C=3, H, W)\n    \n    return img\n\n\n@torch.no_grad()\ndef predict_coordinates(model, img_path, confidence_threshold=0.5, device='cuda'):\n    \"\"\"\n    Predict coordinates for spinal levels with confidence scores\n    Only returns levels where confidence is above threshold\n    \"\"\"\n    # Preprocess image\n    img = preprocess_image(img_path, device)\n    \n    # Set model to evaluation mode\n    model.eval()\n    \n    # Get predictions\n    outputs = model(img)\n    \n    # Process predictions\n    coordinates = {}\n    confidence_scores = {}\n    \n    for level, out in outputs.items():\n        x, y, conf = out.squeeze(0).cpu().numpy()\n        confidence_scores[level] = float(conf)\n        \n        # Only include coordinates with confidence above threshold\n        if conf >= confidence_threshold:","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}