{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":10120999,"sourceType":"datasetVersion","datasetId":6245116}],"dockerImageVersionId":30804,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# # This Python 3 environment comes with many helpful analytics libraries installed\n# # It is defined by the kaggle/python Docker image: https://github.com/kaggle/docker-python\n# # For example, here's several helpful packages to load\n\n# import numpy as np # linear algebra\n# import pandas as pd # data processing, CSV file I/O (e.g. pd.read_csv)\n\n# # Input data files are available in the read-only \"../input/\" directory\n# # For example, running this (by clicking run or pressing Shift+Enter) will list all files under the input directory\n\n# import os\n# for dirname, _, filenames in os.walk('/kaggle/input'):\n#     for filename in filenames:\n#         print(os.path.join(dirname, filename))\n\n# # You can write up to 20GB to the current directory (/kaggle/working/) that gets preserved as output when you create a version using \"Save & Run All\" \n# # You can also write temporary files to /kaggle/temp/, but they won't be saved outside of the current session","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:47.393639Z","iopub.execute_input":"2024-12-06T16:21:47.394258Z","iopub.status.idle":"2024-12-06T16:21:47.399030Z","shell.execute_reply.started":"2024-12-06T16:21:47.394191Z","shell.execute_reply":"2024-12-06T16:21:47.398117Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import 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\n\nimport pydicom as dicom\nimport matplotlib.patches as patches\n\nfrom matplotlib import animation, rc\nimport pandas as pd\n\nimport pydicom as dicom # dicom\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:47.414986Z","iopub.execute_input":"2024-12-06T16:21:47.415280Z","iopub.status.idle":"2024-12-06T16:21:47.420601Z","shell.execute_reply.started":"2024-12-06T16:21:47.415223Z","shell.execute_reply":"2024-12-06T16:21:47.419485Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n\n# 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\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:47.429137Z","iopub.execute_input":"2024-12-06T16:21:47.429444Z","iopub.status.idle":"2024-12-06T16:21:47.503960Z","shell.execute_reply.started":"2024-12-06T16:21:47.429417Z","shell.execute_reply":"2024-12-06T16:21:47.503040Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_desc.head(10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:47.505411Z","iopub.execute_input":"2024-12-06T16:21:47.505678Z","iopub.status.idle":"2024-12-06T16:21:47.514393Z","shell.execute_reply.started":"2024-12-06T16:21:47.505653Z","shell.execute_reply":"2024-12-06T16:21:47.513407Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:47.515664Z","iopub.execute_input":"2024-12-06T16:21:47.516423Z","iopub.status.idle":"2024-12-06T16:21:47.536519Z","shell.execute_reply.started":"2024-12-06T16:21:47.516362Z","shell.execute_reply":"2024-12-06T16:21:47.535592Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:47.538279Z","iopub.execute_input":"2024-12-06T16:21:47.538583Z","iopub.status.idle":"2024-12-06T16:21:52.846538Z","shell.execute_reply.started":"2024-12-06T16:21:47.538556Z","shell.execute_reply":"2024-12-06T16:21:52.845722Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(train)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:52.847558Z","iopub.execute_input":"2024-12-06T16:21:52.847865Z","iopub.status.idle":"2024-12-06T16:21:52.854017Z","shell.execute_reply.started":"2024-12-06T16:21:52.847836Z","shell.execute_reply":"2024-12-06T16:21:52.852960Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\nimport matplotlib.pyplot as plt\n\n# Function to open and display DICOM images\ndef display_dicom_images(image_paths):\n    plt.figure(figsize=(15, 5))  # Adjust figure size if needed\n    for i, path in enumerate(image_paths[:3]):\n        ds = pydicom.dcmread(path)\n        plt.subplot(1, 3, i+1)\n        plt.imshow(ds.pixel_array, cmap=plt.cm.bone)\n        plt.title(f\"Image {i+1}\")\n        plt.axis('off')\n    plt.show()\n\n# Display the first three DICOM images\ndisplay_dicom_images(train_image_paths)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:52.855163Z","iopub.execute_input":"2024-12-06T16:21:52.855489Z","iopub.status.idle":"2024-12-06T16:21:53.379850Z","shell.execute_reply.started":"2024-12-06T16:21:52.855447Z","shell.execute_reply":"2024-12-06T16:21:53.379037Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"This code is a tool for visualizing DICOM (Digital Imaging and Communications in Medicine) images and marking points (coordinates) on the images based on labels provided in a DataFrame. So unifying the coordinates with the DICOM images we get this shit. \nnow with the code below we can do a ton of stuff like in particular\n\nThis code processes and visualizes medical images (DICOM files) with overlaid labels (coordinates) provided in a DataFrame. Here's a summary:\n\n1. **Functions**:\n   - `display_dicom_with_coordinates`: Displays DICOM images and overlays red points for coordinates based on `study_id` and `series_id` in the file path.\n   - `load_dicom_files`: Loads and sorts `.dcm` files numerically from a specified folder.\n\n2. **Directory Structure**:\n   Assumes a folder structure like:  \n   `train_path/train_images/<study_id>/<series_id>/<image_name>.dcm`\n\n3. **Workflow**:\n   - Extract the first DICOM file from each series folder within a study folder.\n   - Use `pydicom` to read the images and overlay points (`x`, `y`) from the label DataFrame.\n\n4. **Output**:\n   A plot displaying DICOM images with overlaid red markers for the labeled coordinates.","metadata":{}},{"cell_type":"code","source":"import os\nimport pydicom\nimport matplotlib.pyplot as plt\nimport pandas as pd\n\n# Function to open and display DICOM images along with coordinates\ndef display_dicom_with_coordinates(image_paths, label_df):\n    fig, axs = plt.subplots(1, len(image_paths), figsize=(18, 6))\n    \n    for idx, path in enumerate(image_paths):  # Display images\n        study_id = int(path.split('/')[-3])\n        series_id = int(path.split('/')[-2])\n        \n        # Filter label coordinates for the current study and series\n        filtered_labels = label_df[(label_df['study_id'] == study_id) & (label_df['series_id'] == series_id)]\n        \n        # Read DICOM image\n        ds = pydicom.dcmread(path)\n        \n        # Plot DICOM image\n        axs[idx].imshow(ds.pixel_array, cmap='gray')\n        axs[idx].set_title(f\"Study ID: {study_id}, Series ID: {series_id}\")\n        axs[idx].axis('off')\n        \n        # Plot coordinates\n        for _, row in filtered_labels.iterrows():\n            axs[idx].plot(row['x'], row['y'], 'ro', markersize=5)\n        \n    plt.tight_layout()\n    plt.show()\n\n# Load DICOM files from a folder\ndef load_dicom_files(path_to_folder):\n    files = [os.path.join(path_to_folder, f) for f in os.listdir(path_to_folder) if f.endswith('.dcm')]\n    files.sort(key=lambda x: int(os.path.splitext(os.path.basename(x))[0].split('-')[-1]))\n    return files\n\n# Display DICOM images with coordinates\nstudy_id = \"100206310\"\nstudy_folder = f'{train_path}/train_images/{study_id}'\n\nimage_paths = []\nfor series_folder in os.listdir(study_folder):\n    series_folder_path = os.path.join(study_folder, series_folder)\n    dicom_files = load_dicom_files(series_folder_path)\n    if dicom_files:\n        image_paths.append(dicom_files[0])  # Add the first image from each series\n\n\ndisplay_dicom_with_coordinates(image_paths, label)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:53.381863Z","iopub.execute_input":"2024-12-06T16:21:53.382155Z","iopub.status.idle":"2024-12-06T16:21:54.041909Z","shell.execute_reply.started":"2024-12-06T16:21:53.382126Z","shell.execute_reply":"2024-12-06T16:21:54.041051Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Preprocessing part","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(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:54.042860Z","iopub.execute_input":"2024-12-06T16:21:54.043112Z","iopub.status.idle":"2024-12-06T16:21:55.195170Z","shell.execute_reply.started":"2024-12-06T16:21:54.043089Z","shell.execute_reply":"2024-12-06T16:21:55.194235Z"}},"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":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.196238Z","iopub.execute_input":"2024-12-06T16:21:55.196571Z","iopub.status.idle":"2024-12-06T16:21:55.202185Z","shell.execute_reply.started":"2024-12-06T16:21:55.196544Z","shell.execute_reply":"2024-12-06T16:21:55.201313Z"}},"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')\n# Merge the dataframes on the common column 'series_id'\nfinal_merged_df = pd.merge(merged_df, train_desc, on='series_id', how='inner')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.203310Z","iopub.execute_input":"2024-12-06T16:21:55.203688Z","iopub.status.idle":"2024-12-06T16:21:55.253387Z","shell.execute_reply.started":"2024-12-06T16:21:55.203660Z","shell.execute_reply":"2024-12-06T16:21:55.252472Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Merge the dataframes on the common column 'series_id'\nfinal_merged_df = pd.merge(merged_df, train_desc, on=['series_id','study_id'], how='inner')\n# Display the first few rows of the final merged dataframe\n\n#okay now we fetch only the L5/S1 \nfinal_merged_df = final_merged_df[final_merged_df['level'] == 'L5/S1'].reset_index(drop=True)\nfinal_merged_df.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.254577Z","iopub.execute_input":"2024-12-06T16:21:55.254856Z","iopub.status.idle":"2024-12-06T16:21:55.283191Z","shell.execute_reply.started":"2024-12-06T16:21:55.254831Z","shell.execute_reply":"2024-12-06T16:21:55.282384Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df['study_id'] == 100206310].sort_values(['x','y'],ascending = True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.284078Z","iopub.execute_input":"2024-12-06T16:21:55.284337Z","iopub.status.idle":"2024-12-06T16:21:55.297682Z","shell.execute_reply.started":"2024-12-06T16:21:55.284307Z","shell.execute_reply":"2024-12-06T16:21:55.296693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df['series_id'] == 1012284084].sort_values(\"instance_number\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.298811Z","iopub.execute_input":"2024-12-06T16:21:55.299157Z","iopub.status.idle":"2024-12-06T16:21:55.313196Z","shell.execute_reply.started":"2024-12-06T16:21:55.299125Z","shell.execute_reply":"2024-12-06T16:21:55.312327Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Filter the dataframe for the given study_id and sort by instance_number\nfiltered_df = final_merged_df[final_merged_df['study_id'] == 1013589491].sort_values(\"instance_number\")\n\n# Display the resulting dataframe\nfiltered_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.314467Z","iopub.execute_input":"2024-12-06T16:21:55.314790Z","iopub.status.idle":"2024-12-06T16:21:55.332889Z","shell.execute_reply.started":"2024-12-06T16:21:55.314762Z","shell.execute_reply":"2024-12-06T16:21:55.331982Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Sort final_merged_df by study_id, series_id, and series_description\nsorted_final_merged_df = final_merged_df[final_merged_df['study_id'] == 1013589491].sort_values(by=['series_id', 'series_description', 'instance_number'])\nsorted_final_merged_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.336039Z","iopub.execute_input":"2024-12-06T16:21:55.336332Z","iopub.status.idle":"2024-12-06T16:21:55.353738Z","shell.execute_reply.started":"2024-12-06T16:21:55.336301Z","shell.execute_reply":"2024-12-06T16:21:55.352901Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# We see that, <br>\n## Saggital T1 images map to Neural Foraminal Narrowing <br>\n## Axial T2 images map to Subarticular Stenosis <br>\n## Sagittal T2/STIR map to Canal Stenosis <br>\n\n--> this is super mega important","metadata":{}},{"cell_type":"code","source":"import pandas as pd\n\n# Create the row_id column\nfinal_merged_df['row_id'] = (\n    final_merged_df['study_id'].astype(str) + '_' +\n    final_merged_df['condition'].str.lower().str.replace(' ', '_') + '_' +\n    final_merged_df['level'].str.lower().str.replace('/', '_')\n)\n\n# Create the image_path column\nfinal_merged_df['image_path'] = (\n    f'{train_path}/train_images/' + \n    final_merged_df['study_id'].astype(str) + '/' +\n    final_merged_df['series_id'].astype(str) + '/' +\n    final_merged_df['instance_number'].astype(str) + '.dcm'\n)\n\n# Note: Check image path, since there's 1 instance id, for 1 image, but there's many more images other than the ones labelled in the instance ID. \n\n# Display the updated dataframe\nfinal_merged_df.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.354678Z","iopub.execute_input":"2024-12-06T16:21:55.354890Z","iopub.status.idle":"2024-12-06T16:21:55.401735Z","shell.execute_reply.started":"2024-12-06T16:21:55.354869Z","shell.execute_reply":"2024-12-06T16:21:55.400850Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Normal/Mild\"].value_counts().sum()","metadata":{"execution":{"iopub.status.busy":"2024-12-06T16:21:55.402897Z","iopub.execute_input":"2024-12-06T16:21:55.403147Z","iopub.status.idle":"2024-12-06T16:21:55.433260Z","shell.execute_reply.started":"2024-12-06T16:21:55.403124Z","shell.execute_reply":"2024-12-06T16:21:55.432556Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Moderate\"].value_counts().sum()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.434165Z","iopub.execute_input":"2024-12-06T16:21:55.434436Z","iopub.status.idle":"2024-12-06T16:21:55.450269Z","shell.execute_reply.started":"2024-12-06T16:21:55.434411Z","shell.execute_reply":"2024-12-06T16:21:55.449393Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"So we have 7233 normal/mild and 1871 with moderate severity only on the L5/S1 ???? isn't this too much?\n\n-> no with all the vertebrae we have 37k mild and 7950 moderate and 734 severe.\n\nWe need to consider even the imbalances in the dataset maybe?","metadata":{}},{"cell_type":"markdown","source":"# The T1, T2 and STIR Explained:\n\nMRI imaging of the spine can be taken in three planes: the axial plane, the sagittal plane, and the coronal plane. The two main image types you'll need for this challenge are the axial and sagittal planes. The axial plane takes images horizonal slices (perpendicular to the spine) across the body from top to bottom. The sagittal plane takes vertical slices (parallel to the spine) going from left to right.\n\nMRI images come in multiple variants. They can generally be classified as either being T1 weighted or T2 weighted. T1 weighted images show fat as being brighter. The inner part of bones would appear brighter on T1 images. T2 images show water as brighter. The spinal canal would appear as brighter on T2 images. MRI images are not standardized with regards to the pixel values that are output from it (unlike CT images). So you'll need to figure out how to standardize these images (or maybe you wont need to at all, we'll leave it up to you).\n\n","metadata":{}},{"cell_type":"markdown","source":"# WTF is a STIR?\n\nIn MRI (Magnetic Resonance Imaging), **T2-weighted (T2) images** and **STIR (Short TI Inversion Recovery)** images are different types of sequences that provide different contrast characteristics, helping to visualize various tissues and abnormalities in the body. Here's a breakdown of what each term means:\n\n### 1. **T2-Weighted (T2) Images**:\n- **T2-weighted imaging** is a specific MRI sequence where the tissue contrast is mainly based on the **T2 relaxation time**, which refers to how quickly the protons in the tissue lose their transverse magnetization after being disturbed by a radiofrequency pulse.\n- **T2 contrast**: In T2 images, **fluid** appears **bright** (high signal), while **fat** and most **soft tissues** appear **darker** (low signal). This makes T2 images particularly useful for detecting **inflammation, edema (swelling), and fluid-filled structures** like cysts or tumors.\n  - **Bright areas**: Fluid, such as in the brain's ventricles or cysts.\n  - **Dark areas**: Fat and tissues like muscles or cartilage.\n  \n**Applications**:\n- Detecting **inflammation**, **swelling**, or **fluid accumulation** in tissues.\n- **Brain imaging** to identify conditions like **stroke, tumors**, or **multiple sclerosis**.\n- **Joint imaging** to detect issues like **ligament damage** or **inflammation**.\n\n### 2. **STIR (Short TI Inversion Recovery)**:\n- **STIR** is a type of MRI sequence that specifically suppresses the signal from **fat** tissue. This is achieved by using a short **Inversion Time (TI)**, which inverts the signal from fat, effectively nullifying it and making fat appear dark.\n  - **Fat suppression**: By suppressing the fat signal, **STIR images** enhance the contrast of other tissues, making **edema** (fluid) and **inflammatory changes** stand out more clearly.\n- STIR images are particularly useful in **musculoskeletal imaging** and when looking for **inflammation or lesions** in areas that contain a lot of fat (like around joints, in the muscles, or in soft tissues).\n\n**Applications**:\n- **Musculoskeletal imaging** (bones, joints, soft tissues) to detect **edema**, **inflammation**, or **soft tissue injuries**.\n- **Detecting tumors** that are adjacent to fatty tissues (e.g., **soft tissue sarcomas**).\n- **Spinal imaging** to assess **disc herniations** or **inflammatory diseases**.\n\n---\n\n### Key Differences:\n- **T2-Weighted** images show **fluid** as bright and are useful for identifying conditions involving **fluid accumulation** or **swelling**.\n- **STIR** images are designed to suppress **fat signals** and are ideal for highlighting **inflammatory or pathological changes** in tissues surrounded by fat.\n\n### Summary:\n- **T2 images** are good for visualizing **fluid**, **inflammation**, and **swelling**.\n- **STIR images** are particularly useful when you need to suppress fat and enhance visibility of **edema** or **inflammation** in tissues that are adjacent to fat.\n\nBoth sequences help provide complementary information to improve diagnosis in various clinical scenarios, particularly when evaluating soft tissue and musculoskeletal conditions.","metadata":{}},{"cell_type":"code","source":"# Define the base path for test images\nbase_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/'\n\n# Function to get image paths for a series\ndef get_image_paths(row):\n    series_path = os.path.join(base_path, str(row['study_id']), str(row['series_id']))\n    if os.path.exists(series_path):\n        return [os.path.join(series_path, f) for f in os.listdir(series_path) if os.path.isfile(os.path.join(series_path, f))]\n    return []\n\n# Mapping of series_description to conditions\ncondition_mapping = {\n    'Sagittal T1': {'left': 'left_neural_foraminal_narrowing', 'right': 'right_neural_foraminal_narrowing'},\n    'Axial T2': {'left': 'left_subarticular_stenosis', 'right': 'right_subarticular_stenosis'},\n    'Sagittal T2/STIR': 'spinal_canal_stenosis'\n}\n\n# Create a list to store the expanded rows\nexpanded_rows = []\n\n# Expand the dataframe by adding new rows for each file path\nfor index, row in test_desc.iterrows():\n    image_paths = get_image_paths(row)\n    conditions = condition_mapping.get(row['series_description'], {})\n    if isinstance(conditions, str):  # Single condition\n        conditions = {'left': conditions, 'right': conditions}\n    for side, condition in conditions.items():\n        for image_path in image_paths:\n            expanded_rows.append({\n                'study_id': row['study_id'],\n                'series_id': row['series_id'],\n                'series_description': row['series_description'],\n                'image_path': image_path,\n                'condition': condition,\n                'row_id': f\"{row['study_id']}_{condition}\"\n            })\n\n# Create a new dataframe from the expanded rows\nexpanded_test_desc = pd.DataFrame(expanded_rows)\n\n# Display the resulting dataframe\nexpanded_test_desc.head(5)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.451213Z","iopub.execute_input":"2024-12-06T16:21:55.451496Z","iopub.status.idle":"2024-12-06T16:21:55.523087Z","shell.execute_reply.started":"2024-12-06T16:21:55.451472Z","shell.execute_reply":"2024-12-06T16:21:55.522270Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# change severity column labels\n#Normal/Mild': 'normal_mild', 'Moderate': 'moderate', 'Severe': 'severe'}\nfinal_merged_df['severity'] = final_merged_df['severity'].map({'Normal/Mild': 'normal_mild', 'Moderate': 'moderate', 'Severe': 'severe'})\nfinal_merged_df = final_merged_df.dropna(subset=['severity'])  # Remove rows where 'severity' is NaN\n\n#Tentative \n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.523999Z","iopub.execute_input":"2024-12-06T16:21:55.524230Z","iopub.status.idle":"2024-12-06T16:21:55.534569Z","shell.execute_reply.started":"2024-12-06T16:21:55.524207Z","shell.execute_reply":"2024-12-06T16:21:55.533705Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_data = expanded_test_desc\ntrain_data = final_merged_df","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.535769Z","iopub.execute_input":"2024-12-06T16:21:55.536099Z","iopub.status.idle":"2024-12-06T16:21:55.548546Z","shell.execute_reply.started":"2024-12-06T16:21:55.536070Z","shell.execute_reply":"2024-12-06T16:21:55.547639Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\n# Define a function to check if a path exists\ndef check_exists(path):\n    return os.path.exists(path)\n\n# Define a function to check if a study ID directory exists\ndef check_study_id(row):\n    study_id = row['study_id']\n    path = f'{train_path}/train_images/{study_id}'\n    return check_exists(path)\n\n# Define a function to check if a series ID directory exists\ndef check_series_id(row):\n    study_id = row['study_id']\n    series_id = row['series_id']\n    path = f'{train_path}/train_images/{study_id}/{series_id}'\n    return check_exists(path)\n\n# Define a function to check if an image file exists\ndef check_image_exists(row):\n    image_path = row['image_path']\n    return check_exists(image_path)\n\n# Apply the functions to the train_data dataframe\ntrain_data['study_id_exists'] = train_data.apply(check_study_id, axis=1)\ntrain_data['series_id_exists'] = train_data.apply(check_series_id, axis=1)\ntrain_data['image_exists'] = train_data.apply(check_image_exists, axis=1)\n\n# Filter train_data\ntrain_data = train_data[(train_data['study_id_exists']) & (train_data['series_id_exists']) & (train_data['image_exists'])]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:21:55.549568Z","iopub.execute_input":"2024-12-06T16:21:55.549860Z","iopub.status.idle":"2024-12-06T16:22:01.756042Z","shell.execute_reply.started":"2024-12-06T16:21:55.549834Z","shell.execute_reply":"2024-12-06T16:22:01.755307Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data.head(3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:01.757043Z","iopub.execute_input":"2024-12-06T16:22:01.757366Z","iopub.status.idle":"2024-12-06T16:22:01.771148Z","shell.execute_reply.started":"2024-12-06T16:22:01.757335Z","shell.execute_reply":"2024-12-06T16:22:01.770284Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(train_data)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:01.772149Z","iopub.execute_input":"2024-12-06T16:22:01.772468Z","iopub.status.idle":"2024-12-06T16:22:01.782942Z","shell.execute_reply.started":"2024-12-06T16:22:01.772440Z","shell.execute_reply":"2024-12-06T16:22:01.782221Z"}},"outputs":[],"execution_count":null},{"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":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:01.783891Z","iopub.execute_input":"2024-12-06T16:22:01.784144Z","iopub.status.idle":"2024-12-06T16:22:01.794561Z","shell.execute_reply.started":"2024-12-06T16:22:01.784120Z","shell.execute_reply":"2024-12-06T16:22:01.793724Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load images randomly  --> this is not working right now\n# probabilmente visto che nei train example ci sono anche le altre vertebre allora\n# va out of bound perché non le trova. Forse la modifica va fatta altrove? prendere tutto il \n# dataset e poi andare a lavorare altrove per eliminarle? oppure possiamo provare a ridurre \n# e splittare il train set e dividercelo in 3\n# ma non so se è una buona idea sarò onesto. \n#--> come sono fatte le train_images?\nimport random\nimages = []\nrow_ids = []\nselected_indices = random.sample(range(len(train_data)), 2)\nfor i in selected_indices:\n    # Read the DICOM file\n    dicom_file = pydicom.dcmread(train_data['image_path'][i])  # Corrected line\n    image = dicom_file.pixel_array  # Extract pixel data from DICOM file\n    images.append(image)\n    row_ids.append(train_data['row_id'][i])\n\n# Plot images\nfig, ax = plt.subplots(1, 2, figsize=(8, 4))\nfor i in range(2):\n    ax[i].imshow(images[i], cmap='gray')\n    ax[i].set_title(f'Row ID: {row_ids[i]}', fontsize=8)\n    ax[i].axis('off')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:01.795688Z","iopub.execute_input":"2024-12-06T16:22:01.796020Z","iopub.status.idle":"2024-12-06T16:22:02.182909Z","shell.execute_reply.started":"2024-12-06T16:22:01.795985Z","shell.execute_reply":"2024-12-06T16:22:02.182078Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"i'ive modified the dicom read_file with teh newer version called dicom.dcread\nand then i had to put into a pixel array\n\nThis snippet of code is iterating over a selection of indices (`selected_indices`) to process DICOM files, extract pixel data (images), and associate them with their respective row IDs from the `train_data` DataFrame. Here's a step-by-step explanation:\n\n### **Explanation of Each Line**\n1. **`for i in selected_indices:`**\n   - Iterates over a list of randomly selected indices (`selected_indices`), which correspond to specific rows in the `train_data` DataFrame.\n\n2. **`dicom_file = pydicom.dcmread(train_data['image_path'][i])`**\n   - Reads a DICOM file at the file path specified by `train_data['image_path'][i]` using the `pydicom.dcmread` method.\n   - `train_data['image_path'][i]` accesses the file path for the i-th selected entry in the DataFrame.\n   - The `dcmread` method loads the DICOM file into a `dicom_file` object, which contains metadata and image data.\n\n3. **`image = dicom_file.pixel_array`**\n   - Extracts the pixel data (image matrix) from the DICOM file. \n   - The `pixel_array` is a NumPy array representing the grayscale image contained in the DICOM file.\n\n4. **`images.append(image)`**\n   - Appends the extracted image data (pixel array) to the `images` list for storage.\n\n5. **`row_ids.append(train_data['row_id'][i])`**\n   - Appends the `row_id` corresponding to the current index `i` from the `train_data` DataFrame to the `row_ids` list.\n   - This keeps track of the identifier associated with each image for future reference.\n\n### **Purpose**\nThe snippet is essentially:\n1. Reading DICOM files from a dataset.\n2. Extracting image data from these files.\n3. Storing the extracted images and their associated identifiers (`row_id`) in separate lists (`images` and `row_ids`).\n\nThis is part of a preprocessing step, often used in medical imaging tasks, to prepare data for analysis, visualization, or model training.","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Loading Data Time","metadata":{}},{"cell_type":"code","source":"train_data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:02.184199Z","iopub.execute_input":"2024-12-06T16:22:02.184595Z","iopub.status.idle":"2024-12-06T16:22:02.201875Z","shell.execute_reply.started":"2024-12-06T16:22:02.184557Z","shell.execute_reply":"2024-12-06T16:22:02.200821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data.dropna()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:02.202820Z","iopub.execute_input":"2024-12-06T16:22:02.203083Z","iopub.status.idle":"2024-12-06T16:22:02.229741Z","shell.execute_reply.started":"2024-12-06T16:22:02.203058Z","shell.execute_reply":"2024-12-06T16:22:02.229012Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"The error occurs because the Lambda transform in your transforms.Compose pipeline is attempting to perform arithmetic on a DICOM file object rather than the pixel data extracted from the file. This happens because in your CustomDataset, the image variable holds the entire DICOM file object returned by pydicom.dcmread, not the pixel data (pixel_array) extracted from it.\n\n\nI've modified the __getitem__ to extract the pixel value\n\n\nadditional info: \n\nExtracting pixel_array:\n\nThe dicom_file.pixel_array retrieves the actual image data from the DICOM file as a NumPy array.\nThis NumPy array represents the grayscale pixel intensity values of the image.\nPassing the Correct Data to the Transform:\n\nThe extracted pixel_array is passed to the transform pipeline instead of the entire DICOM file object.\nThis ensures that operations like multiplication (x * 255) and data type conversion (astype(np.uint8)) are performed on valid image data.\nAdditional Notes:\nThe transforms.Lambda line is converting the image data to uint8 format to make it compatible with ToPILImage.\nThe ToPILImage transformation expects the input to be a NumPy array or a PyTorch tensor in a valid image data format.\nWith this fix, the error should no longer occur, and your dataset and dataloaders should work correctly.","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torch\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:02.230666Z","iopub.execute_input":"2024-12-06T16:22:02.230948Z","iopub.status.idle":"2024-12-06T16:22:02.235696Z","shell.execute_reply.started":"2024-12-06T16:22:02.230923Z","shell.execute_reply":"2024-12-06T16:22:02.234675Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\nclass CustomDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.dataframe)\n\n    def __getitem__(self, index):\n        image_path = self.dataframe['image_path'][index]\n        dicom_file = pydicom.dcmread(image_path)  # Load the DICOM file\n        image = dicom_file.pixel_array  # Extract the pixel array (image data)\n        label = self.dataframe['severity'][index]\n\n        # Check for NaN labels\n        if pd.isnull(label):\n            raise ValueError(f\"Invalid label at index {index}. Label: {label}\")\n        \n        if self.transform:\n            image = self.transform(image)\n\n        return image, label\n\n# Function to create datasets and dataloaders for each series description\ndef create_datasets_and_loaders(df, series_description, transform, batch_size=8):\n    # filtered_df = df[df['series_description'] == series_description]\n    #i'm filtering the creation of the dataset to match only for the L5/S1 stuff\n    #this is highly tentative tbf\n    filtered_df = df[(df['series_description'] == series_description) & (df['level'] == 'L5/S1')]\n\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\n    train_dataset = CustomDataset(train_df, transform)\n    val_dataset = CustomDataset(val_df, 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)\n\n# Define the transforms\ntransform = transforms.Compose([\n    transforms.Lambda(lambda x: (x * 255).astype(np.uint8)),  # Convert back to uint8 for PIL\n    transforms.ToPILImage(),\n    transforms.Resize((224, 224)), # this is for efficientNET B0 actually\n    transforms.Grayscale(num_output_channels=3),\n    transforms.ToTensor(),\n])\n\n# Create dataloaders for each series description\ndataloaders = {}\nlengths = {}\n\ntrainloader_t1, valloader_t1, len_train_t1, len_val_t1 = create_datasets_and_loaders(train_data, 'Sagittal T1', transform)\ntrainloader_t2, valloader_t2, len_train_t2, len_val_t2 = create_datasets_and_loaders(train_data, 'Axial T2', transform)\ntrainloader_t2stir, valloader_t2stir, len_train_t2stir, len_val_t2stir = create_datasets_and_loaders(train_data, 'Sagittal T2/STIR', transform)\n\ndataloaders['Sagittal T1'] = (trainloader_t1, valloader_t1)\ndataloaders['Axial T2'] = (trainloader_t2, valloader_t2)\ndataloaders['Sagittal T2/STIR'] = (trainloader_t2stir, valloader_t2stir)\n\nlengths['Sagittal T1'] = (len_train_t1, len_val_t1)\nlengths['Axial T2'] = (len_train_t2, len_val_t2)\nlengths['Sagittal T2/STIR'] = (len_train_t2stir, len_val_t2stir)\n\n# Dictionary mapping labels to indices\nlabel_map = {'Mild': 0, 'Moderate': 1, 'Severe': 2}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:02.236825Z","iopub.execute_input":"2024-12-06T16:22:02.237131Z","iopub.status.idle":"2024-12-06T16:22:02.265753Z","shell.execute_reply.started":"2024-12-06T16:22:02.237098Z","shell.execute_reply":"2024-12-06T16:22:02.265110Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Function to visualize a batch of images\ndef visualize_batch(dataloader):\n    images, labels = next(iter(dataloader))\n    fig, axes = plt.subplots(1, len(images), figsize=(20, 5))\n    for i, (img, lbl) in enumerate(zip(images, labels)):\n        ax = axes[i]\n        img = img.permute(1, 2, 0)  # Convert to HWC for visualization\n        ax.imshow(img)\n        ax.set_title(f\"Label: {lbl}\")\n        ax.axis('off')\n    plt.show()\n\n# Visualize samples from each dataloader\nprint(\"Visualizing Sagittal T1 samples\")\nvisualize_batch(trainloader_t1)\nprint(\"Visualizing Axial T2 samples\")\nvisualize_batch(trainloader_t2)\nprint(\"Visualizing Sagittal T2/STIR samples\")\nvisualize_batch(trainloader_t2stir)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:02.266754Z","iopub.execute_input":"2024-12-06T16:22:02.267029Z","iopub.status.idle":"2024-12-06T16:22:04.557323Z","shell.execute_reply.started":"2024-12-06T16:22:02.267002Z","shell.execute_reply":"2024-12-06T16:22:04.556435Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"I think we are set and done for what concern the Dataloader -> we have a dataloader for each Series_description -> series description is the unification of the various side of the image into the 3 axis. \n\n\n# THINGS TO CONSIDER AND SEE IF THAT WORKS:\n1) are those images only of the L5/S1 ?  i think yes but i'm not a medic btw\n3) - ","metadata":{}},{"cell_type":"markdown","source":"We are fetching a Batch with this dataloader below\n\nThis retrieves the first batch of images and their corresponding labels from the trainloader_t2 DataLoader.\nimage is a tensor containing the batch of images, typically in the shape (batch_size, channels, height, width).\nlabel is a tensor containing the batch of labels corresponding to the images.\n\nsample= image ...\nThis selects the second image in the batch (image[1]).\nThe .permute(1, 2, 0) operation changes the image's shape from (channels, height, width) to (height, width, channels). This is necessary because PyTorch tensors store channels first, but matplotlib expects channels last for RGB images.\n\n","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nimage, label = next(iter(trainloader_t2))\nsample = image[1].permute(1, 2, 0)  #sample\n\n# Plot images\nplt.figsize=(8, 4)\nplt.imshow(sample, cmap='gray')   #modified images[0] to sample \nplt.title(label[0])\nplt.axis('off')\nplt.tight_layout()\nplt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:04.558790Z","iopub.execute_input":"2024-12-06T16:22:04.559144Z","iopub.status.idle":"2024-12-06T16:22:04.954128Z","shell.execute_reply.started":"2024-12-06T16:22:04.559107Z","shell.execute_reply":"2024-12-06T16:22:04.953302Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# image[0] vs sample\n\nOriginal: image[0]\nDirectly accesses the first image in the batch.\nStored in the original PyTorch format: (channels, height, width).\n\nNew: sample\nA variable used to hold the reformatted version of the image: (height, width, channels) for compatibility with plotting libraries like matplotlib.\nBy assigning the transformed image to sample, you're making it clear that this is the version ready for display, rather than the original tensor from the DataLoader.\n\nWhy?\nReadability: It indicates that this variable is meant for display purposes.\nAvoids Repeated Transformation: If you decide to reuse the reformatted image (e.g., to visualize multiple times or perform additional operations), having it stored in sample avoids recalculating the .permute(1, 2, 0) operation.\n\n\nUsing sample is not strictly necessary but improves code clarity and reusability. It's a good practice, especially when working with transformed data, to assign the result to a meaningful variable.","metadata":{"jp-MarkdownHeadingCollapsed":true}},{"cell_type":"markdown","source":"# Now we need a model ResNET is already done maybe efficientNET b0 should be a good candidate \nalso using the hyperparameter tuning should be a good idea.\n\n\nthis book should be good for the hyperparam tuning but it's very long, just watch some things.\n[The tuning handbook](https://github.com/google-research/tuning_playbook)\n\n\nRESnet is tooooooooooooooo heavy i think with out GPUs\nthis should be a good read\n[Explaining EfficientNET](https://arjun-sarkar786.medium.com/understanding-efficientnet-the-most-powerful-cnn-architecture-eaeb40386fad)\n\n\n[ResNET vs Efficient vs VGG vs NN](https://dev.to/saaransh_gupta_1903/resnet-vs-efficientnet-vs-vgg-vs-nn-2hf5)\n\nthose above are a little bit outdated i think.\n\n\n\n\n# INTERESTING STUFF\n\n[MedFormer](https://github.com/DL4mHealth/Medformer)\nIt's from NeurIPS 2024 so should be lit af.\nthis is the paper https://arxiv.org/pdf/2405.19363\n\nMedicalDETR -> by facebook, still pretty good \n\n\nswinUnet -> https://arxiv.org/abs/2105.05537","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"the model try 1","metadata":{}},{"cell_type":"markdown","source":"","metadata":{}},{"cell_type":"markdown","source":"# THE MODEL\n\nwe need to implement efficienNET","metadata":{}},{"cell_type":"markdown","source":"We need to consider the width and height of the images. \nbecause each EfficienNET have a different resolution\nthe formula is (height, width, 3) the 3 is for the RGB","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nfrom sklearn.model_selection import train_test_split\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torch\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom tqdm import tqdm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:04.955361Z","iopub.execute_input":"2024-12-06T16:22:04.955715Z","iopub.status.idle":"2024-12-06T16:22:04.961566Z","shell.execute_reply.started":"2024-12-06T16:22:04.955677Z","shell.execute_reply":"2024-12-06T16:22:04.960667Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"In the cell Below our problem are the weights of efficientnet_v2_s-dd5fe13b\nHow the hell can i find them? lmao?","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\n\nclass CustomEfficientNetV2(nn.Module):\n    def __init__(self, num_classes=3, pretrained_weights=None):\n        super(CustomEfficientNetV2, self).__init__()\n        self.model = models.efficientnet_v2_s(weights=None)\n        if pretrained_weights:\n            self.model.load_state_dict(torch.load(pretrained_weights, weights_only=True))\n        num_ftrs = self.model.classifier[-1].in_features\n        self.model.classifier[-1] = nn.Linear(num_ftrs, num_classes)\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\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n# Path to the locally uploaded weights file\nweights_path = '/kaggle/input/weight-path/efficientnet_v2_s-dd5fe13b.pth'\n\n# Initialize models\nsagittal_t1_model = CustomEfficientNetV2(num_classes=3, pretrained_weights=weights_path).to(device)\naxial_t2_model = CustomEfficientNetV2(num_classes=3, pretrained_weights=weights_path).to(device)\nsagittal_t2stir_model = CustomEfficientNetV2(num_classes=3, pretrained_weights=weights_path).to(device)\n\n# Optionally freeze initial layers\nfor param in sagittal_t1_model.model.features.parameters():\n    param.requires_grad = False\nfor param in axial_t2_model.model.features.parameters():\n    param.requires_grad = False\nfor param in sagittal_t2stir_model.model.features.parameters():\n    param.requires_grad = False\n\n# Unfreeze the final fully connected layer\nfor param in sagittal_t1_model.model.classifier.parameters():\n    param.requires_grad = True\nfor param in axial_t2_model.model.classifier.parameters():\n    param.requires_grad = True\nfor param in sagittal_t2stir_model.model.classifier.parameters():\n    param.requires_grad = True\n\n# Training parameters\ncriterion = nn.CrossEntropyLoss()\n\n# Initialize separate optimizers for each model\noptimizer_sagittal_t1 = torch.optim.Adam(sagittal_t1_model.model.classifier.parameters(), lr=0.001)\noptimizer_axial_t2 = torch.optim.Adam(axial_t2_model.model.classifier.parameters(), lr=0.001)\noptimizer_sagittal_t2stir = torch.optim.Adam(sagittal_t2stir_model.model.classifier.parameters(), lr=0.001)\n\n# Store the models and optimizers in dictionaries for easy access\nmodels = {\n    'Sagittal T1': sagittal_t1_model,\n    'Axial T2': axial_t2_model,\n    'Sagittal T2/STIR': sagittal_t2stir_model,\n}\noptimizers = {\n    'Sagittal T1': optimizer_sagittal_t1,\n    'Axial T2': optimizer_axial_t2,\n    'Sagittal T2/STIR': optimizer_sagittal_t2stir,\n}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:04.963025Z","iopub.execute_input":"2024-12-06T16:22:04.963419Z","iopub.status.idle":"2024-12-06T16:22:06.835259Z","shell.execute_reply.started":"2024-12-06T16:22:04.963379Z","shell.execute_reply":"2024-12-06T16:22:06.834384Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Count trainable parameters\ntrainable_params = sum(p.numel() for p in sagittal_t1_model.parameters() if p.requires_grad)\nprint(f\"Number of parameters: {trainable_params}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:06.836555Z","iopub.execute_input":"2024-12-06T16:22:06.837190Z","iopub.status.idle":"2024-12-06T16:22:06.845346Z","shell.execute_reply.started":"2024-12-06T16:22:06.837128Z","shell.execute_reply":"2024-12-06T16:22:06.843619Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_map = {'normal_mild': 0, 'moderate': 1, 'severe': 2}","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:06.846576Z","iopub.execute_input":"2024-12-06T16:22:06.846945Z","iopub.status.idle":"2024-12-06T16:22:06.856088Z","shell.execute_reply.started":"2024-12-06T16:22:06.846905Z","shell.execute_reply":"2024-12-06T16:22:06.855214Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# for images, labels in trainloader_t2:\n#     labels = torch.tensor([label_map[label] for label in labels])\n#     labels = labels.to(device)\n#     print(labels)\n#     break\n# -> the labels are already a tensor ???\n# If labels are strings, map them using label_map\n# for images, labels in trainloader_t2:\n#     # Assuming labels are strings and need to be mapped to indices\n#     labels = torch.tensor([label_map[label.item()] for label in labels]).to(device)\n#     print(labels)\n#     break\n# Ensure the labels are not tensors\n# for images, labels in trainloader:\n#     # If labels are tensors and already integers, we can directly use them\n#     if isinstance(labels, torch.Tensor):\n#         labels = labels.to(device)  # Move to device\n#     else:\n#         # If labels are strings, apply the label_map to map strings to integers\n#         labels = torch.tensor([label_map[label] for label in labels]).to(device)\n    \n#     # Now proceed with training\n#     print(labels)  # Debug print to check the labels\n#     break\n\n\n# --> debug statement\n\n# for images, labels in trainloader_t2:\n#     print(\"Labels in batch:\", labels)\n\n#     # Check for NaN labels in the batch\n#     if any(pd.isnull(label) for label in labels):\n#         print(\"Found NaN labels in batch!\")\n#         continue  # Skip this batch\n\n#     # Map string labels to numeric values using label_map\n#     try:\n#         labels = torch.tensor([label_map[label] for label in labels]).to(device)\n#     except KeyError as e:\n#         print(f\"Found an unexpected label: {e}. Skipping this batch.\")\n#         continue  # Skip batch if label is not in label_map\n\n#     # Proceed with training\n#     print(\"Mapped labels:\", labels)\n\n\n# for images, labels in trainloader_t2:\n#     # Filter out NaN labels\n#     valid_indices = ~torch.isnan(labels)\n#     images = images[valid_indices]\n#     labels = labels[valid_indices]\n    \n#     # Map labels if they are valid\n#     labels = torch.tensor([label_map[label.item()] for label in labels]).to(device)\n\n#     # Proceed with training\n\nfor images, labels in trainloader_t2:\n    print(\"Labels in batch:\", labels)\n\n    # Filter out invalid labels\n    if any(pd.isnull(label) for label in labels):\n        print(\"Found NaN labels in batch! Skipping batch...\")\n        continue\n\n    # Map labels to numeric values\n    try:\n        labels = torch.tensor([label_map[label] for label in labels]).to(device)\n    except KeyError as e:\n        print(f\"Found an unexpected label: {e}. Skipping this batch.\")\n        continue  # Skip batch if a label is not in label_map\n\n    # Proceed with training\n    print(\"Mapped labels:\", labels)\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:27:15.262468Z","iopub.execute_input":"2024-12-06T16:27:15.263131Z","iopub.status.idle":"2024-12-06T16:28:02.426544Z","shell.execute_reply.started":"2024-12-06T16:27:15.263096Z","shell.execute_reply":"2024-12-06T16:28:02.425602Z"},"scrolled":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.optim.lr_scheduler as lr_scheduler\nfrom copy import deepcopy\n\ndef 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    \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([label_map[label] for label in 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([label_map[label] for label in 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                    _, predicted = torch.max(probabilities, 1)\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'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":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:28:40.729103Z","iopub.execute_input":"2024-12-06T16:28:40.729930Z","iopub.status.idle":"2024-12-06T16:28:40.741100Z","shell.execute_reply.started":"2024-12-06T16:28:40.729880Z","shell.execute_reply":"2024-12-06T16:28:40.740221Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Okay we have some NaN problem, probably correlated to the ","metadata":{"execution":{"iopub.status.busy":"2024-12-06T15:09:03.687317Z","iopub.execute_input":"2024-12-06T15:09:03.687704Z","iopub.status.idle":"2024-12-06T15:09:03.695757Z","shell.execute_reply.started":"2024-12-06T15:09:03.687670Z","shell.execute_reply":"2024-12-06T15:09:03.694645Z"}}},{"cell_type":"code","source":"# # Training all models\n# for desc, model in models.items():\n#     if desc == 'Sagittal T1':\n#         trainloader, valloader, len_train, len_val = trainloader_t1, valloader_t1, len_train_t1, len_val_t1\n#     elif desc == 'Axial T2':\n#         trainloader, valloader, len_train, len_val = trainloader_t2, valloader_t2, len_train_t2, len_val_t2\n#     elif desc == 'Sagittal T2/STIR':\n#         trainloader, valloader, len_train, len_val = trainloader_t2stir, valloader_t2stir, len_train_t2stir, len_val_t2stir\n    \n#     print(f\"Training model for {desc}\")\n#     train_model(model, trainloader, valloader, len_train, len_val, optimizers[desc])\n\n# Train only the Sagittal T1 model\ndesc = 'Sagittal T1'\nmodel = models[desc]\ntrainloader, valloader, len_train, len_val = trainloader_t1, valloader_t1, len_train_t1, len_val_t1\noptimizer = optimizers[desc]\n\nprint(f\"Training model for {desc}\")\ntrain_model(model, trainloader, valloader, len_train, len_val, optimizer)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:28:48.654737Z","iopub.execute_input":"2024-12-06T16:28:48.655091Z","iopub.status.idle":"2024-12-06T16:34:11.910838Z","shell.execute_reply.started":"2024-12-06T16:28:48.655060Z","shell.execute_reply":"2024-12-06T16:34:11.909927Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"this below is a generic dummy implementation","metadata":{}},{"cell_type":"code","source":"# import tensorflow as tf\n# from tensorflow.keras.layers import Dense, Dropout, GlobalAveragePooling2D\n# from tensorflow.keras.models import Model\n# from tensorflow.keras.applications import EfficientNetB0\n\n# # Load the pre-trained EfficientNetB0 model\n# def build_efficientnet_model(input_shape=(224, 224, 3), num_classes=10):\n#     \"\"\"\n#     Builds an EfficientNet model for classification.\n\n#     Args:\n#     - input_shape (tuple): Shape of input images, e.g., (224, 224, 3).\n#     - num_classes (int): Number of classes for classification.\n\n#     Returns:\n#     - model (tf.keras.Model): Compiled EfficientNet model.\n#     \"\"\"\n#     # Load the base EfficientNetB0 model with pre-trained ImageNet weights\n#     base_model = EfficientNetB0(include_top=False, weights='imagenet', input_shape=input_shape)\n\n#     # Freeze the base model layers (optional, for transfer learning)\n#     base_model.trainable = False\n\n#     # Add custom classification head\n#     inputs = tf.keras.Input(shape=input_shape)\n#     x = base_model(inputs, training=False)  # Use the base model\n#     x = GlobalAveragePooling2D()(x)  # Pool the features\n#     x = Dropout(0.2)(x)  # Add dropout for regularization\n#     outputs = Dense(num_classes, activation='softmax')(x)  # Final classification layer\n\n#     # Create the final model\n#     model = Model(inputs, outputs)\n\n#     # Compile the model\n#     model.compile(\n#         optimizer=tf.keras.optimizers.Adam(learning_rate=0.001),\n#         loss='sparse_categorical_crossentropy',\n#         metrics=['accuracy']\n#     )\n\n#     return model\n\n# # Example usage\n# if __name__ == \"__main__\":\n#     # Define parameters\n#     input_shape = (224, 224, 3)\n#     num_classes = 10  # Adjust according to your dataset\n\n#     # Build and summarize the model\n#     model = build_efficientnet_model(input_shape=input_shape, num_classes=num_classes)\n#     model.summary()\n\n#     # Dummy data for illustration\n#     import numpy as np\n#     X_dummy = np.random.rand(10, 224, 224, 3)\n#     y_dummy = np.random.randint(0, num_classes, size=(10,))\n\n#     # Train the model (example, for real use, replace with actual data)\n#     model.fit(X_dummy, y_dummy, epochs=3, batch_size=2)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-06T16:22:06.899861Z","iopub.status.idle":"2024-12-06T16:22:06.900209Z","shell.execute_reply.started":"2024-12-06T16:22:06.900052Z","shell.execute_reply":"2024-12-06T16:22:06.900068Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"","metadata":{}}]}