{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":11125637,"sourceType":"datasetVersion","datasetId":6938318}],"dockerImageVersionId":30919,"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-03-23T10:30:48.884216Z","iopub.execute_input":"2025-03-23T10:30:48.884656Z","iopub.status.idle":"2025-03-23T10:30:53.270872Z","shell.execute_reply.started":"2025-03-23T10:30:48.884621Z","shell.execute_reply":"2025-03-23T10:30:53.269880Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport cv2\nimport glob\nimport torch\nimport wandb\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 torchvision.transforms.v2 as v2\nfrom torch.optim import AdamW, lr_scheduler\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import train_test_split","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T10:30:53.272233Z","iopub.execute_input":"2025-03-23T10:30:53.272577Z","iopub.status.idle":"2025-03-23T10:31:03.179464Z","shell.execute_reply.started":"2025-03-23T10:30:53.272545Z","shell.execute_reply":"2025-03-23T10:31:03.178854Z"}},"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-03-23T10:31:03.180852Z","iopub.execute_input":"2025-03-23T10:31:03.181283Z","iopub.status.idle":"2025-03-23T10:31:03.259075Z","shell.execute_reply.started":"2025-03-23T10:31:03.181261Z","shell.execute_reply":"2025-03-23T10:31:03.258445Z"}},"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-03-23T10:31:03.260211Z","iopub.execute_input":"2025-03-23T10:31:03.260505Z","iopub.status.idle":"2025-03-23T10:31:04.172107Z","shell.execute_reply.started":"2025-03-23T10:31:03.260476Z","shell.execute_reply":"2025-03-23T10:31:04.171440Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_dummy = train_dummy.fillna(\"Normal/Mild\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T10:31:04.172762Z","iopub.execute_input":"2025-03-23T10:31:04.172962Z","iopub.status.idle":"2025-03-23T10:31:04.181375Z","shell.execute_reply.started":"2025-03-23T10:31:04.172945Z","shell.execute_reply":"2025-03-23T10:31:04.180728Z"}},"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-03-23T10:31:04.182134Z","iopub.execute_input":"2025-03-23T10:31:04.182355Z","iopub.status.idle":"2025-03-23T10:31:04.215551Z","shell.execute_reply.started":"2025-03-23T10:31:04.182332Z","shell.execute_reply":"2025-03-23T10:31:04.214904Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Dataset","metadata":{}},{"cell_type":"code","source":"class SpineDataset(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\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        self.condition = condition\n        self.meta = meta\n        self.mode = mode\n\n\n    def __len__(self):\n        return len(self.id)\n\n    def __getitem__(self, idx):\n        study_id = self.id[idx]\n        \n        if self.condition == 'scs':\n            volume, label = self.volume_scs(study_id)\n            if volume is None or volume.shape[0] == 0:\n                print(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        depth = 32\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    \n        img = [self.load_dicom(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{study_id}/{row.series_id}/{row.instance_number}.dcm\") for _, row in meta.iterrows()]\n        \n        coor = self.coor.loc[(self.coor.study_id == study_id) & (self.coor.series_description == 'Sagittal T2/STIR')]\n        coor_dict = {}\n        \n        for _, row in coor.iterrows():\n            series_id, instance_number = row.series_id, row.instance_number\n            target_row = meta.loc[(meta.series_id == series_id) & (meta.instance_number == instance_number)]\n            if not target_row.empty:\n                idx = target_row.index[0]\n                z = idx/depth if len(img) < depth else idx/len(img)\n                x, y = row.x / img[idx].shape[1], row.y / img[idx].shape[0]\n                coor_dict[row.level] = torch.tensor([x, y, z], dtype=torch.float32)\n    \n        # Ensure all levels have the same shape (Default: [0, 0, 0])\n        coor_dict = {\n            'L1/L2': coor_dict.get('L1/L2', torch.zeros(3)),  \n            'L2/L3': coor_dict.get('L2/L3', torch.zeros(3)),\n            'L3/L4': coor_dict.get('L3/L4', torch.zeros(3)),\n            'L4/L5': coor_dict.get('L4/L5', torch.zeros(3)),\n            'L5/S1': coor_dict.get('L5/S1', torch.zeros(3))\n        }\n    \n        # Resize volume tensor\n        volume = torch.cat([self.resize(torch.tensor(i)[None, ...]).to(torch.float32) for i in img]).contiguous()\n        \n        if volume.shape[0] < depth:\n            volume = torch.cat([volume, torch.zeros(depth - volume.shape[0], volume.shape[1], volume.shape[2])])\n        else:\n            volume = torch.nn.functional.interpolate(volume[None, None, ...], (depth, volume.shape[1], volume.shape[2])).squeeze()\n    \n        return volume, coor_dict\n\n\n    def volume_nfn(self, study_id):\n        pass\n\n    def volume_ss(self, study_id):\n        pass\n\n    def normalize(self, x): # in real world data dicom has extreme pixel value that is why we use normalization\n        upper, lower = torch.quantile(x, torch.tensor([0.99, 0.01]))\n        x = torch.clip(x, lower, upper)  # Remove Extreme outliers\n\n        # x = (x - lower) / (upper - lower)\n\n        # x = x - torch.min(x)\n        # x = x / (torch.max(x)+1e-6)\n\n        x = (x - lower) / (upper - lower + 1e-6)  # Min-max normalization\n        return x\n\n\n    def load_dicom(self, path):\n        return pydicom.dcmread(path).pixel_array\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T10:31:04.216177Z","iopub.execute_input":"2025-03-23T10:31:04.216370Z","iopub.status.idle":"2025-03-23T10:31:04.230206Z","shell.execute_reply.started":"2025-03-23T10:31:04.216354Z","shell.execute_reply":"2025-03-23T10:31:04.229373Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{}},{"cell_type":"code","source":"class ConvNextStem(nn.Sequential):\n    def __init__(self, in_features: int, out_features: int):\n        super().__init__(\n            nn.Conv3d(in_features, out_features, kernel_size=(1, 2, 2), stride=(1, 2, 2)),\n            nn.GroupNorm(num_groups=1, num_channels=out_features)\n        )\n\n\nclass BottleNeckBlock(nn.Module):\n    def __init__(\n        self,\n        in_features: int,\n        out_features: int,\n        expansion: int = 4,\n        drop_p: float = .0,\n        layer_scaler_init_value: float = 1e-6,\n    ):\n        super().__init__()\n        expanded_features = out_features * expansion\n        self.block = nn.Sequential(\n            # narrow -> wide (with depth-wise and bigger kernel)\n            nn.Conv3d(\n                in_features, in_features, kernel_size=(2, 7, 7), padding='same', bias=False, groups=in_features\n            ),\n            # GroupNorm with num_groups=1 is the same as LayerNorm but works for 2D data\n            nn.GroupNorm(num_groups=in_features, num_channels=in_features),\n            # wide -> wide\n            nn.Conv3d(in_features, expanded_features, kernel_size=1),\n            nn.GELU(),\n            # wide -> narrow\n            nn.Conv3d(expanded_features, out_features, kernel_size=1),\n        )\n        #self.layer_scaler = LayerScaler(layer_scaler_init_value, out_features)\n        #self.drop_path = StochasticDepth(drop_p, mode=\"batch\")\n\n\n    def forward(self, x: Tensor) -> Tensor:\n        res = x\n        x = self.block(x)\n        #x = self.layer_scaler(x)\n        #x = self.drop_path(x)\n        x += res\n        return x\n\nclass ConvNexStage(nn.Sequential):\n    def __init__(\n        self, in_features: int, out_features: int, depth: int, **kwargs\n    ):\n        super().__init__(\n            # add the downsampler\n            nn.Sequential(\n                nn.GroupNorm(num_groups=in_features, num_channels=in_features),\n                nn.Conv3d(in_features, out_features, kernel_size=(2, 2, 2), stride=(2, 2, 2))\n            ),\n            *[\n                BottleNeckBlock(out_features, out_features, **kwargs)\n                for _ in range(depth)\n            ],\n        )\n\nclass ConvNextEncoder(nn.Module):\n    def __init__(\n        self,\n        in_channels: int,\n        stem_features: int,\n        depths: List[int],\n        widths: List[int],\n        drop_p: float = .0,\n    ):\n        super().__init__()\n        self.stem = ConvNextStem(in_channels, stem_features)\n\n        in_out_widths = list(zip(widths, widths[1:]))\n        # create drop paths probabilities (one for each stage)\n        drop_probs = [x.item() for x in torch.linspace(0, drop_p, sum(depths))]\n\n        self.stages = nn.ModuleList(\n            [\n                ConvNexStage(stem_features, widths[0], depths[0], drop_p=drop_probs[0]),\n                *[\n                    ConvNexStage(in_features, out_features, depth, drop_p=drop_p)\n                    for (in_features, out_features), depth, drop_p in zip(\n                        in_out_widths, depths[1:], drop_probs[1:]\n                    )\n                ],\n            ]\n        )\n\n\n    def forward(self, x):\n        x = self.stem(x)\n        for stage in self.stages:\n            x = stage(x)\n        return x\n\n\nclass ConvNextSCSDepthDetect(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.encoder = ConvNextEncoder(in_channels=1, stem_features=64, depths=[3,3,9,3], widths=[128, 256, 512, 1024])\n        self.flatten = nn.Sequential(nn.AdaptiveAvgPool3d((1,1,1)),\n                                    nn.Flatten(1),\n                                    nn.LayerNorm(1024),\n                                     )\n        self.l1 = nn.Linear(1024, 3)\n        self.l2 = nn.Linear(1024, 3)\n        self.l3 = nn.Linear(1024, 3)\n        self.l4 = nn.Linear(1024, 3)\n        self.l5 = nn.Linear(1024, 3)\n    def forward(self, x, label=None):\n        x = x.unsqueeze(1)\n        x = self.encoder(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, 'L2/L3': l2, 'L3/L4': l3, 'L4/L5': l4, 'L5/S1': l5}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T10:31:04.232108Z","iopub.execute_input":"2025-03-23T10:31:04.232305Z","iopub.status.idle":"2025-03-23T10:31:04.250571Z","shell.execute_reply.started":"2025-03-23T10:31:04.232288Z","shell.execute_reply":"2025-03-23T10:31:04.249955Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"model = ConvNextSCSDepthDetect().to(device)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T10:31:04.731244Z","iopub.execute_input":"2025-03-23T10:31:04.731559Z","iopub.status.idle":"2025-03-23T10:31:05.437848Z","shell.execute_reply.started":"2025-03-23T10:31:04.731534Z","shell.execute_reply":"2025-03-23T10:31:05.437141Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loss","metadata":{}},{"cell_type":"code","source":"class SCSLoss(nn.Module):\n    def __init__(self):\n        super(SCSLoss, self).__init__()\n\n    def forward(self, outputs, targets):\n        loss = 0\n        for level in ['L1/L2', 'L2/L3', 'L3/L4', 'L4/L5', 'L5/S1']:\n            # Debugging: Print shapes\n            # print(f\"🔹 {level} Output Shape: {outputs[level].shape}, Target Shape: {targets[level].shape}\")\n\n            # # Ensure both tensors have the same shape before L1 loss\n            # if outputs[level].shape != targets[level].shape:\n            #     targets[level] = targets[level].expand_as(outputs[level])  # Expand target if needed\n\n            _loss = nn.functional.l1_loss(outputs[level], targets[level])\n            loss += _loss\n        return loss / 5\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T10:31:05.607900Z","iopub.execute_input":"2025-03-23T10:31:05.608162Z","iopub.status.idle":"2025-03-23T10:31:05.612917Z","shell.execute_reply.started":"2025-03-23T10:31:05.608142Z","shell.execute_reply":"2025-03-23T10:31:05.612119Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Train without augmentation","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-03-23T10:31:07.932217Z","iopub.execute_input":"2025-03-23T10:31:07.932542Z","iopub.status.idle":"2025-03-23T10:31:08.029576Z","shell.execute_reply.started":"2025-03-23T10:31:07.932517Z","shell.execute_reply":"2025-03-23T10:31:08.028662Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# def calculate_accuracy(outputs, targets):\n#     \"\"\"Computes accuracy given model outputs and true labels.\"\"\"\n    \n#     total_correct = 0\n#     total_samples = 0\n    \n#     for level in outputs.keys():\n#         _, preds = torch.max(outputs[level], dim=1)  # Get predicted class indices\n        \n#         # Ensure targets have correct shape\n#         target_labels = targets[level]  # Check if this is one-hot\n#         if target_labels.ndim > 1 and target_labels.shape[1] > 1:\n#             target_labels = target_labels.argmax(dim=1)  # Convert one-hot to class index\n#         correct = (preds == target_labels).sum().item()\n#         total_correct += correct\n#         total_samples += target_labels.size(0)\n\n#     return total_correct, total_samples\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T10:31:10.773807Z","iopub.execute_input":"2025-03-23T10:31:10.774088Z","iopub.status.idle":"2025-03-23T10:31:10.777727Z","shell.execute_reply.started":"2025-03-23T10:31:10.774069Z","shell.execute_reply":"2025-03-23T10:31:10.776821Z"}},"outputs":[],"execution_count":null},{"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    for level in outputs.keys():\n        mse = nn.functional.mse_loss(outputs[level], targets[level], reduction='sum')  # Sum over batch\n        total_mse += mse.item()\n        total_samples += targets[level].numel()  # Total elements\n\n    return total_mse / total_samples  # Average MSE per element\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T10:31:11.163873Z","iopub.execute_input":"2025-03-23T10:31:11.164176Z","iopub.status.idle":"2025-03-23T10:31:11.168469Z","shell.execute_reply.started":"2025-03-23T10:31:11.164150Z","shell.execute_reply":"2025-03-23T10:31:11.167677Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"epochs = 10\n\nfor i in range(1):\n    train_dataset = SpineDataset(train_coor.loc[train_coor.fold!=i], train_meta.loc[train_meta.fold!=i], 'scs', 'train')\n    valid_dataset = SpineDataset(train_coor.loc[train_coor.fold==i], train_meta.loc[train_meta.fold==i], 'scs', 'valid')\n\n    train_loader = DataLoader(train_dataset, batch_size=2, shuffle=True, num_workers=4, pin_memory=False)\n    valid_loader = DataLoader(valid_dataset, batch_size=2, shuffle=False, num_workers=4, pin_memory=False)\n\n    model.to(device)  # Move model to GPU\n    optimizer = AdamW(model.parameters(), lr=0.001, weight_decay=0.0001)\n    scheduler = transformers.get_cosine_schedule_with_warmup(\n        optimizer=optimizer, num_warmup_steps=2 * len(train_loader), num_training_steps=epochs * len(train_loader), num_cycles=0.5\n    )\n    criterion = SCSLoss()\n\n    # train_losses, val_losses, train_accs, val_accs = [], [], [], []\n\n    train_losses, val_losses, train_mses, val_mses = [], [], [], []\n\n    for epoch in range(epochs):\n        print(f\"\\n🚀 Epoch {epoch+1}/{epochs}\")\n\n        # Training phase\n        model.train()\n        # running_loss, correct, total = 0.0, 0, 0\n        running_loss, total_mse = 0.0, 0.0\n        progress_bar = tqdm(train_loader, desc=\"Training Progress\", leave=False)\n\n        for volume, batch in progress_bar:\n            volume = volume.to(device)  # Move input tensors to GPU\n            # Debugging: Check available keys in batch\n            \n            # Ensure all expected keys are present\n            expected_levels = [\"L1/L2\", \"L2/L3\", \"L3/L4\", \"L4/L5\", \"L5/S1\"]\n            batch = {key: value.to(device) for key, value in batch.items() if key in expected_levels}\n\n            optimizer.zero_grad()\n            outputs = model(volume)  \n            loss = criterion(outputs, batch)\n            loss.backward()\n            optimizer.step()\n\n            scheduler.step()  # ✅ Moved outside batch loop\n\n            # running_loss += loss.item()\n            # batch_correct, batch_total = calculate_accuracy(outputs, batch)\n            # correct += batch_correct\n            # total += batch_total\n\n            running_loss += loss.item()\n            batch_mse = calculate_mse_score(outputs, batch)\n            total_mse += batch_mse\n\n            # ✅ Show loss in tqdm bar\n            progress_bar.set_postfix(loss=f\"{loss.item():.4f}\")\n\n            # ✅ Free memory\n            del volume, batch, outputs, loss\n            torch.cuda.empty_cache()\n\n\n        # epoch_train_loss = running_loss / len(train_loader)\n        # epoch_train_acc = correct / total\n        # train_losses.append(epoch_train_loss)\n        # train_accs.append(epoch_train_acc)\n        # print(f\"🔥 Training Loss: {epoch_train_loss:.4f} | Accuracy: {epoch_train_acc:.4%}\")\n\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        print(f\"🔥 Training Loss: {epoch_train_loss:.4f} | MSE Score: {epoch_train_mse:.4f}\")\n\n        # Validation phase\n        model.eval()\n        val_running_loss, val_total_mse = 0.0, 0.0\n        with torch.no_grad():\n            for volume, batch in tqdm(valid_loader, desc=\"Validation Progress\"):\n                volume = volume.to(device)\n                batch = {key: value.to(device) for key, value in batch.items()}\n\n                outputs = model(volume)\n                loss = criterion(outputs, batch)\n                val_running_loss += loss.item()\n\n                # batch_correct, batch_total = calculate_accuracy(outputs, batch)\n                # val_correct += batch_correct\n                # val_total += batch_total\n\n                batch_mse = calculate_mse_score(outputs, batch)\n                val_total_mse += batch_mse\n\n                del volume, batch, outputs, loss\n                torch.cuda.empty_cache()\n\n        # epoch_val_loss = val_running_loss / len(valid_loader)\n        # epoch_val_acc = val_correct / val_total\n        # val_losses.append(epoch_val_loss)\n        # val_accs.append(epoch_val_acc)\n        # print(f\"✅ Validation Loss: {epoch_val_loss:.4f} | Accuracy: {epoch_val_acc:.4%}\")\n\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        print(f\"✅ Validation Loss: {epoch_val_loss:.4f} | MSE Score: {epoch_val_mse:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-23T10:31:16.225390Z","iopub.execute_input":"2025-03-23T10:31:16.225746Z","execution_failed":"2025-03-23T11:48:27.135Z"}},"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\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-03-22T18:58:17.539994Z","iopub.execute_input":"2025-03-22T18:58:17.540327Z","iopub.status.idle":"2025-03-22T18:58:17.786366Z","shell.execute_reply.started":"2025-03-22T18:58:17.540299Z","shell.execute_reply":"2025-03-22T18:58:17.785454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}