{"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"}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import os\nimport pydicom\nimport cv2\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport matplotlib.patches as patches\nfrom matplotlib import animation, rc\nimport pandas as pd\nimport glob\nimport json\nimport collections\nimport seaborn as sns\nfrom sklearn.model_selection import train_test_split\nfrom tensorflow.keras import layers, models\nfrom sklearn.metrics import confusion_matrix, classification_report, roc_auc_score, roc_curve, precision_score, recall_score\nimport tensorflow as tf\nfrom tensorflow.keras import layers, models\nfrom tensorflow.keras.callbacks import EarlyStopping, ReduceLROnPlateau\nfrom tensorflow.keras.applications import ResNet50\nfrom tensorflow.keras.layers import Dense, Flatten\nfrom tensorflow.keras.models import Model","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","execution":{"iopub.status.busy":"2024-12-02T18:41:47.802512Z","iopub.execute_input":"2024-12-02T18:41:47.803096Z","iopub.status.idle":"2024-12-02T18:42:01.452566Z","shell.execute_reply.started":"2024-12-02T18:41:47.803067Z","shell.execute_reply":"2024-12-02T18:42:01.451816Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Loading/Preprocessing","metadata":{}},{"cell_type":"code","source":"train_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\ntest_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/test_images/'\ntrain_images = train_path + \"train_images/\"\ntrain  = pd.read_csv(train_path + 'train.csv')\nlabel = pd.read_csv(train_path + 'train_label_coordinates.csv')\ntrain_desc  = pd.read_csv(train_path + 'train_series_descriptions.csv')\ntest_desc   = pd.read_csv(train_path + 'test_series_descriptions.csv')\nsub         = pd.read_csv(train_path + 'sample_submission.csv')\n# study_id = 1082591956 # 100206310\n","metadata":{"execution":{"iopub.status.busy":"2024-12-02T18:42:01.453888Z","iopub.execute_input":"2024-12-02T18:42:01.454380Z","iopub.status.idle":"2024-12-02T18:42:01.596684Z","shell.execute_reply.started":"2024-12-02T18:42:01.454352Z","shell.execute_reply":"2024-12-02T18:42:01.595747Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## New Loading/Compilation (from EfficientNet)","metadata":{}},{"cell_type":"code","source":"# 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\nnew_train_df.head(5)","metadata":{"execution":{"iopub.status.busy":"2024-12-02T18:42:01.597797Z","iopub.execute_input":"2024-12-02T18:42:01.598151Z","iopub.status.idle":"2024-12-02T18:42:02.691370Z","shell.execute_reply.started":"2024-12-02T18:42:01.598104Z","shell.execute_reply":"2024-12-02T18:42:02.690494Z"},"trusted":true},"outputs":[],"execution_count":null},{"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-12-02T18:42:02.693664Z","iopub.execute_input":"2024-12-02T18:42:02.694416Z","iopub.status.idle":"2024-12-02T18:42:02.699507Z","shell.execute_reply.started":"2024-12-02T18:42:02.694382Z","shell.execute_reply":"2024-12-02T18:42:02.698604Z"},"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 columns 'series_id' and 'study_id' #! (had redundancy in original?)\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-12-02T18:42:02.700483Z","iopub.execute_input":"2024-12-02T18:42:02.700717Z","iopub.status.idle":"2024-12-02T18:42:02.765779Z","shell.execute_reply.started":"2024-12-02T18:42:02.700677Z","shell.execute_reply":"2024-12-02T18:42:02.764947Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 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-12-02T18:42:02.766759Z","iopub.execute_input":"2024-12-02T18:42:02.767111Z","iopub.status.idle":"2024-12-02T18:42:02.903767Z","shell.execute_reply.started":"2024-12-02T18:42:02.767072Z","shell.execute_reply":"2024-12-02T18:42:02.902951Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"### TESTING/CONFIRMATION OF UNDERSTANDING\nunique_condition_plane_combos = final_merged_df[[\"condition\", \"series_description\"]].drop_duplicates()\nprint(unique_condition_plane_combos)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:42:02.904860Z","iopub.execute_input":"2024-12-02T18:42:02.905200Z","iopub.status.idle":"2024-12-02T18:42:02.923883Z","shell.execute_reply.started":"2024-12-02T18:42:02.905174Z","shell.execute_reply":"2024-12-02T18:42:02.923186Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"with pd.option_context('display.max_rows', None, 'display.max_columns', None):  # more options can be specified also\n    print(final_merged_df.groupby([\"condition\", \"level\", \"severity\"]).size())","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:42:02.924724Z","iopub.execute_input":"2024-12-02T18:42:02.924938Z","iopub.status.idle":"2024-12-02T18:42:02.945911Z","shell.execute_reply.started":"2024-12-02T18:42:02.924917Z","shell.execute_reply":"2024-12-02T18:42:02.945140Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"grouped_df = final_merged_df.groupby([\"condition\", \"level\", \"severity\"]).size().reset_index(name=\"count\")\n\n# List of severity levels\nseverity_levels = grouped_df[\"severity\"].unique()\n\n# Create one plot per severity level\nfor severity in severity_levels:\n    subset = grouped_df[grouped_df[\"severity\"] == severity]\n    \n    # Pivot table to prepare for plotting\n    pivot_df = subset.pivot(index=\"condition\", columns=\"level\", values=\"count\").fillna(0)\n    pivot_df = pivot_df[grouped_df[\"level\"].unique()]  # Ensure consistent level order\n    \n    # Plot\n    plt.figure(figsize=(8, 6))\n    sns.barplot(data=subset, x=\"condition\", y=\"count\", hue=\"level\")\n    plt.title(f\"Counts for Severity Level: {severity.capitalize()}\")\n    plt.xlabel(\"Condition\")\n    plt.ylabel(\"Count\")\n    plt.xticks(rotation=45)\n    plt.legend(title=\"Level\")\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:42:02.946973Z","iopub.execute_input":"2024-12-02T18:42:02.947490Z","iopub.status.idle":"2024-12-02T18:42:04.267100Z","shell.execute_reply.started":"2024-12-02T18:42:02.947453Z","shell.execute_reply":"2024-12-02T18:42:04.266123Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"grouped_df = final_merged_df.groupby([\"condition\", \"level\", \"severity\"]).size().reset_index(name=\"count\")\n\n# List of severity levels\nlevels = grouped_df[\"level\"].unique()\n\n# Create one plot per severity level\nfor level in levels:\n    subset = grouped_df[grouped_df[\"level\"] == level]\n    \n    # Pivot table to prepare for plotting\n    pivot_df = subset.pivot(index=\"condition\", columns=\"severity\", values=\"count\").fillna(0)\n    pivot_df = pivot_df[grouped_df[\"severity\"].unique()]  # Ensure consistent level order\n    \n    # Plot\n    plt.figure(figsize=(8, 6))\n    sns.barplot(data=subset, x=\"condition\", y=\"count\", hue=\"severity\")\n    plt.title(f\"Counts for Level: {level.capitalize()}\")\n    plt.xlabel(\"Condition\")\n    plt.ylabel(\"Count\")\n    plt.xticks(rotation=45)\n    plt.legend(title=\"Severity\")\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:42:04.269673Z","iopub.execute_input":"2024-12-02T18:42:04.269944Z","iopub.status.idle":"2024-12-02T18:42:06.287083Z","shell.execute_reply.started":"2024-12-02T18:42:04.269919Z","shell.execute_reply":"2024-12-02T18:42:06.286295Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(final_merged_df[final_merged_df[\"severity\"] == \"Normal/Mild\"].value_counts().sum(), \\\nfinal_merged_df[final_merged_df[\"severity\"] == \"Moderate\"].value_counts().sum(), \\\nfinal_merged_df[final_merged_df[\"severity\"] == \"Severe\"].value_counts().sum())","metadata":{"execution":{"iopub.status.busy":"2024-12-02T18:42:06.288257Z","iopub.execute_input":"2024-12-02T18:42:06.288674Z","iopub.status.idle":"2024-12-02T18:42:06.432077Z","shell.execute_reply.started":"2024-12-02T18:42:06.288630Z","shell.execute_reply":"2024-12-02T18:42:06.431187Z"},"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-12-02T18:42:06.433185Z","iopub.execute_input":"2024-12-02T18:42:06.433477Z","iopub.status.idle":"2024-12-02T18:42:06.565486Z","shell.execute_reply.started":"2024-12-02T18:42:06.433450Z","shell.execute_reply":"2024-12-02T18:42:06.564628Z"},"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-12-02T18:42:06.566637Z","iopub.execute_input":"2024-12-02T18:42:06.567299Z","iopub.status.idle":"2024-12-02T18:42:06.575017Z","shell.execute_reply.started":"2024-12-02T18:42:06.567237Z","shell.execute_reply":"2024-12-02T18:42:06.574200Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"full_test_data = expanded_test_desc\nfull_train_data = final_merged_df","metadata":{"execution":{"iopub.status.busy":"2024-12-02T18:42:06.576150Z","iopub.execute_input":"2024-12-02T18:42:06.576489Z","iopub.status.idle":"2024-12-02T18:42:06.585977Z","shell.execute_reply.started":"2024-12-02T18:42:06.576454Z","shell.execute_reply":"2024-12-02T18:42:06.585250Z"},"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\nfull_train_data['study_id_exists'] = full_train_data.apply(check_study_id, axis=1)\nfull_train_data['series_id_exists'] = full_train_data.apply(check_series_id, axis=1)\nfull_train_data['image_exists'] = full_train_data.apply(check_image_exists, axis=1)\n\n# Filter train_data\nfull_train_data = full_train_data[(full_train_data['study_id_exists']) & (full_train_data['series_id_exists']) & (full_train_data['image_exists'])]","metadata":{"execution":{"iopub.status.busy":"2024-12-02T18:42:06.586896Z","iopub.execute_input":"2024-12-02T18:42:06.587130Z","iopub.status.idle":"2024-12-02T18:43:14.768377Z","shell.execute_reply.started":"2024-12-02T18:42:06.587108Z","shell.execute_reply":"2024-12-02T18:43:14.767605Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"full_train_data.head(10)","metadata":{"execution":{"iopub.status.busy":"2024-12-02T18:43:14.769345Z","iopub.execute_input":"2024-12-02T18:43:14.769593Z","iopub.status.idle":"2024-12-02T18:43:14.783644Z","shell.execute_reply.started":"2024-12-02T18:43:14.769570Z","shell.execute_reply":"2024-12-02T18:43:14.782783Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"full_train_data.head(-10)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:43:14.785027Z","iopub.execute_input":"2024-12-02T18:43:14.785312Z","iopub.status.idle":"2024-12-02T18:43:14.803471Z","shell.execute_reply.started":"2024-12-02T18:43:14.785259Z","shell.execute_reply":"2024-12-02T18:43:14.802654Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#for one hot encoding\n#train_data[['normal_mild', 'severe', 'moderate']] = train_data[['normal_mild', 'severe', 'moderate']].astype(int)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:43:14.804532Z","iopub.execute_input":"2024-12-02T18:43:14.804807Z","iopub.status.idle":"2024-12-02T18:43:14.810876Z","shell.execute_reply.started":"2024-12-02T18:43:14.804784Z","shell.execute_reply":"2024-12-02T18:43:14.810165Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(full_train_data.shape)\nfull_train_data = full_train_data.dropna()\nprint(full_train_data.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:43:14.811710Z","iopub.execute_input":"2024-12-02T18:43:14.811988Z","iopub.status.idle":"2024-12-02T18:43:14.840172Z","shell.execute_reply.started":"2024-12-02T18:43:14.811954Z","shell.execute_reply":"2024-12-02T18:43:14.839349Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#! Can change this as seen fit\nrelevant_train_cols = [\"study_id\", \"condition\", \"severity\", \"level\", \"image_path\", \"x\", \"y\", \"series_description\"]\nrelevant_test_cols = [\"study_id\", \"condition\", \"image_path\", \"series_description\"] # NOT CORRECT?\n# [\"study_id\", \"condition\", \"level\", \"severity\", etc.]\ntrain_data = full_train_data[relevant_train_cols]\ntest_data = full_test_data[relevant_test_cols]\ntrain_data.head(10)","metadata":{"execution":{"iopub.status.busy":"2024-12-02T18:43:14.841153Z","iopub.execute_input":"2024-12-02T18:43:14.841486Z","iopub.status.idle":"2024-12-02T18:43:14.859571Z","shell.execute_reply.started":"2024-12-02T18:43:14.841451Z","shell.execute_reply":"2024-12-02T18:43:14.858801Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_dicom_image(file_path, x_marker, y_marker, plane, condition, img_size=(128, 128)):\n    # Read the DICOM file\n    dicom = pydicom.dcmread(file_path)\n    # dicom = pydicom.read_file(file_path)\n    # Convert the DICOM file to a NumPy array\n    img = dicom.pixel_array\n\n    # scale x_marker and y_marker\n    x_marker /= img.shape[0]\n    y_marker /= img.shape[1]\n    \n    # print(img.shape) # (512, 512)\n    # Cropping region of interest/level based on plane and condition:\n    # if \"Sagittal\" in plane and \"Narrowing\" in condition:\n    #     start_x = max(x_marker - 55, 0) # manually determined these crop constants after\n    #     start_y = max(y_marker - 35, 0) # repeated sampling; need bounding box isolation\n    #     end_x = min(x_marker + 56, img.shape[1]) # can be tweaked further\n    #     end_y = min(y_marker + 36, img.shape[0])\n    #     img = img[start_y:end_y, start_x:end_x]\n    # Normalize the image (to 0-1 range)\n    img = img / np.max(img)\n    # Resize the image\n    img_resized = cv2.resize(img, img_size)\n    # Convert to 3-channel image (optional, as some CNN architectures require 3 channels)\n    img_resized = np.stack((img_resized,)*3, axis=-1)\n    return img_resized, (x_marker, y_marker)\n\ndef load_data_from_folders(data, img_size=(128, 128)):\n    # spinal_canal_stenosis: 0 , left_neural_foraminal_narrowing: 1, \n    # right_neural_foraminal_narrowing: 2, left_subarticular_stenosis: 3, \n    # right_subarticular_stenosis:4\n    X = []  # Image data\n    X_metadata = [] # Meta data (x_marker, y_marker, and level)\n    y = []  # Labels\n    class_map = {\"Spinal Canal Stenosis\": 0 , \"Left Neural Foraminal Narrowing\": 1, \n    \"Right Neural Foraminal Narrowing\": 2, \"Left Subarticular Stenosis\": 3, \n    \"Right Subarticular Stenosis\": 4}\n    severity_map = {\"normal_mild\": 0, \"moderate\": 1, \"severe\": 2}\n    level_map = {'L1/L2': 0, 'L2/L3': 1, 'L3/L4': 2, 'L4/L5': 3, 'L5/S1': 4}\n    crop_square_half_width = 10\n    \n    for index, row in data.iterrows():\n        file_path = row['image_path']\n        class_label = class_map[row['condition']]\n        # print(row['severity'])\n        severity_label = severity_map[row['severity']]\n#         print(index, file_path)\n       \n        if file_path.endswith('.dcm'):\n            try:\n                # Load and preprocess the image\n                x_marker = int(row['x'])\n                y_marker = int(row['y'])\n                level = level_map[row['level']]\n                plane = row['series_description']\n                cond = row['condition']\n                img, (x_marker_scaled, y_marker_scaled) = load_dicom_image(file_path, x_marker, y_marker, plane, cond, img_size)\n                # print(img.shape, x_center_marker, y_center_marker)\n                X.append(img)\n                X_metadata.append((x_marker_scaled, y_marker_scaled, level))\n                y.append([class_label, severity_label])\n            except Exception as e:\n                print(f\"Error loading {file_path}: {e}\")\n        else:\n            break\n    \n    X = np.array(X)\n    X_metadata = np.array(X_metadata)\n    y = np.array(y)\n    return X, X_metadata, y, class_map, severity_map","metadata":{"execution":{"iopub.status.busy":"2024-12-02T18:43:14.860534Z","iopub.execute_input":"2024-12-02T18:43:14.860770Z","iopub.status.idle":"2024-12-02T18:43:14.870106Z","shell.execute_reply.started":"2024-12-02T18:43:14.860747Z","shell.execute_reply":"2024-12-02T18:43:14.869416Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Train Test Split","metadata":{}},{"cell_type":"code","source":"def visualize_images(X, y, class_map, severity_map, num_images=5):\n    cond_labels = {v: k for k, v in class_map.items()} \n    sev_labels = {v: k for k, v in severity_map.items()}\n    plt.figure(figsize=(10, 10))\n    start_r = -min(num_images, len(X)) # 0\n    end_r = 0 # min(num_images, len(X))\n    for i in range(-min(num_images, len(X)), 0):\n        plt.subplot(num_images, 1, i - start_r + 1)\n        plt.imshow(X[i], cmap='gray')\n        plt.title(cond_labels[y[i][0]] + \", \" + sev_labels[y[i][1]])\n        plt.axis('off')\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:43:14.871305Z","iopub.execute_input":"2024-12-02T18:43:14.872003Z","iopub.status.idle":"2024-12-02T18:43:14.883171Z","shell.execute_reply.started":"2024-12-02T18:43:14.871957Z","shell.execute_reply":"2024-12-02T18:43:14.882530Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# NOTE: COULD BE MULTIPLE CONDITIONS FOR EACH SPINE IMAGE;\n# CURRENTLY MADE IT SO THAT EACH LEVEL IS CROPPED (WHEN APPLICABLE) AND THE BASELINE MODEL\n# OUTPUTS THE CONDITION AND SEVERITY FOR THAT LEVEL\ntot_data = 15000\nX, X_metadata, y, class_names, severity_names = load_data_from_folders(train_data.sample(tot_data, random_state=42))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:43:14.884132Z","iopub.execute_input":"2024-12-02T18:43:14.884457Z","iopub.status.idle":"2024-12-02T18:46:46.061164Z","shell.execute_reply.started":"2024-12-02T18:43:14.884432Z","shell.execute_reply":"2024-12-02T18:46:46.060137Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_images(X, y, class_names, severity_names)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:46:46.062421Z","iopub.execute_input":"2024-12-02T18:46:46.062780Z","iopub.status.idle":"2024-12-02T18:46:46.430733Z","shell.execute_reply.started":"2024-12-02T18:46:46.062742Z","shell.execute_reply":"2024-12-02T18:46:46.429684Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"X.shape, y.shape","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:46:46.431882Z","iopub.execute_input":"2024-12-02T18:46:46.432219Z","iopub.status.idle":"2024-12-02T18:46:46.438446Z","shell.execute_reply.started":"2024-12-02T18:46:46.432184Z","shell.execute_reply":"2024-12-02T18:46:46.437331Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# print(y)\nprint(X.shape)\nprint(train_data.shape) # (48657, 7)\ntrain_size = 0.7  #  for training\ntest_val_size = 0.3  #  for testing and validation\nval_size = 0.5\ntest_size = 0.5\n\nX_train, X_test_val, X_metadata_train, X_metadata_test_val, y_train, y_test_val = train_test_split(X, X_metadata, y, test_size=test_val_size, train_size=train_size, random_state=42)\nX_val, X_test, X_metadata_val, X_metadata_test, y_val, y_test = train_test_split(X_test_val, X_metadata_test_val, y_test_val, test_size=test_size, train_size=val_size, random_state=42)","metadata":{"execution":{"iopub.status.busy":"2024-12-02T18:46:46.439660Z","iopub.execute_input":"2024-12-02T18:46:46.439964Z","iopub.status.idle":"2024-12-02T18:46:48.674413Z","shell.execute_reply.started":"2024-12-02T18:46:46.439931Z","shell.execute_reply":"2024-12-02T18:46:48.673449Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"print(type(X_train), X_train.shape, X_metadata_train.shape)\nprint(len(X_train), len(X_val), len(X_test)) # 70-15-15 train-val-test split","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:46:48.675491Z","iopub.execute_input":"2024-12-02T18:46:48.675732Z","iopub.status.idle":"2024-12-02T18:46:48.681161Z","shell.execute_reply.started":"2024-12-02T18:46:48.675709Z","shell.execute_reply":"2024-12-02T18:46:48.680252Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Trial Data Augmentation","metadata":{}},{"cell_type":"markdown","source":"Since we're overfitting on severity on the testing dataset, let's create some more data/images by data augmentaiton. ","metadata":{}},{"cell_type":"code","source":"import albumentations as A\nimport cv2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:46:48.685318Z","iopub.execute_input":"2024-12-02T18:46:48.685654Z","iopub.status.idle":"2024-12-02T18:46:49.370963Z","shell.execute_reply.started":"2024-12-02T18:46:48.685629Z","shell.execute_reply":"2024-12-02T18:46:49.370081Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def normalize_image(image):\n    \"\"\"Normalize image to correct format for albumentations\"\"\"\n    image = np.array(image)\n    \n    if image.dtype != np.uint8:\n        if image.max() <= 1.0:\n            image = (image * 255).astype(np.uint8)\n        else:\n            image = image.astype(np.uint8)\n    if len(image.shape) == 2:\n        image = np.expand_dims(image, axis=-1)\n    \n    return image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:46:49.372260Z","iopub.execute_input":"2024-12-02T18:46:49.373052Z","iopub.status.idle":"2024-12-02T18:46:49.379061Z","shell.execute_reply.started":"2024-12-02T18:46:49.373008Z","shell.execute_reply":"2024-12-02T18:46:49.377979Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_augmentations(X, y, class_map, severity_map, num_images=5, num_augmentations=5):\n    transform = A.Compose([\n        A.RandomBrightnessContrast(brightness_limit=(-0.1, 0.2), contrast_limit=0.2, p=0.6),\n        A.GaussNoise(var_limit=(10.0, 50.0), p=0.6),\n        A.Rotate(limit=10, p=0.6),\n        A.GridDistortion(num_steps=5, distort_limit=0.2, p=0.5),\n        ##A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=0.5)\n    ])\n    cond_labels = {v: k for k, v in class_map.items()}\n    sev_labels = {v: k for k, v in severity_map.items()}\n    fig = plt.figure(figsize=(15, 4*num_images))\n\n    for idx in range(min(num_images, len(X))):\n        original_img = normalize_image(X[idx])\n        \n        plt.subplot(num_images, num_augmentations + 1, idx*(num_augmentations + 1) + 1)\n        plt.imshow(original_img, cmap='gray')\n        plt.title(f\"Original\\n{cond_labels[y[idx][0]]}\\n{sev_labels[y[idx][1]]}\", fontsize=8)\n        plt.axis('off')\n        \n        for aug_idx in range(num_augmentations):\n            augmented = transform(image=original_img)['image']\n            plt.subplot(num_images, num_augmentations + 1, idx*(num_augmentations + 1) + aug_idx + 2)\n            plt.imshow(augmented, cmap='gray')\n            plt.title(f\"Augmented {aug_idx + 1}\", fontsize=8)\n            plt.axis('off')\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:46:49.380226Z","iopub.execute_input":"2024-12-02T18:46:49.380587Z","iopub.status.idle":"2024-12-02T18:46:49.390403Z","shell.execute_reply.started":"2024-12-02T18:46:49.380560Z","shell.execute_reply":"2024-12-02T18:46:49.389600Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"visualize_augmentations(X, y, class_names, severity_names, num_images=5, num_augmentations=3)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:46:49.391221Z","iopub.execute_input":"2024-12-02T18:46:49.391502Z","iopub.status.idle":"2024-12-02T18:46:52.276554Z","shell.execute_reply.started":"2024-12-02T18:46:49.391479Z","shell.execute_reply":"2024-12-02T18:46:52.275423Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:46:52.278302Z","iopub.execute_input":"2024-12-02T18:46:52.278656Z","iopub.status.idle":"2024-12-02T18:46:52.282852Z","shell.execute_reply.started":"2024-12-02T18:46:52.278620Z","shell.execute_reply":"2024-12-02T18:46:52.281976Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"## now that we've tested, going to formally add these augmented images\ndef augment_dataset(X, X_metadata, y, num_augmentations=2):\n    transform = A.Compose([\n        A.RandomBrightnessContrast(brightness_limit=(-0.1, 0.2), contrast_limit=0.2, p=0.6),\n        A.GaussNoise(var_limit=(10.0, 50.0), p=0.6),\n        A.Rotate(limit=10, p=0.6),\n        A.GridDistortion(num_steps=5, distort_limit=0.2, p=0.5),\n        ##A.ElasticTransform(alpha=1, sigma=50, alpha_affine=50, p=0.5)\n    ])\n    augmented_X = []\n    augmented_X_metadata = []\n    augmented_y = []\n\n    for i in range(len(X)):\n        augmented_X.append(X[i])\n        augmented_X_metadata.append(X_metadata[i])\n        augmented_y.append(y[i])\n\n    print(\"Creating augmented images...\")\n    for i in tqdm(range(len(X))):\n        img = normalize_image(X[i])\n        for _ in range(num_augmentations):\n            aug_img = transform(image=img)['image']\n            if X[i].dtype != np.uint8:\n                aug_img = aug_img.astype(X[i].dtype)\n                if X[i].max() <= 1.0:\n                    aug_img = aug_img / 255.0\n            if len(X[i].shape) == 2:\n                aug_img = aug_img.squeze()\n            augmented_X.append(aug_img)\n            augmented_X_metadata.append(X_metadata[i])\n            augmented_y.append(y[i])\n    augmented_X = np.array(augmented_X)\n    augmented_X_metadata = np.array(augmented_X_metadata)\n    augmented_y = np.array(augmented_y)\n\n    return augmented_X, augmented_X_metadata, augmented_y","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:46:52.283645Z","iopub.execute_input":"2024-12-02T18:46:52.283961Z","iopub.status.idle":"2024-12-02T18:46:52.295962Z","shell.execute_reply.started":"2024-12-02T18:46:52.283928Z","shell.execute_reply":"2024-12-02T18:46:52.295064Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"##X_aug, X_metadata_aug, y_aug = augment_dataset(X_train, X_metadata_train, y_train, num_augmentations=2)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:46:52.297389Z","iopub.execute_input":"2024-12-02T18:46:52.297772Z","iopub.status.idle":"2024-12-02T18:46:52.310398Z","shell.execute_reply.started":"2024-12-02T18:46:52.297727Z","shell.execute_reply":"2024-12-02T18:46:52.309430Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"#print(\"Original data size: \", len(X_train))\n#print(\"Augmented data size: \", len(X_aug))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T18:46:52.311359Z","iopub.execute_input":"2024-12-02T18:46:52.311640Z","iopub.status.idle":"2024-12-02T18:46:52.320110Z","shell.execute_reply.started":"2024-12-02T18:46:52.311615Z","shell.execute_reply":"2024-12-02T18:46:52.319355Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"transforms = A.Compose([\n    A.HorizontalFlip(p=0.5),\n    A.RandomBrightnessContrast(p=0.2),\n    A.ShiftScaleRotate(shift_limit=0.05, scale_limit=0.05, rotate_limit=15, p=0.5),\n    A.PixelDropout(dropout_prob=0.1, per_channel=True, p=0.5),\n    A.CoarseDropout(num_holes_range=(3, 6),hole_height_range=(5, 10),\n                    hole_width_range=(5, 10),\n                    fill_value=0,\n                    p=0.25),\n    # A.GaussianBlur(p=0.2)\n])","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T19:08:42.182126Z","iopub.execute_input":"2024-12-02T19:08:42.183065Z","iopub.status.idle":"2024-12-02T19:08:42.189334Z","shell.execute_reply.started":"2024-12-02T19:08:42.183024Z","shell.execute_reply":"2024-12-02T19:08:42.188318Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"y_train_cond = y_train[:, 0]  # Extract condition labels\ny_train_sev = y_train[:, 1]   # Extract severity labels\n\n# Repeat for validation data\ny_val_cond = y_val[:, 0]\ny_val_sev = y_val[:, 1]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T19:08:44.404084Z","iopub.execute_input":"2024-12-02T19:08:44.405008Z","iopub.status.idle":"2024-12-02T19:08:44.409150Z","shell.execute_reply.started":"2024-12-02T19:08:44.404968Z","shell.execute_reply":"2024-12-02T19:08:44.408336Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AugmentedDataGenerator(tf.keras.utils.Sequence):\n    def __init__(self, X, X_metadata, y, batch_size=32, augment=True):\n        \"\"\"\n        Initialize the data generator\n        Args:\n            X: Image data\n            X_metadata: Metadata features\n            y: Labels (should be a numpy array with shape (n_samples, 2))\n            batch_size: Size of each batch\n            augment: Whether to apply augmentation\n        \"\"\"\n        super(AugmentedDataGenerator, self).__init__()\n        self.X = X\n        self.X_metadata = X_metadata\n        self.y = y\n        self.batch_size = batch_size\n        self.augment = augment\n        \n        # Define augmentation pipeline\n        self.transform = transforms\n        \n        self.indexes = np.arange(len(self.X))\n        self.on_epoch_end()\n    \n    def __len__(self):\n        return int(np.floor(len(self.X) / self.batch_size))\n    \n    def __getitem__(self, index):\n        indexes = self.indexes[index * self.batch_size:(index + 1) * self.batch_size]\n        \n        # Preallocate arrays\n        X_batch = np.empty((len(indexes), *self.X[0].shape), dtype=self.X.dtype)\n        X_metadata_batch = np.empty((len(indexes), *self.X_metadata[0].shape), dtype=self.X_metadata.dtype)\n        y_batch_condition = np.empty(len(indexes), dtype=self.y[:, 0].dtype)\n        y_batch_severity = np.empty(len(indexes), dtype=self.y[:, 1].dtype)\n        \n        for i, idx in enumerate(indexes):\n            image = self.X[idx]\n            \n            if self.augment and self.transform:\n                # Normalize and augment\n                image = (image * 255).astype(np.uint8) if image.max() <= 1 else image.astype(np.uint8)\n                augmented = self.transform(image=image)['image']\n                augmented = augmented.astype(self.X.dtype) / 255.0 if augmented.max() >= 1 else augmented\n                X_batch[i] = augmented\n            else:\n                X_batch[i] = image\n            \n            X_metadata_batch[i] = self.X_metadata[idx]\n            y_batch_condition[i] = self.y[idx, 0]\n            y_batch_severity[i] = self.y[idx, 1]\n        \n        # Ensure the output is a tuple of inputs and outputs\n        return (X_batch, X_metadata_batch), {'condition_output': y_batch_condition, 'severity_output': y_batch_severity}\n\n    \n    def on_epoch_end(self):\n        np.random.shuffle(self.indexes)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T19:08:45.038189Z","iopub.execute_input":"2024-12-02T19:08:45.038556Z","iopub.status.idle":"2024-12-02T19:08:45.049082Z","shell.execute_reply.started":"2024-12-02T19:08:45.038527Z","shell.execute_reply":"2024-12-02T19:08:45.048143Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Declaration, Training, and Testing","metadata":{}},{"cell_type":"code","source":"from sklearn.utils.class_weight import compute_class_weight\nimport numpy as np\n\nclass_weights_cond = compute_class_weight(\n    class_weight='balanced',\n    classes=np.unique(y_train_cond),\n    y=y_train_cond\n)\nclass_weights_cond = dict(enumerate(class_weights_cond))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T19:08:46.688993Z","iopub.execute_input":"2024-12-02T19:08:46.689670Z","iopub.status.idle":"2024-12-02T19:08:46.697381Z","shell.execute_reply.started":"2024-12-02T19:08:46.689632Z","shell.execute_reply":"2024-12-02T19:08:46.696685Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tensorflow.keras import layers, models, Input\ndef create_custom_cnn_model_with_severity(input_shape, num_classes_cond, num_classes_sev):\n    input_layer = Input(shape=input_shape)\n    \n    x = layers.Conv2D(32, (5, 5), activation='relu', padding='same')(input_layer)\n    x = layers.MaxPooling2D((2, 2))(x)\n    x = layers.Conv2D(64, (3, 3), activation='relu')(x)\n    x = layers.MaxPooling2D((2, 2))(x)\n    x = layers.Conv2D(128, (3, 3), activation='relu')(x)\n    x = layers.MaxPooling2D((2, 2))(x)\n    # Additional layers (optional/testing)\n    # x = layers.Conv2D(256, (3, 3), activation='relu')(x) # added\n    # x = layers.MaxPooling2D((2, 2))(x) # added\n    # x = layers.Conv2D(512, (3, 3), activation='relu')(x) # added after\n    # x = layers.MaxPooling2D((2, 2))(x) # added after\n    x = layers.Flatten()(x)\n    x = layers.Dense(128, activation='relu')(x)\n    x = layers.Dropout(0.5)(x)\n    x = layers.Dense(64, activation='relu')(x) # added\n    x = layers.Dropout(0.5)(x)\n    condition_output = layers.Dense(num_classes_cond, activation='softmax', name='condition_output')(x)\n    severity_output = layers.Dense(num_classes_sev, activation='softmax', name='severity_output')(x)\n    # Two outputs\n    model = models.Model(inputs=input_layer, outputs=[condition_output, severity_output])\n    optimizer = tf.keras.optimizers.Adam(learning_rate=lr_schedule)\n    lr_schedule = tf.keras.optimizers.schedules.ExponentialDecay(\n        initial_learning_rate=0.01,  # Starting learning rate\n        decay_steps=10000,          # Number of steps before decay\n        decay_rate=0.9              # Decay rate (multiplier for each decay step)\n    )\n    # Compiling with separate losses\n    model.compile(optimizer=optimizer,\n                  loss={'condition_output': 'sparse_categorical_crossentropy', 'severity_output': 'sparse_categorical_crossentropy'},\n                  loss_weights={'condition_output': 2.0, 'severity_output': 1.0},\n                  metrics={'condition_output': 'accuracy', 'severity_output': 'accuracy'})\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T19:08:47.088920Z","iopub.execute_input":"2024-12-02T19:08:47.089934Z","iopub.status.idle":"2024-12-02T19:08:47.097860Z","shell.execute_reply.started":"2024-12-02T19:08:47.089898Z","shell.execute_reply":"2024-12-02T19:08:47.097006Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from tensorflow.keras.applications import VGG16\nfrom tensorflow.keras import layers, models, Input\n\ndef create_vgg_cnn_model_with_separate_branches(input_shape, input_metadata_shape, num_classes_cond, num_classes_sev, fine_tune=True, num_layers_fine_tune=5):\n    # Input layers\n    input_layer = Input(shape=input_shape)\n    input_metadata_layer = Input(shape=input_metadata_shape)\n\n    # Base model for feature extraction\n    base_model = VGG16(weights='imagenet', include_top=False, input_shape=input_shape)\n    if fine_tune:\n        for layer in base_model.layers[:-num_layers_fine_tune]:  # Freeze all but last layers\n            layer.trainable = False\n    else:\n        base_model.trainable = False\n\n    # Shared feature extraction\n    shared_features = base_model(input_layer)\n    shared_features = layers.Flatten()(shared_features)\n    shared_features = layers.Concatenate()([shared_features, input_metadata_layer])\n\n    # Branch for condition-specific features\n    x_cond = layers.Dense(512, activation='relu')(shared_features)\n    x_cond = layers.Dropout(0.3)(x_cond)\n    x_cond = layers.Dense(256, activation='relu')(x_cond)\n    x_cond = layers.Dropout(0.3)(x_cond)\n    condition_output = layers.Dense(num_classes_cond, activation='softmax', name='condition_output')(x_cond)\n\n    # Branch for severity-specific features\n    x_sev = layers.Dense(512, activation='relu')(shared_features)\n    x_sev = layers.Dropout(0.3)(x_sev)\n    x_sev = layers.Dense(256, activation='relu')(x_sev)\n    x_sev = layers.Dropout(0.3)(x_sev)\n    severity_output = layers.Dense(num_classes_sev, activation='softmax', name='severity_output')(x_sev)\n\n    # Create the model\n    model = models.Model(inputs=[input_layer, input_metadata_layer], outputs=[condition_output, severity_output])\n\n    # Compile the model with appropriate loss weights\n    model.compile(\n        optimizer='adam',\n        loss={\n            'condition_output': 'sparse_categorical_crossentropy',\n            'severity_output': 'sparse_categorical_crossentropy',\n        },\n        loss_weights={\n            'condition_output': 1.0,  # Weight condition more heavily\n            'severity_output': 2.0,\n        },\n        metrics={\n            'condition_output': 'accuracy',\n            'severity_output': 'accuracy',\n        }\n    )\n    return model\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T19:14:42.087199Z","iopub.execute_input":"2024-12-02T19:14:42.088110Z","iopub.status.idle":"2024-12-02T19:14:42.097388Z","shell.execute_reply.started":"2024-12-02T19:14:42.088078Z","shell.execute_reply":"2024-12-02T19:14:42.096546Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Define shapes\ninput_shape = (128, 128, 3)  # Example: Images are 128x128 with 3 color channels\ninput_metadata_shape = (3,)  # Example: 3 metadata features\n\n# Define number of classes\nnum_cond_classes = len(class_names)  # Total condition classes\nnum_sev_classes = len(severity_names)  # Total severity classes\n\n# Create model\nmodel = create_vgg_cnn_model_with_separate_branches(\n    input_shape=input_shape,\n    input_metadata_shape=input_metadata_shape,\n    num_classes_cond=num_cond_classes,\n    num_classes_sev=num_sev_classes,\n    fine_tune=True,\n    num_layers_fine_tune=5\n)\n\n# Check model summary\nmodel.summary()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T19:14:42.885238Z","iopub.execute_input":"2024-12-02T19:14:42.886041Z","iopub.status.idle":"2024-12-02T19:14:43.195194Z","shell.execute_reply.started":"2024-12-02T19:14:42.886006Z","shell.execute_reply":"2024-12-02T19:14:43.194478Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from keras.utils import plot_model\n\nplot_model(model, to_file='model.png', show_shapes=True)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T19:14:43.691313Z","iopub.execute_input":"2024-12-02T19:14:43.691629Z","iopub.status.idle":"2024-12-02T19:14:44.088633Z","shell.execute_reply.started":"2024-12-02T19:14:43.691604Z","shell.execute_reply":"2024-12-02T19:14:44.087678Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"tf.config.list_physical_devices('GPU') # Verify that GPU will be used","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T19:14:44.191879Z","iopub.execute_input":"2024-12-02T19:14:44.192243Z","iopub.status.idle":"2024-12-02T19:14:44.198173Z","shell.execute_reply.started":"2024-12-02T19:14:44.192211Z","shell.execute_reply":"2024-12-02T19:14:44.197218Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_generator = AugmentedDataGenerator(\n    X_train,\n    X_metadata_train,\n    y_train,\n    batch_size=16,\n    augment=True\n)\n\nval_generator = AugmentedDataGenerator(\n    X_val,\n    X_metadata_val,\n    y_val,\n    batch_size=16,\n    augment=False,  # No augmentation for validation\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T19:14:44.755594Z","iopub.execute_input":"2024-12-02T19:14:44.755948Z","iopub.status.idle":"2024-12-02T19:14:44.761005Z","shell.execute_reply.started":"2024-12-02T19:14:44.755918Z","shell.execute_reply":"2024-12-02T19:14:44.760009Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"reduce_lr = ReduceLROnPlateau(monitor='val_loss', factor=0.5, patience=3, min_lr=1e-6)\nearly_stopping = EarlyStopping(monitor='val_loss', patience=5)\n\n# y_train_dict = {'condition_output': y_train[:, 0], 'severity_output': y_train[:, 1]}\n# y_val_dict = {'condition_output': y_val[:, 0], 'severity_output': y_val[:, 1]}\n# #history = model.fit([X_train, X_metadata_train], y_train_dict, epochs=40, batch_size=32, validation_data=([X_val, X_metadata_val], y_val_dict), callbacks=[reduce_lr, early_stopping])\n# history = model.fit(\n#     [X_train, X_metadata_train],  # Inputs: images and metadata\n#     {'condition_output': y_train_cond, 'severity_output': y_train_sev},  # Outputs\n#     validation_data=(\n#         [X_val, X_metadata_val], \n#         {'condition_output': y_val_cond, 'severity_output': y_val_sev}\n#     ),\n#     epochs=40,\n#     batch_size=32,\n#     callbacks=[early_stopping, reduce_lr]\n# )\n\nhistory = model.fit(\n    train_generator,\n    validation_data=val_generator,\n    epochs=40,\n    callbacks=[early_stopping, reduce_lr]\n)\n# history = model.fit(train_generator, steps_per_epoch=len(X_train) // batch_size, epochs=20, validation_data=(X_val, y_val), callbacks=[early_stopping]) # required for generator","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T19:14:45.338739Z","iopub.execute_input":"2024-12-02T19:14:45.339460Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_training_history(history):\n    # Extract accuracy for each output\n    cond_acc = history.history['condition_output_accuracy']\n    val_cond_acc = history.history['val_condition_output_accuracy']\n    sev_acc = history.history['severity_output_accuracy']\n    val_sev_acc = history.history['val_severity_output_accuracy']\n    \n    # Extract total loss\n    total_loss = history.history['loss']\n    val_total_loss = history.history['val_loss']\n    \n    epochs = range(1, len(cond_acc) + 1)\n\n    # Plot accuracy for each output\n    plt.figure(figsize=(15, 6))\n    \n    plt.subplot(1, 2, 1)\n    plt.plot(epochs, cond_acc, 'b', label='Condition Training Accuracy')\n    plt.plot(epochs, val_cond_acc, 'r', label='Condition Validation Accuracy')\n    plt.title('Condition Output Accuracy')\n    plt.xlabel('Epochs')\n    plt.ylabel('Accuracy')\n    plt.legend()\n\n    plt.subplot(1, 2, 2)\n    plt.plot(epochs, sev_acc, 'b', label='Severity Training Accuracy')\n    plt.plot(epochs, val_sev_acc, 'r', label='Severity Validation Accuracy')\n    plt.title('Severity Output Accuracy')\n    plt.xlabel('Epochs')\n    plt.ylabel('Accuracy')\n    plt.legend()\n    \n    # Plot total loss\n    plt.figure(figsize=(7, 5))\n    plt.plot(epochs, total_loss, 'b', label='Total Training Loss')\n    plt.plot(epochs, val_total_loss, 'r', label='Total Validation Loss')\n    plt.title('Total Loss')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n    \n    plt.show()","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"plot_training_history(history)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T16:40:56.715964Z","iopub.execute_input":"2024-12-02T16:40:56.716253Z","iopub.status.idle":"2024-12-02T16:40:57.367746Z","shell.execute_reply.started":"2024-12-02T16:40:56.716228Z","shell.execute_reply":"2024-12-02T16:40:57.366882Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Testing Set Predictions and Model Statistics","metadata":{}},{"cell_type":"code","source":"condition_preds, severity_preds = model.predict([X_test, X_metadata_test])\ncondition_label_preds = np.argmax(condition_preds, axis=1)\nseverity_label_preds = np.argmax(severity_preds, axis=1)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T16:40:57.370388Z","iopub.execute_input":"2024-12-02T16:40:57.370651Z","iopub.status.idle":"2024-12-02T16:41:00.338822Z","shell.execute_reply.started":"2024-12-02T16:40:57.370625Z","shell.execute_reply":"2024-12-02T16:41:00.338139Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Sanity checks\nprint(np.unique(condition_label_preds), np.unique(severity_label_preds))\nprint(condition_label_preds.shape)\nprint(severity_label_preds.shape)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T16:41:00.339830Z","iopub.execute_input":"2024-12-02T16:41:00.340146Z","iopub.status.idle":"2024-12-02T16:41:00.345634Z","shell.execute_reply.started":"2024-12-02T16:41:00.340118Z","shell.execute_reply":"2024-12-02T16:41:00.344682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import precision_score, recall_score, f1_score, accuracy_score\n\ndef calculate_precision(y_true, y_pred):\n    return precision_score(y_true, y_pred, average='weighted')\n\ndef calculate_recall(y_true, y_pred):\n    return recall_score(y_true, y_pred, average='weighted')\n\ndef calculate_f1(y_true, y_pred): \n    return f1_score(y_true, y_pred, average='weighted')\n\ndef calculate_accuracy(y_true, y_pred):\n    return accuracy_score(y_true, y_pred)\n\ny_test_condition = y_test[:, 0]\ny_test_severity = y_test[:, 1]\n\ncondition_precision = calculate_precision(y_test_condition, condition_label_preds)\ncondition_recall = calculate_recall(y_test_condition, condition_label_preds)\ncondition_f1 = calculate_f1(y_test_condition, condition_label_preds)\ncondition_accuracy = calculate_accuracy(y_test_condition, condition_label_preds)\n\nseverity_precision = calculate_precision(y_test_severity, severity_label_preds)\nseverity_recall = calculate_recall(y_test_severity, severity_label_preds)\nseverity_f1 = calculate_f1(y_test_severity, severity_label_preds)\nseverity_accuracy = calculate_accuracy(y_test_severity, severity_label_preds)\n\nprint(f\"Condition Accuracy: {condition_accuracy:.4f}\")\nprint(f\"Condition Precision: {condition_precision:.4f}\")\nprint(f\"Condition Recall: {condition_recall:.4f}\")\nprint(f\"Condition F1 Score: {condition_f1:.4f}\")\nprint()\nprint(f\"Severity Accuracy: {severity_accuracy:.4f}\")\nprint(f\"Severity Precision: {severity_precision:.4f}\")\nprint(f\"Severity Recall: {severity_recall:.4f}\")\nprint(f\"Severity F1 Score: {severity_f1:.4f}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T16:41:00.346722Z","iopub.execute_input":"2024-12-02T16:41:00.347076Z","iopub.status.idle":"2024-12-02T16:41:00.370552Z","shell.execute_reply.started":"2024-12-02T16:41:00.347030Z","shell.execute_reply":"2024-12-02T16:41:00.369891Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Visualization","metadata":{}},{"cell_type":"code","source":"def visualize_feature_maps(model, layer_names, input_image, input_metadata):\n    # Extract the VGG16 submodel\n    vgg_model = model.get_layer('vgg16')  # VGG16 base model\n\n    # Get the outputs of the specified layers\n    layer_outputs = [vgg_model.get_layer(name).output for name in layer_names]\n\n    # Create a new model that takes only the image input and outputs the VGG16 feature maps\n    activation_model = Model(inputs=vgg_model.input, outputs=layer_outputs)\n    \n    # Getting the feature maps\n    feature_maps = activation_model.predict(np.expand_dims(input_image, axis=0))\n    \n    for layer_name, feature_map in zip(layer_names, feature_maps):\n        num_filters = feature_map.shape[-1]  # Number of filters in the layer (channel dim)\n        size = feature_map.shape[1]  # Feature map size (height/width)\n        \n        # Plotting each filter in the feature map\n        fig, axes = plt.subplots(1, min(num_filters, 8), figsize=(20, 8))\n        fig.suptitle(f'Layer: {layer_name}')\n\n        # plt.subplots_adjust(top=0.85, bottom=0.1, left=0.1, right=0.9, hspace=0.4, wspace=0.4) # attempt to fix plotting gaps\n        range_lim = min(num_filters, 8)\n        feature_map_inds =  [min(int((i * num_filters) / range_lim), num_filters - 1) for i in range(range_lim)]\n        # ^can manually change this to see different channels\n        # ^currently indices are equally spaced\n        for i in range(range_lim):  # Showing up to 8 filters\n            ax = axes[i]\n            ax.matshow(feature_map[0, :, :, feature_map_inds[i]], cmap='viridis')\n            ax.axis('off')\n        \n        plt.show()\n\ndef visualize_filters(model, layer_name):\n    # Extract the VGG16 submodel\n    vgg_model = model.get_layer('vgg16')  # VGG16 base model\n\n    filters, biases = vgg_model.get_layer(name=layer_name).get_weights()\n    \n    # Normalizing filter values to 0-1 for better visualization\n    f_min, f_max = filters.min(), filters.max()\n    filters = (filters - f_min) / (f_max - f_min)\n    \n    num_filters = filters.shape[-1]  # Number of filters\n    num_channels = filters.shape[-2]  # Input channels (e.g., 3 for RGB)\n\n    fig, axes = plt.subplots(num_channels, min(num_filters, 8), figsize=(20, 8))\n    fig.suptitle(f'Filters of layer: {layer_name}')\n    \n    for i in range(min(num_filters, 8)):  # Showing up to 8 filters\n        for j in range(num_channels):\n            ax = axes[j, i] if num_channels > 1 else axes[i]\n            ax.matshow(filters[:, :, j, i], cmap='viridis')\n            ax.axis('off')\n    \n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T16:50:37.102888Z","iopub.execute_input":"2024-12-02T16:50:37.103238Z","iopub.status.idle":"2024-12-02T16:50:37.114274Z","shell.execute_reply.started":"2024-12-02T16:50:37.103207Z","shell.execute_reply":"2024-12-02T16:50:37.113479Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"conv_layer_names = [layer.name for layer in model.layers[1].layers if isinstance(layer, tf.keras.layers.Conv2D)]\nprint(conv_layer_names)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T16:50:37.776957Z","iopub.execute_input":"2024-12-02T16:50:37.777577Z","iopub.status.idle":"2024-12-02T16:50:37.782272Z","shell.execute_reply.started":"2024-12-02T16:50:37.777542Z","shell.execute_reply":"2024-12-02T16:50:37.781406Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_image_ind = 0\ninput_image = X_test[test_image_ind]\ninput_metadata = tf.convert_to_tensor(X_metadata_test[test_image_ind], dtype=tf.float32)\nvisualize_feature_maps(model, conv_layer_names, input_image, input_metadata)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T16:50:38.441382Z","iopub.execute_input":"2024-12-02T16:50:38.441719Z","iopub.status.idle":"2024-12-02T16:50:53.984635Z","shell.execute_reply.started":"2024-12-02T16:50:38.441687Z","shell.execute_reply":"2024-12-02T16:50:53.983194Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"for layer_name in conv_layer_names:\n    visualize_filters(model, layer_name)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-02T16:50:53.989353Z","iopub.execute_input":"2024-12-02T16:50:53.989949Z","iopub.status.idle":"2024-12-02T16:52:46.478581Z","shell.execute_reply.started":"2024-12-02T16:50:53.989889Z","shell.execute_reply":"2024-12-02T16:52:46.477105Z"}},"outputs":[],"execution_count":null}]}