{"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":31011,"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-05-06T21:07:58.418515Z","iopub.execute_input":"2025-05-06T21:07:58.418710Z","iopub.status.idle":"2025-05-06T21:08:01.402032Z","shell.execute_reply.started":"2025-05-06T21:07:58.418694Z","shell.execute_reply":"2025-05-06T21:08:01.401199Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport timm\nimport torch\nimport random\nimport pydicom\nimport numpy as np\nimport pandas as pd\nimport transformers\nfrom tqdm import tqdm\nimport seaborn as sns\nimport torch.nn as nn\nfrom typing import List\nfrom torch import Tensor\nimport albumentations as A\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\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import classification_report, accuracy_score, f1_score, confusion_matrix, ConfusionMatrixDisplay","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-06T21:08:01.403039Z","iopub.execute_input":"2025-05-06T21:08:01.403260Z","iopub.status.idle":"2025-05-06T21:08:10.011364Z","shell.execute_reply.started":"2025-05-06T21:08:01.403240Z","shell.execute_reply":"2025-05-06T21:08:10.010622Z"}},"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-05-06T21:08:10.012194Z","iopub.execute_input":"2025-05-06T21:08:10.013233Z","iopub.status.idle":"2025-05-06T21:08:10.079309Z","shell.execute_reply.started":"2025-05-06T21:08:10.013203Z","shell.execute_reply":"2025-05-06T21:08:10.078504Z"}},"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-05-06T21:08:10.080111Z","iopub.execute_input":"2025-05-06T21:08:10.080316Z","iopub.status.idle":"2025-05-06T21:08:10.454346Z","shell.execute_reply.started":"2025-05-06T21:08:10.080299Z","shell.execute_reply":"2025-05-06T21:08:10.453750Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dummy = train_dummy.fillna(\"Normal/Mild\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-06T21:08:10.456197Z","iopub.execute_input":"2025-05-06T21:08:10.456413Z","iopub.status.idle":"2025-05-06T21:08:10.465101Z","shell.execute_reply.started":"2025-05-06T21:08:10.456396Z","shell.execute_reply":"2025-05-06T21:08:10.464428Z"}},"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'])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-06T21:08:10.466233Z","iopub.execute_input":"2025-05-06T21:08:10.466450Z","iopub.status.idle":"2025-05-06T21:08:10.484062Z","shell.execute_reply.started":"2025-05-06T21:08:10.466425Z","shell.execute_reply":"2025-05-06T21:08:10.483548Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class SeverityPrediction(Dataset):\n    def __init__(self, df, series, coor, meta, condition, usage='train'):\n        self.series = series\n        self.coor = coor\n        self.meta = meta\n        self.df = df\n        self.condition = condition\n        self.usage = usage\n        self.sag_window = 5\n        self.ax_window = 5\n        \n        # Define labels based on condition\n        if condition == 'scs':\n            self.label = [\n                'spinal_canal_stenosis_l1_l2', 'spinal_canal_stenosis_l2_l3', 'spinal_canal_stenosis_l3_l4', 'spinal_canal_stenosis_l4_l5', 'spinal_canal_stenosis_l5_s1'\n            ]\n        elif condition == 'ss':\n            self.label = [\n                'left_subarticular_stenosis_l1_l2', 'left_subarticular_stenosis_l2_l3', 'left_subarticular_stenosis_l3_l4', 'left_subarticular_stenosis_l4_l5', 'left_subarticular_stenosis_l5_s1',\n                'right_subarticular_stenosis_l1_l2', 'right_subarticular_stenosis_l2_l3', 'right_subarticular_stenosis_l3_l4', 'right_subarticular_stenosis_l4_l5', 'right_subarticular_stenosis_l5_s1'\n            ]\n        elif condition == 'nfn':\n            self.label = [\n                'left_neural_foraminal_narrowing_l1_l2', f'left_neural_foraminal_narrowing_l2_l3', f'left_neural_foraminal_narrowing_l3_l4', f'left_neural_foraminal_narrowing_l4_l5', f'left_neural_foraminal_narrowing_l5_s1',\n                'right_neural_foraminal_narrowing_l1_l2', f'right_neural_foraminal_narrowing_l2_l3', f'right_neural_foraminal_narrowing_l3_l4', f'right_neural_foraminal_narrowing_l4_l5', f'right_neural_foraminal_narrowing_l5_s1'\n            ]\n        \n        # Clean and prepare labels\n        # Get only study IDs with complete label data (no NaN values)\n        self.id = df.loc[~(df[self.label].isna().any(axis=1)), 'study_id'].unique()\n        \n        # Remove specific problematic study ID\n        self.id = list(set(self.id) - set([3637444890]))\n        \n        # Map string labels to numeric values\n        for l in self.label:\n            df[l] = df[l].map({'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2})\n            # Ensure any remaining NaNs are filled with a valid class\n            df[l] = df[l].fillna(0).astype(int)\n\n        # Set up image transformations\n        self.wide_resize = v2.Resize((128, 224))\n        self.rec_resize = v2.Resize((256, 256))\n        self.resize = v2.Resize((128, 128))\n        self.resize_3d = v2.Resize((256, 256))\n        self.pre_resize = v2.Resize((512, 512))\n        \n        # Data augmentation for training\n        self.wide_transforms = A.Compose([\n            A.RandomBrightnessContrast(p=0.25),\n            # A.ShiftScaleRotate(shift_limit=0.1, scale_limit=(-0.1, 0.1), rotate_limit=20, border_mode=0, p=0.5),\n            A.Resize(128, 224),\n        ])\n        self.rec_transforms = A.Compose([\n            A.HorizontalFlip(p=0.5),\n            A.RandomBrightnessContrast(p=0.25),\n            # A.ShiftScaleRotate(shift_limit=0.1, scale_limit=(-0.1, 0.1), rotate_limit=20, border_mode=0, p=0.5),\n            A.Resize(256, 256),\n        ])\n        \n    def __getitem__(self, index):\n        study_id = self.id[index]\n        res = {}\n        \n        # Load images based on condition\n        ax_img, ax_depth = self.volume_ss(study_id)\n        res['ax'] = ax_img.to(torch.float32)\n        res['ax_depth'] = ax_depth\n\n        # Extract labels for this study\n        study_labels = self.df.loc[(self.df.study_id==study_id), self.label]\n        \n        # Verify we have valid labels\n        if len(study_labels) > 0:\n            # Convert to tensor\n            label_values = study_labels.values.squeeze()\n            # Check for NaN or problematic values and replace with 0\n            if isinstance(label_values, np.ndarray):\n                # Replace any NaN or invalid values with 0\n                label_values = np.nan_to_num(label_values, nan=0.0)\n                # Ensure all values are 0, 1, or 2\n                label_values = np.clip(label_values, 0, 2).astype(np.int64)\n            \n            label = torch.tensor(label_values, dtype=torch.long)\n            res['label'] = label\n        else:\n            # Create a fallback label of all zeros\n            num_labels = len(self.label)\n            res['label'] = torch.zeros(num_labels, dtype=torch.long)\n            \n        return res\n    \n    def crop(self, image, x, y, z, x_left, x_right, y_bottom, y_top, wide):\n        size = [image[i].shape for i in z]\n        data = torch.stack([\n            self.pre_resize(torch.tensor(image[i])[None, ...]).squeeze()[\n                max(int((y/shape[0]) * 512 - y_top), 0): int((y/shape[0]) * 512 + y_bottom),\n                max(int((x/shape[1]) * 512 - x_left), 0): int((x/shape[1]) * 512 + x_right)\n            ]\n            for i, shape in zip(z, size)\n        ])\n\n        if wide:\n            data = self.wide_resize(data)\n        else:\n            data = self.rec_resize(data)\n\n        if self.usage == 'train':\n\n            if wide:\n                transformed = self.wide_transforms(\n                    image=data.numpy().transpose((1,2,0)).astype(np.float32)\n                )['image']\n                data = torch.from_numpy(transformed.transpose((2,0,1))).float()\n            else:\n                transformed = self.rec_transforms(\n                    image=data.numpy().transpose((1,2,0)).astype(np.float32)\n                )['image']\n                data = torch.from_numpy(transformed.transpose((2,0,1))).float()\n\n        return data\n\n    def volume_ss(self, study_id):\n        ax_meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Axial T2')]\n        # sagt1_meta = self.meta.loc[(self.meta.study_id==study_id) & (self.meta.series_description=='Sagittal T1')]\n        \n        ax_meta = ax_meta.sort_values('ipp_z', ascending=False).reset_index(drop=True)\n        # sagt1_meta = sagt1_meta.sort_values('ipp_x', ascending=True).reset_index(drop=True)\n        \n        ax_img = [self.normalize(self.load_dicom(f'/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{row.study_id}/{row.series_id}/{row.instance_number}.dcm')) for _, row in ax_meta.iterrows()]\n        # sagt1_img = [self.normalize(self.load_dicom(f'/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{row.study_id}/{row.series_id}/{row.instance_number}.dcm')) for _, row in sagt1_meta.iterrows()]\n        \n        ax_right_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Right Subarticular Stenosis')]\n        ax_left_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Left Subarticular Stenosis')]\n        # sagt1_right_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Right Neural Foraminal Narrowing')]\n        # sagt1_left_sub_coor = self.coor.loc[(self.coor.study_id==study_id) & (self.coor.condition=='Left Neural Foraminal Narrowing')]\n        \n        # AXIAL T2\n        ax_right_dict = {}\n        ax_right_label_dict = {}\n        for _, row in ax_right_sub_coor.iterrows():\n            label = 2\n            u = np.random.uniform(0, 1)\n            if u < 0.5:\n                z_shift = 0\n            elif u > 0.5 and u < 0.85:\n                z_shift = np.random.choice([-1, 1])\n            else:\n                z_shift = np.random.choice([-2, 2])\n            if self.usage == 'train':\n                y_shift = random.randint(-10, 10)\n                x_shift = random.randint(-10, 10)\n            else:\n                y_shift = 0\n                x_shift = 0\n            ax_right_label_dict[row.level] = label + z_shift\n            ax_meta_sub = ax_meta.loc[(ax_meta.series_id==row.series_id)]\n            ax_meta_sub_original_idx = ax_meta_sub.index.tolist()\n            ax_meta_sub = ax_meta_sub.reset_index(drop=True)\n            mid = ax_meta_sub.loc[(ax_meta_sub.instance_number==row.instance_number)].index[0]\n            z = [min(max(mid+w+z_shift, 0), len(ax_meta_sub)-1) for w in range(-(self.ax_window-1)//2, ((self.ax_window-1)//2)+1)]\n            ax_right_dict[row.level] = self.crop([ax_img[i] for i in range(len(ax_img)) if i in ax_meta_sub_original_idx], row.x+x_shift, row.y+y_shift, z, 160-16, 32+16, 64+32, 64+32, wide=False)\n        ax_left_dict = {}\n        ax_left_label_dict = {}\n        for _, row in ax_left_sub_coor.iterrows():\n            label = 2\n            u = np.random.uniform(0, 1)\n            if u < 0.5:\n                z_shift = 0\n            elif u > 0.5 and u < 0.85:\n                z_shift = np.random.choice([-1, 1])\n            else:\n                z_shift = np.random.choice([-2, 2])\n            if self.usage == 'train':\n                y_shift = random.randint(-10, 10)\n                x_shift = random.randint(-10, 10)\n            else:\n                y_shift = 0\n                x_shift = 0\n            ax_left_label_dict[row.level] = label + z_shift\n            ax_meta_sub = ax_meta.loc[(ax_meta.series_id==row.series_id)]\n            ax_meta_sub_original_idx = ax_meta_sub.index.tolist()\n            ax_meta_sub = ax_meta_sub.reset_index(drop=True)\n            mid = ax_meta_sub.loc[(ax_meta_sub.instance_number==row.instance_number)].index[0]\n            z = [min(max(mid+w+z_shift, 0), len(ax_meta_sub)-1) for w in range(-(self.ax_window-1)//2, ((self.ax_window-1)//2)+1)]\n            ax_left_dict[row.level] = self.crop([ax_img[i] for i in range(len(ax_img)) if i in ax_meta_sub_original_idx], row.x+x_shift, row.y+y_shift, z, 32+16, 160-16, 64+32, 64+32, wide=False)\n\n        ax_right_img = [ax_right_dict.get(l, torch.zeros((self.ax_window, 256, 256))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        ax_left_img = [ax_left_dict.get(l, torch.zeros((self.ax_window, 256, 256))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        # sagt1_right_img = [sagt1_right_dict.get(l, torch.zeros((self.sag_window, 128, 224))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        # sagt1_left_img = [sagt1_left_dict.get(l, torch.zeros((self.sag_window, 128, 224))) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']]\n        ax_right_label = {'right_' + l: torch.tensor(ax_right_label_dict.get(l, 2)).to(torch.long) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']}\n        ax_left_label = {'left_' + l: torch.tensor(ax_left_label_dict.get(l, 2)).to(torch.long) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']}\n        # sagt1_right_label = {'right_' + l: torch.tensor(sagt1_right_label_dict.get(l, 2)).to(torch.long) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']}\n        # sagt1_left_label = {'left_' + l: torch.tensor(sagt1_left_label_dict.get(l, 2)).to(torch.long) for l in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']}\n        ax_label = dict(**ax_left_label, **ax_right_label)\n        # sagt1_label = dict(**sagt1_left_label, **sagt1_right_label)\n        ax_img = ax_left_img + ax_right_img\n        # sagt1_img = sagt1_left_img + sagt1_right_img\n        return torch.stack(ax_img).contiguous(), ax_label,\n        \n    def normalize(self, x):\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 __len__(self):\n        return len(self.id)\n\n    def load_dicom(self, path):\n        dicom = pydicom.dcmread(path)      \n        return dicom.pixel_array\n\n    def check_position(self, first, last, pos):\n        first_dcm = pydicom.read_file(first)\n        last_dcm = pydicom.read_file(last)\n        first_dcm = first_dcm.ImagePositionPatient[pos]\n        last_dcm = last_dcm.ImagePositionPatient[pos]\n        if pos == 0:\n            return first_dcm > last_dcm\n        elif pos == 2:\n            return first_dcm < last_dcm\n        else:\n            raise ValueError\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-06T21:08:10.484761Z","iopub.execute_input":"2025-05-06T21:08:10.484975Z","iopub.status.idle":"2025-05-06T21:08:10.511440Z","shell.execute_reply.started":"2025-05-06T21:08:10.484960Z","shell.execute_reply":"2025-05-06T21:08:10.510684Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class LSTMMIL(nn.Module):\n    def __init__(self, input_dim, hidden_dim, num_classes):\n        super(LSTMMIL, self).__init__()\n        self.lstm = nn.LSTM(input_dim, input_dim//2, num_layers=2, batch_first=True, dropout=0.1, bidirectional=True)\n        self.aux_attention = nn.Sequential(\n            nn.Tanh(),\n            nn.Linear(input_dim, 1)\n        )\n        self.attention = nn.Sequential(\n            nn.Tanh(),\n            nn.Linear(input_dim, 1)\n        )\n    def forward(self, bags):\n        \"\"\"\n        Args:\n            bags: (batch_size, num_instances, input_dim)\n\n        Returns:\n            logits: (batch_size, num_classes)\n        \"\"\"\n        batch_size, num_instances, input_dim = bags.size()\n        bags_lstm, _ = self.lstm(bags)\n        attn_scores = self.attention(bags_lstm).squeeze(-1)\n        aux_attn_scores = self.aux_attention(bags_lstm).squeeze(-1)\n        attn_weights = torch.softmax(attn_scores, dim=-1)\n        weighted_instances = torch.bmm(attn_weights.unsqueeze(1), bags_lstm).squeeze(1)  # (batch_size, input_dim)\n        return weighted_instances, aux_attn_scores\n\nclass SSMIL(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.ax_encoder = timm.create_model('tf_efficientnetv2_s.in21k_ft_in1k', in_chans=1, pretrained=True, num_classes=0)\n        self.ax_flatten = nn.Sequential(nn.AdaptiveAvgPool2d((1,1)),\n                                    nn.Flatten(1))\n        self.ax_num_features = self.ax_encoder.num_features\n        self.ax_head = LSTMMIL(self.ax_num_features, 512, 3)\n        self.out = nn.Linear(self.ax_num_features, 3)\n        self.dropout = nn.Dropout(0.1)\n    def forward(self, ax):\n        if isinstance(ax, tuple):\n            ax = ax[0]\n        ax_shape = ax.shape\n        ax = ax.reshape(ax_shape[0]*ax_shape[1]*ax_shape[2], 1, ax_shape[-2], ax_shape[-1])\n        ax = self.ax_encoder.forward_features(ax)\n        ax = self.ax_flatten(ax)\n        ax = ax.reshape(ax_shape[0]*ax_shape[1], ax_shape[2], -1)\n        ax_weighted_sum, ax_attn = self.ax_head(ax)\n        ax_attn = ax_attn.reshape(ax_shape[0], ax_shape[1], -1)\n        out = ax_weighted_sum\n        out = self.dropout(out)\n        out = self.out(out)\n        ax_attn = {'left_L1/L2': ax_attn[:, 0, :],'left_L2/L3': ax_attn[:, 1, :],'left_L3/L4': ax_attn[:, 2, :], 'left_L4/L5': ax_attn[:, 3, :], 'left_L5/S1': ax_attn[:, 4, :],\n                'right_L1/L2': ax_attn[:, 5, :], 'right_L2/L3': ax_attn[:, 6, :], 'right_L3/L4': ax_attn[:, 7, :], 'right_L4/L5': ax_attn[:, 8, :], 'right_L5/S1': ax_attn[:, 9, :]}\n        return out, ax_attn","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-06T21:08:10.512434Z","iopub.execute_input":"2025-05-06T21:08:10.512911Z","iopub.status.idle":"2025-05-06T21:08:10.530127Z","shell.execute_reply.started":"2025-05-06T21:08:10.512879Z","shell.execute_reply":"2025-05-06T21:08:10.529493Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import torch\n# import torch.nn as nn\n# import timm\n\n# class GatedAttentionMIL(nn.Module):\n#     def __init__(self, input_dim, hidden_dim, num_classes):\n#         super(GatedAttentionMIL, self).__init__()\n#         self.attention_V = nn.Linear(input_dim, hidden_dim)\n#         self.attention_U = nn.Linear(input_dim, hidden_dim)\n#         self.attention_weights = nn.Linear(hidden_dim, 1)\n\n#     def forward(self, bags):\n#         \"\"\"\n#         Args:\n#             bags: (batch_size, num_instances, input_dim)\n\n#         Returns:\n#             weighted_instances: (batch_size, input_dim)\n#             attn_scores: (batch_size, num_instances)\n#         \"\"\"\n#         # Apply gating: tanh(Vx) * sigmoid(Ux)\n#         A_V = torch.tanh(self.attention_V(bags))  # (batch_size, num_instances, hidden_dim)\n#         A_U = torch.sigmoid(self.attention_U(bags))  # (batch_size, num_instances, hidden_dim)\n#         A = A_V * A_U  # (batch_size, num_instances, hidden_dim)\n\n#         # Compute attention scores\n#         attn_scores = self.attention_weights(A).squeeze(-1)  # (batch_size, num_instances)\n\n#         # Softmax over instances\n#         attn_weights = torch.softmax(attn_scores, dim=-1)  # (batch_size, num_instances)\n\n#         # Weighted sum of instance features\n#         weighted_instances = torch.bmm(attn_weights.unsqueeze(1), bags).squeeze(1)  # (batch_size, input_dim)\n\n#         return weighted_instances, attn_scores\n\n# class SSMIL(nn.Module):\n#     def __init__(self):\n#         super().__init__()\n#         # Encoder: EfficientNetV2 backbone\n#         self.ax_encoder = timm.create_model(\n#             'tf_efficientnetv2_s.in21k_ft_in1k', in_chans=1, pretrained=True, num_classes=0\n#         )\n#         self.ax_flatten = nn.Sequential(\n#             nn.AdaptiveAvgPool2d((1, 1)),\n#             nn.Flatten(1)\n#         )\n#         self.ax_num_features = self.ax_encoder.num_features\n\n#         # Updated: using GatedAttentionMIL instead of LSTM-MIL\n#         self.ax_head = GatedAttentionMIL(self.ax_num_features, 512, 3)\n#         self.out = nn.Linear(self.ax_num_features, 3)\n#         self.dropout = nn.Dropout(0.1)\n\n#     def forward(self, ax):\n#         \"\"\"\n#         Args:\n#             ax: Tensor of shape (batch_size, num_slices, num_levels, H, W)\n\n#         Returns:\n#             out: (batch_size, num_classes)\n#             ax_attn: dict of attention scores for each vertebral level\n#         \"\"\"\n#         if isinstance(ax, tuple):\n#             ax = ax[0]\n#         ax_shape = ax.shape  # (batch_size, num_slices, num_levels, H, W)\n\n#         # Reshape: (B * S * L, 1, H, W)\n#         ax = ax.reshape(ax_shape[0] * ax_shape[1] * ax_shape[2], 1, ax_shape[-2], ax_shape[-1])\n\n#         # CNN encoder\n#         ax = self.ax_encoder.forward_features(ax)  # (B * S * L, C, H', W')\n#         ax = self.ax_flatten(ax)  # (B * S * L, C)\n\n#         # Reshape back: (B * S, L, C)\n#         ax = ax.reshape(ax_shape[0] * ax_shape[1], ax_shape[2], -1)\n\n#         # Apply MIL head\n#         ax_weighted_sum, ax_attn = self.ax_head(ax)  # (B*S, C), (B*S, L)\n\n#         # Reshape attention: (B, S, L)\n#         ax_attn = ax_attn.reshape(ax_shape[0], ax_shape[1], -1)\n\n#         # Classifier\n#         out = self.dropout(ax_weighted_sum)\n#         out = self.out(out)\n\n#         # Build attention dictionary\n#         ax_attn = {\n#             'left_L1/L2': ax_attn[:, 0, :],\n#             'left_L2/L3': ax_attn[:, 1, :],\n#             'left_L3/L4': ax_attn[:, 2, :],\n#             'left_L4/L5': ax_attn[:, 3, :],\n#             'left_L5/S1': ax_attn[:, 4, :],\n#             'right_L1/L2': ax_attn[:, 5, :],\n#             'right_L2/L3': ax_attn[:, 6, :],\n#             'right_L3/L4': ax_attn[:, 7, :],\n#             'right_L4/L5': ax_attn[:, 8, :],\n#             'right_L5/S1': ax_attn[:, 9, :],\n#         }\n#         return out, ax_attn\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-06T21:08:10.530752Z","iopub.execute_input":"2025-05-06T21:08:10.530955Z","iopub.status.idle":"2025-05-06T21:08:10.544266Z","shell.execute_reply.started":"2025-05-06T21:08:10.530941Z","shell.execute_reply":"2025-05-06T21:08:10.543699Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Train","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-05-06T21:08:10.544928Z","iopub.execute_input":"2025-05-06T21:08:10.545251Z","iopub.status.idle":"2025-05-06T21:08:10.649435Z","shell.execute_reply.started":"2025-05-06T21:08:10.545142Z","shell.execute_reply":"2025-05-06T21:08:10.648588Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loss","metadata":{}},{"cell_type":"code","source":"class NFNSSDepthDetectLoss(nn.Module):\n    def __init__(self):\n        super(NFNSSDepthDetectLoss, self).__init__()\n        self.ce_loss = nn.CrossEntropyLoss()\n\n    def forward(self, outputs, targets):\n        loss = 0\n        for level in ['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            target = targets[level].to(outputs[level].device)  # Move to same device\n            _loss = nn.functional.cross_entropy(outputs[level], target.reshape(-1))\n            loss += _loss\n            \n        return loss / 10  ","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-06T21:08:10.650535Z","iopub.execute_input":"2025-05-06T21:08:10.650819Z","iopub.status.idle":"2025-05-06T21:08:10.655843Z","shell.execute_reply.started":"2025-05-06T21:08:10.650795Z","shell.execute_reply":"2025-05-06T21:08:10.655167Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SCSNFNSSLoss(nn.Module):\n    def __init__(self, is_train=False):\n        super(SCSNFNSSLoss, self).__init__()\n        self.loss = nn.CrossEntropyLoss(weight=torch.tensor([1.0, 2.0, 4.0]).to(device))\n        self.aux_loss_ax = nn.CrossEntropyLoss(weight=torch.tensor([1.0, 2.0, 4.0]).to(device))\n        self.aux_loss_sagt2 = nn.CrossEntropyLoss(weight=torch.tensor([1.0, 2.0, 4.0]).to(device))\n        self.aux_loss_sagt1 = nn.CrossEntropyLoss(weight=torch.tensor([1.0, 2.0, 4.0]).to(device))\n        self.is_train = is_train\n    def forward(self, outputs, targets, ax=None, sagt2=None, sagt1=None):\n        targets = targets.reshape(-1)\n        #print(outputs, targets, outputs.shape, targets.shape)\n        ax_loss = 0\n        sagt2_loss = 0\n        sagt1_loss = 0\n        num_loss = 1\n        loss = self.loss(outputs, targets)\n        if ax is not None:\n            ax_loss = self.aux_loss_ax(ax, targets)\n            if self.is_train:\n                loss += ax_loss*0.5\n                num_loss += 0.5\n        if sagt2 is not None:\n            sagt2_loss = self.aux_loss_sagt2(sagt2, targets)\n            if self.is_train:\n                loss += sagt2_loss*0.5\n                num_loss += 0.5\n        if sagt1 is not None:\n            sagt1_loss = self.aux_loss_sagt1(sagt1, targets)\n            if self.is_train:\n                loss += sagt1_loss*0.5\n                num_loss += 0.5\n        #print(loss)\n        loss = loss/num_loss\n        return loss, ax_loss, sagt2_loss, sagt1_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-06T21:08:10.656559Z","iopub.execute_input":"2025-05-06T21:08:10.656810Z","iopub.status.idle":"2025-05-06T21:08:10.668581Z","shell.execute_reply.started":"2025-05-06T21:08:10.656788Z","shell.execute_reply":"2025-05-06T21:08:10.667691Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Accuracy","metadata":{}},{"cell_type":"code","source":"def calculate_accuracy_score(outputs, targets):\n    \"\"\"Computes accuracy for model predictions across multiple spinal levels.\"\"\"\n    correct_total = 0\n    samples_total = 0\n    \n    pred_classes = {}\n    true_classes = {}\n    \n    for level in outputs.keys():\n        if level in targets:\n            # Get predicted class (argmax along class dimension)\n            preds = torch.argmax(outputs[level], dim=1)\n\n            # Ensure targets are reshaped and on the correct device\n            true_labels = targets[level].reshape(-1).long()\n\n            # Store predictions and ground truth (ensure they're on CPU)\n            pred_classes[level] = preds.detach().cpu().numpy()\n            true_classes[level] = true_labels.detach().cpu().numpy()\n            \n            # Calculate accuracy for this level\n            correct = (preds == true_labels).sum().item()\n            correct_total += correct\n            samples_total += true_labels.size(0)\n    \n    # Avoid division by zero\n    if samples_total == 0:\n        return 0.0, pred_classes, true_classes\n    \n    accuracy = correct_total / samples_total\n    return accuracy, pred_classes, true_classes\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-06T21:08:10.669291Z","iopub.execute_input":"2025-05-06T21:08:10.669512Z","iopub.status.idle":"2025-05-06T21:08:10.682425Z","shell.execute_reply.started":"2025-05-06T21:08:10.669487Z","shell.execute_reply":"2025-05-06T21:08:10.681765Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Tolerance Matrix","metadata":{}},{"cell_type":"code","source":"def calculate_classification_tolerance(outputs, batch, tolerances=[0, 1, 2]):\n    \"\"\"Calculate tolerance metrics for classification\"\"\"\n    tolerance_counts = {f\"±{tol}\": 0 for tol in tolerances}\n    tolerance_counts[\">±2\"] = 0\n    total_predictions = 0\n    \n    predicted_classes = outputs.argmax(dim=1)\n    \n    if 'label' in batch:\n        targets = batch['label']\n        abs_diff = torch.abs(predicted_classes - targets)\n        total_predictions = targets.size(0)\n        \n        # Count predictions within each tolerance\n        for tol in tolerances:\n            tolerance_counts[f\"±{tol}\"] = (abs_diff <= tol).sum().item()\n        \n        # Count predictions beyond the maximum tolerance\n        tolerance_counts[\">±2\"] = (abs_diff > max(tolerances)).sum().item()\n    else:\n        # Process by level\n        levels = [key for key in batch.keys() if key.startswith('spinal_canal_stenosis')]\n        for i, level_key in enumerate(levels):\n            level_preds = outputs[i::len(levels)].argmax(dim=1)\n            level_targets = batch[level_key]\n            abs_diff = torch.abs(level_preds - level_targets)\n            total_predictions += level_targets.size(0)\n            \n            # Count predictions within each tolerance\n            for tol in tolerances:\n                tolerance_counts[f\"±{tol}\"] += (abs_diff <= tol).sum().item()\n            \n            # Count predictions beyond the maximum tolerance\n            tolerance_counts[\">±2\"] += (abs_diff > max(tolerances)).sum().item()\n    \n    return tolerance_counts, total_predictions","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-06T21:08:10.683127Z","iopub.execute_input":"2025-05-06T21:08:10.683349Z","iopub.status.idle":"2025-05-06T21:08:10.699384Z","shell.execute_reply.started":"2025-05-06T21:08:10.683307Z","shell.execute_reply":"2025-05-06T21:08:10.698761Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def train_spine_model(model, train_coor, train_meta, df, series, n_folds=5, epochs=7, batch_size=1,\n                      learning_rate=0.001, weight_decay=0.0001, patience=5,\n                      mixed_precision=True, experiment_name=\"spine_classification\", tolerances=[0, 1, 2],\n                      condition='ss'):\n    \"\"\"\n    Improved training function for spine classification model using advanced loss functions\n\n    Args:\n        model: Your defined model (SCSMIL or equivalent)\n        train_coor: DataFrame containing coordinate annotations\n        train_meta: DataFrame containing metadata\n        df: DataFrame with labels\n        series: Series data\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 values for evaluation\n        condition: Condition type ('scs', 'ss', or 'nfn')\n    \"\"\"\n\n    # Set device\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\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\"Condition: {condition}\")\n    print(f\"==============================\\n\")\n\n    # Initialize scaler for mixed precision\n    scaler = torch.amp.GradScaler() if mixed_precision else None\n\n    # Initialize loss functions from your friend's code\n    loss_module = SCSNFNSSLoss(is_train=True)\n    val_loss_module = SCSNFNSSLoss(is_train=False)\n    ss_depth_loss_module = NFNSSDepthDetectLoss()  # You might want to use NFNSSDepthDetectLoss() for NFN condition\n\n    # Cross-validation loop\n    all_val_accuracies = []\n    all_val_f1_scores = []\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 = SeverityPrediction(\n            df=df,\n            series=series,\n            coor=train_coor.loc[train_coor.fold != fold],\n            meta=train_meta.loc[train_meta.fold != fold],\n            condition=condition,\n            usage='train'\n        )\n\n        valid_dataset = SeverityPrediction(\n            df=df,\n            series=series,\n            coor=train_coor.loc[train_coor.fold == fold],\n            meta=train_meta.loc[train_meta.fold == fold],\n            condition=condition,\n            usage='valid'\n        )\n        \n        train_loader = DataLoader(\n            train_dataset, \n            batch_size=batch_size,\n            shuffle=True, \n            num_workers=0, \n            pin_memory=True\n        )\n\n        valid_loader = DataLoader(\n            valid_dataset, \n            batch_size=batch_size,\n            shuffle=False, \n            num_workers=0, \n            pin_memory=True\n        )\n\n        # Initialize model, optimizer and loss\n        model.to(device)\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        # Initialize tracking variables\n        train_losses, val_losses = [], []\n        train_accs, val_accs = [], []\n        best_val_acc = 0.0\n        counter = 0  # For early stopping\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            running_ax_depth_loss = 0.0\n            total_acc = 0.0\n\n            progress_bar = tqdm(enumerate(train_loader), total=len(train_loader),\n                               desc=\"Training Progress\", leave=False)\n\n            all_train_preds = []\n            all_train_targets = []\n            \n            # Initialize tolerance counts and totals for both train and validation\n            epoch_train_tolerances = {f\"±{tol}\": 0 for tol in tolerances}\n            epoch_train_tolerances[\">±2\"] = 0\n            epoch_val_tolerances = {f\"±{tol}\": 0 for tol in tolerances}\n            epoch_val_tolerances[\">±2\"] = 0\n            train_total_predictions = 0\n            val_total_predictions = 0\n\n            for batch_idx, batch in progress_bar:\n                # Process batch data based on condition\n                ax = batch['ax'].to(device)\n                label = batch['label'].to(device)\n                ax_depth = batch['ax_depth']\n                \n                # Zero gradients before forward pass\n                optimizer.zero_grad()\n\n                # Mixed precision training\n                if mixed_precision and torch.cuda.is_available():\n                    with torch.amp.autocast('cuda'):\n                        # Forward pass with model\n                        preds, ax_depth_pred = model(ax)\n                        loss, _, _, _ = loss_module(preds, label)\n                        \n                        # Additional depth losses\n                        ax_depth_loss = ss_depth_loss_module(ax_depth_pred, ax_depth)\n                        \n                        # Combined loss\n                        total_loss = loss + ax_depth_loss\n\n                    # Scale loss and backward pass\n                    scaler.scale(total_loss).backward()\n\n                    # Unscale before gradient clipping\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n\n                    # Optimizer step\n                    scaler.step(optimizer)\n                    scaler.update()\n                    scheduler.step()\n                else:\n                    # Standard precision training\n                    preds, ax_depth_pred = model(ax)\n                    loss, _, _, _ = loss_module(preds, label)\n                    \n                    # Additional depth losses\n                    ax_depth_loss = ss_depth_loss_module(ax_depth_pred, ax_depth)\n                    \n                    # Combined loss\n                    total_loss = loss + ax_depth_loss\n                    \n                    total_loss.backward()\n\n                    # Gradient clipping\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n\n                    # Optimizer step\n                    optimizer.step()\n                    scheduler.step()\n\n                # Track metrics\n                running_loss += loss.item()\n                running_ax_depth_loss += ax_depth_loss.item() \n\n                # Store predictions for later evaluation\n                all_train_preds.append(preds.detach().cpu())\n                all_train_targets.append(label.detach().cpu())\n                \n                # Calculate batch accuracy\n                # Handle different possible prediction shapes\n                B = preds.size(0)  # batch size\n                \n                # Ground truth should be shape [B], if not already\n                label = label.view(-1)\n                \n                # Convert predictions to class indices\n                pred_classes = preds.argmax(dim=-1)\n                \n                # Compute per-element accuracy\n                correct = (pred_classes == label).float().sum()\n                batch_acc = correct / label.numel()\n                total_acc += batch_acc.item()\n                \n                # Calculate tolerance metrics\n                abs_diff = torch.abs(pred_classes - label)\n                batch_total = label.size(0)\n                train_total_predictions += batch_total\n                \n                # Count predictions within each tolerance\n                for tol in tolerances:\n                    count = (abs_diff <= tol).sum().item()\n                    epoch_train_tolerances[f\"±{tol}\"] += count\n                \n                # Count predictions beyond the maximum tolerance\n                count = (abs_diff > max(tolerances)).sum().item()\n                epoch_train_tolerances[\">±2\"] += count\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                    ax_depth_loss=f\"{ax_depth_loss.item():.4f}\",\n                    total_loss=f\"{total_loss.item():.4f}\",\n                    acc=f\"{batch_acc.item():.4f}\",\n                    lr=f\"{current_lr:.6f}\"\n                )\n\n                # Free memory\n                variables_to_delete = [\n                    'batch', 'loss', 'total_loss',\n                    'ax_depth_loss', 'ax_depth_pred',\n                    'preds', 'label'\n                ]\n                \n                for var in variables_to_delete:\n                    if var in locals():\n                        del locals()[var]\n                \n                torch.cuda.empty_cache()\n                torch.cuda.ipc_collect()  # Optional: helps with inter-process caching\n\n            # Calculate epoch metrics\n            epoch_train_loss = running_loss / len(train_loader)\n            epoch_train_ax_depth_loss = running_ax_depth_loss / len(train_loader) \n            epoch_train_acc = total_acc / len(train_loader)\n            train_losses.append(epoch_train_loss)\n            train_accs.append(epoch_train_acc)\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] = epoch_train_tolerances[k] * 100 / 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('Accuracy/train', epoch_train_acc, epoch)\n            fold_writer.add_scalar('Loss/train_ax_depth', epoch_train_ax_depth_loss, epoch)\n            writer.add_scalar(f'Loss/train/fold_{fold}', epoch_train_loss, epoch)\n            writer.add_scalar(f'Accuracy/train/fold_{fold}', epoch_train_acc, epoch)\n\n            print(f\"🔥 Training Loss: {epoch_train_loss:.4f} | Accuracy: {epoch_train_acc:.4f}\")\n            print(f\"   Depth Losses - AX: {epoch_train_ax_depth_loss:.4f}\")\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            # === Validation phase ===\n            model.eval()\n            val_running_loss = 0.0\n            val_running_ax_depth_loss = 0.0\n            val_total_acc = 0.0\n            \n            all_val_preds = []\n            all_val_targets = []\n\n            with torch.no_grad():\n                for batch in tqdm(valid_loader, desc=\"Validation Progress\", leave=False):\n                    ax = batch['ax'].to(device)\n                    label = batch['label'].to(device)\n                    ax_depth = batch['ax_depth']\n                    \n                    # Forward pass\n                    preds, ax_depth_pred = model(ax)\n                    loss, _, _, _ = val_loss_module(preds, label)\n                    \n                    # Additional depth losses\n                    ax_depth_loss = ss_depth_loss_module(ax_depth_pred, ax_depth)\n                    \n                    val_running_loss += loss.item()\n                    val_running_ax_depth_loss += ax_depth_loss.item()\n                    \n                    # Handle different possible prediction shapes\n                    B = preds.size(0)  # batch size\n                    \n                    # Ground truth should be shape [B], if not already\n                    label = label.view(-1)\n                    \n                    # Convert predictions to class indices\n                    pred_classes = preds.argmax(dim=-1)\n                    \n                    # Compute per-element accuracy\n                    correct = (pred_classes == label).float().sum()\n                    batch_acc = correct / label.numel()\n                    val_total_acc += batch_acc.item()\n\n                    # Store predictions for evaluation\n                    all_val_preds.append(preds.cpu())\n                    all_val_targets.append(label.cpu())\n                    \n                    # Calculate tolerance metrics for validation\n                    abs_diff = torch.abs(pred_classes - label)\n                    batch_total = label.size(0)\n                    val_total_predictions += batch_total\n                    \n                    # Count predictions within each tolerance\n                    for tol in tolerances:\n                        count = (abs_diff <= tol).sum().item()\n                        epoch_val_tolerances[f\"±{tol}\"] += count\n                    \n                    # Count predictions beyond the maximum tolerance\n                    count = (abs_diff > max(tolerances)).sum().item()\n                    epoch_val_tolerances[\">±2\"] += count\n\n                    # Free memory\n                    variables_to_delete = [\n                        'batch', 'loss',\n                        'ax_depth_loss',\n                        'preds', 'label'\n                    ]\n                    \n                    for var in variables_to_delete:\n                        if var in locals():\n                            del locals()[var]\n                    \n                    torch.cuda.empty_cache()\n                    torch.cuda.ipc_collect()\n            \n            # Calculate validation metrics\n            epoch_val_loss = val_running_loss / len(valid_loader)\n            epoch_val_ax_depth_loss = val_running_ax_depth_loss / len(valid_loader) \n            epoch_val_acc = val_total_acc / len(valid_loader)\n            val_losses.append(epoch_val_loss)\n            val_accs.append(epoch_val_acc)\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            # Combine predictions for metrics\n            all_val_preds_tensor = torch.cat(all_val_preds)\n            all_val_targets_tensor = torch.cat(all_val_targets)\n            \n            # Calculate F1 score - FIX THE SHAPE ISSUE HERE\n            pred_classes = all_val_preds_tensor.argmax(dim=-1).cpu().numpy()\n            target_classes = all_val_targets_tensor.cpu().numpy()\n            \n            # Calculate the overall F1 score (no per-column calculation)\n            # This fixes the IndexError by not assuming target_classes has a second dimension\n            avg_f1_score = f1_score(\n                target_classes, \n                pred_classes, \n                average='weighted', \n                zero_division=0\n            )\n\n            # Log validation metrics\n            fold_writer.add_scalar('Loss/val', epoch_val_loss, epoch)\n            fold_writer.add_scalar('Accuracy/val', epoch_val_acc, epoch)\n            fold_writer.add_scalar('F1/val', avg_f1_score, epoch)\n            fold_writer.add_scalar('Loss/val_ax_depth', epoch_val_ax_depth_loss, epoch)\n            writer.add_scalar(f'Loss/val/fold_{fold}', epoch_val_loss, epoch)\n            writer.add_scalar(f'Accuracy/val/fold_{fold}', epoch_val_acc, epoch)\n            writer.add_scalar(f'F1/val/fold_{fold}', avg_f1_score, 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} | Accuracy: {epoch_val_acc:.4f} | F1 Score: {avg_f1_score:.4f}\")\n            print(f\"   Depth Losses - AX: {epoch_val_ax_depth_loss:.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 accuracy\n            if epoch_val_acc > best_val_acc:\n                best_val_acc = epoch_val_acc\n                best_f1 = avg_f1_score\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_acc': epoch_val_acc,\n                    'val_f1': avg_f1_score,\n                    'val_tolerances': epoch_val_tolerances,\n                }, f'models/{experiment_name}_fold_{fold}_best.pt')\n\n                print(f\"📌 New best model saved with Accuracy: {best_val_acc:.4f} and F1: {best_f1:.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 accuracy and F1\n        all_val_accuracies.append(best_val_acc)\n        all_val_f1_scores.append(best_f1)\n        print(f\"Fold {fold+1} completed. Best validation Accuracy: {best_val_acc:.4f}, F1: {best_f1:.4f}\")\n\n        # Print classification report\n        print(\"\\n--- Classification Report ---\")\n        # Fix the classification report call to match the corrected data shapes\n        report = classification_report(\n                target_classes,\n                pred_classes,\n                zero_division=0\n            )\n        print(report)\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_acc': val_accs[-1],\n            'val_f1': avg_f1_score,\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_acc = sum(all_val_accuracies) / len(all_val_accuracies) if all_val_accuracies else 0.0\n    avg_val_f1 = sum(all_val_f1_scores) / len(all_val_f1_scores) if all_val_f1_scores else 0.0\n    print(f\"\\n===== TRAINING COMPLETE =====\")\n    print(f\"Cross-validation results:\")\n    for fold, (acc, f1) in enumerate(zip(all_val_accuracies, all_val_f1_scores)):\n        print(f\"Fold {fold+1}: Accuracy = {acc:.4f}, F1 = {f1:.4f}\")\n    print(f\"Average validation Accuracy: {avg_val_acc:.4f}\")\n    print(f\"Average validation F1 Score: {avg_val_f1:.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\"Condition: {condition}\\n\")\n        f.write(f\"Tolerance thresholds: {tolerances}\\n\")\n        f.write(\"\\nResults:\\n\")\n        for fold, (acc, f1) in enumerate(zip(all_val_accuracies, all_val_f1_scores)):\n            f.write(f\"Fold {fold+1}: Accuracy = {acc:.4f}, F1 = {f1:.4f}\\n\")\n        f.write(f\"Average validation Accuracy: {avg_val_acc:.4f}\\n\")\n        f.write(f\"Average validation F1 Score: {avg_val_f1:.4f}\\n\")\n        \n        # Add tolerance metrics summary across all folds\n        f.write(\"\\nTolerance Metrics Summary:\\n\")\n        for tol in tolerances:\n            f.write(f\"±{tol}: Represents predictions within {tol} grades of ground truth\\n\")\n        f.write(f\">±{max(tolerances)}: Represents predictions more than {max(tolerances)} grades away from ground truth\\n\")\n\n    # Close writer\n    writer.close()\n    \n    return all_val_accuracies, all_val_f1_scores, avg_val_acc, avg_val_f1\n\n# Example usage:\n# model = SSMIL()\n# accuracies, f1_scores, avg_acc, avg_f1 = train_spine_model(model, train_coor, train_meta, train_dummy, train_series, n_folds=5, condition='ss')\n\n# Example usage:\nmodel = SSMIL()\naccuracies, f1_scores, avg_acc, avg_f1 = train_spine_model(model, train_coor, train_meta, train_dummy, train_series, n_folds=5, condition='ss')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-06T21:08:10.700198Z","iopub.execute_input":"2025-05-06T21:08:10.700390Z","execution_failed":"2025-05-06T22:36:28.752Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training with visualization","metadata":{}},{"cell_type":"code","source":"\ndef train_spine_model(model, train_coor, train_meta, df, series, n_folds=5, epochs=7, batch_size=1,\n                      learning_rate=0.001, weight_decay=0.0001, patience=5,\n                      mixed_precision=True, experiment_name=\"spine_classification\", tolerances=[0, 1, 2],\n                      condition='ss'):\n    \"\"\"\n    Improved training function for spine classification model using advanced loss functions\n\n    Args:\n        model: Your defined model (SCSMIL or equivalent)\n        train_coor: DataFrame containing coordinate annotations\n        train_meta: DataFrame containing metadata\n        df: DataFrame with labels\n        series: Series data\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 values for evaluation\n        condition: Condition type ('scs', 'ss', or 'nfn')\n    \"\"\"\n\n    # Set device\n    device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n    # Create output directory for models\n    os.makedirs(\"models\", exist_ok=True)\n\n    plots_dir = f\"plots/{experiment_name}\"\n    os.makedirs(plots_dir, exist_ok=True)\n    os.makedirs(f\"{plots_dir}/confusion_matrices\", exist_ok=True)\n    os.makedirs(f\"{plots_dir}/learning_curves\", exist_ok=True)\n    os.makedirs(f\"{plots_dir}/tolerance_metrics\", exist_ok=True)\n    os.makedirs(f\"{plots_dir}/predictions\", 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\"Condition: {condition}\")\n    print(f\"==============================\\n\")\n\n    # Initialize scaler for mixed precision\n    scaler = torch.amp.GradScaler() if mixed_precision else None\n\n    # Initialize loss functions from your friend's code\n    loss_module = SCSNFNSSLoss(is_train=True)\n    val_loss_module = SCSNFNSSLoss(is_train=False)\n    nfn_depth_loss_module = NFNSSDepthDetectLoss()  # You might want to use NFNSSDepthDetectLoss() for NFN condition\n\n    # Cross-validation loop\n    all_val_accuracies = []\n    all_val_f1_scores = []\n    \n    # Initialize lists to store data for visualization\n    all_fold_train_losses = []\n    all_fold_val_losses = []\n    all_fold_train_accs = []\n    all_fold_val_accs = []\n    all_fold_train_tolerances = []\n    all_fold_val_tolerances = []\n    all_fold_confusion_matrices = []\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 = SeverityPrediction(\n            df=df,\n            series=series,\n            coor=train_coor.loc[train_coor.fold != fold],\n            meta=train_meta.loc[train_meta.fold != fold],\n            condition=condition,\n            usage='train'\n        )\n\n        valid_dataset = SeverityPrediction(\n            df=df,\n            series=series,\n            coor=train_coor.loc[train_coor.fold == fold],\n            meta=train_meta.loc[train_meta.fold == fold],\n            condition=condition,\n            usage='valid'\n        )\n        \n\n        train_loader = DataLoader(\n            train_dataset, \n            batch_size=batch_size,\n            shuffle=True, \n            num_workers=0, \n            pin_memory=True\n        )\n\n        valid_loader = DataLoader(\n            valid_dataset, \n            batch_size=batch_size,\n            shuffle=False, \n            num_workers=0, \n            pin_memory=True\n        )\n\n        # Initialize model, optimizer and loss\n        model.to(device)\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        # Initialize tracking variables\n        train_losses, val_losses = [], []\n        train_accs, val_accs = [], []\n        best_val_acc = 0.0\n        counter = 0  # For early stopping\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            running_ax_depth_loss = 0.0\n            running_sagt1_depth_loss = 0.0\n            total_acc = 0.0\n\n            progress_bar = tqdm(enumerate(train_loader), total=len(train_loader),\n                               desc=\"Training Progress\", leave=False)\n\n            all_train_preds = []\n            all_train_targets = []\n            \n            # Initialize tolerance counts and totals for both train and validation\n            epoch_train_tolerances = {f\"±{tol}\": 0 for tol in tolerances}\n            epoch_train_tolerances[\">±2\"] = 0\n            epoch_val_tolerances = {f\"±{tol}\": 0 for tol in tolerances}\n            epoch_val_tolerances[\">±2\"] = 0\n            train_total_predictions = 0\n            val_total_predictions = 0\n\n            for batch_idx, batch in progress_bar:\n                # Process batch data based on condition\n                \n                ax = batch['ax'].to(device)\n                label = batch['label'].to(device)\n                ax_depth = batch['ax_depth']\n                \n                # Zero gradients before forward pass\n                optimizer.zero_grad()\n\n                # Mixed precision training\n                if mixed_precision and torch.cuda.is_available():\n                    with torch.amp.autocast('cuda'):\n                        # Forward pass with model\n                        \n                        preds, ax_depth_pred = model(ax)\n                        loss, _, _, _ = loss_module(preds, label)\n                        \n                        # Additional depth losses\n                        ax_depth_loss = nfn_depth_loss_module(ax_depth_pred, ax_depth)\n                        \n                        # Combined loss\n                        total_loss = loss + ax_depth_loss\n\n                    # Scale loss and backward pass\n                    scaler.scale(total_loss).backward()\n\n                    # Unscale before gradient clipping\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n\n                    # Optimizer step\n                    scaler.step(optimizer)\n                    scaler.update()\n                    scheduler.step()\n                else:\n                    # Standard precision training\n                    \n                    preds, ax_depth_pred = model(ax)\n                    loss, _, _, _ = loss_module(preds, label)\n                    \n                    # Additional depth losses\n                    ax_depth_loss = nfn_depth_loss_module(ax_depth_pred, ax_depth)\n                    \n                    # Combined loss\n                    total_loss = loss + ax_depth_loss\n                    \n                    total_loss.backward()\n\n                    # Gradient clipping\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n\n                    # Optimizer step\n                    optimizer.step()\n                    scheduler.step()\n\n                # Track metrics\n                running_loss += loss.item()\n                running_ax_depth_loss += ax_depth_loss.item() \n\n                \n                # Store predictions for later evaluation\n                all_train_preds.append(preds.detach().cpu())\n                all_train_targets.append(label.detach().cpu())\n\n                \n                \n                # Calculate batch accuracy\n                \n                # Handle different possible prediction shapes\n                B = preds.size(0)  # batch size\n\n                # # preds: [B, 25, 3] → reshape to [B, 5, 5, 3] (5 disc levels × 5 views each)\n                # preds = preds.view(B, 5, 5, 3)\n                \n                # # Aggregate over 5 instances per disc level (mean pooling) → [B, 5, 3]\n                # preds = preds.mean(dim=2)\n                \n                # Ground truth should be shape [B, 5], if not already\n                label = label.view(B)\n                \n                # Convert predictions to class indices → [B, 5]\n                pred_classes = preds.argmax(dim=-1)\n                \n                # Compute per-element accuracy\n                correct = (pred_classes == label).float().sum()\n                batch_acc = correct / label.numel()\n                total_acc += batch_acc.item()\n                \n                # Calculate tolerance metrics\n                abs_diff = torch.abs(pred_classes - label)\n                batch_total = label.size(0)\n                train_total_predictions += batch_total\n                \n                # Count predictions within each tolerance\n                for tol in tolerances:\n                    count = (abs_diff <= tol).sum().item()\n                    epoch_train_tolerances[f\"±{tol}\"] += count\n                \n                # Count predictions beyond the maximum tolerance\n                count = (abs_diff > max(tolerances)).sum().item()\n                epoch_train_tolerances[\">±2\"] += count\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                    ax_depth_loss=f\"{ax_depth_loss.item():.4f}\",\n                    total_loss=f\"{total_loss.item():.4f}\",\n                    acc=f\"{batch_acc.item():.4f}\",\n                    lr=f\"{current_lr:.6f}\"\n                )\n\n                # Free memory\n                variables_to_delete = [\n                    'batch', 'loss', 'total_loss',\n                    'ax_depth_loss', 'ax_depth_pred',\n                    'preds', 'label'\n                ]\n                \n                for var in variables_to_delete:\n                    if var in locals():\n                        del locals()[var]\n                \n                torch.cuda.empty_cache()\n                torch.cuda.ipc_collect()  # Optional: helps with inter-process caching\n\n            # Calculate epoch metrics\n            epoch_train_loss = running_loss / len(train_loader)\n            epoch_train_ax_depth_loss = running_ax_depth_loss / len(train_loader) \n            epoch_train_acc = total_acc / len(train_loader)\n            train_losses.append(epoch_train_loss)\n            train_accs.append(epoch_train_acc)\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] = epoch_train_tolerances[k] * 100 / 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('Accuracy/train', epoch_train_acc, epoch)\n            \n            fold_writer.add_scalar('Loss/train_ax_depth', epoch_train_ax_depth_loss, epoch)\n            \n            writer.add_scalar(f'Loss/train/fold_{fold}', epoch_train_loss, epoch)\n            writer.add_scalar(f'Accuracy/train/fold_{fold}', epoch_train_acc, epoch)\n\n            print(f\"🔥 Training Loss: {epoch_train_loss:.4f} | Accuracy: {epoch_train_acc:.4f}\")\n            \n            print(f\"   Depth Losses - AX: {epoch_train_ax_depth_loss:.4f}\")\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            # === Validation phase ===\n            model.eval()\n            val_running_loss = 0.0\n            val_running_ax_depth_loss = 0.0\n            val_running_sagt1_depth_loss = 0.0\n            val_total_acc = 0.0\n            \n            all_val_preds = []\n            all_val_targets = []\n\n            with torch.no_grad():\n                for batch in tqdm(valid_loader, desc=\"Validation Progress\", leave=False):\n                    \n                    ax = batch['ax'].to(device)\n                    label = batch['label'].to(device)\n                    ax_depth = batch['ax_depth']\n                    \n                    # Forward pass\n                    preds, ax_depth_pred = model(ax)\n                    loss, _, _, _ = val_loss_module(preds, label)\n                    \n                    # Additional depth losses\n                    ax_depth_loss = nfn_depth_loss_module(ax_depth_pred, ax_depth)\n                    \n                    val_running_loss += loss.item()\n                    \n                    val_running_ax_depth_loss += ax_depth_loss.item()\n                    \n                    \n                    # Handle different possible prediction shapes\n                    B = preds.size(0)  # batch size\n    \n                    # # preds: [B, 25, 3] → reshape to [B, 5, 5, 3] (5 disc levels × 5 views each)\n                    # preds = preds.view(B, 5, 5, 3)\n                    \n                    # # Aggregate over 5 instances per disc level (mean pooling) → [B, 5, 3]\n                    # preds = preds.mean(dim=2)\n                    \n                    # Ground truth should be shape [B, 5], if not already\n                    label = label.view(B)\n                    \n                    # Convert predictions to class indices → [B, 5]\n                    pred_classes = preds.argmax(dim=-1)\n                    \n                    # Compute per-element accuracy\n                    correct = (pred_classes == label).float().sum()\n                    batch_acc = correct / label.numel()\n                    val_total_acc += batch_acc.item()\n\n                    \n                    # Store predictions for evaluation\n                    all_val_preds.append(preds.cpu())\n                    all_val_targets.append(label.cpu())\n                    \n                    # Calculate tolerance metrics for validation\n                    \n                    abs_diff = torch.abs(pred_classes - label)\n                    batch_total = label.size(0)\n                    val_total_predictions += batch_total\n                    \n                    # Count predictions within each tolerance\n                    for tol in tolerances:\n                        count = (abs_diff <= tol).sum().item()\n                        epoch_val_tolerances[f\"±{tol}\"] += count\n                    \n                    # Count predictions beyond the maximum tolerance\n                    count = (abs_diff > max(tolerances)).sum().item()\n                    epoch_val_tolerances[\">±2\"] += count\n\n                    # Free memory\n                    variables_to_delete = [\n                        'batch', 'loss',\n                        'ax_depth_loss',\n                        'preds', 'label'\n                    ]\n                    \n                    for var in variables_to_delete:\n                        if var in locals():\n                            del locals()[var]\n                    \n                    torch.cuda.empty_cache()\n                    torch.cuda.ipc_collect()\n            # Calculate validation metrics\n            epoch_val_loss = val_running_loss / len(valid_loader)\n            epoch_val_ax_depth_loss = val_running_ax_depth_loss / len(valid_loader) \n            epoch_val_acc = val_total_acc / len(valid_loader)\n            val_losses.append(epoch_val_loss)\n            val_accs.append(epoch_val_acc)\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            # Combine predictions for metrics\n            all_val_preds_tensor = torch.cat(all_val_preds)\n            all_val_targets_tensor = torch.cat(all_val_targets)\n            \n            # Calculate F1 score\n            pred_classes = all_val_preds_tensor.argmax(dim=-1).cpu().numpy()\n            target_classes = all_val_targets_tensor.cpu().numpy()\n            \n            # Flatten for confusion matrix later\n            target_classes_flat = target_classes.flatten()\n            pred_classes_flat = pred_classes.flatten()\n            \n            # Compute weighted F1 for each disc level (column)\n            f1_scores_per_level = [\n                f1_score(target_classes[:, i], pred_classes[:, i], average='weighted', zero_division=0)\n                for i in range(target_classes.shape[1])\n            ]\n            \n            avg_f1_score = sum(f1_scores_per_level) / len(f1_scores_per_level)\n\n            # Log validation metrics\n            fold_writer.add_scalar('Loss/val', epoch_val_loss, epoch)\n            fold_writer.add_scalar('Accuracy/val', epoch_val_acc, epoch)\n            fold_writer.add_scalar('F1/val', avg_f1_score, epoch)\n            \n            fold_writer.add_scalar('Loss/val_ax_depth', epoch_val_ax_depth_loss, epoch)\n            \n            writer.add_scalar(f'Loss/val/fold_{fold}', epoch_val_loss, epoch)\n            writer.add_scalar(f'Accuracy/val/fold_{fold}', epoch_val_acc, epoch)\n            writer.add_scalar(f'F1/val/fold_{fold}', avg_f1_score, 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} | Accuracy: {epoch_val_acc:.4f} | F1 Score: {avg_f1_score:.4f}\")\n            \n            print(f\"   Depth Losses - AX: {epoch_val_ax_depth_loss:.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 accuracy\n            if epoch_val_acc > best_val_acc:\n                best_val_acc = epoch_val_acc\n                best_f1 = avg_f1_score\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_acc': epoch_val_acc,\n                    'val_f1': avg_f1_score,\n                    'val_tolerances': epoch_val_tolerances,\n                }, f'models/{experiment_name}_fold_{fold}_best.pt')\n\n                print(f\"📌 New best model saved with Accuracy: {best_val_acc:.4f} and F1: {best_f1:.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 accuracy and F1\n        all_val_accuracies.append(best_val_acc)\n        all_val_f1_scores.append(best_f1)\n        print(f\"Fold {fold+1} completed. Best validation Accuracy: {best_val_acc:.4f}, F1: {best_f1:.4f}\")\n\n        # Print classification report\n        print(\"\\n--- Classification Report ---\")\n        report = classification_report(\n                target_classes_flat,\n                pred_classes_flat,\n                zero_division=0\n            )\n        print(report)\n\n        cm = confusion_matrix(target_classes_flat, pred_classes_flat)\n        fold_cm = {'matrix': cm, 'fold': fold}\n        all_fold_confusion_matrices.append(fold_cm)\n                \n        plt.figure(figsize=(10, 8))\n        disp = ConfusionMatrixDisplay(confusion_matrix=cm)\n        disp.plot(cmap=plt.cm.Blues, values_format='d')\n        plt.title(f'Confusion Matrix - Fold {fold+1}')\n        plt.savefig(f\"{plots_dir}/confusion_matrices/fold_{fold+1}_confusion_matrix.png\", bbox_inches='tight')\n        plt.close()\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_acc': val_accs[-1],\n            'val_f1': avg_f1_score,\n        }, f'models/{experiment_name}_fold_{fold}_final.pt')\n\n        #Later Visualization\n        all_fold_train_losses.append(train_losses)\n        all_fold_val_losses.append(val_losses)\n        all_fold_train_accs.append(train_accs)\n        all_fold_val_accs.append(val_accs)\n        all_fold_train_tolerances.append(epoch_train_tolerances)\n        all_fold_val_tolerances.append(epoch_val_tolerances)\n            \n        # Create fold-specific learning curves\n        plt.figure(figsize=(12, 5))\n            \n        plt.subplot(1, 2, 1)\n        plt.plot(train_losses, label='Train Loss')\n        plt.plot(val_losses, label='Validation Loss')\n        plt.title(f'Fold {fold+1} Loss Curves')\n        plt.xlabel('Epoch')\n        plt.ylabel('Loss')\n        plt.legend()\n        plt.grid(True, alpha=0.3)\n            \n        plt.subplot(1, 2, 2)\n        plt.plot(train_accs, label='Train Accuracy')\n        plt.plot(val_accs, label='Validation Accuracy')\n        plt.title(f'Fold {fold+1} Accuracy Curves')\n        plt.xlabel('Epoch')\n        plt.ylabel('Accuracy')\n        plt.legend()\n        plt.grid(True, alpha=0.3)\n            \n        plt.tight_layout()\n        plt.savefig(f\"{plots_dir}/learning_curves/fold_{fold+1}_learning_curves.png\")\n        plt.close()\n            \n        # Create tolerance visualization for this fold\n        plt.figure(figsize=(10, 6))\n        tol_keys = list(epoch_val_tolerances.keys())\n        tol_values = [epoch_val_tolerances[k] for k in tol_keys]\n            \n        bars = plt.bar(tol_keys, tol_values, color='skyblue')\n        plt.title(f'Validation Tolerance Metrics - Fold {fold+1}')\n        plt.xlabel('Tolerance Level')\n        plt.ylabel('Percentage of Predictions (%)')\n        plt.ylim(0, 100)\n            \n        # Add value labels on top of bars\n        for bar in bars:\n            height = bar.get_height()\n            plt.text(bar.get_x() + bar.get_width()/2., height + 1,\n                    f'{height:.1f}%', ha='center', va='bottom')\n                \n        plt.grid(True, alpha=0.3, axis='y')\n        plt.savefig(f\"{plots_dir}/tolerance_metrics/fold_{fold+1}_tolerance_metrics.png\")\n        plt.close()\n\n        # Close fold writer\n        fold_writer.close()\n\n    # End of cross-validation\n    avg_val_acc = sum(all_val_accuracies) / len(all_val_accuracies) if all_val_accuracies else 0.0\n    avg_val_f1 = sum(all_val_f1_scores) / len(all_val_f1_scores) if all_val_f1_scores else 0.0\n    print(f\"\\n===== TRAINING COMPLETE =====\")\n    print(f\"Cross-validation results:\")\n    for fold, (acc, f1) in enumerate(zip(all_val_accuracies, all_val_f1_scores)):\n        print(f\"Fold {fold+1}: Accuracy = {acc:.4f}, F1 = {f1:.4f}\")\n    print(f\"Average validation Accuracy: {avg_val_acc:.4f}\")\n    print(f\"Average validation F1 Score: {avg_val_f1:.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\"Condition: {condition}\\n\")\n        f.write(f\"Tolerance thresholds: {tolerances}\\n\")\n        f.write(\"\\nResults:\\n\")\n        for fold, (acc, f1) in enumerate(zip(all_val_accuracies, all_val_f1_scores)):\n            f.write(f\"Fold {fold+1}: Accuracy = {acc:.4f}, F1 = {f1:.4f}\\n\")\n        f.write(f\"Average validation Accuracy: {avg_val_acc:.4f}\\n\")\n        f.write(f\"Average validation F1 Score: {avg_val_f1:.4f}\\n\")\n        \n        # Add tolerance metrics summary across all folds\n        f.write(\"\\nTolerance Metrics Summary:\\n\")\n        for tol in tolerances:\n            f.write(f\"±{tol}: Represents predictions within {tol} grades of ground truth\\n\")\n        f.write(f\">±{max(tolerances)}: Represents predictions more than {max(tolerances)} grades away from ground truth\\n\")\n\n    # Create visualizations for the entire experiment\n    \n    # 1. Average learning curves across all folds\n    plt.figure(figsize=(12, 5))\n        \n    # Find the minimum length of epochs (in case early stopping kicked in)\n    min_epochs_train = min([len(losses) for losses in all_fold_train_losses])\n    min_epochs_val = min([len(losses) for losses in all_fold_val_losses])\n        \n    # Prepare data for averaging\n    train_losses_truncated = [losses[:min_epochs_train] for losses in all_fold_train_losses]\n    val_losses_truncated = [losses[:min_epochs_val] for losses in all_fold_val_losses]\n    train_accs_truncated = [accs[:min_epochs_train] for accs in all_fold_train_accs]\n    val_accs_truncated = [accs[:min_epochs_val] for accs in all_fold_val_accs]\n        \n    # Calculate mean and std for each epoch\n    mean_train_loss = np.mean(train_losses_truncated, axis=0)\n    std_train_loss = np.std(train_losses_truncated, axis=0)\n    mean_val_loss = np.mean(val_losses_truncated, axis=0)\n    std_val_loss = np.std(val_losses_truncated, axis=0)\n        \n    mean_train_acc = np.mean(train_accs_truncated, axis=0)\n    std_train_acc = np.std(train_accs_truncated, axis=0)\n    mean_val_acc = np.mean(val_accs_truncated, axis=0)\n    std_val_acc = np.std(val_accs_truncated, axis=0)\n        \n    # Plot loss curves with shaded std dev\n    plt.subplot(1, 2, 1)\n    epochs_range = np.arange(1, min_epochs_train + 1)\n    plt.plot(epochs_range, mean_train_loss, label='Train Loss', color='blue')\n    plt.fill_between(epochs_range, mean_train_loss - std_train_loss, mean_train_loss + std_train_loss, \n                    alpha=0.2, color='blue')\n        \n    plt.plot(epochs_range, mean_val_loss, label='Validation Loss', color='red')\n    plt.fill_between(epochs_range, mean_val_loss - std_val_loss, mean_val_loss + std_val_loss, \n                    alpha=0.2, color='red')\n        \n    plt.title('Average Loss Across All Folds')\n    plt.xlabel('Epoch')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.grid(True, alpha=0.3)\n        \n    plt.tight_layout()\n    plt.savefig(f\"{plots_dir}/average_learning_curves.png\")\n    plt.close()\n        \n    # 2. Combined confusion matrix across all folds\n    if all_fold_confusion_matrices:\n        # Sum up all confusion matrices\n        combined_cm = sum([item['matrix'] for item in all_fold_confusion_matrices])\n        \n        plt.figure(figsize=(10, 8))\n        disp = ConfusionMatrixDisplay(confusion_matrix=combined_cm)\n        disp.plot(cmap=plt.cm.Blues, values_format='d')\n        plt.title('Combined Confusion Matrix Across All Folds')\n        plt.savefig(f\"{plots_dir}/combined_confusion_matrix.png\", bbox_inches='tight')\n        plt.close()\n        \n    # 3. Bar chart comparing fold performances\n    plt.figure(figsize=(10, 6))\n    fold_indices = np.arange(1, n_folds + 1)\n    \n    width = 0.35\n    plt.bar(fold_indices - width/2, all_val_accuracies, width, label='Accuracy', color='royalblue')\n    plt.bar(fold_indices + width/2, all_val_f1_scores, width, label='F1 Score', color='darkslateblue')\n    \n    plt.xlabel('Fold')\n    plt.ylabel('Score')\n    plt.title('Performance Metrics by Fold')\n    plt.xticks(fold_indices)\n    plt.ylim(0, 1.1)\n    plt.legend()\n    plt.grid(True, alpha=0.3, axis='y')\n    \n    # Add value labels\n    for i, (acc, f1) in enumerate(zip(all_val_accuracies, all_val_f1_scores)):\n        plt.text(i + 1 - width/2, acc + 0.02, f'{acc:.3f}', ha='center', va='bottom')\n        plt.text(i + 1 + width/2, f1 + 0.02, f'{f1:.3f}', ha='center', va='bottom')\n        \n        plt.tight_layout()\n        plt.savefig(f\"{plots_dir}/fold_performance_comparison.png\")\n        plt.close()\n        \n        # 4. Average tolerance metrics across folds\n        avg_val_tolerances = {}\n        for tol_key in all_fold_val_tolerances[0].keys():\n            avg_val_tolerances[tol_key] = np.mean([fold_tol[tol_key] for fold_tol in all_fold_val_tolerances])\n            \n        plt.figure(figsize=(10, 6))\n        tol_keys = list(avg_val_tolerances.keys())\n        tol_values = [avg_val_tolerances[k] for k in tol_keys]\n        \n        bars = plt.bar(tol_keys, tol_values, color='skyblue')\n    plt.title('Average Validation Tolerance Metrics Across All Folds')\n    plt.xlabel('Tolerance Level')\n    plt.ylabel('Percentage of Predictions (%)')\n    plt.ylim(0, 100)\n    \n    # Add value labels on top of bars\n    for bar in bars:\n        height = bar.get_height()\n        plt.text(bar.get_x() + bar.get_width()/2., height + 1,\n                    f'{height:.1f}%', ha='center', va='bottom')\n            \n    plt.grid(True, alpha=0.3, axis='y')\n    plt.savefig(f\"{plots_dir}/tolerance_metrics/average_tolerance_metrics.png\")\n    plt.close()\n    \n    # 5. Create radar chart for model performance overview\n    # Prepare data\n    metrics = ['Accuracy', 'F1 Score', '±0 Tolerance', '±1 Tolerance', '±2 Tolerance']\n    values = [\n        avg_val_acc,\n        avg_val_f1,\n        avg_val_tolerances['±0'] / 100,  # Convert percentage to 0-1 scale\n        avg_val_tolerances['±1'] / 100,\n        avg_val_tolerances['±2'] / 100\n    ]\n        \n    # Create radar chart\n    fig = plt.figure(figsize=(8, 8))\n    ax = fig.add_subplot(111, polar=True)\n    \n    # Set the angles for the metrics\n    angles = np.linspace(0, 2*np.pi, len(metrics), endpoint=False).tolist()\n    values.append(values[0])  # Close the loop\n    angles.append(angles[0])  # Close the loop\n    metrics.append(metrics[0])  # For labeling\n    \n    # Plot the values\n    ax.plot(angles, values, 'o-', linewidth=2, color='dodgerblue')\n    ax.fill(angles, values, color='dodgerblue', alpha=0.25)\n    \n    # Set the labels\n    ax.set_thetagrids(np.degrees(angles[:-1]), metrics[:-1])\n    \n    # Set y limits\n    ax.set_ylim(0, 1)\n    \n    # Add gridlines\n    ax.set_rgrids([0.2, 0.4, 0.6, 0.8, 1.0], angle=0)\n    ax.grid(True)\n    \n    plt.title('Model Performance Overview', size=15, y=1.1)\n    plt.tight_layout()\n    plt.savefig(f\"{plots_dir}/model_performance_radar.png\")\n    plt.close()\n    \n    print(f\"\\n✨ All visualizations saved to the '{plots_dir}' directory\")\n\n    # Cose writer\n    writer.close()\n    \n    \n    return all_val_accuracies, all_val_f1_scores, avg_val_acc, avg_val_f1\n# Example usage:\n# model = SSMIL()\n# accuracies, f1_scores, avg_acc, avg_val_acc, avg_val_f1 = train_spine_model(model, train_coor, train_meta, train_dummy, train_series, n_folds=5, condition='ss')","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-06T22:36:28.753Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}