{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.10.12"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":217839134,"sourceType":"kernelVersion"}],"dockerImageVersionId":30823,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true},"papermill":{"default_parameters":{},"duration":112.414966,"end_time":"2025-01-05T22:11:02.912489","environment_variables":{},"exception":null,"input_path":"__notebook__.ipynb","output_path":"__notebook__.ipynb","parameters":{},"start_time":"2025-01-05T22:09:10.497523","version":"2.5.0"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# RSNA 2024 Lumbar Spine Degenerative Classification","metadata":{"papermill":{"duration":0.010299,"end_time":"2025-01-05T22:09:12.775779","exception":false,"start_time":"2025-01-05T22:09:12.76548","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"## Starter Notebook for Pytorch and Deep learning techniques\n\nUsing ResNET","metadata":{"papermill":{"duration":0.008839,"end_time":"2025-01-05T22:09:12.79393","exception":false,"start_time":"2025-01-05T22:09:12.785091","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"What does this notebook contains?\n\n* Data organized in an understandable and easy to use way\n* A pretrained ResNET for inference\n\nI have tried creating a notebook where you can just plug your deep learning models and everything else is sorted. ","metadata":{"papermill":{"duration":0.008888,"end_time":"2025-01-05T22:09:12.811879","exception":false,"start_time":"2025-01-05T22:09:12.802991","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":5.51141,"end_time":"2025-01-05T22:09:18.332201","exception":false,"start_time":"2025-01-05T22:09:12.820791","status":"completed"},"tags":[],"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')\nfake_test_data = False # len(sub) <= 25 # For testing purposes replace False with commented value","metadata":{"papermill":{"duration":0.168498,"end_time":"2025-01-05T22:09:18.510384","exception":false,"start_time":"2025-01-05T22:09:18.341886","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if fake_test_data:\n    # Create Fake Test data\n    n = 100\n    selected_studies = pd.Series(train_desc.study_id.unique()).sample(n).values\n    test_desc = train_desc[train_desc.study_id.isin(selected_studies)]\n    print(f\"New test_desc length: {len(test_desc)}\")\n\n    # Get 25 formats from current sub DataFrame\n    row_label_formats = [(\"{}_\" + row_id.split(\"_\", 1)[1]) for row_id in sub.row_id.tolist()]\n\n    row_ids = [\n        row_format.format(study)\n        for study in selected_studies for row_format in row_label_formats\n    ]\n    sub = pd.DataFrame(dict(row_id=row_ids, normal_mild=1/3, moderate=1/3, severe=1/3))\n    print(f\"New sub length: {len(sub)}\")","metadata":{"papermill":{"duration":0.015664,"end_time":"2025-01-05T22:09:18.535725","exception":false,"start_time":"2025-01-05T22:09:18.520061","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_desc.head(5)","metadata":{"papermill":{"duration":0.024397,"end_time":"2025-01-05T22:09:18.569211","exception":false,"start_time":"2025-01-05T22:09:18.544814","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train.head(5)","metadata":{"papermill":{"duration":0.028677,"end_time":"2025-01-05T22:09:18.607125","exception":false,"start_time":"2025-01-05T22:09:18.578448","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_desc.head(5)","metadata":{"papermill":{"duration":0.018358,"end_time":"2025-01-05T22:09:18.63524","exception":false,"start_time":"2025-01-05T22:09:18.616882","status":"completed"},"tags":[],"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, '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images')\ntest_image_paths = generate_image_paths(test_desc, f'/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/{\"train_images\" if fake_test_data else \"test_images\"}')","metadata":{"papermill":{"duration":54.083596,"end_time":"2025-01-05T22:10:12.728577","exception":false,"start_time":"2025-01-05T22:09:18.644981","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(train_desc)","metadata":{"papermill":{"duration":0.015765,"end_time":"2025-01-05T22:10:12.754557","exception":false,"start_time":"2025-01-05T22:10:12.738792","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(train_image_paths)","metadata":{"papermill":{"duration":0.016297,"end_time":"2025-01-05T22:10:12.780545","exception":false,"start_time":"2025-01-05T22:10:12.764248","status":"completed"},"tags":[],"trusted":true},"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":{"papermill":{"duration":0.533067,"end_time":"2025-01-05T22:10:13.323283","exception":false,"start_time":"2025-01-05T22:10:12.790216","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"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\"/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/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":{"papermill":{"duration":0.847805,"end_time":"2025-01-05T22:10:14.190334","exception":false,"start_time":"2025-01-05T22:10:13.342529","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Data Preprocessing","metadata":{"papermill":{"duration":0.033405,"end_time":"2025-01-05T22:10:14.259207","exception":false,"start_time":"2025-01-05T22:10:14.225802","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":1.103412,"end_time":"2025-01-05T22:10:15.395315","exception":false,"start_time":"2025-01-05T22:10:14.291903","status":"completed"},"tags":[],"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":{"papermill":{"duration":0.039995,"end_time":"2025-01-05T22:10:15.468081","exception":false,"start_time":"2025-01-05T22:10:15.428086","status":"completed"},"tags":[],"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')\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":{"papermill":{"duration":0.087615,"end_time":"2025-01-05T22:10:15.587401","exception":false,"start_time":"2025-01-05T22:10:15.499786","status":"completed"},"tags":[],"trusted":true},"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\nfinal_merged_df.head(5)","metadata":{"papermill":{"duration":0.060897,"end_time":"2025-01-05T22:10:15.680594","exception":false,"start_time":"2025-01-05T22:10:15.619697","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df['study_id'] == 100206310].sort_values(['x','y'],ascending = True)","metadata":{"papermill":{"duration":0.050794,"end_time":"2025-01-05T22:10:15.76452","exception":false,"start_time":"2025-01-05T22:10:15.713726","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df['series_id'] == 1012284084].sort_values(\"instance_number\")","metadata":{"papermill":{"duration":0.049718,"end_time":"2025-01-05T22:10:15.84787","exception":false,"start_time":"2025-01-05T22:10:15.798152","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Now, we can see what the data represents\n\nSeries ID 1012284084 contains 60 images, and how each image maps to each level and condition","metadata":{"papermill":{"duration":0.034067,"end_time":"2025-01-05T22:10:15.915276","exception":false,"start_time":"2025-01-05T22:10:15.881209","status":"completed"},"tags":[]}},{"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":{"papermill":{"duration":0.050642,"end_time":"2025-01-05T22:10:15.999215","exception":false,"start_time":"2025-01-05T22:10:15.948573","status":"completed"},"tags":[],"trusted":true},"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":{"papermill":{"duration":0.052528,"end_time":"2025-01-05T22:10:16.085494","exception":false,"start_time":"2025-01-05T22:10:16.032966","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"We see that, <br>\nSaggital T1 images map to Neural Foraminal Narrowing <br>\nAxial T2 images map to Subarticular Stenosis <br>\nSaggital T2/STIR map to Canal Stenosis <br>","metadata":{"papermill":{"duration":0.03454,"end_time":"2025-01-05T22:10:16.153988","exception":false,"start_time":"2025-01-05T22:10:16.119448","status":"completed"},"tags":[]}},{"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    '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/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":{"papermill":{"duration":0.219456,"end_time":"2025-01-05T22:10:16.407935","exception":false,"start_time":"2025-01-05T22:10:16.188479","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Normal/Mild\"].value_counts().sum()","metadata":{"papermill":{"duration":0.137834,"end_time":"2025-01-05T22:10:16.582231","exception":false,"start_time":"2025-01-05T22:10:16.444397","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Moderate\"].value_counts().sum()","metadata":{"papermill":{"duration":0.068251,"end_time":"2025-01-05T22:10:16.684617","exception":false,"start_time":"2025-01-05T22:10:16.616366","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define the base path for test images\nbase_path = f'/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/{\"train_images\" if fake_test_data else \"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":{"papermill":{"duration":0.158295,"end_time":"2025-01-05T22:10:16.876883","exception":false,"start_time":"2025-01-05T22:10:16.718588","status":"completed"},"tags":[],"trusted":true},"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'})","metadata":{"papermill":{"duration":0.043796,"end_time":"2025-01-05T22:10:16.955299","exception":false,"start_time":"2025-01-05T22:10:16.911503","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_data = expanded_test_desc\ntrain_data = final_merged_df","metadata":{"papermill":{"duration":0.039787,"end_time":"2025-01-05T22:10:17.029964","exception":false,"start_time":"2025-01-05T22:10:16.990177","status":"completed"},"tags":[],"trusted":true},"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'/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/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'/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/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":{"papermill":{"duration":25.893059,"end_time":"2025-01-05T22:10:42.957255","exception":false,"start_time":"2025-01-05T22:10:17.064196","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data.head(3)","metadata":{"papermill":{"duration":0.049236,"end_time":"2025-01-05T22:10:43.043909","exception":false,"start_time":"2025-01-05T22:10:42.994673","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\ndef load_dicom(path):\n    dicom = pydicom.dcmread(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":{"papermill":{"duration":0.04051,"end_time":"2025-01-05T22:10:43.122997","exception":false,"start_time":"2025-01-05T22:10:43.082487","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load images randomly\nimport random\nimages = []\nrow_ids = []\nselected_indices = random.sample(range(len(train_data)), 2)\nfor i in selected_indices:\n    image = load_dicom(train_data['image_path'][i])\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":{"papermill":{"duration":0.304714,"end_time":"2025-01-05T22:10:43.461998","exception":false,"start_time":"2025-01-05T22:10:43.157284","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Loading data","metadata":{"papermill":{"duration":0.037431,"end_time":"2025-01-05T22:10:43.540152","exception":false,"start_time":"2025-01-05T22:10:43.502721","status":"completed"},"tags":[]}},{"cell_type":"code","source":"#for one hot encoding\n#train_data[['normal_mild', 'severe', 'moderate']] = train_data[['normal_mild', 'severe', 'moderate']].astype(int)  ","metadata":{"papermill":{"duration":0.043326,"end_time":"2025-01-05T22:10:43.621161","exception":false,"start_time":"2025-01-05T22:10:43.577835","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data","metadata":{"papermill":{"duration":0.056406,"end_time":"2025-01-05T22:10:43.715414","exception":false,"start_time":"2025-01-05T22:10:43.659008","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = train_data.dropna()","metadata":{"papermill":{"duration":0.064206,"end_time":"2025-01-05T22:10:43.81796","exception":false,"start_time":"2025-01-05T22:10:43.753754","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"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\n# Define a custom dataset class\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        image = load_dicom(image_path)  # Define this function to load your DICOM images\n        label = self.dataframe['severity'][index]\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    \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)),\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":{"papermill":{"duration":1.898991,"end_time":"2025-01-05T22:10:45.755415","exception":false,"start_time":"2025-01-05T22:10:43.856424","status":"completed"},"tags":[],"trusted":true},"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":{"papermill":{"duration":2.288661,"end_time":"2025-01-05T22:10:48.085115","exception":false,"start_time":"2025-01-05T22:10:45.796454","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"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(images[0], cmap='gray')\nplt.title(label[0])\nplt.axis('off')\nplt.tight_layout()\nplt.show()","metadata":{"papermill":{"duration":0.460843,"end_time":"2025-01-05T22:10:48.606534","exception":false,"start_time":"2025-01-05T22:10:48.145691","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Model","metadata":{"papermill":{"duration":0.058312,"end_time":"2025-01-05T22:10:48.726939","exception":false,"start_time":"2025-01-05T22:10:48.668627","status":"completed"},"tags":[]}},{"cell_type":"markdown","source":"ConvNext","metadata":{"papermill":{"duration":0.057014,"end_time":"2025-01-05T22:10:48.842599","exception":false,"start_time":"2025-01-05T22:10:48.785585","status":"completed"},"tags":[]}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\nfrom torchvision import transforms\nfrom torch.utils.data import DataLoader\nfrom sklearn.model_selection import train_test_split\nimport pandas as pd\nfrom tqdm import tqdm\n\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")","metadata":{"papermill":{"duration":0.109842,"end_time":"2025-01-05T22:10:49.010023","exception":false,"start_time":"2025-01-05T22:10:48.900181","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch.nn as nn\nimport torchvision.models as models\n\nclass CustomResNet(nn.Module):\n    def __init__(self, num_classes, pretrained_weights=None):\n        super(CustomResNet, self).__init__()\n        # ResNet-50 modelini yükle\n        self.model = models.resnet50(pretrained=False)\n        \n        # Son katmanı yeniden tanımla\n        self.model.fc = nn.Linear(self.model.fc.in_features, num_classes)\n        \n        # Eğer özel bir ağırlık dosyası varsa yükle\n        if pretrained_weights:\n            state_dict = torch.load(pretrained_weights, map_location=torch.device(\"cpu\"))\n            # Son katmanı çıkartarak ağırlıkları yükle\n            state_dict.pop(\"fc.weight\", None)\n            state_dict.pop(\"fc.bias\", None)\n            self.model.load_state_dict(state_dict, strict=False)\n\n    def forward(self, x):\n        return self.model(x)\n","metadata":{"papermill":{"duration":0.063238,"end_time":"2025-01-05T22:10:49.131654","exception":false,"start_time":"2025-01-05T22:10:49.068416","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Path to the locally uploaded weights file\nweights_path = '/kaggle/input/ibo-birlestirme/combined_model.pth'\n\n# Initialize models\nsagittal_t1_model = CustomResNet(num_classes=3, pretrained_weights=weights_path).to(device)\naxial_t2_model = CustomResNet(num_classes=3, pretrained_weights=weights_path).to(device)\nsagittal_t2stir_model = CustomResNet(num_classes=3, pretrained_weights=weights_path).to(device)\n\n# Optionally freeze initial layers\nfor param in sagittal_t1_model.model.parameters():\n    param.requires_grad = False\nfor param in axial_t2_model.model.parameters():\n    param.requires_grad = False\nfor param in sagittal_t2stir_model.model.parameters():\n    param.requires_grad = False\n\n# Unfreeze the final fully connected layer\nfor param in sagittal_t1_model.model.fc.parameters():\n    param.requires_grad = True\nfor param in axial_t2_model.model.fc.parameters():\n    param.requires_grad = True\nfor param in sagittal_t2stir_model.model.fc.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.fc.parameters(), lr=0.001)\noptimizer_axial_t2 = torch.optim.Adam(axial_t2_model.model.fc.parameters(), lr=0.001)\noptimizer_sagittal_t2stir = torch.optim.Adam(sagittal_t2stir_model.model.fc.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}\n\noptimizers = {\n    'Sagittal T1': optimizer_sagittal_t1,\n    'Axial T2': optimizer_axial_t2,\n    'Sagittal T2/STIR': optimizer_sagittal_t2stir,\n}","metadata":{"papermill":{"duration":2.79806,"end_time":"2025-01-05T22:10:51.987043","exception":false,"start_time":"2025-01-05T22:10:49.188983","status":"completed"},"tags":[],"trusted":true},"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":{"papermill":{"duration":0.06572,"end_time":"2025-01-05T22:10:52.11089","exception":false,"start_time":"2025-01-05T22:10:52.04517","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Training","metadata":{"papermill":{"duration":0.057578,"end_time":"2025-01-05T22:10:52.225722","exception":false,"start_time":"2025-01-05T22:10:52.168144","status":"completed"},"tags":[]}},{"cell_type":"code","source":"label_map = {'normal_mild': 0, 'moderate': 1, 'severe': 2}","metadata":{"papermill":{"duration":0.06347,"end_time":"2025-01-05T22:10:52.346747","exception":false,"start_time":"2025-01-05T22:10:52.283277","status":"completed"},"tags":[],"trusted":true},"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","metadata":{"papermill":{"duration":0.194014,"end_time":"2025-01-05T22:10:52.598459","exception":false,"start_time":"2025-01-05T22:10:52.404445","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"##TODO","metadata":{"papermill":{"duration":0.063894,"end_time":"2025-01-05T22:10:52.721304","exception":false,"start_time":"2025-01-05T22:10:52.65741","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## Inference","metadata":{"papermill":{"duration":0.057003,"end_time":"2025-01-05T22:10:52.836427","exception":false,"start_time":"2025-01-05T22:10:52.779424","status":"completed"},"tags":[]}},{"cell_type":"code","source":"train_data['level'].unique()","metadata":{"papermill":{"duration":0.069673,"end_time":"2025-01-05T22:10:52.964475","exception":false,"start_time":"2025-01-05T22:10:52.894802","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"expanded_test_desc.head(5)","metadata":{"papermill":{"duration":0.071113,"end_time":"2025-01-05T22:10:53.093795","exception":false,"start_time":"2025-01-05T22:10:53.022682","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"levels = ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\n\n# Function to update row_id with levels\ndef update_row_id(row, levels):\n    level = levels[row.name % len(levels)]\n    return f\"{row['study_id']}_{row['condition']}_{level}\"\n\n# Update row_id in expanded_test_desc to include levels\nexpanded_test_desc['row_id'] = expanded_test_desc.apply(lambda row: update_row_id(row, levels), axis=1)","metadata":{"papermill":{"duration":0.070508,"end_time":"2025-01-05T22:10:53.222695","exception":false,"start_time":"2025-01-05T22:10:53.152187","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"expanded_test_desc.head(2)","metadata":{"papermill":{"duration":0.069989,"end_time":"2025-01-05T22:10:53.351227","exception":false,"start_time":"2025-01-05T22:10:53.281238","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define a custom test dataset class\nclass TestDataset(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        image = load_dicom(image_path)  # Define this function to load your DICOM images\n        if self.transform:\n            image = self.transform(image)\n        return image\n\n# Define the transforms\ntransform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((224, 224)),\n    transforms.Grayscale(num_output_channels=3),\n    transforms.ToTensor(),\n])\n\n# Create a test dataset and dataloader\ntest_dataset = TestDataset(expanded_test_desc, transform)\ntestloader = DataLoader(test_dataset, batch_size=1, shuffle=False)","metadata":{"papermill":{"duration":0.067707,"end_time":"2025-01-05T22:10:53.480024","exception":false,"start_time":"2025-01-05T22:10:53.412317","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for image in testloader:\n    print(image.shape)\n    break","metadata":{"papermill":{"duration":0.093497,"end_time":"2025-01-05T22:10:53.631297","exception":false,"start_time":"2025-01-05T22:10:53.5378","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to get the model based on series_description\ndef get_model(series_description):\n    return models.get(series_description, None)\n\n# Function to make predictions on the test data\ndef predict_test_data(testloader, expanded_test_desc):\n    predictions = []\n    normal_mild_probs = []\n    moderate_probs = []\n    severe_probs = []\n    \n    with torch.no_grad():\n        for idx, images in enumerate(tqdm(testloader)):\n            images = images.to(device)\n            series_description = expanded_test_desc.iloc[idx]['series_description']\n            model = get_model(series_description)\n            if model:\n                outputs = model(images)\n                probs = torch.softmax(outputs, dim=1).squeeze(0)\n                normal_mild_probs.append(probs[0].item())\n                moderate_probs.append(probs[1].item())\n                severe_probs.append(probs[2].item())\n                predictions.append(probs)\n            else:\n                normal_mild_probs.append(None)\n                moderate_probs.append(None)\n                severe_probs.append(None)\n                predictions.append(None)\n\n    return normal_mild_probs, moderate_probs, severe_probs, predictions","metadata":{"papermill":{"duration":0.066829,"end_time":"2025-01-05T22:10:53.758662","exception":false,"start_time":"2025-01-05T22:10:53.691833","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Make predictions on the test data\nnormal_mild_probs, moderate_probs, severe_probs, test_predictions = predict_test_data(testloader, expanded_test_desc)","metadata":{"papermill":{"duration":6.024146,"end_time":"2025-01-05T22:10:59.841095","exception":false,"start_time":"2025-01-05T22:10:53.816949","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_predictions[0]","metadata":{"papermill":{"duration":0.319344,"end_time":"2025-01-05T22:11:00.223094","exception":false,"start_time":"2025-01-05T22:10:59.90375","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Add predictions and probabilities to the test DataFrame\nexpanded_test_desc['normal_mild'] = normal_mild_probs\nexpanded_test_desc['moderate'] = moderate_probs\nexpanded_test_desc['severe'] = severe_probs","metadata":{"papermill":{"duration":0.070012,"end_time":"2025-01-05T22:11:00.354363","exception":false,"start_time":"2025-01-05T22:11:00.284351","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission = expanded_test_desc[[\"row_id\",\"normal_mild\",\"moderate\",\"severe\"]]","metadata":{"papermill":{"duration":0.068172,"end_time":"2025-01-05T22:11:00.48393","exception":false,"start_time":"2025-01-05T22:11:00.415758","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"submission.head(10)","metadata":{"papermill":{"duration":0.073706,"end_time":"2025-01-05T22:11:00.620262","exception":false,"start_time":"2025-01-05T22:11:00.546556","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Group by 'row_id' and sum the values\ngrouped_submission = submission.groupby('row_id').sum().reset_index()\n\n# Normalize the columns\ngrouped_submission[['normal_mild', 'moderate', 'severe']] = grouped_submission[['normal_mild', 'moderate', 'severe']].div(grouped_submission[['normal_mild', 'moderate', 'severe']].sum(axis=1), axis=0)\n\n# Check the first 3 rows\ngrouped_submission.head(3)\n\ngrouped_submission.head(3)","metadata":{"papermill":{"duration":0.080812,"end_time":"2025-01-05T22:11:00.763327","exception":false,"start_time":"2025-01-05T22:11:00.682515","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(grouped_submission)","metadata":{"papermill":{"duration":0.069213,"end_time":"2025-01-05T22:11:00.894222","exception":false,"start_time":"2025-01-05T22:11:00.825009","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Save the DataFrame to \"submission.csv\" in the desired directory\ngrouped_submission.to_csv(\"/kaggle/working/submission.csv\", index=False)","metadata":{"papermill":{"duration":0.070561,"end_time":"2025-01-05T22:11:01.02753","exception":false,"start_time":"2025-01-05T22:11:00.956969","status":"completed"},"tags":[],"trusted":true},"outputs":[],"execution_count":null}]}