{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","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":68898,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":57457,"modelId":79334}],"dockerImageVersionId":30732,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"########################################################################\n#                                                                      #\n#                               IMPORTS                                #\n#                                                                      #\n########################################################################\n\nimport random\nimport numpy as np\nimport pandas as pd\nfrom pathlib import Path\nimport matplotlib.pyplot as plt\n\nfrom sklearn.pipeline import Pipeline, FeatureUnion\nfrom sklearn.compose import ColumnTransformer\nfrom sklearn.preprocessing import OneHotEncoder, OrdinalEncoder, FunctionTransformer, LabelEncoder\nfrom sklearn.impute import SimpleImputer\nfrom sklearn.base import BaseEstimator, TransformerMixin\nfrom sklearn import set_config\nfrom sklearn.model_selection import train_test_split\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torchvision import transforms\nfrom sklearn.metrics import accuracy_score, precision_score, recall_score, f1_score\n\nimport cv2\nfrom transformers import ViTModel\n\nfrom torchvision import transforms\nimport torch\nfrom torch.utils.data import Dataset\nimport pydicom\nfrom skimage import exposure, measure, morphology\nfrom skimage.segmentation import watershed\nfrom scipy import ndimage as ndi\nimport numpy as np\n\nfrom transformers import ViTModel, ViTConfig\n\nimport os\nimport pandas as pd\nimport numpy as np\nimport pydicom\nimport torch\nfrom torch.utils.data import Dataset\nfrom torchvision.transforms import ToTensor\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom tqdm import tqdm\nfrom skimage import exposure\n\nset_config(transform_output=\"pandas\")\n\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nimport timm\n\n# !pip install einops\n# from einops import rearrange, repeat","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-08-27T09:27:02.900355Z","iopub.execute_input":"2024-08-27T09:27:02.901177Z","iopub.status.idle":"2024-08-27T09:27:13.317741Z","shell.execute_reply.started":"2024-08-27T09:27:02.901141Z","shell.execute_reply":"2024-08-27T09:27:13.316765Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"########################################################################\n#                                                                      #\n#                            TabularDataset                            #\n#                                                                      #\n########################################################################\n\nclass TabularDataset:\n    def __init__(self, working_directory):\n        self.working_directory = working_directory\n        self.train_directory = self.working_directory / 'train.csv'\n        self.train_label_coordinates_directory = self.working_directory / 'train_label_coordinates.csv'\n        self.train_series_descriptions_directory = self.working_directory / 'train_series_descriptions.csv'\n\n        self.output_columns = [\n            'study_id', 'series_id', 'instance_number',\n            'x', 'y', 'series_description', 'side', 'condition',\n            'level', 'severity', 'condition_level'\n        ]\n              \n    def load_data(self):\n        self.train_df = pd.read_csv(self.train_directory)\n        self.train_label_coordinates_df = pd.read_csv(self.train_label_coordinates_directory)\n        self.train_series_descriptions_df = pd.read_csv(self.train_series_descriptions_directory)\n        \n    def get_data(self):    \n        return self.data\n    \n    def process_data(self):\n        self._create_severity_df()\n        self._create_condition_level()\n        self._merge_data()\n        self._create_side()\n        self._remove_side_condition()\n        self._replace_wrong_coordinates()\n        self._select_output_columns()\n        \n    def _create_severity_df(self):\n        severity_data = []\n        for index, row in self.train_df.iterrows():\n            study_id = row['study_id']\n            for column in self.train_df.columns[1:]:\n                condition_level = column\n                severity = row[column]\n                severity_data.append([study_id, condition_level, severity])\n        \n        self.severity_df = pd.DataFrame(severity_data, columns=['study_id', 'condition_level', 'severity'])\n    \n    def _create_condition_level(self):\n        self.train_label_coordinates_df['condition_level'] = (\n            self.train_label_coordinates_df['condition'].str.lower().str.replace(' ', '_') + '_' + \n            self.train_label_coordinates_df['level'].str.lower().str.replace('/', '_')\n        )\n        \n    def _merge_data(self):\n        self.data = (\n            self.severity_df\n            .merge(self.train_label_coordinates_df, on=['study_id', 'condition_level'])\n            .merge(self.train_series_descriptions_df, on=['study_id', 'series_id'])\n        )\n\n    def _create_side(self):\n        self.data['side'] = self.data['condition'].apply(lambda x: x.split()[0] if x.split()[0] in ('Right', 'Left') else 'Center')\n    \n    def _remove_side_condition(self):\n        self.data['condition'] = self.data['condition'].str.replace('Right', '').str.replace('Left', '').str.strip()\n        \n    def _replace_wrong_coordinates(self, target_cols = ['x', 'y'], fill_value_col='level', threshold=20):\n        for target_col in target_cols:\n            mask = self.data[target_col] < threshold\n            self.data.loc[mask, target_col] = np.nan\n\n            fill_value_dict = (\n                self.data\n                .loc[~mask]\n                .groupby(fill_value_col)[[target_col]]\n                .mean()\n                .to_dict()[target_col]\n            )\n\n            def fill_value(row):\n                return fill_value_dict[row[fill_value_col]] if row[fill_value_col] in fill_value_dict else np.nan\n\n            self.data.loc[mask, target_col] = self.data[mask].apply(fill_value, axis=1)\n        \n    def _select_output_columns(self):\n        self.data = self.data[self.output_columns]#.query(\"series_description != 'Axial T2'\")\n    \nworking_directory = Path('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification')\ndataset = TabularDataset(working_directory)\ndataset.load_data()\ndataset.process_data()\ndata = dataset.get_data()\ndisplay(data)\n\n########################################################################\n#                                                                      #\n#                       TabularDatasetTransform                        #\n#                                                                      #\n########################################################################\n\nclass TabularDatasetTransform:\n    def __init__(self, data):\n        self.data = data.copy()\n        self.data_original = data.copy()\n        self.categorical_features = ['series_description', 'side', 'condition', 'level', 'severity']\n\n        self.output_columns = [\n            'study_id', 'series_id', 'instance_number',\n            'x', 'y', 'series_description', 'side', 'condition',\n            'level', 'severity', 'condition_level'\n        ]\n        \n        self._create_pipeline()\n      \n    def _create_pipeline(self):\n        self.categorical_transformer = Pipeline(steps=[\n            ('impute_mode', SimpleImputer(strategy='most_frequent')),\n            ('ordinal_encode', OrdinalEncoder())\n        ])\n\n        self.preprocessor = ColumnTransformer( \n            transformers=[ \n                ('categorical', self.categorical_transformer, self.categorical_features),\n            ],\n            remainder='passthrough',\n            verbose_feature_names_out=True\n        )\n        \n        self.pipeline = Pipeline(steps=[\n            ('preprocessor', self.preprocessor)\n        ])\n        \n    def fit(self):\n        self.pipeline.fit(self.data)\n        self.transformed_data = self.pipeline.transform(self.data)\n        self.categories_ = self.pipeline.named_steps['preprocessor'].named_transformers_['categorical'].named_steps['ordinal_encode'].categories_\n        \n    def transform(self):\n        self.data = (\n            self.data\n            .assign(\n                series_description=self.transformed_data.filter(like='categorical__series_description').values,\n                side=self.transformed_data.filter(like='categorical__side').values,\n                condition=self.transformed_data.filter(like='categorical__condition').values,\n                level=self.transformed_data.filter(like='categorical__level').values,\n                severity=self.transformed_data.filter(like='categorical__severity').values\n            )[self.output_columns]\n        )\n        return self.data\n\n    def fit_transform(self):\n        self.fit()\n        return self.transform()\n\n    def reverse_ordinal_mapping(self, value, feature):\n        feature_index = self.categorical_features.index(feature)\n        return self.categories_[feature_index][int(value)]\n\n    def inverse_transform(self, X):\n        data = X.copy()\n        for feature in self.categorical_features:\n            if feature in data.columns:\n                data[feature] = data[feature].apply(lambda x: self.reverse_ordinal_mapping(x, feature))\n        return data\n    \n    def create_pivot_table(self):\n#         data_pivoted = self.data_original.query(\"series_description == 'Axial T2'\").pivot_table(\n#             index=['study_id', 'series_id', 'instance_number', 'side', 'series_description', 'condition'],\n#             columns='level',\n#             values=['severity','x', 'y'],\n#             aggfunc=lambda x: x\n#         )\n\n#         col_names = {0.0: 'l1_l2', 1.0: 'l2_l3', 2.0: 'l3_l4', 3.0: 'l4_l5', 4.0: 'l5_s1'}\n#         data_pivoted.rename(columns=col_names, inplace=True)\n        \n#         data_pivoted.columns = ['{}_{}'.format(col[0], col[1]) if col[1] else col[0] for col in data_pivoted.columns]\n#         data_pivoted = data_pivoted.dropna()\n\n        targets = self.data.pivot_table(\n            index=['study_id'],\n            columns = 'condition_level',\n            values='severity',\n            aggfunc=lambda x: x\n        ).reset_index()\n    \n        # inputs\n        inputs = self.data.pivot_table(\n            index=['study_id'],\n            columns=['series_description'],\n            values=['instance_number', 'series_id'],\n            aggfunc= {\n                'instance_number': lambda x: sorted(set(x)),\n                'series_id': lambda x: list(set(x))#str(x.iloc[0])\n\n            }\n        ).reset_index()\n\n        inputs.columns = ['{}_{}'.format(col[0], col[1]) if col[1] else col[0] for col in inputs.columns]\n        \n        col_names = {\n            'series_id': 'series_id_x', 'series_id_1.0': 'series_id_y', 'series_id_2.0': 'series_id_z',\n            'instance_number': 'instance_number_x', 'instance_number_1.0': 'instance_number_y', 'instance_number_2.0': 'instance_number_z'\n        }\n        \n        return inputs.merge(targets, on='study_id').rename(columns=col_names).dropna()\n    \ntransformer = TabularDatasetTransform(data)\ndata_processed = transformer.fit_transform()\ndisplay(data_processed)\n\ndata_pivoted = transformer.create_pivot_table()  # => series_id_max = 10(x) + 10(y) + 5(z)\ndisplay(data_pivoted)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:27:13.319788Z","iopub.execute_input":"2024-08-27T09:27:13.320068Z","iopub.status.idle":"2024-08-27T09:27:16.134455Z","shell.execute_reply.started":"2024-08-27T09:27:13.320044Z","shell.execute_reply":"2024-08-27T09:27:16.133448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data.series_description.value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:27:16.135730Z","iopub.execute_input":"2024-08-27T09:27:16.136037Z","iopub.status.idle":"2024-08-27T09:27:16.149746Z","shell.execute_reply.started":"2024-08-27T09:27:16.136010Z","shell.execute_reply":"2024-08-27T09:27:16.148874Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dF = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv\")\ndF","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:27:16.152149Z","iopub.execute_input":"2024-08-27T09:27:16.152814Z","iopub.status.idle":"2024-08-27T09:27:16.171211Z","shell.execute_reply.started":"2024-08-27T09:27:16.152788Z","shell.execute_reply":"2024-08-27T09:27:16.170359Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_processed.condition_level.value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:27:16.172267Z","iopub.execute_input":"2024-08-27T09:27:16.172552Z","iopub.status.idle":"2024-08-27T09:27:16.185972Z","shell.execute_reply.started":"2024-08-27T09:27:16.172528Z","shell.execute_reply":"2024-08-27T09:27:16.184961Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"########################################################################\n#                                                                      #\n#                             UnetDataset                              #\n#                                                                      #\n########################################################################\n\nclass UnetDataset(Dataset):\n    def __init__(self, data, image_dir, transform=None):\n        self.data = data\n        self.image_dir = image_dir\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        row = self.data.iloc[idx]\n        image_path = os.path.join(self.image_dir, f\"{row['study_id']}/{row['series_id']}/{row['instance_number']}.dcm\")\n        image, metadata = self.load_dicom_image(image_path)\n        mask, (x_min, x_max, y_min, y_max) = self.create_roi_mask(image.shape, row['x'], row['y'], metadata)\n        \n        filtered_data = self.data[(self.data['study_id'] == row['study_id']) & (self.data['series_id'] == row['series_id'])]\n        mask = np.zeros(image.shape)\n\n        for _, roi_row in filtered_data.iterrows():\n            roi_mask, _ = self.create_roi_mask(image.shape, roi_row['x'], roi_row['y'], metadata)\n            mask += roi_mask\n        \n        if self.transform:\n            augmented = self.transform(image=image, mask=mask)\n            image = augmented['image'].float()\n            mask = augmented['mask'].unsqueeze(0).float()\n            \n        return image, torch.clamp(mask, 0, 1)\n        \n    def load_dicom_image(self, image_path):\n        dicom = pydicom.dcmread(image_path)\n        image = dicom.pixel_array\n        metadata = {elem.keyword: elem.value for elem in dicom if elem.VR != 'SQ' and elem.keyword and elem.keyword not in 'PixelData'}\n        \n#         window_level, window_width = metadata['WindowCenter'], metadata['WindowWidth']\n#         lowest_visible_value = window_level - window_width / 2 \n#         highest_visible_value = window_level + window_width / 2\n#         image = image.clip(lowest_visible_value, highest_visible_value)\n\n        min_pixel_value, max_pixel_value = image.min(), image.max()\n        image = (image - min_pixel_value) / (max_pixel_value - min_pixel_value + 1e-10)\n        \n#         image_eq = exposure.equalize_adapthist(image) \n        \n        return image, metadata\n    \n    def create_roi_mask(self, shape, x, y, metadata):\n        x, y = int(x), int(y)\n        \n        rows = metadata['Rows']\n        columns = metadata['Columns']\n        min_dim = min(rows, columns)\n        radius = int(0.05 * min_dim)\n\n        y_min = max(0, y - radius)\n        y_max = min(shape[0], y + radius)\n        x_min = max(0, x - radius)\n        x_max = min(shape[1], x + radius)\n        mask = np.zeros(shape)\n\n        y_grid, x_grid = np.ogrid[y_min:y_max, x_min:x_max]\n        dist_from_center = (x_grid - x)**2 + (y_grid - y)**2\n        mask[y_min:y_max, x_min:x_max] = dist_from_center <= radius**2 \n        \n        return mask, (x_min, x_max, y_min, y_max)\n    \nimage_dir = Path('/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images')\nimage_transforms = A.Compose([\n    A.Resize(height=224, width=224, p=1.0),\n    ToTensorV2(),\n])\n\ndata_unet = UnetDataset(data_processed, image_dir, image_transforms)\ndataloader_unet = DataLoader(data_unet, batch_size=32, shuffle=True)\n\n########################################################################\n#                                                                      #\n#                                U-NET                                 #\n#                                                                      #\n########################################################################\n\ndf = data.query(\"series_description == 'Sagittal T1' or series_description == 'Sagittal T2/STIR'\")#data_processed\ntrain_df, val_df = train_test_split(df, test_size=0.2, shuffle=False)\ntrain_df = train_df.reset_index(drop=True)\nval_df = val_df.reset_index(drop=True)\n\ntrain_df, test_df = train_test_split(train_df, test_size=0.15, shuffle=False)\ntrain_df = train_df.reset_index(drop=True)\n\ntrain_dataset = UnetDataset(df, image_dir, transform=image_transforms)\ntrain_dataloader = DataLoader(train_dataset, batch_size=32, shuffle=True)\n\nval_dataset = UnetDataset(val_df, image_dir, transform=image_transforms)\nval_dataloader = DataLoader(val_dataset, batch_size=32, shuffle=False)\n\ntest_dataset = UnetDataset(test_df, image_dir, transform=image_transforms)\ntest_dataloader = DataLoader(test_dataset, batch_size=32, shuffle=False)\n\ndef double_conv(in_channels, out_channels):\n    return nn.Sequential(\n        nn.Conv2d(in_channels, out_channels, 3, padding=1),\n        nn.ReLU(inplace=True),\n        nn.Conv2d(out_channels, out_channels, 3, padding=1),\n        nn.ReLU(inplace=True))\n\nclass UNet(nn.Module):\n\n    def __init__(self):\n        super().__init__()\n        self.conv_down1 = double_conv(1, 64)\n        self.conv_down2 = double_conv(64, 128)\n        self.conv_down3 = double_conv(128, 256)\n        self.conv_down4 = double_conv(256, 512)\n\n        self.maxpool = nn.MaxPool2d(2)\n        self.upsample = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)\n\n        self.conv_up3 = double_conv(256 + 512, 256)\n        self.conv_up2 = double_conv(128 + 256, 128)\n        self.conv_up1 = double_conv(128 + 64, 64)\n\n        self.last_conv = nn.Conv2d(64, 1, kernel_size=1)\n\n    def forward(self, x):\n        conv1 = self.conv_down1(x)\n        x = self.maxpool(conv1)\n        conv2 = self.conv_down2(x)\n        x = self.maxpool(conv2)\n        conv3 = self.conv_down3(x)\n        x = self.maxpool(conv3)\n        x = self.conv_down4(x)\n\n        x = self.upsample(x)\n        x = torch.cat([x, conv3], dim=1)\n        x = self.conv_up3(x)\n        x = self.upsample(x)\n        x = torch.cat([x, conv2], dim=1)\n        x = self.conv_up2(x)\n        x = self.upsample(x)\n        x = torch.cat([x, conv1], dim=1)\n        x = self.conv_up1(x)\n\n        out = self.last_conv(x)\n        out = torch.sigmoid(out)\n\n        return out\n    \ndevice = torch.device('cuda:0' if torch.cuda.is_available() else 'cpu')    \nunet = UNet().to(device)\noutput = unet(torch.randn(1,1,224,224).to(device))\nprint(\"\",output.shape)\n\ndef dice_coef_metric(inputs, target):\n    intersection = 2.0 * (target * inputs).sum()\n    union = target.sum() + inputs.sum()\n    if target.sum() == 0 and inputs.sum() == 0:\n        return 1.0\n\n    return intersection / union\n\ndef dice_coef_loss(inputs, target):\n    smooth = 1.0\n    intersection = 2.0 * ((target * inputs).sum()) + smooth\n    union = target.sum() + inputs.sum() + smooth\n\n    return 1 - (intersection / union)\n\ndef bce_dice_loss(inputs, target):\n    inputs = inputs.float()\n    target = target.float()\n\n    dicescore = dice_coef_loss(inputs, target)\n    bcescore = nn.BCELoss()\n    bceloss = bcescore(inputs, target)\n\n    return bceloss + dicescore\n\ndef train_model(model_name, model, train_loader, val_loader, train_loss, optimizer, lr_scheduler, num_epochs):\n\n    print(model_name)\n    loss_history = []\n    train_history = []\n    val_history = []\n\n    for epoch in range(num_epochs):\n        model.train()\n\n        losses = []\n        train_iou = []\n\n        if lr_scheduler:\n            warmup_factor = 1.0 / 100\n            warmup_iters = min(100, len(train_loader) - 1)\n            lr_scheduler = warmup_lr_scheduler(optimizer, warmup_iters, warmup_factor)\n\n        with tqdm(train_loader, desc=f\"Training epoch {epoch+1}/{num_epochs}\") as pbar:\n            for i_step, (data, target) in enumerate(pbar):\n                data = data.to(device)\n                target = target.to(device)\n\n                outputs = model(data)\n\n                out_cut = np.copy(outputs.data.cpu().numpy())\n\n                out_cut[np.nonzero(out_cut < 0.5)] = 0.0\n                out_cut[np.nonzero(out_cut >= 0.5)] = 1.0\n\n                train_dice = dice_coef_metric(out_cut, target.data.cpu().numpy())\n\n                loss = train_loss(outputs, target)\n\n                losses.append(loss.item())\n                train_iou.append(train_dice)\n\n                optimizer.zero_grad()\n                loss.backward()\n                optimizer.step()\n\n                if lr_scheduler:\n                    lr_scheduler.step()\n\n                pbar.set_postfix(loss=loss.item())\n\n        val_mean_iou = 0.0 #compute_iou(model, val_loader)\n\n        loss_history.append(np.array(losses).mean())\n        train_history.append(np.array(train_iou).mean())\n        val_history.append(val_mean_iou)\n\n        print(\"Epoch [%d]\" % (epoch))\n        print(\"Mean loss on train:\", np.array(losses).mean(),\n              \"\\nMean DICE on train:\", np.array(train_iou).mean(),\n              \"\\nMean DICE on validation:\", val_mean_iou)\n\n    return loss_history, train_history, val_history\n\ndef compute_iou(model, loader, threshold=0.3):\n    #model.eval()\n    valloss = 0\n\n    with torch.no_grad():\n\n        for i_step, (data, target) in enumerate(loader):\n\n            data = data.to(device)\n            target = target.to(device)\n            #prediction = model(x_gpu)\n\n            outputs = model(data)\n            out_cut = np.copy(outputs.data.cpu().numpy())\n            out_cut[np.nonzero(out_cut < threshold)] = 0.0\n            out_cut[np.nonzero(out_cut >= threshold)] = 1.0\n\n            picloss = dice_coef_metric(out_cut, target.data.cpu().numpy())\n            valloss += picloss\n\n        #print(\"Threshold:  \" + str(threshold) + \"  Validation DICE score:\", valloss / i_step)\n\n    return valloss / i_step\n\nunet_optimizer = torch.optim.Adamax(unet.parameters(), lr=1e-3)\n\n# lr_scheduler\ndef warmup_lr_scheduler(optimizer, warmup_iters, warmup_factor):\n    def f(x):\n        if x >= warmup_iters:\n            return 1\n        alpha = float(x) / warmup_iters\n        return warmup_factor * (1 - alpha) + alpha\n\n    return torch.optim.lr_scheduler.LambdaLR(optimizer, f)\n\n# ########################################################################\n# #                                                                      #\n# #                           U-NET TRAIN/LOAD                           #\n# #                                                                      #\n# ########################################################################\n\nif True:\n    model_path = os.getcwd()\n    model_path = os.path.join(model_path, 'unet_sagittal.pth')\n\n    num_ep = 3\n\n    unet_lh, unet_th, unet_vh = train_model(\"Vanila_UNet\", unet, train_dataloader, val_dataloader, bce_dice_loss, unet_optimizer, False, num_epochs = num_ep)\n    \n    torch.save(unet.state_dict(), model_path)\n\n#     unet = UNet().to(device)\n#     unet.load_state_dict(torch.load(\"unet.pth\",map_location=device))\n# else:\n#     unet_path = Path(\"/kaggle/input/lumbar-spine-degenerative-segregation/pytorch/u-net/1/unet.pth\")\n#     unet = UNet().to(device)\n#     unet.load_state_dict(torch.load(unet_path, map_location=device))","metadata":{"execution":{"iopub.status.busy":"2024-08-27T10:29:35.014676Z","iopub.execute_input":"2024-08-27T10:29:35.015093Z","iopub.status.idle":"2024-08-27T11:33:57.459829Z","shell.execute_reply.started":"2024-08-27T10:29:35.015058Z","shell.execute_reply":"2024-08-27T11:33:57.458854Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df = data.query(\"series_description == 'Sagittal T1' or series_description == 'Sagittal T2/STIR'\")\ndf\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T10:29:14.707291Z","iopub.execute_input":"2024-08-27T10:29:14.707966Z","iopub.status.idle":"2024-08-27T10:29:14.736203Z","shell.execute_reply.started":"2024-08-27T10:29:14.707933Z","shell.execute_reply":"2024-08-27T10:29:14.735337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport pydicom\nimport numpy as np\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom torch.utils.data import Dataset, DataLoader\nfrom pathlib import Path\nimport matplotlib.pyplot as plt\n\n# Definir las transformaciones que se usaron durante el entrenamiento\nimage_transforms = A.Compose([\n    A.Resize(height=224, width=224, p=1.0),\n    ToTensorV2(),\n])\n\n# Definir la función para cargar la imagen DICOM\ndef load_dicom_image(image_path):\n    dicom = pydicom.dcmread(image_path)\n    image = dicom.pixel_array\n    metadata = {elem.keyword: elem.value for elem in dicom if elem.VR != 'SQ' and elem.keyword and elem.keyword not in 'PixelData'}\n    \n    # Normalizar la imagen\n    min_pixel_value, max_pixel_value = image.min(), image.max()\n    image = (image - min_pixel_value) / (max_pixel_value - min_pixel_value + 1e-10)\n    \n    return image, metadata\n\n# Definir la función para predecir la máscara de una imagen DICOM\ndef predict_image(image_path, model, transform, device):\n    # Cargar y preprocesar la imagen\n    image, metadata = load_dicom_image(image_path)\n    \n    if transform:\n        augmented = transform(image=image)\n        image = augmented['image'].unsqueeze(0).to(device)\n    \n    # Convertir la imagen a tipo float32\n    image = image.float()\n    \n    # Predecir la máscara\n    model.eval()\n    with torch.no_grad():\n        output = model(image)\n        output = torch.sigmoid(output)\n        output = output.squeeze().cpu().numpy()\n    \n    return output\n\n# Path de la imagen DICOM\ndicom_image_path = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/44036939/3844393089/10.dcm\"\n\n# Cargar el modelo entrenado\nunet = UNet().to(device)\nmodel_path = '/kaggle/working/unet_sagittal.pth'\nunet.load_state_dict(torch.load(model_path, map_location=device))\n\n# Predecir la máscara de la imagen\npredicted_mask = predict_image(dicom_image_path, unet, image_transforms, device)\n\n# Mostrar la imagen y la máscara predicha\noriginal_image, _ = load_dicom_image(dicom_image_path)\n\nplt.figure(figsize=(10, 5))\nplt.subplot(1, 2, 1)\nplt.imshow(original_image, cmap='gray')\nplt.title(\"Original Image\")\n\nplt.subplot(1, 2, 2)\nplt.imshow(predicted_mask, cmap='gray')\nplt.title(\"Predicted Mask\")\n\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T11:37:34.432410Z","iopub.execute_input":"2024-08-27T11:37:34.432771Z","iopub.status.idle":"2024-08-27T11:37:35.193691Z","shell.execute_reply.started":"2024-08-27T11:37:34.432742Z","shell.execute_reply":"2024-08-27T11:37:35.192703Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# data.series_description","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.048179Z","iopub.status.idle":"2024-08-27T09:36:56.048691Z","shell.execute_reply.started":"2024-08-27T09:36:56.048431Z","shell.execute_reply":"2024-08-27T09:36:56.048450Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# sagital t2/stir = series_description 1\n# sagital t1 = series_description 0\n# axial = series_description 2\n\ndata.query(\"series_description == 'Axial T2'\") # Axial","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.050262Z","iopub.status.idle":"2024-08-27T09:36:56.050947Z","shell.execute_reply.started":"2024-08-27T09:36:56.050705Z","shell.execute_reply":"2024-08-27T09:36:56.050725Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"########################################################################\n#                                                                      #\n#                           U-NET DEPLOYMENT                           #\n#                                                                      #\n########################################################################\n\ndef generate_random_indices(dataset, num_samples):\n    indices = list(range(len(dataset)))\n    random.shuffle(indices)\n    return indices[:num_samples]\n\nnum_samples = 3\nsample_indices = generate_random_indices(test_dataset, num_samples)\n\nfig, axes = plt.subplots(num_samples, 2, figsize=(10, 10))\n\nfor i, idx in enumerate(sample_indices):\n    image, mask = test_dataset[idx]\n    mask = mask[0, :, :]\n    prediction = unet(image.unsqueeze(0).to(device))\n    prediction = prediction[0, 0, :, :].data.cpu().numpy()\n\n    axes[i, 0].imshow(image.permute(1, 2, 0).cpu().numpy(), cmap='gray')\n    axes[i, 0].imshow(mask, cmap='magma', alpha=0.5)\n    axes[i, 0].set_title(f\"Sample {idx}: Ground Truth\")\n    \n    axes[i, 1].imshow(image.permute(1, 2, 0).cpu().numpy(), cmap='gray')\n    axes[i, 1].imshow(prediction, cmap='turbo', alpha=0.5)\n    axes[i, 1].set_title(f\"Sample {idx}: Prediction\")\n    \n    axes[i, 0].axis(\"off\")\n    axes[i, 1].axis(\"off\")\n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.052319Z","iopub.status.idle":"2024-08-27T09:36:56.052776Z","shell.execute_reply.started":"2024-08-27T09:36:56.052533Z","shell.execute_reply":"2024-08-27T09:36:56.052553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"########################################################################\n#                                                                      #\n#                              ViTDataset                              #\n#                                                                      #\n########################################################################\n\n# image_transforms = transforms.Compose([\n#     transforms.Resize((224, 224), antialias=True),\n#     transforms.ConvertImageDtype(torch.float32)\n# ])\n\nimage_transforms = transforms.Compose([\n    transforms.Resize(255, antialias=True),\n    transforms.CenterCrop(224),\n    transforms.RandomRotation(degrees=15),\n    transforms.RandomResizedCrop(224, scale=(0.8, 1.0), antialias=True),\n    transforms.RandomAffine(degrees=0, translate=(0.1, 0.1)),\n    transforms.ConvertImageDtype(torch.float32)\n])\n\nclass ViTDataset(Dataset):\n    def __init__(self, data, image_dir, transform=image_transforms):\n        self.data = data.copy()\n        self.image_dir = Path(image_dir)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        row = self.data.iloc[idx]\n        study_id = row['study_id']\n\n        instance_numbers = [\n            row['instance_number_x'],\n            #row['instance_number_y'],\n            #row['instance_number_z']\n        ]\n        series_ids = [\n            row['series_id_x'],\n            #row['series_id_y'],\n            #row['series_id_z']\n        ]\n\n        volumes = [self.get_volume(instances, study_id, series_id) for instances, series_id in zip(instance_numbers, series_ids)]\n        \n        stack = self.prepare_stack(volumes)\n        labels = self.get_labels(row)\n\n        return {\n            \"stack\":stack,\n            \"labels\":labels\n        }\n    \n    @staticmethod\n    def to_tensor(values):\n        return torch.tensor(values, dtype=torch.float32)\n    \n    def get_volume(self, instance_numbers, study_id, series_id):      \n        image_paths = [\n            self.image_dir / f\"{study_id}/{serie_id}/{instance_number}.dcm\"\n            for instance_number in instance_numbers\n            for serie_id in series_id\n        ]\n        image_paths = [path for path in image_paths if path.exists()]\n        \n        images = [self.load_and_transform_image(image_path) for image_path in image_paths]\n\n        return torch.stack(images)\n\n    def load_and_transform_image(self, image_path):\n        dicom = pydicom.dcmread(image_path)\n        image = dicom.pixel_array\n        image = self.apply_windowing(image, dicom)\n        image = self.normalize_image(image)\n        image = self.to_tensor(image)\n        \n        if self.transform:\n            image = self.transform(image.unsqueeze(0)).squeeze(0)\n\n        return image\n\n    def apply_windowing(self, image, dicom):\n        window_level, window_width = dicom.WindowCenter, dicom.WindowWidth\n        lowest_visible_value = np.minimum(window_level - window_width / 2, 0)\n        highest_visible_value = np.minimum(window_level + window_width / 2, 1000)\n        image = image.clip(lowest_visible_value, highest_visible_value)\n        return image\n\n    def normalize_image(self, image):\n        min_pixel_value, max_pixel_value = image.min(), image.max()\n        return (image - min_pixel_value) / (max_pixel_value - min_pixel_value + 1e-10)\n\n    def prepare_stack(self, volumes):\n        stack = torch.cat(volumes, dim=0)\n        n_slides = stack.shape[0]\n        if n_slides < 15:\n            padding = torch.zeros((15 - n_slides, 224, 224), dtype=torch.float32)\n            stack = torch.cat((stack, padding), dim=0)\n        else:\n            stack = stack[:15]\n        return stack\n    \n    def get_labels(self, row):\n        conditions = row[\n            [\n                'left_neural_foraminal_narrowing_l1_l2',\n                'left_neural_foraminal_narrowing_l2_l3',\n                'left_neural_foraminal_narrowing_l3_l4',\n                'left_neural_foraminal_narrowing_l4_l5',\n                'left_neural_foraminal_narrowing_l5_s1',\n#                 'left_subarticular_stenosis_l1_l2', \n#                 'left_subarticular_stenosis_l2_l3',\n#                 'left_subarticular_stenosis_l3_l4',\n#                 'left_subarticular_stenosis_l4_l5',\n#                 'left_subarticular_stenosis_l5_s1',\n                'right_neural_foraminal_narrowing_l1_l2',\n                'right_neural_foraminal_narrowing_l2_l3',\n                'right_neural_foraminal_narrowing_l3_l4',\n                'right_neural_foraminal_narrowing_l4_l5',\n                'right_neural_foraminal_narrowing_l5_s1',\n#                 'right_subarticular_stenosis_l1_l2',\n#                 'right_subarticular_stenosis_l2_l3',\n#                 'right_subarticular_stenosis_l3_l4',\n#                 'right_subarticular_stenosis_l4_l5',\n#                 'right_subarticular_stenosis_l5_s1',\n                'spinal_canal_stenosis_l1_l2',\n                'spinal_canal_stenosis_l2_l3',\n                'spinal_canal_stenosis_l3_l4',\n                'spinal_canal_stenosis_l4_l5',\n                'spinal_canal_stenosis_l5_s1'\n            ]\n        ].values\n        return torch.tensor(np.array(conditions, dtype=np.float32), dtype=torch.float32) # replace float32 to long format\n\n# dataset = (\n#     data_pivoted.copy()\n#     .assign(\n#         samples=(\n#             data_pivoted.filter(regex='l1_l2$|l2_l3$|l3_l4$|l4_l5$|l5_s1$')\n#             .apply(lambda row: 'True' if any(row != 1.0) else 'False', axis=1)\n#         )\n#     )\n#     .pipe(lambda df: pd.concat([\n#         df.query(\"samples == 'True'\"),\n#         df.query(\"samples == 'False'\")\n#     ], axis=0))\n#     .reset_index(drop=True)\n#     .sample(frac=1, random_state=111)\n#     .drop(columns=['samples'])\n# )\n\ndata_vit = ViTDataset(data_pivoted, image_dir)\ndataloader_vit = DataLoader(data_vit, batch_size=4, shuffle=True)\n\nfor batch in dataloader_vit:\n    print(f'Inputs: {batch[\"stack\"].shape}')\n    print(f'Labels: {batch[\"labels\"].shape}')\n    break","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.054702Z","iopub.status.idle":"2024-08-27T09:36:56.055126Z","shell.execute_reply.started":"2024-08-27T09:36:56.054908Z","shell.execute_reply":"2024-08-27T09:36:56.054925Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\nfrom torchvision import transforms\nimport timm\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nfrom pathlib import Path\nimport pandas as pd\nimport pydicom\n\nclass ClassifierHead(nn.Module):\n    def __init__(self, hidden_size=768, num_targets=3):\n        super().__init__()\n        self.classifier_head = nn.Sequential(\n            nn.Linear(hidden_size, num_targets)\n        )\n    \n    def forward(self, x):\n        return self.classifier_head(x)\n\nclass ViTClassifier(nn.Module):\n    def __init__(self, num_classes=3, num_targets=15, vit_model='vit_base_patch16_224'):\n        super().__init__()\n        self.num_classes = num_classes\n        self.num_targets = num_targets\n        self.vit = timm.create_model(vit_model, pretrained=True)\n        \n        # Modify the first layer to accept num_targets channels\n        self.vit.patch_embed.proj = nn.Conv2d(num_targets, self.vit.patch_embed.proj.out_channels, \n                                              kernel_size=self.vit.patch_embed.proj.kernel_size, \n                                              stride=self.vit.patch_embed.proj.stride, \n                                              padding=self.vit.patch_embed.proj.padding)\n        \n        self.hidden_size = self.vit.head.in_features\n        self.vit.head = nn.Identity()\n\n        self.classifiers = nn.ModuleDict({\n            f\"{name}_{level}\": ClassifierHead(self.hidden_size, num_classes)\n            for name in [\n                'left_neural_foraminal_narrowing',# 'left_subarticular_stenosis', \n                'right_neural_foraminal_narrowing',# 'right_subarticular_stenosis', \n                'spinal_canal_stenosis']\n            for level in ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\n        })\n\n        self.combine_layer = nn.Sequential(\n            nn.Linear(len(self.classifiers) * num_classes, 128),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Linear(128, len(self.classifiers) * num_classes)\n        )\n\n    def forward(self, x):\n        features = self.vit(x)\n        intermediate_outputs = {key: classifier(features) for key, classifier in self.classifiers.items()}\n        concatenated_outputs = torch.cat([intermediate_outputs[key] for key in sorted(intermediate_outputs.keys())], dim=1)\n        combined_outputs = self.combine_layer(concatenated_outputs)\n        final_outputs = {key: combined_outputs[:, i*self.num_classes:(i+1)*self.num_classes] \n                         for i, key in enumerate(sorted(intermediate_outputs.keys()))}\n        return final_outputs\n\ndef criterion(loss_func, outputs, labels):\n    losses = 0.0\n    for idx, key in enumerate(outputs):\n        losses += loss_func(outputs[key], labels[:, idx].to(outputs[key].device).view_as(outputs[key]))\n    return losses\n\nnum_classes = 3\nnum_targets = 15\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = ViTClassifier(num_classes=num_classes, num_targets=num_targets, vit_model='vit_base_patch16_224').to(device)\n\nclass_counts = torch.tensor([7142, 34796, 2762], dtype=torch.float32, device=device)\nclass_weights = 1.0 / (class_counts)\nclass_weights /= class_weights.sum()\nloss_func = nn.BCEWithLogitsLoss(pos_weight=class_weights)  # Ajuste de la función de pérdida\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(optimizer, 'min', patience=2, verbose=True)\n\nnum_epochs = 10\nbest_loss = float('inf')\npatience = 5\ncounter = 0\ntrain_loss_history = []\ntrain_acc_history = []\n\n# Training loop\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    running_corrects = 0\n    total_samples = 0\n    \n    progress_bar = tqdm(dataloader_vit, desc=f'Epoch {epoch+1}/{num_epochs}')\n    \n    for i, batch in enumerate(progress_bar):\n        inputs, labels = batch[\"stack\"].to(device), batch[\"labels\"].to(device)\n        optimizer.zero_grad()\n        \n        outputs = model(inputs)\n        \n        loss = criterion(loss_func, outputs, labels.float())  # Ajuste de la función de pérdida\n        \n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item() * inputs.size(0)\n        \n        batch_corrects = 0\n        for idx, key in enumerate(outputs):\n            preds = (torch.sigmoid(outputs[key]) > 0.5).float()\n            batch_corrects += torch.sum(preds == labels[:, idx].data.view_as(preds))\n        \n        running_corrects += batch_corrects.item()\n        total_samples += labels.size(0) * labels.size(1) \n        \n        progress_bar.set_postfix(loss=running_loss / total_samples, accuracy=running_corrects / total_samples)\n    \n    epoch_loss = running_loss / total_samples\n    epoch_acc = running_corrects / total_samples\n    \n    train_loss_history.append(epoch_loss)\n    train_acc_history.append(epoch_acc)\n    \n    if epoch_loss < best_loss:\n        best_loss = epoch_loss\n        torch.save(model.state_dict(), 'checkpoint.pth')\n        counter = 0\n    else:\n        counter += 1\n        if counter >= patience:\n            print('Early stopping')\n            break\n    \n    scheduler.step(epoch_loss)\n\nwindow_size = 3\nwindow = np.ones(window_size) / window_size\n\nrolling_loss = np.convolve(train_loss_history, window, mode='valid')\nrolling_acc = np.convolve(train_acc_history, window, mode='valid')\n\nrolling_epochs = range(1, len(rolling_loss) + 1)\nepochs = range(1, len(train_loss_history) + 1)\n\nfig, ax = plt.subplots()\nplt.title('Training loss and accuracy')\n\ncolor = 'tab:red'\nax.set_xlabel('Epochs')\nax.set_ylabel('Loss', color=color)\nax.tick_params(axis='y', labelcolor=color)\nax.plot(epochs, train_loss_history, color=color, linestyle='--', alpha=0.5, label='Training loss (original)')\nax.plot(rolling_epochs, rolling_loss, color=color, linestyle='-', label='Training loss (rolling mean)')\n\nax_ = ax.twinx() \ncolor = 'tab:blue'\nax_.set_ylabel('Accuracy', color=color)\nax_.tick_params(axis='y', labelcolor=color)\nax_.plot(epochs, train_acc_history, color=color, linestyle='--', alpha=0.5, label='Training accuracy (original)')\nax_.plot(rolling_epochs, rolling_acc, color=color, linestyle='-', label='Training accuracy (rolling mean)')\n\nfig.tight_layout()\nfig.legend(loc='upper right', bbox_to_anchor=(1,1), bbox_transform=ax.transAxes)\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.056404Z","iopub.status.idle":"2024-08-27T09:36:56.056837Z","shell.execute_reply.started":"2024-08-27T09:36:56.056615Z","shell.execute_reply":"2024-08-27T09:36:56.056633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stop","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.057936Z","iopub.status.idle":"2024-08-27T09:36:56.058253Z","shell.execute_reply.started":"2024-08-27T09:36:56.058096Z","shell.execute_reply":"2024-08-27T09:36:56.058109Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ClassifierHead(nn.Module):\n    def __init__(self, hidden_size=768, num_targets=3):\n        super().__init__()\n#         self.classifier_head = nn.Sequential(\n#             nn.Linear(hidden_size, 512),\n#             nn.BatchNorm1d(512),\n#             nn.ReLU(),\n#             nn.Dropout(0.5),\n#             nn.Linear(512, 128),\n#             nn.BatchNorm1d(128),\n#             nn.ReLU(),\n#             nn.Dropout(0.5),\n#             nn.Linear(128, num_targets),\n#         )\n        self.classifier_head = nn.Sequential(\n            nn.Linear(in_features=hidden_size, out_features=num_targets, bias=True)\n        )\n    \n    def forward(self, x):\n        return self.classifier_head(x)\n\nclass ViTClassifier(nn.Module):\n    def __init__(self, num_classes=3, num_targets=15, vit_model='vit_base_patch16_224'):\n        super().__init__()\n        self.num_classes = num_classes\n        self.num_targets = num_targets\n        self.vit = timm.create_model(vit_model, pretrained=True, in_chans=num_targets)\n        self.hidden_size = self.vit.head.in_features\n        self.vit.head = nn.Identity()\n\n        self.classifiers = nn.ModuleDict({\n            f\"{name}_{level}\": ClassifierHead(self.hidden_size, num_classes)\n            for name in [\n                'left_neural_foraminal_narrowing',# 'left_subarticular_stenosis', \n                'right_neural_foraminal_narrowing',# 'right_subarticular_stenosis', \n                'spinal_canal_stenosis']\n            for level in ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\n        })\n\n        self.combine_layer = nn.Sequential(\n            nn.Linear(len(self.classifiers) * num_classes, 128),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Linear(128, len(self.classifiers) * num_classes)\n        )\n\n    def forward(self, x):\n        features = self.vit(x)\n        intermediate_outputs = {key: classifier(features) for key, classifier in self.classifiers.items()}\n        concatenated_outputs = torch.cat([intermediate_outputs[key] for key in sorted(intermediate_outputs.keys())], dim=1)\n        combined_outputs = self.combine_layer(concatenated_outputs)\n        final_outputs = {key: combined_outputs[:, i*self.num_classes:(i+1)*self.num_classes] \n                         for i, key in enumerate(sorted(intermediate_outputs.keys()))}\n        return final_outputs\n\ndef criterion(loss_func, outputs, labels):\n    losses = 0.0\n    for idx, key in enumerate(outputs):\n        losses += loss_func(outputs[key], labels[:, idx].to(outputs[key].device))\n\n    correlation_penalty = 0.0\n    correlation_matrix = {key: [other_key for other_key in outputs.keys() if other_key != key] for key in outputs.keys()}\n    \n    for key, related_keys in correlation_matrix.items():\n        for related_key in related_keys:\n            correlation_penalty += torch.mean((outputs[key] - outputs[related_key])**2)\n    \n    total_loss = losses + 0.1 * correlation_penalty\n    return total_loss\n\nnum_classes = 3\nnum_targets = 15\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel = ViTClassifier(num_classes=num_classes, num_targets=num_targets, vit_model='vit_base_patch16_224').to(device)\n\nclass_counts = torch.tensor([7142, 34796, 2762], dtype=torch.float32, device=device)\nclass_weights = 1.0 / (class_counts)\nclass_weights /= class_weights.sum()\nloss_func = nn.CrossEntropyLoss(weight=class_weights)  \noptimizer = torch.optim.Adam(model.parameters(), lr=1e-3)\nscheduler = ReduceLROnPlateau(optimizer, 'min', patience=2, verbose=True)\n\nnum_epochs = 10\nbest_loss = float('inf')\npatience = 5\ncounter = 0\ntrain_loss_history = []\ntrain_acc_history = []\n\n# Training loop\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    running_corrects = 0\n    total_samples = 0\n    \n    progress_bar = tqdm(dataloader_vit, desc=f'Epoch {epoch+1}/{num_epochs}')\n    \n    for i, batch in enumerate(progress_bar):\n        inputs, labels = batch[\"stack\"].to(device), batch[\"labels\"].to(device)\n        optimizer.zero_grad()\n        \n        outputs = model(inputs)\n        \n        loss = criterion(loss_func, outputs, labels.long())\n        \n        loss.backward()\n        optimizer.step()\n        \n        running_loss += loss.item() * inputs.size(0)\n        \n        batch_corrects = 0\n        for idx, key in enumerate(outputs):\n            _, preds = torch.max(outputs[key], 1)\n            batch_corrects += torch.sum(preds == labels[:, idx].data)\n        \n        running_corrects += batch_corrects.item()\n        total_samples += labels.size(0) * labels.size(1) \n        \n        progress_bar.set_postfix(loss=running_loss / total_samples, accuracy=running_corrects / total_samples)\n    \n    epoch_loss = running_loss / total_samples\n    epoch_acc = running_corrects / total_samples\n    \n    train_loss_history.append(epoch_loss)\n    train_acc_history.append(epoch_acc)\n    \n    if epoch_loss < best_loss:\n        best_loss = epoch_loss\n        torch.save(model.state_dict(), 'checkpoint.pth')\n        counter = 0\n    else:\n        counter += 1\n        if counter >= patience:\n            print('Early stopping')\n            break\n    \n    scheduler.step(epoch_loss)\n\nwindow_size = 3\nwindow = np.ones(window_size) / window_size\n\nrolling_loss = np.convolve(train_loss_history, window, mode='valid')\nrolling_acc = np.convolve(train_acc_history, window, mode='valid')\n\nrolling_epochs = range(1, len(rolling_loss) + 1)\nepochs = range(1, len(train_loss_history) + 1)\n\nfig, ax = plt.subplots()\nplt.title('Training loss and accuracy')\n\ncolor = 'tab:red'\nax.set_xlabel('Epochs')\nax.set_ylabel('Loss', color=color)\nax.tick_params(axis='y', labelcolor=color)\nax.plot(epochs, train_loss_history, color=color, linestyle='--', alpha=0.5, label='Training loss (original)')\nax.plot(rolling_epochs, rolling_loss, color=color, linestyle='-', label='Training loss (rolling mean)')\n\nax_ = ax.twinx() \ncolor = 'tab:blue'\nax_.set_ylabel('Accuracy', color=color)\nax_.tick_params(axis='y', labelcolor=color)\nax_.plot(epochs, train_acc_history, color=color, linestyle='--', alpha=0.5, label='Training accuracy (original)')\nax_.plot(rolling_epochs, rolling_acc, color=color, linestyle='-', label='Training accuracy (rolling mean)')\n\nfig.tight_layout()\nfig.legend(loc='upper right', bbox_to_anchor=(1,1), bbox_transform=ax.transAxes)\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.059847Z","iopub.status.idle":"2024-08-27T09:36:56.060201Z","shell.execute_reply.started":"2024-08-27T09:36:56.060013Z","shell.execute_reply":"2024-08-27T09:36:56.060026Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stop","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.061174Z","iopub.status.idle":"2024-08-27T09:36:56.061470Z","shell.execute_reply.started":"2024-08-27T09:36:56.061321Z","shell.execute_reply":"2024-08-27T09:36:56.061333Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def infer_example(model, example, device, num_classes=3):\n    x, y = example['stack'], example['labels']\n    y_gt = y\n    model = model.to(device)\n    model.eval()\n\n    with torch.no_grad():\n        y_out = model(x.unsqueeze(0).to(device))\n    \n    # Convertir las salidas del modelo a predicciones\n    y_pred = {key: torch.softmax(output, dim=-1).cpu().numpy() for key, output in y_out.items()}\n    \n    return y_pred, y_gt.numpy()\n\ninference, real = infer_example(model, data_vit[333], device)\nprint(\"Predicciones:\")\nfor key, value in inference.items():\n    print(f\"{key}: {value}\")\n\nprint(\"\\nValores Reales:\")\nprint(real)\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.062834Z","iopub.status.idle":"2024-08-27T09:36:56.063158Z","shell.execute_reply.started":"2024-08-27T09:36:56.062997Z","shell.execute_reply":"2024-08-27T09:36:56.063011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_classes = 3\nnum_targets = 25\nmodel = ViTClassifier(num_classes=num_classes, num_targets=num_targets, vit_model='vit_base_patch16_224').to(device)\n\ninput_images = torch.rand(32, 25, 224, 224).to(device)\noutputs = model(input_images)\nprint([out.shape for out in outputs.values()])","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.064875Z","iopub.status.idle":"2024-08-27T09:36:56.065307Z","shell.execute_reply.started":"2024-08-27T09:36:56.065080Z","shell.execute_reply":"2024-08-27T09:36:56.065099Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stop","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.066604Z","iopub.status.idle":"2024-08-27T09:36:56.067045Z","shell.execute_reply.started":"2024-08-27T09:36:56.066824Z","shell.execute_reply":"2024-08-27T09:36:56.066842Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---","metadata":{}},{"cell_type":"code","source":"stop","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.068407Z","iopub.status.idle":"2024-08-27T09:36:56.068861Z","shell.execute_reply.started":"2024-08-27T09:36:56.068620Z","shell.execute_reply":"2024-08-27T09:36:56.068638Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ClassifierHead(nn.Module):\n    def __init__(self, hidden_size=768, num_targets=3):\n        super().__init__()\n        self.hidden_size = hidden_size\n        self.num_targets = num_targets\n        self.classifier_head = nn.Sequential(\n            nn.Linear(self.hidden_size, 512),\n            nn.BatchNorm1d(512),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(512, 128),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(128, self.num_targets),\n        )\n    \n    def forward(self, x):\n        return self.classifier_head(x)\n\nclass ViTClassifier(nn.Module):\n    def __init__(self, num_classes=3, num_targets=25, vit_model='vit_base_patch16_224'):\n        super().__init__()\n        self.num_classes = num_classes\n        self.num_targets = num_targets\n        self.vit = timm.create_model(vit_model, pretrained=True, in_chans=25)\n        self.hidden_size = self.vit.head.in_features\n        self.vit.head = nn.Identity()\n        \n        self.left_neural_foraminal_narrowing_l1_l2_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.left_neural_foraminal_narrowing_l2_l3_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.left_neural_foraminal_narrowing_l3_l4_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.left_neural_foraminal_narrowing_l4_l5_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.left_neural_foraminal_narrowing_l5_s1_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.left_subarticular_stenosis_l1_l2_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.left_subarticular_stenosis_l2_l3_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.left_subarticular_stenosis_l3_l4_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.left_subarticular_stenosis_l4_l5_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.left_subarticular_stenosis_l5_s1_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.right_neural_foraminal_narrowing_l1_l2_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.right_neural_foraminal_narrowing_l2_l3_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.right_neural_foraminal_narrowing_l3_l4_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.right_neural_foraminal_narrowing_l4_l5_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.right_neural_foraminal_narrowing_l5_s1_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.right_subarticular_stenosis_l1_l2_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.right_subarticular_stenosis_l2_l3_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.right_subarticular_stenosis_l3_l4_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.right_subarticular_stenosis_l4_l5_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.right_subarticular_stenosis_l5_s1_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.spinal_canal_stenosis_l1_l2_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.spinal_canal_stenosis_l2_l3_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.spinal_canal_stenosis_l3_l4_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.spinal_canal_stenosis_l4_l5_classifier = ClassifierHead(self.hidden_size, num_classes)\n        self.spinal_canal_stenosis_l5_s1_classifier = ClassifierHead(self.hidden_size, num_classes)\n        \n    def forward(self, x):\n        features = self.vit(x)[:,0,:]\n        print(features.shape)\n\n        left_neural_foraminal_narrowing_l1_l2_output = self.left_neural_foraminal_narrowing_l1_l2_classifier(features)\n        left_neural_foraminal_narrowing_l2_l3_output = self.left_neural_foraminal_narrowing_l2_l3_classifier(features)\n        left_neural_foraminal_narrowing_l3_l4_output = self.left_neural_foraminal_narrowing_l3_l4_classifier(features)\n        left_neural_foraminal_narrowing_l4_l5_output = self.left_neural_foraminal_narrowing_l4_l5_classifier(features)\n        left_neural_foraminal_narrowing_l5_s1_output = self.left_neural_foraminal_narrowing_l5_s1_classifier(features)\n        left_subarticular_stenosis_l1_l2_output = self.left_subarticular_stenosis_l1_l2_classifier(features)\n        left_subarticular_stenosis_l2_l3_output = self.left_subarticular_stenosis_l2_l3_classifier(features)\n        left_subarticular_stenosis_l3_l4_output = self.left_subarticular_stenosis_l3_l4_classifier(features)\n        left_subarticular_stenosis_l4_l5_output = self.left_subarticular_stenosis_l4_l5_classifier(features)\n        left_subarticular_stenosis_l5_s1_output = self.left_subarticular_stenosis_l5_s1_classifier(features)\n        right_neural_foraminal_narrowing_l1_l2_output = self.right_neural_foraminal_narrowing_l1_l2_classifier(features)\n        right_neural_foraminal_narrowing_l2_l3_output = self.right_neural_foraminal_narrowing_l2_l3_classifier(features)\n        right_neural_foraminal_narrowing_l3_l4_output = self.right_neural_foraminal_narrowing_l3_l4_classifier(features)\n        right_neural_foraminal_narrowing_l4_l5_output = self.right_neural_foraminal_narrowing_l4_l5_classifier(features)\n        right_neural_foraminal_narrowing_l5_s1_output = self.right_neural_foraminal_narrowing_l5_s1_classifier(features)\n        right_subarticular_stenosis_l1_l2_output = self.right_subarticular_stenosis_l1_l2_classifier(features)\n        right_subarticular_stenosis_l2_l3_output = self.right_subarticular_stenosis_l2_l3_classifier(features)\n        right_subarticular_stenosis_l3_l4_output = self.right_subarticular_stenosis_l3_l4_classifier(features)\n        right_subarticular_stenosis_l4_l5_output = self.right_subarticular_stenosis_l4_l5_classifier(features)\n        right_subarticular_stenosis_l5_s1_output = self.right_subarticular_stenosis_l5_s1_classifier(features)\n        spinal_canal_stenosis_l1_l2_output = self.spinal_canal_stenosis_l1_l2_classifier(features)\n        spinal_canal_stenosis_l2_l3_output = self.spinal_canal_stenosis_l2_l3_classifier(features)\n        spinal_canal_stenosis_l3_l4_output = self.spinal_canal_stenosis_l3_l4_classifier(features)\n        spinal_canal_stenosis_l4_l5_output = self.spinal_canal_stenosis_l4_l5_classifier(features)\n        spinal_canal_stenosis_l5_s1_output = self.spinal_canal_stenosis_l5_s1_classifier(features)\n\n    def forward(self, x):\n        features = self.vit(x)\n\n        return {\n            \"left_neural_foraminal_narrowing_l1_l2\": self.left_neural_foraminal_narrowing_l1_l2_classifier(features),\n            \"left_neural_foraminal_narrowing_l2_l3\": self.left_neural_foraminal_narrowing_l2_l3_classifier(features),\n            \"left_neural_foraminal_narrowing_l3_l4\": self.left_neural_foraminal_narrowing_l3_l4_classifier(features),\n            \"left_neural_foraminal_narrowing_l4_l5\": self.left_neural_foraminal_narrowing_l4_l5_classifier(features),\n            \"left_neural_foraminal_narrowing_l5_s1\": self.left_neural_foraminal_narrowing_l5_s1_classifier(features),\n            \"left_subarticular_stenosis_l1_l2\": self.left_subarticular_stenosis_l1_l2_classifier(features),\n            \"left_subarticular_stenosis_l2_l3\": self.left_subarticular_stenosis_l2_l3_classifier(features),\n            \"left_subarticular_stenosis_l3_l4\": self.left_subarticular_stenosis_l3_l4_classifier(features),\n            \"left_subarticular_stenosis_l4_l5\": self.left_subarticular_stenosis_l4_l5_classifier(features),\n            \"left_subarticular_stenosis_l5_s1\": self.left_subarticular_stenosis_l5_s1_classifier(features),\n            \"right_neural_foraminal_narrowing_l1_l2\": self.right_neural_foraminal_narrowing_l1_l2_classifier(features),\n            \"right_neural_foraminal_narrowing_l2_l3\": self.right_neural_foraminal_narrowing_l2_l3_classifier(features),\n            \"right_neural_foraminal_narrowing_l3_l4\": self.right_neural_foraminal_narrowing_l3_l4_classifier(features),\n            \"right_neural_foraminal_narrowing_l4_l5\": self.right_neural_foraminal_narrowing_l4_l5_classifier(features),\n            \"right_neural_foraminal_narrowing_l5_s1\": self.right_neural_foraminal_narrowing_l5_s1_classifier(features),\n            \"right_subarticular_stenosis_l1_l2\": self.right_subarticular_stenosis_l1_l2_classifier(features),\n            \"right_subarticular_stenosis_l2_l3\": self.right_subarticular_stenosis_l2_l3_classifier(features),\n            \"right_subarticular_stenosis_l3_l4\": self.right_subarticular_stenosis_l3_l4_classifier(features),\n            \"right_subarticular_stenosis_l4_l5\": self.right_subarticular_stenosis_l4_l5_classifier(features),\n            \"right_subarticular_stenosis_l5_s1\": self.right_subarticular_stenosis_l5_s1_classifier(features),\n            \"spinal_canal_stenosis_l1_l2\": self.spinal_canal_stenosis_l1_l2_classifier(features),\n            \"spinal_canal_stenosis_l2_l3\": self.spinal_canal_stenosis_l2_l3_classifier(features),\n            \"spinal_canal_stenosis_l3_l4\": self.spinal_canal_stenosis_l3_l4_classifier(features),\n            \"spinal_canal_stenosis_l4_l5\": self.spinal_canal_stenosis_l4_l5_classifier(features),\n            \"spinal_canal_stenosis_l5_s1\": self.spinal_canal_stenosis_l5_s1_classifier(features)\n        }\n\nnum_classes = 3\nnum_targets = 25\nmodel = ViTClassifier(num_classes=num_classes, num_targets=num_targets, vit_model='vit_base_patch16_224').to(device)\n\ninput_images = torch.rand(32, 25, 224, 224).to(device)\noutputs = model(input_images)\nprint([out.shape for out in outputs.values()])","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.071198Z","iopub.status.idle":"2024-08-27T09:36:56.071654Z","shell.execute_reply.started":"2024-08-27T09:36:56.071402Z","shell.execute_reply":"2024-08-27T09:36:56.071421Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stop","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.072628Z","iopub.status.idle":"2024-08-27T09:36:56.073065Z","shell.execute_reply.started":"2024-08-27T09:36:56.072842Z","shell.execute_reply":"2024-08-27T09:36:56.072860Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Model\nnum_classes = 3\nnum_targets = 25\nmodel = ViTClassifier(num_classes=num_classes, num_targets=num_targets, vit_model='vit_base_patch16_224').to(device)\n\n# Loss functions and optimizer\nloss_func = nn.CrossEntropyLoss(weight=torch.tensor([4.872, 1, 12.598], dtype=torch.float32, device=device)) # {1.0: 34796 0.0: 7142 2.0: 2762}\noptimizer = torch.optim.Adam(model.parameters(), lr=1e-4)\nscheduler = ReduceLROnPlateau(optimizer, 'min', patience=3, verbose=True)\n\n# Variables for training\nnum_epochs = 2\nbest_loss = float('inf')\npatience = 5\ncounter = 0\ntrain_loss_history = []\ntrain_acc_history = []\n\n# Criterion function\ndef criterion(loss_func, outputs, labels):\n    losses = 0.0\n    for idx, key in enumerate(outputs):\n        losses += loss_func(outputs[key], labels[:, idx].to(outputs[key].device))\n    return losses\n\n# Training loop\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    running_corrects = 0\n    total_samples = 0\n    \n    progress_bar = tqdm(dataloader_vit, desc=f'Epoch {epoch+1}/{num_epochs}')\n    \n    for i, batch in enumerate(progress_bar):\n        inputs, labels = batch[\"stack\"].to(device), batch[\"labels\"].to(device)\n        optimizer.zero_grad()\n        \n        outputs = model(inputs)\n        \n        loss = criterion(loss_func, outputs, labels.long())\n        \n        loss.backward()\n        optimizer.step()\n        \n        # Update running loss and corrects\n        running_loss += loss.item() * inputs.size(0)\n        _, preds = torch.max(outputs['spinal_canal_stenosis_l5_s1'], 1)\n        running_corrects += torch.sum(preds == labels[:, -1].data)\n        total_samples += inputs.size(0)\n        \n        # Update progress bar\n        progress_bar.set_postfix(loss=running_loss / total_samples, accuracy=running_corrects.double() / total_samples)\n    \n    # Calculate epoch loss and accuracy\n    epoch_loss = running_loss / total_samples\n    epoch_acc = running_corrects.double() / total_samples\n    \n    # Update histories\n    train_loss_history.append(epoch_loss)\n    train_acc_history.append(epoch_acc.item())\n    \n    # Checkpoint and early stopping\n    if epoch_loss < best_loss:\n        best_loss = epoch_loss\n        torch.save(model.state_dict(), 'checkpoint.pth')\n        counter = 0\n    else:\n        counter += 1\n        if counter >= patience:\n            print('Early stopping')\n            break\n    \n    # Step the scheduler\n    scheduler.step(epoch_loss)\n\n# Plotting loss and accuracy\nwindow_size = 3\nwindow = np.ones(window_size) / window_size\n\nrolling_loss = np.convolve(train_loss_history, window, mode='valid')\nrolling_acc = np.convolve(train_acc_history, window, mode='valid')\n\nrolling_epochs = range(1, len(rolling_loss) + 1)\nepochs = range(1, len(train_loss_history) + 1)\n\nfig, ax = plt.subplots()\nplt.title('Training loss and accuracy')\n\ncolor = 'tab:red'\nax.set_xlabel('Epochs')\nax.set_ylabel('Loss', color=color)\nax.tick_params(axis='y', labelcolor=color)\nax.plot(epochs, train_loss_history, color=color, linestyle='--', alpha=0.5, label='Training loss (original)')\nax.plot(rolling_epochs, rolling_loss, color=color, linestyle='-', label='Training loss (rolling mean)')\n\nax_ = ax.twinx() \ncolor = 'tab:blue'\nax_.set_ylabel('Accuracy', color=color)\nax_.tick_params(axis='y', labelcolor=color)\nax_.plot(epochs, train_acc_history, color=color, linestyle='--', alpha=0.5, label='Training accuracy (original)')\nax_.plot(rolling_epochs, rolling_acc, color=color, linestyle='-', label='Training accuracy (rolling mean)')\n\nfig.tight_layout()\nfig.legend(loc='upper right', bbox_to_anchor=(1,1), bbox_transform=ax.transAxes)\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.075041Z","iopub.status.idle":"2024-08-27T09:36:56.075352Z","shell.execute_reply.started":"2024-08-27T09:36:56.075200Z","shell.execute_reply":"2024-08-27T09:36:56.075213Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"aa = data_pivoted.loc[:, \"left_neural_foraminal_narrowing_l1_l2\":\"spinal_canal_stenosis_l5_s1\"]\n\nbb = pd.DataFrame([aa[col].value_counts() for col in aa.columns])\n\nbb.sum()","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.076252Z","iopub.status.idle":"2024-08-27T09:36:56.076557Z","shell.execute_reply.started":"2024-08-27T09:36:56.076399Z","shell.execute_reply":"2024-08-27T09:36:56.076412Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\n\ndef infer(model, dataset, index, device='cpu'):\n    model.to(device).eval()\n    with torch.no_grad():\n        data = dataset[index]\n        inputs = data['stack'].to(device)\n        inputs = inputs.unsqueeze(0)\n        outputs = model(inputs)\n        \n        predicted = {key: torch.max(output, 1)[1].cpu().numpy() for key, output in outputs.items()}\n        softmax_scores = {key: F.softmax(output, dim=1).cpu().numpy() for key, output in outputs.items()}\n        \n    return predicted, softmax_scores, data['labels'].numpy()\n\nindex = 300\npredictions, softmax_scores, actual_labels = infer(model, data_vit, index)\n\npredicted_tensor = torch.tensor([predictions[key][0] for key in predictions.keys()])\nactual_tensor = torch.tensor(actual_labels)\n\nprint(\"Predicciones:\")\nprint(predicted_tensor.view(5, 5))\nprint(\"Etiquetas reales:\")\nprint(actual_tensor.view(5, 5))\n\nprint(\"Softmax Scores:\")\nfor key in softmax_scores:\n    print(f\"Clase {key}: {softmax_scores[key]}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.078012Z","iopub.status.idle":"2024-08-27T09:36:56.078432Z","shell.execute_reply.started":"2024-08-27T09:36:56.078211Z","shell.execute_reply":"2024-08-27T09:36:56.078228Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"torch.save(model.state_dict(), \"/kaggle/working/vit.pth\")\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.079694Z","iopub.status.idle":"2024-08-27T09:36:56.080112Z","shell.execute_reply.started":"2024-08-27T09:36:56.079892Z","shell.execute_reply":"2024-08-27T09:36:56.079909Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_vit[0]","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.081560Z","iopub.status.idle":"2024-08-27T09:36:56.081905Z","shell.execute_reply.started":"2024-08-27T09:36:56.081748Z","shell.execute_reply":"2024-08-27T09:36:56.081762Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def infer_example(model, example, device, num_classes=25):\n    x, y = example['stack'], example['labels']\n    y_gt = y\n    model = model.to(device)\n    model.eval()\n\n    with torch.no_grad():\n        y_out = model(x.unsqueeze(0).to(device))\n\n    return y_out.cpu().numpy(), y_gt\n\ninference, real = infer_example(model, data_vit[333], device)\nprint(np.argmax(inference, axis=-1))\nprint(real)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.083218Z","iopub.status.idle":"2024-08-27T09:36:56.083555Z","shell.execute_reply.started":"2024-08-27T09:36:56.083384Z","shell.execute_reply":"2024-08-27T09:36:56.083398Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"stop","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.085176Z","iopub.status.idle":"2024-08-27T09:36:56.085480Z","shell.execute_reply.started":"2024-08-27T09:36:56.085328Z","shell.execute_reply":"2024-08-27T09:36:56.085341Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.round(inference,2)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.086590Z","iopub.status.idle":"2024-08-27T09:36:56.086919Z","shell.execute_reply.started":"2024-08-27T09:36:56.086765Z","shell.execute_reply":"2024-08-27T09:36:56.086779Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"np.argmax(inference, -1)","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.088559Z","iopub.status.idle":"2024-08-27T09:36:56.089005Z","shell.execute_reply.started":"2024-08-27T09:36:56.088780Z","shell.execute_reply":"2024-08-27T09:36:56.088799Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomHead(nn.Module):\n    def __init__(self, embed_dim, num_classes):\n        super(CustomHead, self).__init__()\n        self.fc1 = nn.Linear(embed_dim, 512)\n        self.relu = nn.ReLU()\n        self.fc2 = nn.Linear(512, num_classes)\n\n    def forward(self, x):\n        x = self.fc1(x)\n        x = self.relu(x)\n        x = self.fc2(x)\n        return x\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nmodel = timm.create_model('vit_base_patch16_224', pretrained=True, in_chans=25)\nmodel.head = nn.Identity()\n\nembed_dim = 768\nnum_classes = 25\ncustom_head = CustomHead(embed_dim, num_classes).to(device)\nmodel.head = custom_head\n\nmodel.to(device)\n\ncriterion = nn.CrossEntropyLoss()\nlearning_rate = 0.01\nnum_epochs = 10\noptimizer = optim.Adam(model.parameters(), lr=learning_rate)\n\n# DataLoader definition\nimage_transforms = transforms.Compose([\n    transforms.Resize((224, 224), antialias=True),\n    transforms.ConvertImageDtype(torch.float32)\n])\n\nclass ViTDataset(Dataset):\n    def __init__(self, data, image_dir, transform=image_transforms):\n        self.data = data.copy()\n        self.image_dir = Path(image_dir)\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data)\n    \n    def __getitem__(self, idx):\n        row = self.data.iloc[idx]\n        study_id = row['study_id']\n\n        instance_numbers = [\n            row['instance_number_x'],\n            row['instance_number_y'],\n            row['instance_number_z']\n        ]\n        series_ids = [\n            row['series_id_x'],\n            row['series_id_y'],\n            row['series_id_z']\n        ]\n\n        volumes = [self.get_volume(instances, study_id, series_id) for instances, series_id in zip(instance_numbers, series_ids)]\n        \n        stack = self.prepare_stack(volumes)\n        labels = self.get_labels(row)\n\n        return {\n            \"stack\": stack.to(device),\n            \"labels\": labels.to(device)\n        }\n    \n    @staticmethod\n    def to_tensor(values):\n        return torch.tensor(values, dtype=torch.float32)\n    \n    def get_volume(self, instance_numbers, study_id, series_id):      \n        image_paths = [\n            self.image_dir / f\"{study_id}/{serie_id}/{instance_number}.dcm\"\n            for instance_number in instance_numbers\n            for serie_id in series_id\n        ]\n        image_paths = [path for path in image_paths if path.exists()]\n        \n        images = [self.load_and_transform_image(image_path) for image_path in image_paths]\n\n        return torch.stack(images)\n\n    def load_and_transform_image(self, image_path):\n        dicom = pydicom.dcmread(image_path)\n        image = dicom.pixel_array\n        image = self.apply_windowing(image, dicom)\n        image = self.normalize_image(image)\n        image = self.to_tensor(image)\n        \n        if self.transform:\n            image = self.transform(image.unsqueeze(0)).squeeze(0)\n\n        return image\n\n    def apply_windowing(self, image, dicom):\n        window_level, window_width = dicom.WindowCenter, dicom.WindowWidth\n        lowest_visible_value = np.minimum(window_level - window_width / 2, 0)\n        highest_visible_value = np.minimum(window_level + window_width / 2, 1000)\n        image = image.clip(lowest_visible_value, highest_visible_value)\n        return image\n\n    def normalize_image(self, image):\n        min_pixel_value, max_pixel_value = image.min(), image.max()\n        return (image - min_pixel_value) / (max_pixel_value - min_pixel_value + 1e-10)\n\n    def prepare_stack(self, volumes):\n        stack = torch.cat(volumes, dim=0)\n        n_slides = stack.shape[0]\n        if n_slides < 25:\n            padding = torch.zeros((25 - n_slides, 224, 224), dtype=torch.float32)\n            stack = torch.cat((stack, padding), dim=0)\n        else:\n            stack = stack[:25]\n        return stack\n    \n    def get_labels(self, row):\n        conditions = row[\n            [\n                'left_neural_foraminal_narrowing_l1_l2',\n                'left_neural_foraminal_narrowing_l2_l3',\n                'left_neural_foraminal_narrowing_l3_l4',\n                'left_neural_foraminal_narrowing_l4_l5',\n                'left_neural_foraminal_narrowing_l5_s1',\n                'left_subarticular_stenosis_l1_l2', \n                'left_subarticular_stenosis_l2_l3',\n                'left_subarticular_stenosis_l3_l4',\n                'left_subarticular_stenosis_l4_l5',\n                'left_subarticular_stenosis_l5_s1',\n                'right_neural_foraminal_narrowing_l1_l2',\n                'right_neural_foraminal_narrowing_l2_l3',\n                'right_neural_foraminal_narrowing_l3_l4',\n                'right_neural_foraminal_narrowing_l4_l5',\n                'right_neural_foraminal_narrowing_l5_s1',\n                'right_subarticular_stenosis_l1_l2',\n                'right_subarticular_stenosis_l2_l3',\n                'right_subarticular_stenosis_l3_l4',\n                'right_subarticular_stenosis_l4_l5',\n                'right_subarticular_stenosis_l5_s1',\n                'spinal_canal_stenosis_l1_l2',\n                'spinal_canal_stenosis_l2_l3',\n                'spinal_canal_stenosis_l3_l4',\n                'spinal_canal_stenosis_l4_l5',\n                'spinal_canal_stenosis_l5_s1'\n            ]\n        ].values\n        return torch.tensor(np.array(conditions, dtype=np.float32), dtype=torch.float32)\n\ndata_vit = ViTDataset(data_pivoted, image_dir)\ndataloader_vit = DataLoader(data_vit, batch_size=4, shuffle=True)\n\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n    progress_bar = tqdm(enumerate(dataloader_vit), total=len(dataloader_vit))\n    for i, batch in progress_bar:\n        inputs, labels = batch['stack'], batch['labels']\n\n        optimizer.zero_grad()\n        outputs = model(inputs)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n\n        running_loss += loss.item()\n        progress_bar.set_description(f'Epoch [{epoch + 1}/{num_epochs}], Loss: {running_loss / (i + 1):.4f}')\n\n    print(f'Epoch [{epoch + 1}/{num_epochs}] complete. Average Loss: {running_loss / len(dataloader_vit):.4f}')\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.090287Z","iopub.status.idle":"2024-08-27T09:36:56.090761Z","shell.execute_reply.started":"2024-08-27T09:36:56.090496Z","shell.execute_reply":"2024-08-27T09:36:56.090514Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---","metadata":{}},{"cell_type":"code","source":"class VisionTransformer(nn.Module):\n    def __init__(self, img_size=224, patch_size=16, num_classes=25, dim=768, num_layers=12, num_heads=12, mlp_dim=3072, channels=1, dropout=0.1, max_num_frames=30):\n        super(VisionTransformer, self).__init__()\n        assert img_size % patch_size == 0, 'Image dimensions must be divisible by the patch size.'\n\n        self.patch_size = patch_size\n        num_patches = (img_size // patch_size) ** 2\n        patch_dim = channels * patch_size ** 2\n\n        self.to_patch_embedding = nn.Linear(patch_dim, dim)\n        self.position_embeddings = nn.Parameter(torch.zeros(1, num_patches + 1, dim))\n        self.temporal_embeddings = nn.Parameter(torch.zeros(1, max_num_frames, dim))\n        self.class_token = nn.Parameter(torch.zeros(1, 1, dim))\n        self.dropout = nn.Dropout(dropout)\n\n        self.transformer_layers = nn.ModuleList([\n            nn.TransformerEncoderLayer(d_model=dim, nhead=num_heads, dim_feedforward=mlp_dim, dropout=dropout)\n            for _ in range(num_layers)\n        ])\n\n        self.mlp_head = nn.Sequential(\n            nn.LayerNorm(dim),\n            nn.Linear(dim, num_classes)\n        )\n\n    def forward(self, x):\n        batch_size, num_frames, _, _, _ = x.shape\n\n        x = self.preprocess_input(x)\n        x = self.add_embeddings(x, batch_size, num_frames)\n        x = self.apply_transformer(x)\n\n        cls_token_final = x[:, 0]\n        output = self.mlp_head(cls_token_final)\n\n        output = output.view(batch_size, -1, output.size(-1)).mean(dim=1)\n\n        return output\n\n    def preprocess_input(self, x):\n        patch_size = self.patch_size\n        patches = rearrange(x, 'b f c (h p1) (w p2) -> (b f) (h w) (p1 p2 c)', p1=patch_size, p2=patch_size)\n        return self.to_patch_embedding(patches)\n\n    def add_embeddings(self, x, batch_size, num_frames):\n\n        cls_tokens = repeat(self.class_token, '() n d -> b n d', b=batch_size * num_frames)\n        x = torch.cat((cls_tokens, x), dim=1)\n        x += self.position_embeddings[:, :(x.size(1))]\n\n        temporal_embeddings = repeat(self.temporal_embeddings[:, :num_frames], '() f d -> (b f) 1 d', b=batch_size)\n        x += temporal_embeddings\n        \n        return self.dropout(x)\n\n    def apply_transformer(self, x):\n        for layer in self.transformer_layers:\n            x = layer(x)\n        return x\n\nvolumen_imagenes = torch.randn(4, 10, 1, 224, 224) # batch_size=1, num_frames=10, channels=1, height=224, width=224\n\nmodel = VisionTransformer()\noutputs = model(volumen_imagenes)\n\nprint(outputs.shape) #[batch_size, num_classes]","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.092035Z","iopub.status.idle":"2024-08-27T09:36:56.092461Z","shell.execute_reply.started":"2024-08-27T09:36:56.092239Z","shell.execute_reply":"2024-08-27T09:36:56.092256Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import DataLoader\n\n# Define device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\n# Initialize model, dataset, and dataloader\nmodel = VisionTransformer().to(device)\ndata_vit = ViTDataset(data_pivoted, image_dir, unet_path)\ndataloader_vit = DataLoader(data_vit, batch_size=1, shuffle=False)\n\n# Loss and optimizer\ncriterion = nn.CrossEntropyLoss()#BCEWithLogitsLoss()  # or any other suitable loss function\noptimizer = optim.Adam(model.parameters(), lr=1e-4)\n\n# Training loop\nnum_epochs = 10  # Define the number of epochs\n\nfor epoch in range(num_epochs):\n    model.train()\n    running_loss = 0.0\n\n    for i, batch in enumerate(dataloader_vit):\n        z, y, inputs = batch['inputs'][0], batch['inputs'][1], batch['inputs'][2]\n        inputs = inputs.unsqueeze(2).to(device)\n        labels = batch['labels'].to(device)\n#         print(labels.shape)\n        # Zero the parameter gradients\n        optimizer.zero_grad()\n\n        # Forward pass\n        outputs = model(inputs)\n\n        # Compute loss\n        loss = criterion(outputs, labels)\n        \n        # Backward pass and optimize\n        loss.backward()\n        optimizer.step()\n\n        # Print statistics\n        running_loss += loss.item()\n        if i % 500 == 9:    # Print every 10 batches\n            print(f'Epoch [{epoch + 1}/{num_epochs}], Step [{i + 1}/{len(dataloader_vit)}], Loss: {running_loss / 10:.4f}')\n            running_loss = 0.0\n\nprint('Finished Training')\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.093675Z","iopub.status.idle":"2024-08-27T09:36:56.094109Z","shell.execute_reply.started":"2024-08-27T09:36:56.093889Z","shell.execute_reply":"2024-08-27T09:36:56.093908Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"dataloader_vit.dataset[0]","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.099495Z","iopub.status.idle":"2024-08-27T09:36:56.099850Z","shell.execute_reply.started":"2024-08-27T09:36:56.099687Z","shell.execute_reply":"2024-08-27T09:36:56.099701Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"---","metadata":{}},{"cell_type":"code","source":"from efficientnet_pytorch import EfficientNet\nclass VisionTransformer(nn.Module):\n    def __init__(self, num_side_classes, num_condition_classes, num_severity_classes, num_levels, freeze_efficientnet=False):\n        super(VisionTransformer, self).__init__()\n        self.efficientnet = EfficientNet.from_pretrained('efficientnet-b0')\n        efficientnet_output_dim = self.efficientnet._fc.in_features\n        if freeze_efficientnet:\n            for param in self.efficientnet.parameters():\n                param.requires_grad = False\n        \n        self.efficientnet._fc = nn.Identity()  # Remove the original FC layer\n        \n        self.side_classifier_head = nn.Sequential(\n            nn.Linear(efficientnet_output_dim, 512),\n            nn.BatchNorm1d(512),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(512, 128),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(128, num_side_classes),\n        )\n\n        self.condition_classifier_head = nn.Sequential(\n            nn.Linear(efficientnet_output_dim, 512),\n            nn.BatchNorm1d(512),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(512, 128),\n            nn.BatchNorm1d(128),\n            nn.ReLU(),\n            nn.Dropout(0.5),\n            nn.Linear(128, num_condition_classes),\n        )\n        \n        self.severity_classifier_heads = nn.ModuleList([\n            nn.Sequential(\n                nn.Linear(efficientnet_output_dim + 2, 512),\n                nn.BatchNorm1d(512),\n                nn.ReLU(),\n                nn.Dropout(0.5),\n                nn.Linear(512, 128),\n                nn.BatchNorm1d(128),\n                nn.ReLU(),\n                nn.Dropout(0.5),\n                nn.Linear(128, num_severity_classes)\n            ) for _ in range(num_levels)\n        ])\n\n    def forward(self, x, x_coords, y_coords):\n        outputs = self.efficientnet(x)\n        side_outputs = self.side_classifier_head(outputs)\n        condition_outputs = self.condition_classifier_head(outputs)\n        \n        severity_outputs = []\n        for i, severity_head in enumerate(self.severity_classifier_heads):\n            level_inputs = torch.cat((outputs, x_coords[:, i].unsqueeze(1), y_coords[:, i].unsqueeze(1)), dim=1)\n            severity_outputs.append(severity_head(level_inputs))\n        \n        return side_outputs, condition_outputs, severity_outputs\n\nnum_side_classes = 3  # Left, Right, Center\nnum_condition_classes = 3  # Neural Foraminal Narrowing, Spinal Canal Stenosis, Subarticular Stenosis\nnum_severity_classes = 3  # Moderate, Normal/Mild, Severe\nnum_levels = 5  # L1/L2, L2/L3, L3/L4, L4/L5, L5/S1\nfreeze_efficientnet = True\n\nmodel_tl = VisionTransformer(num_side_classes, num_condition_classes, num_severity_classes, num_levels, freeze_efficientnet)\nmodel_tl = model_tl.to(device)\n\ninput_tensor = torch.randn(2, 3, 224, 224).to(device)\nx_coords = torch.randn(2, num_levels).to(device)  # Assuming 5 levels for x\ny_coords = torch.randn(2, num_levels).to(device)  # Assuming 5 levels for y\noutput = model_tl(input_tensor, x_coords, y_coords)\n\nfor out in output:\n    if isinstance(out, list):\n        for o in out:\n            print(o.shape)\n    else:\n        print(out.shape)\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.101270Z","iopub.status.idle":"2024-08-27T09:36:56.101729Z","shell.execute_reply.started":"2024-08-27T09:36:56.101476Z","shell.execute_reply":"2024-08-27T09:36:56.101494Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"if True:\n    criterion_side = nn.# Create the list of image paths\nimages_paths = [self.image_dir / f\"{study_id}/{serie_id}/{instance_number}.dcm\"\n                for instance_number in instance_numbers\n                for serie_id in series_id]\n\n# Filter the list to keep only the paths that exist\nexisting_images_paths = [path for path in images_paths if path.exists()]()\n    criterion_condition = nn.CrossEntropyLoss()\n    criterion_severity = nn.CrossEntropyLoss()\n    \n    optimizer = torch.optim.Adam(model_tl.parameters(), lr=1e-4)\n    scheduler = ReduceLROnPlateau(optimizer, 'min', patience=3, verbose=True)\n\n    def train_model(model, dataloader, optimizer, num_epochs=10):\n        best_loss = float('inf')\n        patience = 5\n        counter = 0\n\n        train_loss_history = []\n        train_acc_history = []\n\n        model.train()\n        for epoch in range(num_epochs):\n            running_loss = 0.0\n            running_corrects_side = 0\n            running_corrects_condition = 0\n            running_corrects_severity = [0] * num_levels\n            \n            progress_bar = tqdm(dataloader, desc=f\"Epoch {epoch+1}/{num_epochs}\", unit=\"batch\")\n            for i, data in enumerate(progress_bar):\n                inputs = data['image'].to(device).float()\n                side_labels = data['side'].to(device).long().squeeze()\n                condition_labels = data['condition'].to(device).long().squeeze()\n                severity_labels = data['severity'].to(device).long().squeeze()\n                x_coords = data['x'].to(device).float()\n                y_coords = data['y'].to(device).float()\n\n                optimizer.zero_grad()\n\n                side_outputs, condition_outputs, severity_outputs = model(inputs, x_coords, y_coords)\n             \n                loss_side = criterion_side(side_outputs, side_labels)\n                loss_condition = criterion_condition(condition_outputs, condition_labels)\n                loss_severity = sum([criterion_severity(severity_outputs[j], severity_labels[:, j]) for j in range(num_levels)])\n                \n                loss = (\n                    0.2 * loss_side +\n                    0.2 * loss_condition + \n                    0.6 * loss_severity / num_levels  # Averaging the severity loss\n                )\n\n                loss.backward()\n                optimizer.step()\n\n                running_loss += loss.item()\n\n                _, preds_side = torch.max(side_outputs, 1)\n                running_corrects_side += torch.sum(preds_side == side_labels)\n\n                _, preds_condition = torch.max(condition_outputs, 1)\n                running_corrects_condition += torch.sum(preds_condition == condition_labels)\n\n                for j in range(num_levels):\n                    _, preds_severity = torch.max(severity_outputs[j], 1)\n                    running_corrects_severity[j] += torch.sum(preds_severity == severity_labels[:, j])\n\n                progress_bar.set_postfix(\n                    loss=running_loss / (i + 1),\n                    accuracy_side=(running_corrects_side.double() / ((i + 1) * dataloader.batch_size)).item(),\n                    accuracy_condition=(running_corrects_condition.double() / ((i + 1) * dataloader.batch_size)).item(),\n                    accuracy_severity=(sum(running_corrects_severity).double() / (num_levels * (i + 1) * dataloader.batch_size)).item(),\n                )\n\n            epoch_loss = running_loss / len(dataloader.dataset)\n            epoch_acc_side = running_corrects_side.double() / len(dataloader.dataset)\n            epoch_acc_condition = running_corrects_condition.double() / len(dataloader.dataset)\n            epoch_acc_severity = sum(running_corrects_severity).double() / (num_levels * len(dataloader.dataset))\n            epoch_acc = (epoch_acc_side + epoch_acc_condition + epoch_acc_severity) / 3\n\n            train_loss_history.append(epoch_loss)\n            train_acc_history.append(epoch_acc.item())\n\n            print(f\"Epoch {epoch+1}/{num_epochs} Loss: {epoch_loss:.4f} Accuracy: {epoch_acc:.4f}\")\n            \n            if epoch_loss < best_loss:\n                best_loss = epoch_loss\n                counter = 0\n\n                checkpoint = {\n                    'epoch': epoch,\n                    'model_state_dict': model.state_dict(),\n                    'optimizer_state_dict': optimizer.state_dict()\n                }\n\n                torch.save(checkpoint, 'checkpoint.pth')\n            else:\n                counter += 1\n                if counter >= patience:\n                    print(\"Early stopping\")\n                    break\n\n            scheduler.step(epoch_loss)\n\n        print('Finished Training')\n\n        return train_loss_history, train_acc_history\n\n    train_loss_history, train_acc_history = train_model(model_tl, dataloader_vit, optimizer, num_epochs=5)\n\n    window_size = 3\n    window = np.ones(window_size) / window_size\n\n    rolling_loss = np.convolve(train_loss_history, window, mode='valid')\n    rolling_acc = np.convolve(train_acc_history, window, mode='valid')\n\n    rolling_epochs = range(1, len(rolling_loss) + 1)\n    epochs = range(1, len(train_loss_history) + 1)\n\n    fig, ax = plt.subplots()\n    plt.title('Training loss and accuracy')\n\n    color = 'tab:red'\n    ax.set_xlabel('Epochs')\n    ax.set_ylabel('Loss', color=color)\n    ax.tick_params(axis='y', labelcolor=color)\n    ax.plot(epochs, train_loss_history, color=color, linestyle='--', alpha=0.5, label='Training loss (original)')\n    ax.plot(rolling_epochs, rolling_loss, color=color, linestyle='-', label='Training loss (rolling mean)')\n\n    ax_ = ax.twinx() \n    color = 'tab:blue'\n    ax_.set_ylabel('Accuracy', color=color)\n    ax_.tick_params(axis='y', labelcolor=color)\n    ax_.plot(epochs, train_acc_history, color=color, linestyle='--', alpha=0.5, label='Training accuracy (original)')\n    ax_.plot(rolling_epochs, rolling_acc, color=color, linestyle='-', label='Training accuracy (rolling mean)')\n\n    fig.tight_layout()\n    fig.legend(loc='upper right', bbox_to_anchor=(1,1), bbox_transform=ax.transAxes)\n    plt.show()\nelse:\n    checkpoint = torch.load('/kaggle/working/checkpoint.pth', map_location=torch.device(device))\n    state_dict = checkpoint['model_state_dict']\n    for key in ['vit.pooler.dense.weight', 'vit.pooler.dense.bias']:\n        if key in state_dict:\n            del state_dict[key]\n\n    model_tl.load_state_dict(state_dict, strict=False)\n    optimizer = torch.optim.Adam(model_tl.parameters(), lr=1e-4)\n    scheduler = ReduceLROnPlateau(optimizer, 'min', patience=3, verbose=True)\n    epoch = checkpoint['epoch']\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.103082Z","iopub.status.idle":"2024-08-27T09:36:56.103518Z","shell.execute_reply.started":"2024-08-27T09:36:56.103290Z","shell.execute_reply":"2024-08-27T09:36:56.103310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"########################################################################\n#                                                                      #\n#                            ViT DEPLOYMENT                            #\n#                                                                      #\n########################################################################\n\ndef get_image_from_dataset(dataset, index=0):\n    data = dataset[index]\n    image = data['image'].unsqueeze(0).to(device)\n    side_label = data['side'].to(device)\n    condition_label = data['condition'].to(device)\n    severity_label = data['severity'].to(device)\n    x_coords = data['x'].to(device)\n    y_coords = data['y'].to(device)\n    return image, side_label, condition_label, severity_label, x_coords, y_coords\n\ndef predict_from_dataset(model, dataset, index=0):\n    image, side_label, condition_label, severity_label, x_coords, y_coords = get_image_from_dataset(dataset, index)\n\n    model.eval()\n    with torch.no_grad():\n        side_outputs, condition_outputs, severity_outputs = model(image, x_coords.unsqueeze(0), y_coords.unsqueeze(0))\n\n    return side_outputs, condition_outputs, severity_outputs, side_label, condition_label, severity_label\n\ndef get_predictions_and_labels(outputs, labels):\n    predicted = outputs.argmax(dim=1)\n    actual = labels.squeeze()\n    return predicted, actual\n\nindex = np.random.randint(len(dataloader_vit.dataset))\nside_outputs, condition_outputs, severity_outputs, side_label, condition_label, severity_label = predict_from_dataset(model_tl, dataloader_vit.dataset, index=index)\n\nside_pred, side_actual = get_predictions_and_labels(side_outputs, side_label)\ncondition_pred, condition_actual = get_predictions_and_labels(condition_outputs, condition_label)\n\n# Para severity_outputs y severity_label que son listas\nseverity_preds = [severity_outputs[j].argmax(dim=1).item() for j in range(len(severity_outputs))]\nseverity_actuals = [severity_label[0, j].item() for j in range(len(severity_label[0]))]\n\nprint(\"Side Outputs (Predicted):\", side_pred.item(), \"Actual:\", side_actual.item())\nprint(\"Condition Outputs (Predicted):\", condition_pred.item(), \"Actual:\", condition_actual.item())\n\nfor i, (pred, actual) in enumerate(zip(severity_preds, severity_actuals)):\n    print(f\"Severity Outputs for Level {i} (Predicted): {pred} Actual: {actual}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.104497Z","iopub.status.idle":"2024-08-27T09:36:56.104940Z","shell.execute_reply.started":"2024-08-27T09:36:56.104716Z","shell.execute_reply":"2024-08-27T09:36:56.104734Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"data_pivoted['severity'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-08-27T09:36:56.106443Z","iopub.status.idle":"2024-08-27T09:36:56.106893Z","shell.execute_reply.started":"2024-08-27T09:36:56.106674Z","shell.execute_reply":"2024-08-27T09:36:56.106692Z"},"trusted":true},"execution_count":null,"outputs":[]}]}