{"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"},{"sourceId":7681872,"sourceType":"datasetVersion","datasetId":4481958},{"sourceId":9547858,"sourceType":"datasetVersion","datasetId":5817207},{"sourceId":9548238,"sourceType":"datasetVersion","datasetId":5812788},{"sourceId":6158,"sourceType":"modelInstanceVersion","modelInstanceId":4608,"modelId":2797},{"sourceId":64765,"sourceType":"modelInstanceVersion","modelInstanceId":54020,"modelId":74163},{"sourceId":64795,"sourceType":"modelInstanceVersion","modelInstanceId":54048,"modelId":74163}],"dockerImageVersionId":30733,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## Introduction","metadata":{}},{"cell_type":"markdown","source":"The first challenge in predicting Left Subarticular Stenosis and Right Subarticular Stenosis is to find the appropriate images in the axial plane from the MRI stack at the intervertebral disk levels (five levels in the lumbar spine).\n\nThe training images have been annotated, so identifying the images at the aforementioned five levels is straightforward.\n\nFor a test set, finding the images corresponding to the five levels is no longer straightforward. I have trained a EfficientNet B0 model in PyTorch, with the task to identify five images from the MRI stack which correspond to the five intervertebral disk levels in the lumbar spine.\n\nBecause the anatomy of the spine at the L5/S1 is considerably different from the anatomy at the other four lumbar levels, I have chosen a classifier with three output classes (L5/S1 level, other lumbar IVD level, no IVD level).","metadata":{}},{"cell_type":"code","source":" ###### Set hyperparameters\nhp_model_num_classes = {'ivd': 3, 'sub': 3}\nhp_model_types = hp_model_num_classes.keys()\nhp_new_training = {'ivd': False, 'sub': False}\nhp_number_training_cases = {'ivd': 500}\nhp_learning_rate = {'ivd': 0.005, 'sub': 0.02}\nhp_num_epochs = {'ivd': 10, 'sub': 15}\nhp_apply_gamma = True\nhp_required_img_size = 320\nhp_use_test_images = False\nhp_train_base_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/'\nhp_pretrained_model = {'ivd': 'rsna-trained-values/ivd_best_model_4.pth', 'sub': 'sub-model/sub_best_model_6.pth'}\nseverity_mapping = dict (zip ([ \"Normal/Mild\", \"Moderate\", \"Severe\"], [0, 1, 2]))\nprint (severity_mapping)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from copy import deepcopy\nimport seaborn as sns\n\nimport matplotlib.pyplot as plt\nimport os\nimport time\nimport numpy as np\nimport glob\nimport json\nimport collections\nimport torch\nimport torch.nn as nn\nimport torchvision.models as models\nfrom torchvision import transforms\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom tqdm import tqdm\nfrom PIL import Image  # Importing Image module from PIL (Python Imaging Library) to work with images.\nimport pydicom as dicom\nimport matplotlib.patches as patches\nfrom matplotlib import animation, rc\nimport pandas as pd\n\nfrom scipy import ndimage\nimport pydicom as dicom # dicom\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# read data\ntrain_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\n\ntrain  = pd.read_csv(train_path + 'train.csv')\nlabel = pd.read_csv(train_path + 'train_label_coordinates.csv')\ntrain_desc  = pd.read_csv(train_path + 'train_series_descriptions.csv')\ntest_desc   = pd.read_csv(train_path + 'test_series_descriptions.csv')\nsub         = pd.read_csv(train_path + 'sample_submission.csv')\n\ntrain.head (5)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.537855Z","iopub.execute_input":"2025-01-23T14:31:24.538182Z","iopub.status.idle":"2025-01-23T14:31:24.862877Z","shell.execute_reply.started":"2025-01-23T14:31:24.538154Z","shell.execute_reply":"2025-01-23T14:31:24.857633Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to generate image paths based on directory structure\ndef generate_image_paths(df, data_dir):\n    image_paths = []\n    for study_id, series_id in zip(df['study_id'], df['series_id']):\n        study_dir = os.path.join(data_dir, str(study_id))\n        series_dir = os.path.join(study_dir, str(series_id))\n        images = os.listdir(series_dir)\n        image_paths.extend([os.path.join(series_dir, img) for img in images])\n    return image_paths\n\n# Generate image paths for train and test data\ntrain_image_paths = generate_image_paths(train_desc, f'{train_path}/train_images')\ntest_image_paths = generate_image_paths(test_desc, f'{train_path}/test_images')","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.864061Z","iopub.status.idle":"2025-01-23T14:31:24.864525Z","shell.execute_reply.started":"2025-01-23T14:31:24.864288Z","shell.execute_reply":"2025-01-23T14:31:24.864305Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to display images\ndef display_images(images, title, max_images_per_row=4):\n    # Calculate the number of rows needed\n    num_images = len(images)\n    num_rows = (num_images + max_images_per_row - 1) // max_images_per_row  # Ceiling division\n\n    # Create a subplot grid\n    fig, axes = plt.subplots(num_rows, max_images_per_row, figsize=(10, 3 * num_rows))\n    \n    # Flatten axes array for easier looping if there are multiple rows\n    if num_rows > 1:\n        axes = axes.flatten()\n    else:\n        axes = [axes]  # Make it iterable for consistency\n\n    # Plot each image\n    for idx, image in enumerate(images):\n        ax = axes[idx]\n        ax.imshow(image, cmap=plt.cm.bone)\n        # cmap='gray')  # Assuming grayscale for simplicity, change cmap as needed\n        ax.axis('off')  # Hide axes\n\n    # Turn off unused subplots\n    for idx in range(num_images, len(axes)):\n        axes[idx].axis('off')\n    fig.suptitle(title, fontsize=16)\n\n    plt.tight_layout()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T14:31:24.868535Z","iopub.status.idle":"2025-01-23T14:31:24.869688Z","shell.execute_reply.started":"2025-01-23T14:31:24.869383Z","shell.execute_reply":"2025-01-23T14:31:24.869432Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def axial_image_crop1 (image,  size, ref_size):\n    img_width = image.shape[1]\n    img_height = image.shape[0]\n    if (img_height < ref_size) or (img_width < ref_size):\n        # Create PIL Image object from numpy array\n        image_new = Image.fromarray(image, 'L')\n        image_new = image_new.resize((size, size), Image.Resampling.LANCZOS)\n        return image_new\n    else:\n        crop_size = int ((size / ref_size) * min (img_height, img_width))\n        top = img_height - crop_size\n        left = int ((img_width - crop_size) / 2)\n        bottom = top + crop_size\n        right = left + crop_size\n        IMG_cropped = image [int(top):int(bottom), int(left):int(right)]\n\n        # Create PIL Image object from numpy array\n        image_new = Image.fromarray(IMG_cropped, 'L')\n        image_new = image_new.resize((size, size), Image.Resampling.LANCZOS)\n        return image_new","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.871014Z","iopub.status.idle":"2025-01-23T14:31:24.871490Z","shell.execute_reply.started":"2025-01-23T14:31:24.871256Z","shell.execute_reply":"2025-01-23T14:31:24.871274Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Preprocessing","metadata":{}},{"cell_type":"code","source":"# Define function to reshape a single row of the DataFrame\ndef reshape_row(row):\n    data = {'study_id': [], 'condition': [], 'level': [], 'severity': []}\n    \n    for column, value in row.items():\n        if column not in ['study_id', 'series_id', 'instance_number', 'x', 'y', 'series_description']:\n            parts = column.split('_')\n            condition = ' '.join([word.capitalize() for word in parts[:-2]])\n            level = parts[-2].capitalize() + '/' + parts[-1].capitalize()\n            data['study_id'].append(row['study_id'])\n            data['condition'].append(condition)\n            data['level'].append(level)\n            data['severity'].append(value)\n    \n    return pd.DataFrame(data)\n\n# Reshape the DataFrame for all rows\nnew_train_df = pd.concat([reshape_row(row) for _, row in train.iterrows()], ignore_index=True)\n\n# Display the first few rows of the reshaped dataframe\nnew_train_df.head(25)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.874279Z","iopub.status.idle":"2025-01-23T14:31:24.875014Z","shell.execute_reply.started":"2025-01-23T14:31:24.874653Z","shell.execute_reply":"2025-01-23T14:31:24.874679Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Print columns in a neat way\nprint(\"\\nColumns in new_train_df:\")\nprint(\",\".join(new_train_df.columns))\n\nprint(\"\\nColumns in label:\")\nprint(\",\".join(label.columns))\n\nprint(\"\\nColumns in test_desc:\")\nprint(\",\".join(test_desc.columns))\n\nprint(\"\\nColumns in sub:\")\nprint(\",\".join(sub.columns))","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.877218Z","iopub.status.idle":"2025-01-23T14:31:24.877868Z","shell.execute_reply.started":"2025-01-23T14:31:24.877553Z","shell.execute_reply":"2025-01-23T14:31:24.877581Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"IVD_levels_desc = list (new_train_df[ 'level'].unique())\nIVD_levels_desc","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.880936Z","iopub.status.idle":"2025-01-23T14:31:24.881584Z","shell.execute_reply.started":"2025-01-23T14:31:24.881254Z","shell.execute_reply":"2025-01-23T14:31:24.881281Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Merge the dataframes on the common columns\nmerged_df = pd.merge(new_train_df, label, on=['study_id', 'condition', 'level'], how='inner')","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.884216Z","iopub.status.idle":"2025-01-23T14:31:24.884718Z","shell.execute_reply.started":"2025-01-23T14:31:24.884515Z","shell.execute_reply":"2025-01-23T14:31:24.884535Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"selection_left = merged_df[merged_df['condition'] == 'Left Subarticular Stenosis']\nselection_left.head (5)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.886051Z","iopub.status.idle":"2025-01-23T14:31:24.886633Z","shell.execute_reply.started":"2025-01-23T14:31:24.886318Z","shell.execute_reply":"2025-01-23T14:31:24.886341Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"selection_right = merged_df[merged_df['condition'] == 'Right Subarticular Stenosis']\nselection_right.head (5)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.889078Z","iopub.status.idle":"2025-01-23T14:31:24.889685Z","shell.execute_reply.started":"2025-01-23T14:31:24.889356Z","shell.execute_reply":"2025-01-23T14:31:24.889380Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# ivd_selection_df = pd.merge (selection_left, selection_right, on=['study_id', 'level', 'severity'], how='inner')\nivd_selection_df = pd.concat([ selection_left, selection_right ])\nivd_selection_df = ivd_selection_df.sort_values(by=['study_id', 'level', 'condition'])\nlength_before = len (ivd_selection_df)\nivd_selection_df = ivd_selection_df.dropna()\nprint (length_before, len (ivd_selection_df))\nivd_selection_df.head (10)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.893184Z","iopub.status.idle":"2025-01-23T14:31:24.893732Z","shell.execute_reply.started":"2025-01-23T14:31:24.893516Z","shell.execute_reply":"2025-01-23T14:31:24.893536Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for sev in list (severity_mapping.keys()):\n    count = len (ivd_selection_df[ivd_selection_df['severity'] == sev])\n    print (f\"Number of occurrences of {sev}: {count}\")","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.895451Z","iopub.status.idle":"2025-01-23T14:31:24.896058Z","shell.execute_reply.started":"2025-01-23T14:31:24.895742Z","shell.execute_reply":"2025-01-23T14:31:24.895766Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_paths = []\nfor index, row in ivd_selection_df.iterrows():\n    image_path_short = str(row['study_id']) + '/' + str (row['series_id']) + '/' + str(row['instance_number']) + '.dcm'\n    image_paths.append(image_path_short)\nivd_selection_df['image_path_short'] = image_paths\nivd_selection_df.head (6)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.897752Z","iopub.status.idle":"2025-01-23T14:31:24.898343Z","shell.execute_reply.started":"2025-01-23T14:31:24.898057Z","shell.execute_reply":"2025-01-23T14:31:24.898081Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_ivd_df = ivd_selection_df [['study_id', 'condition', 'series_id', 'level', 'instance_number']]\ntrain_ivd_df.head (5)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.901178Z","iopub.status.idle":"2025-01-23T14:31:24.901630Z","shell.execute_reply.started":"2025-01-23T14:31:24.901436Z","shell.execute_reply":"2025-01-23T14:31:24.901458Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_study_series = train_ivd_df [['study_id', 'series_id']].drop_duplicates()\ntrain_study_series = train_study_series.head (hp_number_training_cases['ivd'])\ntrain_study_series.head (5)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.903637Z","iopub.status.idle":"2025-01-23T14:31:24.904036Z","shell.execute_reply.started":"2025-01-23T14:31:24.903859Z","shell.execute_reply":"2025-01-23T14:31:24.903875Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to get image paths for a series\ndef get_image_filenames (base_path, study_id, series_id):\n    series_path = os.path.join(base_path, str(study_id), str(series_id))\n    if os.path.exists(series_path):\n        files_list = [f for f in os.listdir(series_path) if os.path.isfile(os.path.join(series_path, f))]\n        files_list.sort(key=lambda f: int(f.split('.')[0]))\n        return files_list\n    return []","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.905075Z","iopub.status.idle":"2025-01-23T14:31:24.905521Z","shell.execute_reply.started":"2025-01-23T14:31:24.905281Z","shell.execute_reply":"2025-01-23T14:31:24.905298Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_base_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/'\nprint (get_image_filenames (train_base_path, 208289456, 4140564625))","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.908217Z","iopub.status.idle":"2025-01-23T14:31:24.908832Z","shell.execute_reply.started":"2025-01-23T14:31:24.908527Z","shell.execute_reply":"2025-01-23T14:31:24.908554Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"###### Function to get image paths for a series\ndef get_file_count (base_path, study_id, series_id):\n    file_count = 0\n    series_path = os.path.join(base_path, str(study_id), str(series_id))\n    if os.path.exists(series_path):\n        for path in os.scandir(series_path):\n            if path.is_file():\n                file_count += 1\n    return file_count","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.910746Z","iopub.status.idle":"2025-01-23T14:31:24.911909Z","shell.execute_reply.started":"2025-01-23T14:31:24.911319Z","shell.execute_reply":"2025-01-23T14:31:24.911360Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"###### \ndef create_ivd_training_data (study_series_df, selection_df, base_path):\n    images_labels = []\n    for index, row in study_series_df.iterrows():\n        study_id = row['study_id']\n        series_id = row['series_id']\n        instances_df = selection_df [(selection_df['study_id'] == study_id) & (selection_df['series_id'] == series_id)]\n        L5S1_instances_df = instances_df [(instances_df['level'] == 'L5/S1')]\n        L5S1_instances = list (L5S1_instances_df['instance_number'].drop_duplicates())\n        images_at_L_level = list (instances_df ['instance_number'].drop_duplicates())\n        \n        for filename in get_image_filenames (base_path, study_id, series_id):\n            image_path_short = str(study_id) + '/' + str(series_id) + '/' + filename\n            instance_number = int (filename.split('.')[0])\n            \n            # check for image resolution\n            ds = pydicom.dcmread(base_path + image_path_short)\n            if not ((ds.Rows < hp_required_img_size) or (ds.Columns < hp_required_img_size)):\n                if instance_number in L5S1_instances:\n                    label = 2\n                elif instance_number in images_at_L_level:\n                    label = 1\n                else:\n                    label = 0\n                images_labels.append ([image_path_short, label])\n    return images_labels","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.913593Z","iopub.status.idle":"2025-01-23T14:31:24.914205Z","shell.execute_reply.started":"2025-01-23T14:31:24.913901Z","shell.execute_reply":"2025-01-23T14:31:24.913927Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_images_labels = create_ivd_training_data (train_study_series, train_ivd_df, hp_train_base_path)\nivd_train_data = pd.DataFrame(train_images_labels, columns=['image_path_short', 'level_label'])\nivd_train_data.head (45)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.915369Z","iopub.status.idle":"2025-01-23T14:31:24.916002Z","shell.execute_reply.started":"2025-01-23T14:31:24.915695Z","shell.execute_reply":"2025-01-23T14:31:24.915736Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_dicom(path):\n    dicom = pydicom.read_file(path)\n    image_array = dicom.pixel_array\n\n    # Normalize pixel values to 0-255\n    image_array = ((image_array - image_array.min()) / (image_array.max() - image_array.min())) * 255.0\n    image_array = image_array.astype(np.uint8)\n\n    return image_array","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.917795Z","iopub.status.idle":"2025-01-23T14:31:24.918445Z","shell.execute_reply.started":"2025-01-23T14:31:24.918100Z","shell.execute_reply":"2025-01-23T14:31:24.918125Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def axial_image_crop2 (image,  size, ref_size):\n    # Remove the black band in the lower portion of the image\n    darkness_threshold = 25\n    m, n = image.shape[0], image.shape[1]\n    for j in range (m-1, 0, -1):\n        row_avg = np.average (image [j,:])\n        if row_avg > darkness_threshold:\n            lower_border = j\n            break\n\n    # Find the vertical symmetry axis (x-coordinate only)\n    com_y, com_x = ndimage.center_of_mass( image)\n    com_x = int (round (com_x, 0))\n    # Define crop window borders\n    img_width = int (n/2)\n    left_border = com_x - int (img_width/2)\n    right_border = left_border + img_width\n    top_border = lower_border - img_width\n    #image[:m, com_x:n] = 0\n    IMG_cropped = image[top_border:lower_border, left_border:right_border]\n\n    # Create PIL Image object from numpy array\n    image_new = Image.fromarray(IMG_cropped, 'L')\n    image_new = image_new.resize((size, size), Image.Resampling.LANCZOS)\n    return image_new","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.920793Z","iopub.status.idle":"2025-01-23T14:31:24.921437Z","shell.execute_reply.started":"2025-01-23T14:31:24.921118Z","shell.execute_reply":"2025-01-23T14:31:24.921143Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the transforms\ndef transform (image_array, crop_function):\n    # Calculate parameters for dynamic gamma correction first\n    mid = 0.5\n    mean = np.mean(image_array)\n    gamma = np.log(mid*255)/np.log(mean)\n    # Then crop and resize\n    image_PIL = crop_function (image_array, 224, hp_required_img_size)\n    # and apply gamma correction\n    if hp_apply_gamma:\n        image_PIL = transforms.functional.adjust_gamma (image_PIL, 1/gamma)\n    return image_PIL","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.923283Z","iopub.status.idle":"2025-01-23T14:31:24.923923Z","shell.execute_reply.started":"2025-01-23T14:31:24.923621Z","shell.execute_reply":"2025-01-23T14:31:24.923646Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loading data","metadata":{}},{"cell_type":"code","source":"# Define a custom dataset class\nclass CustomDataset(Dataset):\n    def __init__(self, dataframe, train_path, type_desc, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n        self.base_path = train_path\n        self.type_desc = type_desc\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, index):\n        image_path = self.base_path + self.dataframe['image_path_short'][index]\n        image = load_dicom(image_path)  # Define this function to load your DICOM images\n        if self.type_desc == 'ivd':\n            label = self.dataframe['level_label'][index]            \n            if self.transform:\n                image = self.transform(image, axial_image_crop1)\n        elif self.type_desc == 'sub':\n            if 'Right' in self.dataframe['condition'][index]:\n                # Flip the image around the vertical axis\n                image = np.fliplr (image)\n            label = severity_mapping[self.dataframe['severity'][index]]\n\n            if self.transform:\n                image = self.transform(image, axial_image_crop2)\n\n        trans2 = transforms.Compose([\n            transforms.Grayscale(num_output_channels=3),\n            transforms.ToTensor(),\n        ])\n        image = trans2 (image)\n        return image, label\n\n# Function to create datasets and dataloaders for each model type\ndef create_datasets_and_loaders(filtered_df, type_desc, transform, batch_size=8):\n    train_df, val_df = train_test_split(filtered_df, test_size=0.2, random_state=42)\n    train_df = train_df.reset_index(drop=True)\n    val_df = val_df.reset_index(drop=True)\n    train_dataset = CustomDataset(train_df, hp_train_base_path, type_desc, transform)\n    val_dataset = CustomDataset(val_df, train_base_path, type_desc, transform)\n\n    trainloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True)\n    valloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False)\n    \n    return trainloader, valloader, len(train_df), len(val_df)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.925296Z","iopub.status.idle":"2025-01-23T14:31:24.925898Z","shell.execute_reply.started":"2025-01-23T14:31:24.925621Z","shell.execute_reply":"2025-01-23T14:31:24.925645Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create dataloaders for each model type\ndict_train_dfs = {'ivd': ivd_train_data, 'sub': ivd_selection_df}\ntrain_loaders, val_loaders, train_lengths, val_lengths = {}, {}, {}, {}\nfor model_type in hp_model_types:\n    train_loaders[model_type], val_loaders[model_type], train_lengths[model_type], val_lengths[model_type]= create_datasets_and_loaders(dict_train_dfs[model_type], model_type, transform)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.928087Z","iopub.status.idle":"2025-01-23T14:31:24.928705Z","shell.execute_reply.started":"2025-01-23T14:31:24.928387Z","shell.execute_reply":"2025-01-23T14:31:24.928434Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images = []\nfor index, row in ivd_train_data.head(8).iterrows():\n    image_path_short = row['image_path_short']\n    image_path = train_base_path + image_path_short\n    image = load_dicom(image_path)\n    images.append(image)\ndisplay_images (images, \"IVD training data before transformation\")\nimages, labels = next(iter(train_loaders['ivd']))\nimages = [img.numpy()[0,:,:] for img in images]\ndisplay_images (images, \"IVD training data after transformation\")\nimages, labels = next(iter(train_loaders['sub']))\nimages = [img.numpy()[0,:,:] for img in images]\ndisplay_images (images, \"Subarticular training data\")\nprint (labels)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.930445Z","iopub.status.idle":"2025-01-23T14:31:24.931031Z","shell.execute_reply.started":"2025-01-23T14:31:24.930732Z","shell.execute_reply":"2025-01-23T14:31:24.930755Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Models","metadata":{}},{"cell_type":"code","source":"class CustomEfficientNetV2(nn.Module):\n    def __init__(self, type_desc, num_classes=3, pretrained_weights=None):\n        self.type_desc = type_desc\n        new_training = hp_new_training[type_desc]\n        super(CustomEfficientNetV2, self).__init__()\n        self.model = models.efficientnet_b0(weights=None)\n        if not new_training:\n            num_ftrs = self.model.classifier[-1].in_features\n            self.model.classifier[-1] = nn.Linear(num_ftrs, num_classes)\n        if pretrained_weights:\n            if torch.cuda.is_available():\n                pretrained_dict = torch.load(pretrained_weights)\n            else:\n                pretrained_dict = torch.load(pretrained_weights, map_location=torch.device('cpu'))\n            pretrained_dict = {key.replace(\"model.\", \"\"): value for key, value in pretrained_dict.items()}\n            self.model.load_state_dict(pretrained_dict)\n            \n        if new_training:\n            num_ftrs = self.model.classifier[-1].in_features\n            self.model.classifier[-1] = nn.Linear(num_ftrs, num_classes)\n            \n    def get_type(self):\n        return self.type_desc\n\n    def forward(self, x):\n        return self.model(x)\n\n    def unfreeze_model(self):\n        # Unfreeze the last 20 layers, keeping BatchNorm layers frozen\n        for layer in list(self.model.features.children())[-20:]:\n            if not isinstance(layer, nn.BatchNorm2d):\n                for param in layer.parameters():\n                    param.requires_grad = True\n        \n        # Unfreeze the classifier\n        for param in self.model.classifier.parameters():\n            param.requires_grad = True","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.932513Z","iopub.status.idle":"2025-01-23T14:31:24.933108Z","shell.execute_reply.started":"2025-01-23T14:31:24.932820Z","shell.execute_reply":"2025-01-23T14:31:24.932844Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize models\nmodels_model, models_optimisers = {}, {}\nfor model_type in hp_model_types:\n    # Path to the locally uploaded weights file\n    if hp_new_training[model_type]:\n        weights_path = '/kaggle/input/efficientnet-model-weights/efficientnet_b0_rwightman-7f5810bc.pth'\n    else:\n        weights_path = '/kaggle/input/' + hp_pretrained_model[model_type]\n    models_model[model_type] = CustomEfficientNetV2(model_type, hp_model_num_classes[model_type], pretrained_weights=weights_path).to(device)\n\n    # Unfreeze the final fully connected layer\n    for param in models_model[model_type].model.classifier.parameters():\n        param.requires_grad = True\n    # Initialize separate optimizers for each model\n    models_optimisers[model_type] = torch.optim.Adam(models_model[model_type].model.classifier.parameters(), lr=hp_learning_rate[model_type])\n    criterion = nn.CrossEntropyLoss()","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.935558Z","iopub.status.idle":"2025-01-23T14:31:24.936142Z","shell.execute_reply.started":"2025-01-23T14:31:24.935857Z","shell.execute_reply":"2025-01-23T14:31:24.935882Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Count trainable parameters\n#trainable_params = sum(p.numel() for p in ivd_model.parameters() if p.requires_grad)\n#print(f\"Number of parameters: {trainable_params}\")","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.937748Z","iopub.status.idle":"2025-01-23T14:31:24.938324Z","shell.execute_reply.started":"2025-01-23T14:31:24.938037Z","shell.execute_reply":"2025-01-23T14:31:24.938062Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training","metadata":{}},{"cell_type":"code","source":"def train_model(model, trainloader, valloader, len_train, len_val, optimizer, num_epochs=10, patience=3):\n    # Learning rate scheduler\n    scheduler = lr_scheduler.StepLR(optimizer, step_size=2, gamma=0.1)\n    type_desc = model.get_type()\n    best_val_acc = 0.0\n    best_model_wts = deepcopy(model.state_dict())\n    counter = 0\n    \n    for epoch in range(num_epochs):\n        model.train()\n        train_loss = 0\n        correct_train = 0\n        \n        with tqdm(trainloader, unit=\"batch\") as tepoch:\n            for images, labels in tepoch:\n                images, labels = images.to(device),  torch.tensor(labels).to(device)\n                optimizer.zero_grad()\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                loss.backward()\n                optimizer.step()\n                train_loss += loss.item()\n                \n                probabilities = torch.softmax(outputs, dim=1)\n                _, predicted = torch.max(probabilities, 1)\n                correct_train += (predicted == labels).sum().item()\n                \n                tepoch.set_postfix(epoch=epoch+1)\n        \n        scheduler.step()\n        \n        train_loss /= len(trainloader)\n        train_acc = 100 * correct_train / len_train\n        \n        model.eval()\n        val_loss, correct_val = 0, 0\n        with torch.no_grad():\n            with tqdm(valloader, unit=\"batch\") as vepoch:\n                for images, labels in vepoch:\n                    images, labels = images.to(device),  torch.tensor(labels).to(device)\n                    outputs = model(images)\n                    loss = criterion(outputs, labels)\n                    val_loss += loss.item()\n                    \n                    probabilities = torch.softmax(outputs, dim=1).squeeze(0)\n                    # A batch size of 1 might cause problems with the tensor dimensions\n                    if len (images) > 1:\n                        _, predicted = torch.max(probabilities, 1)\n                    else:\n                        _, predicted = torch.max(probabilities, 0)\n                    correct_val += (predicted == labels).sum().item()\n                    \n                    vepoch.set_postfix(epoch=epoch+1)\n        \n        val_loss /= len(valloader)\n        val_acc = 100 * correct_val / len_val\n        \n        print(f\"Epoch {epoch+1}, Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%, Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        \n        # Save the best model and check for early stopping\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            best_model_wts = deepcopy(model.state_dict())\n            counter = 0\n            torch.save(best_model_wts, f'{type_desc}_best_model_{epoch+1}.pth')\n        else:\n            counter += 1\n        \n        # Early stopping\n        if counter >= patience:\n            print(f\"Early stopping triggered after {epoch+1} epochs\")\n            break\n    \n    # Load best model weights\n    model.load_state_dict(best_model_wts)\n    return model, best_val_acc","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.940218Z","iopub.status.idle":"2025-01-23T14:31:24.940959Z","shell.execute_reply.started":"2025-01-23T14:31:24.940649Z","shell.execute_reply":"2025-01-23T14:31:24.940676Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"raw","source":"","metadata":{}},{"cell_type":"code","source":"for model_type in hp_model_types:\n    if hp_new_training[model_type]:\n        print (f\"Training started for model type {model_type}\")\n        train_model (models_model[model_type], train_loaders[model_type], val_loaders[model_type], train_lengths[model_type], val_lengths[model_type], models_optimisers[model_type], hp_num_epochs[model_type])\n\n## Inference","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.942690Z","iopub.status.idle":"2025-01-23T14:31:24.943281Z","shell.execute_reply.started":"2025-01-23T14:31:24.942974Z","shell.execute_reply":"2025-01-23T14:31:24.943000Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print (test_desc)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.946508Z","iopub.status.idle":"2025-01-23T14:31:24.947118Z","shell.execute_reply.started":"2025-01-23T14:31:24.946824Z","shell.execute_reply":"2025-01-23T14:31:24.946850Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_test_data (study_series_df, base_path):\n    image_paths = []\n    for index, row in study_series_df.iterrows():\n        study_id = row['study_id']\n        series_id = row['series_id']\n        \n        for filename in get_image_filenames (base_path, study_id, series_id):\n            image_path_short = str(study_id) + '/' + str(series_id) + '/' + filename\n            image_paths.append ([image_path_short])\n    return image_paths","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.949123Z","iopub.status.idle":"2025-01-23T14:31:24.949767Z","shell.execute_reply.started":"2025-01-23T14:31:24.949485Z","shell.execute_reply":"2025-01-23T14:31:24.949510Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define a custom test dataset class\nclass TestDataset(Dataset):\n    \n    def __init__(self, dataframe, test_path, transform=None):\n        self.dataframe = dataframe.reset_index()\n        self.transform = transform\n        self.base_path = test_path\n        \n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, index):\n        image_path = self.base_path + self.dataframe.loc[ index, 'image_path_short']\n        image = load_dicom(image_path)  # Define this function to load your DICOM images\n\n        if self.transform:\n            image = self.transform(image, axial_image_crop1)\n\n        trans2 = transforms.Compose([\n            transforms.Grayscale(num_output_channels=3),\n            transforms.ToTensor(),\n        ])\n        image = trans2 (image)\n        return image\n\n# Create a test dataset and dataloader\nif hp_use_test_images:\n    test_base_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/'\n    test_image_paths = create_test_data (test_desc[test_desc['series_description'] == 'Axial T2'], test_base_path)\nelse:\n    test_base_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images/'\n    test_image_paths = create_test_data (train_study_series.tail(1), test_base_path)\nivd_test_data = pd.DataFrame(test_image_paths, columns=['image_path_short'])\nprint (ivd_test_data.head(5))\nivd_test_dataset = TestDataset(ivd_test_data, test_base_path, transform)\nivd_testloader = DataLoader(ivd_test_dataset, batch_size=1, shuffle=False)","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.954681Z","iopub.status.idle":"2025-01-23T14:31:24.955296Z","shell.execute_reply.started":"2025-01-23T14:31:24.954992Z","shell.execute_reply":"2025-01-23T14:31:24.955018Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The first class is shown in orange, the second class in blue. As you can see, there are five clearly distinctive peaks for the probabilities in the graph. Once the appropriate images have been selected from the MRI stack, it is not complicated any more to classify the two conditions at each of the five number levels using another trained model.","metadata":{}},{"cell_type":"code","source":"# Function to make predictions on the test data\ndef predict_test_data(testloader, model):\n    predictions = []\n    \n    with torch.no_grad():\n        for idx, images in enumerate(tqdm(testloader)):\n            images = images.to(device)\n            if model:\n                model.eval()  # Set the model to eval mode\n                outputs = model(images)\n                probs = torch.softmax(outputs, dim=1).squeeze(0)\n                predictions.append(probs)\n            else:\n                predictions.append(None)\n    return predictions","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.957490Z","iopub.status.idle":"2025-01-23T14:31:24.958078Z","shell.execute_reply.started":"2025-01-23T14:31:24.957803Z","shell.execute_reply":"2025-01-23T14:31:24.957826Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Make predictions on the test data\nivd_test_predictions = predict_test_data(ivd_testloader, models_model['ivd'])","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.959538Z","iopub.status.idle":"2025-01-23T14:31:24.960114Z","shell.execute_reply.started":"2025-01-23T14:31:24.959820Z","shell.execute_reply":"2025-01-23T14:31:24.959843Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"probs_1 = [float (ivd_test_predictions[i][1]) for i in range (len(ivd_test_predictions))]\nprobs_2 = [float (ivd_test_predictions[i][2]) for i in range (len(ivd_test_predictions))]\nivd_test_results = ivd_test_data\nivd_test_results['probs_LL'] = probs_1\nivd_test_results['probs_LS'] = probs_2\nprint (ivd_test_results.head(5))","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.962231Z","iopub.status.idle":"2025-01-23T14:31:24.962732Z","shell.execute_reply.started":"2025-01-23T14:31:24.962507Z","shell.execute_reply":"2025-01-23T14:31:24.962525Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig = plt.figure(figsize = (10, 5))\nindices = [i for i in range (1, len (ivd_test_results['probs_LL']) + 1)]\n# creating the bar plot\nbarWidth = 0.2\nplt.bar(indices, ivd_test_results['probs_LL'], width = barWidth)\nplt.bar([i + barWidth*2 for i in indices], ivd_test_results['probs_LS'], width = barWidth)\nplt.xlabel(\"Axial image index number\")\nplt.ylabel(\"Probabilities sliding window average\")\nplt.title(f\"Test results for {hp_number_training_cases}\")\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.964226Z","iopub.status.idle":"2025-01-23T14:31:24.964704Z","shell.execute_reply.started":"2025-01-23T14:31:24.964478Z","shell.execute_reply":"2025-01-23T14:31:24.964496Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def find_highest_peak (probs):\n    max_value = max (probs)\n    for i in range (len (probs)):\n        if probs[i] >= max_value:\n            return i","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.966877Z","iopub.status.idle":"2025-01-23T14:31:24.967694Z","shell.execute_reply.started":"2025-01-23T14:31:24.967362Z","shell.execute_reply":"2025-01-23T14:31:24.967386Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def find_peaks (probs, required_nr_peaks):\n    # create dictionary first\n    probs_dict = {i: probs[i] for i in range (len (probs))}\n\n    peaks_found = []\n    while len (peaks_found) < required_nr_peaks:\n        i = find_highest_peak (list (probs_dict.values()))\n        dict_keys = list (probs_dict.keys())\n        peak_index = dict_keys[i]\n        peaks_found.append(peak_index)\n        # remove elements surrounding the peak\n        for j in range (peak_index - 2, peak_index + 2):\n            if  (j >= 0) and (j < len (probs)):\n                del probs_dict[j]\n    return peaks_found","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.969310Z","iopub.status.idle":"2025-01-23T14:31:24.969756Z","shell.execute_reply.started":"2025-01-23T14:31:24.969566Z","shell.execute_reply":"2025-01-23T14:31:24.969584Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"l5s1_peak = find_peaks (list (ivd_test_results['probs_LS']), 1)\nprobs_LL = list (ivd_test_results['probs_LL'])\nslice_end = l5s1_peak[0] - 2\nall_peaks = l5s1_peak + find_peaks (probs_LL[0:slice_end], 4)\nall_peaks.sort()\ntmp_list = list (ivd_test_data['image_path_short'])\nimage_paths = []\nfor peak_index in all_peaks:\n    image_paths.append (tmp_list [peak_index])\nivd_selected_instances_df = pd.DataFrame(list(zip(IVD_levels_desc, image_paths)), columns =['level', 'image_path_short'])\nivd_selected_instances_df","metadata":{"execution":{"iopub.status.busy":"2025-01-23T14:31:24.971161Z","iopub.status.idle":"2025-01-23T14:31:24.971608Z","shell.execute_reply.started":"2025-01-23T14:31:24.971371Z","shell.execute_reply":"2025-01-23T14:31:24.971388Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"images = []\nfor index, row in ivd_selected_instances_df.iterrows():\n    image_path_short = row['image_path_short']\n    image_path = test_base_path + image_path_short\n    image = load_dicom(image_path)\n    images.append(image)\ndisplay_images (images, \"Selected images for the intervertebral disc levels\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-01-23T14:31:24.972936Z","iopub.status.idle":"2025-01-23T14:31:24.973319Z","shell.execute_reply.started":"2025-01-23T14:31:24.973136Z","shell.execute_reply":"2025-01-23T14:31:24.973152Z"}},"outputs":[],"execution_count":null}]}