{"metadata":{"kaggle":{"accelerator":"gpu","dataSources":[{"sourceType":"competition","sourceId":71549,"databundleVersionId":8561470},{"sourceType":"modelInstanceVersion","sourceId":136172,"isSourceIdPinned":true},{"sourceType":"modelInstanceVersion","sourceId":136186,"isSourceIdPinned":true}],"dockerImageVersionId":30775,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true},"kernelspec":{"display_name":"Python 3 (ipykernel)","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.12.4"}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd \nimport SimpleITK as sitk\nimport numpy as np\nimport torch\nfrom torchvision import transforms\nimport cv2\nimport pydicom\nimport numpy as np\nimport os\nimport glob\nfrom tqdm import tqdm\nimport warnings\nfrom tensorflow.keras.preprocessing.image import load_img, img_to_array\nimport matplotlib.pyplot as plt\nfrom imblearn.over_sampling import SMOTE\n\nfrom sklearn.model_selection import train_test_split\n\n","metadata":{"_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","execution":{"iopub.execute_input":"2024-10-15T04:15:24.995906Z","iopub.status.busy":"2024-10-15T04:15:24.995530Z","iopub.status.idle":"2024-10-15T04:15:37.435838Z","shell.execute_reply":"2024-10-15T04:15:37.434995Z","shell.execute_reply.started":"2024-10-15T04:15:24.995872Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Introduction\n\nLow back pain is the leading cause of disability worldwide, as reported by the World Health Organization, affecting approximately 619 million people in 2020. The prevalence of low back pain increases with age, with most individuals experiencing it at some point in their lives. Pain and limited mobility are common symptoms of spondylosis, a collection of degenerative spine conditions characterized by the degeneration of intervertebral discs and subsequent narrowing of the spinal canal (spinal stenosis). This can lead to compression or irritation of the nerves in the lower back, significantly impacting the quality of life.\n\nMagnetic Resonance Imaging (MRI) offers a detailed visualization of the lumbar spine, including vertebrae, discs, and nerves, enabling radiologists to assess the presence and severity of degenerative conditions accurately. Proper diagnosis and grading of these conditions are crucial for guiding treatment and potential surgical interventions aimed at alleviating back pain and enhancing overall health.\n\nThe Radiological Society of North America (RSNA) has partnered with the American Society of Neuroradiology (ASNR) to conduct a competition exploring the potential of artificial intelligence (AI) in aiding the detection and classification of degenerative spine conditions using lumbar spine MRI images.\n\n## Problem Statement\n\nThe challenge focuses on the classification of five lumbar spine degenerative conditions:\n\n1. **Left Neural Foraminal Narrowing**\n2. **Right Neural Foraminal Narrowing**\n3. **Left Subarticular Stenosis**\n4. **Right Subarticular Stenosis**\n5. **Spinal Canal Stenosis**\n\nFor each imaging study in the dataset, severity scores (Normal/Mild, Moderate, or Severe) have been provided for these five conditions across the intervertebral disc levels L1/L2, L2/L3, L3/L4, L4/L5, and L5/S1.\n\nThe ground truth dataset was created through collaboration between the RSNA challenge planning task force and eight imaging sites across five continents. This expertly curated, multi-institutional dataset aims to enhance standardized classification of degenerative lumbar spine conditions and facilitate the development of tools for accurate and rapid disease classification.\n\n## Evaluation Metrics\n\nSubmissions for the competition will be evaluated based on the average of sample weighted log losses and an `any_severe_spinal` prediction. The specific sample weights are as follows:\n\n- **1** for Normal/Mild\n- **2** for Moderate\n- **4** for Severe\n\nFor each row ID in the test set, predictions must include probabilities for each severity level. The submission file should contain a header and adhere to the following format:\n\n\nIn some instances, the lowest vertebrae may not be visible in the imagery. It is still necessary to make predictions for these cases; however, they will not be scored.\n\nFor this competition, the `any_severe_scalar` is set to **1.0**.\n","metadata":{}},{"cell_type":"code","source":"# Define the path for the dataset\ntrain_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification/'\n\n# Load CSV files\ntrain_df = pd.read_csv(os.path.join(train_path, 'train.csv'))\nlabel_df = pd.read_csv(os.path.join(train_path, 'train_label_coordinates.csv'))\ntrain_description_df = pd.read_csv(os.path.join(train_path, 'train_series_descriptions.csv'))\ntest_description_df = pd.read_csv(os.path.join(train_path, 'test_series_descriptions.csv'))\nsubmission_df = pd.read_csv(os.path.join(train_path, 'sample_submission.csv'))","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:15:52.856651Z","iopub.status.busy":"2024-10-15T04:15:52.856246Z","iopub.status.idle":"2024-10-15T04:15:52.949664Z","shell.execute_reply":"2024-10-15T04:15:52.948680Z","shell.execute_reply.started":"2024-10-15T04:15:52.856611Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Data Overview\n\nIn this section, we will provide an overview of the datasets used in the project. We will detail the characteristics of each dataset, including the number of entries, columns, and data types.\n\n### Training Dataset (`train_df`)\n\nThe training dataset contains the features and target variables for our machine learning model. It has a total of **1975** entries and **26** columns.\n","metadata":{}},{"cell_type":"code","source":"train_df.info()","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:15:54.445746Z","iopub.status.busy":"2024-10-15T04:15:54.445357Z","iopub.status.idle":"2024-10-15T04:15:54.479087Z","shell.execute_reply":"2024-10-15T04:15:54.478167Z","shell.execute_reply.started":"2024-10-15T04:15:54.445708Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"label_df.info()","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:15:55.256673Z","iopub.status.busy":"2024-10-15T04:15:55.256268Z","iopub.status.idle":"2024-10-15T04:15:55.276114Z","shell.execute_reply":"2024-10-15T04:15:55.275007Z","shell.execute_reply.started":"2024-10-15T04:15:55.256637Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"test_description_df.info()","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:15:55.666735Z","iopub.status.busy":"2024-10-15T04:15:55.666352Z","iopub.status.idle":"2024-10-15T04:15:55.676901Z","shell.execute_reply":"2024-10-15T04:15:55.675935Z","shell.execute_reply.started":"2024-10-15T04:15:55.666698Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Image Path Generation\n\nIn this section, we define a function that generates image paths based on the directory structure of the dataset. The function takes in a DataFrame containing study and series IDs, along with the base directory where the images are stored. \n\n### Function: `generate_image_paths`\n\nThis function traverses the directory structure for each study and series ID, retrieves the filenames of the images, and constructs full paths to these images.\n\n#### Parameters:\n- `df` (pd.DataFrame): A DataFrame containing the `study_id` and `series_id`.\n- `data_dir` (str): The base directory path where the images are stored.\n\n#### Returns:\n- `list`: A list of full paths to the images.\n\n### Code Implementation\n","metadata":{}},{"cell_type":"code","source":"def 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        \n        # List images in the series directory\n        images = os.listdir(series_dir)\n        # Create full paths for each image\n        image_paths.extend([os.path.join(series_dir, img) for img in images])\n        \n    return image_paths","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:15:56.539313Z","iopub.status.busy":"2024-10-15T04:15:56.538927Z","iopub.status.idle":"2024-10-15T04:15:56.545565Z","shell.execute_reply":"2024-10-15T04:15:56.544586Z","shell.execute_reply.started":"2024-10-15T04:15:56.539276Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import os\nimport pandas as pd\n\n# Assuming your previous functions and imports are already in place\n\n# Generate image paths for train and test data\ntrain_image_paths = generate_image_paths(train_description_df, os.path.join(train_path, 'train_images'))\ntest_image_paths = generate_image_paths(test_description_df, os.path.join(train_path, 'test_images'))\n\n# Print to verify paths\nprint(\"Train Image Paths:\")\nprint(train_image_paths[:5])  # Print the first 5 train image paths\n\nprint(\"\\nTest Image Paths:\")\nprint(test_image_paths[:5])  # Print the first 5 test image paths\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:15:56.981157Z","iopub.status.busy":"2024-10-15T04:15:56.980297Z","iopub.status.idle":"2024-10-15T04:18:40.025779Z","shell.execute_reply":"2024-10-15T04:18:40.024846Z","shell.execute_reply.started":"2024-10-15T04:15:56.981114Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Image Visualization\n\nIn this section, we visualize a set of images from the training and testing datasets. We utilize the `SimpleITK` library to read DICOM images and display them using `matplotlib`. The goal is to provide a visual understanding of the data we are working with, specifically focusing on lumbar spine images that relate to degenerative conditions.\n\n### Function: `visualize_images`\n\nThis function takes a list of image paths and the number of images to display, then renders these images in a grid format.\n\n#### Parameters:\n- `image_paths` (list): A list of full paths to the images to visualize.\n- `num_images` (int): The number of images to display (default is 5).\n\n### Code Implementation\n","metadata":{}},{"cell_type":"code","source":"import os\nimport matplotlib.pyplot as plt\nimport SimpleITK as sitk\nfrom PIL import Image\n\n# Function to visualize a set of images\ndef visualize_images(image_paths, num_images=5):\n    plt.figure(figsize=(15, 10))\n    \n    # Loop through the number of images to display\n    for i in range(num_images):\n        # Load the DICOM image using SimpleITK\n        image = sitk.ReadImage(image_paths[i])\n        image_array = sitk.GetArrayFromImage(image)  # Convert to NumPy array\n        image_array = np.squeeze(image_array)  # Remove single-dimensional entries\n        \n        # Convert NumPy array to PIL Image\n        pil_image = Image.fromarray(image_array)\n\n        # Plotting the original image\n        plt.subplot(1, num_images, i + 1)\n        plt.imshow(pil_image, cmap='gray')  # Display as grayscale\n        plt.axis('off')  # Hide axis\n        plt.title(f'Image {i + 1}')\n    \n    plt.tight_layout()\n    plt.show()\n\n# Visualize original training images\nprint(\"Visualizing Original Training Images:\")\nvisualize_images(train_image_paths, num_images=5)\n\n# Visualize original testing images\nprint(\"Visualizing Original Testing Images:\")\nvisualize_images(test_image_paths, num_images=5)\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:18:40.027676Z","iopub.status.busy":"2024-10-15T04:18:40.027351Z","iopub.status.idle":"2024-10-15T04:18:42.927400Z","shell.execute_reply":"2024-10-15T04:18:42.926518Z","shell.execute_reply.started":"2024-10-15T04:18:40.027640Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Structuring the DataFrame\n\nIn this section, we transform the original `train_df` DataFrame into a more structured format that is easier to work with for analysis and modeling. The goal is to create a long-form DataFrame where each row corresponds to a specific condition, its level, and severity for a given study.\n\n### Data Structure\n\nThe new structured DataFrame will have the following columns:\n- **study_id**: Identifier for the study.\n- **condition**: The degenerative condition being analyzed (e.g., \"Left Neural Foraminal Narrowing\").\n- **level**: The specific intervertebral disc level being assessed (e.g., \"L1/L2\").\n- **severity**: The severity of the condition, categorized as Normal/Mild, Moderate, or Severe.\n\n### Code Implementation\n","metadata":{}},{"cell_type":"code","source":"# Initialize a list to hold the structured data\nstructured_df = []\n\n# Iterate through each row in the train DataFrame\nfor _, row in train_df.iterrows():\n    # Initialize a dictionary to hold data for the current row\n    df = {\n        'study_id': [],\n        'condition': [],\n        'level': [],\n        'severity': []\n    }\n    \n    # Iterate through the columns in the current row\n    for column, value in row.items():\n        # Skip specific columns that do not contribute to the structured data\n        if column not in ['study_id', 'series_id', 'instance_number', 'x', 'y', 'series_description']:\n            # Split the column name to get condition and level\n            parts = column.split('_')\n            condition = ' '.join([word.capitalize() for word in parts[:-2]])  # Construct condition name\n            level = f\"{parts[-2].capitalize()}/{parts[-1].capitalize()}\"  # Construct level\n\n            # Append the data to the dictionary\n            df['study_id'].append(row['study_id'])\n            df['condition'].append(condition)\n            df['level'].append(level)\n            df['severity'].append(value)\n    \n    # Append the dictionary as a DataFrame to the structured_df list\n    structured_df.append(pd.DataFrame(df))\n\n# all individual DataFrames into a single DataFrame\nstructured_df = pd.concat(structured_df, ignore_index=True)\n\n# Display the structured DataFrame\nstructured_df.head()\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:18:56.849087Z","iopub.status.busy":"2024-10-15T04:18:56.848689Z","iopub.status.idle":"2024-10-15T04:18:58.471965Z","shell.execute_reply":"2024-10-15T04:18:58.470951Z","shell.execute_reply.started":"2024-10-15T04:18:56.849050Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Merging DataFrames\n\nIn this section, we perform two merging operations to combine the structured data with labels and descriptions. The goal is to create a comprehensive DataFrame that contains all relevant information needed for model training.\n\n### Merge Operations\n\n1. **First Merge**: Combine the `structured_df` with `label_df`\n   - **Purpose**: This merge integrates the severity labels for each condition and level associated with a study.\n   - **Key Columns**: The merge is performed on the columns `study_id`, `condition`, and `level`.\n\n2. **Second Merge**: Combine the result of the first merge (`merged_df`) with `train_description_df`\n   - **Purpose**: This final merge adds the series information to our DataFrame, ensuring that all necessary details are included.\n   - **Key Columns**: The merge is performed on `series_id` and `study_id`.\n\n### Code Implementation\n","metadata":{}},{"cell_type":"code","source":"# First merge: structured_df with label\nmerged_df = pd.merge(structured_df, label_df, on=['study_id', 'condition', 'level'], how='inner')\n\n# Second merge: merged_df with train_description\nfinal_df = pd.merge(merged_df, train_description_df, on=['series_id', 'study_id'], how='inner')\n\n# Display the first few rows of the final DataFrame\nfinal_df.head()","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:18:58.780861Z","iopub.status.busy":"2024-10-15T04:18:58.780478Z","iopub.status.idle":"2024-10-15T04:18:58.862039Z","shell.execute_reply":"2024-10-15T04:18:58.861105Z","shell.execute_reply.started":"2024-10-15T04:18:58.780824Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"final_df['severity'].value_counts()\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:18:59.714221Z","iopub.status.busy":"2024-10-15T04:18:59.713842Z","iopub.status.idle":"2024-10-15T04:18:59.730716Z","shell.execute_reply":"2024-10-15T04:18:59.729811Z","shell.execute_reply.started":"2024-10-15T04:18:59.714184Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Handling NaN Values\n\nIn this section, we address the presence of NaN (Not a Number) values in our final DataFrame (`final_df`). Handling missing data is crucial as it can significantly impact the performance of machine learning models.\n\n### Steps to Handle NaN Values\n\n1. **Identify NaN Values**:\n   - We begin by checking for NaN values in each column of the `final_df`. This allows us to understand the extent of missing data and decide on further actions.\n\n2. **Drop Rows with NaN Values**:\n   - After identifying the NaN values, we choose to drop any rows that contain them. This step is optional and should be considered based on the specific dataset and the importance of the affected rows.\n\n### Code Implementation\n","metadata":{}},{"cell_type":"code","source":"# Identify NaN values\nprint(\"NaN values in each column:\")\nprint(final_df.isnull().sum())\n\n\n# Dropping rows with NaN values \nfinal_df = final_df.dropna()\n\nprint(\"After dropping NaN values:\")\nfinal_df.isnull().sum()\n\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:00.992281Z","iopub.status.busy":"2024-10-15T04:19:00.991920Z","iopub.status.idle":"2024-10-15T04:19:01.064412Z","shell.execute_reply":"2024-10-15T04:19:01.063419Z","shell.execute_reply.started":"2024-10-15T04:19:00.992246Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Mapping Severity Levels\n\nIn this section, we categorize the severity levels in our final DataFrame (`final_df`). This process simplifies the representation of severity levels, making them more suitable for analysis and model training.\n\n### Steps to Map Severity Levels\n\n1. **Map Severity Levels**:\n   - We use the `map()` function to replace the existing severity level names with more standardized labels:\n     - 'Normal/Mild' is mapped to 'normal_mild'\n     - 'Moderate' is mapped to 'moderate'\n     - 'Severe' is mapped to 'severe'\n\n2. **Handle NaN Values**:\n   - After mapping, we check for any NaN values that may arise (for example, if any severity levels do not match the mapping). If there are any NaN values, we fill them with 'unknown' to maintain a complete dataset.\n\n### Code Implementation\n","metadata":{}},{"cell_type":"code","source":"# Mapping severity levels\nfinal_df['severity'] = final_df['severity'].map({\n    'Normal/Mild': 'normal_mild', \n    'Moderate': 'moderate', \n    'Severe': 'severe'\n})\n\n# Optionally, handle NaN values that might result from mapping\nfinal_df['severity'] = final_df['severity'].fillna('unknown')  # Fill NaN with 'unknown' if needed\n\n# Display the updated DataFrame with severity levels\nprint(\"Dataframe with severity levels:\")\nfinal_df.head()\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:02.426826Z","iopub.status.busy":"2024-10-15T04:19:02.426461Z","iopub.status.idle":"2024-10-15T04:19:02.454973Z","shell.execute_reply":"2024-10-15T04:19:02.454089Z","shell.execute_reply.started":"2024-10-15T04:19:02.426784Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Creating a Unique Row Identifier\n\nIn this section, we generate a unique identifier for each row in our final DataFrame (`final_df`). This identifier, `row_id`, will facilitate easier referencing of specific records and improve data management.\n\n### Steps to Create a Unique Row Identifier\n\n1. **Generate `row_id`**:\n   - The unique identifier is constructed by concatenating the following components:\n     - The `study_id` converted to a string.\n     - The `condition` string converted to lowercase and spaces replaced with underscores (`_`).\n     - The `level` string converted to lowercase and the slash (`/`) replaced with an underscore (`_`).\n\n### Code Implementation\n\n","metadata":{}},{"cell_type":"code","source":"# Creating a unique row_id\nfinal_df['row_id'] = (\n    final_df['study_id'].astype(str) + '_' +\n    final_df['condition'].str.lower().str.replace(' ', '_') + '_' +\n    final_df['level'].str.lower().str.replace('/', '_')\n)\n\nfinal_df.sample(3)","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:03.776875Z","iopub.status.busy":"2024-10-15T04:19:03.776131Z","iopub.status.idle":"2024-10-15T04:19:03.909572Z","shell.execute_reply":"2024-10-15T04:19:03.908636Z","shell.execute_reply.started":"2024-10-15T04:19:03.776833Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Creating the Image Path Column\n\nIn this section, we generate a column named `image_path` in our final DataFrame (`final_df`). This column will contain the complete file paths to the DICOM images associated with each record, allowing for easy access during analysis or model training.\n\n### Steps to Create the `image_path` Column\n\n1. **Generate `image_path`**:\n   - The complete image path is constructed by concatenating the following components:\n     - The base path to the training images (`train_path`).\n     - The `study_id`, which identifies the study.\n     - The `series_id`, which corresponds to the series of images within the study.\n     - The `instance_number`, which identifies the specific image within the series.\n     - The file extension `.dcm`, indicating that the file is in DICOM format.\n\n### Code Implementation\n","metadata":{}},{"cell_type":"code","source":"# Creating the image_path column\nfinal_df['image_path'] = (\n    train_path + '/train_images/' + \n    final_df['study_id'].astype(str) + '/' +\n    final_df['series_id'].astype(str) + '/' +\n    final_df['instance_number'].astype(str) + '.dcm'\n)\n\n# Display the updated DataFrame with the new image_path\nfinal_df.sample(3)","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:05.015490Z","iopub.status.busy":"2024-10-15T04:19:05.014787Z","iopub.status.idle":"2024-10-15T04:19:05.136016Z","shell.execute_reply":"2024-10-15T04:19:05.135128Z","shell.execute_reply.started":"2024-10-15T04:19:05.015450Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Visualizing DICOM Images with Markers\n\nIn this section, we define a function to visualize a set of DICOM images, overlaying markers at specified coordinates on the images. This is particularly useful for identifying important regions or conditions within the images, such as pathological findings.\n\n### Function: `visualize_images`\n\nThe `visualize_images` function takes several parameters to display images along with corresponding markers:\n\n- **Parameters**:\n  - `image_paths`: A list of file paths to the DICOM images to be visualized.\n  - `final_df`: The DataFrame containing metadata about the images, including marker coordinates.\n  - `study_id`: The identifier for the study to filter the relevant markers.\n  - `series_id`: The identifier for the series within the study to filter the relevant markers.\n  - `num_images`: The number of images to display (default is set to 5).\n\n### Steps in the Function:\n\n1. **Initialize the Plot**:\n   - Create a figure to hold the images using `matplotlib`.\n\n2. **Extract Marker Coordinates**:\n   - Filter the `final_df` DataFrame to get the marker coordinates corresponding to the specified `study_id` and `series_id`.\n\n3. **Load and Display Images**:\n   - Loop through the specified number of images:\n     - Load each DICOM image using `SimpleITK`.\n     - Convert the image to a NumPy array and squeeze it to remove any single-dimensional entries.\n     - Convert the NumPy array to a PIL image for plotting.\n\n4. **Overlay Markers**:\n   - If markers exist for the current `study_id` and `series_id`, overlay them on the image using red circles.\n\n5. **Display the Images**:\n   - Each image is displayed in grayscale with markers, and the axes are hidden for a cleaner view.\n\n### Code Implementation\n","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport SimpleITK as sitk\nfrom PIL import Image\n\n# Function to visualize a set of images with markers\ndef visualize_images(image_paths, final_df, study_id, series_id, num_images=5):\n    plt.figure(figsize=(15, 10))\n    \n    # Extract marker coordinates for the specified study_id and series_id\n    markers = final_df[(final_df['study_id'] == study_id) & (final_df['series_id'] == series_id)]\n    \n    # Loop through the number of images to display\n    for i in range(num_images):\n        # Load the DICOM image using SimpleITK\n        image = sitk.ReadImage(image_paths[i])\n        image_array = sitk.GetArrayFromImage(image)  # Convert to NumPy array\n        image_array = np.squeeze(image_array)  # Remove single-dimensional entries\n        \n        # Convert NumPy array to PIL Image\n        pil_image = Image.fromarray(image_array)\n\n        # Plotting the original image\n        plt.subplot(1, num_images, i + 1)\n        plt.imshow(pil_image, cmap='gray')  # Display as grayscale\n        \n        # Overlay markers\n        if not markers.empty:\n            plt.scatter(markers['x'], markers['y'], c='red', s=100, label='Markers', marker='o')\n        \n        plt.axis('off')  # Hide axis\n        plt.title(f'Image {i + 1}')\n    \n    plt.tight_layout()\n    plt.show()\n\n# Example usage\n# Visualize original training images with markers for a specific study_id and series_id\nstudy_id = 4003253  # Example study_id\nseries_id = 702807833  # Example series_id\n\nprint(\"Visualizing Original Training Images with Markers:\")\nvisualize_images(train_image_paths, final_df, study_id, series_id, num_images=5)\n\n# Visualize original testing images (if desired)\nprint(\"Visualizing Original Testing Images with Markers:\")\nvisualize_images(test_image_paths, final_df, study_id, series_id, num_images=5)\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:06.091741Z","iopub.status.busy":"2024-10-15T04:19:06.090998Z","iopub.status.idle":"2024-10-15T04:19:09.745561Z","shell.execute_reply":"2024-10-15T04:19:09.744589Z","shell.execute_reply.started":"2024-10-15T04:19:06.091701Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Checking for Duplicate Rows in a DataFrame\n\nIn data preprocessing, it's important to ensure that there are no duplicate rows in your DataFrame, as duplicates can affect the results of analysis and modeling. This section contains code to identify and count any duplicate rows in the `final_df` DataFrame.\n\n### Steps in the Code:\n\n1. **Identify Duplicates**:\n   - Use the `duplicated()` method to find duplicate rows in the DataFrame.\n\n2. **Count Duplicate Rows**:\n   - Use the `shape` attribute to determine the number of duplicate rows.\n\n3. **Display Results**:\n   - If duplicate rows are found, display them along with the total count. If none are found, a corresponding message will be printed.\n\n### Code Implementation\n","metadata":{}},{"cell_type":"code","source":"# Check for duplicate rows\nduplicates = final_df[final_df.duplicated()]\n\n# Count the number of duplicate rows\nnum_duplicates = duplicates.shape[0]\n\n# Display the duplicate rows and count\nif num_duplicates > 0:\n    print(\"Duplicate rows found:\")\n    print(duplicates)\n    print(f\"\\nTotal number of duplicate rows: {num_duplicates}\")\nelse:\n    print(\"No duplicate rows found.\")","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:09.747905Z","iopub.status.busy":"2024-10-15T04:19:09.747486Z","iopub.status.idle":"2024-10-15T04:19:09.813334Z","shell.execute_reply":"2024-10-15T04:19:09.812393Z","shell.execute_reply.started":"2024-10-15T04:19:09.747861Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Checking Counts for Each Severity Level\n\nIn the context of medical imaging analysis, understanding the distribution of severity levels in the dataset is crucial. This section contains code to count and display the number of instances for each severity level present in the `final_df` DataFrame.\n\n### Steps in the Code:\n\n1. **Count Severity Levels**:\n   - Utilize the `value_counts()` method to compute the number of occurrences for each unique value in the `severity` column.\n\n2. **Display Results**:\n   - Print the resulting counts to the console for inspection.\n\n### Code Implementation\n","metadata":{}},{"cell_type":"code","source":"# Checking the counts for each severity level\nseverity_counts = final_df['severity'].value_counts()\n\n# Displaying the counts\nprint(\"Severity Counts:\")\nprint(severity_counts)\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:09.815019Z","iopub.status.busy":"2024-10-15T04:19:09.814598Z","iopub.status.idle":"2024-10-15T04:19:09.828267Z","shell.execute_reply":"2024-10-15T04:19:09.827102Z","shell.execute_reply.started":"2024-10-15T04:19:09.814975Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Computing Class Weights for Severity Levels\n\nIn machine learning, class imbalance can lead to models that are biased towards the majority class. To mitigate this issue, we can assign weights to each class during training, making the model pay more attention to the underrepresented classes. This section demonstrates how to compute class weights for the `severity` levels in the `final_df` DataFrame.\n\n### Steps in the Code:\n\n1. **Import Necessary Libraries**:\n   - Import `class_weight` from `sklearn.utils` to compute the class weights.\n\n2. **Get Unique Classes**:\n   - Identify the unique classes present in the `severity` column of the DataFrame.\n\n3. **Compute Class Weights**:\n   - Use the `compute_class_weight` method to calculate the weights for each class, adjusting them inversely to their frequencies.\n\n4. **Create a Class Weight Dictionary**:\n   - Map the class labels to their corresponding weights using a dictionary comprehension.\n\n5. **Display Class Weights**:\n   - Print the dictionary containing the class labels and their computed weights.\n\n### Code Implementation\n","metadata":{}},{"cell_type":"code","source":"from sklearn.utils import class_weight\nimport numpy as np\n\n# Get the unique classes and their corresponding frequencies\nclasses = final_df['severity'].unique()\nclass_weights = class_weight.compute_class_weight(\n    class_weight='balanced',  # Automatically adjusts weights inversely proportional to class frequencies\n    classes=classes,\n    y=final_df['severity']\n)\n\n# Create a dictionary mapping class labels to weights\nclass_weight_dict = dict(zip(classes, class_weights))\n\nprint(\"Class Weights:\", class_weight_dict)\n\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:12.181326Z","iopub.status.busy":"2024-10-15T04:19:12.180485Z","iopub.status.idle":"2024-10-15T04:19:12.208247Z","shell.execute_reply":"2024-10-15T04:19:12.207325Z","shell.execute_reply.started":"2024-10-15T04:19:12.181284Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Image Loading Function for DICOM and Augmented Images\n\nIn medical imaging, DICOM (Digital Imaging and Communications in Medicine) files are widely used to store and transmit medical images. This section provides a function to load images from DICOM files and other image formats, ensuring they are preprocessed appropriately for further analysis or model training.\n\n### Function Overview\n\n- **Function Name**: `load_image`\n- **Parameters**:\n  - `self`: A reference to the instance of the class (used within a class context).\n  - `image_path` (str): The path to the image file that needs to be loaded.\n\n### Function Implementation Steps\n\n1. **Check File Extension**:\n   - The function checks if the provided `image_path` ends with the `.dcm` extension to determine if it's a DICOM file.\n\n2. **Load DICOM Files**:\n   - If the file is a DICOM file, it uses `pydicom.dcmread` to read the file. The `force=True` argument allows reading even if the file doesn't conform to the DICOM standard.\n\n3. **Load Other Image Formats**:\n   - If the file is not a DICOM file, it attempts to load it as an image using OpenCV's `cv2.imread`.\n\n4. **Handle Image Data Types**:\n   - The function checks the data type of the loaded image. If it is not in `uint8` format, it converts it to `uint8`. This is important for compatibility with many image processing libraries and neural networks.\n\n5. **Convert Grayscale to RGB**:\n   - If the loaded image is grayscale (2D array), it is converted to RGB format by stacking the single channel into three identical channels. This ensures that the model receives input images in a consistent format.\n","metadata":{}},{"cell_type":"code","source":"import pydicom\nimport numpy as np\nimport cv2\nimport torch\nfrom torch.utils.data import Dataset\nfrom torchvision import transforms\n\ndef load_image(self, image_path):\n    \"\"\"Load an image from a DICOM file or an augmented image.\"\"\"\n    if image_path.lower().endswith('.dcm'):\n        # Load DICOM file\n        dicom = pydicom.dcmread(image_path, force=True)\n        image = dicom.pixel_array\n    else:\n        # Load augmented image (assumed to be in a format like PNG or JPEG)\n        image = cv2.imread(image_path, cv2.IMREAD_UNCHANGED)\n        if image is None:\n            raise FileNotFoundError(f\"Could not load image from {image_path}\")\n\n    # Convert to a format suitable for processing (e.g., uint8)\n    if image.dtype != np.uint8:\n        image = image.astype(np.uint8)\n\n    # Convert grayscale to RGB\n    if len(image.shape) == 2:  # Grayscale image\n        image = np.stack([image] * 3, axis=-1)  # Repeat the channel\n\n    return image\n\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:13.488379Z","iopub.status.busy":"2024-10-15T04:19:13.487671Z","iopub.status.idle":"2024-10-15T04:19:13.495953Z","shell.execute_reply":"2024-10-15T04:19:13.495001Z","shell.execute_reply.started":"2024-10-15T04:19:13.488338Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Custom Dataset Class for DICOM Image Loading\n\nIn this section, we define a `CustomDataset` class that inherits from PyTorch's `Dataset` class. This class is designed to load DICOM images and their corresponding severity labels from a provided DataFrame. It incorporates image loading, processing, and transformation functionalities necessary for training deep learning models.\n\n## Class Overview\n\nThe `CustomDataset` class serves the following purposes:\n- Loads images from DICOM files or augmented formats.\n- Handles grayscale images by converting them to RGB.\n- Applies transformations to images (e.g., normalization, resizing).\n- Returns the corresponding severity labels for each image.\n\n## Implementation\n","metadata":{}},{"cell_type":"code","source":"class CustomDataset(Dataset):\n    def __init__(self, dataframe, transform=None):\n        self.dataframe = dataframe\n        self.transform = transform\n\n\n    def __len__(self):\n        return len(self.dataframe)  # Ensure this returns the correct length\n\n    def load_image(self, image_path):\n        \"\"\"Load an image from a DICOM file or an augmented image.\"\"\"\n        if image_path.lower().endswith('.dcm'):\n            # Load DICOM file\n            dicom = pydicom.dcmread(image_path, force=True)\n            image = dicom.pixel_array\n        else:\n            # Load augmented image (assumed to be in a format like PNG or JPEG)\n            image = cv2.imread(image_path, cv2.IMREAD_UNCHANGED)\n            if image is None:\n                raise FileNotFoundError(f\"Could not load image from {image_path}\")\n\n        # Convert to a format suitable for processing (e.g., uint8)\n        if image.dtype != np.uint8:\n            image = image.astype(np.uint8)\n\n        # Convert grayscale to RGB if the image has a single channel\n        if len(image.shape) == 2:  # Check if the image is grayscale\n            image = np.stack([image] * 3, axis=-1)  # Repeat the channel to make it RGB\n\n        return image\n\n    def __getitem__(self, index):\n        image_path = self.dataframe['image_path'].iloc[index]\n        label = self.dataframe['severity'].iloc[index]\n        \n        # Load the image\n        image = self.load_image(image_path)\n        if image is None:\n            image = np.zeros((224, 224), dtype=np.uint8)\n\n        # Apply transformations if any\n        if self.transform:\n            image = self.transform(image)\n\n        # Ensure the label is a tensor or return as-is\n        label = torch.tensor(label) if isinstance(label, (int, float)) else label\n\n        return image, label","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:14.665555Z","iopub.status.busy":"2024-10-15T04:19:14.664932Z","iopub.status.idle":"2024-10-15T04:19:14.676140Z","shell.execute_reply":"2024-10-15T04:19:14.675174Z","shell.execute_reply.started":"2024-10-15T04:19:14.665516Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Data Augmentation and Preprocessing\n\nIn this section, we define a series of transformations to augment and preprocess our images before feeding them into a deep learning model. Data augmentation techniques help improve model robustness by artificially expanding the training dataset with transformed versions of the original images.\n\n## Transformations Overview\n\nWe utilize the `torchvision.transforms` module to compose a series of transformations. The transformations applied include random flipping, rotation, shearing, resizing, and normalization. This combination of techniques will help the model generalize better by exposing it to various representations of the data.\n\n## Implementation\n\n","metadata":{}},{"cell_type":"code","source":"import torchvision.transforms as transforms\n\ntransform = transforms.Compose([\n    transforms.ToPILImage(),  # Convert array to PIL Image\n    transforms.RandomHorizontalFlip(),  # Horizontal flip\n    transforms.RandomRotation(degrees=(-15, 15)),  # Random rotation between -15° and 15°\n    transforms.RandomAffine(degrees=0, shear=(15, 15)),  # Apply shear transformation\n    transforms.Resize((224, 224)),  # Resize to 224x224\n    transforms.ToTensor(),  # Convert to tensor and normalize to [0, 1]\n    transforms.Normalize(mean=[0.485, 0.456, 0.406], std=[0.229, 0.224, 0.225]),  # Normalize with ImageNet stats\n])\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:16.026641Z","iopub.status.busy":"2024-10-15T04:19:16.025739Z","iopub.status.idle":"2024-10-15T04:19:16.032904Z","shell.execute_reply":"2024-10-15T04:19:16.032019Z","shell.execute_reply.started":"2024-10-15T04:19:16.026585Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"from sklearn.model_selection import train_test_split\nfrom torch.utils.data import DataLoader","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:16.686153Z","iopub.status.busy":"2024-10-15T04:19:16.685766Z","iopub.status.idle":"2024-10-15T04:19:16.690719Z","shell.execute_reply":"2024-10-15T04:19:16.689688Z","shell.execute_reply.started":"2024-10-15T04:19:16.686114Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Splitting the Dataset into Training and Validation Sets\n\nTo evaluate our model's performance effectively, we need to split our dataset into two distinct subsets: a training set and a validation set. The training set will be used to train the model, while the validation set will allow us to assess the model's performance on unseen data during training.\n\n## Implementation\n\nWe utilize the `train_test_split` function from the `sklearn.model_selection` module to perform the split. This function randomly divides the dataset based on the specified test size, ensuring a representative sample of the data in both subsets.\n","metadata":{}},{"cell_type":"code","source":"train_df, val_df = train_test_split(final_df, test_size=0.2, random_state=42)","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:18.158333Z","iopub.status.busy":"2024-10-15T04:19:18.157932Z","iopub.status.idle":"2024-10-15T04:19:18.181535Z","shell.execute_reply":"2024-10-15T04:19:18.180613Z","shell.execute_reply.started":"2024-10-15T04:19:18.158293Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Balancing the Training Set\n\nIn classification tasks, it's common to encounter imbalanced datasets, where some classes have significantly more samples than others. This imbalance can lead to biased model training, resulting in poor performance on the underrepresented classes. To mitigate this issue, we will balance the classes in our training set.\n\n## Implementation\n\nWe employ two techniques to balance the training data: **Resampling** and **SMOTE (Synthetic Minority Over-sampling Technique)**.\n\n### 1. Resampling\n\nWe will first implement a simple resampling method to increase the number of samples in the minority classes by duplicating existing samples.\n\n### 2. Define the Function\n\nThe following function, `balance_classes`, performs the class balancing:","metadata":{}},{"cell_type":"code","source":"from sklearn.utils import resample\nfrom imblearn.over_sampling import SMOTE\n\n# Function to balance the training set\ndef balance_classes(df):\n    majority_class_count = df['severity'].value_counts().max()\n\n    # Separate classes\n    df_majority = df[df['severity'] == 'normal_mild']\n    df_moderate = df[df['severity'] == 'moderate']\n    df_severe = df[df['severity'] == 'severe']\n\n    # Upsample minority classes to match the majority class count\n    df_moderate_upsampled = resample(df_moderate, \n                                     replace=True, \n                                     n_samples=majority_class_count, \n                                     random_state=42)\n    df_severe_upsampled = resample(df_severe, \n                                   replace=True, \n                                   n_samples=majority_class_count, \n                                   random_state=42)\n\n    # Combine majority and upsampled classes\n    df_balanced = pd.concat([df_majority, df_moderate_upsampled, df_severe_upsampled])\n\n    return df_balanced.sample(frac=1, random_state=42)  # Shuffle the final dataset\n\n# Balance only the training data\ntrain_df_resampled = balance_classes(train_df)\n\n\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:19.646446Z","iopub.status.busy":"2024-10-15T04:19:19.646047Z","iopub.status.idle":"2024-10-15T04:19:19.746760Z","shell.execute_reply":"2024-10-15T04:19:19.745752Z","shell.execute_reply.started":"2024-10-15T04:19:19.646408Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Class Distribution Check\n\nAfter applying the class balancing technique, it’s important to verify the distribution of classes in our training dataset to ensure that the resampling has been effective. This section will compare the original class distribution with the distribution after resampling.\n\n##  Check Original Class Distribution\n\nBefore resampling, we assess the original class distribution in the training DataFrame. This helps us understand the extent of imbalance we are dealing with:\n\n\n","metadata":{}},{"cell_type":"code","source":"# Check class distribution after resampling\nnew_class_counts = train_df_resampled['severity'].value_counts()\nprint(\"Class counts after resampling:\\n\", new_class_counts)\n\n# Check original class distribution before resampling\noriginal_class_counts = train_df['severity'].value_counts()\nprint(\"Original class counts before resampling:\\n\", original_class_counts)","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:22.304034Z","iopub.status.busy":"2024-10-15T04:19:22.303634Z","iopub.status.idle":"2024-10-15T04:19:22.327418Z","shell.execute_reply":"2024-10-15T04:19:22.326288Z","shell.execute_reply.started":"2024-10-15T04:19:22.303992Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Class Distribution Visualization\n\nVisualizing the class distribution before and after resampling is crucial to understand how our balancing strategy has altered the dataset. This section presents a comparative bar plot illustrating the counts of each severity class in both the original and resampled training datasets.\n\n## Setting Up the Visualization\n\nWe use Matplotlib and Seaborn libraries to create the bar plots. The blue bars represent the original class distribution, while the orange bars represent the class distribution after resampling.\n\n","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Setting up the figure for bar plots\nplt.figure(figsize=(12, 6))\nsns.barplot(x=original_class_counts.index.astype(str), y=original_class_counts.values, color='blue', alpha=0.6, label='Original')\nsns.barplot(x=new_class_counts.index.astype(str), y = new_class_counts.values, color='orange', alpha=0.6, label='Resampled')\n\nplt.title('Class Distribution Before and After Resampling')\nplt.xlabel('Classes')\nplt.ylabel('Counts')\nplt.legend()\nplt.xticks(rotation=0)\nplt.show()","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:25.872443Z","iopub.status.busy":"2024-10-15T04:19:25.872014Z","iopub.status.idle":"2024-10-15T04:19:26.489051Z","shell.execute_reply":"2024-10-15T04:19:26.488055Z","shell.execute_reply.started":"2024-10-15T04:19:25.872404Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Creating Datasets and Data Loaders\n\nIn this section, we will create datasets and data loaders that are essential for feeding our data into the model during the training and validation phases. We will use the `CustomDataset` class that was defined earlier, along with the defined transformations.\n\n## Creating Datasets\n\nWe will create two datasets: one for training and one for validation. The training dataset will utilize the resampled DataFrame to ensure balanced class distributions, while the validation dataset will use the original validation DataFrame.\n","metadata":{}},{"cell_type":"code","source":"# Create datasets\ntrain_dataset = CustomDataset(train_df_resampled, transform=transform)\nval_dataset = CustomDataset(val_df, transform=transform)\n\n# Create DataLoaders\ntrain_loader = DataLoader(train_dataset, batch_size=8, shuffle=True)\nval_loader = DataLoader(val_dataset, batch_size=8, shuffle=False)\n\n# train_loader and val_loader are ready for training  model","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:27.786628Z","iopub.status.busy":"2024-10-15T04:19:27.785971Z","iopub.status.idle":"2024-10-15T04:19:27.792298Z","shell.execute_reply":"2024-10-15T04:19:27.791287Z","shell.execute_reply.started":"2024-10-15T04:19:27.786587Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Downloading and Saving Model Weights\n\nIn this section, we will utilize the `torchvision` library to download the pre-trained weights for the EfficientNetV2 model. Pre-trained models are essential for transfer learning, as they allow us to leverage knowledge gained from training on large datasets.\n\n## Importing Required Libraries\n\nWe start by importing the necessary libraries from PyTorch and torchvision.","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torchvision.models as models\nfrom torch.utils.data import DataLoader\nfrom torchvision import transforms\nfrom sklearn.model_selection import train_test_split\nfrom tqdm import tqdm\nimport pandas as pd\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:29.608944Z","iopub.status.busy":"2024-10-15T04:19:29.608252Z","iopub.status.idle":"2024-10-15T04:19:29.613791Z","shell.execute_reply":"2024-10-15T04:19:29.612804Z","shell.execute_reply.started":"2024-10-15T04:19:29.608902Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torchvision.models as models\n\n# Create a model instance to download the weights\nmodel = models.efficientnet_v2_s(weights='DEFAULT')\n\n# Save the model's state_dict to a local file\ntorch.save(model.state_dict(), 'efficientnet_v2_s_weights.pth')\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:30.695980Z","iopub.status.busy":"2024-10-15T04:19:30.695063Z","iopub.status.idle":"2024-10-15T04:19:31.958656Z","shell.execute_reply":"2024-10-15T04:19:31.957620Z","shell.execute_reply.started":"2024-10-15T04:19:30.695937Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Class Definition: UnifiedEfficientNetV2\n\nThe `UnifiedEfficientNetV2` class is a custom neural network model that extends the functionality of the EfficientNetV2 architecture to perform image classification tasks. This class is built using PyTorch, a popular deep learning framework, and encapsulates the EfficientNetV2 model along with additional layers designed for specific classification requirements. Below, we detail the components and features of this class:\n\n### Key Components\n\n1. **Inheritance from nn.Module**: \n   - The class inherits from `torch.nn.Module`, which is a base class for all neural network modules in PyTorch. This inheritance allows the class to leverage built-in functionalities such as parameter management, training, and evaluation.\n\n2. **Constructor (`__init__` method)**:\n   - The constructor takes three parameters:\n     - `num_classes`: Specifies the number of output classes for classification (default is set to 3).\n     - `pretrained`: A boolean flag indicating whether to use pre-trained weights from ImageNet. When set to `True`, it initializes the model with pre-trained weights, helping the model to converge faster during training.\n     - `weights_path`: An optional parameter for loading custom weights from a local file.\n\n   - Inside the constructor:\n     - The EfficientNetV2 model is instantiated using `models.efficientnet_v2_s()`. If `pretrained` is set to `True`, it loads the default pre-trained weights.\n     - If a `weights_path` is provided, the model's weights are loaded from the specified file using `load_state_dict()`.\n     - The input feature size (`in_features`) for the final classifier layer is extracted, which is crucial for designing our custom classifier.\n\n3. **Custom Classifier Layers**:\n   - The original classifier of the EfficientNetV2 model is replaced with an identity function (`nn.Identity()`). This is done because we want to customize the final classification layers according to our specific task.\n   - Three fully connected layers are defined to serve as the new classifier:\n     - **First Layer (`fc1`)**: \n       - A fully connected layer that maps the input features (`in_features`) to 256 output features. \n       - It includes a batch normalization layer and a ReLU activation function for non-linearity.\n     - **Second Layer (`fc2`)**: \n       - Maps the output of `fc1` (256 features) to 128 output features, similarly including batch normalization and ReLU activation.\n     - **Final Layer (`fc3`)**: \n       - Maps the output of `fc2` (128 features) to the number of classes specified by `num_classes`.\n\n4. **Dropout for Regularization**:\n   - Two dropout layers (`dropout1` and `dropout2`) are included to reduce overfitting during training:\n     - `dropout1` is applied after the first fully connected layer, with a dropout probability of 0.5.\n     - `dropout2` is applied after the second fully connected layer, with a dropout probability of 0.3.\n   - Dropout helps the model generalize better by preventing it from relying too heavily on any particular feature.\n\n5. **Forward Method**:\n   - The `forward` method defines the forward pass of the model, which takes an input tensor `x` and performs the following:\n     - It passes the input through the EfficientNetV2 model to obtain embeddings.\n     - The embeddings are then passed through the fully connected layers (`fc1`, `fc2`, and `fc3`) in sequence, applying dropout between the layers.\n     - The final output of the `forward` method is the predicted class scores for the input image.\n\n### Summary\n\nThe `UnifiedEfficientNetV2` class effectively leverages the powerful EfficientNetV2 architecture while allowing for customization of the classification head to suit specific classification tasks. By incorporating dropout layers, batch normalization, and non-linear activation functions, this model is designed to achieve high performance while maintaining robustness against overfitting.\n","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torchvision import models\n\nclass UnifiedEfficientNetV2(nn.Module):\n    def __init__(self, num_classes=3, pretrained=True, weights_path=None):\n        super(UnifiedEfficientNetV2, self).__init__()\n\n        # Load the EfficientNetV2 with optional pretrained weights\n        self.model = models.efficientnet_v2_s(weights=models.EfficientNet_V2_S_Weights.DEFAULT if pretrained else None)\n\n        # Load local weights if a path is provided\n        if weights_path is not None:\n            self.model.load_state_dict(torch.load(weights_path))\n\n        # Capture in_features from the last layer of the original classifier\n        in_features = self.model.classifier[-1].in_features  # Should be 1280 for EfficientNetV2-S\n        \n        # Replace the classifier with an identity function (we'll handle classification manually)\n        self.model.classifier = nn.Identity()\n\n        # Define fully connected layers with BatchNorm and ReLU activation\n        self.fc1 = nn.Sequential(\n            nn.Linear(in_features, 256),  # Update input size to 1280\n            nn.BatchNorm1d(256),\n            nn.ReLU(inplace=True)\n        )\n        self.fc2 = nn.Sequential(\n            nn.Linear(256, 128),\n            nn.BatchNorm1d(128),\n            nn.ReLU(inplace=True)\n        )\n        self.fc3 = nn.Linear(128, num_classes)\n\n        # Dropout layers for regularization\n        self.dropout1 = nn.Dropout(p=0.5)\n        self.dropout2 = nn.Dropout(p=0.3)\n\n    def forward(self, x):\n        # Get embeddings from EfficientNetV2\n        embeddings = self.model(x)  # Should output shape (batch_size, 1280)\n\n        # Fully connected layers with dropout and activations\n        x = self.fc1(embeddings)\n        x = self.dropout1(x)\n        x = self.fc2(x)\n        x = self.dropout2(x)\n        x = self.fc3(x)\n\n        return x\n\n\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:33.830947Z","iopub.status.busy":"2024-10-15T04:19:33.830290Z","iopub.status.idle":"2024-10-15T04:19:33.841731Z","shell.execute_reply":"2024-10-15T04:19:33.840747Z","shell.execute_reply.started":"2024-10-15T04:19:33.830904Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Model Setup: Rationale Behind Choices\n\nIn this section, we elaborate on the reasons for the various choices made during the model setup, particularly concerning the loss function, optimizer, and layer management strategy.\n\n### 1. Criterion: Cross Entropy Loss\n\n**Definition**: Cross Entropy Loss is a widely-used loss function for multi-class classification problems. It measures the dissimilarity between the true label distribution and the predicted probabilities from the model.\n\n**Reasons for Use**:\n- **Multi-Class Classification**: Since our task involves classifying images into multiple severity categories (normal_mild, moderate, severe), Cross Entropy Loss is appropriate as it handles multiple classes effectively.\n- **Probabilistic Interpretation**: The output of the model is passed through a softmax function, producing a probability distribution over classes. Cross Entropy Loss evaluates how well the predicted probabilities align with the true labels, providing a clear metric for improvement during training.\n- **Sensitivity to Class Imbalance**: Although Cross Entropy Loss itself does not inherently address class imbalance, it can be adjusted with class weights to account for uneven class distributions, making it adaptable to our dataset.\n\n### 2. Optimizer: Adam\n\n**Definition**: Adam (Adaptive Moment Estimation) is an adaptive learning rate optimization algorithm that combines the advantages of two other extensions of stochastic gradient descent.\n\n**Reasons for Use**:\n- **Efficiency**: Adam is computationally efficient and has low memory requirements, making it suitable for large datasets and models like EfficientNetV2.\n- **Adaptive Learning Rates**: Adam adjusts the learning rate for each parameter individually based on the first and second moments of the gradients. This adaptability allows the optimizer to converge more quickly and effectively compared to static learning rate methods.\n- **Robustness**: Adam works well in practice across various deep learning tasks, demonstrating robust performance even in situations where other optimizers might struggle.\n- **Regularization**: The inclusion of weight decay (L2 regularization) helps prevent overfitting by penalizing excessively large weights during training, contributing to improved generalization on unseen data.\n\n### 3. Learning Rate Scheduler: StepLR\n\n**Definition**: StepLR is a learning rate scheduler that decreases the learning rate by a factor (gamma) after a specified number of epochs (step_size).\n\n**Reasons for Use**:\n- **Dynamic Learning Rate Adjustment**: Adjusting the learning rate during training can lead to better convergence. A higher learning rate allows for faster training at the start, while a lower learning rate in later stages enables finer adjustments.\n- **Avoiding Overfitting**: By reducing the learning rate, we can prevent the model from making drastic updates to weights when it gets closer to a minimum, thereby promoting stable learning.\n- **Controlled Training Progression**: StepLR provides a structured approach to modify the learning rate, facilitating better control over the training process, especially in complex models.\n\n### 4. Layer Freezing and Unfreezing\n\n**Strategy**: Initially, we freeze the weights of the earlier layers in the model and unfreeze the weights of the final fully connected layer.\n\n**Why Freeze Layers?**\n- **Transfer Learning**: The lower layers of a pre-trained model like EfficientNetV2 typically capture generic features that are applicable across various tasks (e.g., edges, textures). By freezing these layers, we preserve this learned knowledge and avoid overwriting them with potentially noisy updates during the initial training stages.\n- **Focusing Learning on Specific Classes**: By unfreezing only the final layers, we allow the model to adapt its output to the specific classification task (in this case, the severity of conditions) while retaining the beneficial features learned during pre-training.\n\n**Why Unfreeze Final Layers?**\n- **Task-Specific Adaptation**: The final fully connected layers are the most relevant for adapting the model to our specific classification task. Unfreezing them allows the model to fine-tune its weights to better predict the class probabilities for our dataset.\n\n\nThe choices made in setting up the model—Cross Entropy Loss for the criterion, Adam for the optimizer, and the structured layer freezing strategy—are well-suited for multi-class classification tasks. These strategies effectively address challenges like class imbalance, convergence, and the need for specialized feature extraction. The addition of a learning rate scheduler ensures that the training process is optimized over time, leading to improved performance and generalization.\n","metadata":{}},{"cell_type":"code","source":"# Check for GPU availability\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Path to the weights file\nweights_path = '/kaggle/working/efficientnet_v2_s_weights.pth' \n\n# Create an instance of the model\nunified_model = UnifiedEfficientNetV2(num_classes=3, weights_path=weights_path).to(device)\n\n# Forward pass example (ensure image_inputs is defined)\n# outputs = unified_model(image_inputs)\n\n# Freeze initial layers\nfor param in unified_model.model.features.parameters():\n    param.requires_grad = False\n\n# Unfreeze the final fully connected layer\nfor param in unified_model.fc3.parameters():  # Unfreeze the final classifier layer\n    param.requires_grad = True\n\n# Convert class weights to tensor\nclass_weights_tensor = torch.tensor(list(class_weight_dict.values())).float().to(device)\n\n# Define loss function with class weights\ncriterion = nn.CrossEntropyLoss()\n\n# Use Adam optimizer with weight decay for regularization\noptimizer = optim.Adam(unified_model.parameters(), lr=0.001, weight_decay=1e-4)\n\n# Use a learning rate scheduler for dynamic adjustment\nscheduler = optim.lr_scheduler.StepLR(optimizer, step_size=7, gamma=0.1)","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:38.557702Z","iopub.status.busy":"2024-10-15T04:19:38.557052Z","iopub.status.idle":"2024-10-15T04:19:39.440027Z","shell.execute_reply":"2024-10-15T04:19:39.438669Z","shell.execute_reply.started":"2024-10-15T04:19:38.557661Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Counting Trainable Parameters\n\nIn the following code snippet, we calculate the number of trainable parameters in the `UnifiedEfficientNetV2` model. Understanding the number of trainable parameters is crucial for several reasons:","metadata":{}},{"cell_type":"code","source":"# Count trainable parameters\ntrainable_params = sum(p.numel() for p in unified_model.parameters() if p.requires_grad)\nprint(f\"Number of trainable parameters: {trainable_params}\")","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:19:42.279164Z","iopub.status.busy":"2024-10-15T04:19:42.278773Z","iopub.status.idle":"2024-10-15T04:19:42.288209Z","shell.execute_reply":"2024-10-15T04:19:42.287189Z","shell.execute_reply.started":"2024-10-15T04:19:42.279123Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Training Process Overview\n\nThe `train_model` function implements a complete model training and evaluation loop. It ensures the model is trained while tracking its performance on the validation set. In this section, we describe the key steps and logic used in this process, with a focus on model optimization, evaluation, and early stopping mechanisms.\n\n### 1. **Learning Rate Scheduler**\n\nA **learning rate scheduler** is used to adjust the learning rate dynamically during training. The specific scheduler used is `StepLR`, which reduces the learning rate by a factor of `gamma` after every `step_size` number of epochs. This helps in:\n- Allowing faster learning in early epochs.\n- Slowing down the learning rate in later epochs, preventing overshooting the optimal point.\n\n### 2. **Early Stopping Mechanism**\n\nThe training process includes an **early stopping mechanism** based on validation accuracy. Here's how it works:\n- After each epoch, if the validation accuracy improves, the model's weights are saved as the \"best\" weights.\n- If the validation accuracy does not improve for a certain number of epochs (defined by the `patience` parameter), training is stopped early to prevent overfitting. This ensures we do not continue training a model that has started to overfit on the training data.\n\n### 3. **Label Mapping**\n\nA **label mapping dictionary** is used to convert string-based class labels into numerical labels (e.g., `'normal_mild'`, `'moderate'`, `'severe'`). This mapping is essential for calculating losses and metrics in a format suitable for the neural network, as the model's output is numerical.\n\n### 4. **Model Training Loop**\n\nFor each epoch, the following steps are executed:\n\n#### a) **Training Step**\n- The model is set to **training mode**, which activates dropout layers and ensures batch normalization works in training mode.\n- **Images** and **labels** are loaded in batches from the `trainloader`, moved to the GPU (if available), and processed by the model.\n- The **loss** is calculated using the criterion (loss function), and the optimizer updates the model's parameters based on the computed gradients.\n- A **progress bar** (using `tqdm`) is used to track the training process visually, displaying the loss for each batch.\n\n#### b) **Validation Step**\n- After training, the model is switched to **evaluation mode**, where dropout is disabled, and batch normalization uses the running averages calculated during training.\n- **Validation data** is loaded, and the model's predictions are evaluated without updating the weights.\n- The **validation loss** and **validation accuracy** are computed for performance monitoring.\n  \nDuring validation, key metrics such as **probabilities**, **predictions**, and **true labels** are stored to compute evaluation metrics like the ROC curve, confusion matrix, and classification report after training.\n\n### 5. **Evaluation Metrics**\n\nTo thoroughly evaluate the model, multiple metrics are used after training. These metrics provide deeper insights into the model's performance.\n\n#### a) **ROC Curve (Receiver Operating Characteristic)**\nThe ROC curve plots the **true positive rate (TPR)** against the **false positive rate (FPR)** for each class. This is useful for understanding how well the model distinguishes between different classes across varying thresholds.\n- The **AUC (Area Under the Curve)** is calculated for each class, providing a single value that summarizes the performance of the model for that class.\n- The ROC curve for each class is plotted, showing the trade-off between sensitivity and specificity.\n\n#### b) **Confusion Matrix**\nA **confusion matrix** is generated to show the number of correct and incorrect predictions for each class. The matrix helps in:\n- Identifying which classes are most frequently misclassified.\n- Understanding the model's performance on individual classes in more detail.\n\n#### c) **Classification Report**\nThe **classification report** provides a summary of key classification metrics:\n- **Precision**: The proportion of true positive predictions out of all positive predictions.\n- **Recall**: The proportion of true positives out of all actual positives.\n- **F1-Score**: The harmonic mean of precision and recall, providing a balanced measure.\n- **Support**: The number of actual occurrences for each class.\n\nThese metrics are crucial for understanding the balance between precision and recall for each class, especially in cases where the classes might be imbalanced.\n\n### 6. **Saving the Best Model**\n\nThe model weights corresponding to the best validation accuracy are saved to a file. This allows the best-performing version of the model to be preserved and later used for inference or further fine-tuning. The model is saved as `best_model.pth` in the specified output directory.\n\n---\n\n### Summary of Evaluation Metrics:\n\n- **ROC Curve**: Evaluates the performance of the model across varying decision thresholds.\n- **Confusion Matrix**: Provides a visual representation of the model's classification performance for each class.\n- **Classification Report**: Summarizes precision, recall, F1-score, and support for each class.\n\nThese metrics provide a well-rounded evaluation of the model’s classification performance, going beyond simple accuracy and allowing for a more nuanced analysis of the model's behavior.\n","metadata":{}},{"cell_type":"code","source":"import torch.optim.lr_scheduler as lr_scheduler\nfrom copy import deepcopy\nfrom tqdm import tqdm\nimport os\nimport torch\nfrom sklearn.metrics import roc_curve, auc, confusion_matrix, classification_report, precision_recall_fscore_support\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport numpy as np\n\n# Define the train_model function with ROC, Confusion Matrix, and Classification Report\ndef train_model(model, trainloader, valloader, len_train, len_val, optimizer, criterion, num_epochs=10, patience=3, output_dir='./'):\n    # Learning rate scheduler\n    scheduler = lr_scheduler.StepLR(optimizer, step_size=2, gamma=0.1)\n    \n    # Initialize variables to track the best validation accuracy and weights\n    best_val_acc = 0.0\n    best_model_wts = deepcopy(model.state_dict())\n    counter = 0  # Early stopping counter\n\n    # Ensure output directory exists\n    os.makedirs(output_dir, exist_ok=True)\n    \n    # Label mapping dictionary\n    mapping_labels = {\n        'normal_mild': 0,\n        'moderate': 1,\n        'severe': 2\n    }\n    \n    # ROC data variables\n    all_val_labels = []\n    all_val_probs = []\n    all_val_preds = []  # To store predicted class labels for confusion matrix\n\n    # Training loop\n    for epoch in range(num_epochs):\n        model.train()\n        train_loss = 0.0\n        correct_train = 0\n        \n        # Training step\n        with tqdm(trainloader, unit=\"batch\") as tepoch:\n            for images, labels in tepoch:\n                images = images.to(device)\n                \n                # Map string labels to numerical labels and ensure they are long tensors\n                labels = torch.tensor([mapping_labels[label] for label in labels], dtype=torch.long).to(device)\n                \n                optimizer.zero_grad()\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                loss.backward()\n                optimizer.step()\n                \n                train_loss += loss.item()\n                probabilities = torch.softmax(outputs, dim=1)\n                _, predicted = torch.max(probabilities, 1)\n                correct_train += (predicted == labels).sum().item()\n                \n                tepoch.set_postfix(epoch=epoch+1, loss=train_loss / len_train)\n        \n        scheduler.step()  # Update the learning rate\n        \n        # Calculate average training loss and accuracy\n        train_loss /= len(trainloader)\n        train_acc = 100 * correct_train / len_train\n        \n        # Validation step\n        model.eval()\n        val_loss = 0.0\n        correct_val = 0\n        \n        all_val_labels_epoch = []\n        all_val_probs_epoch = []\n        all_val_preds_epoch = []\n\n        with torch.no_grad():  # Disable gradient calculations during validation\n            with tqdm(valloader, unit=\"batch\") as vepoch:\n                for images, labels in vepoch:\n                    images = images.to(device)\n                    labels = torch.tensor([mapping_labels[label] for label in labels], dtype=torch.long).to(device)\n                    \n                    outputs = model(images)\n                    loss = criterion(outputs, labels)\n                    val_loss += loss.item()\n                    \n                    probabilities = torch.softmax(outputs, dim=1)\n                    _, predicted = torch.max(probabilities, 1)\n                    correct_val += (predicted == labels).sum().item()\n\n                    # Collect probabilities and true labels for ROC and metrics\n                    all_val_labels_epoch.append(labels.cpu().numpy())\n                    all_val_probs_epoch.append(probabilities.cpu().numpy())\n                    all_val_preds_epoch.append(predicted.cpu().numpy())\n                    \n                    vepoch.set_postfix(epoch=epoch+1, val_loss=val_loss / len_val)\n        \n        # Aggregate ROC data for the epoch\n        all_val_labels.append(np.concatenate(all_val_labels_epoch))\n        all_val_probs.append(np.concatenate(all_val_probs_epoch))\n        all_val_preds.append(np.concatenate(all_val_preds_epoch))\n\n        # Calculate validation loss and accuracy\n        val_loss /= len(valloader)\n        val_acc = 100 * correct_val / len_val\n        \n        print(f\"Epoch {epoch+1}, Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%, Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        \n        # Save the best model if the validation accuracy improves\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            best_model_wts = deepcopy(model.state_dict())\n            model_save_path = os.path.join(output_dir, 'best_model.pth')\n            torch.save(best_model_wts, model_save_path)\n            print(f\"Best model saved to {model_save_path}\")\n            counter = 0  # Reset early stopping counter\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 the best model weights after training\n    model.load_state_dict(best_model_wts)\n    \n    # Convert ROC data to numpy arrays\n    all_val_labels = np.concatenate(all_val_labels)\n    all_val_probs = np.concatenate(all_val_probs)\n    all_val_preds = np.concatenate(all_val_preds)\n    \n    # Plot ROC Curve\n    plot_roc_curve(all_val_labels, all_val_probs, output_dir)\n\n    # Plot Confusion Matrix\n    plot_confusion_matrix(all_val_labels, all_val_preds, output_dir)\n\n    # Print Classification Report\n    print_classification_report(all_val_labels, all_val_preds)\n\n    # Return the trained model and best validation accuracy\n    return model, best_val_acc\n\n# Function to plot the ROC curve\ndef plot_roc_curve(true_labels, probs, output_dir):\n    num_classes = probs.shape[1]\n    fpr = {}\n    tpr = {}\n    roc_auc = {}\n\n    # Compute ROC curve and ROC area for each class\n    for i in range(num_classes):\n        fpr[i], tpr[i], _ = roc_curve(true_labels == i, probs[:, i])\n        roc_auc[i] = auc(fpr[i], tpr[i])\n\n    # Plot ROC curve for each class\n    plt.figure()\n    colors = ['blue', 'green', 'red']\n    class_names = ['normal_mild', 'moderate', 'severe']\n    \n    for i, color in enumerate(colors):\n        plt.plot(fpr[i], tpr[i], color=color, lw=2,\n                 label=f'ROC curve for {class_names[i]} (area = {roc_auc[i]:.2f})')\n    \n    plt.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--')\n    plt.xlim([0.0, 1.0])\n    plt.ylim([0.0, 1.05])\n    plt.xlabel('False Positive Rate')\n    plt.ylabel('True Positive Rate')\n    plt.title('Receiver Operating Characteristic (ROC) Curve')\n    plt.legend(loc=\"lower right\")\n    plt.show()\n    \n    # Save the ROC curve\n    roc_path = os.path.join(output_dir, 'roc_curve.png')\n    plt.savefig(roc_path)\n    plt.close()\n    print(f\"ROC curve saved to {roc_path}\")\n\n# Function to plot confusion matrix\ndef plot_confusion_matrix(true_labels, preds, output_dir):\n    cm = confusion_matrix(true_labels, preds)\n    plt.figure(figsize=(8, 6))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=['normal_mild', 'moderate', 'severe'], \n                yticklabels=['normal_mild', 'moderate', 'severe'])\n    plt.ylabel('True label')\n    plt.xlabel('Predicted label')\n    plt.title('Confusion Matrix')\n    plt.show()\n    \n    # Save the confusion matrix plot\n    cm_path = os.path.join(output_dir, 'confusion_matrix.png')\n    plt.savefig(cm_path)\n    plt.close()\n    print(f\"Confusion matrix saved to {cm_path}\")\n\n# Function to print classification report\ndef print_classification_report(true_labels, preds):\n    class_names = ['normal_mild', 'moderate', 'severe']\n    report = classification_report(true_labels, preds, target_names=class_names)\n    print(\"Classification Report:\\n\", report)\n","metadata":{"execution":{"iopub.execute_input":"2024-10-14T15:34:23.258836Z","iopub.status.busy":"2024-10-14T15:34:23.258112Z","iopub.status.idle":"2024-10-14T15:34:23.296651Z","shell.execute_reply":"2024-10-14T15:34:23.295737Z","shell.execute_reply.started":"2024-10-14T15:34:23.258800Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Training Summary\n\n### Arguments Passed to `train_model` Function:\n\n1. **`unified_model`**:\n   - A pre-initialized instance of **UnifiedEfficientNetV2**, designed for classification into 3 classes (`normal_mild`, `moderate`, `severe`). The model is loaded with pre-trained weights and ready for further training.\n\n2. **`train_loader`**:\n   - A `DataLoader` object that feeds the training dataset in batches. It handles batching, shuffling, and parallel loading of input data (`images`, `labels`) during each epoch.\n\n3. **`val_loader`**:\n   - Similar to `train_loader`, but for the validation dataset. It provides validation data in batches for performance evaluation after each training epoch.\n\n4. **`len(train_loader.dataset)`**:\n   - Specifies the number of examples in the training dataset, used to compute metrics like average loss and accuracy over the full dataset.\n\n5. **`len(val_loader.dataset)`**:\n   - Specifies the number of examples in the validation dataset, used for computing validation metrics such as accuracy and loss.\n\n6. **`optimizer`**:\n   - The **Adam optimizer** is used for updating the model's weights based on the gradients computed during backpropagation. The initial learning rate is set to `0.0001` to facilitate fine-tuning.\n\n7. **`criterion`**:\n   - A **cross-entropy loss function**, used to measure the difference between predicted probabilities and the true class labels. It is adjusted with class weights to handle imbalanced datasets.\n\n8. **`num_epochs=30`**:\n   - Specifies that training will run for a maximum of **30 epochs**, unless early stopping is triggered.\n\n9. **`patience=3`**:\n   - **Early stopping** will occur if validation accuracy does not improve for **3 consecutive epochs**. This prevents overfitting and reduces unnecessary computations.\n\n---\n\n### Output:\n\n1. **`trained_model`**:\n   - After training, this variable holds the best-performing model (i.e., the model that achieved the highest validation accuracy).\n\n2. **`best_accuracy`**:\n   - The best validation accuracy achieved during training, used to assess how well the model generalizes to unseen data.\n\n---\n\n### Summary of the Training Process:\n\n- The model will be trained for up to **30 epochs** on the training dataset using **Adam optimization** and **cross-entropy loss**.\n- After each epoch, the model will be evaluated on the **validation dataset** to track performance on unseen data.\n- If validation accuracy does not improve for **3 consecutive epochs**, **early stopping** will halt the training to avoid overfitting.\n- The **best-performing model** (the one with the highest validation accuracy) will be saved, and the corresponding **best accuracy** will be returned for further use and reporting.\n","metadata":{}},{"cell_type":"code","source":"trained_model, best_accuracy = train_model(\n    unified_model, \n    train_loader, \n    val_loader, \n    len(train_loader.dataset), \n    len(val_loader.dataset), \n    optimizer, \n    criterion, \n    num_epochs=10, \n    patience=3\n)\n","metadata":{"execution":{"iopub.execute_input":"2024-10-13T13:10:47.731799Z","iopub.status.busy":"2024-10-13T13:10:47.731044Z","iopub.status.idle":"2024-10-13T19:10:15.991527Z","shell.execute_reply":"2024-10-13T19:10:15.990391Z","shell.execute_reply.started":"2024-10-13T13:10:47.731753Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Fine-tuning and Resuming Training: Summary of Steps\n\n### 1. **Decreasing Learning Rate for Fine-Tuning**\n   - **New Learning Rate**: The learning rate is set to **0.0001**, a lower value for fine-tuning the model after initial training. A smaller learning rate helps the model make smaller updates to weights, improving performance without overshooting optimal solutions.\n\n### 2. **Updating the Optimizer**\n   - The **Adam optimizer** is reconfigured with the updated learning rate of **0.0001** and weight decay of **1e-4** to regularize the model by preventing overfitting.\n\n\n### 3. **Resuming Training from Epoch 11**\n   - The model resumes training from **epoch 11**, continuing for **10 more epochs**. Fine-tuning will help the model refine its weights for better performance.\n   - Early stopping is configured with **patience = 3**, meaning the training will stop if the validation accuracy doesn't improve for 3 consecutive epochs.\n\n---\n\n### Arguments Passed to the `resume_training` Function:\n\n- **`optimizer`**: Adam optimizer with a lower learning rate.\n- **`criterion`**: Cross-entropy loss function with class weights.\n- **`start_epoch=11`**: Resuming from epoch 11.\n- **`num_epochs=10`**: Fine-tuning for 10 more epochs.\n- **`patience=3`**: Early stopping after 3 consecutive epochs of no improvement.\n #### Rest of the parameters are same as we passed in train_model function\n\n---\n\n### Output:\n- **`trained_model`**: The fine-tuned model after training.\n- **`best_accuracy`**: The best validation accuracy achieved during the fine-tuning process.\n","metadata":{}},{"cell_type":"code","source":"# Check for GPU availability\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Path to the saved model\nmodel_path = '/kaggle/input/full/pytorch/default/1/full_model.pth'  \n  \n\n# Load the entire model\nunified_model = torch.load(model_path)\nunified_model.to(device)\nprint(\"Full model loaded.\")","metadata":{"execution":{"iopub.execute_input":"2024-10-15T04:20:06.549671Z","iopub.status.busy":"2024-10-15T04:20:06.549277Z","iopub.status.idle":"2024-10-15T04:20:07.815188Z","shell.execute_reply":"2024-10-15T04:20:07.814280Z","shell.execute_reply.started":"2024-10-15T04:20:06.549635Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import torch\nimport torch.optim as optim\nimport torch.optim.lr_scheduler as lr_scheduler\nfrom copy import deepcopy\nfrom tqdm import tqdm\nimport os\nimport numpy as np\nfrom sklearn.metrics import roc_curve, auc, confusion_matrix, classification_report\n\n# Check for GPU availability\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n\n# Adjust learning rate if needed\nnew_learning_rate = 0.0001  \noptimizer = optim.Adam(unified_model.parameters(), lr=new_learning_rate, weight_decay=1e-4)\n \n# Define loss function with class weights (make sure class weights are defined)\nclass_weights_tensor = torch.tensor(list(class_weight_dict.values())).float().to(device)\ncriterion = nn.CrossEntropyLoss()\n\n\n# Function to resume training from a specific epoch\ndef resume_training(model, trainloader, valloader, len_train, len_val, optimizer, criterion, start_epoch=11, num_epochs=10, patience=3, output_dir='./'):\n    # Learning rate scheduler\n    scheduler = lr_scheduler.ReduceLROnPlateau(optimizer, mode='max', factor=0.1, patience=2, verbose=True)\n    \n    best_val_acc = 0.0  # To track the best validation accuracy\n    best_model_wts = deepcopy(model.state_dict())  # Save the best model weights\n    counter = 0  # Early stopping counter\n    \n    # Label mapping dictionary\n    mapping_labels = {\n        'normal_mild': 0,\n        'moderate': 1,\n        'severe': 2\n    }\n    \n    # Variables to store ROC data across epochs\n    all_val_labels = []\n    all_val_probs = []\n    all_val_preds = []  # Store predicted class labels for confusion matrix\n\n    # Training loop\n    for epoch in range(start_epoch, start_epoch + num_epochs):\n        model.train()  # Set model to training mode\n        train_loss = 0.0\n        correct_train = 0\n        \n        # Training step\n        with tqdm(trainloader, unit=\"batch\") as tepoch:\n            for images, labels in tepoch:\n                images = images.to(device)\n                \n                # Map string labels to numerical labels and ensure they are long tensors\n                labels = torch.tensor([mapping_labels[label] for label in labels], dtype=torch.long).to(device)\n                \n                optimizer.zero_grad()\n                outputs = model(images)\n                loss = criterion(outputs, labels)\n                loss.backward()\n                optimizer.step()\n                \n                train_loss += loss.item()\n                probabilities = torch.softmax(outputs, dim=1)\n                _, predicted = torch.max(probabilities, 1)\n                correct_train += (predicted == labels).sum().item()\n                \n                tepoch.set_postfix(epoch=epoch + 1, loss=train_loss / len_train)\n        \n        # Validation step\n        model.eval()  # Set model to evaluation mode\n        val_loss = 0.0\n        correct_val = 0\n        \n        # Store data for ROC and Confusion Matrix\n        all_val_labels_epoch = []\n        all_val_probs_epoch = []\n        all_val_preds_epoch = []\n\n        with torch.no_grad():  # Disable gradient calculations during validation\n            with tqdm(valloader, unit=\"batch\") as vepoch:\n                for images, labels in vepoch:\n                    images = images.to(device)\n                    labels = torch.tensor([mapping_labels[label] for label in labels], dtype=torch.long).to(device)\n                    \n                    outputs = model(images)\n                    loss = criterion(outputs, labels)\n                    val_loss += loss.item()\n                    \n                    probabilities = torch.softmax(outputs, dim=1)\n                    _, predicted = torch.max(probabilities, 1)\n                    correct_val += (predicted == labels).sum().item()\n\n                    # Collect probabilities and true labels for ROC and metrics\n                    all_val_labels_epoch.append(labels.cpu().numpy())\n                    all_val_probs_epoch.append(probabilities.cpu().numpy())\n                    all_val_preds_epoch.append(predicted.cpu().numpy())\n                    \n                    vepoch.set_postfix(epoch=epoch + 1, val_loss=val_loss / len_val)\n        \n        # Aggregate ROC data for the epoch\n        all_val_labels.append(np.concatenate(all_val_labels_epoch))\n        all_val_probs.append(np.concatenate(all_val_probs_epoch))\n        all_val_preds.append(np.concatenate(all_val_preds_epoch))\n\n        # Calculate average training and validation loss, and accuracy\n        train_loss /= len(trainloader)\n        train_acc = 100 * correct_train / len_train\n        \n        val_loss /= len(valloader)\n        val_acc = 100 * correct_val / len_val\n        \n        # Print training and validation results\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        # Adjust learning rate based on validation loss\n        scheduler.step(val_loss)\n        \n        # Save the best model if the validation accuracy improves\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            best_model_wts = deepcopy(model.state_dict())\n            model_save_path = os.path.join(output_dir, 'best_model.pth')\n            torch.save(best_model_wts, model_save_path)\n            print(f\"Best model saved to {model_save_path}\")\n            counter = 0  # Reset early stopping counter\n        else:\n            counter += 1\n        \n        # Early stopping based on patience\n        if counter >= patience:\n            print(f\"Early stopping triggered after {epoch + 1 - start_epoch} epochs\")\n            break\n    \n    # Load the best model weights after training\n    model.load_state_dict(best_model_wts)\n    \n    # Convert ROC data to numpy arrays for plotting\n    all_val_labels = np.concatenate(all_val_labels)\n    all_val_probs = np.concatenate(all_val_probs)\n    all_val_preds = np.concatenate(all_val_preds)\n    \n    # Plot ROC Curve\n    plot_roc_curve(all_val_labels, all_val_probs, output_dir)\n\n    # Plot Confusion Matrix\n    plot_confusion_matrix(all_val_labels, all_val_preds, output_dir)\n\n    # Print Classification Report\n    print_classification_report(all_val_labels, all_val_preds)\n\n    return model, best_val_acc\n\n\n# Function to plot the ROC curve\ndef plot_roc_curve(true_labels, probs, output_dir):\n    num_classes = probs.shape[1]\n    fpr = {}\n    tpr = {}\n    roc_auc = {}\n\n    # Compute ROC curve and ROC area for each class\n    for i in range(num_classes):\n        fpr[i], tpr[i], _ = roc_curve(true_labels == i, probs[:, i])\n        roc_auc[i] = auc(fpr[i], tpr[i])\n\n    # Plot ROC curve for each class\n    plt.figure()\n    colors = ['blue', 'green', 'red']\n    class_names = ['normal_mild', 'moderate', 'severe']\n    \n    for i, color in enumerate(colors):\n        plt.plot(fpr[i], tpr[i], color=color, lw=2,\n                 label=f'ROC curve for {class_names[i]} (area = {roc_auc[i]:.2f})')\n    \n    plt.plot([0, 1], [0, 1], color='navy', lw=2, linestyle='--')\n    plt.xlim([0.0, 1.0])\n    plt.ylim([0.0, 1.05])\n    plt.xlabel('False Positive Rate')\n    plt.ylabel('True Positive Rate')\n    plt.title('Receiver Operating Characteristic (ROC) Curve')\n    plt.legend(loc=\"lower right\")\n    plt.show()\n    \n    # Save the ROC curve\n    roc_path = os.path.join(output_dir, 'roc_curve.png')\n    plt.savefig(roc_path)\n    plt.close()\n    print(f\"ROC curve saved to {roc_path}\")\n\n# Function to plot confusion matrix\ndef plot_confusion_matrix(true_labels, preds, output_dir):\n    cm = confusion_matrix(true_labels, preds)\n    plt.figure(figsize=(8, 6))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', xticklabels=['normal_mild', 'moderate', 'severe'], \n                yticklabels=['normal_mild', 'moderate', 'severe'])\n    plt.ylabel('True label')\n    plt.xlabel('Predicted label')\n    plt.title('Confusion Matrix')\n    plt.show()\n    \n    # Save the confusion matrix plot\n    cm_path = os.path.join(output_dir, 'confusion_matrix.png')\n    plt.savefig(cm_path)\n    plt.close()\n    print(f\"Confusion matrix saved to {cm_path}\")\n\n# Function to print classification report\ndef print_classification_report(true_labels, preds):\n    class_names = ['normal_mild', 'moderate', 'severe']\n    report = classification_report(true_labels, preds, target_names=class_names)\n    print(\"Classification Report:\\n\", report)\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T05:09:49.256023Z","iopub.status.busy":"2024-10-15T05:09:49.255603Z","iopub.status.idle":"2024-10-15T05:09:49.301067Z","shell.execute_reply":"2024-10-15T05:09:49.300081Z","shell.execute_reply.started":"2024-10-15T05:09:49.255986Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Resume training from epoch 12\nfine_tuned_model, best_val_accuracy = resume_training(\n    unified_model, \n    train_loader, \n    val_loader, \n    len(train_loader.dataset), \n    len(val_loader.dataset), \n    optimizer, \n    criterion, \n    start_epoch=11, \n    num_epochs=10, \n    patience=3, \n    output_dir='./')\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T05:09:51.123996Z","iopub.status.busy":"2024-10-15T05:09:51.123116Z","iopub.status.idle":"2024-10-15T10:13:13.409729Z","shell.execute_reply":"2024-10-15T10:13:13.408793Z","shell.execute_reply.started":"2024-10-15T05:09:51.123957Z"}},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"# Check for GPU availability\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\n\n# Path to the best model weights\nbest_model_path = './best_model.pth'  \n\n# Create an instance of the model\nunified_model = UnifiedEfficientNetV2(num_classes=3, weights_path='/kaggle/working/efficientnet_v2_s_weights.pth').to(device)\n\n# Load the best model weights\nunified_model.load_state_dict(torch.load(best_model_path))\nprint(\"Best model weights loaded.\")\n\n# Adjust learning rate if needed\nnew_learning_rate = 0.00001  # Set this to your desired learning rate\noptimizer = optim.Adam(unified_model.parameters(), lr=new_learning_rate, weight_decay=1e-4)\n \n# Define loss function with class weights (make sure class weights are defined)\nclass_weights_tensor = torch.tensor(list(class_weight_dict.values())).float().to(device)\ncriterion = nn.CrossEntropyLoss()\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T10:14:18.303104Z","iopub.status.busy":"2024-10-15T10:14:18.302471Z","iopub.status.idle":"2024-10-15T10:14:19.217215Z","shell.execute_reply":"2024-10-15T10:14:19.216297Z","shell.execute_reply.started":"2024-10-15T10:14:18.303060Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Resume Training from Epoch 20\n\n\n### 1. **Decreasing Learning Rate again**\n   - **New Learning Rate**: The learning rate is set to **0.00001**, \n\n### 2. **Updating the Optimizer**\n   - The **Adam optimizer** is reconfigured with the updated learning rate of **0.0001** and weight decay of **1e-4** to regularize the model by preventing overfitting.\n\n\n### 3. **Resuming Training from Epoch 11**\n   - The model resumes training from **epoch 20**, continuing for **10 more epochs**. Fine-tuning will help the model refine its weights for better performance.\n   - Early stopping is configured with **patience = 3**, meaning the training will stop if the validation accuracy doesn't improve for 3 consecutive epochs.\n\n---\n\n### Arguments Passed to the `resume_training` Function:\n\n- **`optimizer`**: Adam optimizer with a lower learning rate.\n- **`criterion`**: Cross-entropy loss function with class weights.\n- **`start_epoch=11`**: Resuming from epoch 21.\n- **`num_epochs=10`**: Fine-tuning for 10 more epochs.\n- **`patience=3`**: Early stopping after 3 consecutive epochs of no improvement.\n #### Rest of the parameters are same as we passed in train_model function\n\n---\n\n### Output:\n- **`trained_model`**: The fine-tuned model after training.\n- **`best_accuracy`**: The best validation accuracy achieved during the fine-tuning process.\n\n\n","metadata":{}},{"cell_type":"code","source":"# Resume training from epoch 20\nfine_tuned_model, best_val_accuracy = resume_training(\n    unified_model, \n    train_loader, \n    val_loader, \n    len(train_loader.dataset), \n    len(val_loader.dataset), \n    optimizer, \n    criterion, \n    start_epoch=20, \n    num_epochs=10, \n    patience=3, \n    output_dir='./')\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T10:15:19.745515Z","iopub.status.busy":"2024-10-15T10:15:19.744804Z","iopub.status.idle":"2024-10-15T15:01:07.561189Z","shell.execute_reply":"2024-10-15T15:01:07.560284Z","shell.execute_reply.started":"2024-10-15T10:15:19.745471Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"###  **Model Evaluation**\n\n#### **ROC Curve and AUC Analysis**:\n\n   **Class-wise AUC scores**:\n   - **normal_mild**: The ROC AUC for the `normal_mild` class is **0.80**, indicating that the model is good at distinguishing this class from the others. An AUC of 0.80 reflects a strong ability to differentiate between true positive and false positive cases.\n   - **moderate**: The ROC AUC for the `moderate` class is **0.67**, which suggests a moderate performance. This lower AUC score implies that the model has difficulty in distinguishing `moderate` cases from other classes, possibly due to overlapping features or imbalanced data.\n   - **severe**: The ROC AUC for the `severe` class is **0.82**, indicating strong performance. The model is effective in identifying `severe` cases, with a high true positive rate relative to the false positive rate.\n\n#### **Overall Analysis**:\n   - The **normal_mild** and **severe** classes have relatively high AUC values, showing the model’s competence in predicting these categories.\n   - The **moderate** class has a lower AUC, indicating room for improvement. This may be caused by class imbalances or less distinctive feature representation for this category.\n ","metadata":{}},{"cell_type":"markdown","source":"# Confusion Matrix Analysis\n### Correct Predictions (True Positives):\n- **normal_mild**: The model correctly classified **23,823** instances as \"normal_mild.\"\n- **moderate**: The model correctly classified **3,734** instances as \"moderate.\"\n- **severe**: The model correctly classified **1,715** instances as \"severe.\"\n\n### Incorrect Predictions (False Positives and False Negatives):\n- **normal_mild**: \n    - **9,981** instances of \"normal_mild\" were misclassified as \"moderate.\"\n    - **3,967** instances of \"normal_mild\" were misclassified as \"severe.\"\n- **moderate**: \n    - **1,777** instances of \"moderate\" were misclassified as \"normal_mild.\"\n    - **2,299** instances of \"moderate\" were misclassified as \"severe.\"\n- **severe**: \n    - **270** instances of \"severe\" were misclassified as \"normal_mild.\"\n    - **1,085** instances of \"severe\" were misclassified as \"moderate.\"\n\n### Class-Level Performance:\n\n### normal_mild:\n- **Correctly Classified**: 23,823 out of 37,771 actual instances (63%).\n- **Misclassification**: A significant number of instances were misclassified as \"moderate\" (9,981) and \"severe\" (3,967).\n  \n### moderate:\n- **Correctly Classified**: 3,734 out of 7,810 actual instances (48%).\n- **Misclassification**: High confusion with both \"normal_mild\" (1,777) and \"severe\" (2,299).\n\n### severe:\n- **Correctly Classified**: 1,715 out of 3,070 actual instances (56%).\n- **Misclassification**: Moderate confusion with \"moderate\" (1,085) and some with \"normal_mild\" (270).\n\n### Key Insights:\n- **normal_mild**: The model performs best on this class, with a 63% accuracy. However, there is still a significant amount of confusion between \"normal_mild\" and \"moderate.\"\n  \n- **moderate**: This is the hardest class for the model to predict, with only 48% accuracy. The model struggles to differentiate between \"moderate\" and the other classes, especially \"severe.\"\n\n- **severe**: The model does reasonably well here, with a 56% accuracy. There is, however, some confusion with the \"moderate\" class.\n\n-------------\n","metadata":{}},{"cell_type":"markdown","source":"## Classification Report Analysis\n\n### Key Metrics:\n- **Precision**: Measures how many of the predicted positive instances are correct (TP / (TP + FP)). \n  - High precision means a low false positive rate.\n- **Recall**: Measures how many of the actual positive instances were correctly identified (TP / (TP + FN)). \n  - High recall means a low false negative rate.\n- **F1-Score**: Harmonic mean of precision and recall. Useful when precision and recall need to be balanced.\n- **Support**: The number of actual occurrences of the class in the dataset.\n\n---\n\n### Class: \"normal_mild\"\n- **Precision**: 0.92\n  - 92% of the instances predicted as \"normal_mild\" were correct.\n- **Recall**: 0.63\n  - The model correctly identified 63% of the actual \"normal_mild\" cases.\n- **F1-Score**: 0.75\n  - The balance between precision and recall shows decent performance but with room for improvement in recall.\n- **Support**: 37,780\n  - This is the most frequent class in the dataset.\n\n---\n\n### Class: \"moderate\"\n- **Precision**: 0.25\n  - Only 25% of the predicted \"moderate\" labels were correct, indicating significant misclassification.\n- **Recall**: 0.48\n  - The model correctly identified 48% of the actual \"moderate\" cases.\n- **F1-Score**: 0.33\n  - Low F1-score due to low precision, showing poor overall performance in identifying \"moderate\" cases.\n- **Support**: 7,810\n  - Less frequent than \"normal_mild,\" which could contribute to misclassifications.\n\n---\n\n### Class: \"severe\"\n- **Precision**: 0.21\n  - Only 21% of the predicted \"severe\" instances were correct.\n- **Recall**: 0.56\n  - The model correctly identified 56% of the actual \"severe\" cases.\n- **F1-Score**: 0.31\n  - Poor performance in identifying \"severe\" cases, reflected in the low F1-score.\n- **Support**: 3,070\n  - This class is the least frequent in the dataset, which may cause difficulties in classification.\n\n---\n\n### Overall Metrics:\n- **Accuracy**: 60%\n  - The model correctly classified 60% of all instances across all classes.\n  \n- **Macro Average**:\n  - **Precision**: 0.46\n  - **Recall**: 0.56\n  - **F1-Score**: 0.46\n  - These averages treat all classes equally and show that the model has better recall than precision, but struggles in precision for \"moderate\" and \"severe\" classes.\n  \n- **Weighted Average**:\n  - **Precision**: 0.77\n  - **Recall**: 0.60\n  - **F1-Score**: 0.65\n  - These averages account for the frequency of each class. Since \"normal_mild\" is the most frequent, it heavily influences these averages.\n\n---\n\n- The model performs well on the dominant class (\"normal_mild\") but struggles with \"moderate\" and \"severe\" cases, especially in terms of precision.\n- Low precision for \"moderate\" and \"severe\" indicates the model often incorrectly predicts these classes.\n- Further improvements could include addressing class imbalance or refining features to help the model better distinguish between \"moderate\" and \"severe\" cases.\n","metadata":{}},{"cell_type":"code","source":"# Get the current learning rate from the optimizer after training\nfor param_group in optimizer.param_groups:\n    print(f\"Final Learning Rate: {param_group['lr']}\")\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T15:04:51.253651Z","iopub.status.busy":"2024-10-15T15:04:51.252639Z","iopub.status.idle":"2024-10-15T15:04:51.259141Z","shell.execute_reply":"2024-10-15T15:04:51.258127Z","shell.execute_reply.started":"2024-10-15T15:04:51.253596Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## Saving Model and Training Configuration in Kaggle\n\nIn this section, we describe the steps used to save the trained model, training configuration, and best accuracy to the appropriate directory in Kaggle. These steps ensure that the model and related information can be easily accessed or reused later for inference, evaluation, or retraining.\n\n### 1. **Saving Model Weights**\n\nThe first step involves saving the **trained model weights**. Only the model's parameters (weights and biases) are saved, which can later be loaded into the same model architecture for inference or fine-tuning. This is an efficient way to store the model without saving the full architecture, which allows for flexibility in using different environments.\n\n- **Why?** Saving the model weights ensures that we preserve the learned parameters, which can be loaded into a new instance of the same model architecture for future tasks (e.g., inference, further training).\n\n### 2. **Saving the Entire Model (Optional)**\n\nAs an optional step, the entire model (including the architecture) can be saved. This includes both the weights and the model's structure, which means that the model can be reloaded without needing to define the architecture separately.\n\n- **Why?** Saving the entire model allows for easy reuse in environments where the exact architecture needs to be preserved. This is particularly useful when transferring models between different environments or platforms.\n\n### 3. **Saving Training Configuration**\n\nThe training configuration, including important hyperparameters and settings, is saved in a JSON file. The following information is stored:\n- **Number of epochs**: Specifies how many iterations the model went through during training.\n- **Patience**: The patience parameter used for early stopping.\n- **Learning rate**: The learning rate used by the optimizer during training.\n- **Loss function**: The type of criterion (loss function) used to calculate the error between predicted and actual values.\n\n- **Why?** Saving the training configuration ensures that the model can be reproduced with the same settings if needed. This is particularly important when retraining or fine-tuning the model, or for hyperparameter optimization.\n\n### 4. **Saving the Best Accuracy**\n\nThe best validation accuracy achieved during the training process is saved to a text file. This provides a quick reference for how well the model performed on the validation set.\n\n- **Why?** Keeping track of the best accuracy allows us to compare the performance of different models or versions. It provides a simple performance metric for evaluating how well the model generalizes to unseen data.\n\n---\n\n### Summary of Steps:\n\n- **Model Weights**: Saves the trained parameters of the model for reuse.\n- **Entire Model** (Optional): Saves the full model (architecture + weights).\n- **Training Configuration**: Saves the training setup (epochs, learning rate, criterion).\n- **Best Accuracy**: Saves the best validation accuracy for evaluation.\n\nThese steps are crucial for ensuring that the model and its training process are well-documented and can be easily reproduced, evaluated, or deployed in different environments.\n","metadata":{}},{"cell_type":"code","source":"\n# 1. Save the model weights\ntorch.save(fine_tuned_model.state_dict(), os.path.join(output_dir, 'fine_tuned_model.pth'))\n\n# 2. Save the entire model \ntorch.save(fine_tuned_model, os.path.join(output_dir, 'full_fine_tuned_model.pth'))\n\n# 3. Save training configuration\nconfig = {\n    'num_epochs': 30,\n    'patience': 3,\n    'learning_rate': optimizer.param_groups[0]['lr'],\n    'criterion': str(criterion),\n}\nwith open(os.path.join(output_dir, 'training_config.json'), 'w') as f:\n    json.dump(config, f)\n\n# 4. Save the best accuracy\nwith open(os.path.join(output_dir, 'best_val_accuracy.txt'), 'w') as f:\n    f.write(f'Best Accuracy: {best_val_accuracy:.2f}%')\n\n\n\nprint(\"All files saved to Kaggle output directory.\")\n","metadata":{"execution":{"iopub.execute_input":"2024-10-15T15:11:37.525032Z","iopub.status.busy":"2024-10-15T15:11:37.524162Z","iopub.status.idle":"2024-10-15T15:11:38.077139Z","shell.execute_reply":"2024-10-15T15:11:38.075534Z","shell.execute_reply.started":"2024-10-15T15:11:37.524992Z"}},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# Notebook Summary: Fine-Tuning, Training, and Model Evaluation\n\n### 1. **Model Preparation**\n   - **Model Used**: UnifiedEfficientNetV2 with a classification head for 3 classes (`normal_mild`, `moderate`, `severe`).\n   - **Weights Initialization**: Pre-trained weights loaded to enhance the starting performance of the model.\n   - **Optimization Setup**: Initially used the Adam optimizer with a learning rate of `0.001` and weight decay of `1e-4` to manage overfitting.\n\n### 2. **Training Phases**\n\n#### **Phase 1: Initial Training**\n   - **Epochs**: Trained for **10 epochs** with a learning rate of `0.001`.\n   - **Loss Function**: Cross-Entropy Loss with class weighting to handle class imbalances.\n   - **Scheduler**: Used **StepLR** scheduler to reduce the learning rate after every 2 epochs.\n   - **Early Stopping**: Implemented with a patience of **3**, stopping if no improvement in validation accuracy was observed.\n   - **Best Accuracy**: Saved model checkpoints and recorded best accuracy after each epoch with improvement.\n\n#### **Phase 2: Fine-Tuning After Early Stopping**\n   - **Resumed Training**: After early stopping, resumed training for **10 more epochs** with a reduced learning rate of **0.0001**.\n   - **Scheduler**: Switched to **ReduceLROnPlateau** scheduler to reduce the learning rate when validation loss stopped improving.\n\n#### **Phase 3: Further Fine-Tuning**\n   - **Additional Epochs**: Trained for another **10 epochs** with an even lower learning rate of **0.00001** for more precise fine-tuning.\n   - **Early Stopping**: Applied early stopping as before with the same patience criteria.\n\n### 3. **Model Evaluation**\n\n\n#### ROC Curve and AUC Analysis:\n- **normal_mild**: AUC of **0.80**, indicating strong ability to distinguish this class.\n- **moderate**: AUC of **0.67**, suggesting moderate performance and significant overlap with other classes.\n- **severe**: AUC of **0.82**, showing the model's effective identification of severe cases.\n- **Overall**: The model performs well on `normal_mild` and `severe`, but struggles with `moderate` due to class imbalance and overlapping features.\n\n#### Confusion Matrix Analysis:\n- **Correct Predictions**: High accuracy for `normal_mild` (63%) and `severe` (56%), with moderate accuracy for `moderate` (48%).\n- **Misclassification**: Significant confusion between `moderate` and other classes, with many `moderate` instances misclassified as `severe`.\n- **Key Insights**: The model performs best on the `normal_mild` class but struggles to accurately predict `moderate` cases, indicating the need for further improvement in distinguishing between these classes.\n\n#### Classification Report Analysis:\n- **normal_mild**: High precision (92%) but moderate recall (63%), reflecting strong performance.\n- **moderate**: Low precision (25%) and F1-score (0.33), indicating poor performance in identifying this class.\n- **severe**: Precision (21%) and recall (56%) show difficulties in classification.\n- **Overall Accuracy**: 60%.\n- **Macro Average**: Precision (0.46), Recall (0.56), F1-Score (0.46), showing the model struggles with precision but performs better in recall.\n- **Weighted Average**: Precision (0.77), F1-Score (0.65), heavily influenced by the `normal_mild` class due to its frequency.\n\n\n### 5. **Next Steps and Improvements**\n\n#### **Potential Improvements**:\n   - **Data Augmentation**: Use techniques like random cropping, flipping, and color jittering to improve generalization.\n   - **Hyperparameter Tuning**: Adjust batch size, optimizer types (e.g., AdamW), and learning rate decay strategies.\n   - **Ensemble Methods**: Combine predictions from multiple models to boost overall performance.\n   - **Advanced Schedulers**: Experiment with **CosineAnnealingLR** or **ReduceLROnPlateau** for more adaptive learning rate scheduling.\n\n### 6. **Concluding Remarks**\n   - This notebook demonstrated how to fine-tune an EfficientNet model for a multi-class classification task.\n   - Effective evaluation techniques (ROC curve, confusion matrix, and classification report) provided in-depth insights into the model's performance.\n   - Early stopping and adaptive learning rate strategies played a crucial role in optimizing the model and preventing overfitting.\n","metadata":{}},{"cell_type":"markdown","source":"----","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]}]}