{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":52254,"databundleVersionId":9674523,"sourceType":"competition"}],"dockerImageVersionId":30804,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\nimport os","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:53:06.345234Z","iopub.execute_input":"2024-12-13T03:53:06.345581Z","iopub.status.idle":"2024-12-13T03:53:07.247197Z","shell.execute_reply.started":"2024-12-13T03:53:06.345548Z","shell.execute_reply":"2024-12-13T03:53:07.246518Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"train_df = pd.read_csv(\"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_2024.csv\")\nimage_level_df = pd.read_csv(\"/kaggle/input/rsna-2023-abdominal-trauma-detection/image_level_labels_2024.csv\")\n\ntrain_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:53:09.367169Z","iopub.execute_input":"2024-12-13T03:53:09.367607Z","iopub.status.idle":"2024-12-13T03:53:09.41863Z","shell.execute_reply.started":"2024-12-13T03:53:09.367577Z","shell.execute_reply":"2024-12-13T03:53:09.417784Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"image_level_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:53:12.673503Z","iopub.execute_input":"2024-12-13T03:53:12.674117Z","iopub.status.idle":"2024-12-13T03:53:12.685404Z","shell.execute_reply.started":"2024-12-13T03:53:12.674071Z","shell.execute_reply":"2024-12-13T03:53:12.684232Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# CSV Preprocessing","metadata":{}},{"cell_type":"code","source":"series_meta_df = pd.read_csv(\"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_series_meta.csv\")\nseries_meta_df.head()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:53:16.053097Z","iopub.execute_input":"2024-12-13T03:53:16.053631Z","iopub.status.idle":"2024-12-13T03:53:16.079512Z","shell.execute_reply.started":"2024-12-13T03:53:16.053599Z","shell.execute_reply":"2024-12-13T03:53:16.078712Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess_data():\n    \"\"\"CSV preprocessing function - keeping the exact same code we verified works\"\"\"\n    try:\n        # Load data with corrected column names\n        train_df = pd.read_csv(\"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_2024.csv\")\n        train_df.columns = [\n            'patient_id', 'bowel_healthy', 'bowel_injury', 'extravasation_healthy',\n            'extravasation_injury', 'kidney_healthy', 'kidney_low', 'kidney_high',\n            'liver_healthy', 'liver_low', 'liver_high', 'spleen_healthy',\n            'spleen_low', 'spleen_high', 'any_injury'\n        ]\n\n        image_labels_df = pd.read_csv(\"/kaggle/input/rsna-2023-abdominal-trauma-detection/image_level_labels_2024.csv\")\n        image_labels_df.columns = ['patient_id', 'series_id', 'instance_number', 'injury_name']\n\n        series_meta_df = pd.read_csv(\"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_series_meta.csv\")\n        series_meta_df.columns = ['patient_id', 'series_id', 'aortic_hu', 'incomplete_organ']\n\n        print(\"\\nProcessing series metadata...\")\n        # Filter out incomplete scans\n        complete_series = series_meta_df[~series_meta_df['incomplete_organ'].astype(bool)]\n        print(f\"Complete series: {len(complete_series)} out of {len(series_meta_df)}\")\n\n        # Create enhanced dataset with series information\n        print(\"\\nCreating enhanced dataset...\")\n        enhanced_df = train_df.merge(\n            complete_series[['patient_id', 'series_id', 'aortic_hu']],\n            on='patient_id',\n            how='inner'\n        )\n\n        # Process image-level labels\n        print(\"\\nProcessing image-level labels...\")\n        injury_locations = image_labels_df.groupby(\n            ['patient_id', 'series_id', 'injury_name']\n        )['instance_number'].agg(list).reset_index()\n        \n        print(f\"Patients with labeled injuries: {len(injury_locations['patient_id'].unique())}\")\n\n        # Add injury location information\n        enhanced_df = enhanced_df.merge(\n            injury_locations,\n            on=['patient_id', 'series_id'],\n            how='left'\n        )\n\n        # Print dataset statistics\n        total_patients = len(enhanced_df['patient_id'].unique())\n        print(\"\\nDataset Statistics:\")\n        print(f\"Total patients: {total_patients}\")\n\n        # Binary injuries\n        print(\"\\nBinary Injury Statistics:\")\n        for injury in ['bowel_injury', 'extravasation_injury']:\n            pos_count = enhanced_df[injury].sum()\n            total = len(enhanced_df)\n            print(f\"{injury}: {pos_count} ({pos_count/total*100:.2f}%)\")\n\n        # Multi-class injuries\n        print(\"\\nMulti-level Injury Statistics:\")\n        for organ in ['kidney', 'liver', 'spleen']:\n            print(f\"\\n{organ.capitalize()}:\")\n            healthy = enhanced_df[f'{organ}_healthy'].sum()\n            low = enhanced_df[f'{organ}_low'].sum()\n            high = enhanced_df[f'{organ}_high'].sum()\n            total = len(enhanced_df)\n            print(f\"- Healthy: {healthy} ({healthy/total*100:.2f}%)\")\n            print(f\"- Low-grade injury: {low} ({low/total*100:.2f}%)\")\n            print(f\"- High-grade injury: {high} ({high/total*100:.2f}%)\")\n\n        return enhanced_df, image_labels_df\n\n    except Exception as e:\n        print(f\"Error during preprocessing: {str(e)}\")\n        import traceback\n        traceback.print_exc()\n        return None, None\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:53:19.730625Z","iopub.execute_input":"2024-12-13T03:53:19.731445Z","iopub.status.idle":"2024-12-13T03:53:19.741645Z","shell.execute_reply.started":"2024-12-13T03:53:19.731414Z","shell.execute_reply":"2024-12-13T03:53:19.740799Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    enhanced_df = preprocess_data()\n    if enhanced_df is not None:\n        print(\"\\nPreprocessing completed successfully!\")\n        #print(f\"Final dataset shape: {enhanced_df.shape}\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:53:22.245453Z","iopub.execute_input":"2024-12-13T03:53:22.245874Z","iopub.status.idle":"2024-12-13T03:53:22.304605Z","shell.execute_reply.started":"2024-12-13T03:53:22.245842Z","shell.execute_reply":"2024-12-13T03:53:22.303703Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# EDA Analysis","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport seaborn as sns","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:53:27.712561Z","iopub.execute_input":"2024-12-13T03:53:27.712914Z","iopub.status.idle":"2024-12-13T03:53:28.9396Z","shell.execute_reply.started":"2024-12-13T03:53:27.712885Z","shell.execute_reply":"2024-12-13T03:53:28.938715Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def analyze_injury_distribution(train_df):\n    \"\"\"Analyze and print injury distribution statistics\"\"\"\n    print(\"\\n=== Injury Distribution Analysis ===\")\n    total_patients = len(train_df['patient_id'].unique())\n    print(f\"Total number of patients: {total_patients}\")\n\n    # Binary injuries (bowel and extravasation)\n    print(\"\\nBinary Injury Statistics:\")\n    binary_injuries = ['bowel', 'extravasation']\n    for injury in binary_injuries:\n        injury_count = train_df[f'{injury}_injury'].sum()\n        percentage = (injury_count / len(train_df)) * 100\n        print(f\"{injury.capitalize()} injuries: {injury_count} ({percentage:.2f}%)\")\n\n    # Multi-level injuries (kidney, liver, spleen)\n    print(\"\\nMulti-level Injury Statistics:\")\n    organs = ['kidney', 'liver', 'spleen']\n    for organ in organs:\n        healthy = train_df[f'{organ}_healthy'].sum()\n        low = train_df[f'{organ}_low'].sum()\n        high = train_df[f'{organ}_high'].sum()\n        total = len(train_df)\n        print(f\"\\n{organ.capitalize()}:\")\n        print(f\"- Healthy: {healthy} ({(healthy/total)*100:.2f}%)\")\n        print(f\"- Low-grade injury: {low} ({(low/total)*100:.2f}%)\")\n        print(f\"- High-grade injury: {high} ({(high/total)*100:.2f}%)\")\n\n    # Any injury statistics\n    any_injury_count = train_df['any_injury'].sum()\n    print(f\"\\nPatients with any injury: {any_injury_count} ({(any_injury_count/len(train_df))*100:.2f}%)\")\n\ndef analyze_image_level_injuries(image_df):\n    \"\"\"Analyze image-level injury patterns\"\"\"\n    print(\"\\n=== Image-Level Injury Analysis ===\")\n    \n    # Injury counts by type\n    print(\"Injury instances in images:\")\n    injury_counts = image_df['injury_name'].value_counts()\n    print(injury_counts)\n\n    # Images per patient statistics\n    images_per_patient = image_df.groupby('patient_id').size()\n    print(\"\\nImages per patient statistics:\")\n    print(images_per_patient.describe())\n\n    # Series per patient statistics\n    series_per_patient = image_df.groupby('patient_id')['series_id'].nunique()\n    print(\"\\nSeries per patient statistics:\")\n    print(series_per_patient.describe())\n\ndef create_visualizations(train_df, image_df):\n    \"\"\"Create visualization plots\"\"\"\n    print(\"\\n=== Creating Visualization Plots ===\")\n    \n    # Class balance information\n    print(\"\\n=== Class Balance Information ===\")\n    print(\"Overall injury prevalence:\")\n    \n    # Binary classes\n    print(f\"bowel_healthy: {train_df['bowel_healthy'].sum()} cases ({train_df['bowel_healthy'].mean()*100:.2f}%)\")\n    print(f\"bowel_injury: {train_df['bowel_injury'].sum()} cases ({train_df['bowel_injury'].mean()*100:.2f}%)\")\n    print(f\"extravasation_healthy: {train_df['extravasation_healthy'].sum()} cases ({train_df['extravasation_healthy'].mean()*100:.2f}%)\")\n    print(f\"extravasation_injury: {train_df['extravasation_injury'].sum()} cases ({train_df['extravasation_injury'].mean()*100:.2f}%)\")\n    \n    # Multi-class\n    for organ in ['kidney', 'liver', 'spleen']:\n        print(f\"{organ}_healthy: {train_df[f'{organ}_healthy'].sum()} cases ({train_df[f'{organ}_healthy'].mean()*100:.2f}%)\")\n        print(f\"{organ}_low: {train_df[f'{organ}_low'].sum()} cases ({train_df[f'{organ}_low'].mean()*100:.2f}%)\")\n        print(f\"{organ}_high: {train_df[f'{organ}_high'].sum()} cases ({train_df[f'{organ}_high'].mean()*100:.2f}%)\")\n    \n    # Any injury\n    print(f\"any_injury: {train_df['any_injury'].sum()} cases ({train_df['any_injury'].mean()*100:.2f}%)\")\n\n    # Create correlation heatmap\n    plt.figure(figsize=(12, 8))\n    injury_cols = [col for col in train_df.columns if col not in ['patient_id', 'series_id', 'aortic_hu', 'injury_name', 'instance_number']]\n    correlation_matrix = train_df[injury_cols].corr()\n    sns.heatmap(correlation_matrix, annot=True, cmap='coolwarm', center=0, fmt='.2f')\n    plt.title('Correlation Between Different Injury Types')\n    plt.tight_layout()\n    plt.show()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:53:30.943984Z","iopub.execute_input":"2024-12-13T03:53:30.944466Z","iopub.status.idle":"2024-12-13T03:53:30.959598Z","shell.execute_reply.started":"2024-12-13T03:53:30.94443Z","shell.execute_reply":"2024-12-13T03:53:30.958705Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    # Get enhanced_df and image_df from CSV preprocessing\n    enhanced_df, image_df = preprocess_data()\n    \n    # Run EDA analysis on the DataFrame, not the tuple\n    analyze_injury_distribution(enhanced_df)\n    analyze_image_level_injuries(image_df)\n    create_visualizations(enhanced_df, image_df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:53:33.045072Z","iopub.execute_input":"2024-12-13T03:53:33.045396Z","iopub.status.idle":"2024-12-13T03:53:33.974231Z","shell.execute_reply.started":"2024-12-13T03:53:33.045368Z","shell.execute_reply":"2024-12-13T03:53:33.973351Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Improved Image Visualization","metadata":{}},{"cell_type":"code","source":"import os\nimport pydicom\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom pathlib import Path\n\ndef apply_windowing(img, window_center, window_width):\n    \"\"\"Apply windowing to better visualize different tissues\"\"\"\n    img_min = window_center - window_width // 2\n    img_max = window_center + window_width // 2\n    img = np.clip(img, img_min, img_max)\n    img = ((img - img_min) / (window_width) * 255.0)\n    return np.clip(img, 0, 255).astype('uint8')\n\ndef get_window_parameters(injury_type):\n    \"\"\"Get appropriate window parameters for different injury types\"\"\"\n    if injury_type in ['liver', 'spleen', 'kidney']:\n        return 40, 400  # Soft tissue window\n    elif injury_type == 'extravasation':\n        return 50, 150  # Narrower window for better contrast\n    elif injury_type == 'bowel':\n        return -50, 250  # Window for bowel pathology\n    return 40, 400  # Default abdominal window\n\ndef load_dicom(path, injury_type):\n    \"\"\"Load and preprocess DICOM image with injury-specific windowing\"\"\"\n    try:\n        dicom = pydicom.dcmread(path)\n        img = dicom.pixel_array\n        \n        # Convert to Hounsfield Units (HU)\n        if hasattr(dicom, 'RescaleIntercept') and hasattr(dicom, 'RescaleSlope'):\n            img = img * float(dicom.RescaleSlope) + float(dicom.RescaleIntercept)\n        \n        # Apply specific windowing\n        window_center, window_width = get_window_parameters(injury_type)\n        img = apply_windowing(img, window_center, window_width)\n        return img\n        \n    except Exception as e:\n        print(f\"Error loading {path}: {str(e)}\")\n        return None\n\ndef get_sample_cases(train_df, image_df, base_path, injury_type, severity=None, n_samples=5):\n    \"\"\"Get sample images with specific injury type and severity\"\"\"\n    images = []\n    labels = []\n    \n    try:\n        if injury_type in ['bowel', 'extravasation']:\n            # For binary injuries, use image-level labels\n            if injury_type == 'extravasation':\n                samples = image_df[image_df['injury_name'] == 'Active_Extravasation']\n            else:\n                samples = image_df[image_df['injury_name'] == 'Bowel']\n            \n            samples = samples.head(n_samples)\n            for _, row in samples.iterrows():\n                path = os.path.join(base_path, \n                                  str(row['patient_id']),\n                                  str(row['series_id']), \n                                  f\"{row['instance_number']}.dcm\")\n                if os.path.exists(path):\n                    img = load_dicom(path, injury_type)\n                    if img is not None:\n                        images.append(img)\n                        labels.append(f\"{injury_type}\\nPatient {row['patient_id']}\")\n        else:\n            # For organ injuries with severity levels\n            if severity == 'high':\n                samples = train_df[train_df[f'{injury_type}_high'] == 1]\n            elif severity == 'low':\n                samples = train_df[train_df[f'{injury_type}_low'] == 1]\n            else:\n                samples = train_df[train_df[f'{injury_type}_healthy'] == 1]\n            \n            samples = samples.head(n_samples)\n            for _, row in samples.iterrows():\n                patient_path = os.path.join(base_path, str(row['patient_id']))\n                if os.path.exists(patient_path):\n                    series_folders = os.listdir(patient_path)\n                    if series_folders:\n                        series_path = os.path.join(patient_path, series_folders[0])\n                        dcm_files = sorted([f for f in os.listdir(series_path) if f.endswith('.dcm')])\n                        if dcm_files:\n                            middle_slice = len(dcm_files) // 2\n                            path = os.path.join(series_path, dcm_files[middle_slice])\n                            img = load_dicom(path, injury_type)\n                            if img is not None:\n                                images.append(img)\n                                severity_label = severity if severity else 'healthy'\n                                labels.append(f\"{injury_type} ({severity_label})\\nPatient {row['patient_id']}\")\n                                \n        return images, labels\n        \n    except Exception as e:\n        print(f\"Error processing {injury_type} images: {str(e)}\")\n        return [], []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:53:40.461274Z","iopub.execute_input":"2024-12-13T03:53:40.461575Z","iopub.status.idle":"2024-12-13T03:53:40.668469Z","shell.execute_reply.started":"2024-12-13T03:53:40.46155Z","shell.execute_reply":"2024-12-13T03:53:40.667821Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Set paths\n    base_path = \"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images\"\n    \n    # Load dataframes\n    train_df = pd.read_csv(\"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_2024.csv\")\n    image_df = pd.read_csv(\"/kaggle/input/rsna-2023-abdominal-trauma-detection/image_level_labels_2024.csv\")\n    \n    # Define visualization structure\n    injury_configs = [\n        ('liver', ['low', 'high']),\n        ('kidney', ['low', 'high']),\n        ('spleen', ['low', 'high']),\n        ('extravasation', [None]),\n        ('bowel', [None])\n    ]\n    \n    total_rows = sum(len(severities) for _, severities in injury_configs)\n    \n    # Create figure\n    fig, axes = plt.subplots(total_rows, 5, figsize=(20, 4*total_rows))\n    plt.suptitle(\"Sample Images of Different Injury Types and Severities\", fontsize=16)\n    \n    current_row = 0\n    for injury_type, severities in injury_configs:\n        for severity in severities:\n            print(f\"Processing {injury_type} images{' ('+severity+')' if severity else ''}...\")\n            images, labels = get_sample_cases(train_df, image_df, base_path, injury_type, severity)\n\n            for j in range(5):\n                if j < len(images):\n                    axes[current_row, j].imshow(images[j], cmap='gray')\n                    axes[current_row, j].set_title(labels[j])\n                axes[current_row, j].axis('off')\n            \n            current_row += 1\n    \n    plt.tight_layout()\n    plt.savefig('injury_samples_comprehensive.png')\n    plt.show()\n    \n    # Print summary\n    print(\"\\n=== Visualization Summary ===\")\n    for injury_type, severities in injury_configs:\n        for severity in severities:\n            images, _ = get_sample_cases(train_df, image_df, base_path, injury_type, severity)\n            severity_str = f\" ({severity})\" if severity else \"\"\n            print(f\"{injury_type}{severity_str}: Found {len(images)} sample images\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:53:42.444455Z","iopub.execute_input":"2024-12-13T03:53:42.445246Z","iopub.status.idle":"2024-12-13T03:53:51.133987Z","shell.execute_reply.started":"2024-12-13T03:53:42.445213Z","shell.execute_reply":"2024-12-13T03:53:51.133153Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Each injury type shows distinct imaging characteristics:\n\n* Organ injuries (liver, kidney, spleen) show different patterns based on severity\n* Extravasation cases show active contrast extravasation\n* Bowel injuries have their unique appearance\n\nGood representation across all categories:\n\n* 5 samples for each injury type and severity level\n* Total of 40 images (5 images × 8 categories)\n* Even distribution across severity levels","metadata":{}},{"cell_type":"markdown","source":"# Data Preprocessing","metadata":{}},{"cell_type":"code","source":"import os\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom pathlib import Path\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:00.508376Z","iopub.execute_input":"2024-12-13T03:54:00.509204Z","iopub.status.idle":"2024-12-13T03:54:00.513466Z","shell.execute_reply.started":"2024-12-13T03:54:00.509171Z","shell.execute_reply":"2024-12-13T03:54:00.512693Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def apply_windowing(img, window_center=40, window_width=400):\n    \"\"\"Apply windowing to better visualize different tissues\"\"\"\n    img_min = window_center - window_width // 2\n    img_max = window_center + window_width // 2\n    img = np.clip(img, img_min, img_max)\n    img = ((img - img_min) / (window_width) * 255.0)\n    return np.clip(img, 0, 255).astype('uint8')\n\ndef get_window_parameters(injury_type):\n    \"\"\"Get appropriate window parameters for different injury types\"\"\"\n    if injury_type in ['liver', 'spleen', 'kidney']:\n        return 40, 400  # Soft tissue window\n    elif injury_type == 'extravasation':\n        return 50, 150  # Narrower window for better contrast\n    elif injury_type == 'bowel':\n        return -50, 250  # Window for bowel pathology\n    return 40, 400  # Default abdominal window\n\ndef load_dicom(path, injury_type):\n    \"\"\"Load and preprocess DICOM image with injury-specific windowing\"\"\"\n    try:\n        dicom = pydicom.dcmread(path)\n        img = dicom.pixel_array\n        \n        # Convert to Hounsfield Units (HU)\n        if hasattr(dicom, 'RescaleIntercept') and hasattr(dicom, 'RescaleSlope'):\n            img = img * float(dicom.RescaleSlope) + float(dicom.RescaleIntercept)\n            \n        # Apply specific windowing\n        window_center, window_width = get_window_parameters(injury_type)\n        img = apply_windowing(img, window_center, window_width)\n        \n        return img\n    except Exception as e:\n        print(f\"Error loading {path}: {str(e)}\")\n        return None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:02.717579Z","iopub.execute_input":"2024-12-13T03:54:02.717965Z","iopub.status.idle":"2024-12-13T03:54:02.730312Z","shell.execute_reply.started":"2024-12-13T03:54:02.717936Z","shell.execute_reply":"2024-12-13T03:54:02.729238Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_sample_cases(enhanced_df, image_df, base_path, injury_type, severity=None, n_samples=5):\n    \"\"\"Get sample images with specific injury type and severity\"\"\"\n    images = []\n    labels = []\n    try:\n        if injury_type in ['bowel', 'extravasation']:\n            # For binary injuries, use image-level labels\n            if injury_type == 'extravasation':\n                samples = image_df[image_df['injury_name'] == 'Active_Extravasation']\n            else:\n                samples = image_df[image_df['injury_name'] == 'Bowel']\n            samples = samples.head(n_samples)\n            \n            for _, row in samples.iterrows():\n                path = os.path.join(base_path,\n                                  str(row['patient_id']),\n                                  str(row['series_id']),\n                                  f\"{row['instance_number']}.dcm\")\n                if os.path.exists(path):\n                    img = load_dicom(path, injury_type)\n                    if img is not None:\n                        images.append(img)\n                        labels.append(f\"{injury_type}\\nPatient {row['patient_id']}\")\n        else:\n            # For organ injuries with severity levels\n            if severity == 'high':\n                samples = enhanced_df[enhanced_df[f'{injury_type}_high'] == 1]\n            elif severity == 'low':\n                samples = enhanced_df[enhanced_df[f'{injury_type}_low'] == 1]\n            else:\n                samples = enhanced_df[enhanced_df[f'{injury_type}_healthy'] == 1]\n            \n            samples = samples.head(n_samples)\n            for _, row in samples.iterrows():\n                patient_path = os.path.join(base_path, str(row['patient_id']))\n                if os.path.exists(patient_path):\n                    series_folders = os.listdir(patient_path)\n                    if series_folders:\n                        series_path = os.path.join(patient_path, series_folders[0])\n                        dcm_files = sorted([f for f in os.listdir(series_path) if f.endswith('.dcm')])\n                        if dcm_files:\n                            middle_slice = len(dcm_files) // 2\n                            path = os.path.join(series_path, dcm_files[middle_slice])\n                            img = load_dicom(path, injury_type)\n                            if img is not None:\n                                images.append(img)\n                                severity_label = severity if severity else 'healthy'\n                                labels.append(f\"{injury_type} ({severity_label})\\nPatient {row['patient_id']}\")\n        \n        return images, labels\n    except Exception as e:\n        print(f\"Error processing {injury_type} images: {str(e)}\")\n        return [], []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:05.0156Z","iopub.execute_input":"2024-12-13T03:54:05.016169Z","iopub.status.idle":"2024-12-13T03:54:05.026099Z","shell.execute_reply.started":"2024-12-13T03:54:05.016137Z","shell.execute_reply":"2024-12-13T03:54:05.025153Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    \"\"\"Main function to run both CSV and Data preprocessing\"\"\"\n    print(\"Starting Data Preprocessing...\")\n    \n    # First get enhanced_df from CSV preprocessing\n    enhanced_df, image_df = preprocess_data()\n    if enhanced_df is None:\n        return\n    \n    # Set paths\n    base_path = \"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images\"\n\n    # Define visualization structure\n    injury_configs = [\n        ('liver', ['low', 'high']),\n        ('kidney', ['low', 'high']),\n        ('spleen', ['low', 'high']),\n        ('extravasation', [None]),\n        ('bowel', [None])\n    ]\n\n    # Create figure for visualization\n    total_rows = sum(len(severities) for _, severities in injury_configs)\n    fig, axes = plt.subplots(total_rows, 5, figsize=(20, 4*total_rows))\n    plt.suptitle(\"Sample Images of Different Injury Types and Severities\", fontsize=16)\n\n    current_row = 0\n    for injury_type, severities in injury_configs:\n        for severity in severities:\n            print(f\"Processing {injury_type} images{' ('+severity+')' if severity else ''}...\")\n            images, labels = get_sample_cases(enhanced_df, image_df, base_path, \n                                           injury_type, severity)\n            \n            for j in range(5):\n                if j < len(images):\n                    axes[current_row, j].imshow(images[j], cmap='gray')\n                    axes[current_row, j].set_title(labels[j])\n                    axes[current_row, j].axis('off')\n            current_row += 1\n\n    plt.tight_layout()\n    plt.savefig('injury_samples_visualization.png')\n    plt.show()\n\n    # Print processing summary\n    print(\"\\n=== Preprocessing Summary ===\")\n    for injury_type, severities in injury_configs:\n        for severity in severities:\n            images, _ = get_sample_cases(enhanced_df, image_df, base_path, \n                                       injury_type, severity)\n            severity_str = f\" ({severity})\" if severity else \"\"\n            print(f\"{injury_type}{severity_str}: Processed {len(images)} sample images\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:07.309252Z","iopub.execute_input":"2024-12-13T03:54:07.310173Z","iopub.status.idle":"2024-12-13T03:54:15.114912Z","shell.execute_reply.started":"2024-12-13T03:54:07.310139Z","shell.execute_reply":"2024-12-13T03:54:15.114011Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Model Development Pipeline","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nimport pandas as pd\nimport numpy as np\nimport pydicom\nimport timm\nfrom sklearn.model_selection import train_test_split\nfrom sklearn.metrics import roc_auc_score, confusion_matrix, auc, accuracy_score, roc_curve, classification_report, f1_score\nimport time\nfrom tqdm import tqdm\nimport gc\nimport os\nimport cv2\nfrom torch.utils.data import Sampler\nfrom torch.optim.lr_scheduler import OneCycleLR\nimport torch.nn.functional as F\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom collections import OrderedDict\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:29.919013Z","iopub.execute_input":"2024-12-13T03:54:29.919736Z","iopub.status.idle":"2024-12-13T03:54:29.925222Z","shell.execute_reply.started":"2024-12-13T03:54:29.919701Z","shell.execute_reply":"2024-12-13T03:54:29.92436Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BinaryTraumaModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        self.backbone = timm.create_model('mobilenetv2_100', pretrained=True, in_chans=1)\n        n_features = self.backbone.classifier.in_features\n\n        self.shared = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.BatchNorm1d(n_features),\n            nn.Dropout(0.5),\n            nn.Linear(n_features, 256),\n            nn.ReLU(),\n            nn.BatchNorm1d(256),\n            nn.Dropout(0.3)\n        )\n\n        # Match classifier names with target names\n        self.classifiers = nn.ModuleDict({\n            'bowel_injury': nn.Linear(256, 2),\n            'extravasation_injury': nn.Linear(256, 2),\n            'any_injury': nn.Linear(256, 2)\n        })\n\n        self._initialize_weights()\n\n    def _initialize_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Linear):\n                nn.init.xavier_normal_(m.weight)\n                if m.bias is not None:\n                    nn.init.constant_(m.bias, 0)\n\n    def forward(self, x):\n        x = self.backbone.forward_features(x)\n        x = self.shared(x)\n        return {k: classifier(x) for k, classifier in self.classifiers.items()}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:31.886003Z","iopub.execute_input":"2024-12-13T03:54:31.88634Z","iopub.status.idle":"2024-12-13T03:54:31.894017Z","shell.execute_reply.started":"2024-12-13T03:54:31.886309Z","shell.execute_reply":"2024-12-13T03:54:31.893105Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class MultiClassTraumaModel(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # Use ResNet34 for more capacity\n        self.backbone = timm.create_model('resnet34', pretrained=True, in_chans=1)\n        n_features = self.backbone.fc.in_features\n        \n        self.dropout1 = nn.Dropout(0.5)\n        self.dropout2 = nn.Dropout(0.3)\n        \n        # Shared feature processing with more capacity\n        self.shared_features = nn.Sequential(\n            nn.AdaptiveAvgPool2d(1),\n            nn.Flatten(),\n            nn.Linear(n_features, 512),\n            nn.BatchNorm1d(512),\n            nn.ReLU(),\n            self.dropout1\n        )\n        \n        # Separate classifiers for each organ\n        self.classifiers = nn.ModuleDict({\n            organ: nn.Sequential(\n                nn.Linear(512, 256),\n                nn.BatchNorm1d(256),\n                nn.ReLU(),\n                self.dropout2,\n                nn.Linear(256, 3)\n            ) for organ in ['kidney', 'liver', 'spleen']\n        })\n    \n    def forward(self, x):\n        features = self.backbone.forward_features(x)\n        shared = self.shared_features(features)\n        return {k: classifier(shared) for k, classifier in self.classifiers.items()}\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:34.021275Z","iopub.execute_input":"2024-12-13T03:54:34.021587Z","iopub.status.idle":"2024-12-13T03:54:34.028767Z","shell.execute_reply.started":"2024-12-13T03:54:34.02156Z","shell.execute_reply":"2024-12-13T03:54:34.027851Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class BinaryFocalLoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n    \n    def forward(self, pred, target, task_type):\n        # Skip if no positive samples in batch\n        if target.sum() == 0:\n            return pred.sum() * 0\n        \n        # Define weights based on task type\n        if task_type == 'bowel':\n            weights = torch.tensor([1.0, 2.0])  # healthy=1, injury=2\n        elif task_type == 'extravasation':\n            weights = torch.tensor([1.0, 6.0])  # healthy=1, injury=6\n        elif task_type == 'any_injury':\n            weights = torch.tensor([1.0, 6.0])  # healthy=1, injury=6\n        else:\n            weights = torch.tensor([1.0, 1.0])  # default case\n            \n        weights = weights.to(pred.device)\n        \n        return F.cross_entropy(pred, target, weight=weights)\n        \n\nclass MultiCELoss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        # Weight pattern: healthy=1, low_grade=2, high_grade=4\n        self.weights = {\n            'kidney': torch.tensor([1.0, 2.0, 4.0]),\n            'liver': torch.tensor([1.0, 2.0, 4.0]), \n            'spleen': torch.tensor([1.0, 2.0, 4.0])\n        }\n    \n    def forward(self, pred, target, organ_type):\n        # Get the weights for the specific organ\n        weights = self.weights[organ_type].to(pred.device)\n        return F.cross_entropy(pred, target, weight=weights)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:36.367618Z","iopub.execute_input":"2024-12-13T03:54:36.36849Z","iopub.status.idle":"2024-12-13T03:54:36.375949Z","shell.execute_reply.started":"2024-12-13T03:54:36.368455Z","shell.execute_reply":"2024-12-13T03:54:36.374951Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_valid_slices(patient_id, series_id, injury_type, enhanced_df):\n    \"\"\"\n    Get valid slices for a given patient and injury type.\n    \n    Args:\n        patient_id (str): Patient identifier\n        series_id (str): Series identifier\n        injury_type (str): Type of injury ('bowel' or 'extravasation')\n        enhanced_df (pd.DataFrame): Enhanced dataframe with injury information\n        \n    Returns:\n        list: List of valid slice indices for the specified injury type\n    \"\"\"\n    try:\n        # Filter data for specific patient and series\n        patient_data = enhanced_df[\n            (enhanced_df['patient_id'] == patient_id) &\n            (enhanced_df['series_id'] == series_id)\n        ]\n        \n        if patient_data.empty:\n            return []\n\n        if injury_type == 'bowel':\n            injury_name = 'Bowel'\n        elif injury_type == 'extravasation':\n            injury_name = 'Active_Extravasation'\n        else:\n            return []\n            \n        # Get injury slices from the filtered data\n        injury_rows = patient_data[patient_data['injury_name'] == injury_name]\n        \n        if injury_rows.empty:\n            return []\n            \n        # Get instance numbers (slice indices)\n        slices = injury_rows['instance_number'].iloc[0]\n        \n        # Handle both list and single value cases\n        if isinstance(slices, list):\n            return slices\n        elif isinstance(slices, (int, float)):\n            return [int(slices)]\n        else:\n            return []\n            \n    except Exception as e:\n        print(f\"Error getting valid slices for patient {patient_id}: {str(e)}\")\n        return []","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:38.372232Z","iopub.execute_input":"2024-12-13T03:54:38.37298Z","iopub.status.idle":"2024-12-13T03:54:38.379324Z","shell.execute_reply.started":"2024-12-13T03:54:38.372948Z","shell.execute_reply":"2024-12-13T03:54:38.378395Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_organ_label(row, organ):\n    \"\"\"\n    Get the label for an organ (helper function).\n    \n    Args:\n        row (pd.Series): Row from the dataframe\n        organ (str): Organ name ('kidney', 'liver', or 'spleen')\n        \n    Returns:\n        int: Label indicating injury severity (0: healthy, 1: low-grade, 2: high-grade)\n    \"\"\"\n    if row[f'{organ}_healthy'] == 1:\n        return 0\n    elif row[f'{organ}_low'] == 1:\n        return 1\n    elif row[f'{organ}_high'] == 1:\n        return 2\n    return 0  # Default to healthy if no label is found\n\nclass EnhancedDataset(Dataset):\n    def __init__(self, enhanced_df, image_dir, transform=None):\n        self.df = enhanced_df\n        self.image_dir = image_dir\n        self.transform = transform\n        \n        if not os.path.exists(image_dir):\n            raise ValueError(f\"Image directory {image_dir} does not exist\")\n    \n    def load_and_preprocess_dicom(self, path):\n        try:\n            dicom = pydicom.dcmread(path)\n            img = dicom.pixel_array.astype(np.float32)\n            \n            if hasattr(dicom, 'RescaleIntercept') and hasattr(dicom, 'RescaleSlope'):\n                slope = float(dicom.RescaleSlope)\n                intercept = float(dicom.RescaleIntercept)\n                img = (img * slope + intercept).astype(np.float32)\n            \n            # Apply windowing\n            window_center, window_width = 40, 400\n            img_min = window_center - window_width // 2\n            img_max = window_center + window_width // 2\n            img = np.clip(img, img_min, img_max)\n            img = ((img - img_min) / (img_max - img_min) * 255.0).astype(np.uint8)\n            \n            # Convert to float and normalize to [0, 1]\n            img = img.astype(np.float32) / 255.0\n            \n            if self.transform:\n                transformed = self.transform(image=img)\n                img = transformed['image']\n            \n            if len(img.shape) == 2:\n                img = torch.from_numpy(img).unsqueeze(0)\n                \n            return img\n            \n        except Exception as e:\n            print(f\"Error processing DICOM {path}: {str(e)}\")\n            return torch.zeros(1, 256, 256)\n\n    def load_patient_images(self, patient_id, series_id, bowel_slices=None, extravasation_slices=None):\n        try:\n            patient_path = os.path.join(self.image_dir, str(patient_id), str(series_id))\n            if not os.path.exists(patient_path):\n                raise FileNotFoundError(f\"Patient path {patient_path} not found\")\n                \n            dcm_files = sorted([f for f in os.listdir(patient_path) if f.endswith('.dcm')])\n            if not dcm_files:\n                raise ValueError(\"No DICOM files found\")\n                \n            if not bowel_slices and not extravasation_slices:\n                middle_idx = len(dcm_files) // 2\n                image_path = os.path.join(patient_path, dcm_files[middle_idx])\n                img = self.load_and_preprocess_dicom(image_path)\n                if img is None:\n                    return torch.zeros(1, 256, 256)\n                return img\n                \n            images = []\n            slice_indices = set()\n            if bowel_slices:\n                slice_indices.update(bowel_slices)\n            if extravasation_slices:\n                slice_indices.update(extravasation_slices)\n                \n            for idx in sorted(slice_indices):\n                if 0 <= idx < len(dcm_files):\n                    image_path = os.path.join(patient_path, dcm_files[idx])\n                    img = self.load_and_preprocess_dicom(image_path)\n                    if img is not None:\n                        images.append(img)\n            \n            if not images:\n                return torch.zeros(1, 256, 256)\n                \n            return torch.stack(images)\n            \n        except Exception as e:\n            print(f\"Error loading images for patient {patient_id}: {str(e)}\")\n            return torch.zeros(1, 256, 256)\n\n    def __getitem__(self, idx):\n        row = self.df.iloc[idx]\n        \n        patient_id = str(row['patient_id'])\n        series_id = str(row['series_id'])\n        \n        # Get valid slices\n        bowel_slices = get_valid_slices(patient_id, series_id, 'bowel', self.df)\n        extravasation_slices = get_valid_slices(patient_id, series_id, 'extravasation', self.df)\n        \n        # Load images\n        images = self.load_patient_images(patient_id, series_id, bowel_slices, extravasation_slices)\n        \n        # Create labels dictionary\n        labels = {\n            'bowel_injury': torch.tensor(row['bowel_injury'], dtype=torch.long),\n            'extravasation_injury': torch.tensor(row['extravasation_injury'], dtype=torch.long),\n            'any_injury': torch.tensor(row['any_injury'], dtype=torch.long),\n            'kidney': torch.tensor(get_organ_label(row, 'kidney'), dtype=torch.long),\n            'liver': torch.tensor(get_organ_label(row, 'liver'), dtype=torch.long),\n            'spleen': torch.tensor(get_organ_label(row, 'spleen'), dtype=torch.long)\n        }\n        \n        return images, labels\n\n    def __len__(self):\n        return len(self.df)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:52.652881Z","iopub.execute_input":"2024-12-13T03:54:52.653226Z","iopub.status.idle":"2024-12-13T03:54:52.669071Z","shell.execute_reply.started":"2024-12-13T03:54:52.653194Z","shell.execute_reply":"2024-12-13T03:54:52.668108Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def prepare_enhanced_datasets():\n    \"\"\"Prepare train and validation datasets using enhanced_df\"\"\"\n    try:\n        # Get enhanced_df from CSV preprocessing\n        enhanced_df, _ = preprocess_data()  \n        \n        # Prepare train/validation splits\n        train_idx, val_idx = train_test_split(\n            range(len(enhanced_df)),\n            test_size=0.2,\n            stratify=enhanced_df['any_injury'],\n            random_state=42\n        )\n        \n        # Define transforms\n        train_transform = A.Compose([\n            A.Resize(256, 256),\n            A.HorizontalFlip(p=0.5),\n            A.VerticalFlip(p=0.5),\n            A.RandomRotate90(p=0.5),\n            A.ShiftScaleRotate(shift_limit=0.0625, scale_limit=0.1, rotate_limit=45, p=0.5),\n            A.OneOf([\n                A.GaussNoise(var_limit=[10, 50]),\n                A.GaussianBlur(),\n                A.MotionBlur(),\n            ], p=0.3),\n            A.OneOf([\n                A.OpticalDistortion(),\n                A.GridDistortion(),\n                A.ElasticTransform(),\n            ], p=0.3),\n            A.OneOf([\n                A.CLAHE(),\n                A.RandomBrightnessContrast(),\n                A.RandomGamma(),\n            ], p=0.3),\n            ToTensorV2(),\n        ])\n        \n        val_transform = A.Compose([\n            A.Resize(256, 256),\n            ToTensorV2(),\n        ])\n        \n        # Create datasets\n        train_dataset = EnhancedDataset(\n            enhanced_df.iloc[train_idx],\n            \"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images\",\n            transform=train_transform\n        )\n        \n        val_dataset = EnhancedDataset(\n            enhanced_df.iloc[val_idx],\n            \"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images\",\n            transform=val_transform\n        )\n        \n        return train_dataset, val_dataset\n        \n    except Exception as e:\n        print(f\"Error preparing datasets: {str(e)}\")\n        return None, None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:54.598866Z","iopub.execute_input":"2024-12-13T03:54:54.599211Z","iopub.status.idle":"2024-12-13T03:54:54.607055Z","shell.execute_reply.started":"2024-12-13T03:54:54.599183Z","shell.execute_reply":"2024-12-13T03:54:54.606172Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_model(model, train_loader, criterion, optimizer, config, device, model_type):\n    \"\"\"Streamlined training function with minimal output\"\"\"\n    model.train()\n    total_loss = 0\n    start_time = time.time()\n\n    for batch_idx, (images, targets) in enumerate(train_loader):\n        # Check time limit\n        if time.time() - start_time > config['time_limit']:\n            break\n            \n        images = images.to(device)\n        targets = {k: v.to(device) for k, v in targets.items()}\n        \n        optimizer.zero_grad()\n        outputs = model(images)\n        \n        # Calculate loss based on model type\n        if model_type == 'binary':\n            losses = [criterion(outputs[k], targets[k], k) for k in outputs]\n        else:  # multi-class\n            losses = [criterion(outputs[k], targets[k], k) for k in outputs]\n            \n        loss = sum(losses) / len(losses)\n        loss.backward()\n        \n        torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)\n        optimizer.step()\n        \n        total_loss += loss.item()\n\n    final_avg_loss = total_loss / (batch_idx + 1)\n    print(f\"Training Loss: {final_avg_loss:.4f}\")\n    return final_avg_loss","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:56.778075Z","iopub.execute_input":"2024-12-13T03:54:56.778412Z","iopub.status.idle":"2024-12-13T03:54:56.785617Z","shell.execute_reply.started":"2024-12-13T03:54:56.778381Z","shell.execute_reply":"2024-12-13T03:54:56.784744Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate_model(model, val_loader, criterion, device, model_type):\n    \"\"\"Streamlined validation function with minimal output\"\"\"\n    model.eval()\n    total_loss = 0\n    all_predictions = {}\n    all_targets = {}\n\n    with torch.no_grad():\n        for images, targets in val_loader:\n            images = images.to(device)\n            targets = {k: v.to(device) for k, v in targets.items()}\n            outputs = model(images)\n            \n            # Calculate loss\n            if model_type == 'binary':\n                losses = [criterion(outputs[k], targets[k], k) for k in outputs]\n            else:\n                losses = [criterion(outputs[k], targets[k], k) for k in outputs]\n            loss = sum(losses) / len(losses)\n            total_loss += loss.item()\n\n            # Store predictions and targets\n            for k in outputs:\n                if k not in all_predictions:\n                    all_predictions[k] = []\n                    all_targets[k] = []\n                \n                if model_type == 'binary':\n                    probs = torch.softmax(outputs[k], dim=1)\n                    all_predictions[k].extend(probs[:, 1].cpu().numpy())\n                else:\n                    preds = outputs[k].argmax(dim=1)\n                    all_predictions[k].extend(preds.cpu().numpy())\n                all_targets[k].extend(targets[k].cpu().numpy())\n\n    # Calculate and print metrics\n    avg_loss = total_loss / len(val_loader)\n    print(f\"Validation Loss: {avg_loss:.4f}\\n\")\n    print(\"Validation Metrics:\")\n    \n    for k in all_predictions:\n        preds = np.array(all_predictions[k])\n        targets = np.array(all_targets[k])\n        accuracy = accuracy_score(targets, (preds > 0.5).astype(int) if model_type == 'binary' else preds)\n        \n        if model_type == 'binary':\n            try:\n                auc = roc_auc_score(targets, preds)\n                print(f\"{k}: Accuracy = {accuracy:.4f}, AUC = {auc:.4f}\")\n            except:\n                print(f\"{k}: Accuracy = {accuracy:.4f}\")\n        else:\n            print(f\"{k}: Accuracy = {accuracy:.4f}\")\n\n    return avg_loss\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:54:58.624881Z","iopub.execute_input":"2024-12-13T03:54:58.6252Z","iopub.status.idle":"2024-12-13T03:54:58.634906Z","shell.execute_reply.started":"2024-12-13T03:54:58.625174Z","shell.execute_reply":"2024-12-13T03:54:58.633928Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main_training():\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"Using device: {device}\")\n\n    # Configuration\n    config = {\n        'batch_size': 16,\n        'num_epochs': 15,\n        'time_limit': 800,\n        'binary_lr': 1e-4,\n        'multi_lr': 2e-4,\n        'weight_decay': 0.01,\n    }\n\n    # Get enhanced datasets\n    train_dataset, val_dataset = prepare_enhanced_datasets()\n\n    # Create data loaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config['batch_size'],\n        shuffle=True,\n        num_workers=2,\n        pin_memory=True\n    )\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=config['batch_size'],\n        shuffle=False,\n        num_workers=2,\n        pin_memory=True\n    )\n\n    # Initialize binary model (MobileNetV2)\n    binary_model = BinaryTraumaModel().to(device)\n    binary_criterion = BinaryFocalLoss()\n    binary_optimizer = optim.AdamW(\n        binary_model.parameters(),\n        lr=config['binary_lr'],\n        weight_decay=config['weight_decay']\n    )\n    binary_scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        binary_optimizer,\n        mode='min',\n        factor=0.5,\n        patience=2,\n        verbose=True\n    )\n\n    # Initialize multi-class model (ResNet34)\n    multi_model = MultiClassTraumaModel().to(device)\n    multi_criterion = MultiCELoss()\n    multi_optimizer = optim.AdamW(\n        multi_model.parameters(),\n        lr=config['multi_lr'],\n        weight_decay=config['weight_decay']\n    )\n    multi_scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        multi_optimizer,\n        mode='min',\n        factor=0.5,\n        patience=2,\n        verbose=True\n    )\n\n    # Training binary model\n    print(\"Training binary classification model (MobileNetV2)...\")\n    best_binary_loss = float('inf')\n    binary_patience = 0\n\n    for epoch in range(config['num_epochs']):\n        print(f\"\\nEpoch {epoch+1}/{config['num_epochs']}\")\n\n        train_loss = train_model(\n            binary_model,\n            train_loader,  \n            binary_criterion,\n            binary_optimizer,\n            config,\n            device,\n            'binary'\n        )\n\n        val_loss = validate_model(\n            binary_model,\n            val_loader,  \n            binary_criterion,\n            device,\n            'binary'\n        )\n\n        binary_scheduler.step(val_loss)\n\n        if val_loss < best_binary_loss:\n            best_binary_loss = val_loss\n            torch.save(binary_model.state_dict(), 'best_binary_model.pth')\n            print(\"Saved new best binary model!\")\n            binary_patience = 0\n        else:\n            binary_patience += 1\n            if binary_patience >= 5:\n                print(\"Early stopping triggered for binary model\")\n                break\n\n    # Training multi-class model\n    print(\"\\nTraining multi-class model (ResNet34)...\")\n    best_multi_loss = float('inf')\n    multi_patience = 0\n\n    for epoch in range(config['num_epochs']):\n        print(f\"\\nEpoch {epoch+1}/{config['num_epochs']}\")\n\n        train_loss = train_model(\n            multi_model,\n            train_loader,  \n            multi_criterion,\n            multi_optimizer,\n            config,\n            device,\n            'multi'\n        )\n\n        val_loss = validate_model(\n            multi_model,\n            val_loader, \n            multi_criterion,\n            device,\n            'multi'\n        )\n\n        multi_scheduler.step(val_loss)\n\n        if val_loss < best_multi_loss:\n            best_multi_loss = val_loss\n            torch.save(multi_model.state_dict(), 'best_multi_model.pth')\n            print(\"Saved new best multi-class model!\")\n            multi_patience = 0\n        else:\n            multi_patience += 1\n            if multi_patience >= 5:\n                print(\"Early stopping triggered for multi-class model\")\n                break\n\n    print(\"\\nTraining completed!\")\n    print(f\"Best binary model validation loss: {best_binary_loss:.4f}\")\n    print(f\"Best multi-class model validation loss: {best_multi_loss:.4f}\")\n\n    # Memory cleanup\n    gc.collect()\n    if torch.cuda.is_available():\n        torch.cuda.empty_cache()\n\nif __name__ == \"__main__\":\n    main_training()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T03:55:00.724436Z","iopub.execute_input":"2024-12-13T03:55:00.724873Z","iopub.status.idle":"2024-12-13T04:12:13.100387Z","shell.execute_reply.started":"2024-12-13T03:55:00.724842Z","shell.execute_reply":"2024-12-13T04:12:13.099364Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prediction Visualization","metadata":{}},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport pydicom\nimport os\nfrom torchvision import transforms\nimport matplotlib.pyplot as plt\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T04:17:24.712137Z","iopub.execute_input":"2024-12-13T04:17:24.712499Z","iopub.status.idle":"2024-12-13T04:17:24.717284Z","shell.execute_reply.started":"2024-12-13T04:17:24.712468Z","shell.execute_reply":"2024-12-13T04:17:24.716363Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_and_preprocess_dicom(path):\n    \"\"\"Load and preprocess a single DICOM image\"\"\"\n    try:\n        dicom = pydicom.dcmread(path)\n        img = dicom.pixel_array.astype(np.float32)\n        \n        # Convert to HU\n        if hasattr(dicom, 'RescaleIntercept') and hasattr(dicom, 'RescaleSlope'):\n            slope = float(dicom.RescaleSlope)\n            intercept = float(dicom.RescaleIntercept)\n            img = (img * slope + intercept).astype(np.float32)\n        \n        # Apply windowing\n        window_center, window_width = 40, 400\n        img_min = window_center - window_width // 2\n        img_max = window_center + window_width // 2\n        img = np.clip(img, img_min, img_max)\n        img = ((img - img_min) / (img_max - img_min) * 255.0).astype(np.uint8)\n        \n        # Convert to float and normalize to [0, 1]\n        img = img.astype(np.float32) / 255.0\n        return img\n        \n    except Exception as e:\n        print(f\"Error processing DICOM {path}: {str(e)}\")\n        return None\n\ndef predict_case(binary_model, multi_model, image_paths, device):\n    \"\"\"Make predictions for a single case with multiple images\"\"\"\n    transform = A.Compose([\n        A.Resize(256, 256),\n        ToTensorV2(),\n    ])\n    \n    binary_model.eval()\n    multi_model.eval()\n    \n    predictions = {\n        'binary': {\n            'bowel_injury': [],\n            'extravasation_injury': [],\n            'any_injury': []\n        },\n        'multi': {\n            'kidney': [],\n            'liver': [],\n            'spleen': []\n        }\n    }\n    \n    with torch.no_grad():\n        for img_path in image_paths:\n            # Load and preprocess image\n            img = load_and_preprocess_dicom(img_path)\n            if img is None:\n                continue\n                \n            # Apply transform\n            img = transform(image=img)['image']\n            \n            # Add batch dimension and move to device\n            img = img.unsqueeze(0).to(device)\n            \n            # Binary predictions\n            binary_outputs = binary_model(img)\n            for k, v in binary_outputs.items():\n                probs = torch.softmax(v, dim=1)\n                pred_prob = probs[:, 1].cpu().numpy()[0]  # Probability of positive class\n                predictions['binary'][k].append(pred_prob)\n            \n            # Multi-class predictions\n            multi_outputs = multi_model(img)\n            for k, v in multi_outputs.items():\n                pred_class = torch.argmax(v, dim=1).cpu().numpy()[0]\n                predictions['multi'][k].append(pred_class)\n                \n    return predictions\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T04:18:06.1487Z","iopub.execute_input":"2024-12-13T04:18:06.149043Z","iopub.status.idle":"2024-12-13T04:18:06.160274Z","shell.execute_reply.started":"2024-12-13T04:18:06.149013Z","shell.execute_reply":"2024-12-13T04:18:06.159501Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def display_predictions(images, predictions, case_id, save_path=None):\n    \"\"\"Display and optionally save predictions with images\"\"\"\n    n_images = len(images)\n    fig, axes = plt.subplots(2, n_images, figsize=(4*n_images, 8))\n    \n    # Plot images\n    for i, img_path in enumerate(images):\n        img = load_and_preprocess_dicom(img_path)\n        if img is not None:\n            axes[0, i].imshow(img, cmap='gray')\n            axes[0, i].axis('off')\n            axes[0, i].set_title(f'Image {i+1}')\n    \n    # Plot predictions\n    for i in range(n_images):\n        text = f\"Binary Predictions:\\n\"\n        for k, v in predictions['binary'].items():\n            if i < len(v):\n                text += f\"{k}: {v[i]:.3f}\\n\"\n        \n        text += \"\\nMulti-class Predictions:\\n\"\n        for k, v in predictions['multi'].items():\n            if i < len(v):\n                text += f\"{k}: {v[i]}\\n\"\n                \n        axes[1, i].text(0.1, 0.5, text, transform=axes[1, i].transAxes, \n                       verticalalignment='center')\n        axes[1, i].axis('off')\n    \n    plt.suptitle(f'Case {case_id} Predictions')\n    plt.tight_layout()\n    \n    if save_path:\n        plt.savefig(save_path)\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T04:18:34.961443Z","iopub.execute_input":"2024-12-13T04:18:34.961794Z","iopub.status.idle":"2024-12-13T04:18:34.969216Z","shell.execute_reply.started":"2024-12-13T04:18:34.961763Z","shell.execute_reply":"2024-12-13T04:18:34.968379Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_valid_cases(base_path, enhanced_df):\n    \"\"\"Get valid cases with their image paths from the dataset\"\"\"\n    valid_cases = []\n    \n    # Get cases with active extravasation or bowel injuries\n    injury_cases = enhanced_df[\n        (enhanced_df['extravasation_injury'] == 1) | \n        (enhanced_df['bowel_injury'] == 1)\n    ]\n    \n    for _, row in injury_cases.iterrows():\n        patient_id = str(row['patient_id'])\n        series_id = str(row['series_id'])\n        \n        patient_path = os.path.join(base_path, patient_id, series_id)\n        if not os.path.exists(patient_path):\n            continue\n            \n        # Get all DICOM files in the directory\n        dcm_files = sorted([f for f in os.listdir(patient_path) if f.endswith('.dcm')])\n        if len(dcm_files) < 5:\n            continue\n            \n        # Get middle 5 slices\n        middle_idx = len(dcm_files) // 2\n        slice_indices = [\n            middle_idx - 2,\n            middle_idx - 1,\n            middle_idx,\n            middle_idx + 1,\n            middle_idx + 2\n        ]\n        \n        image_paths = [\n            os.path.join(patient_path, dcm_files[idx]) \n            for idx in slice_indices\n            if 0 <= idx < len(dcm_files)\n        ]\n        \n        if len(image_paths) == 5:\n            case_info = {\n                'case_id': patient_id,  # Using patient_id as case_id\n                'image_paths': image_paths,\n                'labels': {\n                    'bowel_injury': row['bowel_injury'],\n                    'extravasation_injury': row['extravasation_injury'],\n                    'any_injury': row['any_injury'],\n                    'kidney': 2 if row['kidney_high'] else 1 if row['kidney_low'] else 0,\n                    'liver': 2 if row['liver_high'] else 1 if row['liver_low'] else 0,\n                    'spleen': 2 if row['spleen_high'] else 1 if row['spleen_low'] else 0\n                }\n            }\n            valid_cases.append(case_info)\n            \n            if len(valid_cases) >= 5:  # Limit to 5 cases for now\n                break\n                \n    return valid_cases","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T04:24:17.663512Z","iopub.execute_input":"2024-12-13T04:24:17.664267Z","iopub.status.idle":"2024-12-13T04:24:17.672717Z","shell.execute_reply.started":"2024-12-13T04:24:17.66423Z","shell.execute_reply":"2024-12-13T04:24:17.671738Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Device configuration\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"Using device: {device}\")\n    \n    # Load trained models\n    binary_model = BinaryTraumaModel().to(device)\n    binary_model.load_state_dict(torch.load('best_binary_model.pth'))\n    \n    multi_model = MultiClassTraumaModel().to(device)\n    multi_model.load_state_dict(torch.load('best_multi_model.pth'))\n    \n    # Get enhanced_df from CSV preprocessing\n    enhanced_df, _ = preprocess_data()\n    \n    # Get valid cases\n    base_path = \"/kaggle/input/rsna-2023-abdominal-trauma-detection/train_images\"\n    test_cases = get_valid_cases(base_path, enhanced_df)\n    \n    if not test_cases:\n        print(\"No valid cases found!\")\n        return\n        \n    print(f\"Found {len(test_cases)} valid cases\")\n    \n    # Create a directory for saving predictions\n    os.makedirs('predictions', exist_ok=True)\n    \n    # Make predictions for all cases\n    all_predictions = {}\n    \n    for case in test_cases:\n        case_id = case['case_id']\n        print(f\"\\nProcessing case {case_id}...\")\n        \n        predictions = predict_case(\n            binary_model, \n            multi_model, \n            case['image_paths'], \n            device\n        )\n        \n        all_predictions[case_id] = predictions\n        \n        # Display and save results for individual case\n        save_path = f\"predictions/case_{case_id}_predictions.png\"\n        display_predictions(\n            case['image_paths'],\n            predictions,\n            case_id,\n            save_path\n        )\n        \n        # Print numerical results\n        print(f\"\\nResults for Case {case_id}:\")\n        print(\"Ground Truth Labels:\")\n        for k, v in case['labels'].items():\n            print(f\"{k}: {v}\")\n            \n        print(\"\\nPredictions:\")\n        print(\"Binary Predictions:\")\n        for k, v in predictions['binary'].items():\n            mean_pred = np.mean(v)\n            if not np.isnan(mean_pred):\n                print(f\"{k}: {mean_pred:.3f} (mean probability)\")\n            \n        print(\"\\nMulti-class Predictions:\")\n        for k, v in predictions['multi'].items():\n            if len(v) > 0:\n                print(f\"{k}: Most common class = {np.bincount(v).argmax()}\")\n    \n    # Create summary visualization\n    plt.figure(figsize=(20, 4*len(test_cases)))\n    \n    for idx, case in enumerate(test_cases):\n        case_id = case['case_id']\n        predictions = all_predictions[case_id]\n        \n        for i, img_path in enumerate(case['image_paths']):\n            plt.subplot(len(test_cases), 5, idx*5 + i + 1)\n            img = load_and_preprocess_dicom(img_path)\n            if img is not None:\n                plt.imshow(img, cmap='gray')\n                \n                text = f\"Case {case_id}\\nImage {i+1}\\n\"\n                text += \"Ground Truth:\\n\"\n                for k, v in case['labels'].items():\n                    text += f\"{k}: {v}\\n\"\n                text += \"Predictions:\\n\"\n                for k, v in predictions['binary'].items():\n                    if i < len(v):\n                        text += f\"{k}: {v[i]:.2f}\\n\"\n                for k, v in predictions['multi'].items():\n                    if i < len(v):\n                        text += f\"{k}: {v[i]}\\n\"\n                        \n                plt.title(text, fontsize=8)\n            plt.axis('off')\n    \n    plt.tight_layout()\n    plt.savefig('predictions/all_cases_summary.png')\n    plt.show()\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-12-13T04:24:41.403545Z","iopub.execute_input":"2024-12-13T04:24:41.403913Z","iopub.status.idle":"2024-12-13T04:24:53.891416Z","shell.execute_reply.started":"2024-12-13T04:24:41.403883Z","shell.execute_reply":"2024-12-13T04:24:53.890494Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}