{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.14","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":71549,"databundleVersionId":8561470,"sourceType":"competition"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport pydicom\nimport cv2\nimport numpy as np\nimport os\nfrom pathlib import Path\nimport warnings\nfrom typing import List, Dict, Tuple","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:30:21.192259Z","iopub.execute_input":"2024-11-21T07:30:21.193083Z","iopub.status.idle":"2024-11-21T07:30:22.501208Z","shell.execute_reply.started":"2024-11-21T07:30:21.193042Z","shell.execute_reply":"2024-11-21T07:30:22.500383Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Data Visualization","metadata":{}},{"cell_type":"markdown","source":"I adapted and restructured some code from the https://www.kaggle.com/code/abhinavsuri/anatomy-image-visualization-overview-rsna-raids while adding more comprehensive visualization and analysis capabilities in data visualization task. ","metadata":{}},{"cell_type":"code","source":"def main():\n    \"\"\"Main function for data exploration and visualization\"\"\"\n    # Set up paths\n    base_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\n    train_path = f\"{base_path}/train_images\"\n    \n    # Load metadata\n    train_df = pd.read_csv(f\"{base_path}/train.csv\")\n    coordinates_df = pd.read_csv(f\"{base_path}/train_label_coordinates.csv\")\n    series_desc_df = pd.read_csv(f\"{base_path}/train_series_descriptions.csv\")\n    \n    # Print basic dataset information\n    print_dataset_stats(train_df)\n    \n    # Visualize distribution of conditions\n    plot_condition_distributions(train_df)\n    \n    # Example: Load and display images for one patient\n    example_patient_id = train_df['study_id'].iloc[0]\n    patient_images = load_patient_images(example_patient_id, train_path, series_desc_df)\n    visualize_patient_images(patient_images)\n    \n    # Show pathology locations for the example patient\n    show_pathology_locations(example_patient_id, patient_images, coordinates_df, train_df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:30:28.804935Z","iopub.execute_input":"2024-11-21T07:30:28.805860Z","iopub.status.idle":"2024-11-21T07:30:28.811347Z","shell.execute_reply.started":"2024-11-21T07:30:28.805827Z","shell.execute_reply":"2024-11-21T07:30:28.810455Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def print_dataset_stats(train_df: pd.DataFrame) -> None:\n    \"\"\"Print basic statistics about the dataset\"\"\"\n    print(f\"Total number of cases: {len(train_df)}\")\n    print(\"\\nColumns in dataset:\")\n    for col in train_df.columns:\n        print(f\"- {col}\")\n    \n    # Count conditions by type\n    condition_types = ['foraminal', 'subarticular', 'canal']\n    print(\"\\nCondition counts:\")\n    for condition in condition_types:\n        cols = [col for col in train_df.columns if condition in col]\n        total_cases = train_df[cols].notna().sum().sum()\n        print(f\"{condition}: {total_cases} annotations\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:30:31.566231Z","iopub.execute_input":"2024-11-21T07:30:31.567080Z","iopub.status.idle":"2024-11-21T07:30:31.572317Z","shell.execute_reply.started":"2024-11-21T07:30:31.567044Z","shell.execute_reply":"2024-11-21T07:30:31.571365Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def plot_condition_distributions(train_df: pd.DataFrame) -> None:\n    \"\"\"Plot distribution of different conditions\"\"\"\n    condition_types = ['foraminal', 'subarticular', 'canal']\n    \n    fig, axes = plt.subplots(1, 3, figsize=(20, 5))\n    \n    for idx, condition in enumerate(condition_types):\n        # Get columns for this condition\n        condition_cols = [col for col in train_df.columns if condition in col]\n        condition_data = train_df[condition_cols]\n        \n        # Count value distributions\n        with warnings.catch_warnings():\n            warnings.simplefilter(action='ignore', category=FutureWarning)\n            value_counts = condition_data.apply(pd.value_counts).fillna(0).T\n        \n        # Plot\n        value_counts.plot(kind='bar', stacked=True, ax=axes[idx])\n        axes[idx].set_title(f'{condition} Distribution')\n        axes[idx].set_xlabel('Vertebral Level')\n        axes[idx].set_ylabel('Count')\n        plt.xticks(rotation=45)\n    \n    plt.tight_layout()\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:30:34.282129Z","iopub.execute_input":"2024-11-21T07:30:34.282471Z","iopub.status.idle":"2024-11-21T07:30:34.289392Z","shell.execute_reply.started":"2024-11-21T07:30:34.282440Z","shell.execute_reply":"2024-11-21T07:30:34.288349Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_patient_images(study_id: int, base_path: str, series_desc_df: pd.DataFrame) -> Dict:\n    \"\"\"Load all images for a single patient\"\"\"\n    study_path = Path(base_path) / str(study_id)\n    \n    image_data = {}\n    for series_id in os.listdir(study_path):\n        if series_id.startswith('.'):\n            continue\n            \n        # Get series description\n        series_desc = series_desc_df[\n            (series_desc_df['study_id'] == study_id) & \n            (series_desc_df['series_id'] == int(series_id))\n        ]['series_description'].iloc[0]\n        \n        # Load all DICOM images in this series\n        series_path = study_path / series_id\n        image_data[series_id] = {\n            'description': series_desc,\n            'images': []\n        }\n        \n        for dcm_file in sorted(os.listdir(series_path)):\n            if dcm_file.endswith('.dcm'):\n                dcm_path = series_path / dcm_file\n                dcm = pydicom.dcmread(str(dcm_path))\n                image_data[series_id]['images'].append({\n                    'instance_number': dcm_file.replace('.dcm', ''),\n                    'dicom': dcm\n                })\n    \n    return image_data","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:30:36.956892Z","iopub.execute_input":"2024-11-21T07:30:36.957478Z","iopub.status.idle":"2024-11-21T07:30:36.963943Z","shell.execute_reply.started":"2024-11-21T07:30:36.957434Z","shell.execute_reply":"2024-11-21T07:30:36.962963Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_patient_images(image_data: Dict) -> None:\n    \"\"\"Display all images for a patient organized by series\"\"\"\n    for series_id, series_data in image_data.items():\n        images = [img['dicom'].pixel_array for img in series_data['images']]\n        \n        # Calculate grid dimensions\n        n_images = len(images)\n        n_cols = min(4, n_images)\n        n_rows = (n_images + n_cols - 1) // n_cols\n        \n        # Create subplot grid\n        fig, axes = plt.subplots(n_rows, n_cols, figsize=(15, 3*n_rows))\n        fig.suptitle(f\"Series: {series_data['description']}\")\n        \n        if n_rows == 1:\n            axes = [axes]\n        \n        # Plot each image\n        for idx, img in enumerate(images):\n            row = idx // n_cols\n            col = idx % n_cols\n            axes[row][col].imshow(img, cmap='gray')\n            axes[row][col].axis('off')\n            axes[row][col].set_title(f\"Image {idx+1}\")\n        \n        # Turn off empty subplots\n        for idx in range(len(images), n_rows * n_cols):\n            row = idx // n_cols\n            col = idx % n_cols\n            axes[row][col].axis('off')\n        \n        plt.tight_layout()\n        plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:30:40.077917Z","iopub.execute_input":"2024-11-21T07:30:40.078568Z","iopub.status.idle":"2024-11-21T07:30:40.085724Z","shell.execute_reply.started":"2024-11-21T07:30:40.078507Z","shell.execute_reply":"2024-11-21T07:30:40.084800Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def show_pathology_locations(\n    study_id: int, \n    image_data: Dict, \n    coordinates_df: pd.DataFrame,\n    train_df: pd.DataFrame\n) -> None:\n    \"\"\"Display images with annotated pathology locations\"\"\"\n    # Get coordinates for this study\n    study_coords = coordinates_df[coordinates_df['study_id'] == study_id]\n    \n    # Get conditions for this study\n    study_conditions = train_df[train_df['study_id'] == study_id]\n    \n    for _, coord in study_coords.iterrows():\n        series_id = str(coord['series_id'])\n        instance_num = str(coord['instance_number'])\n        \n        # Find the corresponding image\n        for img in image_data[series_id]['images']:\n            if img['instance_number'] == instance_num:\n                # Create a copy of the image for annotation\n                pixel_array = img['dicom'].pixel_array\n                normalized_img = cv2.normalize(\n                    pixel_array, \n                    None, \n                    alpha=0,\n                    beta=255, \n                    norm_type=cv2.NORM_MINMAX, \n                    dtype=cv2.CV_8U\n                )\n                \n                # Draw circle at pathology location\n                annotated_img = cv2.circle(\n                    normalized_img.copy(),\n                    (int(coord['x']), int(coord['y'])),\n                    radius=10,\n                    color=(255, 0, 0),\n                    thickness=2\n                )\n                \n                # Get severity for this condition/level\n                condition_col = f\"{coord['condition'].lower().replace(' ', '_')}_{coord['level'].lower().replace('/', '_')}\"\n                severity = study_conditions[condition_col].iloc[0] if condition_col in study_conditions else \"Unknown\"\n                \n                # Display image\n                plt.figure(figsize=(8, 8))\n                plt.imshow(annotated_img, cmap='gray')\n                plt.title(f\"Level: {coord['level']}\\nCondition: {coord['condition']}\\nSeverity: {severity}\")\n                plt.axis('off')\n                plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:30:42.983847Z","iopub.execute_input":"2024-11-21T07:30:42.984663Z","iopub.status.idle":"2024-11-21T07:30:42.992135Z","shell.execute_reply.started":"2024-11-21T07:30:42.984627Z","shell.execute_reply":"2024-11-21T07:30:42.991124Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:30:45.746718Z","iopub.execute_input":"2024-11-21T07:30:45.747332Z","iopub.status.idle":"2024-11-21T07:31:03.201088Z","shell.execute_reply.started":"2024-11-21T07:30:45.747299Z","shell.execute_reply":"2024-11-21T07:31:03.200074Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Pattern Analysis","metadata":{}},{"cell_type":"markdown","source":"Analyze four key aspects of the data:\n\n1. Class distributions: Understanding how severity levels are distributed across conditions\n2. Condition co-occurrence: Analyzing how different conditions appear together\n3. Level patterns: Understanding how conditions vary across vertebral levels\n4. Series patterns: Analyzing the relationship between image series types and conditions","metadata":{}},{"cell_type":"markdown","source":"Generate visualizations and print detailed statistics about patterns in the data. This will help us:\n\n1. Identify class imbalances\n2. Understand which conditions tend to occur together\n3. See if certain vertebral levels are more prone to specific conditions\n4. Determine which image series are most useful for detecting each condition","metadata":{}},{"cell_type":"code","source":"import pandas as pd\nimport matplotlib.pyplot as plt\nimport numpy as np\nimport seaborn as sns\nfrom pathlib import Path\nimport warnings\nwarnings.filterwarnings('ignore')","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:31:22.248708Z","iopub.execute_input":"2024-11-21T07:31:22.249284Z","iopub.status.idle":"2024-11-21T07:31:23.476135Z","shell.execute_reply.started":"2024-11-21T07:31:22.249250Z","shell.execute_reply":"2024-11-21T07:31:23.475446Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    \"\"\"Main function for analyzing patterns in the lumbar spine data\"\"\"\n    base_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\n    \n    # Load necessary data\n    train_df = pd.read_csv(f\"{base_path}/train.csv\")\n    coords_df = pd.read_csv(f\"{base_path}/train_label_coordinates.csv\")\n    series_df = pd.read_csv(f\"{base_path}/train_series_descriptions.csv\")\n    \n    # Analyze patterns\n    print(\"1. Analyzing class distribution patterns...\")\n    analyze_class_distributions(train_df)\n    \n    print(\"\\n2. Analyzing condition co-occurrence...\")\n    analyze_condition_cooccurrence(train_df)\n    \n    print(\"\\n3. Analyzing condition patterns across vertebral levels...\")\n    analyze_level_patterns(train_df)\n    \n    print(\"\\n4. Analyzing series descriptions and their relationships to conditions...\")\n    analyze_series_patterns(series_df, coords_df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:31:26.009157Z","iopub.execute_input":"2024-11-21T07:31:26.009983Z","iopub.status.idle":"2024-11-21T07:31:26.015040Z","shell.execute_reply.started":"2024-11-21T07:31:26.009947Z","shell.execute_reply":"2024-11-21T07:31:26.014105Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def analyze_class_distributions(train_df):\n    \"\"\"Analyze the distribution of severity classes across different conditions\"\"\"\n    # Get condition columns\n    condition_cols = [col for col in train_df.columns \n                     if any(cond in col for cond in ['foraminal', 'subarticular', 'canal'])]\n    \n    # Create figure for distribution plots\n    plt.figure(figsize=(15, 8))\n    severity_counts = {}\n    \n    # Count severity distributions for each condition type\n    for col in condition_cols:\n        severity_counts[col] = train_df[col].value_counts()\n    \n    # Convert to DataFrame for easier plotting\n    severity_df = pd.DataFrame(severity_counts).fillna(0)\n    \n    # Plot overall severity distribution\n    severity_df.transpose().plot(kind='bar', stacked=True)\n    plt.title('Severity Distribution Across All Conditions and Levels')\n    plt.xlabel('Condition and Level')\n    plt.ylabel('Count')\n    plt.xticks(rotation=45, ha='right')\n    plt.tight_layout()\n    plt.show()\n    \n    # Print summary statistics\n    print(\"\\nSeverity Distribution Summary:\")\n    total_annotations = severity_df.sum().sum()\n    for severity in ['Normal/Mild', 'Moderate', 'Severe']:\n        count = severity_df.loc[severity].sum()\n        percentage = (count/total_annotations) * 100\n        print(f\"{severity}: {count:.0f} cases ({percentage:.1f}%)\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:31:29.008442Z","iopub.execute_input":"2024-11-21T07:31:29.009286Z","iopub.status.idle":"2024-11-21T07:31:29.015986Z","shell.execute_reply.started":"2024-11-21T07:31:29.009251Z","shell.execute_reply":"2024-11-21T07:31:29.015005Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def analyze_condition_cooccurrence(train_df):\n    \"\"\"Analyze how different conditions co-occur\"\"\"\n    # Create separate DataFrames for each condition type\n    conditions = {\n        'canal': [col for col in train_df.columns if 'canal' in col],\n        'foraminal': [col for col in train_df.columns if 'foraminal' in col],\n        'subarticular': [col for col in train_df.columns if 'subarticular' in col]\n    }\n    \n    # Create co-occurrence matrix\n    plt.figure(figsize=(12, 8))\n    cooccurrence_matrix = np.zeros((3, 3))\n    condition_types = list(conditions.keys())\n    \n    for i, cond1 in enumerate(condition_types):\n        for j, cond2 in enumerate(condition_types):\n            # Count cases where both conditions are severe\n            severe_cases1 = train_df[conditions[cond1]] == 'Severe'\n            severe_cases2 = train_df[conditions[cond2]] == 'Severe'\n            cooccurrence = (severe_cases1.any(axis=1) & severe_cases2.any(axis=1)).sum()\n            cooccurrence_matrix[i, j] = cooccurrence\n    \n    # Plot co-occurrence heatmap\n    sns.heatmap(cooccurrence_matrix, \n                annot=True, \n                fmt='g',\n                xticklabels=condition_types,\n                yticklabels=condition_types)\n    plt.title('Co-occurrence of Severe Cases Between Conditions')\n    plt.tight_layout()\n    plt.show()\n    \n    print(\"\\nKey Co-occurrence Patterns:\")\n    for i, cond1 in enumerate(condition_types):\n        for j, cond2 in enumerate(condition_types):\n            if i < j:\n                print(f\"{cond1.capitalize()} + {cond2.capitalize()}: {cooccurrence_matrix[i,j]:.0f} severe cases\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:31:31.529463Z","iopub.execute_input":"2024-11-21T07:31:31.530401Z","iopub.status.idle":"2024-11-21T07:31:31.537364Z","shell.execute_reply.started":"2024-11-21T07:31:31.530361Z","shell.execute_reply":"2024-11-21T07:31:31.536488Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def analyze_level_patterns(train_df):\n    \"\"\"Analyze patterns across vertebral levels\"\"\"\n    levels = ['l1_l2', 'l2_l3', 'l3_l4', 'l4_l5', 'l5_s1']\n    conditions = ['canal', 'foraminal', 'subarticular']\n    \n    plt.figure(figsize=(15, 6))\n    \n    # Create matrix of severity by level\n    level_severity = np.zeros((len(conditions), len(levels)))\n    \n    for i, condition in enumerate(conditions):\n        for j, level in enumerate(levels):\n            # Get columns for this condition and level\n            cols = [col for col in train_df.columns if condition in col and level in col]\n            # Calculate percentage of severe cases\n            severe_cases = (train_df[cols] == 'Severe').sum().sum()\n            total_cases = train_df[cols].notna().sum().sum()\n            level_severity[i, j] = (severe_cases / total_cases * 100) if total_cases > 0 else 0\n    \n    # Plot level severity patterns\n    sns.heatmap(level_severity,\n                annot=True,\n                fmt='.1f',\n                xticklabels=levels,\n                yticklabels=conditions,\n                cmap='YlOrRd')\n    plt.title('Percentage of Severe Cases by Level and Condition')\n    plt.xlabel('Vertebral Level')\n    plt.ylabel('Condition Type')\n    plt.tight_layout()\n    plt.show()\n    \n    print(\"\\nLevel-wise Pattern Summary:\")\n    for condition in conditions:\n        print(f\"\\n{condition.capitalize()} patterns:\")\n        for level in levels:\n            cols = [col for col in train_df.columns if condition in col and level in col]\n            severe_count = (train_df[cols] == 'Severe').sum().sum()\n            total_count = train_df[cols].notna().sum().sum()\n            if total_count > 0:\n                percentage = (severe_count / total_count) * 100\n                print(f\"  {level}: {percentage:.1f}% severe cases\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:31:34.078274Z","iopub.execute_input":"2024-11-21T07:31:34.078625Z","iopub.status.idle":"2024-11-21T07:31:34.086890Z","shell.execute_reply.started":"2024-11-21T07:31:34.078596Z","shell.execute_reply":"2024-11-21T07:31:34.086031Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def analyze_series_patterns(series_df, coords_df):\n    \"\"\"Analyze patterns in series descriptions and their relationship to conditions\"\"\"\n    # Count series descriptions\n    series_counts = series_df['series_description'].value_counts()\n    \n    plt.figure(figsize=(12, 6))\n    series_counts.plot(kind='bar')\n    plt.title('Distribution of Series Types')\n    plt.xlabel('Series Description')\n    plt.ylabel('Count')\n    plt.xticks(rotation=45, ha='right')\n    plt.tight_layout()\n    plt.show()\n    \n    # Analyze which series types are used for different conditions\n    condition_series = coords_df.merge(series_df, \n                                     left_on=['study_id', 'series_id'],\n                                     right_on=['study_id', 'series_id'])\n    \n    print(\"\\nSeries Usage by Condition:\")\n    for condition in condition_series['condition'].unique():\n        print(f\"\\n{condition}:\")\n        series_for_condition = condition_series[\n            condition_series['condition'] == condition\n        ]['series_description'].value_counts()\n        \n        for series_type, count in series_for_condition.items():\n            print(f\"  {series_type}: {count} annotations\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:31:36.527139Z","iopub.execute_input":"2024-11-21T07:31:36.527946Z","iopub.status.idle":"2024-11-21T07:31:36.534249Z","shell.execute_reply.started":"2024-11-21T07:31:36.527908Z","shell.execute_reply":"2024-11-21T07:31:36.533273Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:31:39.530424Z","iopub.execute_input":"2024-11-21T07:31:39.530780Z","iopub.status.idle":"2024-11-21T07:31:41.049889Z","shell.execute_reply.started":"2024-11-21T07:31:39.530750Z","shell.execute_reply":"2024-11-21T07:31:41.049002Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Severe Class Imbalance:\n\n* Normal/Mild: 77.4%\n* Moderate: 16.3%\n* Severe: 6.3%\nThis suggests we'll need class balancing techniques.\n\nCondition-Level Patterns:\n\nDifferent conditions are best visible in different image types:\n\n* Canal Stenosis → Sagittal T2/STIR\n* Neural Foraminal Narrowing → Sagittal T1\n* Subarticular Stenosis → Axial T2\n\nLevel-wise Distribution:\n\n* L4/L5 has highest severity rates across conditions\n* Lower levels (L3-S1) show more severe cases than upper levels","metadata":{}},{"cell_type":"markdown","source":"# Preprocessing Pipeline","metadata":{}},{"cell_type":"markdown","source":"Key preprocessing strategies implemented based on my analysis:\n\nStratified Fold Creation:\n\n* Uses presence of severe cases for stratification\n* Ensures similar distribution of severe cases across folds\n* Groups by study_id to prevent data leakage\n\nCondition-Specific Processing:\n\nMaps each condition to its most informative series type:\n\n* Canal Stenosis → Sagittal T2/STIR\n* Foraminal Narrowing → Sagittal T1\n* Subarticular Stenosis → Axial T2\n\nImage Processing:\n\n* Normalizes pixel values\n* Extracts patches around annotation points\n* Maintains spatial context around pathologies\n\nClass Balancing:\n\n* Applies augmentation only to severe cases\n* Uses rotation and flipping to increase severe class representation\n* Helps address the 77.4% vs 6.3% class imbalance","metadata":{}},{"cell_type":"code","source":"from pathlib import Path\nfrom typing import Dict, List, Tuple\nfrom sklearn.model_selection import StratifiedGroupKFold\nfrom collections import defaultdict","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:32:10.712818Z","iopub.execute_input":"2024-11-21T07:32:10.713162Z","iopub.status.idle":"2024-11-21T07:32:10.884907Z","shell.execute_reply.started":"2024-11-21T07:32:10.713133Z","shell.execute_reply":"2024-11-21T07:32:10.883972Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def save_processed_data(samples: List[Dict], filename: str):\n    \"\"\"Save processed samples to a numpy file\"\"\"\n    # Convert samples to numpy arrays for efficient storage\n    processed_data = {\n        'images': [],\n        'conditions': [],\n        'levels': [],\n        'severities': [],\n        'study_ids': []\n    }\n    \n    # Only save samples with valid severity\n    valid_samples = [\n        sample for sample in samples \n        if isinstance(sample.get('severity'), str) and not pd.isna(sample.get('severity'))\n    ]\n    \n    for sample in valid_samples:\n        processed_data['images'].append(sample['image'])\n        processed_data['conditions'].append(sample['condition'])\n        processed_data['levels'].append(sample['level'])\n        processed_data['severities'].append(sample['severity'])\n        processed_data['study_ids'].append(sample['study_id'])\n    \n    # Convert lists to numpy arrays\n    for key in processed_data:\n        processed_data[key] = np.array(processed_data[key])\n    \n    # Save to file\n    np.save(filename, processed_data)\n    print(f\"Saved {len(valid_samples)} valid samples to {filename}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:32:13.426821Z","iopub.execute_input":"2024-11-21T07:32:13.427165Z","iopub.status.idle":"2024-11-21T07:32:13.434298Z","shell.execute_reply.started":"2024-11-21T07:32:13.427135Z","shell.execute_reply":"2024-11-21T07:32:13.433324Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_stratified_folds(train_df: pd.DataFrame, n_splits: int = 5) -> List[Dict]:\n    \"\"\"Create stratified folds ensuring similar distribution of severe cases\"\"\"\n    # Create severity indicator for stratification\n    severity_cols = [col for col in train_df.columns \n                    if any(c in col for c in ['canal', 'foraminal', 'subarticular'])]\n    \n    # Create binary indicator for having any severe case\n    train_df['has_severe'] = (train_df[severity_cols] == 'Severe').any(axis=1)\n    \n    # Initialize fold splitter\n    skf = StratifiedGroupKFold(n_splits=n_splits, shuffle=True, random_state=42)\n    \n    # Create folds\n    folds = []\n    for train_idx, val_idx in skf.split(\n        train_df, \n        train_df['has_severe'], \n        groups=train_df['study_id']\n    ):\n        fold = {\n            'train': train_df.iloc[train_idx]['study_id'].tolist(),\n            'val': train_df.iloc[val_idx]['study_id'].tolist()\n        }\n        folds.append(fold)\n    \n    return folds\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:32:15.628654Z","iopub.execute_input":"2024-11-21T07:32:15.629056Z","iopub.status.idle":"2024-11-21T07:32:15.635459Z","shell.execute_reply.started":"2024-11-21T07:32:15.629026Z","shell.execute_reply":"2024-11-21T07:32:15.634563Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_fold_data(\n    study_ids: List[int],\n    base_path: str,\n    coords_df: pd.DataFrame,\n    series_df: pd.DataFrame,\n    train_df: pd.DataFrame,\n    augment: bool = False\n) -> List[Dict]:\n    \"\"\"Process data for a set of studies\"\"\"\n    processed_data = []\n    \n    for study_id in study_ids:\n        # Get series information for this study\n        study_series = series_df[series_df['study_id'] == study_id]\n        \n        # Process each condition type with its corresponding series\n        processed_data.extend(\n            process_study_images(\n                study_id=study_id,\n                base_path=base_path,\n                study_series=study_series,\n                coords_df=coords_df,\n                train_df=train_df,\n                augment=augment\n            )\n        )\n    \n    return processed_data\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:32:17.966561Z","iopub.execute_input":"2024-11-21T07:32:17.967309Z","iopub.status.idle":"2024-11-21T07:32:17.972457Z","shell.execute_reply.started":"2024-11-21T07:32:17.967277Z","shell.execute_reply":"2024-11-21T07:32:17.971580Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_study_images(\n    study_id: int,\n    base_path: str,\n    study_series: pd.DataFrame,\n    coords_df: pd.DataFrame,\n    train_df: pd.DataFrame,\n    augment: bool = False\n) -> List[Dict]:\n    \"\"\"Process images for a single study\"\"\"\n    processed_samples = []\n    study_path = Path(base_path) / 'train_images' / str(study_id)\n    \n    # Get study data\n    study_coords = coords_df[coords_df['study_id'] == study_id]\n    study_labels = train_df[train_df['study_id'] == study_id]\n    \n    # Process each type of condition with its corresponding series type\n    condition_series_map = {\n        'Spinal Canal Stenosis': 'Sagittal T2',\n        'Neural Foraminal Narrowing': 'Sagittal T1',\n        'Subarticular Stenosis': 'Axial T2'\n    }\n    \n    for condition, series_type in condition_series_map.items():\n        # Get relevant series\n        series = study_series[\n            study_series['series_description'].str.contains(series_type, na=False)\n        ]\n        \n        for _, series_row in series.iterrows():\n            series_id = series_row['series_id']\n            series_path = study_path / str(series_id)\n            \n            # Get coordinates for this series\n            series_coords = study_coords[\n                (study_coords['series_id'] == series_id) &\n                (study_coords['condition'].str.contains(condition, na=False))\n            ]\n            \n            # Process each image in the series\n            for _, coord_row in series_coords.iterrows():\n                # Get severity from labels\n                label_col = f\"{coord_row['condition'].lower().replace(' ', '_')}_{coord_row['level'].lower().replace('/', '_')}\"\n                if label_col in study_labels.columns:\n                    severity = study_labels[label_col].iloc[0]\n                    \n                    # Process image\n                    processed = process_single_image(\n                        series_path=series_path,\n                        coord_row=coord_row,\n                        severity=severity,\n                        augment=augment\n                    )\n                    if processed:\n                        processed_samples.extend(processed)\n    \n    return processed_samples\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:32:20.600319Z","iopub.execute_input":"2024-11-21T07:32:20.601136Z","iopub.status.idle":"2024-11-21T07:32:20.608744Z","shell.execute_reply.started":"2024-11-21T07:32:20.601098Z","shell.execute_reply":"2024-11-21T07:32:20.607880Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def process_single_image(\n    series_path: Path,\n    coord_row: pd.Series,\n    severity: str,\n    augment: bool = False\n) -> List[Dict]:\n    \"\"\"Process a single image and its annotation\"\"\"\n    processed_samples = []\n    \n    # Load DICOM image\n    dcm_path = series_path / f\"{int(coord_row['instance_number'])}.dcm\"\n    if not dcm_path.exists():\n        return None\n    \n    try:\n        dcm = pydicom.dcmread(str(dcm_path))\n        image = dcm.pixel_array\n        \n        # Preprocess image\n        processed_image = preprocess_image(\n            image,\n            int(coord_row['x']),\n            int(coord_row['y'])\n        )\n        \n        # Create base sample\n        sample = {\n            'image': processed_image,\n            'study_id': coord_row['study_id'],\n            'condition': coord_row['condition'],\n            'level': coord_row['level'].replace('/', '_'),\n            'severity': severity,\n            'coordinates': (coord_row['x'], coord_row['y'])\n        }\n        \n        processed_samples.append(sample)\n        \n        # Apply augmentations if required\n        if augment and severity == 'Severe':\n            augmented_samples = apply_augmentations(sample)\n            processed_samples.extend(augmented_samples)\n            \n        return processed_samples\n    except:\n        return None\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:32:22.991377Z","iopub.execute_input":"2024-11-21T07:32:22.992292Z","iopub.status.idle":"2024-11-21T07:32:23.001244Z","shell.execute_reply.started":"2024-11-21T07:32:22.992241Z","shell.execute_reply":"2024-11-21T07:32:23.000259Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess_image(\n    image: np.ndarray,\n    x: int,\n    y: int,\n    patch_size: int = 224\n) -> np.ndarray:\n    \"\"\"Preprocess a single image\"\"\"\n    # Normalize pixel values\n    normalized = cv2.normalize(\n        image,\n        None,\n        alpha=0,\n        beta=255,\n        norm_type=cv2.NORM_MINMAX,\n        dtype=cv2.CV_8U\n    )\n    \n    # Extract patch around the annotation point\n    half_size = patch_size // 2\n    \n    # Pad image if necessary\n    padded = np.pad(\n        normalized,\n        ((half_size, half_size), (half_size, half_size)),\n        mode='constant',\n        constant_values=0\n    )\n    \n    # Extract patch\n    x_start = x + half_size - half_size\n    x_end = x + half_size + half_size\n    y_start = y + half_size - half_size\n    y_end = y + half_size + half_size\n    \n    patch = padded[y_start:y_end, x_start:x_end]\n    \n    # Ensure patch size\n    if patch.shape != (patch_size, patch_size):\n        patch = cv2.resize(patch, (patch_size, patch_size))\n    \n    # Add channel dimension\n    patch = np.expand_dims(patch, axis=-1)\n    \n    return patch","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:32:25.431297Z","iopub.execute_input":"2024-11-21T07:32:25.431692Z","iopub.status.idle":"2024-11-21T07:32:25.438252Z","shell.execute_reply.started":"2024-11-21T07:32:25.431659Z","shell.execute_reply":"2024-11-21T07:32:25.437357Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def apply_augmentations(sample: Dict) -> List[Dict]:\n    \"\"\"Apply augmentations to balance severe cases\"\"\"\n    augmented_samples = []\n    \n    # Add rotated versions\n    for angle in [90, 180, 270]:\n        aug_sample = sample.copy()\n        aug_sample['image'] = np.rot90(sample['image'], k=angle//90)\n        augmented_samples.append(aug_sample)\n    \n    # Add flipped versions\n    aug_sample = sample.copy()\n    aug_sample['image'] = np.flip(sample['image'], axis=1)\n    augmented_samples.append(aug_sample)\n    \n    return augmented_samples\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:32:27.797805Z","iopub.execute_input":"2024-11-21T07:32:27.798491Z","iopub.status.idle":"2024-11-21T07:32:27.803935Z","shell.execute_reply.started":"2024-11-21T07:32:27.798458Z","shell.execute_reply":"2024-11-21T07:32:27.802908Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    \"\"\"Main function for preprocessing pipeline\"\"\"\n    # Set paths\n    base_path = '/kaggle/input/rsna-2024-lumbar-spine-degenerative-classification'\n    \n    # Load data\n    train_df = pd.read_csv(f\"{base_path}/train.csv\")\n    coords_df = pd.read_csv(f\"{base_path}/train_label_coordinates.csv\")\n    series_df = pd.read_csv(f\"{base_path}/train_series_descriptions.csv\")\n    \n    # Create train-validation split\n    print(\"Creating stratified fold splits...\")\n    folds = create_stratified_folds(train_df)\n    \n    # Process example fold\n    fold_num = 0\n    train_studies = folds[fold_num]['train']\n    val_studies = folds[fold_num]['val']\n    \n    print(f\"\\nProcessing fold {fold_num}...\")\n    # Process training data\n    train_samples = process_fold_data(\n        study_ids=train_studies,\n        base_path=base_path,\n        coords_df=coords_df,\n        series_df=series_df,\n        train_df=train_df,\n        augment=True\n    )\n    \n    # Process validation data\n    val_samples = process_fold_data(\n        study_ids=val_studies,\n        base_path=base_path,\n        coords_df=coords_df,\n        series_df=series_df,\n        train_df=train_df,\n        augment=False\n    )\n    \n    print(f\"Processed {len(train_samples)} training samples and {len(val_samples)} validation samples\")\n    \n    # Save processed data\n    print(\"\\nSaving processed data...\")\n    save_processed_data(train_samples, '/kaggle/working/train_processed.npy')\n    save_processed_data(val_samples, '/kaggle/working/val_processed.npy')\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:32:29.924072Z","iopub.execute_input":"2024-11-21T07:32:29.924744Z","iopub.status.idle":"2024-11-21T07:32:29.930992Z","shell.execute_reply.started":"2024-11-21T07:32:29.924708Z","shell.execute_reply":"2024-11-21T07:32:29.930031Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:32:33.027140Z","iopub.execute_input":"2024-11-21T07:32:33.027853Z","iopub.status.idle":"2024-11-21T07:48:22.066015Z","shell.execute_reply.started":"2024-11-21T07:32:33.027817Z","shell.execute_reply":"2024-11-21T07:48:22.064973Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Classification Model","metadata":{}},{"cell_type":"markdown","source":"I start with the classification model first. Here's why:\n\n* The main evaluation metric emphasizes detecting \"any_severe_spinal\" conditions, which is a classification task\n* Classification model will help us understand the key features for identifying severity levels\n* The learned features from classification can inform the regression model development later","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim.lr_scheduler \nimport timm\nimport numpy as np\nfrom typing import Dict, List, Tuple\nimport pandas as pd\nimport pydicom\nfrom pathlib import Path\nimport cv2","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:48:43.838801Z","iopub.execute_input":"2024-11-21T07:48:43.839378Z","iopub.status.idle":"2024-11-21T07:48:49.642867Z","shell.execute_reply.started":"2024-11-21T07:48:43.839342Z","shell.execute_reply":"2024-11-21T07:48:49.641464Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class LumbarSpineDataset(Dataset):\n    \"\"\"Dataset class for Lumbar Spine Classification\"\"\"\n    def __init__(self, samples: List[Dict], augment: bool = False):\n        # Filter out samples with nan severity\n        self.samples = [\n            sample for sample in samples \n            if isinstance(sample['severity'], str) and not pd.isna(sample['severity'])\n        ]\n        self.augment = augment\n        \n        # Map conditions to indices\n        self.condition_map = {\n            'Spinal Canal Stenosis': 0,\n            'Left Neural Foraminal Narrowing': 1,\n            'Right Neural Foraminal Narrowing': 2,\n            'Left Subarticular Stenosis': 3,\n            'Right Subarticular Stenosis': 4\n        }\n        \n        # Map levels to indices\n        self.level_map = {\n            'L1_L2': 0, 'L2_L3': 1, 'L3_L4': 2, 'L4_L5': 3, 'L5_S1': 4\n        }\n        \n        # Map severity to indices\n        self.severity_map = {\n            'Normal/Mild': 0,\n            'Moderate': 1,\n            'Severe': 2\n        }\n        \n        print(f\"After filtering nan values: {len(self.samples)} valid samples\")\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        \n        # Convert image to torch tensor\n        image = torch.from_numpy(sample['image']).float()\n        image = image.permute(2, 0, 1)  # CHW format\n        \n        # Create condition one-hot encoding\n        condition = torch.zeros(len(self.condition_map))\n        condition[self.condition_map[sample['condition']]] = 1\n        \n        # Create level one-hot encoding\n        level = torch.zeros(len(self.level_map))\n        level[self.level_map[sample['level']]] = 1\n        \n        # Create severity label\n        severity = torch.tensor(self.severity_map[sample['severity']])\n        \n        return {\n            'image': image,\n            'condition': condition,\n            'level': level,\n            'severity': severity,\n            'study_id': sample['study_id']\n        }","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:48:49.644690Z","iopub.execute_input":"2024-11-21T07:48:49.645064Z","iopub.status.idle":"2024-11-21T07:48:49.659798Z","shell.execute_reply.started":"2024-11-21T07:48:49.645023Z","shell.execute_reply":"2024-11-21T07:48:49.658497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AttentionBlock(nn.Module):\n    def __init__(self, in_channels):\n        super().__init__()\n        self.attention = nn.Sequential(\n            nn.Conv2d(in_channels, in_channels // 8, 1),\n            nn.ReLU(),\n            nn.Conv2d(in_channels // 8, in_channels, 1),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        attention_weights = self.attention(x)\n        return x * attention_weights\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:48:52.683557Z","iopub.execute_input":"2024-11-21T07:48:52.683913Z","iopub.status.idle":"2024-11-21T07:48:52.689304Z","shell.execute_reply.started":"2024-11-21T07:48:52.683883Z","shell.execute_reply":"2024-11-21T07:48:52.688454Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class LumbarClassifier(nn.Module):\n    def __init__(self, num_classes=3):\n        super().__init__()\n        \n        # Load pretrained EfficientNetV2 backbone\n        self.backbone = timm.create_model(\n            'tf_efficientnetv2_s',\n            pretrained=True,\n            in_chans=1,\n            features_only=True\n        )\n        \n        # Get backbone feature dimensions\n        dummy_input = torch.randn(1, 1, 224, 224)\n        features = self.backbone(dummy_input)\n        feature_dims = [f.shape[1] for f in features]\n        \n        # Attention blocks for each feature level\n        self.attention_blocks = nn.ModuleList([\n            AttentionBlock(dims) for dims in feature_dims\n        ])\n        \n        # Global Average Pooling\n        self.gap = nn.AdaptiveAvgPool2d(1)\n        \n        # Calculate total feature dimensions\n        total_dims = sum(feature_dims)\n        \n        # BiLSTM for sequential feature processing\n        self.bilstm = nn.LSTM(\n            input_size=total_dims,\n            hidden_size=512,\n            num_layers=2,\n            bidirectional=True,\n            batch_first=True\n        )\n        \n        # Condition and level embedding\n        self.condition_embed = nn.Linear(5, 64)  # 5 conditions\n        self.level_embed = nn.Linear(5, 64)      # 5 levels\n        \n        # Final classification layers\n        self.classifier = nn.Sequential(\n            nn.Linear(2*512 + 64 + 64, 512),  # BiLSTM + condition + level\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(256, num_classes)\n        )\n        \n        # Initialize weights\n        self._initialize_weights()\n    \n    def _initialize_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight)\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.Linear):\n                nn.init.kaiming_normal_(m.weight)\n                nn.init.constant_(m.bias, 0)\n    \n    def forward(self, x, condition, level):\n        # Get backbone features\n        features = self.backbone(x)\n        \n        # Apply attention to each feature level\n        attended_features = [\n            att(feat) for feat, att in zip(features, self.attention_blocks)\n        ]\n        \n        # Global average pooling on each feature map\n        pooled_features = [self.gap(feat) for feat in attended_features]\n        \n        # Concatenate features\n        concat_features = torch.cat([\n            feat.view(feat.size(0), -1) for feat in pooled_features\n        ], dim=1)\n        \n        # Reshape for LSTM\n        lstm_out, _ = self.bilstm(\n            concat_features.unsqueeze(1)\n        )\n        lstm_out = lstm_out[:, -1, :]  # Take last output\n        \n        # Embed condition and level\n        condition_embedding = self.condition_embed(condition)\n        level_embedding = self.level_embed(level)\n        \n        # Concatenate all features\n        combined_features = torch.cat([\n            lstm_out,\n            condition_embedding,\n            level_embedding\n        ], dim=1)\n        \n        # Final classification\n        logits = self.classifier(combined_features)\n        \n        return logits","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:48:54.701540Z","iopub.execute_input":"2024-11-21T07:48:54.702350Z","iopub.status.idle":"2024-11-21T07:48:54.714229Z","shell.execute_reply.started":"2024-11-21T07:48:54.702317Z","shell.execute_reply":"2024-11-21T07:48:54.713226Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_epoch(model, train_loader, criterion, optimizer, device):\n    model.train()\n    total_loss = 0\n    correct = 0\n    total = 0\n    \n    for batch in train_loader:\n        images = batch['image'].to(device)\n        conditions = batch['condition'].to(device)\n        levels = batch['level'].to(device)\n        labels = batch['severity'].to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(images, conditions, levels)\n        loss = criterion(outputs, labels)\n        \n        loss.backward()\n        optimizer.step()\n        \n        total_loss += loss.item()\n        _, predicted = outputs.max(1)\n        total += labels.size(0)\n        correct += predicted.eq(labels).sum().item()\n    \n    return total_loss / len(train_loader), 100. * correct / total","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:48:58.753626Z","iopub.execute_input":"2024-11-21T07:48:58.753978Z","iopub.status.idle":"2024-11-21T07:48:58.760189Z","shell.execute_reply.started":"2024-11-21T07:48:58.753950Z","shell.execute_reply":"2024-11-21T07:48:58.759292Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(model, val_loader, criterion, device):\n    model.eval()\n    total_loss = 0\n    correct = 0\n    total = 0\n    \n    with torch.no_grad():\n        for batch in val_loader:\n            images = batch['image'].to(device)\n            conditions = batch['condition'].to(device)\n            levels = batch['level'].to(device)\n            labels = batch['severity'].to(device)\n            \n            outputs = model(images, conditions, levels)\n            loss = criterion(outputs, labels)\n            \n            total_loss += loss.item()\n            _, predicted = outputs.max(1)\n            total += labels.size(0)\n            correct += predicted.eq(labels).sum().item()\n    \n    return total_loss / len(val_loader), 100. * correct / total","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:49:01.044836Z","iopub.execute_input":"2024-11-21T07:49:01.045175Z","iopub.status.idle":"2024-11-21T07:49:01.051309Z","shell.execute_reply.started":"2024-11-21T07:49:01.045147Z","shell.execute_reply":"2024-11-21T07:49:01.050382Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_processed_data(filename):\n    \"\"\"Load processed samples from numpy file\"\"\"\n    processed_data = np.load(filename, allow_pickle=True).item()\n    \n    # Convert back to list of dictionaries format\n    samples = []\n    for i in range(len(processed_data['images'])):\n        sample = {\n            'image': processed_data['images'][i],\n            'condition': processed_data['conditions'][i],\n            'level': processed_data['levels'][i],\n            'severity': processed_data['severities'][i],\n            'study_id': processed_data['study_ids'][i]\n        }\n        samples.append(sample)\n    \n    return samples","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:49:03.552132Z","iopub.execute_input":"2024-11-21T07:49:03.552482Z","iopub.status.idle":"2024-11-21T07:49:03.558115Z","shell.execute_reply.started":"2024-11-21T07:49:03.552449Z","shell.execute_reply":"2024-11-21T07:49:03.557021Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Set device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"Using device: {device}\")\n    \n    print(\"Loading processed data...\")\n    train_samples = load_processed_data('/kaggle/working/train_processed.npy')\n    val_samples = load_processed_data('/kaggle/working/val_processed.npy')\n    print(f\"Loaded {len(train_samples)} training samples and {len(val_samples)} validation samples\")\n    \n    # Create datasets and dataloaders\n    train_dataset = LumbarSpineDataset(train_samples, augment=True)\n    val_dataset = LumbarSpineDataset(val_samples, augment=False)\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=32,\n        shuffle=True,\n        num_workers=4\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=32,\n        shuffle=False,\n        num_workers=4\n    )\n    \n    # Create model\n    model = LumbarClassifier().to(device)\n    print(\"Created model and moved to device\")\n    \n    # Define loss function with class weights to handle imbalance\n    # Weights are inversely proportional to class frequencies\n    # [Normal/Mild: 77.4%, Moderate: 16.3%, Severe: 6.3%]\n    weights = torch.tensor([1.0, 4.75, 12.29]).to(device)  # Normalized inverse frequencies\n    criterion = nn.CrossEntropyLoss(weight=weights)\n    \n    # Define optimizer and scheduler\n    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.01)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=30, eta_min=1e-6)\n    \n    # Training parameters\n    num_epochs = 30\n    best_val_acc = 0\n    best_val_loss = float('inf')\n    patience = 5\n    patience_counter = 0\n    \n    print(\"Starting training...\")\n    for epoch in range(num_epochs):\n        print(f\"\\nEpoch {epoch+1}/{num_epochs}\")\n        print(\"-\" * 20)\n        \n        # Training phase\n        train_loss, train_acc = train_epoch(\n            model=model,\n            train_loader=train_loader,\n            criterion=criterion,\n            optimizer=optimizer,\n            device=device\n        )\n        \n        # Validation phase\n        val_loss, val_acc = validate(\n            model=model,\n            val_loader=val_loader,\n            criterion=criterion,\n            device=device\n        )\n        \n        # Learning rate scheduling\n        scheduler.step()\n        current_lr = scheduler.get_last_lr()[0]\n        \n        # Print epoch results\n        print(f\"Train Loss: {train_loss:.4f}, Train Acc: {train_acc:.2f}%\")\n        print(f\"Val Loss: {val_loss:.4f}, Val Acc: {val_acc:.2f}%\")\n        print(f\"Learning Rate: {current_lr:.6f}\")\n        \n        # Save best model\n        if val_acc > best_val_acc:\n            best_val_acc = val_acc\n            best_val_loss = val_loss\n            print(\"Saving best model...\")\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'scheduler_state_dict': scheduler.state_dict(),\n                'val_acc': val_acc,\n                'val_loss': val_loss,\n            }, 'best_model.pth')\n            patience_counter = 0\n        else:\n            patience_counter += 1\n        \n        # Early stopping\n        if patience_counter >= patience:\n            print(f\"\\nEarly stopping triggered after {epoch+1} epochs\")\n            break\n    \n    print(\"\\nTraining completed!\")\n    print(f\"Best Validation Accuracy: {best_val_acc:.2f}%\")\n    print(f\"Best Validation Loss: {best_val_loss:.4f}\")\n    \n    # Load best model for final evaluation\n    checkpoint = torch.load('best_model.pth')\n    model.load_state_dict(checkpoint['model_state_dict'])\n    \n    # Final validation pass\n    final_val_loss, final_val_acc = validate(model, val_loader, criterion, device)\n    print(\"\\nFinal Model Performance:\")\n    print(f\"Validation Loss: {final_val_loss:.4f}\")\n    print(f\"Validation Accuracy: {final_val_acc:.2f}%\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:49:06.423538Z","iopub.execute_input":"2024-11-21T07:49:06.423970Z","iopub.status.idle":"2024-11-21T07:49:06.440940Z","shell.execute_reply.started":"2024-11-21T07:49:06.423928Z","shell.execute_reply":"2024-11-21T07:49:06.439835Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T07:49:32.571252Z","iopub.execute_input":"2024-11-21T07:49:32.571632Z","iopub.status.idle":"2024-11-21T09:53:04.706046Z","shell.execute_reply.started":"2024-11-21T07:49:32.571600Z","shell.execute_reply":"2024-11-21T09:53:04.704945Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Classification Model Analysis","metadata":{}},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import confusion_matrix, classification_report\nfrom typing import Dict, List, Tuple","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T09:54:50.584580Z","iopub.execute_input":"2024-11-21T09:54:50.584986Z","iopub.status.idle":"2024-11-21T09:54:50.590020Z","shell.execute_reply.started":"2024-11-21T09:54:50.584948Z","shell.execute_reply":"2024-11-21T09:54:50.589176Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def analyze_classification_model(model_path: str, val_loader: DataLoader, device: torch.device):\n    \"\"\"Analyze classification model performance in detail\"\"\"\n    # Load the best model\n    model = LumbarClassifier().to(device)\n    checkpoint = torch.load(model_path)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    model.eval()\n    \n    # Initialize lists to store predictions and true labels\n    all_preds = []\n    all_labels = []\n    all_conditions = []\n    all_levels = []\n    study_ids = []\n    \n    # Get predictions\n    with torch.no_grad():\n        for batch in val_loader:\n            images = batch['image'].to(device)\n            conditions = batch['condition'].to(device)\n            levels = batch['level'].to(device)\n            labels = batch['severity']\n            \n            outputs = model(images, conditions, levels)\n            _, preds = outputs.max(1)\n            \n            all_preds.extend(preds.cpu().numpy())\n            all_labels.extend(labels.numpy())\n            \n            # Get condition and level names\n            condition_indices = torch.argmax(conditions, dim=1).cpu().numpy()\n            level_indices = torch.argmax(levels, dim=1).cpu().numpy()\n            \n            for c_idx, l_idx in zip(condition_indices, level_indices):\n                all_conditions.append(list(val_loader.dataset.condition_map.keys())[c_idx])\n                all_levels.append(list(val_loader.dataset.level_map.keys())[l_idx])\n            \n            study_ids.extend(batch['study_id'])\n    \n    # Create confusion matrix\n    severity_names = ['Normal/Mild', 'Moderate', 'Severe']\n    cm = confusion_matrix(all_labels, all_preds)\n    \n    plt.figure(figsize=(10, 8))\n    sns.heatmap(cm, annot=True, fmt='d', xticklabels=severity_names, yticklabels=severity_names)\n    plt.title('Confusion Matrix')\n    plt.xlabel('Predicted')\n    plt.ylabel('True')\n    plt.savefig('confusion_matrix.png')\n    plt.close()\n    \n    # Generate classification report\n    report = classification_report(all_labels, all_preds, target_names=severity_names)\n    print(\"\\nClassification Report:\")\n    print(report)\n    \n    # Analyze performance by condition and level\n    results_df = pd.DataFrame({\n        'Study_ID': study_ids,\n        'True_Label': [severity_names[l] for l in all_labels],\n        'Predicted': [severity_names[p] for p in all_preds],\n        'Condition': all_conditions,\n        'Level': all_levels,\n        'Correct': [1 if p == l else 0 for p, l in zip(all_preds, all_labels)]\n    })\n    \n    # Performance by condition\n    print(\"\\nAccuracy by Condition:\")\n    condition_acc = results_df.groupby('Condition')['Correct'].mean() * 100\n    print(condition_acc)\n    \n    # Performance by level\n    print(\"\\nAccuracy by Level:\")\n    level_acc = results_df.groupby('Level')['Correct'].mean() * 100\n    print(level_acc)\n    \n    # Save detailed results\n    results_df.to_csv('classification_results.csv', index=False)\n    \n    # Plot accuracy by condition and level\n    plt.figure(figsize=(12, 5))\n    plt.subplot(1, 2, 1)\n    condition_acc.plot(kind='bar')\n    plt.title('Accuracy by Condition')\n    plt.xticks(rotation=45)\n    plt.tight_layout()\n    \n    plt.subplot(1, 2, 2)\n    level_acc.plot(kind='bar')\n    plt.title('Accuracy by Level')\n    plt.xticks(rotation=45)\n    plt.tight_layout()\n    \n    plt.savefig('accuracy_analysis.png')\n    plt.close()\n    \n    return results_df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T09:54:53.479311Z","iopub.execute_input":"2024-11-21T09:54:53.479968Z","iopub.status.idle":"2024-11-21T09:54:53.492852Z","shell.execute_reply.started":"2024-11-21T09:54:53.479934Z","shell.execute_reply":"2024-11-21T09:54:53.492026Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Load the same data and model setup as before\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Load validation data\n    val_samples = load_processed_data('/kaggle/working/val_processed.npy')\n    val_dataset = LumbarSpineDataset(val_samples, augment=False)\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=32,\n        shuffle=False,\n        num_workers=4\n    )\n    \n    # Analyze model\n    results_df = analyze_classification_model('best_model.pth', val_loader, device)\n    print(\"\\nAnalysis completed and saved to files.\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T09:54:56.426269Z","iopub.execute_input":"2024-11-21T09:54:56.426625Z","iopub.status.idle":"2024-11-21T09:55:24.033403Z","shell.execute_reply.started":"2024-11-21T09:54:56.426593Z","shell.execute_reply":"2024-11-21T09:55:24.032440Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Based on the results, we can observe several important insights before moving to the regression model:\n\nClass Imbalance:\n\n* Normal/Mild: 7607 samples (78.3%)\n* Moderate: 1528 samples (15.7%)\n* Severe: 586 samples (6%)\n\nPerformance by Severity:\n\n* Normal/Mild: Strong performance (F1: 0.92)\n* Moderate: Struggles most (F1: 0.52)\n* Severe: Better than moderate but still challenging (F1: 0.60)\n\nPerformance by Condition:\n\n* Best: Spinal Canal Stenosis (89.64%)\n* Others: Relatively consistent (~78-84%)\n\nPerformance by Level:\n\n* Best: Upper levels (L1_L2: 95.96%, L2_L3: 89.06%)\n* Worst: L4_L5 (74.32%)\n* Gradient of difficulty from top to bottom","metadata":{}},{"cell_type":"markdown","source":"# Regression Model","metadata":{}},{"cell_type":"code","source":"import torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nimport torch.optim.lr_scheduler\nimport timm\nimport numpy as np\nfrom typing import Dict, List, Tuple","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T09:58:37.355607Z","iopub.execute_input":"2024-11-21T09:58:37.356507Z","iopub.status.idle":"2024-11-21T09:58:37.361148Z","shell.execute_reply.started":"2024-11-21T09:58:37.356455Z","shell.execute_reply":"2024-11-21T09:58:37.360315Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class LumbarSpineRegDataset(Dataset):\n    \"\"\"Dataset class for Lumbar Spine Regression\"\"\"\n    def __init__(self, samples: List[Dict], augment: bool = False):\n        # Filter out samples with nan severity\n        self.samples = [\n            sample for sample in samples \n            if isinstance(sample['severity'], str) and not pd.isna(sample['severity'])\n        ]\n        self.augment = augment\n        \n        # Map conditions to indices\n        self.condition_map = {\n            'Spinal Canal Stenosis': 0,\n            'Left Neural Foraminal Narrowing': 1,\n            'Right Neural Foraminal Narrowing': 2,\n            'Left Subarticular Stenosis': 3,\n            'Right Subarticular Stenosis': 4\n        }\n        \n        # Map levels to indices\n        self.level_map = {\n            'L1_L2': 0, 'L2_L3': 1, 'L3_L4': 2, 'L4_L5': 3, 'L5_S1': 4\n        }\n        \n        # Map severity to continuous values\n        self.severity_map = {\n            'Normal/Mild': 0.0,\n            'Moderate': 1.0,\n            'Severe': 2.0\n        }\n        \n        print(f\"After filtering nan values: {len(self.samples)} valid samples\")\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        sample = self.samples[idx]\n        \n        # Convert image to torch tensor\n        image = torch.from_numpy(sample['image']).float()\n        image = image.permute(2, 0, 1)  # CHW format\n        \n        # Create condition one-hot encoding\n        condition = torch.zeros(len(self.condition_map))\n        condition[self.condition_map[sample['condition']]] = 1\n        \n        # Create level one-hot encoding\n        level = torch.zeros(len(self.level_map))\n        level[self.level_map[sample['level']]] = 1\n        \n        # Create severity value\n        severity = torch.tensor([self.severity_map[sample['severity']]], dtype=torch.float)\n        \n        # Create importance weight based on level and severity\n        weight = 1.0\n        if sample['level'] in ['L4_L5', 'L5_S1']:\n            weight *= 1.5\n        if sample['severity'] in ['Moderate', 'Severe']:\n            weight *= 2.0\n        \n        return {\n            'image': image,\n            'condition': condition,\n            'level': level,\n            'severity': severity,\n            'weight': torch.tensor([weight], dtype=torch.float),\n            'study_id': sample['study_id']\n        }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T09:58:39.692174Z","iopub.execute_input":"2024-11-21T09:58:39.693024Z","iopub.status.idle":"2024-11-21T09:58:39.702397Z","shell.execute_reply.started":"2024-11-21T09:58:39.692984Z","shell.execute_reply":"2024-11-21T09:58:39.701502Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class LumbarRegressor(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n        # Load pretrained EfficientNetV2 backbone\n        self.backbone = timm.create_model(\n            'tf_efficientnetv2_s',\n            pretrained=True,\n            in_chans=1,\n            features_only=True\n        )\n        \n        # Get backbone feature dimensions\n        dummy_input = torch.randn(1, 1, 224, 224)\n        features = self.backbone(dummy_input)\n        feature_dims = [f.shape[1] for f in features]\n        \n        # Attention blocks for each feature level\n        self.attention_blocks = nn.ModuleList([\n            AttentionBlock(dims) for dims in feature_dims\n        ])\n        \n        # Global Average Pooling\n        self.gap = nn.AdaptiveAvgPool2d(1)\n        \n        # Calculate total feature dimensions\n        total_dims = sum(feature_dims)\n        \n        # BiLSTM for sequential feature processing\n        self.bilstm = nn.LSTM(\n            input_size=total_dims,\n            hidden_size=512,\n            num_layers=2,\n            bidirectional=True,\n            batch_first=True\n        )\n        \n        # Condition and level embedding\n        self.condition_embed = nn.Linear(5, 64)  # 5 conditions\n        self.level_embed = nn.Linear(5, 64)      # 5 levels\n        \n        # Regression head\n        self.regressor = nn.Sequential(\n            nn.Linear(2*512 + 64 + 64, 512),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(512, 256),\n            nn.ReLU(),\n            nn.Dropout(0.2),\n            nn.Linear(256, 1),\n            nn.Sigmoid()  # Output between 0 and 1\n        )\n        \n        self._initialize_weights()\n    \n    def _initialize_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv2d):\n                nn.init.kaiming_normal_(m.weight)\n            elif isinstance(m, nn.BatchNorm2d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.Linear):\n                nn.init.kaiming_normal_(m.weight)\n                nn.init.constant_(m.bias, 0)\n    \n    def forward(self, x, condition, level):\n        # Get backbone features\n        features = self.backbone(x)\n        \n        # Apply attention to each feature level\n        attended_features = [\n            att(feat) for feat, att in zip(features, self.attention_blocks)\n        ]\n        \n        # Global average pooling on each feature map\n        pooled_features = [self.gap(feat) for feat in attended_features]\n        \n        # Concatenate features\n        concat_features = torch.cat([\n            feat.view(feat.size(0), -1) for feat in pooled_features\n        ], dim=1)\n        \n        # Reshape for LSTM\n        lstm_out, _ = self.bilstm(\n            concat_features.unsqueeze(1)\n        )\n        lstm_out = lstm_out[:, -1, :]  # Take last output\n        \n        # Embed condition and level\n        condition_embedding = self.condition_embed(condition)\n        level_embedding = self.level_embed(level)\n        \n        # Concatenate all features\n        combined_features = torch.cat([\n            lstm_out,\n            condition_embedding,\n            level_embedding\n        ], dim=1)\n        \n        # Regression output\n        return self.regressor(combined_features) * 2.0  # Scale to 0-2 range\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T09:58:41.912166Z","iopub.execute_input":"2024-11-21T09:58:41.912537Z","iopub.status.idle":"2024-11-21T09:58:41.925679Z","shell.execute_reply.started":"2024-11-21T09:58:41.912475Z","shell.execute_reply":"2024-11-21T09:58:41.924757Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class WeightedL1Loss(nn.Module):\n    def __init__(self):\n        super().__init__()\n        \n    def forward(self, pred, target, weight):\n        return (weight * torch.abs(pred - target)).mean()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T09:58:44.317749Z","iopub.execute_input":"2024-11-21T09:58:44.318465Z","iopub.status.idle":"2024-11-21T09:58:44.322897Z","shell.execute_reply.started":"2024-11-21T09:58:44.318429Z","shell.execute_reply":"2024-11-21T09:58:44.322006Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def train_epoch(model, train_loader, criterion, optimizer, device):\n    model.train()\n    total_loss = 0\n    \n    for batch in train_loader:\n        images = batch['image'].to(device)\n        conditions = batch['condition'].to(device)\n        levels = batch['level'].to(device)\n        targets = batch['severity'].to(device)\n        weights = batch['weight'].to(device)\n        \n        optimizer.zero_grad()\n        outputs = model(images, conditions, levels)\n        loss = criterion(outputs, targets, weights)\n        \n        loss.backward()\n        optimizer.step()\n        \n        total_loss += loss.item()\n    \n    return total_loss / len(train_loader)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T09:58:46.532639Z","iopub.execute_input":"2024-11-21T09:58:46.533495Z","iopub.status.idle":"2024-11-21T09:58:46.539200Z","shell.execute_reply.started":"2024-11-21T09:58:46.533452Z","shell.execute_reply":"2024-11-21T09:58:46.538188Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def validate(model, val_loader, criterion, device):\n    model.eval()\n    total_loss = 0\n    all_preds = []\n    all_targets = []\n    \n    with torch.no_grad():\n        for batch in val_loader:\n            images = batch['image'].to(device)\n            conditions = batch['condition'].to(device)\n            levels = batch['level'].to(device)\n            targets = batch['severity'].to(device)\n            weights = batch['weight'].to(device)\n            \n            outputs = model(images, conditions, levels)\n            loss = criterion(outputs, targets, weights)\n            \n            total_loss += loss.item()\n            all_preds.extend(outputs.cpu().numpy())\n            all_targets.extend(targets.cpu().numpy())\n    \n    mae = np.mean(np.abs(np.array(all_preds) - np.array(all_targets)))\n    return total_loss / len(val_loader), mae","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T09:58:49.047692Z","iopub.execute_input":"2024-11-21T09:58:49.048083Z","iopub.status.idle":"2024-11-21T09:58:49.054396Z","shell.execute_reply.started":"2024-11-21T09:58:49.048052Z","shell.execute_reply":"2024-11-21T09:58:49.053567Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Set device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"Using device: {device}\")\n    \n    # Load data\n    print(\"Loading processed data...\")\n    train_samples = load_processed_data('/kaggle/working/train_processed.npy')\n    val_samples = load_processed_data('/kaggle/working/val_processed.npy')\n    \n    # Create datasets and dataloaders\n    train_dataset = LumbarSpineRegDataset(train_samples, augment=True)\n    val_dataset = LumbarSpineRegDataset(val_samples, augment=False)\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=32,\n        shuffle=True,\n        num_workers=4\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=32,\n        shuffle=False,\n        num_workers=4\n    )\n    \n    # Create model\n    model = LumbarRegressor().to(device)\n    criterion = WeightedL1Loss()\n    optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4, weight_decay=0.01)\n    scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max=30, eta_min=1e-6)\n    \n    # Training parameters\n    num_epochs = 30\n    best_val_loss = float('inf')\n    patience = 5\n    patience_counter = 0\n    \n    print(\"Starting training...\")\n    for epoch in range(num_epochs):\n        print(f\"\\nEpoch {epoch+1}/{num_epochs}\")\n        print(\"-\" * 20)\n        \n        train_loss = train_epoch(model, train_loader, criterion, optimizer, device)\n        val_loss, val_mae = validate(model, val_loader, criterion, device)\n        scheduler.step()\n        \n        print(f\"Train Loss: {train_loss:.4f}\")\n        print(f\"Val Loss: {val_loss:.4f}, Val MAE: {val_mae:.4f}\")\n        \n        if val_loss < best_val_loss:\n            best_val_loss = val_loss\n            print(\"Saving best model...\")\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_loss': val_loss,\n                'val_mae': val_mae\n            }, 'best_regression_model.pth')\n            patience_counter = 0\n        else:\n            patience_counter += 1\n        \n        if patience_counter >= patience:\n            print(f\"\\nEarly stopping triggered after {epoch+1} epochs\")\n            break\n    \n    print(\"\\nTraining completed!\")\n    print(f\"Best Validation Loss: {best_val_loss:.4f}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T09:58:51.981326Z","iopub.execute_input":"2024-11-21T09:58:51.982081Z","iopub.status.idle":"2024-11-21T09:58:51.991024Z","shell.execute_reply.started":"2024-11-21T09:58:51.982035Z","shell.execute_reply":"2024-11-21T09:58:51.990143Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"if __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T09:58:54.935936Z","iopub.execute_input":"2024-11-21T09:58:54.936257Z","iopub.status.idle":"2024-11-21T12:21:31.265660Z","shell.execute_reply.started":"2024-11-21T09:58:54.936229Z","shell.execute_reply":"2024-11-21T12:21:31.264504Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Regression Model Analysis","metadata":{}},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport pandas as pd\nimport seaborn as sns\nimport matplotlib.pyplot as plt\nfrom sklearn.metrics import mean_absolute_error, mean_squared_error, r2_score\nfrom typing import Dict, List, Tuple","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T12:21:47.009638Z","iopub.execute_input":"2024-11-21T12:21:47.010019Z","iopub.status.idle":"2024-11-21T12:21:47.014987Z","shell.execute_reply.started":"2024-11-21T12:21:47.009986Z","shell.execute_reply":"2024-11-21T12:21:47.014174Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def analyze_regression_model(model_path: str, val_loader: DataLoader, device: torch.device):\n    \"\"\"Analyze regression model performance in detail\"\"\"\n    # Load the best model\n    model = LumbarRegressor().to(device)\n    checkpoint = torch.load(model_path)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    model.eval()\n    \n    # Initialize lists to store predictions and true values\n    predictions = []\n    true_values = []\n    conditions = []\n    levels = []\n    study_ids = []\n    weights = []\n    \n    # Get predictions\n    with torch.no_grad():\n        for batch in val_loader:\n            images = batch['image'].to(device)\n            condition_tensor = batch['condition'].to(device)\n            level_tensor = batch['level'].to(device)\n            targets = batch['severity']\n            \n            outputs = model(images, condition_tensor, level_tensor)\n            \n            predictions.extend(outputs.cpu().numpy())\n            true_values.extend(targets.numpy())\n            weights.extend(batch['weight'].numpy())\n            \n            # Get condition and level names\n            condition_indices = torch.argmax(condition_tensor, dim=1).cpu().numpy()\n            level_indices = torch.argmax(level_tensor, dim=1).cpu().numpy()\n            \n            for c_idx, l_idx in zip(condition_indices, level_indices):\n                conditions.append(list(val_loader.dataset.condition_map.keys())[c_idx])\n                levels.append(list(val_loader.dataset.level_map.keys())[l_idx])\n            \n            study_ids.extend(batch['study_id'])\n    \n    # Convert to numpy arrays\n    predictions = np.array(predictions)\n    true_values = np.array(true_values)\n    \n    # Calculate metrics\n    mae = mean_absolute_error(true_values, predictions)\n    rmse = np.sqrt(mean_squared_error(true_values, predictions))\n    r2 = r2_score(true_values, predictions)\n    \n    print(\"\\nOverall Metrics:\")\n    print(f\"Mean Absolute Error: {mae:.4f}\")\n    print(f\"Root Mean Squared Error: {rmse:.4f}\")\n    print(f\"R² Score: {r2:.4f}\")\n    \n    # Create DataFrame for detailed analysis\n    results_df = pd.DataFrame({\n        'Study_ID': study_ids,\n        'True_Value': true_values.flatten(),\n        'Predicted': predictions.flatten(),\n        'Condition': conditions,\n        'Level': levels,\n        'Weight': weights,\n        'Absolute_Error': np.abs(predictions.flatten() - true_values.flatten())\n    })\n    \n    # Analyze by condition\n    print(\"\\nMAE by Condition:\")\n    condition_mae = results_df.groupby('Condition')['Absolute_Error'].mean()\n    print(condition_mae)\n    \n    # Analyze by level\n    print(\"\\nMAE by Level:\")\n    level_mae = results_df.groupby('Level')['Absolute_Error'].mean()\n    print(level_mae)\n    \n    # Plot prediction vs true values\n    plt.figure(figsize=(10, 8))\n    plt.scatter(true_values, predictions, alpha=0.5)\n    plt.plot([0, 2], [0, 2], 'r--')  # Perfect prediction line\n    plt.xlabel('True Values')\n    plt.ylabel('Predictions')\n    plt.title('Prediction vs True Values')\n    plt.savefig('regression_scatter.png')\n    plt.close()\n    \n    # Plot error distribution\n    plt.figure(figsize=(10, 6))\n    sns.histplot(data=results_df['Absolute_Error'], bins=50)\n    plt.title('Error Distribution')\n    plt.xlabel('Absolute Error')\n    plt.ylabel('Count')\n    plt.savefig('error_distribution.png')\n    plt.close()\n    \n    # Plot MAE by condition\n    plt.figure(figsize=(12, 5))\n    plt.subplot(1, 2, 1)\n    condition_mae.plot(kind='bar')\n    plt.title('MAE by Condition')\n    plt.xticks(rotation=45)\n    plt.tight_layout()\n    \n    # Plot MAE by level\n    plt.subplot(1, 2, 2)\n    level_mae.plot(kind='bar')\n    plt.title('MAE by Level')\n    plt.xticks(rotation=45)\n    plt.tight_layout()\n    plt.savefig('mae_analysis.png')\n    plt.close()\n    \n    # Error analysis by severity range\n    def get_severity_range(value):\n        if value <= 0.5:\n            return 'Normal/Mild'\n        elif value <= 1.5:\n            return 'Moderate'\n        else:\n            return 'Severe'\n    \n    results_df['Severity_Range'] = results_df['True_Value'].apply(get_severity_range)\n    \n    print(\"\\nMAE by Severity Range:\")\n    severity_mae = results_df.groupby('Severity_Range')['Absolute_Error'].mean()\n    print(severity_mae)\n    \n    # Save detailed results\n    results_df.to_csv('regression_results.csv', index=False)\n    \n    # Compare with classification model performance\n    print(\"\\nSeverity Range Distribution:\")\n    print(results_df['Severity_Range'].value_counts(normalize=True) * 100)\n    \n    return results_df\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T12:21:49.862503Z","iopub.execute_input":"2024-11-21T12:21:49.862863Z","iopub.status.idle":"2024-11-21T12:21:49.878359Z","shell.execute_reply.started":"2024-11-21T12:21:49.862834Z","shell.execute_reply":"2024-11-21T12:21:49.877367Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Set device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Load validation data\n    val_samples = load_processed_data('/kaggle/working/val_processed.npy')\n    val_dataset = LumbarSpineRegDataset(val_samples, augment=False)\n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=32,\n        shuffle=False,\n        num_workers=4\n    )\n    \n    # Analyze model\n    results_df = analyze_regression_model('best_regression_model.pth', val_loader, device)\n    print(\"\\nAnalysis completed and saved to files.\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T12:21:53.017399Z","iopub.execute_input":"2024-11-21T12:21:53.018274Z","iopub.status.idle":"2024-11-21T12:22:22.143159Z","shell.execute_reply.started":"2024-11-21T12:21:53.018237Z","shell.execute_reply":"2024-11-21T12:22:22.142142Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Evaluation Metrics","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import log_loss\nfrom typing import Dict, List, Tuple","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T12:26:18.104954Z","iopub.execute_input":"2024-11-21T12:26:18.105324Z","iopub.status.idle":"2024-11-21T12:26:18.109859Z","shell.execute_reply.started":"2024-11-21T12:26:18.105293Z","shell.execute_reply":"2024-11-21T12:26:18.108890Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_processed_data(filename: str) -> List[Dict]:\n    \"\"\"Load processed samples from numpy file\"\"\"\n    processed_data = np.load(filename, allow_pickle=True).item()\n    \n    # Convert back to list of dictionaries format\n    samples = []\n    for i in range(len(processed_data['images'])):\n        sample = {\n            'image': processed_data['images'][i],\n            'condition': processed_data['conditions'][i],\n            'level': processed_data['levels'][i],\n            'severity': processed_data['severities'][i],\n            'study_id': processed_data['study_ids'][i]\n        }\n        samples.append(sample)\n    \n    return samples","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T12:46:11.391488Z","iopub.execute_input":"2024-11-21T12:46:11.391878Z","iopub.status.idle":"2024-11-21T12:46:11.397456Z","shell.execute_reply.started":"2024-11-21T12:46:11.391848Z","shell.execute_reply":"2024-11-21T12:46:11.396574Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def compute_competition_metric(predictions: np.ndarray,\n                             true_labels: np.ndarray,\n                             study_ids: np.ndarray,\n                             conditions: List[str],\n                             levels: List[str]) -> Tuple[float, Dict]:\n    \"\"\"\n    Compute the competition metric: average of sample weighted log losses and any_severe_spinal\n    \"\"\"\n    # Create DataFrame with all information\n    results_df = pd.DataFrame({\n        'study_id': study_ids,\n        'condition': conditions,\n        'level': levels\n    })\n\n    # Convert true labels to one-hot encoding\n    true_labels_onehot = np.zeros((len(true_labels), 3))\n    for i, label in enumerate(true_labels):\n        true_labels_onehot[i, int(label)] = 1\n\n    # Add true labels and predictions to DataFrame\n    results_df['true_normal_mild'] = true_labels_onehot[:, 0]\n    results_df['true_moderate'] = true_labels_onehot[:, 1]\n    results_df['true_severe'] = true_labels_onehot[:, 2]\n    results_df['pred_normal_mild'] = predictions[:, 0]\n    results_df['pred_moderate'] = predictions[:, 1]\n    results_df['pred_severe'] = predictions[:, 2]\n\n    # Initialize metrics\n    metrics = {}\n    total_loss = 0\n    total_weight = 0\n\n    # Define condition weights\n    condition_weights = {\n        'Spinal Canal Stenosis': {\n            'L1_L2': 0.08, 'L2_L3': 0.14, 'L3_L4': 0.22, 'L4_L5': 0.35, 'L5_S1': 0.21\n        },\n        'Left Neural Foraminal Narrowing': {\n            'L1_L2': 0.07, 'L2_L3': 0.13, 'L3_L4': 0.21, 'L4_L5': 0.33, 'L5_S1': 0.26\n        },\n        'Right Neural Foraminal Narrowing': {\n            'L1_L2': 0.07, 'L2_L3': 0.13, 'L3_L4': 0.21, 'L4_L5': 0.33, 'L5_S1': 0.26\n        },\n        'Left Subarticular Stenosis': {\n            'L1_L2': 0.07, 'L2_L3': 0.13, 'L3_L4': 0.21, 'L4_L5': 0.33, 'L5_S1': 0.26\n        },\n        'Right Subarticular Stenosis': {\n            'L1_L2': 0.07, 'L2_L3': 0.13, 'L3_L4': 0.21, 'L4_L5': 0.33, 'L5_S1': 0.26\n        }\n    }\n\n    # Compute weighted log loss for each condition and level\n    for condition in condition_weights.keys():\n        for level in condition_weights[condition].keys():\n            mask = (results_df['condition'] == condition) & (results_df['level'] == level)\n            if mask.any():\n                y_true = results_df.loc[mask, ['true_normal_mild', 'true_moderate', 'true_severe']].values\n                y_pred = results_df.loc[mask, ['pred_normal_mild', 'pred_moderate', 'pred_severe']].values\n\n                # Add small epsilon to avoid log(0)\n                y_pred = np.clip(y_pred, 1e-7, 1-1e-7)\n                \n                try:\n                    ll = log_loss(y_true, y_pred)\n                    weight = condition_weights[condition][level]\n                    total_loss += ll * weight\n                    total_weight += weight\n                    metrics[f'{condition}_{level}_log_loss'] = ll\n                except ValueError as e:\n                    print(f\"Warning: Error computing log loss for {condition} {level}: {e}\")\n                    continue\n\n    # Compute weighted average log loss\n    if total_weight > 0:\n        weighted_log_loss = total_loss / total_weight\n    else:\n        weighted_log_loss = 0\n    metrics['weighted_log_loss'] = weighted_log_loss\n\n    # Compute any_severe_spinal metric\n    study_severe_true = []\n    study_severe_pred = []\n\n    for study_id in results_df['study_id'].unique():\n        study_mask = results_df['study_id'] == study_id\n        study_data = results_df[study_mask]\n\n        # True severe cases\n        has_severe = (study_data['true_severe'] == 1).any()\n        study_severe_true.append(has_severe)\n\n        # Predicted severe cases\n        pred_severe = study_data['pred_severe'].max()\n        study_severe_pred.append(pred_severe)\n\n    # Add small epsilon to avoid log(0)\n    study_severe_pred = np.clip(study_severe_pred, 1e-7, 1-1e-7)\n    \n    try:\n        any_severe_loss = log_loss(study_severe_true, study_severe_pred)\n    except ValueError as e:\n        print(f\"Warning: Error computing any_severe_loss: {e}\")\n        any_severe_loss = 0\n\n    metrics['any_severe_loss'] = any_severe_loss\n\n    # Compute final score\n    final_score = 0.7 * weighted_log_loss + 0.3 * any_severe_loss\n    metrics['final_score'] = final_score\n\n    return final_score, metrics\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T12:46:41.926369Z","iopub.execute_input":"2024-11-21T12:46:41.926764Z","iopub.status.idle":"2024-11-21T12:46:41.942113Z","shell.execute_reply.started":"2024-11-21T12:46:41.926730Z","shell.execute_reply":"2024-11-21T12:46:41.941004Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_model(model, val_loader, device):\n    \"\"\"Evaluate model using competition metrics\"\"\"\n    model.eval()\n    all_predictions = []\n    all_labels = []\n    all_study_ids = []\n    all_conditions = []\n    all_levels = []\n\n    with torch.no_grad():\n        for batch in val_loader:\n            images = batch['image'].to(device)\n            conditions = batch['condition'].to(device)\n            levels = batch['level'].to(device)\n\n            # Get model predictions\n            outputs = model(images, conditions, levels)\n            probs = F.softmax(outputs, dim=1)\n\n            # Store results\n            all_predictions.append(probs.cpu().numpy())\n            all_labels.append(batch['severity'].numpy())\n            all_study_ids.extend(batch['study_id'])\n\n            # Get condition and level names\n            for c_idx, l_idx in zip(torch.argmax(conditions, dim=1).cpu().numpy(),\n                                  torch.argmax(levels, dim=1).cpu().numpy()):\n                all_conditions.append(list(val_loader.dataset.condition_map.keys())[c_idx])\n                all_levels.append(list(val_loader.dataset.level_map.keys())[l_idx])\n\n    # Concatenate all predictions\n    predictions = np.concatenate(all_predictions)\n    labels = np.concatenate(all_labels)\n\n    # Compute metrics\n    score, metrics = compute_competition_metric(\n        predictions,\n        labels,\n        np.array(all_study_ids),\n        all_conditions,\n        all_levels\n    )\n\n    return score, metrics\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T12:47:05.399036Z","iopub.execute_input":"2024-11-21T12:47:05.399767Z","iopub.status.idle":"2024-11-21T12:47:05.407327Z","shell.execute_reply.started":"2024-11-21T12:47:05.399731Z","shell.execute_reply":"2024-11-21T12:47:05.406547Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def evaluate_regression_model(model, val_loader, device):\n    \"\"\"Evaluate regression model with proper probability conversion\"\"\"\n    model.eval()\n    all_predictions = []\n    all_labels = []\n    all_study_ids = []\n    all_conditions = []\n    all_levels = []\n\n    with torch.no_grad():\n        for batch in val_loader:\n            images = batch['image'].to(device)\n            conditions = batch['condition'].to(device)\n            levels = batch['level'].to(device)\n\n            # Get regression predictions\n            reg_output = model(images, conditions, levels)\n            reg_output = reg_output.squeeze()  # Remove all extra dimensions\n\n            # Convert regression values to probabilities\n            probs = torch.zeros((reg_output.shape[0], 3), device=device)\n\n            # Normal/Mild: values <= 0.5\n            normal_mask = reg_output <= 0.5\n            probs[normal_mask, 0] = 1.0\n\n            # Moderate: values between 0.5 and 1.5\n            moderate_mask = (reg_output > 0.5) & (reg_output <= 1.5)\n            probs[moderate_mask, 1] = 1.0\n\n            # Severe: values > 1.5\n            severe_mask = reg_output > 1.5\n            probs[severe_mask, 2] = 1.0\n\n            # Add smoothing to avoid zero probabilities\n            probs = probs + 1e-7\n            probs = probs / probs.sum(dim=1, keepdim=True)\n\n            # Store results\n            all_predictions.append(probs.cpu().numpy())\n            all_labels.append(batch['severity'].numpy())\n            all_study_ids.extend(batch['study_id'])\n\n            # Get condition and level names\n            for c_idx, l_idx in zip(torch.argmax(conditions, dim=1).cpu().numpy(),\n                                  torch.argmax(levels, dim=1).cpu().numpy()):\n                all_conditions.append(list(val_loader.dataset.condition_map.keys())[c_idx])\n                all_levels.append(list(val_loader.dataset.level_map.keys())[l_idx])\n\n    # Concatenate all predictions\n    predictions = np.concatenate(all_predictions)\n    labels = np.concatenate(all_labels)\n\n    # Compute metrics\n    score, metrics = compute_competition_metric(\n        predictions,\n        labels,\n        np.array(all_study_ids),\n        all_conditions,\n        all_levels\n    )\n\n    return score, metrics","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T12:47:29.186746Z","iopub.execute_input":"2024-11-21T12:47:29.187488Z","iopub.status.idle":"2024-11-21T12:47:29.197737Z","shell.execute_reply.started":"2024-11-21T12:47:29.187453Z","shell.execute_reply":"2024-11-21T12:47:29.196789Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    \"\"\"Main evaluation function\"\"\"\n    # Set device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"Using device: {device}\")\n\n    # Load validation data\n    print(\"Loading validation data...\")\n    val_samples = load_processed_data('/kaggle/working/val_processed.npy')\n\n    # Create dataloaders\n    val_dataset_class = LumbarSpineDataset(val_samples, augment=False)\n    val_loader_class = DataLoader(\n        val_dataset_class,\n        batch_size=32,\n        shuffle=False,\n        num_workers=4\n    )\n\n    val_dataset_reg = LumbarSpineRegDataset(val_samples, augment=False)\n    val_loader_reg = DataLoader(\n        val_dataset_reg,\n        batch_size=32,\n        shuffle=False,\n        num_workers=4\n    )\n\n    try:\n        # Evaluate classification model\n        print(\"\\nEvaluating Classification Model...\")\n        class_model = LumbarClassifier().to(device)\n        class_model.load_state_dict(torch.load('best_model.pth')['model_state_dict'])\n        class_model.eval()\n\n        class_score, class_metrics = evaluate_model(class_model, val_loader_class, device)\n\n        print(\"\\nClassification Model Results:\")\n        print(f\"Final Score: {class_score:.4f}\")\n        print(\"\\nDetailed Metrics:\")\n        for key, value in class_metrics.items():\n            if not key.endswith('_log_loss'):\n                print(f\"{key}: {value:.4f}\")\n\n        print(\"\\nPer-condition Log Losses:\")\n        for key, value in class_metrics.items():\n            if key.endswith('_log_loss') and key != 'weighted_log_loss':\n                print(f\"{key}: {value:.4f}\")\n\n        # Evaluate regression model\n        print(\"\\nEvaluating Regression Model...\")\n        reg_model = LumbarRegressor().to(device)\n        reg_model.load_state_dict(torch.load('best_regression_model.pth')['model_state_dict'])\n        reg_model.eval()\n\n        reg_score, reg_metrics = evaluate_regression_model(reg_model, val_loader_reg, device)\n\n        print(\"\\nRegression Model Results:\")\n        print(f\"Final Score: {reg_score:.4f}\")\n        print(\"\\nDetailed Metrics:\")\n        for key, value in reg_metrics.items():\n            if not key.endswith('_log_loss'):\n                print(f\"{key}: {value:.4f}\")\n\n        print(\"\\nPer-condition Log Losses:\")\n        for key, value in reg_metrics.items():\n            if key.endswith('_log_loss') and key != 'weighted_log_loss':\n                print(f\"{key}: {value:.4f}\")\n\n        # Compare models\n        print(\"\\nModel Comparison:\")\n        print(f\"Classification Model Score: {class_score:.4f}\")\n        print(f\"Regression Model Score: {reg_score:.4f}\")\n\n        # Save results\n        results = {\n            'classification': {\n                'score': class_score,\n                'metrics': class_metrics\n            },\n            'regression': {\n                'score': reg_score,\n                'metrics': reg_metrics\n            }\n        }\n\n        np.save('evaluation_results.npy', results)\n        print(\"\\nResults saved to evaluation_results.npy\")\n        \n    except Exception as e:\n        print(f\"Error during evaluation: {e}\")\n        raise\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T12:48:06.705297Z","iopub.execute_input":"2024-11-21T12:48:06.706003Z","iopub.status.idle":"2024-11-21T12:48:54.107262Z","shell.execute_reply.started":"2024-11-21T12:48:06.705967Z","shell.execute_reply":"2024-11-21T12:48:54.106153Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Overall Performance:\n\n* Classification Model Score: 0.7661\n* Regression Model Score: 2.9375\n* The classification model significantly outperforms the regression model\n\nDetailed Analysis:\na) Classification Model:\n\n* Strong performance on Spinal Canal Stenosis (lowest log losses)\n* Better at upper levels (L1_L2, L2_L3)\n* any_severe_loss: 0.6037\n\nb) Regression Model:\n\n* Higher log losses across all conditions\n* Particularly struggles with lower levels\n* Extremely high any_severe_loss: 2.5813\n\nPattern Analysis:\n\nBoth models show similar patterns:\n\n* Better performance on upper levels\n* Struggles with L4_L5 and L5_S1\n* Spinal Canal Stenosis is easiest to predict\n* Neural Foraminal Narrowing and Subarticular Stenosis are more challenging","metadata":{}},{"cell_type":"markdown","source":"# Prediction Pipeline","metadata":{}},{"cell_type":"code","source":"class OptimizedPredictionPipeline:\n    def __init__(self, \n                 classification_model: nn.Module,\n                 device: torch.device):\n        \"\"\"\n        Initialize prediction pipeline with optimized rules based on evaluation results\n        \"\"\"\n        self.classification_model = classification_model\n        self.device = device\n        self.classification_model.eval()\n        \n        # Define condition-specific weights based on evaluation results\n        self.condition_weights = {\n            'Spinal Canal Stenosis': 1.0,  # Best performing condition\n            'Left Neural Foraminal Narrowing': 0.95,\n            'Right Neural Foraminal Narrowing': 0.95,\n            'Left Subarticular Stenosis': 0.93,\n            'Right Subarticular Stenosis': 0.93\n        }\n        \n        # Define level-specific weights based on evaluation results\n        self.level_weights = {\n            'L1_L2': 1.0,  # Best performing level\n            'L2_L3': 0.98,\n            'L3_L4': 0.95,\n            'L4_L5': 0.90,  # Most challenging level\n            'L5_S1': 0.92\n        }\n        \n        # Define log loss thresholds for confidence adjustment\n        self.log_loss_thresholds = {\n            'Spinal Canal Stenosis': {\n                'L1_L2': 0.1067,  # Using actual log loss values from evaluation\n                'L2_L3': 0.4514,\n                'L3_L4': 0.5757,\n                'L4_L5': 0.8928,\n                'L5_S1': 0.2164\n            },\n            'Neural Foraminal Narrowing': {\n                'L1_L2': 0.2200,  # Average of left and right\n                'L2_L3': 0.4635,\n                'L3_L4': 0.9959,\n                'L4_L5': 1.1089,\n                'L5_S1': 1.0936\n            },\n            'Subarticular Stenosis': {\n                'L1_L2': 0.2285,\n                'L2_L3': 0.5337,\n                'L3_L4': 0.9766,\n                'L4_L5': 1.1551,\n                'L5_S1': 0.7856\n            }\n        }\n    \n    def preprocess_input(self, image: torch.Tensor, condition: str, level: str) -> Tuple[torch.Tensor, torch.Tensor, torch.Tensor]:\n        \"\"\"\n        Preprocess inputs for model prediction\n        \"\"\"\n        # Ensure image is on correct device\n        if isinstance(image, np.ndarray):\n            image = torch.from_numpy(image).float()\n        image = image.to(self.device)\n        \n        # Create condition one-hot encoding\n        condition_map = {\n            'Spinal Canal Stenosis': 0,\n            'Left Neural Foraminal Narrowing': 1,\n            'Right Neural Foraminal Narrowing': 2,\n            'Left Subarticular Stenosis': 3,\n            'Right Subarticular Stenosis': 4\n        }\n        condition_tensor = torch.zeros(5, device=self.device)\n        condition_tensor[condition_map[condition]] = 1\n        \n        # Create level one-hot encoding\n        level_map = {\n            'L1_L2': 0, 'L2_L3': 1, 'L3_L4': 2, 'L4_L5': 3, 'L5_S1': 4\n        }\n        level_tensor = torch.zeros(5, device=self.device)\n        level_tensor[level_map[level]] = 1\n        \n        return image, condition_tensor, level_tensor\n    \n    def apply_condition_level_adjustments(self, \n                                        predictions: torch.Tensor,\n                                        condition: str,\n                                        level: str) -> torch.Tensor:\n        \"\"\"\n        Apply condition and level-specific adjustments to predictions\n        \"\"\"\n        # Get base weights\n        condition_weight = self.condition_weights.get(condition, 0.95)\n        level_weight = self.level_weights.get(level, 0.90)\n        \n        # Get log loss threshold\n        if 'Neural Foraminal Narrowing' in condition:\n            condition_key = 'Neural Foraminal Narrowing'\n        elif 'Subarticular Stenosis' in condition:\n            condition_key = 'Subarticular Stenosis'\n        else:\n            condition_key = condition\n            \n        log_loss_threshold = self.log_loss_thresholds[condition_key][level]\n        \n        # Apply weights\n        adjusted_predictions = predictions * (condition_weight * level_weight)\n        \n        # Adjust based on log loss threshold\n        if log_loss_threshold > 0.8:  # High uncertainty\n            # More conservative predictions for high uncertainty cases\n            adjusted_predictions[:, 2] *= 0.9  # Reduce severe predictions\n            adjusted_predictions[:, 1] *= 0.95  # Slightly reduce moderate predictions\n            adjusted_predictions[:, 0] += 0.1  # Bias toward normal/mild\n        \n        # Normalize predictions\n        adjusted_predictions = F.normalize(adjusted_predictions, p=1, dim=1)\n        \n        return adjusted_predictions\n    \n    def get_prediction_confidence(self, predictions: torch.Tensor) -> torch.Tensor:\n        \"\"\"\n        Calculate prediction confidence\n        \"\"\"\n        # Get max probability and entropy\n        max_prob = predictions.max(dim=1)[0]\n        entropy = -(predictions * torch.log(predictions + 1e-7)).sum(dim=1)\n        \n        # Combine max probability and entropy for confidence score\n        confidence = max_prob * (1 - entropy/np.log(3))  # Normalize entropy by max possible value\n        \n        return confidence\n    \n    def predict(self, \n                image: torch.Tensor,\n                condition: str,\n                level: str) -> Dict[str, torch.Tensor]:\n        \"\"\"\n        Generate predictions with confidence scores\n        \"\"\"\n        with torch.no_grad():\n            # Preprocess inputs\n            image, condition_tensor, level_tensor = self.preprocess_input(image, condition, level)\n            \n            # Get base predictions\n            base_predictions = self.classification_model(image, condition_tensor, level_tensor)\n            base_probabilities = F.softmax(base_predictions, dim=1)\n            \n            # Apply adjustments\n            adjusted_predictions = self.apply_condition_level_adjustments(\n                base_probabilities,\n                condition,\n                level\n            )\n            \n            # Calculate confidence\n            confidence = self.get_prediction_confidence(adjusted_predictions)\n            \n            return {\n                'probabilities': adjusted_predictions,\n                'confidence': confidence,\n                'severity_prediction': torch.argmax(adjusted_predictions, dim=1),\n                'original_probabilities': base_probabilities\n            }\n    \n    def batch_predict(self, \n                     dataloader: torch.utils.data.DataLoader) -> List[Dict[str, torch.Tensor]]:\n        \"\"\"\n        Generate predictions for a batch of data\n        \"\"\"\n        predictions = []\n        \n        with torch.no_grad():\n            for batch in dataloader:\n                images = batch['image'].to(self.device)\n                conditions = batch['condition'].to(self.device)\n                levels = batch['level'].to(self.device)\n                \n                # Get predictions for batch\n                outputs = self.classification_model(images, conditions, levels)\n                probs = F.softmax(outputs, dim=1)\n                \n                # Process each sample in batch\n                for i in range(len(images)):\n                    condition_idx = torch.argmax(conditions[i]).item()\n                    level_idx = torch.argmax(levels[i]).item()\n                    \n                    condition = list(self.condition_weights.keys())[condition_idx]\n                    level = list(self.level_weights.keys())[level_idx]\n                    \n                    # Apply adjustments\n                    adjusted_probs = self.apply_condition_level_adjustments(\n                        probs[i].unsqueeze(0),\n                        condition,\n                        level\n                    )\n                    \n                    confidence = self.get_prediction_confidence(adjusted_probs)\n                    \n                    predictions.append({\n                        'probabilities': adjusted_probs,\n                        'confidence': confidence,\n                        'severity_prediction': torch.argmax(adjusted_probs, dim=1),\n                        'original_probabilities': probs[i].unsqueeze(0)\n                    })\n        \n        return predictions\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T13:01:58.992295Z","iopub.execute_input":"2024-11-21T13:01:58.992701Z","iopub.status.idle":"2024-11-21T13:01:59.012363Z","shell.execute_reply.started":"2024-11-21T13:01:58.992662Z","shell.execute_reply":"2024-11-21T13:01:59.011410Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    \"\"\"Example usage of the prediction pipeline\"\"\"\n    # Set device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Load the classification model\n    model = LumbarClassifier().to(device)\n    model.load_state_dict(torch.load('best_model.pth')['model_state_dict'])\n    \n    # Create prediction pipeline\n    pipeline = OptimizedPredictionPipeline(model, device)\n    \n    # Load validation data\n    val_samples = load_processed_data('/kaggle/working/val_processed.npy')\n    val_dataset = LumbarSpineDataset(val_samples, augment=False)\n    val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)\n    \n    # Generate predictions\n    print(\"Generating predictions...\")\n    predictions = pipeline.batch_predict(val_loader)\n    \n    # Analyze results\n    confidences = torch.cat([p['confidence'] for p in predictions])\n    print(f\"\\nAverage prediction confidence: {confidences.mean():.4f}\")\n    print(f\"Minimum confidence: {confidences.min():.4f}\")\n    print(f\"Maximum confidence: {confidences.max():.4f}\")\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T13:02:38.374420Z","iopub.execute_input":"2024-11-21T13:02:38.375077Z","iopub.status.idle":"2024-11-21T13:03:03.770740Z","shell.execute_reply.started":"2024-11-21T13:02:38.375040Z","shell.execute_reply":"2024-11-21T13:03:03.769763Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Initialize pipeline\n#pipeline = OptimizedPredictionPipeline(classification_model, device)\n\n# Single prediction\n#result = pipeline.predict(image, condition, level)\n\n# Batch predictions\n#results = pipeline.batch_predict(dataloader)","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Prediction Visualization","metadata":{}},{"cell_type":"code","source":"class PredictionVisualizer:\n    \"\"\"Visualization tools for the prediction pipeline\"\"\"\n    def __init__(self):\n        self.severity_classes = ['Normal/Mild', 'Moderate', 'Severe']\n        self.colors = {\n            'Normal/Mild': '#2ecc71',  # Green\n            'Moderate': '#f1c40f',     # Yellow\n            'Severe': '#e74c3c'        # Red\n        }\n        \n    def plot_prediction_distribution(self, predictions: List[Dict[str, torch.Tensor]], \n                                   save_path: str = None):\n        \"\"\"Plot distribution of predictions across severity classes\"\"\"\n        plt.figure(figsize=(10, 6))\n        \n        # Convert predictions to numpy arrays\n        all_probs = torch.cat([p['probabilities'] for p in predictions]).cpu().numpy()\n        \n        # Create boxplot\n        bp = plt.boxplot([all_probs[:, i] for i in range(3)], \n                        labels=self.severity_classes,\n                        patch_artist=True)\n        \n        # Color boxes\n        for patch, color in zip(bp['boxes'], self.colors.values()):\n            patch.set_facecolor(color)\n            patch.set_alpha(0.6)\n        \n        plt.title('Distribution of Prediction Probabilities by Class')\n        plt.ylabel('Probability')\n        plt.grid(True, alpha=0.3)\n        \n        if save_path:\n            plt.savefig(f\"{save_path}/prediction_distribution.png\")\n        plt.show()\n    \n    def plot_confidence_heatmap(self, predictions: List[Dict[str, torch.Tensor]], \n                              conditions: List[str], \n                              levels: List[str],\n                              save_path: str = None):\n        \"\"\"Create heatmap of prediction confidence by condition and level\"\"\"\n        confidence_matrix = np.zeros((len(set(conditions)), len(set(levels))))\n        count_matrix = np.zeros_like(confidence_matrix)\n        \n        unique_conditions = list(set(conditions))\n        unique_levels = list(set(levels))\n        \n        # Aggregate confidences\n        for pred, cond, lvl in zip(predictions, conditions, levels):\n            i = unique_conditions.index(cond)\n            j = unique_levels.index(lvl)\n            confidence_matrix[i, j] += pred['confidence'].item()\n            count_matrix[i, j] += 1\n        \n        # Calculate average confidence\n        avg_confidence = np.divide(confidence_matrix, count_matrix, \n                                 where=count_matrix != 0)\n        \n        # Create heatmap\n        plt.figure(figsize=(12, 8))\n        sns.heatmap(avg_confidence, \n                   xticklabels=unique_levels,\n                   yticklabels=unique_conditions,\n                   annot=True, \n                   fmt='.3f',\n                   cmap='YlOrRd')\n        \n        plt.title('Average Prediction Confidence by Condition and Level')\n        plt.xlabel('Spinal Level')\n        plt.ylabel('Condition')\n        \n        if save_path:\n            plt.savefig(f\"{save_path}/confidence_heatmap.png\")\n        plt.show()\n    \n    def plot_severity_distribution(self, predictions: List[Dict[str, torch.Tensor]], \n                                 conditions: List[str],\n                                 save_path: str = None):\n        \"\"\"Plot severity distribution by condition\"\"\"\n        severity_counts = {cond: [0, 0, 0] for cond in set(conditions)}\n        \n        # Count predictions for each severity level\n        for pred, cond in zip(predictions, conditions):\n            severity = pred['severity_prediction'].item()\n            severity_counts[cond][severity] += 1\n        \n        # Create stacked bar chart\n        df = pd.DataFrame(severity_counts, index=self.severity_classes).T\n        \n        plt.figure(figsize=(12, 6))\n        df.plot(kind='bar', stacked=True, color=[self.colors[c] for c in self.severity_classes])\n        \n        plt.title('Predicted Severity Distribution by Condition')\n        plt.xlabel('Condition')\n        plt.ylabel('Count')\n        plt.legend(title='Severity', bbox_to_anchor=(1.05, 1))\n        plt.tight_layout()\n        \n        if save_path:\n            plt.savefig(f\"{save_path}/severity_distribution.png\")\n        plt.show()\n    \n    def plot_confidence_histogram(self, predictions: List[Dict[str, torch.Tensor]], \n                                save_path: str = None):\n        \"\"\"Plot histogram of prediction confidences\"\"\"\n        confidences = [pred['confidence'].item() for pred in predictions]\n        \n        plt.figure(figsize=(10, 6))\n        plt.hist(confidences, bins=50, color='skyblue', alpha=0.7, edgecolor='black')\n        plt.axvline(np.mean(confidences), color='red', linestyle='--', \n                   label=f'Mean: {np.mean(confidences):.3f}')\n        \n        plt.title('Distribution of Prediction Confidences')\n        plt.xlabel('Confidence Score')\n        plt.ylabel('Count')\n        plt.legend()\n        plt.grid(True, alpha=0.3)\n        \n        if save_path:\n            plt.savefig(f\"{save_path}/confidence_histogram.png\")\n        plt.show()\n    \n    def plot_prediction_changes(self, predictions: List[Dict[str, torch.Tensor]], \n                              save_path: str = None):\n        \"\"\"Compare original vs adjusted predictions\"\"\"\n        orig_probs = torch.cat([p['original_probabilities'] for p in predictions]).cpu().numpy()\n        adj_probs = torch.cat([p['probabilities'] for p in predictions]).cpu().numpy()\n        \n        fig, axes = plt.subplots(1, 3, figsize=(15, 5))\n        \n        for i, severity in enumerate(self.severity_classes):\n            axes[i].scatter(orig_probs[:, i], adj_probs[:, i], \n                          alpha=0.5, color=list(self.colors.values())[i])\n            axes[i].plot([0, 1], [0, 1], 'r--', alpha=0.5)\n            axes[i].set_title(f'{severity} Probabilities')\n            axes[i].set_xlabel('Original')\n            axes[i].set_ylabel('Adjusted')\n            axes[i].grid(True, alpha=0.3)\n        \n        plt.tight_layout()\n        if save_path:\n            plt.savefig(f\"{save_path}/prediction_changes.png\")\n        plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T13:09:50.390331Z","iopub.execute_input":"2024-11-21T13:09:50.390726Z","iopub.status.idle":"2024-11-21T13:09:50.410206Z","shell.execute_reply.started":"2024-11-21T13:09:50.390689Z","shell.execute_reply":"2024-11-21T13:09:50.409234Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_predictions(pipeline, dataloader, save_path=None):\n    \"\"\"Generate comprehensive visualization of predictions\"\"\"\n    # Get predictions\n    predictions = pipeline.batch_predict(dataloader)\n    \n    # Get conditions and levels from dataloader\n    conditions = []\n    levels = []\n    for batch in dataloader:\n        for c, l in zip(batch['condition'], batch['level']):\n            c_idx = torch.argmax(c).item()\n            l_idx = torch.argmax(l).item()\n            conditions.append(list(pipeline.condition_weights.keys())[c_idx])\n            levels.append(list(pipeline.level_weights.keys())[l_idx])\n    \n    # Create visualizer\n    visualizer = PredictionVisualizer()\n    \n    # Generate all plots\n    print(\"Generating visualization plots...\")\n    \n    print(\"\\n1. Prediction Distribution\")\n    visualizer.plot_prediction_distribution(predictions, save_path)\n    \n    print(\"\\n2. Confidence Heatmap\")\n    visualizer.plot_confidence_heatmap(predictions, conditions, levels, save_path)\n    \n    print(\"\\n3. Severity Distribution\")\n    visualizer.plot_severity_distribution(predictions, conditions, save_path)\n    \n    print(\"\\n4. Confidence Histogram\")\n    visualizer.plot_confidence_histogram(predictions, save_path)\n    \n    print(\"\\n5. Prediction Adjustments\")\n    visualizer.plot_prediction_changes(predictions, save_path)\n    \n    print(\"\\nVisualization completed!\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T13:10:22.678846Z","iopub.execute_input":"2024-11-21T13:10:22.679725Z","iopub.status.idle":"2024-11-21T13:10:22.686548Z","shell.execute_reply.started":"2024-11-21T13:10:22.679681Z","shell.execute_reply":"2024-11-21T13:10:22.685548Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    \"\"\"Example usage of visualization tools\"\"\"\n    # Set device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    \n    # Load model and create pipeline\n    model = LumbarClassifier().to(device)\n    model.load_state_dict(torch.load('best_model.pth')['model_state_dict'])\n    pipeline = OptimizedPredictionPipeline(model, device)\n    \n    # Load validation data\n    val_samples = load_processed_data('/kaggle/working/val_processed.npy')\n    val_dataset = LumbarSpineDataset(val_samples, augment=False)\n    val_loader = DataLoader(val_dataset, batch_size=32, shuffle=False)\n    \n    # Generate visualizations\n    visualize_predictions(pipeline, val_loader)#, save_path='/kaggle/working/visualizations')\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T13:16:53.277764Z","iopub.execute_input":"2024-11-21T13:16:53.278816Z","iopub.status.idle":"2024-11-21T13:17:21.743902Z","shell.execute_reply.started":"2024-11-21T13:16:53.278768Z","shell.execute_reply":"2024-11-21T13:17:21.742967Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Sample Image Prediction","metadata":{}},{"cell_type":"code","source":"class ModifiedOptimizedPredictionPipeline:\n    def __init__(self, classification_model, device):\n        self.classification_model = classification_model\n        self.device = device\n        self.classification_model.eval()\n        \n        # Define maps\n        self.condition_map = {\n            'Spinal Canal Stenosis': 0,\n            'Left Neural Foraminal Narrowing': 1,\n            'Right Neural Foraminal Narrowing': 2,\n            'Left Subarticular Stenosis': 3,\n            'Right Subarticular Stenosis': 4\n        }\n        \n        self.level_map = {\n            'L1_L2': 0, 'L2_L3': 1, 'L3_L4': 2, 'L4_L5': 3, 'L5_S1': 4\n        }\n    \n    def prepare_input_tensors(self, image, condition, level):\n        \"\"\"Prepare input tensors with correct dimensions\"\"\"\n        # Image tensor\n        if isinstance(image, np.ndarray):\n            image = torch.from_numpy(image).float()\n        if len(image.shape) == 3:  # (H, W, C)\n            image = image.permute(2, 0, 1)  # (C, H, W)\n        if len(image.shape) == 3:\n            image = image.unsqueeze(0)  # (1, C, H, W)\n        \n        # Condition tensor\n        condition_tensor = torch.zeros(1, len(self.condition_map))\n        condition_tensor[0, self.condition_map[condition]] = 1\n        \n        # Level tensor\n        level_tensor = torch.zeros(1, len(self.level_map))\n        level_tensor[0, self.level_map[level]] = 1\n        \n        # Move to device\n        image = image.to(self.device)\n        condition_tensor = condition_tensor.to(self.device)\n        level_tensor = level_tensor.to(self.device)\n        \n        return image, condition_tensor, level_tensor\n    \n    def predict(self, image, condition, level):\n        \"\"\"Make prediction with proper tensor handling\"\"\"\n        with torch.no_grad():\n            # Prepare inputs\n            image_tensor, condition_tensor, level_tensor = self.prepare_input_tensors(\n                image, condition, level\n            )\n            \n            # Get predictions\n            outputs = self.classification_model(image_tensor, condition_tensor, level_tensor)\n            probabilities = F.softmax(outputs, dim=1)\n            \n            # Get prediction class and confidence\n            pred_class = torch.argmax(probabilities, dim=1)[0]\n            confidence = torch.max(probabilities, dim=1)[0][0]\n            \n            return {\n                'probabilities': probabilities,\n                'severity_prediction': pred_class,\n                'confidence': confidence,\n            }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T13:43:00.282296Z","iopub.execute_input":"2024-11-21T13:43:00.282977Z","iopub.status.idle":"2024-11-21T13:43:00.292602Z","shell.execute_reply.started":"2024-11-21T13:43:00.282941Z","shell.execute_reply":"2024-11-21T13:43:00.291661Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_preprocessed_samples(pipeline, preprocessed_path: str, num_samples: int = 5):\n    \"\"\"Make predictions on randomly selected samples\"\"\"\n    # Load preprocessed data\n    print(\"Loading preprocessed data...\")\n    preprocessed_data = np.load(preprocessed_path, allow_pickle=True).item()\n    \n    # Get random indices\n    total_samples = len(preprocessed_data['images'])\n    sample_indices = np.random.choice(total_samples, num_samples, replace=False)\n    \n    severity_classes = ['Normal/Mild', 'Moderate', 'Severe']\n    colors = ['green', 'yellow', 'red']\n    \n    for idx in sample_indices:\n        try:\n            # Get sample data\n            image = preprocessed_data['images'][idx]\n            condition = preprocessed_data['conditions'][idx]\n            level = preprocessed_data['levels'][idx]\n            true_severity = preprocessed_data['severities'][idx]\n            study_id = preprocessed_data['study_ids'][idx]\n            \n            print(f\"\\nProcessing Sample {idx} (Study ID: {study_id})\")\n            \n            # Get prediction\n            prediction = pipeline.predict(image, condition, level)\n            \n            # Plot results\n            plt.figure(figsize=(10, 5))\n            \n            # Plot image\n            plt.subplot(1, 2, 1)\n            plt.imshow(image.squeeze(), cmap='gray')\n            plt.title(f\"Study ID: {study_id}\\n{condition}\\nLevel: {level}\")\n            plt.axis('off')\n            \n            # Plot prediction results\n            plt.subplot(1, 2, 2)\n            probs = prediction['probabilities'].squeeze().cpu().numpy()\n            pred_class = prediction['severity_prediction'].item()\n            confidence = prediction['confidence'].item()\n            \n            # Create bar plot\n            bars = plt.bar(severity_classes, probs, color=colors, alpha=0.6)\n            plt.ylim(0, 1)\n            plt.title(f\"Predictions\\nTrue: {true_severity}\\n\" + \n                     f\"Predicted: {severity_classes[pred_class]}\\n\" +\n                     f\"Confidence: {confidence:.2f}\")\n            \n            # Add value labels on bars\n            for bar in bars:\n                height = bar.get_height()\n                plt.text(bar.get_x() + bar.get_width()/2., height,\n                        f'{height:.2f}',\n                        ha='center', va='bottom')\n            \n            plt.tight_layout()\n            plt.show()\n            \n            # Print detailed prediction information\n            print(\"\\nDetailed Prediction Information:\")\n            print(f\"Condition: {condition}\")\n            print(f\"Level: {level}\")\n            print(f\"True Severity: {true_severity}\")\n            print(f\"Predicted Severity: {severity_classes[pred_class]}\")\n            print(f\"Confidence: {confidence:.4f}\")\n            print(\"\\nClass Probabilities:\")\n            for cls, prob in zip(severity_classes, probs):\n                print(f\"{cls}: {prob:.4f}\")\n            print(\"-\" * 50)\n            \n        except Exception as e:\n            print(f\"Error processing sample {idx}: {str(e)}\")\n            continue","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T13:43:44.173067Z","iopub.execute_input":"2024-11-21T13:43:44.173846Z","iopub.status.idle":"2024-11-21T13:43:44.184796Z","shell.execute_reply.started":"2024-11-21T13:43:44.173809Z","shell.execute_reply":"2024-11-21T13:43:44.183765Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Set device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"Using device: {device}\")\n    \n    # Load the model\n    model = LumbarClassifier().to(device)\n    model.load_state_dict(torch.load('best_model.pth')['model_state_dict'])\n    model.eval()\n    \n    # Create prediction pipeline\n    pipeline = ModifiedOptimizedPredictionPipeline(model, device)\n    \n    # Make predictions on preprocessed samples\n    predict_preprocessed_samples(\n        pipeline,\n        preprocessed_path='/kaggle/working/val_processed.npy',\n        num_samples=10\n    )\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T13:49:33.955380Z","iopub.execute_input":"2024-11-21T13:49:33.956205Z","iopub.status.idle":"2024-11-21T13:49:39.303080Z","shell.execute_reply.started":"2024-11-21T13:49:33.956165Z","shell.execute_reply":"2024-11-21T13:49:39.302206Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Advanced Analysis","metadata":{}},{"cell_type":"code","source":"from sklearn.metrics import confusion_matrix","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T14:01:20.747103Z","iopub.execute_input":"2024-11-21T14:01:20.747765Z","iopub.status.idle":"2024-11-21T14:01:20.751983Z","shell.execute_reply.started":"2024-11-21T14:01:20.747716Z","shell.execute_reply":"2024-11-21T14:01:20.751073Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class AdvancedAnalysis:\n    def __init__(self, pipeline, preprocessed_path: str):\n        self.pipeline = pipeline\n        self.preprocessed_data = np.load(preprocessed_path, allow_pickle=True).item()\n        self.severity_classes = ['Normal/Mild', 'Moderate', 'Severe']\n    \n    def analyze_prediction_patterns(self, num_samples: int = 1000):\n        \"\"\"Perform statistical analysis of prediction patterns with error handling\"\"\"\n        total_samples = len(self.preprocessed_data['images'])\n        sample_indices = np.random.choice(total_samples, min(num_samples, total_samples), replace=False)\n        \n        results = {\n            'condition_level_accuracy': defaultdict(lambda: defaultdict(list)),\n            'confidence_scores': [],\n            'true_vs_pred': [],\n            'level_difficulty': defaultdict(list),\n            'condition_difficulty': defaultdict(list)\n        }\n        \n        print(\"Analyzing prediction patterns...\")\n        processed_samples = 0\n        \n        for idx in sample_indices:\n            try:\n                # Get sample data\n                image = self.preprocessed_data['images'][idx]\n                condition = self.preprocessed_data['conditions'][idx]\n                level = self.preprocessed_data['levels'][idx]\n                true_severity = self.preprocessed_data['severities'][idx]\n                \n                # Get prediction\n                prediction = self.pipeline.predict(image, condition, level)\n                pred_class = prediction['severity_prediction'].item()\n                confidence = prediction['confidence'].item()\n                \n                # Store results\n                results['confidence_scores'].append(confidence)\n                results['true_vs_pred'].append((true_severity, self.severity_classes[pred_class]))\n                results['condition_level_accuracy'][condition][level].append(\n                    true_severity == self.severity_classes[pred_class]\n                )\n                results['level_difficulty'][level].append(confidence)\n                results['condition_difficulty'][condition].append(confidence)\n                \n                processed_samples += 1\n                if processed_samples % 100 == 0:\n                    print(f\"Processed {processed_samples}/{num_samples} samples...\")\n                \n            except Exception as e:\n                print(f\"Error processing sample {idx}: {str(e)}\")\n                continue\n        \n        print(f\"\\nCompleted analysis of {processed_samples} samples\")\n        self._visualize_statistics(results)\n        return results\n    \n    def _visualize_statistics(self, results: Dict):\n        \"\"\"Visualize statistical analysis results with fixed accuracy calculations\"\"\"\n        # 1. Overall Confidence Distribution\n        plt.figure(figsize=(15, 5))\n        plt.subplot(1, 3, 1)\n        sns.histplot(results['confidence_scores'], bins=30)\n        plt.title('Confidence Score Distribution')\n        plt.xlabel('Confidence')\n        plt.ylabel('Count')\n        \n        # 2. Level-wise Accuracy\n        level_acc = {}\n        for level in set([l for d in results['condition_level_accuracy'].values() for l in d.keys()]):\n            level_values = []\n            for cond_data in results['condition_level_accuracy'].values():\n                if level in cond_data:\n                    level_values.extend(cond_data[level])\n            if level_values:\n                level_acc[level] = np.mean(level_values)\n        \n        plt.subplot(1, 3, 2)\n        if level_acc:\n            plt.bar(level_acc.keys(), level_acc.values())\n            plt.title('Accuracy by Spinal Level')\n            plt.xticks(rotation=45)\n        \n        # 3. Condition-wise Accuracy\n        condition_acc = {}\n        for condition, level_data in results['condition_level_accuracy'].items():\n            condition_values = []\n            for level_list in level_data.values():\n                condition_values.extend(level_list)\n            if condition_values:\n                condition_acc[condition] = np.mean(condition_values)\n        \n        plt.subplot(1, 3, 3)\n        if condition_acc:\n            plt.bar(range(len(condition_acc)), condition_acc.values())\n            plt.xticks(range(len(condition_acc)), \n                      [c.split()[0] for c in condition_acc.keys()], \n                      rotation=45)\n            plt.title('Accuracy by Condition')\n        \n        plt.tight_layout()\n        plt.show()\n        \n        # 4. Confusion Matrix\n        true_labels = [t for t, _ in results['true_vs_pred']]\n        pred_labels = [p for _, p in results['true_vs_pred']]\n        \n        cm = confusion_matrix(true_labels, pred_labels, \n                            labels=self.severity_classes)\n        plt.figure(figsize=(8, 6))\n        sns.heatmap(cm, annot=True, fmt='d', \n                   xticklabels=self.severity_classes,\n                   yticklabels=self.severity_classes)\n        plt.title('Confusion Matrix')\n        plt.xlabel('Predicted')\n        plt.ylabel('True')\n        plt.show()\n        \n        # Print summary statistics\n        print(\"\\nSummary Statistics:\")\n        print(f\"Total Samples Analyzed: {len(results['confidence_scores'])}\")\n        print(f\"Average Confidence: {np.mean(results['confidence_scores']):.4f}\")\n        print(f\"Confidence Std Dev: {np.std(results['confidence_scores']):.4f}\")\n        \n        print(\"\\nAccuracy by Level:\")\n        for level, acc in level_acc.items():\n            print(f\"{level}: {acc:.4f}\")\n        \n        print(\"\\nAccuracy by Condition:\")\n        for cond, acc in condition_acc.items():\n            print(f\"{cond}: {acc:.4f}\")\n            \n        # Add detailed analysis for challenging cases\n        print(\"\\nDetailed Analysis of Challenging Cases:\")\n        correct_predictions = sum(t == p for t, p in results['true_vs_pred'])\n        total_predictions = len(results['true_vs_pred'])\n        overall_accuracy = correct_predictions / total_predictions if total_predictions > 0 else 0\n        \n        print(f\"Overall Accuracy: {overall_accuracy:.4f}\")\n        \n        # Analyze accuracy by severity\n        severity_acc = defaultdict(lambda: {'correct': 0, 'total': 0})\n        for true, pred in results['true_vs_pred']:\n            severity_acc[true]['total'] += 1\n            if true == pred:\n                severity_acc[true]['correct'] += 1\n        \n        print(\"\\nAccuracy by Severity:\")\n        for severity, counts in severity_acc.items():\n            acc = counts['correct'] / counts['total'] if counts['total'] > 0 else 0\n            print(f\"{severity}: {acc:.4f} ({counts['correct']}/{counts['total']})\")\n    \n    def analyze_challenging_cases(self, num_cases: int = 5):\n        \"\"\"Find and analyze particularly challenging cases\"\"\"\n        print(\"\\nAnalyzing challenging cases...\")\n        challenging_cases = []\n        \n        # Define challenging criteria\n        challenging_criteria = [\n            ('L4_L5', 'Severe'),    # Known difficult level with severe cases\n            ('L5_S1', 'Moderate'),  # Transition cases at difficult level\n            ('L3_L4', 'Severe'),    # Higher level severe cases\n            ('L4_L5', 'Moderate')   # Moderate cases at difficult level\n        ]\n        \n        for criterion_level, criterion_severity in challenging_criteria:\n            # Find matching cases\n            for idx in range(len(self.preprocessed_data['images'])):\n                if len(challenging_cases) >= num_cases:\n                    break\n                    \n                level = self.preprocessed_data['levels'][idx]\n                severity = self.preprocessed_data['severities'][idx]\n                \n                if level == criterion_level and severity == criterion_severity:\n                    challenging_cases.append(idx)\n        \n        # Analyze challenging cases\n        print(f\"\\nAnalyzing {len(challenging_cases)} challenging cases:\")\n        for idx in challenging_cases:\n            image = self.preprocessed_data['images'][idx]\n            condition = self.preprocessed_data['conditions'][idx]\n            level = self.preprocessed_data['levels'][idx]\n            true_severity = self.preprocessed_data['severities'][idx]\n            \n            # Get prediction\n            prediction = self.pipeline.predict(image, condition, level)\n            probs = prediction['probabilities'].squeeze().cpu().numpy()\n            pred_class = prediction['severity_prediction'].item()\n            confidence = prediction['confidence'].item()\n            \n            # Plot results\n            plt.figure(figsize=(10, 5))\n            \n            # Plot image\n            plt.subplot(1, 2, 1)\n            plt.imshow(image.squeeze(), cmap='gray')\n            plt.title(f\"Challenging Case\\n{condition}\\nLevel: {level}\")\n            plt.axis('off')\n            \n            # Plot prediction probabilities\n            plt.subplot(1, 2, 2)\n            colors = ['green', 'yellow', 'red']\n            bars = plt.bar(self.severity_classes, probs, color=colors, alpha=0.6)\n            plt.ylim(0, 1)\n            plt.title(f\"Predictions\\nTrue: {true_severity}\\n\" +\n                     f\"Predicted: {self.severity_classes[pred_class]}\\n\" +\n                     f\"Confidence: {confidence:.2f}\")\n            \n            # Add value labels\n            for bar in bars:\n                height = bar.get_height()\n                plt.text(bar.get_x() + bar.get_width()/2., height,\n                        f'{height:.2f}',\n                        ha='center', va='bottom')\n            \n            plt.tight_layout()\n            plt.show()\n            \n            # Print analysis\n            print(f\"\\nChallenging Case Analysis:\")\n            print(f\"Condition: {condition}\")\n            print(f\"Level: {level}\")\n            print(f\"True Severity: {true_severity}\")\n            print(f\"Predicted: {self.severity_classes[pred_class]}\")\n            print(f\"Confidence: {confidence:.4f}\")\n            print(\"\\nProbability Distribution:\")\n            for cls, prob in zip(self.severity_classes, probs):\n                print(f\"{cls}: {prob:.4f}\")\n            print(\"-\" * 50)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T14:15:02.463083Z","iopub.execute_input":"2024-11-21T14:15:02.463469Z","iopub.status.idle":"2024-11-21T14:15:02.491873Z","shell.execute_reply.started":"2024-11-21T14:15:02.463436Z","shell.execute_reply":"2024-11-21T14:15:02.491039Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # Set device\n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"Using device: {device}\")\n    \n    # Load model and create pipeline\n    model = LumbarClassifier().to(device)\n    model.load_state_dict(torch.load('best_model.pth')['model_state_dict'])\n    model.eval()\n    \n    pipeline = ModifiedOptimizedPredictionPipeline(model, device)\n    \n    # Create analyzer\n    analyzer = AdvancedAnalysis(\n        pipeline=pipeline,\n        preprocessed_path='/kaggle/working/val_processed.npy'\n    )\n    \n    # Run statistical analysis\n    print(\"\\nRunning Statistical Analysis...\")\n    stats_results = analyzer.analyze_prediction_patterns(num_samples=500)\n    \n    # Analyze challenging cases\n    print(\"\\nAnalyzing Challenging Cases...\")\n    analyzer.analyze_challenging_cases(num_cases=5)\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2024-11-21T14:15:08.488842Z","iopub.execute_input":"2024-11-21T14:15:08.489587Z","iopub.status.idle":"2024-11-21T14:15:23.813250Z","shell.execute_reply.started":"2024-11-21T14:15:08.489548Z","shell.execute_reply":"2024-11-21T14:15:23.812413Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"Overall Model Performance:\n\n* Overall Accuracy: 81.80%\n* High Average Confidence: 0.9445 (94.45%)\n* Reasonable Confidence Std Dev: 0.1143\n\nLevel-wise Performance (from best to worst):\n* L1_L2: 96.39%   (Best - Upper level)\n* L2_L3: 87.91%   (Good - Upper level)\n* L3_L4: 79.81%   (Moderate)\n* L5_S1: 77.78%   (Challenging)\n* L4_L5: 71.93%   (Most challenging)\n\nCondition-wise Performance:\n* Spinal Canal Stenosis:          86.52% (Best)\n* Right Neural Foraminal:         84.26%\n* Left Neural Foraminal:          83.17%\n* Right Subarticular Stenosis:    79.80%\n* Left Subarticular Stenosis:     75.73% (Most challenging)\n\nSeverity-wise Performance:\n* Normal/Mild: 89.87% (337/375 cases)\n* Severe:      65.62% (21/32 cases)\n* Moderate:    54.84% (51/93 cases)\n\nChallenging Cases Analysis:\na) Successful Severe Predictions:\n\nCase 2: Spinal Canal Stenosis (L4_L5)\n* Confidence: 0.9985\n* Very clear prediction\n\nCase 3: Left Subarticular Stenosis (L4_L5)\n* Confidence: 0.9958\n* Strong prediction\n\nCase 5: Right Neural Foraminal (L4_L5)\n* Confidence: 0.9973\n* Excellent prediction\n\nb) Misclassifications:\n\nCase 1: Right Neural Foraminal (L4_L5)\n* True: Severe, Predicted: Moderate\n* High confidence but wrong (0.9841)\n\nCase 4: Spinal Canal Stenosis (L4_L5)\n* True: Severe, Predicted: Normal/Mild\n* Lower confidence (0.7331)\n\nKey Insights:\n\nStrong Points:\n* Excellent at Normal/Mild cases (89.87%)\n* Very good at upper levels (L1_L2, L2_L3)\n* High confidence in predictions\n\nAreas for Improvement:\n* Moderate cases (54.84% accuracy)\n* L4_L5 level performance (71.93%)\n* Left Subarticular Stenosis (75.73%)\n\nCritical Patterns:\n* Model struggles with borderline cases\n* High confidence doesn't always mean correct prediction\n* L4_L5 level is consistently challenging\n\nRecommendations:\n\nModel Improvements:\n* Focus on better moderate case detection\n* Add more attention to L4_L5 level cases\n* Consider ensemble approach for borderline cases\n\nClinical Application:\n* High confidence in normal/mild cases\n* Double-check high-confidence predictions at L4_L5\n* Use confidence scores for decision support\n","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}