{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":39272,"databundleVersionId":4629629,"sourceType":"competition"},{"sourceId":4619805,"sourceType":"datasetVersion","datasetId":2688675}],"dockerImageVersionId":30787,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# **RSNA Dataset Description**\n\nThe dataset contains 54,706 entries with 14 columns, each representing a specific attribute related to patient and medical imaging information.","metadata":{}},{"cell_type":"markdown","source":"## Exploratory Data Analysis","metadata":{}},{"cell_type":"code","source":"import pandas as pd \n\n# Path to the CSV file\nrsna_path = '/kaggle/input/rsna-breast-cancer-detection/train.csv'\n\n# Read the CSV file into a DataFrame\ndf_rsna = pd.read_csv(rsna_path)\n\ndf_rsna.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:30.268746Z","iopub.execute_input":"2024-11-05T03:47:30.269143Z","iopub.status.idle":"2024-11-05T03:47:30.366864Z","shell.execute_reply.started":"2024-11-05T03:47:30.269105Z","shell.execute_reply":"2024-11-05T03:47:30.365751Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Display basic info\ndf_rsna.info()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:30.68371Z","iopub.execute_input":"2024-11-05T03:47:30.684311Z","iopub.status.idle":"2024-11-05T03:47:30.710126Z","shell.execute_reply.started":"2024-11-05T03:47:30.684272Z","shell.execute_reply":"2024-11-05T03:47:30.709078Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Missing Values Analysis\n","metadata":{}},{"cell_type":"code","source":"# Checking missing values in each column\nmissing_values = df_rsna.isnull().sum()\nprint(\"Missing Values in each column:\\n\", missing_values)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:31.37375Z","iopub.execute_input":"2024-11-05T03:47:31.374553Z","iopub.status.idle":"2024-11-05T03:47:31.394912Z","shell.execute_reply.started":"2024-11-05T03:47:31.374513Z","shell.execute_reply":"2024-11-05T03:47:31.39395Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Summary statistics for both numerical and categorical columns\nsummary_statistics = df_rsna.describe(include='all')\nprint(\"Summary Statistics:\\n\", summary_statistics)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:31.572837Z","iopub.execute_input":"2024-11-05T03:47:31.573166Z","iopub.status.idle":"2024-11-05T03:47:31.649541Z","shell.execute_reply.started":"2024-11-05T03:47:31.573131Z","shell.execute_reply":"2024-11-05T03:47:31.64868Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\nimport numpy as np\n\n# Set a high-quality style for the plot\nsns.set(style=\"whitegrid\")\n\n# Plotting the distribution of 'age' with advanced styling\nplt.figure(figsize=(14, 8))\nn, bins, patches = plt.hist(df_rsna['age'].dropna(), bins=20, edgecolor='black', color='lightpink', alpha=1.0)\n\n# Title and labels\nplt.title('Age Distribution in Dataset', fontsize=20, fontweight='bold', color='midnightblue', pad=20)\nplt.xlabel('Age', fontsize=16, labelpad=10)\nplt.ylabel('Frequency', fontsize=16, labelpad=10)\n\n# Apply a pastel color gradient to the bars\nbin_centers = 0.5 * (bins[:-1] + bins[1:])\ncolormap = plt.cm.Pastel1  # Choose a pastel color map\nfor count, patch in zip(n, patches):\n    plt.setp(patch, 'facecolor', colormap((count - np.min(n)) / np.ptp(n)))\n\n# Adding a grid and text annotations for visual appeal\nplt.grid(axis='y', linestyle='--', alpha=0.5)\n\n# Add annotations on top of each bar\nfor count, x in zip(n, bin_centers):\n    plt.text(x, count + 2, f'{int(count)}', ha='center', va='bottom', fontsize=10, color='dimgray')\n\n# Customizing tick parameters\nplt.xticks(fontsize=12, color='slategray')\nplt.yticks(fontsize=12, color='slategray')\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:31.711039Z","iopub.execute_input":"2024-11-05T03:47:31.711552Z","iopub.status.idle":"2024-11-05T03:47:32.363879Z","shell.execute_reply.started":"2024-11-05T03:47:31.711519Z","shell.execute_reply":"2024-11-05T03:47:32.362765Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Set a high-quality style with pastel tones\nsns.set(style=\"whitegrid\")\n\n# Plotting the distribution of 'density' with enhancements\nplt.figure(figsize=(12, 8))\ndensity_counts = df_rsna['density'].value_counts()\nbars = density_counts.plot(kind='bar', edgecolor='black', color='coral', alpha=1.0)\n\n# Title and labels\nplt.title('Distribution of Density in Dataset', fontsize=20, fontweight='bold', color='midnightblue', pad=20)\nplt.xlabel('Density', fontsize=16, labelpad=10)\nplt.ylabel('Frequency', fontsize=16, labelpad=10)\n\n# Add pastel colors to the bars\ncolormap = plt.cm.Pastel2  # Pastel colormap\nfor i, bar in enumerate(bars.containers[0]):  # Apply color gradient\n    bar.set_color(colormap(i / len(density_counts)))\n\n# Adding grid lines for better readability\nplt.grid(axis='y', linestyle='--', alpha=0.5)\n\n# Adding annotations on top of each bar\nfor i, (density, count) in enumerate(density_counts.items()):\n    plt.text(i, count + 20, f'{count}', ha='center', va='bottom', fontsize=12, color='dimgray')\n\n# Customizing tick parameters\nplt.xticks(fontsize=14, color='slategray')\nplt.yticks(fontsize=14, color='slategray')\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:32.365871Z","iopub.execute_input":"2024-11-05T03:47:32.366314Z","iopub.status.idle":"2024-11-05T03:47:32.801799Z","shell.execute_reply.started":"2024-11-05T03:47:32.366266Z","shell.execute_reply":"2024-11-05T03:47:32.800747Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Take Away from dataset Exploration: \n1. **Missing Values Analysis:** Remove rows where density is missing\n2. **Remove Implants:** Filter out rows where implants == 1\n3. **Dataset Imbalance:** A and D have least amount of images\n4. **Dataset Split:** Split dataset into train, val, test\n5. **Label Encoding:** Encode categorical labels into numeric e.g. A=0, B=1, C=2, D=3\n6. **Data Augmentation and Batch Generation:** Handling image data and feed it into deep learning model","metadata":{}},{"cell_type":"markdown","source":"## Preprocessing\n","metadata":{}},{"cell_type":"code","source":"# Count the total number of rows before removing missing density values\ntotal_rows_before = len(df_rsna)\nprint(\"Total rows before removing missing density values:\", total_rows_before)\n\n# Step 1: Drop rows with missing target variable (density)\ndf_rsna_preprocessed = df_rsna.dropna(subset=['density'])\n\n# Count the total number of rows after removing missing density values\ntotal_rows_after = len(df_rsna_preprocessed)\nprint(\"Total rows after removing missing density values:\", total_rows_after)\n\n# Display the first few rows of the DataFrame after removal\ndf_rsna_preprocessed.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:32.80382Z","iopub.execute_input":"2024-11-05T03:47:32.804186Z","iopub.status.idle":"2024-11-05T03:47:32.838291Z","shell.execute_reply.started":"2024-11-05T03:47:32.804146Z","shell.execute_reply":"2024-11-05T03:47:32.837309Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Count the total number of rows before removal\ntotal_rows_before = len(df_rsna_preprocessed)\nprint(\"Total rows before removing implants:\", total_rows_before)\n\n# Count the number of rows with implants\nimplant_count = df_rsna_preprocessed[df_rsna_preprocessed['implant'] == 1].shape[0]\nprint(\"Number of rows with implants (implant == 1):\", implant_count)\n\n# Remove rows with implants\ndf_rsna_preprocessed = df_rsna_preprocessed[df_rsna_preprocessed['implant'] == 0]\n\n# Count the total number of rows after removal\ntotal_rows_after = len(df_rsna_preprocessed)\nprint(\"Total rows after removing implants:\", total_rows_after)\n\n# Display the first few rows to confirm changes\ndf_rsna_preprocessed.head()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:32.839557Z","iopub.execute_input":"2024-11-05T03:47:32.839949Z","iopub.status.idle":"2024-11-05T03:47:32.869758Z","shell.execute_reply.started":"2024-11-05T03:47:32.839906Z","shell.execute_reply":"2024-11-05T03:47:32.868872Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.preprocessing import LabelEncoder\n\n# Step 3: Encode Density (if density is categorical)\n# Assuming density categories are \"A\", \"B\", \"C\", etc.\nlabel_encoder = LabelEncoder()\ndf_rsna_preprocessed['density_encoded'] = label_encoder.fit_transform(df_rsna_preprocessed['density'])\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:33.066043Z","iopub.execute_input":"2024-11-05T03:47:33.066729Z","iopub.status.idle":"2024-11-05T03:47:33.082708Z","shell.execute_reply.started":"2024-11-05T03:47:33.066669Z","shell.execute_reply":"2024-11-05T03:47:33.081754Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"df_rsna_preprocessed.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:33.487465Z","iopub.execute_input":"2024-11-05T03:47:33.487852Z","iopub.status.idle":"2024-11-05T03:47:33.505119Z","shell.execute_reply.started":"2024-11-05T03:47:33.487813Z","shell.execute_reply":"2024-11-05T03:47:33.504068Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Saving df_rsna_preprocessed to a CSV file\noutput_path = '/kaggle/working/df_rsna_preprocessed.csv'\ndf_rsna_preprocessed.to_csv(output_path, index=False)\nprint(\"File saved to:\", output_path)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:34.040438Z","iopub.execute_input":"2024-11-05T03:47:34.041245Z","iopub.status.idle":"2024-11-05T03:47:34.232579Z","shell.execute_reply.started":"2024-11-05T03:47:34.041198Z","shell.execute_reply":"2024-11-05T03:47:34.231521Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport matplotlib.image as mpimg\n\n# Specify the patient ID to display\npatient_id_to_display = 10095  \n\n# Filter for the specific patient and required views\npatient_data = df_rsna_preprocessed[(df_rsna_preprocessed['patient_id'] == patient_id_to_display) &\n                                    (((df_rsna_preprocessed['laterality'] == 'R') & (df_rsna_preprocessed['view'] == 'CC')) |\n                                     ((df_rsna_preprocessed['laterality'] == 'L') & (df_rsna_preprocessed['view'] == 'CC')) |\n                                     ((df_rsna_preprocessed['laterality'] == 'R') & (df_rsna_preprocessed['view'] == 'MLO')) |\n                                     ((df_rsna_preprocessed['laterality'] == 'L') & (df_rsna_preprocessed['view'] == 'MLO')))]\n\n# Select the first image for each unique combination of `laterality` and `view`\npatient_data_unique_views = patient_data.drop_duplicates(subset=['laterality', 'view'])\n\n# Check if we have exactly 4 images (one for each combination of view and laterality)\nif len(patient_data_unique_views) == 4:\n    fig, axes = plt.subplots(2, 2, figsize=(12, 12))\n    fig.suptitle(f\"Images for Patient ID: {patient_id_to_display}\", fontsize=16)\n    \n    for i, (idx, row) in enumerate(patient_data_unique_views.iterrows()):\n        # Load image based on image_id\n        image_path = f\"/kaggle/input/rsna-breast-cancer-512-pngs/{row['patient_id']}_{row['image_id']}.png\"  # Update path and extension as needed\n        img = mpimg.imread(image_path)\n        \n        # Position in 2x2 grid\n        ax = axes[i // 2, i % 2]\n        \n        # Display image and metadata\n        ax.imshow(img, cmap='gray')\n        ax.axis('off')\n        ax.set_title(f\"Image ID: {row['image_id']}\\n\"\n                     f\"Laterality: {row['laterality']}\\n\"\n                     f\"View: {row['view']}\\n\"\n                     f\"Density: {row['density']}\", fontsize=12)\n\n    plt.tight_layout(rect=[0, 0, 1, 0.96])  # Adjust layout to fit title\n    plt.show()\nelse:\n    print(\"Could not find exactly 4 unique images (one per view) for this patient.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:34.556392Z","iopub.execute_input":"2024-11-05T03:47:34.557225Z","iopub.status.idle":"2024-11-05T03:47:35.94867Z","shell.execute_reply.started":"2024-11-05T03:47:34.557184Z","shell.execute_reply":"2024-11-05T03:47:35.947745Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### Dataset Splitting","metadata":{}},{"cell_type":"code","source":"import matplotlib.pyplot as plt\nimport seaborn as sns\n\n# Set a high-quality style with pastel tones\nsns.set(style=\"whitegrid\")\n\n# Plotting the distribution of 'density' categories with enhancements\nplt.figure(figsize=(10, 7))\ndensity_counts = df_rsna_preprocessed['density'].value_counts()\nbars = density_counts.plot(kind='bar', color='lightsteelblue', edgecolor='black', alpha=0.85)\n\n# Title and labels\nplt.title('Distribution of Density Categories', fontsize=20, fontweight='bold', color='navy', pad=20)\nplt.xlabel('Density Category', fontsize=16, labelpad=10)\nplt.ylabel('Frequency', fontsize=16, labelpad=10)\n\n# Adding a pastel color gradient to the bars\ncolormap = plt.cm.Pastel2\nfor i, bar in enumerate(bars.containers[0]):\n    bar.set_color(colormap(i / len(density_counts)))\n\n# Adding a dashed grid for readability\nplt.grid(axis='y', linestyle='--', alpha=0.6)\n\n# Adding annotations on top of each bar for counts\nfor i, (density, count) in enumerate(density_counts.items()):\n    plt.text(i, count + 5, f'{count}', ha='center', va='bottom', fontsize=12, color='navy')\n\n# Customizing ticks\nplt.xticks(fontsize=14, color='slategray', rotation=0)\nplt.yticks\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:35.968533Z","iopub.execute_input":"2024-11-05T03:47:35.968874Z","iopub.status.idle":"2024-11-05T03:47:36.35145Z","shell.execute_reply.started":"2024-11-05T03:47:35.96884Z","shell.execute_reply":"2024-11-05T03:47:36.350537Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Display the count of images in each 'density' category\nimage_count_per_density = df_rsna_preprocessed['density'].value_counts()\n\n# Print the count of images per density category\nprint(\"Count of images in each density category:\")\nprint(image_count_per_density)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:36.627722Z","iopub.execute_input":"2024-11-05T03:47:36.628083Z","iopub.status.idle":"2024-11-05T03:47:36.638091Z","shell.execute_reply.started":"2024-11-05T03:47:36.628048Z","shell.execute_reply":"2024-11-05T03:47:36.637109Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import pandas as pd\nfrom sklearn.model_selection import train_test_split\n\n# Load the dataset (if not already loaded)\ndf = df_rsna_preprocessed  # Assuming df_rsna_preprocessed is your main DataFrame\n\n# Step 1: Group by patient_id and sample to ensure unique patients in each set\n# We'll use `density` for stratification since that’s our target\n\n# Extract unique patients with their associated densities (majority class per patient)\npatient_data = df.groupby('patient_id')['density'].agg(lambda x: x.mode()[0]).reset_index()\n\n# Step 2: Perform stratified sampling to split patients into train, val, and test\n\n# First, split into 80% training and 20% temp (which will later be split into val and test)\ntrain_patients, temp_patients = train_test_split(\n    patient_data, test_size=0.2, stratify=patient_data['density'], random_state=42\n)\n\n# Next, split the temp set into 50% validation and 50% test (which is 10% each of the original dataset)\nval_patients, test_patients = train_test_split(\n    temp_patients, test_size=0.5, stratify=temp_patients['density'], random_state=42\n)\n\n# Step 3: Map the patient splits back to the main DataFrame\n\n# Filter the main DataFrame for each set\ntrain_df = df[df['patient_id'].isin(train_patients['patient_id'])]\nval_df = df[df['patient_id'].isin(val_patients['patient_id'])]\ntest_df = df[df['patient_id'].isin(test_patients['patient_id'])]\n\n# Display the sizes of each split\nprint(\"Training set size:\", len(train_df))\nprint(\"Validation set size:\", len(val_df))\nprint(\"Test set size:\", len(test_df))\n\n# Optionally, save these sets to CSVs\ntrain_df.to_csv('/kaggle/working/train_df.csv', index=False)\nval_df.to_csv('/kaggle/working/val_df.csv', index=False)\ntest_df.to_csv('/kaggle/working/test_df.csv', index=False)\n\nprint(\"Datasets saved successfully.\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:47:39.274135Z","iopub.execute_input":"2024-11-05T03:47:39.274773Z","iopub.status.idle":"2024-11-05T03:47:40.230937Z","shell.execute_reply.started":"2024-11-05T03:47:39.274719Z","shell.execute_reply":"2024-11-05T03:47:40.229995Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Category-wise image counts in each set\ntrain_category_counts = train_df['density'].value_counts()\nval_category_counts = val_df['density'].value_counts()\ntest_category_counts = test_df['density'].value_counts()\n\n# Display category-wise counts\nprint(\"Category-wise image counts in each set:\")\nprint(\"Training set:\\n\", train_category_counts)\nprint(\"\\nValidation set:\\n\", val_category_counts)\nprint(\"\\nTest set:\\n\", test_category_counts)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-04T13:15:08.291328Z","iopub.execute_input":"2024-11-04T13:15:08.292024Z","iopub.status.idle":"2024-11-04T13:15:08.307266Z","shell.execute_reply.started":"2024-11-04T13:15:08.291978Z","shell.execute_reply":"2024-11-04T13:15:08.306137Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nfrom torch.utils.data import Dataset\nfrom torchvision import transforms\nfrom PIL import Image\nimport os\nimport pandas as pd\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nimport cv2\n\nclass BreastDensityDataset(Dataset):\n    def __init__(self, csv_file, img_dir, transform=None):\n        \"\"\"\n        Args:\n            csv_file (string): Path to the CSV file with annotations.\n            img_dir (string): Directory with all the images.\n            transform (callable, optional): Optional transform to be applied on an image.\n        \"\"\"\n        self.data = pd.read_csv(csv_file)\n        self.img_dir = img_dir\n        self.transform = transform\n\n    def __len__(self):\n        return len(self.data)\n\n    def __getitem__(self, idx):\n        # Get patient_id, image_id, and density label\n        patient_id = self.data.iloc[idx]['patient_id']\n        image_id = self.data.iloc[idx]['image_id']\n        label = self.data.iloc[idx]['density_encoded']  # Using the existing encoded label directly\n        \n        # Construct the file path with {patient_id}_{image_id}.png\n        img_path = os.path.join(self.img_dir, f\"{patient_id}_{image_id}.png\")\n        \n        # Load the image\n        image = cv2.imread(img_path)\n        image = cv2.cvtColor(image, cv2.COLOR_BGR2RGB)  # Convert to RGB format\n        \n        # Apply transformations if specified\n        if self.transform:\n            # Apply Albumentations transform, which requires a dictionary\n            image = self.transform(image=image)[\"image\"]\n        \n        return image, label, patient_id, image_id","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:56:25.668557Z","iopub.execute_input":"2024-11-05T03:56:25.668953Z","iopub.status.idle":"2024-11-05T03:56:31.39046Z","shell.execute_reply.started":"2024-11-05T03:56:25.668917Z","shell.execute_reply":"2024-11-05T03:56:31.389541Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Directory containing images\nimg_dir = \"/kaggle/input/rsna-breast-cancer-512-pngs\"  # Update this path to your image folder\n\n# Paths to CSV files for each dataset split\ntrain_csv = '/kaggle/working/train_df.csv'\nval_csv = '/kaggle/working/val_df.csv'\ntest_csv = '/kaggle/working/test_df.csv'","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:56:31.392166Z","iopub.execute_input":"2024-11-05T03:56:31.392666Z","iopub.status.idle":"2024-11-05T03:56:31.397157Z","shell.execute_reply.started":"2024-11-05T03:56:31.392628Z","shell.execute_reply":"2024-11-05T03:56:31.396224Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Checking for duplicates in each of the train, validation, and test sets based on 'patient_id' and 'image_id' combination\n\ntrain_duplicates = train_df.duplicated(subset=['patient_id', 'image_id']).sum()\nval_duplicates = val_df.duplicated(subset=['patient_id', 'image_id']).sum()\ntest_duplicates = test_df.duplicated(subset=['patient_id', 'image_id']).sum()\n\n# Displaying the results\nprint(\"Number of duplicates in Training set:\", train_duplicates)\nprint(\"Number of duplicates in Validation set:\", val_duplicates)\nprint(\"Number of duplicates in Test set:\", test_duplicates)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:56:31.398188Z","iopub.execute_input":"2024-11-05T03:56:31.398507Z","iopub.status.idle":"2024-11-05T03:56:31.414621Z","shell.execute_reply.started":"2024-11-05T03:56:31.398473Z","shell.execute_reply":"2024-11-05T03:56:31.41366Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom torchvision import transforms\nimport cv2\n\n# Albumentations transformations for training\ntrain_transform = A.Compose([\n    A.Resize(224, 224),  # Resize to 299x299 for InceptionV3\n    A.HorizontalFlip(p=0.5),\n    A.VerticalFlip(p=0.5),\n    A.RandomBrightnessContrast(p=0.5),\n    A.ShiftScaleRotate(shift_limit=0.2, scale_limit=0.2, rotate_limit=30, p=0.5),\n    A.MedianBlur(blur_limit=5, p=0.1),\n    A.GaussianBlur(blur_limit=(3, 7), p=0.1),\n    A.GaussNoise(p=0.2),\n    A.ElasticTransform(p=0.1),\n    A.GridDistortion(p=0.1),\n    A.OpticalDistortion(p=0.1),\n    A.CoarseDropout(max_holes=8, max_height=32, max_width=32, p=0.5),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n])\n\n# Albumentations transformations for validation and test (without augmentation)\nval_test_transform = A.Compose([\n    A.Resize(224, 224),\n    A.Normalize(mean=(0.485, 0.456, 0.406), std=(0.229, 0.224, 0.225)),\n    ToTensorV2()\n])\n\n\n\n# Create dataset instances\ntrain_dataset = BreastDensityDataset(csv_file=train_csv, img_dir=img_dir, transform=train_transform)\nval_dataset = BreastDensityDataset(csv_file=val_csv, img_dir=img_dir, transform=val_test_transform)\ntest_dataset = BreastDensityDataset(csv_file=test_csv, img_dir=img_dir, transform=val_test_transform)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T03:56:32.274495Z","iopub.execute_input":"2024-11-05T03:56:32.275402Z","iopub.status.idle":"2024-11-05T03:56:32.350133Z","shell.execute_reply.started":"2024-11-05T03:56:32.275351Z","shell.execute_reply":"2024-11-05T03:56:32.349024Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import matplotlib.pyplot as plt\n\n# Function to display a few images from the dataset with {patient_id}_{image_id} and density\ndef display_sample_images(dataset, num_images=4):\n    plt.figure(figsize=(10, 10))\n    \n    for i in range(num_images):\n        # Get a sample from the dataset\n        image, label, patient_id, image_id = dataset[i]\n        \n        # Convert to grayscale by taking only one channel, e.g., the first channel\n        image = image[0]  # Assuming the image is (3, H, W), take the first channel\n        \n        # Plot the image in grayscale\n        ax = plt.subplot(1, num_images, i + 1)\n        plt.imshow(image, cmap='gray')  # Use 'gray' colormap for grayscale display\n        plt.axis(\"off\")\n        ax.set_title(f\"{patient_id}_{image_id}\\nDensity: {label}\")  # Adding density label\n\n    plt.tight_layout()\n    plt.show()\n\n# Display sample images from test_dataset\ndisplay_sample_images(test_dataset, num_images=4)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T04:24:04.80022Z","iopub.execute_input":"2024-11-05T04:24:04.800605Z","iopub.status.idle":"2024-11-05T04:24:05.84676Z","shell.execute_reply.started":"2024-11-05T04:24:04.800568Z","shell.execute_reply":"2024-11-05T04:24:05.845751Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"## **DeiT (Data-efficient Image Transformers) - 2020**\n","metadata":{}},{"cell_type":"code","source":"pip install torchsummary\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T04:24:14.475878Z","iopub.execute_input":"2024-11-05T04:24:14.476262Z","iopub.status.idle":"2024-11-05T04:24:26.258129Z","shell.execute_reply.started":"2024-11-05T04:24:14.476224Z","shell.execute_reply":"2024-11-05T04:24:26.256799Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport timm\nimport torch.nn as nn\nfrom torchsummary import summary\n\n# Initialize the model\nmodel = timm.create_model('deit_base_patch16_224', pretrained=True)\n\n# Adjust the final layer for your specific number of classes\nnum_classes = 4  # Modify as needed for your dataset\nin_features = model.get_classifier().in_features\nmodel.head = nn.Linear(in_features, num_classes)\n\n# Move model to GPU if available and wrap with DataParallel for multiple GPUs\ndevice = torch.device(\"cuda\" if torch.cuda.is_available() else \"cpu\")\nif torch.cuda.device_count() > 1:\n    print(f\"Using {torch.cuda.device_count()} GPUs!\")\n    model = nn.DataParallel(model)  # This will parallelize across multiple GPUs\nmodel = model.to(device)\n\n\n# Calculate total and trainable parameters\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\n# Print model architecture and parameter details\n#print(f\"Model Architecture:\\n{model}\")\nprint(f\"\\nTotal Parameters: {total_params:,}\")\nprint(f\"Trainable Parameters: {trainable_params:,}\")\n\n# Calculate model size in MB\nparam_size = 4  # Size of each parameter in bytes (float32 takes 4 bytes)\nmodel_size_mb = total_params * param_size / (1024 ** 2)\nprint(f\"Model Size: {model_size_mb:.2f} MB\")\n\n# Print a detailed summary of the model layers (optional, may be long)\n#print(\"\\nDetailed Model Summary:\")\n#summary(model, (3, 224, 224))\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T04:00:21.269917Z","iopub.execute_input":"2024-11-05T04:00:21.270249Z","iopub.status.idle":"2024-11-05T04:00:25.681612Z","shell.execute_reply.started":"2024-11-05T04:00:21.270214Z","shell.execute_reply":"2024-11-05T04:00:25.680649Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"### **Training Loop**","metadata":{}},{"cell_type":"code","source":"import time\nimport torch\nimport timm\nimport torch.nn as nn\nfrom torch.optim import AdamW\nfrom torch.utils.data import DataLoader\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nimport matplotlib.pyplot as plt\nimport numpy as np\nfrom tqdm import tqdm\n\n# Hyperparameters\nbatch_size = 32\nlearning_rate = 1e-4\nepochs = 50  # Increased to allow early stopping\npatience = 10  # Early stopping patience\n\n# Define data loaders\ntrain_loader = DataLoader(train_dataset, batch_size=batch_size, shuffle=True, num_workers=4)\nval_loader = DataLoader(val_dataset, batch_size=batch_size, shuffle=False, num_workers=4)\n\n# Define optimizer and adaptive learning rate scheduler\noptimizer = AdamW(model.parameters(), lr=learning_rate, weight_decay=1e-4)\nscheduler = ReduceLROnPlateau(optimizer, mode='min', factor=0.5, patience=3, verbose=True)\n\n# Loss function\ncriterion = nn.CrossEntropyLoss()\n\n# Training tracking variables\nbest_val_accuracy = 0.0\nbest_model_weights = None\ntrain_losses, val_losses, train_accuracies, val_accuracies = [], [], [], []\nno_improve_epochs = 0  # Track epochs without improvement for early stopping\n\n# Track total training time\nstart_time = time.time()\n\n# Training and validation loops\nfor epoch in range(epochs):\n    epoch_start_time = time.time()\n\n    # Training loop with progress bar\n    model.train()\n    total_train_loss = 0.0\n    correct_train = 0\n    train_bar = tqdm(train_loader, desc=f\"Epoch {epoch+1}/{epochs} - Training\", leave=False)\n    for images, labels, _, _ in train_bar:\n        images, labels = images.to(device), labels.to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        loss = criterion(outputs, labels)\n        loss.backward()\n        optimizer.step()\n        \n        total_train_loss += loss.item()\n        _, preds = torch.max(outputs, 1)\n        correct_train += (preds == labels).sum().item()  # Calculate train accuracy\n        train_bar.set_postfix({\"Batch Loss\": loss.item()})\n\n    # Calculate training loss and accuracy\n    train_loss = total_train_loss / len(train_loader)\n    train_accuracy = correct_train / len(train_loader.dataset)\n    train_losses.append(train_loss)\n    train_accuracies.append(train_accuracy)\n\n    # Validation loop with progress bar\n    model.eval()\n    total_val_loss = 0.0\n    correct_val = 0\n    val_bar = tqdm(val_loader, desc=f\"Epoch {epoch+1}/{epochs} - Validation\", leave=False)\n    with torch.no_grad():\n        for images, labels, _, _ in val_bar:\n            images, labels = images.to(device), labels.to(device)\n            outputs = model(images)\n            loss = criterion(outputs, labels)\n            total_val_loss += loss.item()\n\n            _, preds = torch.max(outputs, 1)\n            correct_val += (preds == labels).sum().item()\n            val_bar.set_postfix({\"Batch Loss\": loss.item()})\n\n    # Calculate validation loss and accuracy\n    val_loss = total_val_loss / len(val_loader)\n    val_accuracy = correct_val / len(val_loader.dataset)\n    val_losses.append(val_loss)\n    val_accuracies.append(val_accuracy)\n    \n    # Adaptive learning rate scheduler step\n    scheduler.step(val_loss)\n\n    # Print statistics for the epoch\n    print(f\"Epoch {epoch+1}/{epochs}, Train Loss: {train_loss:.4f}, Train Acc: {train_accuracy:.4f}, Val Loss: {val_loss:.4f}, Val Acc: {val_accuracy:.4f}\")\n\n    # Check if current model is the best\n    if val_accuracy > best_val_accuracy:\n        best_val_accuracy = val_accuracy\n        best_model_weights = model.state_dict()  # Save best model weights\n        no_improve_epochs = 0  # Reset early stopping counter\n        print(\"Best model updated.\")\n    else:\n        no_improve_epochs += 1  # Increase early stopping counter\n\n    # Early stopping check\n    if no_improve_epochs >= patience:\n        print(\"Early stopping triggered.\")\n        break\n\n    # Track time per epoch\n    epoch_time = time.time() - epoch_start_time\n    print(f\"Epoch {epoch+1} completed in {epoch_time:.2f} seconds.\")\n\n# Calculate total training time\ntotal_time = time.time() - start_time\nprint(f\"\\nTotal Training Time: {total_time // 60:.0f} minutes, {total_time % 60:.2f} seconds.\")\n\n# Load best model weights\nmodel.load_state_dict(best_model_weights)\n\n# Save best model weights\ntorch.save(best_model_weights, 'best_deit_model.pth')\nprint(\"Best model weights saved to 'best_deit_model.pth'.\")\n\n# Plot training and validation curves\nepochs_range = range(1, len(train_losses) + 1)\nplt.figure(figsize=(14, 6))\n\n# Plot Loss\nplt.subplot(1, 2, 1)\nplt.plot(epochs_range, train_losses, label='Train Loss')\nplt.plot(epochs_range, val_losses, label='Val Loss')\nplt.xlabel('Epochs')\nplt.ylabel('Loss')\nplt.legend()\nplt.title('Training and Validation Loss')\n\n# Plot Accuracy\nplt.subplot(1, 2, 2)\nplt.plot(epochs_range, train_accuracies, label='Train Accuracy')\nplt.plot(epochs_range, val_accuracies, label='Val Accuracy')\nplt.xlabel('Epochs')\nplt.ylabel('Accuracy')\nplt.legend()\nplt.title('Training and Validation Accuracy')\n\nplt.tight_layout()\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T04:31:09.133511Z","iopub.execute_input":"2024-11-05T04:31:09.134268Z","iopub.status.idle":"2024-11-05T06:39:14.773236Z","shell.execute_reply.started":"2024-11-05T04:31:09.134227Z","shell.execute_reply":"2024-11-05T06:39:14.772189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"test_loader = DataLoader(test_dataset, batch_size=batch_size, shuffle=False, num_workers=4)\nmodel.eval()\ncorrect = 0\nwith torch.no_grad():\n    for images, labels, _, _ in test_loader:\n        images, labels = images.to(device), labels.to(device)\n        outputs = model(images)\n        _, preds = torch.max(outputs, 1)\n        correct += (preds == labels).sum().item()\n\ntest_accuracy = correct / len(test_loader.dataset)\nprint(f\"Test Accuracy: {test_accuracy:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T06:39:14.775302Z","iopub.execute_input":"2024-11-05T06:39:14.775639Z","iopub.status.idle":"2024-11-05T06:39:34.939097Z","shell.execute_reply.started":"2024-11-05T06:39:14.775602Z","shell.execute_reply":"2024-11-05T06:39:34.93797Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"from sklearn.metrics import classification_report, confusion_matrix, ConfusionMatrixDisplay\nimport matplotlib.pyplot as plt\n\n# After training, evaluate on the test or validation set\nmodel.eval()\nall_preds = []\nall_labels = []\n\nwith torch.no_grad():\n    for images, labels, _, _ in val_loader:  # Use validation or test loader\n        images, labels = images.to(device), labels.to(device)\n        outputs = model(images)\n        _, preds = torch.max(outputs, 1)\n        \n        all_preds.extend(preds.cpu().numpy())\n        all_labels.extend(labels.cpu().numpy())\n\n# Generate Classification Report\nprint(\"Classification Report:\")\nprint(classification_report(all_labels, all_preds, target_names=[f\"Class {i}\" for i in range(num_classes)]))\n\n# Compute Sensitivity (Recall) for Each Class\n# Sensitivity is already included in the classification report as recall.\n\n# Generate and Plot Confusion Matrix\nconf_matrix = confusion_matrix(all_labels, all_preds)\ndisp = ConfusionMatrixDisplay(confusion_matrix=conf_matrix, display_labels=[f\"Class {i}\" for i in range(num_classes)])\ndisp.plot(cmap=plt.cm.Blues)\nplt.title(\"Confusion Matrix\")\nplt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T06:59:58.549065Z","iopub.execute_input":"2024-11-05T06:59:58.549807Z","iopub.status.idle":"2024-11-05T07:00:19.343355Z","shell.execute_reply.started":"2024-11-05T06:59:58.549763Z","shell.execute_reply":"2024-11-05T07:00:19.342329Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Assuming you have the necessary imports\nimport random\nimport torch.nn.functional as F  # For softmax function\n\ndef display_predictions_random_patients(dataset, model, num_images=4, num_patients=2):\n    model.eval()\n    plt.figure(figsize=(10, 10))\n    \n    # Get unique patient IDs from the dataset\n    unique_patients = list(set([dataset[i][2] for i in range(len(dataset))]))  # Assuming patient_id is the 3rd item in the tuple\n    selected_patients = random.sample(unique_patients, num_patients)  # Randomly select patient IDs\n    \n    # Initialize image index\n    img_count = 0\n\n    for patient_id in selected_patients:\n        # Get images for the selected patient\n        patient_images = [i for i in range(len(dataset)) if dataset[i][2] == patient_id]\n        \n        for idx in patient_images:\n            if img_count >= num_images:\n                break  # Stop if we have reached the desired number of images\n            \n            # Get image and label from dataset\n            image, label, _, image_id = dataset[idx]\n            \n            # Move image to GPU and add batch dimension\n            with torch.no_grad():\n                image = image.unsqueeze(0).to(device)\n                output = model(image)\n                \n                # Calculate confidence score\n                probabilities = F.softmax(output, dim=1)  # Apply softmax to get probabilities\n                confidence, pred = torch.max(probabilities, 1)  # Get predicted class and its confidence score\n            \n            # Rearrange image dimensions for Matplotlib\n            image = image.squeeze().cpu().numpy()  # Remove batch dimension and convert to numpy array\n            if image.ndim == 2:  # If the image is already in HxW format\n                image_to_show = image\n            elif image.shape[0] == 3:  # If the image is in CxHxW format\n                image_to_show = image.transpose(1, 2, 0)  # Convert to HxWxC\n                # Convert to grayscale if needed\n                if image_to_show.shape[2] == 3:\n                    image_to_show = image_to_show.mean(axis=2)  # Convert RGB to grayscale\n\n            # Plotting\n            ax = plt.subplot(1, num_images, img_count + 1)\n            plt.imshow(image_to_show, cmap='gray')  # Force grayscale\n            plt.title(f\"Pred: {pred.item()} (Conf: {confidence.item() * 100:.2f}%)\\nTrue: {label}\")\n            plt.axis(\"off\")\n            \n            img_count += 1\n        \n    plt.show()\n\n# Call function to display predictions for 2 random patients\ndisplay_predictions_random_patients(test_dataset, model, num_images=4, num_patients=2)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-05T07:02:21.419654Z","iopub.execute_input":"2024-11-05T07:02:21.420334Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}