{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.13","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"},{"sourceId":117001,"sourceType":"modelInstanceVersion","modelInstanceId":98356,"modelId":122533},{"sourceId":117611,"sourceType":"modelInstanceVersion","modelInstanceId":98892,"modelId":123062}],"dockerImageVersionId":30746,"isInternetEnabled":false,"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 cv2\nimport os\nimport time\nimport numpy as np\nimport glob\nimport json\nimport collections\nimport torch\nimport torch.nn as nn\nimport seaborn as sns\nfrom PIL import Image\n\nimport pydicom as dicom\nimport matplotlib.patches as patches\n\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import precision_score, recall_score, f1_score\nfrom sklearn.metrics import roc_auc_score\n\nfrom torch.utils.data import Dataset, DataLoader\nimport torchvision.transforms as transforms\nimport torch\nimport torch.optim.lr_scheduler as lr_scheduler\nimport torch.nn as nn\nimport torchvision.models as models\n\nfrom tqdm import tqdm\n\nfrom matplotlib import animation, rc\nimport pandas as pd\n\n# import pydicom as dicom # dicom\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport pydicom\nimport matplotlib.pyplot as plt\nimport warnings\nimport random\nfrom copy import deepcopy","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-10-08T17:10:44.977816Z","iopub.execute_input":"2024-10-08T17:10:44.978069Z","iopub.status.idle":"2024-10-08T17:10:52.683718Z","shell.execute_reply.started":"2024-10-08T17:10:44.978044Z","shell.execute_reply":"2024-10-08T17:10:52.682739Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# read data\ntrain_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\n\ntrain  = pd.read_csv(train_path + 'train.csv')\nlabel = pd.read_csv(train_path + 'train_label_coordinates.csv')\ntrain_desc  = pd.read_csv(train_path + 'train_series_descriptions.csv')\ntest_desc   = pd.read_csv(train_path + 'test_series_descriptions.csv')\nsub         = pd.read_csv(train_path + 'sample_submission.csv')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:10:52.685237Z","iopub.execute_input":"2024-10-08T17:10:52.685652Z","iopub.status.idle":"2024-10-08T17:10:52.848340Z","shell.execute_reply.started":"2024-10-08T17:10:52.685627Z","shell.execute_reply":"2024-10-08T17:10:52.847521Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(test_desc)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:10:52.853792Z","iopub.execute_input":"2024-10-08T17:10:52.854336Z","iopub.status.idle":"2024-10-08T17:10:52.868920Z","shell.execute_reply.started":"2024-10-08T17:10:52.854301Z","shell.execute_reply":"2024-10-08T17:10:52.868064Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data visualization","metadata":{}},{"cell_type":"code","source":"test_desc.head(20)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:10:52.869958Z","iopub.execute_input":"2024-10-08T17:10:52.870434Z","iopub.status.idle":"2024-10-08T17:10:52.881963Z","shell.execute_reply.started":"2024-10-08T17:10:52.870398Z","shell.execute_reply":"2024-10-08T17:10:52.881166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"Total cases: \", len(train))","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:10:52.883320Z","iopub.execute_input":"2024-10-08T17:10:52.883577Z","iopub.status.idle":"2024-10-08T17:10:52.888856Z","shell.execute_reply.started":"2024-10-08T17:10:52.883555Z","shell.execute_reply":"2024-10-08T17:10:52.887869Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train.columns","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:10:52.889908Z","iopub.execute_input":"2024-10-08T17:10:52.890203Z","iopub.status.idle":"2024-10-08T17:10:52.901424Z","shell.execute_reply.started":"2024-10-08T17:10:52.890180Z","shell.execute_reply":"2024-10-08T17:10:52.900634Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"figure, axis = plt.subplots(1, 3, figsize=(20, 5))\nfor idx, d in enumerate(['foraminal', 'subarticular', 'canal']):\n    diagnosis = list(filter(lambda x: x.find(d) > -1, train.columns))\n    dff = train[diagnosis]\n    with warnings.catch_warnings():\n        warnings.simplefilter(action='ignore', category = FutureWarning)\n        value_counts = dff.apply(pd.value_counts).fillna(0).T\n    value_counts.plot(kind='bar', stacked=True, ax=axis[idx])\n    axis[idx].set_title(f'{d} distribution')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:10:52.902511Z","iopub.execute_input":"2024-10-08T17:10:52.902805Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# generate image paths 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","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:10:54.263956Z","iopub.execute_input":"2024-10-08T17:10:54.264272Z","iopub.status.idle":"2024-10-08T17:10:54.272023Z","shell.execute_reply.started":"2024-10-08T17:10:54.264245Z","shell.execute_reply":"2024-10-08T17:10:54.271134Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 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-10-08T17:10:54.275512Z","iopub.execute_input":"2024-10-08T17:10:54.275877Z","iopub.status.idle":"2024-10-08T17:11:29.745633Z","shell.execute_reply.started":"2024-10-08T17:10:54.275852Z","shell.execute_reply":"2024-10-08T17:11:29.744807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_desc)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:29.746783Z","iopub.execute_input":"2024-10-08T17:11:29.747080Z","iopub.status.idle":"2024-10-08T17:11:29.752723Z","shell.execute_reply.started":"2024-10-08T17:11:29.747055Z","shell.execute_reply":"2024-10-08T17:11:29.751890Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_image_paths)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:29.753823Z","iopub.execute_input":"2024-10-08T17:11:29.754430Z","iopub.status.idle":"2024-10-08T17:11:29.764520Z","shell.execute_reply.started":"2024-10-08T17:11:29.754406Z","shell.execute_reply":"2024-10-08T17:11:29.763655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_image_paths_2 = []\nfor l_in in train_image_paths:\n    for element in l_in:\n        train_image_paths_2.append(element)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:29.765929Z","iopub.execute_input":"2024-10-08T17:11:29.766298Z","iopub.status.idle":"2024-10-08T17:11:29.811937Z","shell.execute_reply.started":"2024-10-08T17:11:29.766267Z","shell.execute_reply":"2024-10-08T17:11:29.811243Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len(train_image_paths_2)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:29.812850Z","iopub.execute_input":"2024-10-08T17:11:29.813098Z","iopub.status.idle":"2024-10-08T17:11:29.818311Z","shell.execute_reply.started":"2024-10-08T17:11:29.813077Z","shell.execute_reply":"2024-10-08T17:11:29.817488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#function to open and display dicom images\ndef display_dicom_images(image_paths):\n    plt.figure(figsize=(15, 5))\n    for i, path in enumerate(image_paths[:3][: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()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:29.819287Z","iopub.execute_input":"2024-10-08T17:11:29.819522Z","iopub.status.idle":"2024-10-08T17:11:29.828911Z","shell.execute_reply.started":"2024-10-08T17:11:29.819501Z","shell.execute_reply":"2024-10-08T17:11:29.828170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"display_dicom_images(train_image_paths_2)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:29.829954Z","iopub.execute_input":"2024-10-08T17:11:29.830289Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 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):\n        study_id = int(path.split('/')[-3])\n        series_id = int(path.split('/')[-2])\n        \n        filtered_labels = label_df[(label_df['study_id'] == study_id) & (label_df['series_id'] == series_id)]\n        \n        \n        ds = pydicom.dcmread(path)\n        \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        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()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def display_and_crop_dicom(image_paths, label_df, crop_size=50):\n    fig, axs = plt.subplots(1 ,len(image_paths), figsize=(18, 6))\n    \n    for idx, path in enumerate(image_paths):\n        study_id = int(path.split('/')[-3])\n        series_id = int(path.split('/')[-2])\n        \n        filtered_labels = label_df[(label_df['study_id'] == study_id) & (label_df['series_id'] == series_id)]\n        \n        ds = pydicom.dcmread(path)\n        img = ds.pixel_array\n        \n        axs[idx].imshow(img, cmap='gray')\n        axs[idx].set_title(f'Study ID: {study_id}, Series ID: {series_id}')\n        axs[idx].axis('off')\n        \n        for _, row in filtered_labels.iterrows():\n            x, y = int(row['x']), int(row['y'])\n            \n            x1, y1 = max(0, x - crop_size // 2), max(0, y - crop_size // 2)\n            x2, y2 = min(img.shape[1], x + crop_size // 2), min(img.shape[0], y + crop_size // 2)\n            \n            cropped_img = img[y1:y2, x1:x2]\n            \n            axs[idx].plot(row['x'], row['y'], 'ro', markersize=5)\n            fig_crop, ax_crop = plt.subplots()\n            ax_crop.imshow(cropped_img, cmap='gray')\n            ax_crop.set_title(f'cropped roi at ({x}, {y})')\n            ax_crop.axis('off')\n        \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#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","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"#display dicom image with coordinates example\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])\n        \ndisplay_dicom_with_coordinates(image_paths, label)","metadata":{"execution":{"iopub.status.idle":"2024-10-08T17:11:31.418821Z","shell.execute_reply.started":"2024-10-08T17:11:30.539087Z","shell.execute_reply":"2024-10-08T17:11:31.417960Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"study_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])\n\ndisplay_and_crop_dicom(image_paths, label, crop_size=50)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:31.419932Z","iopub.execute_input":"2024-10-08T17:11:31.420269Z","iopub.status.idle":"2024-10-08T17:11:35.879340Z","shell.execute_reply.started":"2024-10-08T17:11:31.420240Z","shell.execute_reply":"2024-10-08T17:11:35.878127Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"image_paths","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:35.880672Z","iopub.execute_input":"2024-10-08T17:11:35.881027Z","iopub.status.idle":"2024-10-08T17:11:35.888288Z","shell.execute_reply.started":"2024-10-08T17:11:35.880994Z","shell.execute_reply":"2024-10-08T17:11:35.887117Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data preparation and preprocessing","metadata":{}},{"cell_type":"code","source":"def 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)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:35.889686Z","iopub.execute_input":"2024-10-08T17:11:35.890052Z","iopub.status.idle":"2024-10-08T17:11:35.899564Z","shell.execute_reply.started":"2024-10-08T17:11:35.890020Z","shell.execute_reply":"2024-10-08T17:11:35.898609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"new_train_df = pd.concat([reshape_row(row) for _, row in train.iterrows()], ignore_index=True)\nnew_train_df.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:35.900676Z","iopub.execute_input":"2024-10-08T17:11:35.901084Z","iopub.status.idle":"2024-10-08T17:11:37.206133Z","shell.execute_reply.started":"2024-10-08T17:11:35.901052Z","shell.execute_reply":"2024-10-08T17:11:37.205242Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(\"\\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-10-08T17:11:37.207457Z","iopub.execute_input":"2024-10-08T17:11:37.207817Z","iopub.status.idle":"2024-10-08T17:11:37.214271Z","shell.execute_reply.started":"2024-10-08T17:11:37.207787Z","shell.execute_reply":"2024-10-08T17:11:37.213408Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"merged_df = pd.merge(new_train_df, label, on=['study_id', 'condition', 'level'], how='inner')\nfinal_merged_df = pd.merge(merged_df, train_desc, on='series_id', how='inner')","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:37.215560Z","iopub.execute_input":"2024-10-08T17:11:37.215933Z","iopub.status.idle":"2024-10-08T17:11:37.297077Z","shell.execute_reply.started":"2024-10-08T17:11:37.215899Z","shell.execute_reply":"2024-10-08T17:11:37.296034Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df = pd.merge(merged_df, train_desc, on=['series_id', 'study_id'], how='inner')\nfinal_merged_df.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:37.298470Z","iopub.execute_input":"2024-10-08T17:11:37.298761Z","iopub.status.idle":"2024-10-08T17:11:37.328318Z","shell.execute_reply.started":"2024-10-08T17:11:37.298736Z","shell.execute_reply":"2024-10-08T17:11:37.327182Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-10-08T17:11:37.335812Z","iopub.execute_input":"2024-10-08T17:11:37.336090Z","iopub.status.idle":"2024-10-08T17:11:37.360106Z","shell.execute_reply.started":"2024-10-08T17:11:37.336066Z","shell.execute_reply":"2024-10-08T17:11:37.359259Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df[final_merged_df['series_id'] == 1012284084].sort_values(\"instance_number\")","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:37.361234Z","iopub.execute_input":"2024-10-08T17:11:37.361613Z","iopub.status.idle":"2024-10-08T17:11:37.378283Z","shell.execute_reply.started":"2024-10-08T17:11:37.361577Z","shell.execute_reply":"2024-10-08T17:11:37.377171Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"filtered_df = final_merged_df[final_merged_df['study_id'] == 1013589491].sort_values(\"instance_number\")\n\nfiltered_df","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:37.379340Z","iopub.execute_input":"2024-10-08T17:11:37.379615Z","iopub.status.idle":"2024-10-08T17:11:37.402390Z","shell.execute_reply.started":"2024-10-08T17:11:37.379591Z","shell.execute_reply":"2024-10-08T17:11:37.401462Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sorted_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-10-08T17:11:37.403696Z","iopub.execute_input":"2024-10-08T17:11:37.404038Z","iopub.status.idle":"2024-10-08T17:11:37.427843Z","shell.execute_reply.started":"2024-10-08T17:11:37.404008Z","shell.execute_reply":"2024-10-08T17:11:37.426945Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df['row_id'] = (final_merged_df['study_id'].astype(str) + '_' + final_merged_df['condition'].str.lower().str.replace(' ','_') + '_' + final_merged_df['level'].str.lower().str.replace('/', '_'))\n\nfinal_merged_df['image_path'] = (f'{train_path}/train_images/' + final_merged_df['study_id'].astype(str) + '/' + final_merged_df['series_id'].astype(str) + '/' + final_merged_df['instance_number'].astype(str) + '.dcm')\n\nfinal_merged_df.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:37.429001Z","iopub.execute_input":"2024-10-08T17:11:37.429348Z","iopub.status.idle":"2024-10-08T17:11:37.670653Z","shell.execute_reply.started":"2024-10-08T17:11:37.429316Z","shell.execute_reply":"2024-10-08T17:11:37.669563Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Normal/Mild\"].value_counts().sum()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:37.671959Z","iopub.execute_input":"2024-10-08T17:11:37.672340Z","iopub.status.idle":"2024-10-08T17:11:37.831859Z","shell.execute_reply.started":"2024-10-08T17:11:37.672306Z","shell.execute_reply":"2024-10-08T17:11:37.830978Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Moderate\"].value_counts().sum()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:37.833474Z","iopub.execute_input":"2024-10-08T17:11:37.833829Z","iopub.status.idle":"2024-10-08T17:11:37.883788Z","shell.execute_reply.started":"2024-10-08T17:11:37.833796Z","shell.execute_reply":"2024-10-08T17:11:37.882885Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df[final_merged_df[\"severity\"] == \"Severe\"].value_counts().sum()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:37.884959Z","iopub.execute_input":"2024-10-08T17:11:37.885382Z","iopub.status.idle":"2024-10-08T17:11:37.917580Z","shell.execute_reply.started":"2024-10-08T17:11:37.885349Z","shell.execute_reply":"2024-10-08T17:11:37.916672Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df[final_merged_df['series_description'] == 'Sagittal T1'].value_counts().sum()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:37.918626Z","iopub.execute_input":"2024-10-08T17:11:37.918890Z","iopub.status.idle":"2024-10-08T17:11:38.005067Z","shell.execute_reply.started":"2024-10-08T17:11:37.918866Z","shell.execute_reply":"2024-10-08T17:11:38.004199Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df[final_merged_df['series_description'] == 'Axial T2'].value_counts().sum()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:38.006275Z","iopub.execute_input":"2024-10-08T17:11:38.006602Z","iopub.status.idle":"2024-10-08T17:11:38.094002Z","shell.execute_reply.started":"2024-10-08T17:11:38.006565Z","shell.execute_reply":"2024-10-08T17:11:38.093073Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df[final_merged_df['series_description'] == 'Sagittal T2/STIR'].value_counts().sum()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:38.095246Z","iopub.execute_input":"2024-10-08T17:11:38.095550Z","iopub.status.idle":"2024-10-08T17:11:38.143415Z","shell.execute_reply.started":"2024-10-08T17:11:38.095526Z","shell.execute_reply":"2024-10-08T17:11:38.142436Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# base path for test images\nbase_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/'","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:38.144460Z","iopub.execute_input":"2024-10-08T17:11:38.144732Z","iopub.status.idle":"2024-10-08T17:11:38.148969Z","shell.execute_reply.started":"2024-10-08T17:11:38.144708Z","shell.execute_reply":"2024-10-08T17:11:38.147929Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# 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\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\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):\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            \nexpanded_test_desc = pd.DataFrame(expanded_rows)\n\nexpanded_test_desc.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:38.150197Z","iopub.execute_input":"2024-10-08T17:11:38.150527Z","iopub.status.idle":"2024-10-08T17:11:38.233712Z","shell.execute_reply.started":"2024-10-08T17:11:38.150495Z","shell.execute_reply":"2024-10-08T17:11:38.232858Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(expanded_test_desc['image_path'][150])","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:38.234943Z","iopub.execute_input":"2024-10-08T17:11:38.235317Z","iopub.status.idle":"2024-10-08T17:11:38.241446Z","shell.execute_reply.started":"2024-10-08T17:11:38.235289Z","shell.execute_reply":"2024-10-08T17:11:38.240410Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_merged_df['severity'] = final_merged_df['severity'].map({'Normal/Mild': 'normal_mild', 'Moderate': 'moderate', 'Severe': 'severe'})","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:38.242847Z","iopub.execute_input":"2024-10-08T17:11:38.243389Z","iopub.status.idle":"2024-10-08T17:11:38.258809Z","shell.execute_reply.started":"2024-10-08T17:11:38.243355Z","shell.execute_reply":"2024-10-08T17:11:38.257785Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_data = expanded_test_desc\ntrain_data = final_merged_df","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:38.260134Z","iopub.execute_input":"2024-10-08T17:11:38.260512Z","iopub.status.idle":"2024-10-08T17:11:38.268028Z","shell.execute_reply.started":"2024-10-08T17:11:38.260480Z","shell.execute_reply":"2024-10-08T17:11:38.267150Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def check_exists(path):\n    return os.path.exists(path)\n\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\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\ndef check_image_exists(row):\n    image_path = row['image_path']\n    return check_exists(image_path)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:38.269141Z","iopub.execute_input":"2024-10-08T17:11:38.269423Z","iopub.status.idle":"2024-10-08T17:11:38.278390Z","shell.execute_reply.started":"2024-10-08T17:11:38.269399Z","shell.execute_reply":"2024-10-08T17:11:38.277517Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_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\ntrain_data = train_data[(train_data['study_id_exists']) & (train_data['series_id_exists']) & (train_data['image_exists'])]","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:38.279382Z","iopub.execute_input":"2024-10-08T17:11:38.279700Z","iopub.status.idle":"2024-10-08T17:11:56.069447Z","shell.execute_reply.started":"2024-10-08T17:11:38.279677Z","shell.execute_reply":"2024-10-08T17:11:56.068609Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.head(50)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.070709Z","iopub.execute_input":"2024-10-08T17:11:56.071332Z","iopub.status.idle":"2024-10-08T17:11:56.106024Z","shell.execute_reply.started":"2024-10-08T17:11:56.071299Z","shell.execute_reply":"2024-10-08T17:11:56.105100Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data['image_path'][10]","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.107179Z","iopub.execute_input":"2024-10-08T17:11:56.107431Z","iopub.status.idle":"2024-10-08T17:11:56.113193Z","shell.execute_reply.started":"2024-10-08T17:11:56.107408Z","shell.execute_reply":"2024-10-08T17:11:56.112241Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_dicom(path):\n    dicom = pydicom.read_file(path)\n    data = dicom.pixel_array\n    data = data - np.min(data)\n    if np.max(data) != 0:\n        data = data / np.max(data)\n    data = (data * 255).astype(np.uint8)\n    return data","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.114317Z","iopub.execute_input":"2024-10-08T17:11:56.114580Z","iopub.status.idle":"2024-10-08T17:11:56.121997Z","shell.execute_reply.started":"2024-10-08T17:11:56.114552Z","shell.execute_reply":"2024-10-08T17:11:56.121249Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images = []\nrow_ids = []\nselected_indices = random.sample(range(len(train_data)), 2)\n\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    \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')\n    \nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.123066Z","iopub.execute_input":"2024-10-08T17:11:56.123359Z","iopub.status.idle":"2024-10-08T17:11:56.473039Z","shell.execute_reply.started":"2024-10-08T17:11:56.123337Z","shell.execute_reply":"2024-10-08T17:11:56.472170Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data.head(25)\nlen(train_data)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.474363Z","iopub.execute_input":"2024-10-08T17:11:56.474798Z","iopub.status.idle":"2024-10-08T17:11:56.481356Z","shell.execute_reply.started":"2024-10-08T17:11:56.474762Z","shell.execute_reply":"2024-10-08T17:11:56.480168Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data = train_data.dropna()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.482593Z","iopub.execute_input":"2024-10-08T17:11:56.482927Z","iopub.status.idle":"2024-10-08T17:11:56.524563Z","shell.execute_reply.started":"2024-10-08T17:11:56.482896Z","shell.execute_reply":"2024-10-08T17:11:56.523819Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"train_data[train_data['series_description'] == 'Sagittal T1'].value_counts().sum()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.525716Z","iopub.execute_input":"2024-10-08T17:11:56.525987Z","iopub.status.idle":"2024-10-08T17:11:56.610408Z","shell.execute_reply.started":"2024-10-08T17:11:56.525963Z","shell.execute_reply":"2024-10-08T17:11:56.609459Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class CroppedDataset(Dataset):\n    def __init__(self,\n                image_paths,\n                label_df,\n                crop_size=50,\n                transform=None):\n        \n        self.image_paths = image_paths\n        self.label_df = label_df\n        self.crop_size = crop_size\n        self.transform = transform\n        self.cropped_images = []\n        self.labels = []\n        \n        self.load_and_crop_images()\n        \n    def load_and_crop_images(self):\n        print('Cropping images, please wait...')\n        for idx, path in enumerate(tqdm(self.image_paths, desc='Cropping progress', unit='image')):\n            study_id = int(path.split('/')[-3])\n            series_id = int(path.split('/')[-2])\n            \n            filtered_labels = self.label_df[\n                (self.label_df['study_id'] == study_id) & (self.label_df['series_id'] == series_id)\n            ]\n            \n            img = load_dicom(path)\n            \n            for _, row in filtered_labels.iterrows():\n                x, y = int(row['x']), int(row['y'])\n                x1, y1 = max(0, x - self.crop_size // 2), max(0, y - self.crop_size // 2)\n                x2, y2 = min(img.shape[1], x + self.crop_size // 2), min(img.shape[0], y + self.crop_size // 2)\n                \n                # Crop the image\n                cropped_img = img[y1:y2, x1:x2]\n                \n                \n                self.cropped_images.append(cropped_img)\n                self.labels.append(row['severity'])  # Assuming label column exists\n                \n    def __len__(self):\n        return len(self.cropped_images)\n    \n    def __getitem__(self, idx):\n        cropped_img = self.cropped_images[idx]\n        label = self.labels[idx]\n        \n        if self.transform:\n            cropped_img = self.transform(cropped_img)\n            \n        return cropped_img, label","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.611395Z","iopub.execute_input":"2024-10-08T17:11:56.611652Z","iopub.status.idle":"2024-10-08T17:11:56.623548Z","shell.execute_reply.started":"2024-10-08T17:11:56.611629Z","shell.execute_reply":"2024-10-08T17:11:56.622589Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class 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)\n        label = self.dataframe['severity'][index]\n        \n        if self.transform:\n            image = self.transform(image)\n            \n        return image, label","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.624660Z","iopub.execute_input":"2024-10-08T17:11:56.624934Z","iopub.status.idle":"2024-10-08T17:11:56.636208Z","shell.execute_reply.started":"2024-10-08T17:11:56.624911Z","shell.execute_reply":"2024-10-08T17:11:56.635503Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_datasets_and_loaders(df, series_description, transform, \n#                                 batch_size=8):\n                                batch_size=32):\n#                                 batch_size=64):\n    filtered_df = df[df['series_description'] == series_description]\n\n    train_df, val_df = train_test_split(filtered_df, test_size=0.2, random_state=42)\n    train_df = train_df.reset_index(drop=True)\n    val_df = val_df.reset_index(drop=True)\n\n    train_dataset = CustomDataset(train_df, transform)\n    val_dataset = CustomDataset(val_df, transform)\n\n    trainloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=4, pin_memory=True)\n    valloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=4, pin_memory=True)\n\n    return trainloader, valloader, len(train_df), len(val_df)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.637143Z","iopub.execute_input":"2024-10-08T17:11:56.637396Z","iopub.status.idle":"2024-10-08T17:11:56.651060Z","shell.execute_reply.started":"2024-10-08T17:11:56.637374Z","shell.execute_reply":"2024-10-08T17:11:56.650163Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def create_cropped_datasets_and_loaders(df, series_description, transform, \n#                                         image_paths, \n#                                 batch_size=8):\n                                batch_size=32):\n#                                 batch_size=64):\n    filtered_df = df[df['series_description'] == series_description]\n\n    train_df, val_df = train_test_split(filtered_df, test_size=0.2, random_state=42)\n    \n    train_df = train_df.reset_index(drop=True)\n    val_df = val_df.reset_index(drop=True)\n    \n    train_image_path = train_df[train_df['series_description'] == series_description]['image_path'].tolist()\n    val_image_path = val_df[val_df['series_description'] == series_description]['image_path'].tolist()\n    \n    unique_train_image_path = list(set(train_image_path))\n    unique_val_image_path = list(set(val_image_path))\n    \n    \n#     train_dataset = CroppedDataset(train_df, transform)\n#     val_dataset = CroppedDataset(val_df, transform)\n\n    train_dataset = CroppedDataset(unique_train_image_path, train_df, 50, transform)\n    val_dataset = CroppedDataset(unique_val_image_path, val_df, 50, transform)\n\n    trainloader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=4, pin_memory=True)\n    valloader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=4, pin_memory=True)\n\n    return trainloader, valloader, len(train_df), len(val_df)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.652071Z","iopub.execute_input":"2024-10-08T17:11:56.652362Z","iopub.status.idle":"2024-10-08T17:11:56.661315Z","shell.execute_reply.started":"2024-10-08T17:11:56.652340Z","shell.execute_reply":"2024-10-08T17:11:56.660475Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# image_paths_t1 = train_data[train_data['series_description'] == 'Sagittal T1']['image_path'].tolist()\n# image_paths_t2 = train_data[train_data['series_description'] == 'Axial T2']['image_path'].tolist()\n# image_paths_t2stir = train_data[train_data['series_description'] == 'Sagittal T2/STIR']['image_path'].tolist()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.662527Z","iopub.execute_input":"2024-10-08T17:11:56.663052Z","iopub.status.idle":"2024-10-08T17:11:56.674123Z","shell.execute_reply.started":"2024-10-08T17:11:56.663021Z","shell.execute_reply":"2024-10-08T17:11:56.673383Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = transforms.Compose([\n#     transforms.Lambda(lambda x: (x * 255).astype(np.uint8)),\n#     transforms.Lambda(lambda x: (x - np.min(x)) / (np.max(x) - np.min(x))),\n    transforms.ToPILImage(),\n    transforms.Resize((224, 224)),\n#     transforms.Grayscale(num_output_channels=3),\n    transforms.ToTensor(),\n])\n\ndataloaders = {}\nlenghts = {}\n\n\ntrainloader_t1, valloader_t1, len_train_t1, len_val_t1 = create_cropped_datasets_and_loaders(train_data, 'Sagittal T1', transform)\ntrainloader_t2, valloader_t2, len_train_t2, len_val_t2 = create_cropped_datasets_and_loaders(train_data, 'Axial T2', transform)\ntrainloader_t2stir, valloader_t2stir, len_train_t2stir, len_val_t2stir = create_cropped_datasets_and_loaders(train_data, 'Sagittal T2/STIR', transform)\n\ndataloaders['Sagittal T1'] = (trainloader_t1, valloader_t1)\ndataloaders['Axial T2'] = (trainloader_t2, valloader_t2)\ndataloaders['Sagittal T2/STIR'] = (trainloader_t2stir, valloader_t2stir)\n\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:11:56.675028Z","iopub.execute_input":"2024-10-08T17:11:56.675288Z","iopub.status.idle":"2024-10-08T17:22:40.966027Z","shell.execute_reply.started":"2024-10-08T17:11:56.675266Z","shell.execute_reply":"2024-10-08T17:22:40.965003Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"lengths = {}\n\nlengths['Sagittal T1'] = (len_train_t1, len_val_t1)\nlengths['Axial T2'] = (len_train_t2, len_val_t2)\nlengths['Sagittal T2/STIR'] = (len_train_t2stir, len_val_t2stir)\n\nlabel_map = {'Mild': 0, 'Moderate': 1, 'Severe': 2}","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:40.967550Z","iopub.execute_input":"2024-10-08T17:22:40.968009Z","iopub.status.idle":"2024-10-08T17:22:40.974146Z","shell.execute_reply.started":"2024-10-08T17:22:40.967971Z","shell.execute_reply":"2024-10-08T17:22:40.973185Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"len_train_t2 = len(trainloader_t2.dataset)\nlen_val_t2 = len(valloader_t2.dataset)\nlen_train_t1 = len(trainloader_t1.dataset)\nlen_val_t1 = len(valloader_t1.dataset)\nlen_train_t2stir = len(trainloader_t2stir.dataset)\nlen_val_t2stir = len(valloader_t2stir.dataset)\n\nprint(f\"Number of training samples: {len_train_t2}\")\nprint(f\"Number of validation samples: {len_val_t2}\")\n\nprint(f\"Number of validation samples: {len_train_t1}\")\nprint(f\"Number of validation samples: {len_val_t1}\")\n\nprint(f\"Number of validation samples: {len_train_t2stir}\")\nprint(f\"Number of validation samples: {len_val_t2stir}\")","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:40.975418Z","iopub.execute_input":"2024-10-08T17:22:40.975757Z","iopub.status.idle":"2024-10-08T17:22:40.988825Z","shell.execute_reply.started":"2024-10-08T17:22:40.975730Z","shell.execute_reply":"2024-10-08T17:22:40.987796Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# NO CROP PROCEDURE\n\n# #define transforms\n# transform = transforms.Compose([\n#     transforms.Lambda(lambda x: (x * 255).astype(np.uint8)),\n#     transforms.ToPILImage(),\n#     transforms.Resize((224, 224)),\n#     transforms.Grayscale(num_output_channels=3),\n#     transforms.ToTensor(),\n# ])\n\n# dataloaders = {}\n# lengths  = {}\n\n# trainloader_t1, valloader_t1, len_train_t1, len_val_t1 = create_datasets_and_loaders(train_data, 'Sagittal T1', transform)\n# trainloader_t2, valloader_t2, len_train_t2, len_val_t2 = create_datasets_and_loaders(train_data, 'Axial T2', transform)\n# trainloader_t2stir, valloader_t2stir, len_train_t2stir, len_val_t2stir = create_datasets_and_loaders(train_data, 'Sagittal T2/STIR', transform)\n\n# dataloaders['Sagittal T1'] = (trainloader_t1, valloader_t1)\n# dataloaders['Axial T2'] = (trainloader_t2, valloader_t2)\n# dataloaders['Sagittal T2/STIR'] = (trainloader_t2stir, valloader_t2stir)\n\n# lengths['Sagittal T1'] = (len_train_t1, len_val_t1)\n# lengths['Axial T2'] = (len_train_t2, len_val_t2)\n# lengths['Sagittal T2/STIR'] = (len_train_t2stir, len_val_t2stir)\n\n# label_map = {'Mild': 0, 'Moderate': 1, 'Severe': 2}","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:40.989697Z","iopub.execute_input":"2024-10-08T17:22:40.989961Z","iopub.status.idle":"2024-10-08T17:22:41.001898Z","shell.execute_reply.started":"2024-10-08T17:22:40.989939Z","shell.execute_reply":"2024-10-08T17:22:41.000962Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def visualize_batch(dataloader, num_images=6):\n    images, labels = next(iter(dataloader))\n    \n    images = images[:num_images]\n    labesl = labels[:num_images]\n    \n    fig, axes = plt.subplots(1, len(images), figsize=(15, 5))\n    \n    if len(images) == 1:\n        axes = [axes]\n        \n    for i, (img, lbl) in enumerate(zip(images, labels)):\n        ax = axes[i]\n        \n        img = img.permute(1, 2, 0)\n        \n        img = (img - img.min()) / (img.max() - img.min())\n        \n        ax.imshow(img, cmap= 'gray')\n        ax.set_title(f\"Label: {lbl}\", fontsize=10)\n        ax.axis('off')\n        \n    plt.tight_layout()\n    plt.show()\n\nprint(\"Visualizing Sagittal T1 samples\")\nvisualize_batch(trainloader_t1, 10)\nprint(\"Visualizing Axial T2 samples\")\nvisualize_batch(trainloader_t2, 10)\nprint(\"Visualizing Sagittal T2/STIR samples\")\nvisualize_batch(trainloader_t2stir, 10)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:41.003032Z","iopub.execute_input":"2024-10-08T17:22:41.003367Z","iopub.status.idle":"2024-10-08T17:22:45.503083Z","shell.execute_reply.started":"2024-10-08T17:22:41.003337Z","shell.execute_reply":"2024-10-08T17:22:45.502096Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"images, labels = next(iter(trainloader_t2))\n\nimage_to_display = images[0]\n\nif image_to_display.dim() == 3:\n    if image_to_display.shape[0] == 1:\n        image_to_display = image_to_display.squeeze(0)\n    else:\n        image_to_display = image_to_display.permute(1, 2, 0)\n        \n\nimage_to_display = (image_to_display - image_to_display.min()) / (image_to_display.max() - image_to_display.min())\n\nplt.figure(figsize=(8, 4))\nplt.imshow(image_to_display, cmap='gray')\nplt.title(f'Label: {labels[0]}')\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:45.504443Z","iopub.execute_input":"2024-10-08T17:22:45.504786Z","iopub.status.idle":"2024-10-08T17:22:46.370017Z","shell.execute_reply.started":"2024-10-08T17:22:45.504752Z","shell.execute_reply":"2024-10-08T17:22:46.368947Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Building models","metadata":{}},{"cell_type":"code","source":"class CustomModel(nn.Module):\n    def __init__(self, num_classes=3, pretrained_weights=None):\n        super(CustomModel, self).__init__()\n        self.model = models.resnet18(weights=None)\n        \n        #adjust for grayscale images\n#         self.model.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False)\n        \n        if pretrained_weights:\n            state_dict = torch.load(pretrained_weights)\n            self.model.load_state_dict(state_dict)\n#             self.model.load_state_dict(torch.load(pretrained_weights))\n\n            with torch.no_grad():\n                # conv1 weights are of shape (64, 3, 7, 7) -> to (64, 1, 7, 7)\n                # average across 3 channels\n                self.model.conv1.weight = nn.Parameter(self.model.conv1.weight.mean(dim=1, keepdim=True))\n                self.model.conv1 = nn.Conv2d(1, 64, kernel_size=7, stride=2, padding=3, bias=False)\n                \n        \n        num_ftrs = self.model.fc.in_features\n        self.model.fc = nn.Linear(num_ftrs, num_classes)\n            \n    def forward(self, x):\n        return self.model(x)\n    \n    def unfreeze_model(self):\n        for name, param in self.model.named_parameters():\n            if \"bn\" not in name:\n                param.requires_grad = True\n        \n        for param in self.model.fc.parameters():\n            param.requires_grad = True","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:46.371660Z","iopub.execute_input":"2024-10-08T17:22:46.372090Z","iopub.status.idle":"2024-10-08T17:22:46.382758Z","shell.execute_reply.started":"2024-10-08T17:22:46.372031Z","shell.execute_reply":"2024-10-08T17:22:46.381853Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class ResidualBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, ks1, ks2, stride=1, drop=False):\n        self.conv1 = nn.Conv2d(in_channels, out_channels, kernel_size=ks1, stride=stride, padding=(ks1-1) // 2)\n        self.bn1 = nn.BatchNorm2d(out_channels)\n        self.elu = nn.ELU(inplace=True)\n        self.dropout = nn.Dropout(p=0.2) if drop else nn.Dropout(p=0.0)\n        self.conv2 = nn.Conv2d(out_channels, out_channels, kernel_size=ks2, stride=1, padding=(ks2-1)//2, bias=False)\n        self.bn2 = nn.BatchNorm2d(out_channels)\n        \n        self.shortcut = nn.Sequential()\n        if stride != 1 or in_channels != out_channels:\n            self.shortcut = nn.Sequential(\n                nn.Conv2d(in_channels, out_channels, kernel_size=1, stride=stride, bias=False),\n                nn.BatchNorm2d(out_channels)\n            )\n            \n        \n    \n    def forward(self, x):\n        out = self.elu(self.bn1(self.conv1(x)))\n        out = self.dropout(out)\n        out = self.bn2(self.conv2(out))\n        out += self.shortcut(x)\n        out = self.elu(out)\n        return out ","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:46.383935Z","iopub.execute_input":"2024-10-08T17:22:46.384287Z","iopub.status.idle":"2024-10-08T17:22:46.396580Z","shell.execute_reply.started":"2024-10-08T17:22:46.384253Z","shell.execute_reply":"2024-10-08T17:22:46.395820Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# class residualModel(nn.Module):\n#     def __init__(self):\n#         super(residualModel, self).__init__()\n#         self.conv1 = nn.Conv2d(3, 32, kernel_size=7, stride=1, padding=3)\n#         self.bn1 = nn.BatchNorm2d(32)\n#         self.elu = nn.ELU(inplace=True)\n        \n#         self.resblock1 = ResidualBlock(32, 32, stride=1, ks1=5, ks2=3, drop=True)\n#         self.resblock2 = ResidualBlock(32, 32, stride=1, ks1=3, ks2=5, drop=True)        \n#         self.resblock3 = ResidualBlock(32, 32, stride=1, ks1=5, ks2=3, drop=True)\n        \n#         self.resblock4 = ResidualBlock(32, 64, stride=2, ks1=3, ks2=5, drop=True)        \n#         self.resblock5 = ResidualBlock(64, 64, stride=1, ks1=5, ks2=3, drop=True)\n#         self.resblock6 = ResidualBlock(64, 128, stride=2, ks1=3, ks2=5, drop=True)        \n#         self.resblock7 = ResidualBlock(128, 128, stride=1, ks1=5, ks2=3, drop=True)        \n        \n#         self.resblock8 = ResidualBlock(128, 256, stride=2, ks1=3, ks2=5, drop=True)        \n#         self.resblock9 = ResidualBlock(256, 256, stride=1, ks1=5, ks2=3, drop=True)        \n#         self.resblock10 = ResidualBlock(512, 512, stride=2, ks1=3, ks2=5, drop=True)\n#         self.resblock11 = ResidualBlock(512, 512, stride=1, ks1=5, ks2=3, drop=True)        \n#         self.resblock12 = ResidualBlock(1024, 1024, stride=2, ks1=3, ks2=5, drop=True)\n        \n#         self.pool = nn.AdaptiveAvgPool2d((4, 4))\n#         self.fc = nn.Linear(1024 * 4 * 4, 256)\n#         self.fc2 = nn.Linear(256, 3)\n        \n#         self.dropout1 = nn.Dropout(p=0.2)\n#         self.dropout2 = nn.Dropout(p=0.2)\n        \n#     def forward(self, x):\n#         x = self.elu(self.bn1(self.conv1(x)))\n#         x = self.resblock1(x)\n# #         x = self.resblock2(x)\n# #         x = self.resblock3(x)\n\n#         x = self.resblock4(x)\n#         x = self.resblock5(x)\n#         x = self.resblock6(x)\n#         x = self.resblock7(x)\n\n#         x = self.resblock8(x)\n#         x = self.resblock9(x)\n#         x = self.resblock10(x)\n#         x = self.resblock11(x)\n#         x = self.resblock12(x)\n\n#         x = self.pool(x)\n#         x = x.view(-1, 1024 * 4 * 4)\n#         x = self.dropout1(self.elu(self.fc(x)))\n#         x = self.dropout2(self.elu(self.fc2(x)))\n#         return x\n        ","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:46.397962Z","iopub.execute_input":"2024-10-08T17:22:46.398414Z","iopub.status.idle":"2024-10-08T17:22:46.411799Z","shell.execute_reply.started":"2024-10-08T17:22:46.398381Z","shell.execute_reply":"2024-10-08T17:22:46.411028Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# define the Grad-CAM class\nclass GradCAM:\n    def __init__(self, model, target_layer):\n        self.model = model\n        self.target_layer = target_layer\n        self.gradients = None\n        self.activation = None\n        \n        self.target_layer.register_forward_hook(self.save_activation)\n        self.target_layer.register_backward_hook(self.save_gradient)\n        \n    def save_activation(self, module, input, output):\n        self.activation = output\n        \n    def save_gradient(self, module, grad_input, grad_output):\n        self.gradients = grad_output[0]\n        \n    def generate(self, input_image, target_class=None):\n        input_image.requires_grad = True\n        model_output = self.model(input_image)\n        \n        if target_class is None:\n            target_class = torch.argmax(model_output)\n            \n        self.model.zero_grad()\n        model_output[:, target_class].backward(retain_graph=True)\n        \n        gradients = self.gradients.cpu().data.numpy()\n        activations = self.activation.cpu().data.numpy()\n        \n        weights = np.mean(gradients, axis=(2, 3)) #average gradients over spatial dimensions\n        cam = np.zeros(activations.shape[2:], dtype=np.float32)\n        \n        for i, w in enumerate(weights[0]):\n            cam += w * activations[0, i, :, :]\n            \n        cam = np.maximum(cam, 0)\n        cam = cv2.resize(cam, (input_image.shape[2], input_image.shape[3]))\n        \n        cam = cam - np.min(cam)\n        cam = cam / np.max(cam)\n        \n        return cam","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:46.412805Z","iopub.execute_input":"2024-10-08T17:22:46.413141Z","iopub.status.idle":"2024-10-08T17:22:46.426292Z","shell.execute_reply.started":"2024-10-08T17:22:46.413092Z","shell.execute_reply":"2024-10-08T17:22:46.425488Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def show_cam_on_image(img, cam):\n    heatmap = cv2.applyColorMap(np.uint8(255 * cam), cv2.COLORMAP_JET)\n    heatmap = np.float32(heatmap) / 255\n    cam_image = heatmap + np.float32(img)\n    cam_image = cam_image / np.max(cam_image)\n    return np.uint8(255 * cam_image)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:46.427379Z","iopub.execute_input":"2024-10-08T17:22:46.427740Z","iopub.status.idle":"2024-10-08T17:22:46.438328Z","shell.execute_reply.started":"2024-10-08T17:22:46.427708Z","shell.execute_reply":"2024-10-08T17:22:46.437542Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def save_checkpoint(state, filename):\n    torch.save(state, filename)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:46.439247Z","iopub.execute_input":"2024-10-08T17:22:46.439514Z","iopub.status.idle":"2024-10-08T17:22:46.447424Z","shell.execute_reply.started":"2024-10-08T17:22:46.439491Z","shell.execute_reply":"2024-10-08T17:22:46.446523Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"device = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nweights_path = '/kaggle/input/resnet18/pytorch/resnet18-lumbar/1/resnet18-f37072fd.pth'\n\n\nsagittal_t1_model = CustomModel(num_classes=3, pretrained_weights=weights_path).to(device)\naxial_t2_model = CustomModel(num_classes=3, pretrained_weights=weights_path).to(device)\nsagittal_t2stir_model = CustomModel(num_classes=3, pretrained_weights=weights_path).to(device)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:46.448583Z","iopub.execute_input":"2024-10-08T17:22:46.448905Z","iopub.status.idle":"2024-10-08T17:22:48.012610Z","shell.execute_reply.started":"2024-10-08T17:22:46.448875Z","shell.execute_reply":"2024-10-08T17:22:48.011587Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Freeze initial layers\nfor name, param in sagittal_t1_model.model.named_parameters():\n    if any(layer in name for layer in ['layer1', 'layer2', 'layer3', 'layer4']):\n        param.requires_grad = False\n        \nfor name, param in axial_t2_model.model.named_parameters():\n    if any(layer in name for layer in ['layer1', 'layer2', 'layer3', 'layer4']):\n        param.requires_grad = False\n        \nfor name, param in sagittal_t2stir_model.model.named_parameters():\n    if any(layer in name for layer in ['layer1', 'layer2', 'layer3', 'layer4']):\n        param.requires_grad = False\n\n# Unfreeze the final fully connected layer\nfor param in sagittal_t1_model.model.fc.parameters():\n    param.requires_grad = True\n    \nfor param in axial_t2_model.model.fc.parameters():\n    param.requires_grad = True\n    \nfor param in sagittal_t2stir_model.model.fc.parameters():\n    param.requires_grad = True","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:48.014245Z","iopub.execute_input":"2024-10-08T17:22:48.014783Z","iopub.status.idle":"2024-10-08T17:22:48.025246Z","shell.execute_reply.started":"2024-10-08T17:22:48.014748Z","shell.execute_reply":"2024-10-08T17:22:48.024319Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from torchinfo import summary\nsummary(sagittal_t1_model)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:48.026447Z","iopub.execute_input":"2024-10-08T17:22:48.026775Z","iopub.status.idle":"2024-10-08T17:22:48.053800Z","shell.execute_reply.started":"2024-10-08T17:22:48.026750Z","shell.execute_reply":"2024-10-08T17:22:48.052892Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Hyperparameter settings","metadata":{}},{"cell_type":"code","source":"# criterion = nn.CrossEntropyLoss()\nweights = torch.tensor([1.0, 2.0, 4.0])\ncriterion = nn.CrossEntropyLoss(weight=weights.to(device))\n\noptimizer_sagittal_t1 = torch.optim.Adam(sagittal_t1_model.model.fc.parameters(), lr=0.001)\noptimizer_axial_t2 = torch.optim.Adam(axial_t2_model.model.fc.parameters(), lr=0.001)\noptimizer_sagittal_t2stir = torch.optim.Adam(sagittal_t2stir_model.model.fc.parameters(), lr=0.001)\n\nmodel_dict = {\n    'Sagittal T1': sagittal_t1_model,\n    'Axial T2': axial_t2_model,\n    'Sagittal T2/STIR': sagittal_t2stir_model,\n}\n\noptimizers = {\n    'Sagittal T1': optimizer_sagittal_t1,\n    'Axial T2': optimizer_axial_t2,\n    'Sagittal T2/STIR': optimizer_sagittal_t2stir\n}","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:48.054794Z","iopub.execute_input":"2024-10-08T17:22:48.055050Z","iopub.status.idle":"2024-10-08T17:22:48.063796Z","shell.execute_reply.started":"2024-10-08T17:22:48.055027Z","shell.execute_reply":"2024-10-08T17:22:48.062916Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"trainable_params = sum(p.numel() for p in sagittal_t1_model.parameters() if p.requires_grad)\nprint(f\"Number of parameters: {trainable_params}\")","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:48.064969Z","iopub.execute_input":"2024-10-08T17:22:48.065252Z","iopub.status.idle":"2024-10-08T17:22:48.075900Z","shell.execute_reply.started":"2024-10-08T17:22:48.065229Z","shell.execute_reply":"2024-10-08T17:22:48.075069Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_map = {'normal_mild': 0, 'moderate': 1, 'severe': 2}","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:22:48.076868Z","iopub.execute_input":"2024-10-08T17:22:48.077108Z","iopub.status.idle":"2024-10-08T17:22:48.086425Z","shell.execute_reply.started":"2024-10-08T17:22:48.077086Z","shell.execute_reply":"2024-10-08T17:22:48.085497Z"},"trusted":true},"execution_count":null,"outputs":[]},{"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-10-08T17:22:48.087563Z","iopub.execute_input":"2024-10-08T17:22:48.087916Z","iopub.status.idle":"2024-10-08T17:22:48.604754Z","shell.execute_reply.started":"2024-10-08T17:22:48.087886Z","shell.execute_reply":"2024-10-08T17:22:48.603655Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def train_model(model, trainloader, valloader, len_train, len_val, optimizer, model_name,\n                num_epochs=10, \n#                 num_epochs=1, \n                patience=5):\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    train_losses, val_losses = [], []\n    train_accuracies, val_accuracies = [], []\n    train_precisions, val_precisions = [], []\n    train_recalls, val_recalls = [], []\n    train_f1s, val_f1s = [], []\n    \n    for epoch in range(num_epochs):\n        model.train()\n        train_loss = 0\n        correct_train = 0\n        all_train_labels = []\n        all_train_preds = []\n        \n        with tqdm(trainloader, unit=\"batch\") as tepoch:\n            for images, labels in tepoch:\n                images, labels = images.to(device), torch.tensor([label_map[label] for label in labels]).to(device)\n                optimizer.zero_grad()\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                loss.backward()\n                optimizer.step()\n                train_loss += loss.item()\n                \n                probabilities = torch.softmax(outputs, dim=1)\n                _, predicted = torch.max(probabilities, 1)\n                correct_train += (predicted == labels).sum().item()\n                \n                #collect predictions and true labels for metrics\n                all_train_preds.extend(predicted.cpu().detach().numpy())\n                all_train_labels.extend(labels.cpu().detach().numpy())\n                \n                tepoch.set_postfix(epoch=epoch+1)\n        \n        scheduler.step()\n        \n        # compute training metrics\n        train_loss /= len(trainloader)\n#         train_acc = 100 * correct_train / len_train\n        train_acc = correct_train / len_train\n        train_precision = precision_score(all_train_labels, all_train_preds, average=\"weighted\")\n        train_recall = recall_score(all_train_labels, all_train_preds, average=\"weighted\")\n        train_f1 = f1_score(all_train_labels, all_train_preds, average=\"weighted\")\n        \n        train_losses.append(train_loss)\n        train_accuracies.append(train_acc)\n        train_precisions.append(train_precision)\n        train_recalls.append(train_recall)\n        train_f1s.append(train_f1)\n        \n#         print(len_train)\n#         print(correct_train)\n        print(f\"Epoch {epoch+1}, Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.4f}, Precision: {train_precision:.4f}, Recall: {train_recall:.4f}, F1 Score: {train_f1:.4f}\")\n        \n        model.eval()\n        val_loss, correct_val = 0, 0\n        all_val_labels = []\n        all_val_preds = []\n        \n        with torch.no_grad():\n            with tqdm(valloader, unit=\"batch\") as vepoch:\n                for images, labels in vepoch:\n                    images, labels = images.to(device), torch.tensor([label_map[label] for label in labels]).to(device)\n                    outputs = model(images)\n                    loss = criterion(outputs, labels)\n                    val_loss += loss.item()\n                    \n                    probabilities = torch.softmax(outputs, dim=1).squeeze(0)\n                    _, predicted = torch.max(probabilities, 1)\n                    correct_val += (predicted == labels).sum().item()\n                    \n                    \n                    all_val_preds.extend(predicted.cpu().detach().numpy())\n                    all_val_labels.extend(labels.cpu().detach().numpy())\n                    vepoch.set_postfix(epoch=epoch+1)\n        \n        val_loss /= len(valloader)\n#         val_acc = 100 * correct_val / len_val\n        val_acc = correct_val / len_val\n        val_precision = precision_score(all_val_labels, all_val_preds, average=\"weighted\")\n        val_recall = recall_score(all_val_labels, all_val_preds, average=\"weighted\")\n        val_f1 = f1_score(all_val_labels, all_val_preds, average=\"weighted\")\n        \n        val_losses.append(val_loss)\n        val_accuracies.append(val_acc)\n        val_precisions.append(val_precision)\n        val_recalls.append(val_recall)\n        val_f1s.append(val_f1)\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        print(f\"Epoch {epoch+1}, Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.4f}, Val Precision: {val_precision:.4f}, Val Recall: {val_recall:.4f}, Val F1 Score: {val_f1:.4f}\")\n        \n        # Save the best model and check for early stopping\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            best_model_wts = deepcopy(model.state_dict())\n            counter = 0\n#             torch.save(best_model_wts, f'best_model_{epoch+1}.pth')\n\n            checkpoint = save_checkpoint({\n                'epoch': epoch + 1,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'scheduler_state_dict': scheduler.state_dict(),\n                'best_val_acc': best_val_acc   \n            }, filename=f'{model_name}.pth')\n    \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    return model, best_val_acc, train_losses, val_losses, train_accuracies, val_accuracies, train_precisions, val_precisions, train_recalls, val_recalls, train_f1s, val_f1s, checkpoint","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:32:48.283900Z","iopub.execute_input":"2024-10-08T17:32:48.284383Z","iopub.status.idle":"2024-10-08T17:32:48.315381Z","shell.execute_reply.started":"2024-10-08T17:32:48.284340Z","shell.execute_reply":"2024-10-08T17:32:48.314337Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_metrics(train_losses, val_losses, train_accuracies, val_accuracies, train_precisions, val_precisions, train_recalls, val_recalls, train_f1s, val_f1s):\n    \n\n    loc_pos = 'lower right'\n    \n    train_color = '#ED4C5C'\n    valid_color = '#0091E1'\n    \n    \n    \n    epochs = range(1, len(train_losses) + 1)\n\n    # Plot Loss\n    plt.figure(figsize=(12, 4))\n    plt.subplot(1, 2, 1)\n    plt.plot(epochs, train_losses, color=train_color, label='Training Loss')\n    plt.plot(epochs, val_losses, color=valid_color, label='Validation Loss')\n    plt.title('Training and Validation Loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.grid(True)\n    plt.legend(loc=loc_pos)\n\n    # Plot Accuracy\n    plt.subplot(1, 2, 2)\n    plt.plot(epochs, train_accuracies, color=train_color, label='Training Accuracy')\n    plt.plot(epochs, val_accuracies, color=valid_color, label='Validation Accuracy')\n    plt.title('Training and Validation Accuracy')\n    plt.xlabel('Epochs')\n    plt.ylabel('Accuracy (%)')\n    plt.legend()\n    plt.grid(True)\n    plt.show()\n\n    # Plot Precision, Recall, and F1 Score\n    plt.figure(figsize=(20, 4))\n    \n    plt.subplot(1, 3, 1)\n    plt.plot(epochs, train_precisions, color=train_color, label='Training Precision')\n    plt.plot(epochs, val_precisions, color=valid_color, label='Validation Precision')\n    plt.title('Training and Validation Precision')\n    plt.xlabel('Epochs')\n    plt.ylabel('Precision')\n    plt.grid(True)\n    plt.legend()\n\n    plt.subplot(1, 3, 2)\n    plt.plot(epochs, train_recalls, color=train_color, label='Training Recall')\n    plt.plot(epochs, val_recalls, color=valid_color, label='Validation Recall')\n    plt.title('Training and Validation Recall')\n    plt.xlabel('Epochs')\n    plt.ylabel('Recall')\n    plt.legend()\n    plt.grid(True)\n\n    plt.subplot(1, 3, 3)\n    plt.plot(epochs, train_f1s, color=train_color, label='Training F1 Score')\n    plt.plot(epochs, val_f1s, color=valid_color, label='Validation F1 Score')\n    plt.title('Training and Validation F1 Score')\n    plt.xlabel('Epochs')\n    plt.ylabel('F1 Score')\n    plt.legend()\n    plt.grid(True)\n\n    plt.show()\n","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:32:48.871095Z","iopub.execute_input":"2024-10-08T17:32:48.871974Z","iopub.status.idle":"2024-10-08T17:32:48.886736Z","shell.execute_reply.started":"2024-10-08T17:32:48.871941Z","shell.execute_reply":"2024-10-08T17:32:48.884553Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Model","metadata":{}},{"cell_type":"code","source":"# # training all models\n# for desc, model in models.items():\n#     if desc == 'Sagittal T1':\n#         trainloader, valloader, len_train, len_val = trainloader_t1, valloader_t1, len_train_t1, len_val_t1\n        \n#     elif desc == 'Axial T2':\n#         trainloader, valloader, len_train, len_val = trainloader_t2, valloader_t2, len_train_t2, len_val_t2\n        \n#     elif desc == 'Sagittal T2/STIR':\n#         trainloader, valloader, len_train, len_val = trainloader_t2stir, valloader_t2stir, len_train_t2stir, len_val_t2stir\n        \n#     print(f\"Training model for {desc}\")\n#     train_model(model, trainloader, valloader, len_train, len_val, optimizers[desc])","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:32:48.928324Z","iopub.execute_input":"2024-10-08T17:32:48.928611Z","iopub.status.idle":"2024-10-08T17:32:48.932702Z","shell.execute_reply.started":"2024-10-08T17:32:48.928587Z","shell.execute_reply":"2024-10-08T17:32:48.931807Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len_train_t1)\nprint(len_val_t1)\n\nprint(len_train_t2)\nprint(len_val_t2)\n\nprint(len_train_t2stir)\nprint(len_val_t2stir)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:32:48.956522Z","iopub.execute_input":"2024-10-08T17:32:48.956768Z","iopub.status.idle":"2024-10-08T17:32:48.961650Z","shell.execute_reply.started":"2024-10-08T17:32:48.956747Z","shell.execute_reply":"2024-10-08T17:32:48.960792Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for desc, model in model_dict.items():\n    if desc == 'Sagittal T1':\n        trainloader, valloader, len_train, len_val = trainloader_t1, valloader_t1, len_train_t1, len_val_t1\n    \n    elif desc == 'Axial T2':\n        trainloader, valloader, len_train, len_val = trainloader_t2, valloader_t2, len_train_t2, len_val_t2\n        \n    elif desc == 'Sagittal T2/STIR':\n        trainloader, valloader, len_train, len_val = trainloader_t2stir, valloader_t2stir, len_train_t2stir, len_val_t2stir\n        \n    print(f\"Training model for {desc}\")\n\n    model, best_val_acc, train_losses, val_losses, train_accuracies, val_accuracies, train_precisions, val_precisions, train_recalls, val_recalls, train_f1s, val_f1s, checkpoint = train_model(\n        model, trainloader, valloader, len_train, len_val, optimizers[desc], \"-\".join(desc.replace(\"/\", \"_\").split()))   \n    print(f\"Plotting metrics for {desc}\")\n    plot_metrics(train_losses, val_losses, train_accuracies, val_accuracies, train_precisions, val_precisions, train_recalls, val_recalls, train_f1s, val_f1s)","metadata":{"execution":{"iopub.status.busy":"2024-10-08T17:32:49.005082Z","iopub.execute_input":"2024-10-08T17:32:49.005374Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Model Inference","metadata":{}},{"cell_type":"code","source":"train_data['level'].unique()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"expanded_test_desc.head(5)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"levels = ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\n\ndef update_row_id(row, levels):\n    level = levels[row.name % len(levels)]\n    return f\"{row['study_id']}_{row['condition']}_{level}\"\n\nexpanded_test_desc['row_id'] = expanded_test_desc.apply(lambda row: update_row_id(row, levels), axis=1)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"expanded_test_desc.head(2)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"print(len(expanded_test_desc))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_checkpoint(model, optimizer, scheduler, checkpoint_path):\n    checkpoint = torch.load(checkpoint_path)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    optimizer.load_state_dict(checkpoint['optimizer_state_dict'])\n    scheduler.load_state_dict(checkpoint['scheduler_state_dict'])\n    epoch = checkpoint['epoch']\n    best_val_acc = checkpoint['best_val_acc']\n    \n    return model, optimizer, scheduler, epoch, best_val_acc","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"class TestDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n        \n    def __len__(self):\n        return len(self.dataframe)\n    \n    def __getitem__(self, index):\n        image_path = self.dataframe['image_path'][index]\n        image = load_dicom(image_path)\n        if self.transform:\n            image = self.transform(image)\n        return image","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"transform = transforms.Compose([\n    transforms.ToPILImage(),\n    transforms.Resize((224, 224)),\n#     transforms.Grayscale(num_output_channels=3),\n    transforms.ToTensor(),\n])\n\ntest_dataset = TestDataset(expanded_test_desc, transform)\ntestloader = DataLoader(test_dataset, batch_size=1, shuffle=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"for image in testloader:\n    print(image.shape)\n    break","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# def get_model(series_description):\n#     return models.get(series_description, None)\n\n# def predict_test_data(testloader, expanded_test_desc):\n#     predictions = []\n#     normal_mild_probs = []\n#     moderate_probs = []\n#     severe_probs = []\n    \n#     for model in models.values():\n#         model.eval()\n        \n#     with torch.no_grad():\n#         for idx, images in enumerate(tqdm(testloader)):\n#             images = images.to(device)\n#             series_description = expanded_test_desc.iloc[idx]['series_description']\n#             model = get_model(series_description)\n#             if model:\n#                 model.eval()\n#                 outputs = model(images)\n#                 probs = torch.softmax(outputs, dim=1).squeeze(0)\n#                 normal_mild_probs.append(probs[0].item())\n#                 moderate_probs.append(probs[1].item())\n#                 severe_probs.append(probs[2].item())\n#                 predictions.append(probs)\n#             else:\n#                 normal_mild_probs.append(None)\n#                 moderate_probs.append(None)\n#                 severe_probs.append(None)\n#                 predictions.append(None)\n        \n#     return normal_mild_probs, moderate_probs, severe_probs, predictions","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def load_saved_model(model_dict, series_description, model_paths, checkpoint):\n    model = model_dict.get(series_description, None)\n    if model and series_description in model_paths:\n        checkpoint = torch.load(model_paths[series_description])\n        model.load_state_dict(checkpoint['model_state_dict'])\n        model.eval()\n        return model\n    return None","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_test_data(test_loader, expanded_test_desc, model_paths, model_dict, checkpoint):\n    predictions = []\n    normal_mild_probs = []\n    moderate_probs = []\n    severe_probs = []\n    \n    with torch.no_grad():\n        for idx, images in enumerate(tqdm(testloader)):\n            images = images.to(device)\n            series_description = expanded_test_desc.iloc[idx]['series_description']\n            \n            model = load_saved_model(model_dict, series_description, model_paths, checkpoint)\n            \n            if model:\n                outputs = model(images)\n                probs = torch.softmax(outputs, dim=1).squeeze(0)\n                normal_mild_probs.append(probs[0].item())\n                moderate_probs.append(probs[1].item())\n                severe_probs.append(probs[2].item())\n                predictions.append(probs.argmax().item())\n                \n            else:\n                normal_mild_probs.append(None)\n                moderate_probs.append(None)\n                severe_probs.append(None)\n                predictions.append(None)\n                \n    return normal_mild_probs, moderate_probs, severe_probs, predictions","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def predict_test_data_with_gradcam(testloader, expanded_test_desc, model_paths, model_dict, checkpoint):\n    predictions = []\n    normal_mild_probs = []\n    moderate_probs = []\n    severe_probs = []\n    gradcam_images = []\n    \n#     with torch.no_grad():\n    for idx, images in enumerate(tqdm(testloader)):\n        images = images.to(device)\n        series_description = expanded_test_desc.iloc[idx]['series_description']\n\n        model = load_saved_model(model_dict, series_description, model_paths, checkpoint)\n        target_layer = model.model.layer4[0].conv2\n\n        if model:\n            gradcam = GradCAM(model=model, target_layer=target_layer)\n\n            outputs = model(images)\n            probs = torch.softmax(outputs, dim=1).squeeze(0)\n\n            normal_mild_probs.append(probs[0].item())\n            moderate_probs.append(probs[1].item())\n            severe_probs.append(probs[2].item())\n\n            predictions.append(probs.argmax().item())\n\n            cam = gradcam.generate(images, target_class=probs.argmax().item())\n            gradcam_images.append(cam)\n\n        else:\n            normal_mild_probs.append(None)\n            moderate_probs.append(None)\n            severe_probs.append(None)\n            predictions.append(None)\n            gradcam_images.append(None)\n                \n    return normal_mild_probs, moderate_probs, severe_probs, predictions, gradcam_images","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def overlay_heatmap(image, heatmap, alpha=0.35):\n    heatmap = cv2.applyColorMap(np.uint8(255 * heatmap), cv2.COLORMAP_JET)\n    image = convert_to_rgb(image)\n    heatmap = cv2.resize(heatmap, (image.shape[1], image.shape[0]))\n    overlay = cv2.addWeighted(heatmap, alpha, image, 1 - alpha, 0)\n    \n    return overlay","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"model_paths = {\n    'Sagittal T1': '/kaggle/working/Sagittal-T1.pth',\n    'Axial T2': '/kaggle/working/Axial-T2.pth',\n    'Sagittal T2/STIR':'/kaggle/working/Sagittal-T2_STIR.pth'\n}","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# normal_mild_probs, moderate_probs, severe_probs, test_predictions = predict_test_data(testloader, expanded_test_desc, model_paths, model_dict, checkpoint)\nnormal_mild_probs, moderate_probs, severe_probs, test_predictions, gradcam_images = predict_test_data_with_gradcam(testloader, expanded_test_desc, model_paths, model_dict, checkpoint)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"heatmap_image = testloader.dataset[0][0].cpu().numpy()\ncam = gradcam_images[0]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def convert_to_rgb(image):\n    if len(image.shape) == 2:\n        image = cv2.cvtColor(np.uint8(255 * image), cv2.COLOR_GRAY2RGB)\n    elif image.shape[2] == 1:\n        image = cv2.cvtColor(np.uint8(255 * image), cv2.COLOR_GRAY2RGB)\n    return image","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"heatmap_image = convert_to_rgb(heatmap_image)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"overlay_image = overlay_heatmap(heatmap_image, cam, alpha=0.35)\nplt.imshow(overlay_image)\nplt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def gradcam_batch(testloader, gradcam_images, num_images=5):\n\n    # Counter to keep track of how many images have been visualized\n    counter = 0\n    \n    # Iterate over the dataloader\n#     for batch_idx, (images, labels) in enumerate(testloader):\n    for batch_idx, (images) in enumerate(testloader):\n        images = images.to(device)\n        \n        # Iterate over each image in the batch\n        for i in range(images.size(0)):\n            if counter >= num_images:\n                return  # Stop once num_images is reached\n\n            # Get the original image (PyTorch format: C x H x W)\n            image = images[i].cpu().numpy()\n            image = convert_to_rgb(image)\n            image = np.transpose(image, (1, 2, 0))  # Convert to H x W x C\n\n            # Normalize the image if necessary (to 0-255)\n            if image.max() <= 1:\n                image = np.uint8(255 * image)\n\n            # Get the precomputed Grad-CAM heatmap for the current image\n            cam = gradcam_images[counter]\n            \n            # If the Grad-CAM heatmap is None, skip this image\n            if cam is None:\n                counter += 1\n                continue\n            \n            # Overlay the heatmap on the original image\n            overlay_image = overlay_heatmap(image, cam)\n\n            # Plot the original image, heatmap, and the overlay\n            plt.figure(figsize=(10, 5))\n            \n            # Original Image\n            plt.subplot(1, 3, 1)\n            plt.imshow(image, cmap='gray')\n            plt.title(\"Original Image\")\n            plt.axis('off')\n            \n            # Grad-CAM Heatmap\n            plt.subplot(1, 3, 2)\n            plt.imshow(cam, cmap='jet')\n            plt.title(\"Grad-CAM Heatmap\")\n            plt.axis('off')\n            \n            # Overlay Image\n            plt.subplot(1, 3, 3)\n            plt.imshow(overlay_image)\n            plt.title(\"Overlay Heatmap\")\n            plt.axis('off')\n            \n            plt.tight_layout()\n            plt.show()\n            \n            # Increment the counter\n            counter += 1\n            if counter >= num_images:\n                return  # Stop when the specified number of images is reached","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"gradcam_batch(testloader, gradcam_images)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def plot_probability_distributions(normal_mild_probs, moderate_probs, severe_probs):\n    plt.figure(figsize=(12, 5))\n    \n    plt.subplot(1, 3, 1)\n    sns.histplot(normal_mild_probs, bins=30, kde=True)\n    plt.grid(True)\n    plt.title('Normal/Mild probabilities')\n    \n    plt.subplot(1,3,2)\n    sns.histplot(moderate_probs, bins=30, kde=True)\n    plt.grid(True)\n    plt.title('Moderate probabilities')\n    \n    plt.subplot(1, 3, 3)\n    sns.histplot(severe_probs, bins=30, kde=True)\n    plt.grid(True)\n    plt.title('Severe probabiblities')\n    \n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plot_probability_distributions(normal_mild_probs, moderate_probs, severe_probs)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"def calculate_pseudo_auc(normal_mild_probs, moderate_probs, severe_probs, pseudo_labels):\n    y_pred_probs = np.vstack((normal_mild_probs, moderate_probs, severe_probs)).T\n    auc = roc_auc_score(pseudo_labels, y_pred_probs, multi_class='ovr')\n    return auc","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"expanded_test_desc['normal_mild'] = normal_mild_probs\nexpanded_test_desc['moderate'] = moderate_probs\nexpanded_test_desc['severe'] = severe_probs","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission = expanded_test_desc[[\"row_id\", \"normal_mild\", \"moderate\", \"severe\"]]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"submission.head(10)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"grouped_submission = submission.groupby('row_id').max().reset_index()\ngrouped_submission","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub[['normal_mild', 'moderate', 'severe']] = grouped_submission[['normal_mild', 'moderate', 'severe']]","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.to_csv(\"/kaggle/working/submission.csv\", index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"sub.head(25)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}