{"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":"none","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"dockerImageVersionId":30733,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport pydicom as dicom\nimport matplotlib.pyplot as plt\nfrom glob import glob\nimport seaborn as sns\nimport numpy as np\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torchvision.transforms as transforms\nfrom torch.utils.data import DataLoader, Dataset, WeightedRandomSampler, Subset\nfrom transformers import ViTForImageClassification, ViTFeatureExtractor\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.model_selection import KFold\nfrom torch.utils.data import ConcatDataset\nimport torchvision.models as models\nfrom PIL import Image\nimport timm\nimport os\nimport pydicom\nfrom tqdm import tqdm\nimport random\nimport gc\nimport warnings\nwarnings.filterwarnings(\"ignore\")\n%matplotlib inline","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-06-21T07:49:54.280904Z","iopub.execute_input":"2024-06-21T07:49:54.281470Z","iopub.status.idle":"2024-06-21T07:50:13.672892Z","shell.execute_reply.started":"2024-06-21T07:49:54.281435Z","shell.execute_reply":"2024-06-21T07:50:13.672113Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def seeding(SEED):\n    np.random.seed(SEED)\n    random.seed(SEED)\n    os.environ['PYTHONHASHSEED'] = str(SEED)\n    torch.manual_seed(SEED)\n    if torch.cuda.is_available(): \n        torch.cuda.manual_seed(SEED)\n        torch.cuda.manual_seed_all(SEED)\n        torch.backends.cudnn.deterministic = True\n        torch.backends.cudnn.benchmark = False\n    print('seeding done!!!')\n\ndef flush():\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n        torch.cuda.reset_peak_memory_stats()\n    print('Memory Flushed')\nseeding(2024)\nflush()","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:57.269742Z","iopub.execute_input":"2024-06-21T07:50:57.270372Z","iopub.status.idle":"2024-06-21T07:50:57.629581Z","shell.execute_reply.started":"2024-06-21T07:50:57.270339Z","shell.execute_reply":"2024-06-21T07:50:57.628618Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom(path):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:57.631415Z","iopub.execute_input":"2024-06-21T07:50:57.631868Z","iopub.status.idle":"2024-06-21T07:50:57.637295Z","shell.execute_reply.started":"2024-06-21T07:50:57.631831Z","shell.execute_reply":"2024-06-21T07:50:57.636399Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train.csv\")\nprint(df_train.shape)\ndf_train.head()","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:57.638434Z","iopub.execute_input":"2024-06-21T07:50:57.638713Z","iopub.status.idle":"2024-06-21T07:50:57.699009Z","shell.execute_reply.started":"2024-06-21T07:50:57.638690Z","shell.execute_reply":"2024-06-21T07:50:57.698058Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train.isna().sum().sort_values(ascending=False)","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:57.748839Z","iopub.execute_input":"2024-06-21T07:50:57.749097Z","iopub.status.idle":"2024-06-21T07:50:57.766815Z","shell.execute_reply.started":"2024-06-21T07:50:57.749074Z","shell.execute_reply":"2024-06-21T07:50:57.766011Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Train/Test Series Descriptions","metadata":{}},{"cell_type":"code","source":"df_train_series_desc = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_series_descriptions.csv\")\nprint(df_train_series_desc.shape)\nprint(\"Unique study ids in train description: \", len(df_train_series_desc['study_id'].unique()))\ndf_train_series_desc.head()","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:58.137377Z","iopub.execute_input":"2024-06-21T07:50:58.137667Z","iopub.status.idle":"2024-06-21T07:50:58.159160Z","shell.execute_reply.started":"2024-06-21T07:50:58.137643Z","shell.execute_reply":"2024-06-21T07:50:58.158312Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_categories = {\n    'Sagittal T2/STIR': {\n        'series': ['spinal_canal_stenosis'],\n        'series_ids': [],\n        'images': []\n    },\n    'Sagittal T1': {\n        'series': ['right_neural_foraminal_narrowing', 'left_neural_foraminal_narrowing'],\n        'series_ids': [],\n        'images': []\n    },\n    'Axial T2': {\n        'series': ['right_subarticular_stenosis', 'left_subarticular_stenosis'],\n        'series_ids': [],\n        'images': []\n    }\n}","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:58.222814Z","iopub.execute_input":"2024-06-21T07:50:58.223072Z","iopub.status.idle":"2024-06-21T07:50:58.228552Z","shell.execute_reply.started":"2024-06-21T07:50:58.223051Z","shell.execute_reply":"2024-06-21T07:50:58.227315Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_test_series_desc = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_series_descriptions.csv\")\nprint(df_test_series_desc.shape)\ndf_test_series_desc.head()","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:58.325315Z","iopub.execute_input":"2024-06-21T07:50:58.325637Z","iopub.status.idle":"2024-06-21T07:50:58.340398Z","shell.execute_reply.started":"2024-06-21T07:50:58.325611Z","shell.execute_reply":"2024-06-21T07:50:58.339554Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"#### Training Label Coordinates","metadata":{}},{"cell_type":"code","source":"df_train_label_coordinates = pd.read_csv(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_label_coordinates.csv\")\nprint(df_train_label_coordinates.shape)\nprint(\"Unique Study Ids in label coordinates: \", len(df_train_label_coordinates['study_id'].unique()))\ndf_train_label_coordinates.head()","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:58.550231Z","iopub.execute_input":"2024-06-21T07:50:58.550663Z","iopub.status.idle":"2024-06-21T07:50:58.672517Z","shell.execute_reply.started":"2024-06-21T07:50:58.550638Z","shell.execute_reply":"2024-06-21T07:50:58.671633Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train_label_coordinates['condition'].unique()","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:58.686398Z","iopub.execute_input":"2024-06-21T07:50:58.686664Z","iopub.status.idle":"2024-06-21T07:50:58.697087Z","shell.execute_reply.started":"2024-06-21T07:50:58.686643Z","shell.execute_reply":"2024-06-21T07:50:58.696216Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# # Total number of images in training\n# count = 0\n# study_ids = glob(\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/*\")\n# for study_id in study_ids:\n#     series_ids = glob(f\"{study_id}/*\")\n#     for series_id in series_ids:\n#         images = glob(f\"{series_id}/*\")\n#         count += len(images)\n        \n# print(\"Total number if images:\", count) # 147218","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:58.813796Z","iopub.execute_input":"2024-06-21T07:50:58.814054Z","iopub.status.idle":"2024-06-21T07:50:58.818034Z","shell.execute_reply.started":"2024-06-21T07:50:58.814032Z","shell.execute_reply":"2024-06-21T07:50:58.817122Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_id_paths = glob(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/*\")\nunique_study_ids = [i.split(\"/\")[-1] for i in study_id_paths]\n\n# Matches 'df_train' unique study ids and 'df_test_series_desc' unique study ids\nprint(\"Unique Study Ids: \", len(unique_study_ids))","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:58.948820Z","iopub.execute_input":"2024-06-21T07:50:58.949082Z","iopub.status.idle":"2024-06-21T07:50:59.016598Z","shell.execute_reply.started":"2024-06-21T07:50:58.949059Z","shell.execute_reply":"2024-06-21T07:50:59.015749Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# number_of_images = {}\n# number_of_series_for_study_id = {}\n# for study_id in unique_study_ids:\n#     series_list = glob(f\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/{str(study_id)}/*\")\n#     number_of_images[study_id] = {}\n#     for series in series_list:\n#         series_id = series.split(\"/\")[-1]\n#         files = glob(f\"{series}/*\")\n#         number_of_images[study_id][series_id] = len(files)\n\n# print(\"For each study id, given all the series and number of images:\\n\")\n# list(number_of_images.items())[:5]","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:50:59.069270Z","iopub.execute_input":"2024-06-21T07:50:59.069727Z","iopub.status.idle":"2024-06-21T07:50:59.073529Z","shell.execute_reply.started":"2024-06-21T07:50:59.069703Z","shell.execute_reply":"2024-06-21T07:50:59.072518Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"### Loading Images of each 3 classes","metadata":{}},{"cell_type":"code","source":"categories = list(image_categories.keys())\n\npath = \"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images\"\nfor study_id in unique_study_ids:\n    study_id_path = os.path.join(path, study_id)\n    for category in categories:\n        try:\n            series_ids = df_train_series_desc[\n                (df_train_series_desc['study_id'] == int(study_id)) &\n                (df_train_series_desc['series_description'] == category)\n            ]['series_id'].tolist()\n            image_categories[category]['series_ids'].extend(series_ids)\n            \n            for series_id in series_ids:\n                series_id_path = os.path.join(study_id_path, str(series_id))\n                image_paths = glob(f\"{series_id_path}/*\")\n                image_categories[category]['images'].extend(image_paths)\n\n        except Exception as e:\n            print(study_id)\n            print(category)\n            continue","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:51:16.452006Z","iopub.execute_input":"2024-06-21T07:51:16.452345Z","iopub.status.idle":"2024-06-21T07:52:03.996182Z","shell.execute_reply.started":"2024-06-21T07:51:16.452315Z","shell.execute_reply":"2024-06-21T07:52:03.995395Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# print(\"Number of Images in 'Sagittal T1' are: \", len(image_categories['Sagittal T1']['images']))\n# print(\"Number of Images in 'Sagittal T2/STIR' are: \", len(image_categories['Sagittal T2/STIR']['images']))\n# print(\"Number of Images in 'Axial T2' are: \", len(image_categories['Axial T2']['images']))","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:03.998005Z","iopub.execute_input":"2024-06-21T07:52:03.998359Z","iopub.status.idle":"2024-06-21T07:52:04.002697Z","shell.execute_reply.started":"2024-06-21T07:52:03.998325Z","shell.execute_reply":"2024-06-21T07:52:04.001663Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Convert 'level' column to lowercase and replace '/' with '_'\ndf_train_label_coordinates['level'] = df_train_label_coordinates['level'].str.lower().str.replace('/', '_')\n\n# Convert 'condition' column to lowercase and replace spaces with '_'\ndf_train_label_coordinates['condition'] = df_train_label_coordinates['condition'].str.lower().str.replace(' ', '_')\n\ndf_train_label_coordinates['label'] = df_train_label_coordinates.apply(\n    lambda x: df_train.loc[df_train['study_id'] == x['study_id'], f\"{x['condition']}_{x['level']}\"].values[0],\n    axis=1\n)","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:04.003836Z","iopub.execute_input":"2024-06-21T07:52:04.004125Z","iopub.status.idle":"2024-06-21T07:52:19.890764Z","shell.execute_reply.started":"2024-06-21T07:52:04.004101Z","shell.execute_reply":"2024-06-21T07:52:19.890005Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train_label_coordinates = df_train_label_coordinates.dropna(subset=['label'])","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:19.893380Z","iopub.execute_input":"2024-06-21T07:52:19.893750Z","iopub.status.idle":"2024-06-21T07:52:19.909002Z","shell.execute_reply.started":"2024-06-21T07:52:19.893717Z","shell.execute_reply":"2024-06-21T07:52:19.908267Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"df_train_label_coordinates[df_train_label_coordinates['label'] == 'Moderate']['condition'].value_counts()","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:19.910028Z","iopub.execute_input":"2024-06-21T07:52:19.910292Z","iopub.status.idle":"2024-06-21T07:52:19.930345Z","shell.execute_reply.started":"2024-06-21T07:52:19.910268Z","shell.execute_reply":"2024-06-21T07:52:19.929310Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Building Dataset","metadata":{}},{"cell_type":"code","source":"class SpineDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None):\n        self.df = df\n        self.img_dir = img_dir\n        self.transform = transform\n        self.class_to_index = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def get_labels(self):\n        row = self.df\n        label = row['label']\n        label = label.map(self.class_to_index)\n        return label.tolist()\n        \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        study_id = row['study_id']\n        series_id = row['series_id']\n        instance_number = row['instance_number']\n        label = row['label']\n        image_path = os.path.join(self.img_dir, f'{study_id}/{series_id}/{instance_number}.dcm')\n        \n#         dicom_array = dicom.dcmread(image_path)\n#         image = dicom_array.pixel_array\n#         image = cv2.resize(image, (224, 224))\n#         image = Image.fromarray(image).convert(\"RGB\")\n        image = load_dicom(image_path)\n        \n        if self.transform:\n            image = cv2.cvtColor(image, cv2.COLOR_GRAY2BGR)\n            image = self.transform(image)\n#             image = image.transpose(2, 0, 1).astype(np.float32) / 255.\n        \n        label_idx = self.class_to_index[label]\n        label_tensor = torch.tensor(label_idx, dtype=torch.long)\n        \n        return {'image': image, 'label': label_tensor}","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:19.931292Z","iopub.execute_input":"2024-06-21T07:52:19.931565Z","iopub.status.idle":"2024-06-21T07:52:19.941055Z","shell.execute_reply.started":"2024-06-21T07:52:19.931535Z","shell.execute_reply":"2024-06-21T07:52:19.940066Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Grayscale(num_output_channels=3),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(1),\n    transforms.Resize((512, 512)),\n    transforms.ToTensor(),\n#     transforms.Normalize(mean=[0.0, 0.0, 0.0], std=[1.0, 1.0, 1.0])\n    transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])\n])\n\nval_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Grayscale(num_output_channels=3),\n    transforms.Resize((512, 512)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])\n])\n\n\naugmented_transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Grayscale(num_output_channels=3),\n    transforms.RandomHorizontalFlip(p=0.5),\n    transforms.RandomVerticalFlip(p=0.5),\n    transforms.RandomRotation(1),\n    transforms.ColorJitter(brightness=0.1, contrast=0.1, saturation=0.1, hue=0.1),\n    transforms.Resize((512, 512)),\n    transforms.ToTensor(),\n    transforms.Normalize(mean=[0.5, 0.5, 0.5], std=[0.5, 0.5, 0.5])\n#     transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225])\n])","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:19.942226Z","iopub.execute_input":"2024-06-21T07:52:19.942627Z","iopub.status.idle":"2024-06-21T07:52:19.953917Z","shell.execute_reply.started":"2024-06-21T07:52:19.942601Z","shell.execute_reply":"2024-06-21T07:52:19.953175Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class AugmentedSpineDataset(Dataset):\n    def __init__(self, df, img_dir, transform=None):\n        self.df = df\n        self.img_dir = img_dir\n        self.transform = transform\n        self.class_to_index = {'Normal/Mild': 0, 'Moderate': 1, 'Severe': 2}\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def get_labels(self):\n        row = self.df\n        label = row['label']\n        label = label.map(self.class_to_index)\n        return label.tolist()\n        \n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        study_id = row['study_id']\n        series_id = row['series_id']\n        instance_number = row['instance_number']\n        label = row['label']\n        image_path = os.path.join(self.img_dir, f'{study_id}/{series_id}/{instance_number}.dcm')\n        image = load_dicom(image_path)\n        \n        if self.transform:\n            image = cv2.cvtColor(image, cv2.COLOR_GRAY2BGR)\n            image = self.transform(image)\n        \n        label_idx = self.class_to_index[label]\n        label_tensor = torch.tensor(label_idx, dtype=torch.long)\n        \n        return {'image': image, 'label': label_tensor}","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:19.954809Z","iopub.execute_input":"2024-06-21T07:52:19.955047Z","iopub.status.idle":"2024-06-21T07:52:19.966910Z","shell.execute_reply.started":"2024-06-21T07:52:19.955023Z","shell.execute_reply":"2024-06-21T07:52:19.966168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CustomConcatDataset(ConcatDataset):\n    def __init__(self, datasets):\n        super(CustomConcatDataset, self).__init__(datasets)\n    \n    def get_labels(self):\n        # Example of aggregating custom functions from each dataset\n        results = []\n        for dataset in self.datasets:\n            if hasattr(dataset, 'get_labels'):\n                results.append(dataset.get_labels())\n        return results","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:19.967900Z","iopub.execute_input":"2024-06-21T07:52:19.968220Z","iopub.status.idle":"2024-06-21T07:52:19.978367Z","shell.execute_reply.started":"2024-06-21T07:52:19.968188Z","shell.execute_reply":"2024-06-21T07:52:19.977535Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spinal_df_train = df_train_label_coordinates[df_train_label_coordinates['condition'] == 'spinal_canal_stenosis']\nforaminal_df_train = df_train_label_coordinates[df_train_label_coordinates['condition'].isin(['right_neural_foraminal_narrowing', 'left_neural_foraminal_narrowing'])]\nsabarticular_df_train = df_train_label_coordinates[df_train_label_coordinates['condition'].isin(['right_subarticular_stenosis', 'left_subarticular_stenosis'])]\nspinal_df_train.shape, foraminal_df_train.shape, sabarticular_df_train.shape","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:19.982664Z","iopub.execute_input":"2024-06-21T07:52:19.982915Z","iopub.status.idle":"2024-06-21T07:52:20.011429Z","shell.execute_reply.started":"2024-06-21T07:52:19.982893Z","shell.execute_reply":"2024-06-21T07:52:20.010549Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"spinal_train_data, spinal_val_data = train_test_split(spinal_df_train, test_size=0.2, random_state=42)\n\nspinal_train_dataset = SpineDataset(spinal_train_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=train_transform)\nspinal_val_dataset = SpineDataset(spinal_val_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=val_transform)\n\nmoderate_severe_spinal_train_data = spinal_train_data[spinal_train_data['label'] != 'Normal/Mild']\nmoderate_severe_spinal_val_data = spinal_val_data[spinal_val_data['label'] != 'Normal/Mild']\naugmented_spinal_train_dataset = AugmentedSpineDataset(moderate_severe_spinal_train_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=augmented_transform)\naugmented_spinal_val_dataset = AugmentedSpineDataset(moderate_severe_spinal_val_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=val_transform)\ncombined_spinal_train_dataset = CustomConcatDataset([spinal_train_dataset, augmented_spinal_train_dataset])\ncombined_spinal_val_dataset = CustomConcatDataset([spinal_val_dataset, augmented_spinal_val_dataset])\n\nclass_weights = torch.tensor([1, 2, 4])\ntrain_labels = combined_spinal_train_dataset.get_labels()\nlabels_list = []\nfor i in train_labels:\n    labels_list.extend(i)\nsamples_weights = class_weights[labels_list]\nsampler = WeightedRandomSampler(weights=samples_weights, \n                                num_samples=len(samples_weights), \n                                replacement=True)\n\nspinal_train_loader = DataLoader(\n    combined_spinal_train_dataset,\n    batch_size=32,\n    sampler=sampler\n)\n\nval_labels = combined_spinal_val_dataset.get_labels()\nlabels_list = []\nfor i in val_labels:\n    labels_list.extend(i)\nsamples_weights = class_weights[labels_list]\nsampler = WeightedRandomSampler(weights=samples_weights, \n                                num_samples=len(samples_weights), \n                                replacement=True)\n\nspinal_val_loader = DataLoader(\n    combined_spinal_val_dataset,\n    batch_size=32,\n    sampler=sampler\n)","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:20.012446Z","iopub.execute_input":"2024-06-21T07:52:20.012868Z","iopub.status.idle":"2024-06-21T07:52:20.074952Z","shell.execute_reply.started":"2024-06-21T07:52:20.012793Z","shell.execute_reply":"2024-06-21T07:52:20.074067Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"foraminal_train_data, foraminal_val_data = train_test_split(foraminal_df_train, test_size=0.2, random_state=42)\n\nforaminal_train_dataset = SpineDataset(foraminal_train_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=train_transform)\nforaminal_val_dataset = SpineDataset(foraminal_val_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=val_transform)\n\nmoderate_severe_foraminal_train_data = foraminal_train_data[foraminal_train_data['label'] != 'Normal/Mild']\nmoderate_severe_foraminal_val_data = foraminal_val_data[foraminal_val_data['label'] != 'Normal/Mild']\naugmented_foraminal_train_dataset = AugmentedSpineDataset(moderate_severe_foraminal_train_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=augmented_transform)\naugmented_foraminal_val_dataset = AugmentedSpineDataset(moderate_severe_foraminal_val_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=val_transform)\ncombined_foraminal_train_dataset = CustomConcatDataset([foraminal_train_dataset, augmented_foraminal_train_dataset])\ncombined_foraminal_val_dataset = CustomConcatDataset([foraminal_val_dataset, augmented_foraminal_val_dataset])\n\nclass_weights = torch.tensor([1, 2, 4])\ntrain_labels = combined_foraminal_train_dataset.get_labels()\nlabels_list = []\nfor i in train_labels:\n    labels_list.extend(i)\nsamples_weights = class_weights[labels_list]\nsampler = WeightedRandomSampler(weights=samples_weights, \n                                num_samples=len(samples_weights), \n                                replacement=True)\n\nforaminal_train_loader = DataLoader(\n    combined_foraminal_train_dataset,\n    batch_size=32,\n    sampler=sampler\n)\n\nval_labels = combined_foraminal_val_dataset.get_labels()\nlabels_list = []\nfor i in val_labels:\n    labels_list.extend(i)\nsamples_weights = class_weights[labels_list]\nsampler = WeightedRandomSampler(weights=samples_weights, \n                                num_samples=len(samples_weights), \n                                replacement=True)\n\nforaminal_val_loader = DataLoader(\n    combined_foraminal_val_dataset,\n    batch_size=32,\n    sampler=sampler\n)","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:20.076283Z","iopub.execute_input":"2024-06-21T07:52:20.076650Z","iopub.status.idle":"2024-06-21T07:52:20.111163Z","shell.execute_reply.started":"2024-06-21T07:52:20.076617Z","shell.execute_reply":"2024-06-21T07:52:20.110437Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sabarticular_train_data, sabarticular_val_data = train_test_split(sabarticular_df_train, test_size=0.2, random_state=42)\n\nsabarticular_train_dataset = SpineDataset(sabarticular_train_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=train_transform)\nsabarticular_val_dataset = SpineDataset(sabarticular_val_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=val_transform)\n\nmoderate_severe_sabarticular_train_data = sabarticular_train_data[sabarticular_train_data['label'] != 'Normal/Mild']\nmoderate_severe_sabarticular_val_data = sabarticular_val_data[sabarticular_val_data['label'] != 'Normal/Mild']\naugmented_sabarticular_train_dataset = AugmentedSpineDataset(moderate_severe_sabarticular_train_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=augmented_transform)\naugmented_sabarticular_val_dataset = AugmentedSpineDataset(moderate_severe_sabarticular_val_data, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=val_transform)\ncombined_sabarticular_train_dataset = CustomConcatDataset([sabarticular_train_dataset, augmented_sabarticular_train_dataset])\ncombined_sabarticular_val_dataset = CustomConcatDataset([sabarticular_val_dataset, augmented_sabarticular_val_dataset])\n\nclass_weights = torch.tensor([1, 2, 4])\ntrain_labels = combined_sabarticular_train_dataset.get_labels()\nlabels_list = []\nfor i in train_labels:\n    labels_list.extend(i)\nsamples_weights = class_weights[labels_list]\nsampler = WeightedRandomSampler(weights=samples_weights, \n                                num_samples=len(samples_weights), \n                                replacement=True)\n\nsabarticular_train_loader = DataLoader(\n    combined_sabarticular_train_dataset,\n    batch_size=32,\n    sampler=sampler\n)\n\nval_labels = combined_sabarticular_val_dataset.get_labels()\nlabels_list = []\nfor i in val_labels:\n    labels_list.extend(i)\nsamples_weights = class_weights[labels_list]\nsampler = WeightedRandomSampler(weights=samples_weights, \n                                num_samples=len(samples_weights), \n                                replacement=True)\n\nsabarticular_val_loader = DataLoader(\n    combined_sabarticular_val_dataset,\n    batch_size=32,\n    sampler=sampler\n)","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:20.112310Z","iopub.execute_input":"2024-06-21T07:52:20.112596Z","iopub.status.idle":"2024-06-21T07:52:20.147012Z","shell.execute_reply.started":"2024-06-21T07:52:20.112573Z","shell.execute_reply":"2024-06-21T07:52:20.146018Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Modeling","metadata":{}},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nprint(\"Device: \", device)\nnum_classes = 3  # Normal/Mild, Moderate, Severe\n\nspinal_model = models.efficientnet_v2_s(pretrained=True)\nspinal_model.fc = nn.Sequential(\n    nn.Linear(spinal_model.classifier[-1].in_features, 512),\n    nn.ReLU(),\n    nn.Dropout(0.5),\n    nn.Linear(512, num_classes)\n)\n\n# spinal_model.fc = nn.Linear(spinal_model.classifier[-1].in_features, num_classes)\n# spinal_model = nn.DataParallel(spinal_model)\nspinal_model.to(device)\nspinal_optimizer = torch.optim.AdamW(spinal_model.parameters(), lr=1e-4, weight_decay=1e-4)\nspinal_scheduler = torch.optim.lr_scheduler.StepLR(spinal_optimizer, step_size=1, gamma=0.2)\nspinal_loss_fn = nn.CrossEntropyLoss()\n\nforaminal_model = models.efficientnet_v2_s(pretrained=True)\nforaminal_model.fc = nn.Sequential(\n    nn.Linear(foraminal_model.classifier[-1].in_features, 512),\n    nn.ReLU(),\n    nn.Dropout(0.5),\n    nn.Linear(512, num_classes)\n)\n# foraminal_model.fc = nn.Linear(foraminal_model.classifier[-1].in_features, num_classes)\n# foraminal_model = nn.DataParallel(foraminal_model)\nforaminal_model.to(device)\nforaminal_optimizer = torch.optim.AdamW(foraminal_model.parameters(), lr=1e-4, weight_decay=1e-4)\nforaminal_scheduler = torch.optim.lr_scheduler.StepLR(foraminal_optimizer, step_size=1, gamma=0.2)\nforaminal_loss_fn = nn.CrossEntropyLoss()\n\n\nsabarticular_model = models.efficientnet_v2_s(pretrained=True)\nsabarticular_model.fc = nn.Sequential(\n    nn.Linear(sabarticular_model.classifier[-1].in_features, 512),\n    nn.ReLU(),\n    nn.Dropout(0.5),\n    nn.Linear(512, num_classes)\n)\n# sabarticular_model.fc = nn.Linear(sabarticular_model.classifier[-1].in_features, num_classes)\n# sabarticular_model = nn.DataParallel(sabarticular_model)\nsabarticular_model.to(device)\nsabarticular_optimizer = torch.optim.AdamW(sabarticular_model.parameters(), lr=1e-4, weight_decay=1e-4)\nsabarticular_scheduler = torch.optim.lr_scheduler.StepLR(sabarticular_optimizer, step_size=1, gamma=0.2)\nsabarticular_loss_fn = nn.CrossEntropyLoss()","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:20.148180Z","iopub.execute_input":"2024-06-21T07:52:20.148517Z","iopub.status.idle":"2024-06-21T07:52:23.734845Z","shell.execute_reply.started":"2024-06-21T07:52:20.148474Z","shell.execute_reply":"2024-06-21T07:52:23.734024Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Optionally freeze initial layers\nfor param in spinal_model.parameters():\n    param.requires_grad = False\nfor param in foraminal_model.parameters():\n    param.requires_grad = False\nfor param in sabarticular_model.parameters():\n    param.requires_grad = False\n\n# Unfreeze the final fully connected layer\n# for param in spinal_model.fc.parameters():\n#     param.requires_grad = True\n# for param in foraminal_model.fc.parameters():\n#     param.requires_grad = True\n# for param in sabarticular_model.fc.parameters():\n#     param.requires_grad = True\n\nfor layer in list(spinal_model.features.children())[-1:]:\n    if not isinstance(layer, nn.BatchNorm2d):\n        for param in layer.parameters():\n            param.requires_grad = True\n\nfor param in spinal_model.classifier.parameters():\n    param.requires_grad = True\n\n    \nfor layer in list(foraminal_model.features.children())[-1:]:\n    if not isinstance(layer, nn.BatchNorm2d):\n        for param in layer.parameters():\n            param.requires_grad = True\n\nfor param in foraminal_model.classifier.parameters():\n    param.requires_grad = True\n    \n    \nfor layer in list(sabarticular_model.features.children())[-1:]:\n    if not isinstance(layer, nn.BatchNorm2d):\n        for param in layer.parameters():\n            param.requires_grad = True\n\nfor param in sabarticular_model.classifier.parameters():\n    param.requires_grad = True","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:23.735952Z","iopub.execute_input":"2024-06-21T07:52:23.736253Z","iopub.status.idle":"2024-06-21T07:52:23.759642Z","shell.execute_reply.started":"2024-06-21T07:52:23.736227Z","shell.execute_reply":"2024-06-21T07:52:23.758647Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 10\n\nbest_val_loss = float('inf')\npatience, trials = 2, 0\n\nfor epoch in tqdm(range(num_epochs)):\n    spinal_model.train()\n    running_loss = 0.0\n    correct_predictions = 0\n    total_samples = 0\n\n    for batch in spinal_train_loader:\n        images = batch['image'].to(device)\n        labels = batch['label'].to(device)\n        \n        spinal_optimizer.zero_grad()\n        outputs = spinal_model(images)\n        loss = spinal_loss_fn(outputs, labels)\n        loss.backward()\n        spinal_optimizer.step()\n\n        running_loss += loss.item()\n        _, predicted = torch.max(outputs, 1)\n        correct_predictions += (predicted == labels).sum().item()\n        total_samples += labels.size(0)\n\n    train_loss = running_loss / len(spinal_train_loader)\n    train_accuracy = correct_predictions / total_samples\n\n    # Validation loop\n    spinal_model.eval()\n    val_running_loss = 0.0\n    val_correct_predictions = 0\n    val_total_samples = 0\n\n    with torch.no_grad():\n        for batch in spinal_val_loader:\n            images = batch['image'].to(device)\n            labels = batch['label'].to(device)\n\n            outputs = spinal_model(images)\n            loss = spinal_loss_fn(outputs, labels)\n\n            val_running_loss += loss.item()\n            _, predicted = torch.max(outputs, 1)\n            val_correct_predictions += (predicted == labels).sum().item()\n            val_total_samples += labels.size(0)\n\n    val_loss = val_running_loss / len(spinal_val_loader)\n    val_accuracy = val_correct_predictions / val_total_samples\n\n    print(f\"Epoch {epoch+1}/{num_epochs}, \"\n          f\"Train Loss: {train_loss:.4f}, Train Accuracy: {train_accuracy:.4f}, \"\n          f\"Val Loss: {val_loss:.4f}, Val Accuracy: {val_accuracy:.4f}\")\n\n    spinal_scheduler.step()\n\n    # Check for early stopping\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        trials = 0\n        torch.save(spinal_model.state_dict(), 'spinal_best_spine_condition_classifier.pth')\n    else:\n        trials += 1\n        if trials >= patience:\n            print('Early stopping')\n            break","metadata":{"execution":{"iopub.status.busy":"2024-06-21T07:52:23.762808Z","iopub.execute_input":"2024-06-21T07:52:23.763115Z","iopub.status.idle":"2024-06-21T08:37:07.618336Z","shell.execute_reply.started":"2024-06-21T07:52:23.763090Z","shell.execute_reply":"2024-06-21T08:37:07.617375Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 10\n\nbest_val_loss = float('inf')\npatience, trials = 2, 0\n\nfor epoch in tqdm(range(num_epochs)):\n    foraminal_model.train()\n    running_loss = 0.0\n    correct_predictions = 0\n    total_samples = 0\n\n    for batch in foraminal_train_loader:\n        images = batch['image'].to(device)\n        labels = batch['label'].to(device)\n        \n        foraminal_optimizer.zero_grad()\n        outputs = foraminal_model(images)\n        loss = foraminal_loss_fn(outputs, labels)\n        loss.backward()\n        foraminal_optimizer.step()\n\n        running_loss += loss.item()\n        _, predicted = torch.max(outputs, 1)\n        correct_predictions += (predicted == labels).sum().item()\n        total_samples += labels.size(0)\n\n    train_loss = running_loss / len(foraminal_train_loader)\n    train_accuracy = correct_predictions / total_samples\n\n    # Validation loop\n    foraminal_model.eval()\n    val_running_loss = 0.0\n    val_correct_predictions = 0\n    val_total_samples = 0\n\n    with torch.no_grad():\n        for batch in foraminal_val_loader:\n            images = batch['image'].to(device)\n            labels = batch['label'].to(device)\n\n            outputs = foraminal_model(images)\n            loss = foraminal_loss_fn(outputs, labels)\n\n            val_running_loss += loss.item()\n            _, predicted = torch.max(outputs, 1)\n            val_correct_predictions += (predicted == labels).sum().item()\n            val_total_samples += labels.size(0)\n\n    val_loss = val_running_loss / len(foraminal_val_loader)\n    val_accuracy = val_correct_predictions / val_total_samples\n\n    print(f\"Epoch {epoch+1}/{num_epochs}, \"\n          f\"Train Loss: {train_loss:.4f}, Train Accuracy: {train_accuracy:.4f}, \"\n          f\"Val Loss: {val_loss:.4f}, Val Accuracy: {val_accuracy:.4f}\")\n\n    foraminal_scheduler.step()\n\n    # Check for early stopping\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        trials = 0\n        torch.save(foraminal_model.state_dict(), 'foraminal_best_spine_condition_classifier.pth')\n    else:\n        trials += 1\n        if trials >= patience:\n            print('Early stopping')\n            break","metadata":{"execution":{"iopub.status.busy":"2024-06-21T08:37:07.619522Z","iopub.execute_input":"2024-06-21T08:37:07.619814Z","iopub.status.idle":"2024-06-21T09:37:32.429397Z","shell.execute_reply.started":"2024-06-21T08:37:07.619788Z","shell.execute_reply":"2024-06-21T09:37:32.428429Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"num_epochs = 10\n\nbest_val_loss = float('inf')\npatience, trials = 2, 0\n\nfor epoch in tqdm(range(num_epochs)):\n    sabarticular_model.train()\n    running_loss = 0.0\n    correct_predictions = 0\n    total_samples = 0\n\n    for batch in sabarticular_train_loader:\n        images = batch['image'].to(device)\n        labels = batch['label'].to(device)\n        \n        sabarticular_optimizer.zero_grad()\n        outputs = sabarticular_model(images)\n        loss = sabarticular_loss_fn(outputs, labels)\n        loss.backward()\n        sabarticular_optimizer.step()\n\n        running_loss += loss.item()\n        _, predicted = torch.max(outputs, 1)\n        correct_predictions += (predicted == labels).sum().item()\n        total_samples += labels.size(0)\n\n    train_loss = running_loss / len(sabarticular_train_loader)\n    train_accuracy = correct_predictions / total_samples\n\n    # Validation loop\n    sabarticular_model.eval()\n    val_running_loss = 0.0\n    val_correct_predictions = 0\n    val_total_samples = 0\n\n    with torch.no_grad():\n        for batch in sabarticular_val_loader:\n            images = batch['image'].to(device)\n            labels = batch['label'].to(device)\n\n            outputs = sabarticular_model(images)\n            loss = sabarticular_loss_fn(outputs, labels)\n\n            val_running_loss += loss.item()\n            _, predicted = torch.max(outputs, 1)\n            val_correct_predictions += (predicted == labels).sum().item()\n            val_total_samples += labels.size(0)\n\n    val_loss = val_running_loss / len(sabarticular_val_loader)\n    val_accuracy = val_correct_predictions / val_total_samples\n\n    print(f\"Epoch {epoch+1}/{num_epochs}, \"\n          f\"Train Loss: {train_loss:.4f}, Train Accuracy: {train_accuracy:.4f}, \"\n          f\"Val Loss: {val_loss:.4f}, Val Accuracy: {val_accuracy:.4f}\")\n\n    sabarticular_scheduler.step()\n\n    # Check for early stopping\n    if val_loss < best_val_loss:\n        best_val_loss = val_loss\n        trials = 0\n        torch.save(sabarticular_model.state_dict(), 'sabarticular_best_spine_condition_classifier.pth')\n    else:\n        trials += 1\n        if trials >= patience:\n            print('Early stopping')\n            break","metadata":{"execution":{"iopub.status.busy":"2024-06-21T09:37:32.430770Z","iopub.execute_input":"2024-06-21T09:37:32.431143Z","iopub.status.idle":"2024-06-21T10:35:35.835061Z","shell.execute_reply.started":"2024-06-21T09:37:32.431101Z","shell.execute_reply":"2024-06-21T10:35:35.834152Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def train(model, train_loader, criterion, optimizer):\n#     model.train()\n#     running_loss = 0.0\n#     for inputs, labels in train_loader:\n#         optimizer.zero_grad()\n#         outputs = model(inputs)\n#         loss = criterion(outputs, labels)\n#         loss.backward()\n#         optimizer.step()\n#         running_loss += loss.item()\n#     return running_loss / len(train_loader)\n\n# def evaluate(model, val_loader, criterion):\n#     model.eval()\n#     val_loss = 0.0\n#     with torch.no_grad():\n#         for inputs, labels in val_loader:\n#             outputs = model(inputs)\n#             loss = criterion(outputs, labels)\n#             val_loss += loss.item()\n#     return val_loss / len(val_loader)","metadata":{"execution":{"iopub.status.busy":"2024-06-21T10:35:35.836430Z","iopub.execute_input":"2024-06-21T10:35:35.836799Z","iopub.status.idle":"2024-06-21T10:35:35.842213Z","shell.execute_reply.started":"2024-06-21T10:35:35.836767Z","shell.execute_reply":"2024-06-21T10:35:35.841338Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# k_folds = 5\n# num_epochs = 10\n# kf = KFold(n_splits=k_folds, shuffle=True)\n\n# # Cross-validation loop\n# results = {}\n\n# spinal_train_dataset = SpineDataset(spinal_df_train, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images', transform=train_transform)\n\n# for fold, (train_idx, val_idx) in enumerate(kf.split(sabarticular_train_dataset)):\n#     print(f'FOLD {fold}')\n#     print('--------------------------------')\n\n#     # Sample elements randomly from a given list of indices, no replacement.\n#     train_subset = Subset(spinal_train_dataset, train_idx)\n#     val_subset = Subset(spinal_train_dataset, val_idx)\n    \n#     class_weights = torch.tensor([1, 2, 4])\n#     train_labels = spinal_train_dataset[train_idx].get_labels()\n#     labels_list = []\n#     for i in train_labels:\n#         labels_list.extend(i)\n#     samples_weights = class_weights[labels_list]\n#     sampler = WeightedRandomSampler(weights=samples_weights, \n#                                     num_samples=len(samples_weights), \n#                                     replacement=True)\n\n#     spinal_train_loader = DataLoader(\n#         train_subset,\n#         batch_size=32,\n#         sampler=sampler\n#     )\n    \n#     val_labels = spinal_train_dataset[val_idx].get_labels()\n#     labels_list = []\n#     for i in val_labels:\n#         labels_list.extend(i)\n#     samples_weights = class_weights[labels_list]\n#     sampler = WeightedRandomSampler(weights=samples_weights, \n#                                     num_samples=len(samples_weights), \n#                                     replacement=True)\n\n#     spinal_val_loader = DataLoader(\n#         val_subset,\n#         batch_size=32,\n#         sampler=sampler\n#     )\n    \n\n#     # Run the training loop for defined number of epochs\n#     for epoch in range(num_epochs):\n#         train_loss = train(spinal_model, spinal_train_loader, spinal_criterion, spinal_optimizer)\n#         val_loss = evaluate(spinal_model, spinal_val_loader, spinal_criterion)\n#         print(f'Epoch {epoch+1}/{num_epochs}, Train Loss: {train_loss}, Validation Loss: {val_loss}')\n    \n#     # Save the model if needed or keep track of validation metrics\n#     results[fold] = val_loss\n\n# print(f'K-FOLD CROSS VALIDATION RESULTS FOR {k_folds} FOLDS')\n# print('--------------------------------')\n# sum_val_loss = 0.0\n# for key, value in results.items():\n#     print(f'Fold {key}: Validation Loss: {value}')\n#     sum_val_loss += value\n# print(f'Average Validation Loss: {sum_val_loss / len(results)}')","metadata":{"execution":{"iopub.status.busy":"2024-06-21T10:35:35.843874Z","iopub.execute_input":"2024-06-21T10:35:35.844236Z","iopub.status.idle":"2024-06-21T10:35:35.852213Z","shell.execute_reply.started":"2024-06-21T10:35:35.844203Z","shell.execute_reply":"2024-06-21T10:35:35.851431Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# torch.save(sabarticular_model.state_dict(), 'sabarticular_best_spine_condition_classifier.pth')\n# torch.save(foraminal_model.state_dict(), 'foraminal_best_spine_condition_classifier.pth')\n# torch.save(spinal_model.state_dict(), 'spinal_best_spine_condition_classifier.pth')","metadata":{"execution":{"iopub.status.busy":"2024-06-21T10:35:35.853276Z","iopub.execute_input":"2024-06-21T10:35:35.853564Z","iopub.status.idle":"2024-06-21T10:35:35.862058Z","shell.execute_reply.started":"2024-06-21T10:35:35.853533Z","shell.execute_reply":"2024-06-21T10:35:35.861345Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}