{"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":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":163026,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":138643,"modelId":161263},{"sourceId":163101,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":138709,"modelId":161328},{"sourceId":163103,"sourceType":"modelInstanceVersion","isSourceIdPinned":true,"modelInstanceId":138711,"modelId":161330}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"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\n# import torch\n# import 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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-11-13T04:16:24.082836Z","iopub.execute_input":"2024-11-13T04:16:24.083213Z","iopub.status.idle":"2024-11-13T04:16:25.483788Z","shell.execute_reply.started":"2024-11-13T04:16:24.083175Z","shell.execute_reply":"2024-11-13T04:16:25.482520Z"},"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')\n\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')","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:16:25.485824Z","iopub.execute_input":"2024-11-13T04:16:25.486959Z","iopub.status.idle":"2024-11-13T04:16:25.680357Z","shell.execute_reply.started":"2024-11-13T04:16:25.486913Z","shell.execute_reply":"2024-11-13T04:16:25.679115Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_desc.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:16:25.681945Z","iopub.execute_input":"2024-11-13T04:16:25.682341Z","iopub.status.idle":"2024-11-13T04:16:25.704344Z","shell.execute_reply.started":"2024-11-13T04:16:25.682294Z","shell.execute_reply":"2024-11-13T04:16:25.703040Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\n\ndirectory = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/train_images'\nfolders = [d for d in os.listdir(directory) if os.path.isdir(os.path.join(directory, d))]\nprint(len(folders))\n","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:16:25.708349Z","iopub.execute_input":"2024-11-13T04:16:25.708810Z","iopub.status.idle":"2024-11-13T04:16:26.794894Z","shell.execute_reply.started":"2024-11-13T04:16:25.708752Z","shell.execute_reply":"2024-11-13T04:16:26.793940Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"so, there are 1975 study ids (1975 different people scan folders)                                                and each of it contain 3 diffrent series ids (3 types of scans folders),                                each of series id folder contain different scans with instance numbers","metadata":{}},{"cell_type":"markdown","source":"of these instance number's images will be choosen from \"train_label_coordinates.csv\", by mapping series ids.","metadata":{}},{"cell_type":"markdown","source":"so, each series id folder may contain, min of 1 image to max of any.","metadata":{}},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:16:26.797804Z","iopub.execute_input":"2024-11-13T04:16:26.798117Z","iopub.status.idle":"2024-11-13T04:16:26.821785Z","shell.execute_reply.started":"2024-11-13T04:16:26.798084Z","shell.execute_reply":"2024-11-13T04:16:26.820877Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_desc.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:16:26.823011Z","iopub.execute_input":"2024-11-13T04:16:26.823519Z","iopub.status.idle":"2024-11-13T04:16:26.834587Z","shell.execute_reply.started":"2024-11-13T04:16:26.823451Z","shell.execute_reply":"2024-11-13T04:16:26.833652Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_desc.shape","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:16:26.835605Z","iopub.execute_input":"2024-11-13T04:16:26.835892Z","iopub.status.idle":"2024-11-13T04:16:26.845562Z","shell.execute_reply.started":"2024-11-13T04:16:26.835853Z","shell.execute_reply":"2024-11-13T04:16:26.844738Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Function to generate image paths based on directory structure\ndef generate_image_paths(df, data_dir):\n    image_paths = []\n    for study_id, series_id in zip(df['study_id'], df['series_id']):\n        study_dir = os.path.join(data_dir, str(study_id))\n        series_dir = os.path.join(study_dir, str(series_id))\n        images = os.listdir(series_dir)\n        image_paths.extend([os.path.join(series_dir, img) for img in images])\n    return image_paths\n\n# Generate image paths for train and test data\ntrain_image_paths = generate_image_paths(train_desc, f'{train_path}/train_images')\ntest_image_paths = generate_image_paths(test_desc, f'{train_path}/test_images')","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:16:26.846845Z","iopub.execute_input":"2024-11-13T04:16:26.847192Z","iopub.status.idle":"2024-11-13T04:17:31.041299Z","shell.execute_reply.started":"2024-11-13T04:16:26.847152Z","shell.execute_reply":"2024-11-13T04:17:31.040544Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(train_image_paths[2])\n","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:31.042327Z","iopub.execute_input":"2024-11-13T04:17:31.042633Z","iopub.status.idle":"2024-11-13T04:17:31.047463Z","shell.execute_reply.started":"2024-11-13T04:17:31.042601Z","shell.execute_reply":"2024-11-13T04:17:31.046488Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"len(train_desc)","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:31.048536Z","iopub.execute_input":"2024-11-13T04:17:31.048785Z","iopub.status.idle":"2024-11-13T04:17:31.060246Z","shell.execute_reply.started":"2024-11-13T04:17:31.048757Z","shell.execute_reply":"2024-11-13T04:17:31.059388Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_desc","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:31.061269Z","iopub.execute_input":"2024-11-13T04:17:31.061558Z","iopub.status.idle":"2024-11-13T04:17:31.075611Z","shell.execute_reply.started":"2024-11-13T04:17:31.061528Z","shell.execute_reply":"2024-11-13T04:17:31.074725Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"(train_desc['study_id']).shape[0]/3","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:31.076667Z","iopub.execute_input":"2024-11-13T04:17:31.077013Z","iopub.status.idle":"2024-11-13T04:17:31.084455Z","shell.execute_reply.started":"2024-11-13T04:17:31.076980Z","shell.execute_reply":"2024-11-13T04:17:31.083595Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_desc['study_id'].unique().shape","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:31.085607Z","iopub.execute_input":"2024-11-13T04:17:31.085955Z","iopub.status.idle":"2024-11-13T04:17:31.096048Z","shell.execute_reply.started":"2024-11-13T04:17:31.085908Z","shell.execute_reply":"2024-11-13T04:17:31.095183Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"getting 123 duplicate records(means 123 people scan folders were extra)","metadata":{}},{"cell_type":"code","source":"len(train_image_paths)","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:31.100506Z","iopub.execute_input":"2024-11-13T04:17:31.100788Z","iopub.status.idle":"2024-11-13T04:17:31.106303Z","shell.execute_reply.started":"2024-11-13T04:17:31.100752Z","shell.execute_reply":"2024-11-13T04:17:31.105410Z"},"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:31.107477Z","iopub.execute_input":"2024-11-13T04:17:31.107815Z","iopub.status.idle":"2024-11-13T04:17:31.701671Z","shell.execute_reply.started":"2024-11-13T04:17:31.107775Z","shell.execute_reply":"2024-11-13T04:17:31.700686Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# visualizing 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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:31.703239Z","iopub.execute_input":"2024-11-13T04:17:31.703807Z","iopub.status.idle":"2024-11-13T04:17:32.491460Z","shell.execute_reply.started":"2024-11-13T04:17:31.703750Z","shell.execute_reply":"2024-11-13T04:17:32.490561Z"},"trusted":true},"outputs":[],"execution_count":null},{"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:32.492853Z","iopub.execute_input":"2024-11-13T04:17:32.493152Z","iopub.status.idle":"2024-11-13T04:17:33.887419Z","shell.execute_reply.started":"2024-11-13T04:17:32.493118Z","shell.execute_reply":"2024-11-13T04:17:33.886476Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Print columns in a neat way\nprint(\"\\nColumns in new_train_df:\")\nprint(\",\".join(new_train_df.columns))\n\nprint(\"\\nColumns in label:\")\nprint(\",\".join(label.columns))\n\nprint(\"\\nColumns in test_desc:\")\nprint(\",\".join(test_desc.columns))\n\nprint(\"\\nColumns in sub:\")\nprint(\",\".join(sub.columns))","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:33.888755Z","iopub.execute_input":"2024-11-13T04:17:33.889168Z","iopub.status.idle":"2024-11-13T04:17:33.895631Z","shell.execute_reply.started":"2024-11-13T04:17:33.889134Z","shell.execute_reply":"2024-11-13T04:17:33.894588Z"},"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:33.896838Z","iopub.execute_input":"2024-11-13T04:17:33.897148Z","iopub.status.idle":"2024-11-13T04:17:33.974921Z","shell.execute_reply.started":"2024-11-13T04:17:33.897114Z","shell.execute_reply":"2024-11-13T04:17:33.973876Z"},"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:33.976333Z","iopub.execute_input":"2024-11-13T04:17:33.976833Z","iopub.status.idle":"2024-11-13T04:17:34.002225Z","shell.execute_reply.started":"2024-11-13T04:17:33.976798Z","shell.execute_reply":"2024-11-13T04:17:34.001305Z"},"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.003365Z","iopub.execute_input":"2024-11-13T04:17:34.003742Z","iopub.status.idle":"2024-11-13T04:17:34.025621Z","shell.execute_reply.started":"2024-11-13T04:17:34.003709Z","shell.execute_reply":"2024-11-13T04:17:34.024673Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df['series_id'] == 1012284084].sort_values(\"instance_number\")","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.026912Z","iopub.execute_input":"2024-11-13T04:17:34.027509Z","iopub.status.idle":"2024-11-13T04:17:34.045750Z","shell.execute_reply.started":"2024-11-13T04:17:34.027462Z","shell.execute_reply":"2024-11-13T04:17:34.044658Z"},"trusted":true},"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.046892Z","iopub.execute_input":"2024-11-13T04:17:34.047181Z","iopub.status.idle":"2024-11-13T04:17:34.068337Z","shell.execute_reply.started":"2024-11-13T04:17:34.047147Z","shell.execute_reply":"2024-11-13T04:17:34.067253Z"},"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.069714Z","iopub.execute_input":"2024-11-13T04:17:34.070083Z","iopub.status.idle":"2024-11-13T04:17:34.096680Z","shell.execute_reply.started":"2024-11-13T04:17:34.070044Z","shell.execute_reply":"2024-11-13T04:17:34.095686Z"},"trusted":true},"outputs":[],"execution_count":null},{"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.097856Z","iopub.execute_input":"2024-11-13T04:17:34.098191Z","iopub.status.idle":"2024-11-13T04:17:34.343956Z","shell.execute_reply.started":"2024-11-13T04:17:34.098152Z","shell.execute_reply":"2024-11-13T04:17:34.343087Z"},"trusted":true},"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-11-13T04:17:34.345581Z","iopub.execute_input":"2024-11-13T04:17:34.345968Z","iopub.status.idle":"2024-11-13T04:17:34.498469Z","shell.execute_reply.started":"2024-11-13T04:17:34.345923Z","shell.execute_reply":"2024-11-13T04:17:34.497560Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Moderate\"].value_counts().sum()","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.499733Z","iopub.execute_input":"2024-11-13T04:17:34.500104Z","iopub.status.idle":"2024-11-13T04:17:34.547842Z","shell.execute_reply.started":"2024-11-13T04:17:34.500062Z","shell.execute_reply":"2024-11-13T04:17:34.546982Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Severe\"].value_counts().sum()","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.549077Z","iopub.execute_input":"2024-11-13T04:17:34.549689Z","iopub.status.idle":"2024-11-13T04:17:34.580034Z","shell.execute_reply.started":"2024-11-13T04:17:34.549644Z","shell.execute_reply":"2024-11-13T04:17:34.579211Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"final_merged_df.shape","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.581016Z","iopub.execute_input":"2024-11-13T04:17:34.581275Z","iopub.status.idle":"2024-11-13T04:17:34.587029Z","shell.execute_reply.started":"2024-11-13T04:17:34.581246Z","shell.execute_reply":"2024-11-13T04:17:34.586183Z"},"trusted":true},"outputs":[],"execution_count":null},{"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.588251Z","iopub.execute_input":"2024-11-13T04:17:34.588596Z","iopub.status.idle":"2024-11-13T04:17:34.653670Z","shell.execute_reply.started":"2024-11-13T04:17:34.588562Z","shell.execute_reply":"2024-11-13T04:17:34.652763Z"},"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.655164Z","iopub.execute_input":"2024-11-13T04:17:34.655551Z","iopub.status.idle":"2024-11-13T04:17:34.666418Z","shell.execute_reply.started":"2024-11-13T04:17:34.655508Z","shell.execute_reply":"2024-11-13T04:17:34.665586Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_data = expanded_test_desc\ntrain_data = final_merged_df","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.667507Z","iopub.execute_input":"2024-11-13T04:17:34.667795Z","iopub.status.idle":"2024-11-13T04:17:34.675712Z","shell.execute_reply.started":"2024-11-13T04:17:34.667760Z","shell.execute_reply":"2024-11-13T04:17:34.674907Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data.head(10)","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.676702Z","iopub.execute_input":"2024-11-13T04:17:34.676970Z","iopub.status.idle":"2024-11-13T04:17:34.699223Z","shell.execute_reply.started":"2024-11-13T04:17:34.676940Z","shell.execute_reply":"2024-11-13T04:17:34.698265Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_data.head(10)","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.700366Z","iopub.execute_input":"2024-11-13T04:17:34.700678Z","iopub.status.idle":"2024-11-13T04:17:34.713701Z","shell.execute_reply.started":"2024-11-13T04:17:34.700647Z","shell.execute_reply":"2024-11-13T04:17:34.712843Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data.shape","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.714926Z","iopub.execute_input":"2024-11-13T04:17:34.715268Z","iopub.status.idle":"2024-11-13T04:17:34.722561Z","shell.execute_reply.started":"2024-11-13T04:17:34.715228Z","shell.execute_reply":"2024-11-13T04:17:34.721673Z"},"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'{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'])]\ntrain_data.shape","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:34.723824Z","iopub.execute_input":"2024-11-13T04:17:34.724116Z","iopub.status.idle":"2024-11-13T04:17:51.205383Z","shell.execute_reply.started":"2024-11-13T04:17:34.724069Z","shell.execute_reply":"2024-11-13T04:17:51.204505Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data.head(3)","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:51.206370Z","iopub.execute_input":"2024-11-13T04:17:51.206650Z","iopub.status.idle":"2024-11-13T04:17:51.220625Z","shell.execute_reply.started":"2024-11-13T04:17:51.206620Z","shell.execute_reply":"2024-11-13T04:17:51.219776Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pydicom\ndef load_dicom(path):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:51.221758Z","iopub.execute_input":"2024-11-13T04:17:51.222059Z","iopub.status.idle":"2024-11-13T04:17:51.231166Z","shell.execute_reply.started":"2024-11-13T04:17:51.222028Z","shell.execute_reply":"2024-11-13T04:17:51.230372Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import random\nimport matplotlib.pyplot as plt  # Make sure to import matplotlib\nimport pydicom\nimport numpy as np\n\n# Define the load_dicom function\ndef load_dicom(path):\n    dicom = pydicom.dcmread(path)  # Correct method to read DICOM files\n    data = dicom.pixel_array\n    data = data - np.min(data)  # Normalization (if needed)\n    return data\n\n# Load images randomly\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()\n","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:51.232287Z","iopub.execute_input":"2024-11-13T04:17:51.232717Z","iopub.status.idle":"2024-11-13T04:17:51.544777Z","shell.execute_reply.started":"2024-11-13T04:17:51.232676Z","shell.execute_reply":"2024-11-13T04:17:51.543940Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:51.545887Z","iopub.execute_input":"2024-11-13T04:17:51.546164Z","iopub.status.idle":"2024-11-13T04:17:51.565579Z","shell.execute_reply.started":"2024-11-13T04:17:51.546135Z","shell.execute_reply":"2024-11-13T04:17:51.564581Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_data = train_data.dropna()","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:51.571651Z","iopub.execute_input":"2024-11-13T04:17:51.572131Z","iopub.status.idle":"2024-11-13T04:17:51.610578Z","shell.execute_reply.started":"2024-11-13T04:17:51.572098Z","shell.execute_reply":"2024-11-13T04:17:51.609817Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# train_test_splits and loaders ","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 numpy as np\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, temp_df = train_test_split(filtered_df, test_size=0.3, random_state=42)  # 70% train\n    val_df, test_df = train_test_split(temp_df, test_size=0.5, random_state=42)  # 15% val, 15% test\n    train_df = train_df.reset_index(drop=True)\n    val_df = val_df.reset_index(drop=True)\n    test_df = test_df.reset_index(drop=True)\n\n    train_dataset = CustomDataset(train_df, transform)\n    val_dataset = CustomDataset(val_df, transform)\n    test_dataset = CustomDataset(test_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    testloader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False)  # Create the test DataLoader\n    \n    return trainloader, valloader, testloader, train_df, val_df, test_df, len(train_df), len(val_df), len(test_df)  # Return test loader and lengths\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\n# Create loaders for Sagittal T1\ntrainloader_t1, valloader_t1, testloader_t1,train_df_t1, val_df_t1, test_df_t1,len_train_t1, len_val_t1, len_test_t1 = create_datasets_and_loaders(train_data, 'Sagittal T1', transform)\n# Create loaders for Axial T2\ntrainloader_t2, valloader_t2, testloader_t2,train_df_t2, val_df_t2, test_df_t2, len_train_t2, len_val_t2, len_test_t2 = create_datasets_and_loaders(train_data, 'Axial T2', transform)\n# Create loaders for Sagittal T2/STIR\ntrainloader_t2stir, valloader_t2stir, testloader_t2stir, train_df_t2stir, val_df_t2stir, test_df_t2stir, len_train_t2stir, len_val_t2stir, len_test_t2stir = create_datasets_and_loaders(train_data, 'Sagittal T2/STIR', transform)\n\n# Store the loaders in the dataloaders dictionary\ndataloaders['Sagittal T1'] = (trainloader_t1, valloader_t1, testloader_t1)\ndataloaders['Axial T2'] = (trainloader_t2, valloader_t2, testloader_t2)\ndataloaders['Sagittal T2/STIR'] = (trainloader_t2stir, valloader_t2stir, testloader_t2stir)\n\n# Store the lengths in the lengths dictionary\nlengths['Sagittal T1'] = (len_train_t1, len_val_t1, len_test_t1)\nlengths['Axial T2'] = (len_train_t2, len_val_t2, len_test_t2)\nlengths['Sagittal T2/STIR'] = (len_train_t2stir, len_val_t2stir, len_test_t2stir)\n\n# Dictionary mapping labels to indices\nlabel_map = {'Mild': 0, 'Moderate': 1, 'Severe': 2}","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:51.611894Z","iopub.execute_input":"2024-11-13T04:17:51.612190Z","iopub.status.idle":"2024-11-13T04:17:56.285143Z","shell.execute_reply.started":"2024-11-13T04:17:51.612157Z","shell.execute_reply":"2024-11-13T04:17:56.284304Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_df_t2stir","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:17:56.286242Z","iopub.execute_input":"2024-11-13T04:17:56.286751Z","iopub.status.idle":"2024-11-13T04:17:56.306853Z","shell.execute_reply.started":"2024-11-13T04:17:56.286715Z","shell.execute_reply":"2024-11-13T04:17:56.305885Z"},"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:56.308154Z","iopub.execute_input":"2024-11-13T04:17:56.308496Z","iopub.status.idle":"2024-11-13T04:17:59.213006Z","shell.execute_reply.started":"2024-11-13T04:17:56.308456Z","shell.execute_reply":"2024-11-13T04:17:59.212079Z"},"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":{"execution":{"iopub.status.busy":"2024-11-13T04:17:59.214471Z","iopub.execute_input":"2024-11-13T04:17:59.214839Z","iopub.status.idle":"2024-11-13T04:17:59.639278Z","shell.execute_reply.started":"2024-11-13T04:17:59.214800Z","shell.execute_reply":"2024-11-13T04:17:59.638332Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# model training","metadata":{}},{"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":{"execution":{"iopub.status.busy":"2024-11-13T04:27:13.382080Z","iopub.execute_input":"2024-11-13T04:27:13.382481Z","iopub.status.idle":"2024-11-13T04:27:13.412942Z","shell.execute_reply.started":"2024-11-13T04:27:13.382442Z","shell.execute_reply":"2024-11-13T04:27:13.411957Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\n\nclass EfficientNetV2(nn.Module):\n    def __init__(self, num_classes=3):\n        super(EfficientNetV2, self).__init__()\n        self.model = models.efficientnet_v2_s(weights=None) \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\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Initialize models without pretrained weights\nsagittal_t1_model = EfficientNetV2(num_classes=3).to(device)\naxial_t2_model = EfficientNetV2(num_classes=3).to(device)\nsagittal_t2stir_model = EfficientNetV2(num_classes=3).to(device)\n\n# No layer freezing since we're training from scratch\n# All layers are trainable by default.\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}\n","metadata":{"execution":{"iopub.status.busy":"2024-11-02T12:47:21.501672Z","iopub.execute_input":"2024-11-02T12:47:21.502026Z","iopub.status.idle":"2024-11-02T12:47:23.100369Z","shell.execute_reply.started":"2024-11-02T12:47:21.501992Z","shell.execute_reply":"2024-11-02T12:47:23.099312Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"label_map = {'normal_mild': 0, 'moderate': 1, 'severe': 2}","metadata":{"execution":{"iopub.status.busy":"2024-11-02T12:47:23.101904Z","iopub.execute_input":"2024-11-02T12:47:23.102589Z","iopub.status.idle":"2024-11-02T12:47:23.107015Z","shell.execute_reply.started":"2024-11-02T12:47:23.102554Z","shell.execute_reply":"2024-11-02T12:47:23.106073Z"},"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":{"execution":{"iopub.status.busy":"2024-11-02T12:47:23.108319Z","iopub.execute_input":"2024-11-02T12:47:23.108688Z","iopub.status.idle":"2024-11-02T12:47:23.301035Z","shell.execute_reply.started":"2024-11-02T12:47:23.108655Z","shell.execute_reply":"2024-11-02T12:47:23.300021Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom copy import deepcopy\nimport torch.nn.functional as F\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 = F.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 = F.softmax(outputs, dim=1)\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            # Save the model architecture and weights\n            torch.save({\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'best_val_acc': best_val_acc,\n            }, f'best_model_epoch_{epoch+1}.pth')\n            counter = 0\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\n","metadata":{"execution":{"iopub.status.busy":"2024-11-02T12:47:23.302692Z","iopub.execute_input":"2024-11-02T12:47:23.303373Z","iopub.status.idle":"2024-11-02T12:47:23.320080Z","shell.execute_reply.started":"2024-11-02T12:47:23.303326Z","shell.execute_reply":"2024-11-02T12:47:23.319134Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define a function to test the model\n# def test_model(model, testloader, criterion):\n#     model.eval()  # Set the model to evaluation mode\n#     total_loss = 0.0\n#     correct_predictions = 0\n#     total_samples = 0\n    \n#     with torch.no_grad():  # Disable gradient calculation for testing\n#         for batch in testloader:\n#             # Check if the batch is a tuple and unpack accordingly\n#             if isinstance(batch, tuple) and len(batch) == 2:\n#                 images, labels = batch\n#             else:\n#                 continue  # Skip if the batch structure is not as expected\n\n#             images, labels = images.to(device), labels.to(device)  # Move to device\n            \n#             outputs = model(images)  # Get model predictions\n#             loss = criterion(outputs, labels)  # Calculate loss\n            \n#             total_loss += loss.item()\n#             _, predicted = torch.max(outputs.data, 1)  # Get the index of the max log-probability\n#             correct_predictions += (predicted == labels).sum().item()\n#             total_samples += labels.size(0)\n    \n#     avg_loss = total_loss / len(testloader)\n#     accuracy = correct_predictions / total_samples\n#     print(f'Test Loss: {avg_loss:.4f}, Test Accuracy: {accuracy:.4f}')\n\n# Training and testing all models\nfor desc, model in models.items():\n    if desc == 'Sagittal T1':\n        trainloader, valloader, testloader, len_train, len_val = trainloader_t1, valloader_t1, testloader_t1, len_train_t1, len_val_t1\n    elif desc == 'Axial T2':\n        trainloader, valloader, testloader, len_train, len_val = trainloader_t2, valloader_t2, testloader_t2, len_train_t2, len_val_t2\n    elif desc == 'Sagittal T2/STIR':\n        trainloader, valloader, testloader, len_train, len_val = trainloader_t2stir, valloader_t2stir, testloader_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#     print(f\"Testing model for {desc}\")\n#     test_model(model, testloader, criterion)  # Test the model after training\n","metadata":{"execution":{"iopub.status.busy":"2024-11-02T12:47:23.321564Z","iopub.execute_input":"2024-11-02T12:47:23.322011Z","iopub.status.idle":"2024-11-02T15:07:38.763584Z","shell.execute_reply.started":"2024-11-02T12:47:23.321969Z","shell.execute_reply":"2024-11-02T15:07:38.762381Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import os\nimport shutil\n\n# Path to the working directory\nworking_dir = '/kaggle/working/'\n\n# Delete all files and folders in the working directory\nfor filename in os.listdir(working_dir):\n    file_path = os.path.join(working_dir, filename)\n    try:\n        if os.path.isdir(file_path):\n            shutil.rmtree(file_path)  # Remove directory\n        else:\n            os.remove(file_path)  # Remove file\n    except Exception as e:\n        print(f\"Error deleting {file_path}: {e}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-11-02T15:32:36.194259Z","iopub.execute_input":"2024-11-02T15:32:36.194659Z","iopub.status.idle":"2024-11-02T15:32:36.242318Z","shell.execute_reply.started":"2024-11-02T15:32:36.194622Z","shell.execute_reply":"2024-11-02T15:32:36.241540Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# saving .pth","metadata":{}},{"cell_type":"code","source":"torch.save(sagittal_t1_model.state_dict(), 'model1_weights.pth')\ntorch.save(axial_t2_model.state_dict(), 'model2_weights.pth')\ntorch.save(sagittal_t2stir_model.state_dict(), 'model3_weights.pth')","metadata":{"execution":{"iopub.status.busy":"2024-11-02T16:22:07.373749Z","iopub.execute_input":"2024-11-02T16:22:07.374172Z","iopub.status.idle":"2024-11-02T16:22:07.918605Z","shell.execute_reply.started":"2024-11-02T16:22:07.374134Z","shell.execute_reply":"2024-11-02T16:22:07.917802Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_df_t1.head(20)","metadata":{"execution":{"iopub.status.busy":"2024-11-02T16:19:28.638356Z","iopub.execute_input":"2024-11-02T16:19:28.639031Z","iopub.status.idle":"2024-11-02T16:19:28.661863Z","shell.execute_reply.started":"2024-11-02T16:19:28.638991Z","shell.execute_reply":"2024-11-02T16:19:28.660831Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Testing","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport torch\nfrom torchvision import transforms\nfrom PIL import Image\nimport pydicom\nimport numpy as np\n\n# Step 1: Mapping severity to integers\nseverity_mapping = {'normal_mild': 0, 'moderate': 1, 'severe': 2}  # Adjust based on your severity levels\ntest_df_t1['mapped_severity'] = test_df_t1['severity'].map(severity_mapping)\n\n# Step 2: Filter DataFrame for a specific series description\nspecific_series_description = 'Sagittal T1'  # Replace with your criteria\nfiltered_df = test_df_t1[test_df_t1['series_description'] == specific_series_description]\n\n# Step 3: Load images using the generator\nimages = []\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),  # Adjust size as needed\n    transforms.ToTensor()            # Convert image to tensor\n])\n\n\nclass DICOMDataset(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, idx):\n        image_path = self.dataframe.iloc[idx]['image_path']\n        image = load_dicom_image(image_path)\n        if self.transform:\n            image = self.transform(image)\n        return image\n\n# Define your transformations\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),  # Adjust size as needed\n    transforms.ToTensor()            # Convert image to tensor\n])\n\n# Create dataset and dataloader\ndataset = DICOMDataset(filtered_df, transform=transform)\ndataloader = DataLoader(dataset, batch_size=16, shuffle=False)  # Adjust batch size as needed\n\n# Step 6: Run the model to get predictions in batches\nsagittal_t1_model.eval()  # Set model to evaluation mode\n\npredictions = []\nwith torch.no_grad():\n    for images_tensor in dataloader:\n        images_tensor = images_tensor.to(device)  # Move batch to device\n        outputs = sagittal_t1_model(images_tensor)  # Get model predictions\n        probabilities = torch.softmax(outputs, dim=1)  # Convert logits to probabilities\n        predicted_classes = torch.argmax(probabilities, dim=1)  # Get predicted class labels\n        predictions.append(predicted_classes.cpu().numpy())\n\n# Combine all predictions\npredicted_labels = np.concatenate(predictions)\n\n# Calculate accuracy\ntrue_labels = filtered_df['mapped_severity'].values  # Extract true severity labels\naccuracy = (predicted_labels == true_labels).mean()  # Calculate accuracy\nprint(f'Accuracy: {accuracy:.4f}')\n","metadata":{"execution":{"iopub.status.busy":"2024-11-02T16:19:42.674872Z","iopub.execute_input":"2024-11-02T16:19:42.675277Z","iopub.status.idle":"2024-11-02T16:20:44.496898Z","shell.execute_reply.started":"2024-11-02T16:19:42.675238Z","shell.execute_reply":"2024-11-02T16:20:44.495935Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport torch\nfrom torchvision import transforms\nfrom PIL import Image\nimport pydicom\nimport numpy as np\n\n# Step 1: Mapping severity to integers\nseverity_mapping = {'normal_mild': 0, 'moderate': 1, 'severe': 2}  # Adjust based on your severity levels\ntest_df_t2['mapped_severity'] = test_df_t2['severity'].map(severity_mapping)\n\n# Step 2: Filter DataFrame for a specific series description\nspecific_series_description = 'Axial T2'  # Replace with your criteria\nfiltered_df = test_df_t2[test_df_t2['series_description'] == specific_series_description]\n\n# Step 3: Load images using the generator\nimages = []\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),  # Adjust size as needed\n    transforms.ToTensor()            # Convert image to tensor\n])\n\n\nclass DICOMDataset(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, idx):\n        image_path = self.dataframe.iloc[idx]['image_path']\n        image = load_dicom_image(image_path)\n        if self.transform:\n            image = self.transform(image)\n        return image\n\n# Define your transformations\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),  # Adjust size as needed\n    transforms.ToTensor()            # Convert image to tensor\n])\n\n# Create dataset and dataloader\ndataset = DICOMDataset(filtered_df, transform=transform)\ndataloader = DataLoader(dataset, batch_size=16, shuffle=False)  # Adjust batch size as needed\n\n# Step 6: Run the model to get predictions in batches\naxial_t2_model.eval()  # Set model to evaluation mode\n\npredictions = []\nwith torch.no_grad():\n    for images_tensor in dataloader:\n        images_tensor = images_tensor.to(device)  # Move batch to device\n        outputs = axial_t2_model(images_tensor)  # Get model predictions\n        probabilities = torch.softmax(outputs, dim=1)  # Convert logits to probabilities\n        predicted_classes = torch.argmax(probabilities, dim=1)  # Get predicted class labels\n        predictions.append(predicted_classes.cpu().numpy())\n\n# Combine all predictions\npredicted_labels = np.concatenate(predictions)\n\n# Calculate accuracy\ntrue_labels = filtered_df['mapped_severity'].values  # Extract true severity labels\naccuracy = (predicted_labels == true_labels).mean()  # Calculate accuracy\nprint(f'Accuracy: {accuracy:.4f}')\n","metadata":{"execution":{"iopub.status.busy":"2024-11-02T16:31:18.520537Z","iopub.execute_input":"2024-11-02T16:31:18.520935Z","iopub.status.idle":"2024-11-02T16:32:08.792865Z","shell.execute_reply.started":"2024-11-02T16:31:18.520899Z","shell.execute_reply":"2024-11-02T16:32:08.791902Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport torch\nfrom torchvision import transforms\nfrom PIL import Image\nimport pydicom\nimport numpy as np\n\n# Step 1: Mapping severity to integers\nseverity_mapping = {'normal_mild': 0, 'moderate': 1, 'severe': 2}  # Adjust based on your severity levels\ntest_df_t2stir['mapped_severity'] = test_df_t2stir['severity'].map(severity_mapping)\n\n# Step 2: Filter DataFrame for a specific series description\nspecific_series_description = 'Sagittal T2/STIR'  # Replace with your criteria\nfiltered_df = test_df_t2stir[test_df_t2stir['series_description'] == specific_series_description]\n\n# Step 3: Load images using the generator\nimages = []\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),  # Adjust size as needed\n    transforms.ToTensor()            # Convert image to tensor\n])\n\n\nclass DICOMDataset(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, idx):\n        image_path = self.dataframe.iloc[idx]['image_path']\n        image = load_dicom_image(image_path)\n        if self.transform:\n            image = self.transform(image)\n        return image\n\n# Define your transformations\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),  # Adjust size as needed\n    transforms.ToTensor()            # Convert image to tensor\n])\n\n# Create dataset and dataloader\ndataset = DICOMDataset(filtered_df, transform=transform)\ndataloader = DataLoader(dataset, batch_size=16, shuffle=False)  # Adjust batch size as needed\n\n# Step 6: Run the model to get predictions in batches\nsagittal_t2stir_model.eval()  # Set model to evaluation mode\n\npredictions = []\nwith torch.no_grad():\n    for images_tensor in dataloader:\n        images_tensor = images_tensor.to(device)  # Move batch to device\n        outputs = sagittal_t2stir_model(images_tensor)  # Get model predictions\n        probabilities = torch.softmax(outputs, dim=1)  # Convert logits to probabilities\n        predicted_classes = torch.argmax(probabilities, dim=1)  # Get predicted class labels\n        predictions.append(predicted_classes.cpu().numpy())\n\n# Combine all predictions\npredicted_labels = np.concatenate(predictions)\n\n# Calculate accuracy\ntrue_labels = filtered_df['mapped_severity'].values  # Extract true severity labels\naccuracy = (predicted_labels == true_labels).mean()  # Calculate accuracy\nprint(f'Accuracy: {accuracy:.4f}')\n","metadata":{"execution":{"iopub.status.busy":"2024-11-02T16:44:15.762850Z","iopub.execute_input":"2024-11-02T16:44:15.763643Z","iopub.status.idle":"2024-11-02T16:44:44.776433Z","shell.execute_reply.started":"2024-11-02T16:44:15.763598Z","shell.execute_reply":"2024-11-02T16:44:44.775490Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_df_t2.iloc[14]","metadata":{"execution":{"iopub.status.busy":"2024-11-11T12:58:13.758075Z","iopub.execute_input":"2024-11-11T12:58:13.758915Z","iopub.status.idle":"2024-11-11T12:58:13.766381Z","shell.execute_reply.started":"2024-11-11T12:58:13.758874Z","shell.execute_reply":"2024-11-11T12:58:13.765449Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\nimage_path = test_df_t2.iloc[14]['image_path']  # Update 'image_path' to match your DataFrame column name\n\n# Load and transform the image\nimage = load_dicom_image(image_path)\n\n# Display the image\nplt.imshow(image, cmap='gray')  # Assuming the image is grayscale; adjust if needed\nplt.axis('off')  # Hide axes for better visualization\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-11-11T12:58:27.780274Z","iopub.execute_input":"2024-11-11T12:58:27.781147Z","iopub.status.idle":"2024-11-11T12:58:27.994412Z","shell.execute_reply.started":"2024-11-11T12:58:27.781106Z","shell.execute_reply":"2024-11-11T12:58:27.993432Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Feature maps","metadata":{}},{"cell_type":"code","source":"import torch\nimport matplotlib.pyplot as plt\nfrom torchvision import transforms\nfrom PIL import Image\nimport pydicom\nimport timm\n\n# Load your model\nmodel = timm.create_model('efficientnetv2_s', pretrained=False)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel.load_state_dict(torch.load('/kaggle/input/model2_weights.pth/pytorch/default/1/model2_weights.pth', map_location=device), strict=False)\nmodel = model.to(device)\nmodel.eval()\n\n# Load the test dataframe (assuming you have 'test_df_t2')\n# Adjust according to your DataFrame structure\nimage_idx = 0  # Specify the index of the image you want to visualize\nimage_path = test_df_t2.iloc[image_idx]['image_path']  # Update 'image_path' to match your DataFrame column name\n\n# Function to load DICOM images\ndef load_dicom_image(image_path):\n    dicom = pydicom.dcmread(image_path)\n    image = dicom.pixel_array\n\n    # Normalize the image array to the range [0, 255] if needed\n    if image.dtype != np.uint8:\n        image = (image / np.max(image) * 255).astype(np.uint8)\n\n    # Convert to PIL Image\n    image = Image.fromarray(image)\n\n    # Convert to RGB as needed for the model\n    if image.mode != 'RGB':\n        image = image.convert('RGB')  # Convert to RGB\n\n    return image\n\n# Load and transform the image\nimage = load_dicom_image(image_path)\n\n# Define transformations for the image\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor()\n])\n\n# Apply transformations to the image\nimage_tensor = transform(image).unsqueeze(0)  # Add batch dimension\n\n# Hook to get the feature maps\nfeature_maps = []\n\ndef get_feature_maps(module, input, output):\n    feature_maps.append(output)\n\n# Register hooks for specific layers or blocks (assuming layers are named 'blocks')\nfor name, layer in model.named_modules():\n    if 'blocks' in name:  # Focus on blocks or specific layers\n        layer.register_forward_hook(get_feature_maps)\n\n# Forward pass through the model to capture feature maps\nwith torch.no_grad():\n    output = model(image_tensor.to(device))  # Ensure image is on the correct device\n\n# Plotting the feature maps for all layers\nnum_layers = len(feature_maps)\nbatch_size = 5  # Set batch size for number of layers per plot\n\n# Loop through the layers in batches\nfor start in range(0, num_layers, batch_size):\n    end = min(start + batch_size, num_layers)  # Ensure we don't exceed the total number of layers\n\n    # Create a new set of subplots for this batch of layers\n    fig, axes = plt.subplots(end - start, 1, figsize=(10, (end - start) * 2))  # Adjust the number of rows dynamically\n\n    # If there is only one subplot, make sure axes is iterable\n    if end - start == 1:\n        axes = [axes]\n\n    # Plot the feature maps for each layer in this batch\n    for layer_index in range(start, end):\n        feature_map = feature_maps[layer_index]\n        num_feature_maps = feature_map.shape[1]\n\n        # Loop through feature maps in this layer (limit to 5 per layer)\n        for i in range(min(num_feature_maps, 5)):  # Limit to 5 feature maps per layer\n            axes[layer_index - start].imshow(feature_map[0, i].cpu().numpy(), cmap='gray')\n            axes[layer_index - start].axis('off')  # Turn off the axis\n\n        axes[layer_index - start].set_title(f'Layer {layer_index + 1} - {num_feature_maps} feature maps')\n\n    # Display the plot\n    plt.tight_layout()\n    plt.show()\n","metadata":{"scrolled":true,"execution":{"iopub.status.busy":"2024-11-11T12:51:16.967362Z","iopub.execute_input":"2024-11-11T12:51:16.967760Z","iopub.status.idle":"2024-11-11T12:52:32.420705Z","shell.execute_reply.started":"2024-11-11T12:51:16.967723Z","shell.execute_reply":"2024-11-11T12:52:32.419680Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport matplotlib.pyplot as plt\nfrom torchvision import transforms\nfrom torch.utils.data import Dataset\nimport numpy as np\nimport pydicom\nfrom PIL import Image\nimport timm\n\n# Load your model\nmodel = timm.create_model('efficientnetv2_s', pretrained=False)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel.load_state_dict(torch.load('/kaggle/input/model2_weights.pth/pytorch/default/1/model2_weights.pth', map_location=device), strict=False)\nmodel = model.to(device)\nmodel.eval()\n\n# Load the test dataframe (assuming you have 'test_df_t2')\n# Adjust according to your DataFrame structure\nimage_idx = 0  # Specify the index of the image you want to visualize\nimage_path = test_df_t2.iloc[image_idx]['image_path']  # Update 'image_path' to match your DataFrame column name\n\n# Function to load DICOM images\ndef load_dicom_image(image_path):\n    dicom = pydicom.dcmread(image_path)\n    image = dicom.pixel_array\n\n    # Normalize the image array to the range [0, 255] if needed\n    if image.dtype != np.uint8:\n        image = (image / np.max(image) * 255).astype(np.uint8)\n\n    # Convert to PIL Image\n    image = Image.fromarray(image)\n\n    # Convert to RGB as needed for the model\n    if image.mode != 'RGB':\n        image = image.convert('RGB')  # Convert to RGB\n\n    return image\n\n# Load and transform the image\nimage = load_dicom_image(image_path)\n\n# Define transformations for the image\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor()\n])\n\n# Apply transformations to the image\nimage_tensor = transform(image).unsqueeze(0)  # Add batch dimension\n\n# Hook to get the feature maps\nfeature_maps = []\n\ndef get_feature_maps(module, input, output):\n    feature_maps.append(output)\n\n# Register hooks for specific layers or blocks (assuming layers are named 'blocks')\nfor name, layer in model.named_modules():\n    if 'blocks' in name:  # Focus on blocks or specific layers\n        layer.register_forward_hook(get_feature_maps)\n\n# Forward pass through the model to capture feature maps\nwith torch.no_grad():\n    output = model(image_tensor.to(device))  # Ensure image is on the correct device\n\n# Check if feature maps are being collected\nif len(feature_maps) == 0:\n    print(\"No feature maps captured. Please check the hook registration and layer names.\")\nelse:\n    print(f\"Captured {len(feature_maps)} feature maps\")\n\n# Plotting the feature maps for all layers\nnum_layers = len(feature_maps)\nprint(f\"Total number of layers captured: {num_layers}\")\n\n# Set a batch size to limit the number of layers per plot\nbatch_size = 5\n\n# Loop through the layers in batches\nfor start in range(0, num_layers, batch_size):\n    end = min(start + batch_size, num_layers)  # Ensure we don't exceed the total number of layers\n\n    # Create a new set of subplots for this batch of layers\n    fig, axes = plt.subplots(end - start, 1, figsize=(10, (end - start) * 2))  # Dynamically adjust the number of rows\n\n    # If there is only one subplot, make sure axes is iterable\n    if end - start == 1:\n        axes = [axes]\n\n    # Plot the feature maps for each layer in this batch\n    for layer_index in range(start, end):\n        feature_map = feature_maps[layer_index]\n        num_feature_maps = feature_map.shape[1]\n\n        # Loop through feature maps in this layer (limit to 5 per layer)\n        for i in range(min(num_feature_maps, 5)):  # Limit to 5 feature maps per layer\n            axes[layer_index - start].imshow(feature_map[0, i].cpu().numpy(), cmap='gray')\n            axes[layer_index - start].axis('off')  # Turn off the axis\n\n        axes[layer_index - start].set_title(f'Layer {layer_index + 1} - {num_feature_maps} feature maps')\n\n    # Display the plot\n    plt.tight_layout()\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-11-11T13:00:08.193871Z","iopub.execute_input":"2024-11-11T13:00:08.194319Z","iopub.status.idle":"2024-11-11T13:01:30.666657Z","shell.execute_reply.started":"2024-11-11T13:00:08.194280Z","shell.execute_reply":"2024-11-11T13:01:30.665648Z"},"scrolled":true,"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Print all layers of the model\nfor name, layer in model.named_modules():\n    print(f\"{name}: {layer}\")\n","metadata":{"execution":{"iopub.status.busy":"2024-11-11T12:30:13.061709Z","iopub.execute_input":"2024-11-11T12:30:13.062477Z","iopub.status.idle":"2024-11-11T12:30:13.106820Z","shell.execute_reply.started":"2024-11-11T12:30:13.062434Z","shell.execute_reply":"2024-11-11T12:30:13.105820Z"},"_kg_hide-output":false,"_kg_hide-input":false,"scrolled":true,"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for layer_index, feature_map in enumerate(feature_maps):\n    print(f\"Layer {layer_index+1} feature map shape: {feature_map.shape}\")","metadata":{"execution":{"iopub.status.busy":"2024-11-11T12:41:49.799020Z","iopub.execute_input":"2024-11-11T12:41:49.799446Z","iopub.status.idle":"2024-11-11T12:41:49.813353Z","shell.execute_reply.started":"2024-11-11T12:41:49.799406Z","shell.execute_reply":"2024-11-11T12:41:49.812423Z"},"scrolled":true,"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# GradCAM","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport numpy as np\nimport cv2\nimport matplotlib.pyplot as plt\nfrom torchvision import transforms\nfrom PIL import Image\nimport timm\nimport pydicom\n\n# Load your model (adjust as needed)\nmodel = timm.create_model('efficientnetv2_s', pretrained=False)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel.load_state_dict(torch.load('/kaggle/input/model2_weights.pth/pytorch/default/1/model2_weights.pth', map_location=device), strict=False)\nmodel = model.to(device)\nmodel.eval()\n\n# Hook to capture feature maps and gradients\nfeature_maps = []\ngradients = []\n\ndef save_feature_maps(module, input, output):\n    feature_maps.append(output)\n\ndef save_gradient(grad):\n    gradients.append(grad)\n\n# Choose the layer (convolutional layer) where Grad-CAM will be computed\ntarget_layer = model.blocks[-1]  # Example: last block, adjust based on your model architecture\n\n# Register the hook for capturing the feature maps\ntarget_layer.register_forward_hook(save_feature_maps)\n# Register the hook for capturing the gradients\ntarget_layer.register_backward_hook(lambda self, grad_input, grad_output: save_gradient(grad_output[0]))\n\n# Load and preprocess the image (adjust paths)\nimage_idx = 0\nimage_path = test_df_t2.iloc[image_idx]['image_path']\nimage = load_dicom_image(image_path)\n\n# Define transformations\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor()\n])\n\nimage_tensor = transform(image).unsqueeze(0).to(device)\n\n# Forward pass to get output (predictions)\noutput = model(image_tensor)\n\n# Get the predicted class\npredicted_class = torch.argmax(output, dim=1)\n\n# Backward pass to get gradients\nmodel.zero_grad()\noutput[0, predicted_class].backward()\n\n# Get the gradients and activations from the selected layer\ngradient = gradients[0]\nactivation = feature_maps[0]\n\n# Global average pooling on the gradients to get the weights\nweights = torch.mean(gradient, dim=(2, 3), keepdim=True)\n\n# Compute Grad-CAM\ngrad_cam_map = F.relu(torch.sum(weights * activation, dim=1)).squeeze()\n\n# Resize the heatmap to match the input image size\ngrad_cam_map = grad_cam_map.cpu().detach().numpy()\ngrad_cam_map = cv2.resize(grad_cam_map, (224, 224))  # Resize to input image size\n\n# Normalize the heatmap\ngrad_cam_map = np.maximum(grad_cam_map, 0)\ngrad_cam_map = grad_cam_map / grad_cam_map.max()\n\n# Convert to RGB and apply colormap\nheatmap = cv2.applyColorMap(np.uint8(255 * grad_cam_map), cv2.COLORMAP_JET)\n\n# Convert image to numpy for overlaying (resize it to match the heatmap size)\nimage_numpy = np.array(image.convert(\"RGB\"))\nimage_numpy_resized = cv2.resize(image_numpy, (224, 224))  # Resize to 224x224\n\n# Overlay the heatmap on the resized original image\nsuperimposed_image = heatmap * 0.4 + image_numpy_resized  # Adjust alpha for heatmap intensity\nsuperimposed_image = np.uint8(np.clip(superimposed_image, 0, 255))\n\n# Display the results\nplt.figure(figsize=(10, 10))\nplt.imshow(superimposed_image)\nplt.axis('off')\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-11-11T13:10:44.418574Z","iopub.execute_input":"2024-11-11T13:10:44.419349Z","iopub.status.idle":"2024-11-11T13:10:45.332446Z","shell.execute_reply.started":"2024-11-11T13:10:44.419306Z","shell.execute_reply":"2024-11-11T13:10:45.331397Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport numpy as np\nimport cv2\nimport matplotlib.pyplot as plt\nfrom torchvision import transforms\nfrom PIL import Image\nimport timm\nimport pydicom\n\n# Load your model (adjust as needed)\nmodel = timm.create_model('efficientnetv2_s', pretrained=False)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel.load_state_dict(torch.load('/kaggle/input/model3_weights.pth/pytorch/default/1/model3_weights.pth', map_location=device), strict=False)\nmodel = model.to(device)\nmodel.eval()\n\n# Hook to capture feature maps and gradients\nfeature_maps = []\ngradients = []\n\ndef save_feature_maps(module, input, output):\n    feature_maps.append(output)\n\ndef save_gradient(grad):\n    gradients.append(grad)\n\n# Choose the layer (convolutional layer) where Grad-CAM will be computed\ntarget_layer = model.blocks[-1]  # Example: last block, adjust based on your model architecture\n\n# Register the hook for capturing the feature maps\ntarget_layer.register_forward_hook(save_feature_maps)\n# Register the hook for capturing the gradients\ntarget_layer.register_backward_hook(lambda self, grad_input, grad_output: save_gradient(grad_output[0]))\n\n# Load and preprocess the image (adjust paths)\nimage_idx = 0\nimage_path = test_df_t2stir.iloc[image_idx]['image_path']\nimage = load_dicom_image(image_path)\n\n# Define transformations\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor()\n])\n\nimage_tensor = transform(image).unsqueeze(0).to(device)\n\n# Forward pass to get output (predictions)\noutput = model(image_tensor)\n\n# Get the predicted class\npredicted_class = torch.argmax(output, dim=1)\n\n# Backward pass to get gradients\nmodel.zero_grad()\noutput[0, predicted_class].backward()\n\n# Get the gradients and activations from the selected layer\ngradient = gradients[0]\nactivation = feature_maps[0]\n\n# Global average pooling on the gradients to get the weights\nweights = torch.mean(gradient, dim=(2, 3), keepdim=True)\n\n# Compute Grad-CAM\ngrad_cam_map = F.relu(torch.sum(weights * activation, dim=1)).squeeze()\n\n# Resize the heatmap to match the input image size\ngrad_cam_map = grad_cam_map.cpu().detach().numpy()\ngrad_cam_map = cv2.resize(grad_cam_map, (224, 224))  # Resize to input image size\n\n# Normalize the heatmap\ngrad_cam_map = np.maximum(grad_cam_map, 0)\ngrad_cam_map = grad_cam_map / grad_cam_map.max()\n\n# Convert to RGB and apply colormap\nheatmap = cv2.applyColorMap(np.uint8(255 * grad_cam_map), cv2.COLORMAP_JET)\n\n# Convert image to numpy for overlaying (resize it to match the heatmap size)\nimage_numpy = np.array(image.convert(\"RGB\"))\nimage_numpy_resized = cv2.resize(image_numpy, (224, 224))  # Resize to 224x224\n\n# Overlay the heatmap on the resized original image\nsuperimposed_image = heatmap * 0.4 + image_numpy_resized  # Adjust alpha for heatmap intensity\nsuperimposed_image = np.uint8(np.clip(superimposed_image, 0, 255))\n\n# Display the results\nplt.figure(figsize=(10, 10))\nplt.imshow(superimposed_image)\nplt.axis('off')\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-11-11T13:20:20.911491Z","iopub.execute_input":"2024-11-11T13:20:20.912280Z","iopub.status.idle":"2024-11-11T13:20:22.917226Z","shell.execute_reply.started":"2024-11-11T13:20:20.912227Z","shell.execute_reply":"2024-11-11T13:20:22.916296Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport torch.nn.functional as F\nimport numpy as np\nimport cv2\nimport matplotlib.pyplot as plt\nfrom torchvision import transforms\nfrom PIL import Image\nimport timm\nimport pydicom\n\n# Load your model (adjust as needed)\nmodel = timm.create_model('efficientnetv2_s', pretrained=False)\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel.load_state_dict(torch.load('/kaggle/input/model1_weights.pth/pytorch/default/1/model1_weights.pth', map_location=device), strict=False)\nmodel = model.to(device)\nmodel.eval()\n\n# Hook to capture feature maps and gradients\nfeature_maps = []\ngradients = []\n\ndef save_feature_maps(module, input, output):\n    feature_maps.append(output)\n\ndef save_gradient(grad):\n    gradients.append(grad)\n\n# Choose the layer (convolutional layer) where Grad-CAM will be computed\ntarget_layer = model.blocks[-1]  # Example: last block, adjust based on your model architecture\n\n# Register the hook for capturing the feature maps\ntarget_layer.register_forward_hook(save_feature_maps)\n# Register the hook for capturing the gradients\ntarget_layer.register_backward_hook(lambda self, grad_input, grad_output: save_gradient(grad_output[0]))\n\n# Load and preprocess the image (adjust paths)\nimage_idx = 0\nimage_path = test_df_t1.iloc[image_idx]['image_path']\nimage = load_dicom_image(image_path)\n\n# Define transformations\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor()\n])\n\nimage_tensor = transform(image).unsqueeze(0).to(device)\n\n# Forward pass to get output (predictions)\noutput = model(image_tensor)\n\n# Get the predicted class\npredicted_class = torch.argmax(output, dim=1)\n\n# Backward pass to get gradients\nmodel.zero_grad()\noutput[0, predicted_class].backward()\n\n# Get the gradients and activations from the selected layer\ngradient = gradients[0]\nactivation = feature_maps[0]\n\n# Global average pooling on the gradients to get the weights\nweights = torch.mean(gradient, dim=(2, 3), keepdim=True)\n\n# Compute Grad-CAM\ngrad_cam_map = F.relu(torch.sum(weights * activation, dim=1)).squeeze()\n\n# Resize the heatmap to match the input image size\ngrad_cam_map = grad_cam_map.cpu().detach().numpy()\ngrad_cam_map = cv2.resize(grad_cam_map, (224, 224))  # Resize to input image size\n\n# Normalize the heatmap\ngrad_cam_map = np.maximum(grad_cam_map, 0)\ngrad_cam_map = grad_cam_map / grad_cam_map.max()\n\n# Convert to RGB and apply colormap\nheatmap = cv2.applyColorMap(np.uint8(255 * grad_cam_map), cv2.COLORMAP_JET)\n\n# Convert image to numpy for overlaying (resize it to match the heatmap size)\nimage_numpy = np.array(image.convert(\"RGB\"))\nimage_numpy_resized = cv2.resize(image_numpy, (224, 224))  # Resize to 224x224\n\n# Overlay the heatmap on the resized original image\nsuperimposed_image = heatmap * 0.4 + image_numpy_resized  # Adjust alpha for heatmap intensity\nsuperimposed_image = np.uint8(np.clip(superimposed_image, 0, 255))\n\n# Display the results\nplt.figure(figsize=(10, 10))\nplt.imshow(superimposed_image)\nplt.axis('off')\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-11-11T13:22:18.240121Z","iopub.execute_input":"2024-11-11T13:22:18.241009Z","iopub.status.idle":"2024-11-11T13:22:20.278994Z","shell.execute_reply.started":"2024-11-11T13:22:18.240966Z","shell.execute_reply":"2024-11-11T13:22:20.276900Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# actual vs predicted","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom torchvision import transforms\nfrom PIL import Image\nimport timm\nimport pydicom\n\n# Load your model (adjust as needed)\nmodel = timm.create_model('efficientnetv2_s', pretrained=False, num_classes=3)  # Ensure num_classes is set correctly\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel.load_state_dict(torch.load('/kaggle/input/model2_weights.pth/pytorch/default/1/model2_weights.pth', map_location=device), strict=False)\nmodel = model.to(device)\nmodel.eval()\n\n# Function to load DICOM images\ndef load_dicom_image(image_path):\n    dicom = pydicom.dcmread(image_path)\n    image = dicom.pixel_array\n\n    # Normalize the image array to the range [0, 255] if needed\n    if image.dtype != np.uint8:\n        image = (image / np.max(image) * 255).astype(np.uint8)\n\n    # Convert to PIL Image\n    image = Image.fromarray(image)\n\n    # Convert to RGB as needed for the model\n    if image.mode != 'RGB':\n        image = image.convert('RGB')  # Convert to RGB\n\n    return image\n\n# Transformation to apply to each image\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor()\n])\n\n# Class names corresponding to your output classes\nclass_names = [\"normal_mild\", \"moderate\", \"severe\"]\n\n# Sample 10 images from the DataFrame (replace `test_df_t2` with your DataFrame)\nsample_df = test_df_t2.sample(10)\n\n# Create subplots to display the images\nfig, axes = plt.subplots(5, 2, figsize=(15, 25))\naxes = axes.flatten()\n\n# Iterate over the sample data and display images with actual and predicted labels\nfor idx, (index, row) in enumerate(sample_df.iterrows()):\n    image_path = row['image_path']\n    actual_label = row['severity']\n\n    # Load and process the image\n    image = load_dicom_image(image_path)\n    image_tensor = transform(image).unsqueeze(0).to(device)\n\n    # Get model prediction\n    with torch.no_grad():\n        output = model(image_tensor)\n        predicted_idx = torch.argmax(output, dim=1).item()\n\n        # Print the predicted index and check it\n        print(f\"Predicted Index: {predicted_idx}, Output Shape: {output.shape}\")\n\n        # Ensure the predicted index is within range\n        if predicted_idx < len(class_names):\n            predicted_label = class_names[predicted_idx]\n        else:\n            predicted_label = \"Unknown\"  # If there's an out-of-range index, assign a fallback label\n\n    # Display the image and labels\n    axes[idx].imshow(image)\n    axes[idx].axis('off')\n    axes[idx].set_title(f\"Actual: {actual_label}\\nPredicted: {predicted_label}\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:28:00.992955Z","iopub.execute_input":"2024-11-13T04:28:00.993341Z","iopub.status.idle":"2024-11-13T04:28:07.144647Z","shell.execute_reply.started":"2024-11-13T04:28:00.993304Z","shell.execute_reply":"2024-11-13T04:28:07.143708Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom torchvision import transforms\nfrom PIL import Image\nimport timm\nimport pydicom\n\n# Load your model (adjust as needed)\nmodel = timm.create_model('efficientnetv2_s', pretrained=False, num_classes=3)  # Ensure num_classes is set correctly\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel.load_state_dict(torch.load('/kaggle/input/model1_weights.pth/pytorch/default/1/model1_weights.pth', map_location=device), strict=False)\nmodel = model.to(device)\nmodel.eval()\n\n# Function to load DICOM images\ndef load_dicom_image(image_path):\n    dicom = pydicom.dcmread(image_path)\n    image = dicom.pixel_array\n\n    # Normalize the image array to the range [0, 255] if needed\n    if image.dtype != np.uint8:\n        image = (image / np.max(image) * 255).astype(np.uint8)\n\n    # Convert to PIL Image\n    image = Image.fromarray(image)\n\n    # Convert to RGB as needed for the model\n    if image.mode != 'RGB':\n        image = image.convert('RGB')  # Convert to RGB\n\n    return image\n\n# Transformation to apply to each image\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor()\n])\n\n# Class names corresponding to your output classes\nclass_names = [\"normal_mild\", \"moderate\", \"severe\"]\n\n# Sample 10 images from the DataFrame (replace `test_df_t2` with your DataFrame)\nsample_df = test_df_t1.sample(10)\n\n# Create subplots to display the images\nfig, axes = plt.subplots(5, 2, figsize=(15, 25))\naxes = axes.flatten()\n\n# Iterate over the sample data and display images with actual and predicted labels\nfor idx, (index, row) in enumerate(sample_df.iterrows()):\n    image_path = row['image_path']\n    actual_label = row['severity']\n\n    # Load and process the image\n    image = load_dicom_image(image_path)\n    image_tensor = transform(image).unsqueeze(0).to(device)\n\n    # Get model prediction\n    with torch.no_grad():\n        output = model(image_tensor)\n        predicted_idx = torch.argmax(output, dim=1).item()\n\n        # Print the predicted index and check it\n        print(f\"Predicted Index: {predicted_idx}, Output Shape: {output.shape}\")\n\n        # Ensure the predicted index is within range\n        if predicted_idx < len(class_names):\n            predicted_label = class_names[predicted_idx]\n        else:\n            predicted_label = \"Unknown\"  # If there's an out-of-range index, assign a fallback label\n\n    # Display the image and labels\n    axes[idx].imshow(image)\n    axes[idx].axis('off')\n    axes[idx].set_title(f\"Actual: {actual_label}\\nPredicted: {predicted_label}\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:41:05.731052Z","iopub.execute_input":"2024-11-13T04:41:05.731927Z","iopub.status.idle":"2024-11-13T04:41:09.623075Z","shell.execute_reply.started":"2024-11-13T04:41:05.731884Z","shell.execute_reply":"2024-11-13T04:41:09.622136Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nimport torch\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom torchvision import transforms\nfrom PIL import Image\nimport timm\nimport pydicom\n\n# Load your model (adjust as needed)\nmodel = timm.create_model('efficientnetv2_s', pretrained=False, num_classes=3)  # Ensure num_classes is set correctly\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nmodel.load_state_dict(torch.load('/kaggle/input/model3_weights.pth/pytorch/default/1/model3_weights.pth', map_location=device), strict=False)\nmodel = model.to(device)\nmodel.eval()\n\n# Function to load DICOM images\ndef load_dicom_image(image_path):\n    dicom = pydicom.dcmread(image_path)\n    image = dicom.pixel_array\n\n    # Normalize the image array to the range [0, 255] if needed\n    if image.dtype != np.uint8:\n        image = (image / np.max(image) * 255).astype(np.uint8)\n\n    # Convert to PIL Image\n    image = Image.fromarray(image)\n\n    # Convert to RGB as needed for the model\n    if image.mode != 'RGB':\n        image = image.convert('RGB')  # Convert to RGB\n\n    return image\n\n# Transformation to apply to each image\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),\n    transforms.ToTensor()\n])\n\n# Class names corresponding to your output classes\nclass_names = [\"normal_mild\", \"moderate\", \"severe\"]\n\n# Sample 10 images from the DataFrame (replace `test_df_t2` with your DataFrame)\nsample_df = test_df_t2stir.sample(10)\n\n# Create subplots to display the images\nfig, axes = plt.subplots(5, 2, figsize=(15, 25))\naxes = axes.flatten()\n\n# Iterate over the sample data and display images with actual and predicted labels\nfor idx, (index, row) in enumerate(sample_df.iterrows()):\n    image_path = row['image_path']\n    actual_label = row['severity']\n\n    # Load and process the image\n    image = load_dicom_image(image_path)\n    image_tensor = transform(image).unsqueeze(0).to(device)\n\n    # Get model prediction\n    with torch.no_grad():\n        output = model(image_tensor)\n        predicted_idx = torch.argmax(output, dim=1).item()\n\n        # Print the predicted index and check it\n        print(f\"Predicted Index: {predicted_idx}, Output Shape: {output.shape}\")\n\n        # Ensure the predicted index is within range\n        if predicted_idx < len(class_names):\n            predicted_label = class_names[predicted_idx]\n        else:\n            predicted_label = \"Unknown\"  # If there's an out-of-range index, assign a fallback label\n\n    # Display the image and labels\n    axes[idx].imshow(image)\n    axes[idx].axis('off')\n    axes[idx].set_title(f\"Actual: {actual_label}\\nPredicted: {predicted_label}\")\n\nplt.tight_layout()\nplt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:41:45.011984Z","iopub.execute_input":"2024-11-13T04:41:45.012362Z","iopub.status.idle":"2024-11-13T04:41:49.089648Z","shell.execute_reply.started":"2024-11-13T04:41:45.012327Z","shell.execute_reply":"2024-11-13T04:41:49.088377Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# B7","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport timm\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom copy import deepcopy\nimport pydicom\nfrom torchvision import transforms\nfrom PIL import Image\n\nclass EfficientNetB7(nn.Module):\n    def __init__(self, num_classes=3):\n        super(EfficientNetB7, self).__init__()\n        # Load EfficientNet B7 with pretrained weights from ImageNet-21k\n        self.model = timm.create_model('tf_efficientnet_b7_ns', pretrained=True)\n\n        # Freeze the first 40% of the layers\n        total_layers = len(list(self.model.parameters()))\n        freeze_up_to = int(total_layers * 0.4)\n        for i, param in enumerate(self.model.parameters()):\n            if i < freeze_up_to:\n                param.requires_grad = False\n\n        # Replace the final classifier layer to match the number of custom classes\n        num_ftrs = self.model.classifier.in_features\n        self.model.classifier = nn.Linear(num_ftrs, num_classes)\n\n    def forward(self, x):\n        return self.model(x)\n\n# Instantiate and move the model to the device\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nsagittal_t1_model = EfficientNetB7(num_classes=3).to(device)\naxial_t2_model = EfficientNetB7(num_classes=3).to(device)\nsagittal_t2stir_model = EfficientNetB7(num_classes=3).to(device)\n\n# Initialize separate optimizers for each model\noptimizer_sagittal_t1 = torch.optim.Adam(filter(lambda p: p.requires_grad, sagittal_t1_model.parameters()), lr=0.001)\noptimizer_axial_t2 = torch.optim.Adam(filter(lambda p: p.requires_grad, axial_t2_model.parameters()), lr=0.001)\noptimizer_sagittal_t2stir = torch.optim.Adam(filter(lambda p: p.requires_grad, sagittal_t2stir_model.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}\n\n# Label map for class indices\nlabel_map = {'normal_mild': 0, 'moderate': 1, 'severe': 2}\n\n# Define a transform to convert images into tensors\ntransform = transforms.Compose([\n    transforms.Resize((224, 224)),  # Resize images to 224x224\n    transforms.ToTensor(),  # Convert images to PyTorch tensors\n])\n\n# Function to load a DICOM image from a file path\ndef load_dicom_image(file_path):\n    dicom_data = pydicom.dcmread(file_path)  # Read the DICOM file\n    image = dicom_data.pixel_array  # Get the pixel data as a NumPy array\n    image = Image.fromarray(image)  # Convert to a PIL image\n    return transform(image)  # Apply the transform (resize + to tensor)\n\n# Training function\n# Updated train_model function with additional debugging\n# Updated train_model function with additional handling for batches\n# Updated train_model function\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    criterion = nn.CrossEntropyLoss()  # Ensure the criterion is defined\n\n    # Loop through epochs\n    for epoch in range(num_epochs):\n        print(f\"Epoch {epoch + 1}/{num_epochs}\")\n        model.train()\n        train_loss = 0\n        correct_train = 0\n\n        # Training loop\n        for batch in trainloader:\n            print(f\"Batch: {batch}\")  # Debugging line to check the structure of the batch\n            \n            # Unpack the batch\n            if isinstance(batch, tuple):\n                image_paths, labels = batch  # Assuming batch is a tuple (image_paths, labels)\n            else:\n                image_paths = batch  # In case the batch only contains images (no labels)\n\n            print(f\"Image Paths: {image_paths}\")  # Debugging line\n\n            # Handle image loading: if paths are strings, load DICOM images, otherwise treat them as tensors\n            # Handle image loading: if paths are strings, load DICOM images, otherwise treat them as tensors\n            images = []\n            for path in image_paths:\n                if isinstance(path, str):  # If the path is a string (i.e., a file path), load the image\n                    image = load_dicom_image(path).to(device)  # Load and move to device\n                    images.append(image)\n                elif isinstance(path, torch.Tensor):  # If the path is already a tensor, move it to the device\n                    images.append(path.to(device))\n                else:\n                    raise ValueError(f\"Unexpected type for path: {type(path)}\")  # Add this to catch unexpected types\n\n            # Stack the images into a tensor (assuming the images are now all tensors)\n            images = torch.stack(images)\n\n\n            # Move labels to device\n            labels = labels.to(device)\n\n            optimizer.zero_grad()\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            optimizer.step()\n\n            train_loss += loss.item()\n            _, preds = torch.max(outputs, 1)\n            correct_train += torch.sum(preds == labels).item()\n\n        scheduler.step()\n\n        # Print training progress\n        train_acc = correct_train / len_train\n        print(f'Epoch {epoch + 1}/{num_epochs}, Loss: {train_loss:.4f}, Training Accuracy: {train_acc:.4f}')\n\n        # Check for early stopping condition (optional)\n        if train_acc > best_val_acc:\n            best_val_acc = train_acc\n            best_model_wts = deepcopy(model.state_dict())\n            counter = 0\n        else:\n            counter += 1\n            if counter >= patience:\n                print(\"Early stopping triggered.\")\n                break\n\n    # Load best model weights\n    model.load_state_dict(best_model_wts)\n    return model\n\n# Ensure that trainloader is not empty\nprint(f\"Number of batches in trainloader: {len(trainloader)}\")\n\n# Ensure the model is being trained\ntrained_sagittal_t1_model = train_model(\n    model=models['Sagittal T1'], \n    trainloader=trainloader, \n    valloader=valloader, \n    len_train=len_train, \n    len_val=len_val, \n    optimizer=optimizers['Sagittal T1'],\n    num_epochs=10\n)","metadata":{"execution":{"iopub.status.busy":"2024-11-13T04:39:38.986783Z","iopub.execute_input":"2024-11-13T04:39:38.987225Z","iopub.status.idle":"2024-11-13T04:39:48.830171Z","shell.execute_reply.started":"2024-11-13T04:39:38.987188Z","shell.execute_reply":"2024-11-13T04:39:48.828825Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# import torch\n# from tqdm import tqdm\n# from sklearn.metrics import accuracy_score\n\n# # Define device\n# device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# # Load each saved model for testing\n# # Paths to each saved best model (update these paths as needed)\n# saved_model_paths = {\n#     'Sagittal T1': 'path_to_best_model_sagittal_t1.pth',\n#     'Axial T2': 'path_to_best_model_axial_t2.pth',\n#     'Sagittal T2/STIR': 'path_to_best_model_sagittal_t2stir.pth'\n# }\n\n# # Reload the models and load weights\n# for desc, model in models.items():\n#     model.load_state_dict(torch.load(saved_model_paths[desc]))\n#     model.to(device)\n#     model.eval()  # Set to evaluation mode\n\n# # Function to test a single model\n# def test_model(model, testloader):\n#     all_labels = []\n#     all_predictions = []\n#     with torch.no_grad():\n#         for images, labels in tqdm(testloader, desc=\"Testing\"):\n#             images, labels = images.to(device), torch.tensor([label_map[label] for label in labels]).to(device)\n#             outputs = model(images)\n#             probabilities = torch.softmax(outputs, dim=1)\n#             _, predicted = torch.max(probabilities, 1)\n#             all_predictions.extend(predicted.cpu().numpy())\n#             all_labels.extend(labels.cpu().numpy())\n\n#     # Calculate accuracy\n#     accuracy = accuracy_score(all_labels, all_predictions) * 100\n#     return accuracy\n\n# # Test each model on its respective test set\n# for desc, model in models.items():\n#     if desc == 'Sagittal T1':\n#         testloader = testloader_t1\n#     elif desc == 'Axial T2':\n#         testloader = testloader_t2\n#     elif desc == 'Sagittal T2/STIR':\n#         testloader = testloader_t2stir\n\n#     print(f\"Testing model for {desc} images...\")\n#     accuracy = test_model(model, testloader)\n#     print(f\"Accuracy for {desc}: {accuracy:.2f}%\")\n","metadata":{"execution":{"iopub.status.busy":"2024-10-31T14:12:42.395356Z","iopub.execute_input":"2024-10-31T14:12:42.395720Z","iopub.status.idle":"2024-10-31T14:12:42.401947Z","shell.execute_reply.started":"2024-10-31T14:12:42.395685Z","shell.execute_reply":"2024-10-31T14:12:42.400914Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{},"outputs":[],"execution_count":null}]}