{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.12.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":36363,"databundleVersionId":4050810,"sourceType":"competition"},{"sourceId":4264054,"sourceType":"datasetVersion","datasetId":2406209}],"dockerImageVersionId":31236,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Install DICOM decompression libraries\n!pip install -q gdcm\n!pip install -q pylibjpeg pylibjpeg-libjpeg\n\n# Then restart the kernel and run again","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-12-30T17:05:49.779233Z","iopub.execute_input":"2025-12-30T17:05:49.779576Z","iopub.status.idle":"2025-12-30T17:05:56.554356Z","shell.execute_reply.started":"2025-12-30T17:05:49.779544Z","shell.execute_reply":"2025-12-30T17:05:56.553302Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nComplete Preprocessing Pipeline for RSNA 2022 Cervical Spine Fracture Detection\nThis script handles DICOM loading, HU conversion, windowing, resampling, and standardization\n\"\"\"\n\nimport os\nimport numpy as np\nimport pydicom\nfrom glob import glob\nfrom scipy.ndimage import zoom\nimport matplotlib.pyplot as plt\nfrom tqdm import tqdm\nimport pandas as pd\n\n# ============================================================================\n# STEP 1: DICOM LOADING AND METADATA EXTRACTION\n# ============================================================================\n\ndef load_dicom_series(patient_folder):\n    \"\"\"\n    Load all DICOM slices for a patient and stack into 3D volume\n    \n    Args:\n        patient_folder: Path to folder containing .dcm files\n        \n    Returns:\n        volume: 3D numpy array (num_slices, height, width)\n        slice_thickness: Slice thickness in mm\n        pixel_spacing: (row_spacing, col_spacing) in mm\n        metadata: First DICOM slice metadata\n    \"\"\"\n    # Get all .dcm files\n    dicom_files = glob(os.path.join(patient_folder, \"*.dcm\"))\n    \n    if len(dicom_files) == 0:\n        raise ValueError(f\"No DICOM files found in {patient_folder}\")\n    \n    # Read all slices\n    slices = []\n    for dcm_file in dicom_files:\n        try:\n            ds = pydicom.dcmread(dcm_file)\n            slices.append(ds)\n        except Exception as e:\n            print(f\"Error reading {dcm_file}: {e}\")\n            continue\n    \n    if len(slices) == 0:\n        raise ValueError(f\"Could not read any DICOM files from {patient_folder}\")\n    \n    # Sort by ImagePositionPatient (Z coordinate - superior/inferior position)\n    slices.sort(key=lambda x: float(x.ImagePositionPatient[2]))\n    \n    # Stack into 3D array\n    volume = np.stack([s.pixel_array for s in slices])\n    \n    # Get metadata from first slice\n    metadata = slices[0]\n    \n    # Extract spacing information\n    try:\n        slice_thickness = float(metadata.SliceThickness)\n    except:\n        # If SliceThickness not available, calculate from positions\n        if len(slices) > 1:\n            slice_thickness = abs(\n                float(slices[1].ImagePositionPatient[2]) - \n                float(slices[0].ImagePositionPatient[2])\n            )\n        else:\n            slice_thickness = 1.0  # Default\n    \n    pixel_spacing = [float(x) for x in metadata.PixelSpacing]\n    \n    return volume, slice_thickness, pixel_spacing, metadata\n\n\n# ============================================================================\n# STEP 2: HOUNSFIELD UNIT (HU) CONVERSION\n# ============================================================================\n\ndef apply_hu_conversion(volume, metadata):\n    \"\"\"\n    Convert raw pixel values to Hounsfield Units (HU)\n    HU = pixel_value * slope + intercept\n    \n    Args:\n        volume: 3D numpy array of raw pixel values\n        metadata: DICOM metadata containing RescaleSlope and RescaleIntercept\n        \n    Returns:\n        volume_hu: 3D numpy array in Hounsfield Units\n    \"\"\"\n    try:\n        intercept = float(metadata.RescaleIntercept)\n        slope = float(metadata.RescaleSlope)\n    except:\n        intercept = 0.0\n        slope = 1.0\n        print(\"Warning: RescaleIntercept/Slope not found, using defaults\")\n    \n    volume_hu = volume.astype(np.float32) * slope + intercept\n    return volume_hu\n\n\n# ============================================================================\n# STEP 3: WINDOWING FOR BONE VISUALIZATION\n# ============================================================================\n\ndef apply_window(volume_hu, window_center=400, window_width=1800):\n    \"\"\"\n    Apply windowing to enhance bone structures\n    Standard bone window: WC=400, WW=1800 (range: -500 to 1300 HU)\n    \n    Args:\n        volume_hu: 3D numpy array in Hounsfield Units\n        window_center: Center of the window (HU)\n        window_width: Width of the window (HU)\n        \n    Returns:\n        volume_windowed: 3D numpy array normalized to [0, 1]\n    \"\"\"\n    lower = window_center - window_width // 2\n    upper = window_center + window_width // 2\n    \n    # Clip values to window range\n    volume_windowed = np.clip(volume_hu, lower, upper)\n    \n    # Normalize to [0, 1]\n    volume_normalized = (volume_windowed - lower) / (upper - lower)\n    \n    return volume_normalized.astype(np.float32)\n\n\n# ============================================================================\n# STEP 4: RESAMPLING TO ISOTROPIC SPACING (IMPROVED)\n# ============================================================================\n\ndef resample_volume(volume, current_spacing, target_spacing=(1.5, 1.0, 1.0)):\n    \"\"\"\n    Resample volume to target spacing with better depth preservation\n    Uses slightly larger Z-spacing (1.5mm) to preserve more slices\n    \n    Args:\n        volume: 3D numpy array (D, H, W)\n        current_spacing: (z_spacing, y_spacing, x_spacing) in mm\n        target_spacing: Desired spacing in mm (default: 1.5mm z, 1mm x,y)\n        \n    Returns:\n        resampled_volume: 3D numpy array with target spacing\n    \"\"\"\n    # Calculate resize factors for each dimension\n    resize_factor = np.array(current_spacing) / np.array(target_spacing)\n    \n    # Resample using trilinear interpolation (order=1)\n    # order=0: nearest neighbor, order=1: bilinear, order=3: cubic\n    resampled_volume = zoom(volume, resize_factor, order=1)\n    \n    return resampled_volume.astype(np.float32)\n\n\n# ============================================================================\n# STEP 5: CROP OR PAD TO TARGET SHAPE\n# ============================================================================\n\n# ============================================================================\n# STEP 5: IMPROVED CROP OR PAD WITH CERVICAL SPINE FOCUS\n# ============================================================================\n\ndef crop_or_pad_cervical(volume, target_shape=(96, 320, 320), cervical_focus=True):\n    \"\"\"\n    Crop or pad volume to target shape with focus on cervical spine region\n    Cervical spine is typically in the upper portion of CT scans\n    \n    Args:\n        volume: 3D numpy array (D, H, W)\n        target_shape: Desired output shape (D, H, W)\n        cervical_focus: If True, crop from top of volume (cervical region)\n        \n    Returns:\n        output: 3D numpy array with target shape\n    \"\"\"\n    current_shape = np.array(volume.shape)\n    target_shape_arr = np.array(target_shape)\n    \n    # Initialize output with zeros (black background)\n    output = np.zeros(target_shape, dtype=volume.dtype)\n    \n    # Calculate crop/pad slices for each dimension\n    slices_vol = []\n    slices_out = []\n    \n    for i in range(3):\n        if current_shape[i] >= target_shape_arr[i]:\n            # Crop\n            if i == 0 and cervical_focus:\n                # For depth (Z-axis): take from TOP (cervical region)\n                # Cervical spine is in upper ~30-40% of typical CT scan\n                start = int(current_shape[i] * 0.15)  # Skip very top (air)\n                start = max(0, min(start, current_shape[i] - target_shape_arr[i]))\n            else:\n                # For height/width: take from center\n                start = (current_shape[i] - target_shape_arr[i]) // 2\n            \n            slices_vol.append(slice(start, start + target_shape_arr[i]))\n            slices_out.append(slice(0, target_shape_arr[i]))\n        else:\n            # Pad - place in center\n            start = (target_shape_arr[i] - current_shape[i]) // 2\n            slices_vol.append(slice(0, current_shape[i]))\n            slices_out.append(slice(start, start + current_shape[i]))\n    \n    # Apply cropping/padding in one operation\n    output[slices_out[0], slices_out[1], slices_out[2]] = \\\n        volume[slices_vol[0], slices_vol[1], slices_vol[2]]\n    \n    return output\n\n\ndef crop_or_pad(volume, target_shape=(96, 320, 320)):\n    \"\"\"\n    Standard crop or pad (center-based) - kept for backward compatibility\n    For cervical spine, use crop_or_pad_cervical instead\n    \n    Args:\n        volume: 3D numpy array (D, H, W)\n        target_shape: Desired output shape (D, H, W)\n        \n    Returns:\n        output: 3D numpy array with target shape\n    \"\"\"\n    current_shape = np.array(volume.shape)\n    target_shape_arr = np.array(target_shape)\n    \n    # Initialize output with zeros (black background)\n    output = np.zeros(target_shape, dtype=volume.dtype)\n    \n    # Calculate crop/pad slices for each dimension\n    slices_vol = []\n    slices_out = []\n    \n    for i in range(3):\n        if current_shape[i] >= target_shape_arr[i]:\n            # Crop - take from the center\n            start = (current_shape[i] - target_shape_arr[i]) // 2\n            slices_vol.append(slice(start, start + target_shape_arr[i]))\n            slices_out.append(slice(0, target_shape_arr[i]))\n        else:\n            # Pad - place in center\n            start = (target_shape_arr[i] - current_shape[i]) // 2\n            slices_vol.append(slice(0, current_shape[i]))\n            slices_out.append(slice(start, start + current_shape[i]))\n    \n    # Apply cropping/padding in one operation\n    output[slices_out[0], slices_out[1], slices_out[2]] = \\\n        volume[slices_vol[0], slices_vol[1], slices_vol[2]]\n    \n    return output\n\n\n# ============================================================================\n# STEP 6: IMPROVED PREPROCESSING PIPELINE WITH BETTER RESOLUTION\n# ============================================================================\n\ndef preprocess_patient(patient_folder, target_shape=(96, 320, 320), \n                       window_center=400, window_width=1800,\n                       target_spacing=(1.5, 1.0, 1.0),\n                       cervical_focus=True):\n    \"\"\"\n    Complete preprocessing pipeline for one patient (IMPROVED VERSION)\n    \n    Improvements:\n    - Larger target shape (96, 320, 320) for better spatial resolution\n    - Less aggressive depth resampling (1.5mm vs 1.0mm)\n    - Cervical spine focused cropping\n    - Better preservation of anatomical details\n    \n    Steps:\n    1. Load DICOM series\n    2. Convert to Hounsfield Units\n    3. Apply bone windowing\n    4. Resample to target spacing (1.5mm z, 1mm x,y)\n    5. Crop/pad to target shape with cervical focus\n    \n    Args:\n        patient_folder: Path to patient's DICOM folder\n        target_shape: Output shape (D, H, W) - default (96, 320, 320)\n        window_center: HU window center for bone\n        window_width: HU window width for bone\n        target_spacing: Target voxel spacing (z, y, x) in mm\n        cervical_focus: Whether to focus on cervical region when cropping\n        \n    Returns:\n        volume_final: Preprocessed 3D volume (D, H, W) in range [0, 1]\n    \"\"\"\n    try:\n        # Step 1: Load DICOM series\n        volume, slice_thickness, pixel_spacing, metadata = load_dicom_series(patient_folder)\n        \n        # Step 2: Convert to Hounsfield Units\n        volume_hu = apply_hu_conversion(volume, metadata)\n        \n        # Step 3: Apply bone windowing\n        volume_windowed = apply_window(volume_hu, window_center, window_width)\n        \n        # Step 4: Resample to target spacing (less aggressive)\n        current_spacing = (slice_thickness, pixel_spacing[0], pixel_spacing[1])\n        volume_resampled = resample_volume(volume_windowed, current_spacing, target_spacing)\n        \n        # Step 5: Crop or pad to target shape (with cervical focus)\n        if cervical_focus:\n            volume_final = crop_or_pad_cervical(volume_resampled, target_shape, cervical_focus=True)\n        else:\n            volume_final = crop_or_pad(volume_resampled, target_shape)\n        \n        return volume_final\n        \n    except Exception as e:\n        print(f\"Error preprocessing {patient_folder}: {e}\")\n        # Return zero volume on error\n        return np.zeros(target_shape, dtype=np.float32)\n\n\ndef preprocess_patient_memory_efficient(patient_folder, target_shape=(64, 256, 256),\n                                        window_center=400, window_width=1800,\n                                        target_spacing=(2.0, 1.0, 1.0)):\n    \"\"\"\n    Memory-efficient version for Kaggle with limited GPU memory\n    Uses smaller target shape and more aggressive depth resampling\n    \n    Use this if you get CUDA out of memory errors\n    \n    Args:\n        patient_folder: Path to patient's DICOM folder\n        target_shape: Smaller output shape (D, H, W) - default (64, 256, 256)\n        window_center: HU window center for bone\n        window_width: HU window width for bone\n        target_spacing: More aggressive spacing (2mm z, 1mm x,y)\n        \n    Returns:\n        volume_final: Preprocessed 3D volume (D, H, W) in range [0, 1]\n    \"\"\"\n    return preprocess_patient(\n        patient_folder=patient_folder,\n        target_shape=target_shape,\n        window_center=window_center,\n        window_width=window_width,\n        target_spacing=target_spacing,\n        cervical_focus=True\n    )\n\n\n# ============================================================================\n# STEP 7: BATCH PREPROCESSING WITH MULTIPLE RESOLUTION OPTIONS\n# ============================================================================\n\ndef preprocess_all_patients(train_csv_path, train_images_root, output_dir,\n                            target_shape=(96, 320, 320), save_npy=True,\n                            resolution_mode='high'):\n    \"\"\"\n    Preprocess all patients and optionally save to disk\n    \n    Resolution modes:\n    - 'high': (96, 320, 320) - Best quality, requires more memory\n    - 'medium': (80, 256, 256) - Balanced quality and memory\n    - 'low': (64, 224, 224) - Memory efficient\n    \n    Args:\n        train_csv_path: Path to train.csv\n        train_images_root: Root directory of train_images\n        output_dir: Directory to save preprocessed volumes\n        target_shape: Target volume shape (overrides resolution_mode if specified)\n        save_npy: Whether to save preprocessed volumes as .npy files\n        resolution_mode: 'high', 'medium', or 'low'\n        \n    Returns:\n        None (saves files to disk)\n    \"\"\"\n    # Set target shape and spacing based on resolution mode\n    if resolution_mode == 'high':\n        target_shape = (96, 320, 320)\n        target_spacing = (1.5, 1.0, 1.0)\n        print(\"Using HIGH RESOLUTION mode: (96, 320, 320)\")\n    elif resolution_mode == 'medium':\n        target_shape = (80, 256, 256)\n        target_spacing = (1.75, 1.0, 1.0)\n        print(\"Using MEDIUM RESOLUTION mode: (80, 256, 256)\")\n    elif resolution_mode == 'low':\n        target_shape = (64, 224, 224)\n        target_spacing = (2.0, 1.25, 1.25)\n        print(\"Using LOW RESOLUTION mode: (64, 224, 224)\")\n    \n    # Load training CSV\n    train_df = pd.read_csv(train_csv_path)\n    patient_ids = train_df['StudyInstanceUID'].unique()\n    \n    print(f\"Preprocessing {len(patient_ids)} patients...\")\n    print(f\"Target shape: {target_shape}\")\n    print(f\"Target spacing: {target_spacing}\")\n    \n    # Create output directory\n    if save_npy:\n        os.makedirs(output_dir, exist_ok=True)\n    \n    # Process each patient\n    successful = 0\n    failed = 0\n    \n    for patient_id in tqdm(patient_ids, desc=\"Processing patients\"):\n        patient_folder = os.path.join(train_images_root, str(patient_id))\n        \n        if not os.path.exists(patient_folder):\n            print(f\"Warning: Patient folder not found: {patient_folder}\")\n            failed += 1\n            continue\n        \n        try:\n            # Preprocess\n            volume = preprocess_patient(\n                patient_folder, \n                target_shape=target_shape,\n                target_spacing=target_spacing,\n                cervical_focus=True\n            )\n            \n            # Save if requested\n            if save_npy:\n                output_path = os.path.join(output_dir, f\"{patient_id}.npy\")\n                np.save(output_path, volume)\n            \n            successful += 1\n            \n        except Exception as e:\n            print(f\"Failed to process {patient_id}: {e}\")\n            failed += 1\n    \n    print(f\"\\nPreprocessing complete!\")\n    print(f\"Successful: {successful}\")\n    print(f\"Failed: {failed}\")\n\n\n# ============================================================================\n# STEP 8: IMPROVED VISUALIZATION UTILITIES\n# ============================================================================\n\ndef visualize_preprocessing_steps(patient_folder, save_path=None, \n                                  target_shape=(96, 320, 320)):\n    \"\"\"\n    Visualize each step of the preprocessing pipeline (IMPROVED)\n    Shows better spatial resolution with new settings\n    \"\"\"\n    # Load original\n    volume, thickness, spacing, metadata = load_dicom_series(patient_folder)\n    \n    # Apply each step\n    volume_hu = apply_hu_conversion(volume, metadata)\n    volume_windowed = apply_window(volume_hu)\n    current_spacing = (thickness, spacing[0], spacing[1])\n    volume_resampled = resample_volume(volume_windowed, current_spacing, \n                                       target_spacing=(1.5, 1.0, 1.0))\n    volume_final = crop_or_pad_cervical(volume_resampled, target_shape, cervical_focus=True)\n    \n    # Select middle slices\n    mid_original = len(volume) // 2\n    mid_final = volume_final.shape[0] // 2\n    \n    # Create visualization\n    fig, axes = plt.subplots(2, 3, figsize=(18, 12))\n    \n    # Original\n    axes[0, 0].imshow(volume[mid_original], cmap='gray')\n    axes[0, 0].set_title(f'Original\\nShape: {volume.shape}\\nSpacing: {current_spacing}', fontsize=12)\n    axes[0, 0].axis('off')\n    \n    # HU converted\n    axes[0, 1].imshow(volume_hu[mid_original], cmap='gray', vmin=-1000, vmax=1000)\n    axes[0, 1].set_title(f'HU Converted\\nRange: [{volume_hu.min():.0f}, {volume_hu.max():.0f}] HU', fontsize=12)\n    axes[0, 1].axis('off')\n    \n    # Windowed\n    axes[0, 2].imshow(volume_windowed[mid_original], cmap='gray')\n    axes[0, 2].set_title(f'Bone Windowed\\nWC=400, WW=1800\\nRange: [0, 1]', fontsize=12)\n    axes[0, 2].axis('off')\n    \n    # Resampled\n    axes[1, 0].imshow(volume_resampled[volume_resampled.shape[0]//2], cmap='gray')\n    axes[1, 0].set_title(f'Resampled to 1.5mm isotropic\\nShape: {volume_resampled.shape}', fontsize=12)\n    axes[1, 0].axis('off')\n    \n    # Final\n    axes[1, 1].imshow(volume_final[mid_final], cmap='gray')\n    axes[1, 1].set_title(f'Final (Cervical Focused)\\nShape: {volume_final.shape}', fontsize=12)\n    axes[1, 1].axis('off')\n    \n    # 3D view (MIP - Maximum Intensity Projection)\n    mip = np.max(volume_final, axis=0)\n    axes[1, 2].imshow(mip, cmap='gray')\n    axes[1, 2].set_title('MIP (Max Projection)\\nAxial View', fontsize=12)\n    axes[1, 2].axis('off')\n    \n    plt.suptitle('Preprocessing Pipeline - Improved Resolution', fontsize=16, fontweight='bold')\n    plt.tight_layout()\n    \n    if save_path:\n        plt.savefig(save_path, dpi=150, bbox_inches='tight')\n        print(f\"Visualization saved to {save_path}\")\n    \n    plt.show()\n    \n    # Print statistics\n    print(f\"\\n{'='*60}\")\n    print(f\"PREPROCESSING STATISTICS\")\n    print(f\"{'='*60}\")\n    print(f\"Original shape:        {volume.shape}\")\n    print(f\"Original spacing:      {current_spacing} mm\")\n    print(f\"Resampled shape:       {volume_resampled.shape}\")\n    print(f\"Final shape:           {volume_final.shape}\")\n    print(f\"Compression ratio:     {volume.size / volume_final.size:.2f}x\")\n    print(f\"Memory (original):     {volume.nbytes / 1024 / 1024:.2f} MB\")\n    print(f\"Memory (final):        {volume_final.nbytes / 1024 / 1024:.2f} MB\")\n    print(f\"{'='*60}\\n\")\n\n\ndef visualize_volume_slices(volume, num_slices=9, title=\"Volume Slices\"):\n    \"\"\"\n    Visualize multiple slices from a 3D volume\n    \"\"\"\n    fig, axes = plt.subplots(3, 3, figsize=(12, 12))\n    axes = axes.flatten()\n    \n    # Select evenly spaced slices\n    slice_indices = np.linspace(0, volume.shape[0]-1, num_slices, dtype=int)\n    \n    for idx, slice_idx in enumerate(slice_indices):\n        axes[idx].imshow(volume[slice_idx], cmap='gray')\n        axes[idx].set_title(f'Slice {slice_idx}/{volume.shape[0]}')\n        axes[idx].axis('off')\n    \n    plt.suptitle(title, fontsize=16)\n    plt.tight_layout()\n    plt.show()\n\n\n# ============================================================================\n# EXAMPLE USAGE WITH IMPROVED SETTINGS\n# ============================================================================\n\nif __name__ == \"__main__\":\n    \n    print(\"=\"*80)\n    print(\"IMPROVED PREPROCESSING PIPELINE - BETTER SPATIAL RESOLUTION\")\n    print(\"=\"*80)\n    \n    # Example: Preprocess a single patient with HIGH RESOLUTION\n    patient_id = \"1.2.826.0.1.3680043.10001\"\n    patient_folder = f\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images/{patient_id}\"\n    \n    print(\"\\n1. HIGH RESOLUTION MODE (96, 320, 320) - RECOMMENDED\")\n    print(\"-\" * 60)\n    volume_high = preprocess_patient(\n        patient_folder, \n        target_shape=(96, 320, 320),\n        target_spacing=(1.5, 1.0, 1.0),\n        cervical_focus=True\n    )\n    print(f\"✓ Preprocessed volume shape: {volume_high.shape}\")\n    print(f\"✓ Memory usage: {volume_high.nbytes / 1024 / 1024:.2f} MB\")\n    print(f\"✓ Value range: [{volume_high.min():.3f}, {volume_high.max():.3f}]\")\n    \n    print(\"\\n2. MEDIUM RESOLUTION MODE (80, 256, 256) - BALANCED\")\n    print(\"-\" * 60)\n    volume_medium = preprocess_patient(\n        patient_folder, \n        target_shape=(80, 256, 256),\n        target_spacing=(1.75, 1.0, 1.0),\n        cervical_focus=True\n    )\n    print(f\"✓ Preprocessed volume shape: {volume_medium.shape}\")\n    print(f\"✓ Memory usage: {volume_medium.nbytes / 1024 / 1024:.2f} MB\")\n    \n    print(\"\\n3. LOW RESOLUTION MODE (64, 224, 224) - MEMORY EFFICIENT\")\n    print(\"-\" * 60)\n    volume_low = preprocess_patient_memory_efficient(\n        patient_folder, \n        target_shape=(64, 224, 224),\n        target_spacing=(2.0, 1.25, 1.25)\n    )\n    print(f\"✓ Preprocessed volume shape: {volume_low.shape}\")\n    print(f\"✓ Memory usage: {volume_low.nbytes / 1024 / 1024:.2f} MB\")\n    \n    # Visualize preprocessing steps with improved resolution\n    print(\"\\n4. VISUALIZING PREPROCESSING STEPS...\")\n    print(\"-\" * 60)\n    visualize_preprocessing_steps(patient_folder, target_shape=(96, 320, 320))\n    \n    # Visualize volume slices\n    print(\"\\n5. VISUALIZING VOLUME SLICES...\")\n    print(\"-\" * 60)\n    visualize_volume_slices(volume_high, num_slices=9, \n                           title=f\"HIGH RES - Patient {patient_id}\")\n    \n    # Compare resolutions side by side\n    print(\"\\n6. RESOLUTION COMPARISON\")\n    print(\"-\" * 60)\n    fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n    \n    mid_slice = volume_high.shape[0] // 2\n    axes[0].imshow(volume_high[mid_slice], cmap='gray')\n    axes[0].set_title(f'HIGH: (96, 320, 320)\\n{volume_high.nbytes/1024/1024:.1f} MB', \n                     fontsize=14, fontweight='bold')\n    axes[0].axis('off')\n    \n    mid_slice = volume_medium.shape[0] // 2\n    axes[1].imshow(volume_medium[mid_slice], cmap='gray')\n    axes[1].set_title(f'MEDIUM: (80, 256, 256)\\n{volume_medium.nbytes/1024/1024:.1f} MB', \n                     fontsize=14, fontweight='bold')\n    axes[1].axis('off')\n    \n    mid_slice = volume_low.shape[0] // 2\n    axes[2].imshow(volume_low[mid_slice], cmap='gray')\n    axes[2].set_title(f'LOW: (64, 224, 224)\\n{volume_low.nbytes/1024/1024:.1f} MB', \n                     fontsize=14, fontweight='bold')\n    axes[2].axis('off')\n    \n    plt.suptitle('Resolution Comparison - Middle Slice', fontsize=16, fontweight='bold')\n    plt.tight_layout()\n    plt.show()\n    \n    # Recommendations\n    print(\"\\n\" + \"=\"*80)\n    print(\"RECOMMENDATIONS FOR YOUR PROJECT\")\n    print(\"=\"*80)\n    print(\"\"\"\n    ✓ HIGH RESOLUTION (96, 320, 320):\n      - Best for fracture detection accuracy\n      - Good balance between detail and memory\n      - Recommended if you have GPU with 16GB+ VRAM\n      - Batch size: 1-2\n    \n    ✓ MEDIUM RESOLUTION (80, 256, 256):\n      - Good compromise for most systems\n      - Suitable for Kaggle free tier (16GB GPU)\n      - Batch size: 2-4\n    \n    ✓ LOW RESOLUTION (64, 224, 224):\n      - Use only if memory is very limited\n      - May lose some fracture details\n      - Batch size: 4-8\n    \n    RECOMMENDED: Start with HIGH resolution and reduce if you get OOM errors\n    \"\"\")\n    \n    # Example: Batch preprocessing (commented out to avoid long execution)\n    print(\"\\n7. BATCH PREPROCESSING EXAMPLE (commented out):\")\n    print(\"-\" * 60)\n    print(\"\"\"preprocess_all_patients(\n        train_csv_path=\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv\",\n        train_images_root=\"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images\",\n        output_dir=\"/kaggle/working/preprocessed_volumes_high\",\n        resolution_mode='high',  # or 'medium' or 'low'\n        save_npy=True\n    )\"\"\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T17:05:56.556321Z","iopub.execute_input":"2025-12-30T17:05:56.556611Z","iopub.status.idle":"2025-12-30T17:06:06.317269Z","shell.execute_reply.started":"2025-12-30T17:05:56.556579Z","shell.execute_reply":"2025-12-30T17:06:06.316442Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nPyTorch Dataset and DataLoader for RSNA 2022 Cervical Spine Fracture Detection\nSupports on-the-fly preprocessing or loading pre-saved .npy files\nIncludes data augmentation for training\n\"\"\"\n\nimport os\nimport numpy as np\nimport pandas as pd\nimport torch\nfrom torch.utils.data import Dataset, DataLoader\nfrom sklearn.model_selection import StratifiedKFold, train_test_split\nimport random\nimport pydicom\nfrom glob import glob\nfrom scipy.ndimage import zoom\n\n# ============================================================================\n# PREPROCESSING FUNCTIONS (EMBEDDED TO AVOID IMPORT ISSUES)\n# ============================================================================\n\ndef load_dicom_series(patient_folder):\n    \"\"\"Load all DICOM slices for a patient and stack into 3D volume\"\"\"\n    dicom_files = glob(os.path.join(patient_folder, \"*.dcm\"))\n    if len(dicom_files) == 0:\n        raise ValueError(f\"No DICOM files found in {patient_folder}\")\n    \n    slices = []\n    for dcm_file in dicom_files:\n        try:\n            ds = pydicom.dcmread(dcm_file)\n            slices.append(ds)\n        except:\n            continue\n    \n    if len(slices) == 0:\n        raise ValueError(f\"Could not read any DICOM files from {patient_folder}\")\n    \n    slices.sort(key=lambda x: float(x.ImagePositionPatient[2]))\n    volume = np.stack([s.pixel_array for s in slices])\n    \n    metadata = slices[0]\n    try:\n        slice_thickness = float(metadata.SliceThickness)\n    except:\n        if len(slices) > 1:\n            slice_thickness = abs(float(slices[1].ImagePositionPatient[2]) - \n                                 float(slices[0].ImagePositionPatient[2]))\n        else:\n            slice_thickness = 1.0\n    \n    pixel_spacing = [float(x) for x in metadata.PixelSpacing]\n    return volume, slice_thickness, pixel_spacing, metadata\n\n\ndef apply_hu_conversion(volume, metadata):\n    \"\"\"Convert raw pixel values to Hounsfield Units\"\"\"\n    try:\n        intercept = float(metadata.RescaleIntercept)\n        slope = float(metadata.RescaleSlope)\n    except:\n        intercept = 0.0\n        slope = 1.0\n    \n    volume_hu = volume.astype(np.float32) * slope + intercept\n    return volume_hu\n\n\ndef apply_window(volume_hu, window_center=400, window_width=1800):\n    \"\"\"Apply windowing to enhance bone structures\"\"\"\n    lower = window_center - window_width // 2\n    upper = window_center + window_width // 2\n    volume_windowed = np.clip(volume_hu, lower, upper)\n    volume_normalized = (volume_windowed - lower) / (upper - lower)\n    return volume_normalized.astype(np.float32)\n\n\ndef resample_volume(volume, current_spacing, target_spacing=(1.5, 1.0, 1.0)):\n    \"\"\"Resample volume to target spacing\"\"\"\n    resize_factor = np.array(current_spacing) / np.array(target_spacing)\n    resampled_volume = zoom(volume, resize_factor, order=1)\n    return resampled_volume.astype(np.float32)\n\n\ndef crop_or_pad_cervical(volume, target_shape=(96, 320, 320), cervical_focus=True):\n    \"\"\"Crop or pad volume with focus on cervical spine region\"\"\"\n    current_shape = np.array(volume.shape)\n    target_shape_arr = np.array(target_shape)\n    output = np.zeros(target_shape, dtype=volume.dtype)\n    \n    slices_vol = []\n    slices_out = []\n    \n    for i in range(3):\n        if current_shape[i] >= target_shape_arr[i]:\n            if i == 0 and cervical_focus:\n                start = int(current_shape[i] * 0.15)\n                start = max(0, min(start, current_shape[i] - target_shape_arr[i]))\n            else:\n                start = (current_shape[i] - target_shape_arr[i]) // 2\n            \n            slices_vol.append(slice(start, start + target_shape_arr[i]))\n            slices_out.append(slice(0, target_shape_arr[i]))\n        else:\n            start = (target_shape_arr[i] - current_shape[i]) // 2\n            slices_vol.append(slice(0, current_shape[i]))\n            slices_out.append(slice(start, start + current_shape[i]))\n    \n    output[slices_out[0], slices_out[1], slices_out[2]] = \\\n        volume[slices_vol[0], slices_vol[1], slices_vol[2]]\n    \n    return output\n\n\ndef preprocess_patient(patient_folder, target_shape=(96, 320, 320), \n                       window_center=400, window_width=1800,\n                       target_spacing=(1.5, 1.0, 1.0),\n                       cervical_focus=True):\n    \"\"\"Complete preprocessing pipeline for one patient\"\"\"\n    try:\n        volume, slice_thickness, pixel_spacing, metadata = load_dicom_series(patient_folder)\n        volume_hu = apply_hu_conversion(volume, metadata)\n        volume_windowed = apply_window(volume_hu, window_center, window_width)\n        current_spacing = (slice_thickness, pixel_spacing[0], pixel_spacing[1])\n        volume_resampled = resample_volume(volume_windowed, current_spacing, target_spacing)\n        \n        if cervical_focus:\n            volume_final = crop_or_pad_cervical(volume_resampled, target_shape, cervical_focus=True)\n        else:\n            volume_final = crop_or_pad_cervical(volume_resampled, target_shape, cervical_focus=False)\n        \n        return volume_final\n    except Exception as e:\n        print(f\"Error preprocessing {patient_folder}: {e}\")\n        return np.zeros(target_shape, dtype=np.float32)\n\n\n# ============================================================================\n# DATASET CLASS - OPTION 1: ON-THE-FLY PREPROCESSING\n# ============================================================================\n\nclass SpineFractureDataset(Dataset):\n    \"\"\"\n    Dataset that preprocesses DICOM files on-the-fly\n    Use this if you don't want to save all preprocessed volumes to disk\n    \"\"\"\n    \n    def __init__(self, df, image_root, target_shape=(96, 320, 320),\n                 target_spacing=(1.5, 1.0, 1.0), transform=None,\n                 cache_data=False):\n        \"\"\"\n        Args:\n            df: DataFrame with columns [StudyInstanceUID, patient_overall, C1-C7]\n            image_root: Root directory for train_images\n            target_shape: Target volume shape (D, H, W)\n            target_spacing: Target voxel spacing (z, y, x) in mm\n            transform: Optional augmentation function\n            cache_data: Whether to cache preprocessed volumes in memory\n        \"\"\"\n        self.df = df.reset_index(drop=True)\n        self.image_root = image_root\n        self.target_shape = target_shape\n        self.target_spacing = target_spacing\n        self.transform = transform\n        self.cache_data = cache_data\n        \n        # Label columns\n        self.label_cols = ['patient_overall'] + [f'C{i}' for i in range(1, 8)]\n        \n        # Cache for preprocessed volumes\n        self.cache = {} if cache_data else None\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        # Get patient ID\n        patient_id = self.df.loc[idx, 'StudyInstanceUID']\n        \n        # Check cache first\n        if self.cache_data and patient_id in self.cache:\n            volume = self.cache[patient_id]\n        else:\n            # Load and preprocess volume\n            patient_folder = os.path.join(self.image_root, str(patient_id))\n            \n            try:\n                # Use embedded preprocessing function\n                volume = preprocess_patient(\n                    patient_folder,\n                    target_shape=self.target_shape,\n                    target_spacing=self.target_spacing,\n                    cervical_focus=True\n                )\n                \n                # Cache if enabled\n                if self.cache_data:\n                    self.cache[patient_id] = volume\n                    \n            except Exception as e:\n                print(f\"Error loading patient {patient_id}: {e}\")\n                # Return zero volume on error\n                volume = np.zeros(self.target_shape, dtype=np.float32)\n        \n        # Add channel dimension (1, D, H, W)\n        volume = volume[np.newaxis, ...].astype(np.float32)\n        \n        # Get labels\n        labels = self.df.loc[idx, self.label_cols].values.astype(np.float32)\n        \n        # Apply augmentation if provided\n        if self.transform:\n            volume = self.transform(volume)\n        \n        return torch.from_numpy(volume), torch.from_numpy(labels), patient_id\n\n\n# ============================================================================\n# DATASET CLASS - OPTION 2: LOAD PRE-SAVED NPY FILES\n# ============================================================================\n\nclass SpineFractureDatasetNPY(Dataset):\n    \"\"\"\n    Dataset that loads pre-saved .npy files\n    Much faster than on-the-fly preprocessing\n    Use this if you've already preprocessed and saved all volumes\n    \"\"\"\n    \n    def __init__(self, df, npy_root, transform=None):\n        \"\"\"\n        Args:\n            df: DataFrame with columns [StudyInstanceUID, patient_overall, C1-C7]\n            npy_root: Directory containing .npy files\n            transform: Optional augmentation function\n        \"\"\"\n        self.df = df.reset_index(drop=True)\n        self.npy_root = npy_root\n        self.transform = transform\n        \n        # Label columns\n        self.label_cols = ['patient_overall'] + [f'C{i}' for i in range(1, 8)]\n        \n    def __len__(self):\n        return len(self.df)\n    \n    def __getitem__(self, idx):\n        # Get patient ID\n        patient_id = self.df.loc[idx, 'StudyInstanceUID']\n        \n        # Load preprocessed volume\n        npy_path = os.path.join(self.npy_root, f\"{patient_id}.npy\")\n        \n        try:\n            volume = np.load(npy_path)\n        except Exception as e:\n            print(f\"Error loading {npy_path}: {e}\")\n            # Return zero volume on error\n            volume = np.zeros((96, 320, 320), dtype=np.float32)\n        \n        # Add channel dimension (1, D, H, W)\n        volume = volume[np.newaxis, ...].astype(np.float32)\n        \n        # Get labels\n        labels = self.df.loc[idx, self.label_cols].values.astype(np.float32)\n        \n        # Apply augmentation if provided\n        if self.transform:\n            volume = self.transform(volume)\n        \n        return torch.from_numpy(volume), torch.from_numpy(labels), patient_id\n\n\n# ============================================================================\n# DATA AUGMENTATION\n# ============================================================================\n\nclass SpineAugmentation:\n    \"\"\"\n    Data augmentation for 3D CT volumes\n    Includes flips, rotations, noise, and intensity adjustments\n    \"\"\"\n    \n    def __init__(self, flip_prob=0.5, rotate_prob=0.3, noise_prob=0.2,\n                 intensity_prob=0.3):\n        self.flip_prob = flip_prob\n        self.rotate_prob = rotate_prob\n        self.noise_prob = noise_prob\n        self.intensity_prob = intensity_prob\n    \n    def __call__(self, volume):\n        \"\"\"\n        Apply augmentations to volume\n        Args:\n            volume: numpy array (1, D, H, W)\n        Returns:\n            augmented volume\n        \"\"\"\n        # Random horizontal flip (left-right)\n        if random.random() < self.flip_prob:\n            volume = np.flip(volume, axis=3).copy()  # Flip width\n        \n        # Random rotation (small angles)\n        if random.random() < self.rotate_prob:\n            angle = random.uniform(-10, 10)\n            volume = self._rotate_volume(volume, angle)\n        \n        # Random noise\n        if random.random() < self.noise_prob:\n            noise = np.random.normal(0, 0.01, volume.shape).astype(np.float32)\n            volume = np.clip(volume + noise, 0, 1)\n        \n        # Random intensity adjustment\n        if random.random() < self.intensity_prob:\n            factor = random.uniform(0.9, 1.1)\n            volume = np.clip(volume * factor, 0, 1)\n        \n        return volume\n    \n    def _rotate_volume(self, volume, angle):\n        \"\"\"\n        Rotate volume by small angle (in XY plane)\n        \"\"\"\n        from scipy.ndimage import rotate\n        # Rotate each slice in the axial plane\n        rotated = rotate(volume, angle, axes=(2, 3), reshape=False, order=1)\n        return rotated.astype(np.float32)\n\n\n# ============================================================================\n# TRAIN/VALIDATION SPLIT\n# ============================================================================\n\ndef create_train_val_split(train_csv_path, val_size=0.40, random_state=42):\n    \"\"\"\n    Create stratified train/validation split\n    Stratifies by patient_overall to ensure balanced fracture distribution\n    \n    Args:\n        train_csv_path: Path to train.csv\n        val_size: Fraction of data for validation (0.15 = 15%)\n        random_state: Random seed for reproducibility\n        \n    Returns:\n        train_df, val_df: DataFrames for training and validation\n    \"\"\"\n    # Load CSV\n    train_df = pd.read_csv(train_csv_path)\n    \n    # Get unique patient IDs\n    patient_ids = train_df['StudyInstanceUID'].unique()\n    \n    # Get fracture status for each patient (for stratification)\n    patient_fractures = train_df.groupby('StudyInstanceUID')['patient_overall'].first().values\n    \n    # Split patient IDs (not rows)\n    train_ids, val_ids = train_test_split(\n        patient_ids,\n        test_size=val_size,\n        stratify=patient_fractures,\n        random_state=random_state\n    )\n    \n    # Create DataFrames\n    train_df_split = train_df[train_df['StudyInstanceUID'].isin(train_ids)].reset_index(drop=True)\n    val_df_split = train_df[train_df['StudyInstanceUID'].isin(val_ids)].reset_index(drop=True)\n    \n    print(f\"Train patients: {len(train_ids)} ({len(train_df_split)} rows)\")\n    print(f\"Val patients: {len(val_ids)} ({len(val_df_split)} rows)\")\n    print(f\"Train fracture rate: {train_df_split['patient_overall'].mean():.2%}\")\n    print(f\"Val fracture rate: {val_df_split['patient_overall'].mean():.2%}\")\n    \n    return train_df_split, val_df_split\n\n\ndef create_kfold_splits(train_csv_path, n_folds=5, random_state=42):\n    \"\"\"\n    Create K-Fold cross-validation splits\n    Useful for more robust evaluation\n    \n    Args:\n        train_csv_path: Path to train.csv\n        n_folds: Number of folds\n        random_state: Random seed\n        \n    Returns:\n        List of (train_df, val_df) tuples for each fold\n    \"\"\"\n    train_df = pd.read_csv(train_csv_path)\n    patient_ids = train_df['StudyInstanceUID'].unique()\n    patient_fractures = train_df.groupby('StudyInstanceUID')['patient_overall'].first().values\n    \n    skf = StratifiedKFold(n_splits=n_folds, shuffle=True, random_state=random_state)\n    \n    folds = []\n    for fold, (train_idx, val_idx) in enumerate(skf.split(patient_ids, patient_fractures)):\n        train_ids = patient_ids[train_idx]\n        val_ids = patient_ids[val_idx]\n        \n        train_df_fold = train_df[train_df['StudyInstanceUID'].isin(train_ids)].reset_index(drop=True)\n        val_df_fold = train_df[train_df['StudyInstanceUID'].isin(val_ids)].reset_index(drop=True)\n        \n        folds.append((train_df_fold, val_df_fold))\n        print(f\"Fold {fold+1}: Train={len(train_ids)}, Val={len(val_ids)}\")\n    \n    return folds\n\n\n# ============================================================================\n# CREATE DATALOADERS\n# ============================================================================\n\ndef create_dataloaders(train_df, val_df, image_root, batch_size=2,\n                      num_workers=2, use_augmentation=True,\n                      target_shape=(96, 320, 320)):\n    \"\"\"\n    Create PyTorch DataLoaders for training and validation\n    \n    Args:\n        train_df: Training DataFrame\n        val_df: Validation DataFrame\n        image_root: Root directory of train_images\n        batch_size: Batch size (keep small for 3D data)\n        num_workers: Number of worker processes\n        use_augmentation: Whether to use data augmentation for training\n        target_shape: Target volume shape\n        \n    Returns:\n        train_loader, val_loader: PyTorch DataLoaders\n    \"\"\"\n    # Create augmentation\n    train_transform = SpineAugmentation() if use_augmentation else None\n    \n    # Create datasets\n    train_dataset = SpineFractureDataset(\n        df=train_df,\n        image_root=image_root,\n        target_shape=target_shape,\n        transform=train_transform,\n        cache_data=False  # Set to True if you have enough RAM\n    )\n    \n    val_dataset = SpineFractureDataset(\n        df=val_df,\n        image_root=image_root,\n        target_shape=target_shape,\n        transform=None,  # No augmentation for validation\n        cache_data=False\n    )\n    \n    # Create dataloaders\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=batch_size,\n        shuffle=True,\n        num_workers=num_workers,\n        pin_memory=True,\n        drop_last=True  # Drop incomplete batches\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=batch_size,\n        shuffle=False,\n        num_workers=num_workers,\n        pin_memory=True\n    )\n    \n    return train_loader, val_loader\n\n\ndef create_dataloaders_npy(train_df, val_df, npy_root, batch_size=4,\n                           num_workers=2, use_augmentation=True):\n    \"\"\"\n    Create DataLoaders for pre-saved .npy files\n    Much faster than on-the-fly preprocessing\n    \n    Args:\n        train_df: Training DataFrame\n        val_df: Validation DataFrame\n        npy_root: Directory containing .npy files\n        batch_size: Batch size (can be larger with .npy)\n        num_workers: Number of worker processes\n        use_augmentation: Whether to use augmentation\n        \n    Returns:\n        train_loader, val_loader: PyTorch DataLoaders\n    \"\"\"\n    train_transform = SpineAugmentation() if use_augmentation else None\n    \n    train_dataset = SpineFractureDatasetNPY(\n        df=train_df,\n        npy_root=npy_root,\n        transform=train_transform\n    )\n    \n    val_dataset = SpineFractureDatasetNPY(\n        df=val_df,\n        npy_root=npy_root,\n        transform=None\n    )\n    \n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=batch_size,\n        shuffle=True,\n        num_workers=num_workers,\n        pin_memory=True\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=batch_size,\n        shuffle=False,\n        num_workers=num_workers,\n        pin_memory=True\n    )\n    \n    return train_loader, val_loader\n\n\n# ============================================================================\n# EXAMPLE USAGE\n# ============================================================================\n\nif __name__ == \"__main__\":\n    \n    print(\"=\"*80)\n    print(\"CREATING DATALOADERS FOR TRAINING\")\n    print(\"=\"*80)\n    \n    # Paths\n    train_csv_path = \"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv\"\n    image_root = \"/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train_images\"\n    \n    # Create train/val split\n    print(\"\\n1. Creating train/validation split...\")\n    print(\"-\" * 60)\n    train_df, val_df = create_train_val_split(train_csv_path, val_size=0.40)\n    \n    # Option 1: On-the-fly preprocessing (slower but no disk space needed)\n    print(\"\\n2. Creating DataLoaders (on-the-fly preprocessing)...\")\n    print(\"-\" * 60)\n    train_loader, val_loader = create_dataloaders(\n        train_df=train_df,\n        val_df=val_df,\n        image_root=image_root,\n        batch_size=2,  # Small batch for memory efficiency\n        num_workers=2,\n        use_augmentation=True,\n        target_shape=(96, 320, 320)\n    )\n    \n    print(f\"✓ Train batches: {len(train_loader)}\")\n    print(f\"✓ Val batches: {len(val_loader)}\")\n    \n    # Test dataloader\n    print(\"\\n3. Testing DataLoader...\")\n    print(\"-\" * 60)\n    for volumes, labels, patient_ids in train_loader:\n        print(f\"✓ Batch volume shape: {volumes.shape}\")  # (batch, 1, 96, 320, 320)\n        print(f\"✓ Batch labels shape: {labels.shape}\")   # (batch, 8)\n        print(f\"✓ Patient IDs: {patient_ids}\")\n        print(f\"✓ Volume range: [{volumes.min():.3f}, {volumes.max():.3f}]\")\n        print(f\"✓ Labels (first sample): {labels[0].numpy()}\")\n        break\n    \n    # Calculate class weights for loss function\n    print(\"\\n4. Calculating class weights for loss function...\")\n    print(\"-\" * 60)\n    label_cols = ['patient_overall'] + [f'C{i}' for i in range(1, 8)]\n    \n    pos_weights = []\n    for col in label_cols:\n        pos_rate = train_df[col].mean()\n        # Weight for positive class (higher if rare)\n        weight = (1 - pos_rate) / (pos_rate + 1e-6)\n        pos_weights.append(weight)\n        print(f\"{col}: pos_rate={pos_rate:.3f}, weight={weight:.3f}\")\n    \n    pos_weights_tensor = torch.tensor(pos_weights, dtype=torch.float32)\n    print(f\"\\n✓ Positive class weights: {pos_weights_tensor}\")\n    \n    # Memory estimation\n    print(\"\\n5. Memory Estimation...\")\n    print(\"-\" * 60)\n    batch_size = 2\n    volume_memory = batch_size * 1 * 96 * 320 * 320 * 4 / (1024**3)  # 4 bytes per float32\n    print(f\"✓ Memory per batch (batch_size={batch_size}): {volume_memory:.2f} GB\")\n    print(f\"✓ Recommended GPU memory: {volume_memory * 4:.2f} GB (for model + gradients)\")\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"DATALOADER SETUP COMPLETE!\")\n    print(\"=\"*80)\n    print(\"\"\"\nNext steps:\n1. Build 3D CNN model\n2. Define loss function (use pos_weights for class imbalance)\n3. Set up optimizer and training loop\n4. Start training!\n    \"\"\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T17:06:06.318623Z","iopub.execute_input":"2025-12-30T17:06:06.318988Z","iopub.status.idle":"2025-12-30T17:06:24.494938Z","shell.execute_reply.started":"2025-12-30T17:06:06.318965Z","shell.execute_reply":"2025-12-30T17:06:24.493768Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nVisualization Tools for Preprocessed Spine Fracture Data\nView CT volumes, labels, and explore the dataset interactively\n\"\"\"\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nimport pandas as pd\nimport torch\nfrom matplotlib.patches import Rectangle\nimport seaborn as sns\n\n\n# ============================================================================\n# 1. VISUALIZE SINGLE VOLUME WITH LABELS\n# ============================================================================\n\ndef visualize_volume_with_labels(volume, labels, patient_id, num_slices=12):\n    \"\"\"\n    Visualize a 3D volume with its fracture labels\n    \n    Args:\n        volume: numpy array (D, H, W) or (1, D, H, W) or torch tensor\n        labels: numpy array (8,) [overall, C1-C7]\n        patient_id: Patient ID string\n        num_slices: Number of slices to display\n    \"\"\"\n    # Convert torch tensor to numpy if needed\n    if torch.is_tensor(volume):\n        volume = volume.cpu().numpy()\n    \n    # Remove channel dimension if present\n    if volume.ndim == 4:\n        volume = volume[0]\n    \n    # Extract label information\n    label_names = ['Overall'] + [f'C{i}' for i in range(1, 8)]\n    fracture_status = {label_names[i]: bool(labels[i]) for i in range(len(labels))}\n    \n    # Create color for title (red if fracture, green if no fracture)\n    title_color = 'red' if fracture_status['Overall'] else 'green'\n    \n    # Select evenly spaced slices\n    depth = volume.shape[0]\n    slice_indices = np.linspace(0, depth-1, num_slices, dtype=int)\n    \n    # Create subplot grid\n    rows = 3\n    cols = 4\n    fig, axes = plt.subplots(rows, cols, figsize=(16, 12))\n    axes = axes.flatten()\n    \n    # Plot each slice\n    for idx, slice_idx in enumerate(slice_indices):\n        axes[idx].imshow(volume[slice_idx], cmap='gray', vmin=0, vmax=1)\n        axes[idx].set_title(f'Slice {slice_idx}/{depth}', fontsize=10)\n        axes[idx].axis('off')\n    \n    # Overall title\n    overall_status = \"FRACTURE DETECTED\" if fracture_status['Overall'] else \"NO FRACTURE\"\n    fig.suptitle(f'Patient: {patient_id} - {overall_status}', \n                 fontsize=16, fontweight='bold', color=title_color)\n    \n    # Add label information as text\n    label_text = \"Fracture Labels:\\n\"\n    for name, has_fracture in fracture_status.items():\n        status = \"✓ FRACTURE\" if has_fracture else \"✗ No fracture\"\n        color = \"red\" if has_fracture else \"black\"\n        label_text += f\"{name}: {status}\\n\"\n    \n    # Add text box with labels\n    fig.text(0.02, 0.5, label_text, fontsize=11, verticalalignment='center',\n             bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.5))\n    \n    plt.tight_layout(rect=[0.1, 0, 1, 0.96])\n    plt.show()\n    \n    # Print volume statistics\n    print(f\"\\n{'='*60}\")\n    print(f\"Volume Statistics for Patient {patient_id}\")\n    print(f\"{'='*60}\")\n    print(f\"Shape: {volume.shape}\")\n    print(f\"Value range: [{volume.min():.3f}, {volume.max():.3f}]\")\n    print(f\"Mean: {volume.mean():.3f}\")\n    print(f\"Std: {volume.std():.3f}\")\n    print(f\"Non-zero voxels: {np.count_nonzero(volume):,} ({np.count_nonzero(volume)/volume.size*100:.1f}%)\")\n    print(f\"{'='*60}\\n\")\n\n\n# ============================================================================\n# 2. COMPARE MULTIPLE PATIENTS SIDE BY SIDE\n# ============================================================================\n\ndef compare_patients(volumes_list, labels_list, patient_ids_list, slice_position=0.5):\n    \"\"\"\n    Compare multiple patients side by side\n    \n    Args:\n        volumes_list: List of volumes\n        labels_list: List of label arrays\n        patient_ids_list: List of patient IDs\n        slice_position: Position to slice (0.0 to 1.0, 0.5 = middle)\n    \"\"\"\n    num_patients = len(volumes_list)\n    \n    fig, axes = plt.subplots(2, num_patients, figsize=(5*num_patients, 10))\n    \n    if num_patients == 1:\n        axes = axes.reshape(-1, 1)\n    \n    for i, (volume, labels, patient_id) in enumerate(zip(volumes_list, labels_list, patient_ids_list)):\n        # Convert if needed\n        if torch.is_tensor(volume):\n            volume = volume.cpu().numpy()\n        if volume.ndim == 4:\n            volume = volume[0]\n        \n        # Get slice\n        slice_idx = int(volume.shape[0] * slice_position)\n        \n        # Axial view\n        axes[0, i].imshow(volume[slice_idx], cmap='gray')\n        fracture_status = \"FRACTURE\" if labels[0] == 1 else \"NO FRACTURE\"\n        color = 'red' if labels[0] == 1 else 'green'\n        axes[0, i].set_title(f'{patient_id}\\n{fracture_status}', \n                            fontsize=12, fontweight='bold', color=color)\n        axes[0, i].axis('off')\n        \n        # Sagittal view (middle slice)\n        sagittal = volume[:, volume.shape[1]//2, :]\n        axes[1, i].imshow(sagittal, cmap='gray', aspect='auto')\n        axes[1, i].set_title('Sagittal View', fontsize=10)\n        axes[1, i].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\n\n# ============================================================================\n# 3. EXPLORE SLICES INTERACTIVELY\n# ============================================================================\n\ndef explore_volume_slices(volume, labels, patient_id, view='axial'):\n    \"\"\"\n    Show all slices in a grid for detailed exploration\n    \n    Args:\n        volume: 3D volume\n        labels: Fracture labels\n        patient_id: Patient ID\n        view: 'axial', 'sagittal', or 'coronal'\n    \"\"\"\n    # Convert if needed\n    if torch.is_tensor(volume):\n        volume = volume.cpu().numpy()\n    if volume.ndim == 4:\n        volume = volume[0]\n    \n    # Select view\n    if view == 'axial':\n        slices = volume  # (D, H, W)\n        num_slices = volume.shape[0]\n    elif view == 'sagittal':\n        slices = np.transpose(volume, (2, 0, 1))  # (W, D, H)\n        num_slices = volume.shape[2]\n    elif view == 'coronal':\n        slices = np.transpose(volume, (1, 0, 2))  # (H, D, W)\n        num_slices = volume.shape[1]\n    \n    # Calculate grid size\n    cols = 8\n    rows = (num_slices + cols - 1) // cols\n    \n    fig, axes = plt.subplots(rows, cols, figsize=(20, rows*2.5))\n    axes = axes.flatten()\n    \n    for i in range(len(axes)):\n        if i < num_slices:\n            axes[i].imshow(slices[i], cmap='gray')\n            axes[i].set_title(f'{i}', fontsize=8)\n        axes[i].axis('off')\n    \n    fracture_status = \"FRACTURE\" if labels[0] == 1 else \"NO FRACTURE\"\n    color = 'red' if labels[0] == 1 else 'green'\n    fig.suptitle(f'Patient {patient_id} - {view.upper()} View - {fracture_status}', \n                 fontsize=16, fontweight='bold', color=color)\n    \n    plt.tight_layout()\n    plt.show()\n\n\n# ============================================================================\n# 4. VISUALIZE BATCH FROM DATALOADER\n# ============================================================================\n\ndef visualize_batch(train_loader, num_samples=4):\n    \"\"\"\n    Visualize a batch from the DataLoader\n    \n    Args:\n        train_loader: PyTorch DataLoader\n        num_samples: Number of samples to show\n    \"\"\"\n    # Get one batch\n    volumes, labels, patient_ids = next(iter(train_loader))\n    \n    # Limit to num_samples\n    num_samples = min(num_samples, volumes.shape[0])\n    \n    fig, axes = plt.subplots(2, num_samples, figsize=(4*num_samples, 8))\n    \n    if num_samples == 1:\n        axes = axes.reshape(-1, 1)\n    \n    for i in range(num_samples):\n        volume = volumes[i].cpu().numpy()[0]  # Remove channel dim\n        label = labels[i].cpu().numpy()\n        patient_id = patient_ids[i]\n        \n        # Get middle slice\n        mid_slice = volume.shape[0] // 2\n        \n        # Axial view\n        axes[0, i].imshow(volume[mid_slice], cmap='gray')\n        fracture_status = \"FRACTURE\" if label[0] == 1 else \"NO FRACTURE\"\n        color = 'red' if label[0] == 1 else 'green'\n        axes[0, i].set_title(f'{patient_id}\\n{fracture_status}', \n                            fontsize=10, color=color, fontweight='bold')\n        axes[0, i].axis('off')\n        \n        # MIP (Maximum Intensity Projection)\n        mip = np.max(volume, axis=0)\n        axes[1, i].imshow(mip, cmap='gray')\n        axes[1, i].set_title('MIP', fontsize=10)\n        axes[1, i].axis('off')\n    \n    plt.suptitle('Batch Visualization from DataLoader', fontsize=14, fontweight='bold')\n    plt.tight_layout()\n    plt.show()\n    \n    # Print batch info\n    print(f\"\\n{'='*60}\")\n    print(f\"Batch Information\")\n    print(f\"{'='*60}\")\n    print(f\"Batch shape: {volumes.shape}\")\n    print(f\"Labels shape: {labels.shape}\")\n    print(f\"Number of samples: {len(patient_ids)}\")\n    print(f\"Fractures in batch: {labels[:, 0].sum().item()}/{len(patient_ids)}\")\n    print(f\"{'='*60}\\n\")\n\n\n# ============================================================================\n# 5. VISUALIZE DATASET STATISTICS\n# ============================================================================\n\ndef visualize_dataset_statistics(train_df):\n    \"\"\"\n    Visualize overall dataset statistics\n    \n    Args:\n        train_df: Training DataFrame\n    \"\"\"\n    fig, axes = plt.subplots(2, 3, figsize=(18, 10))\n    \n    # 1. Overall fracture distribution\n    fracture_counts = train_df['patient_overall'].value_counts()\n    axes[0, 0].bar(['No Fracture', 'Fracture'], \n                   [fracture_counts[0], fracture_counts[1]],\n                   color=['green', 'red'])\n    axes[0, 0].set_title('Overall Fracture Distribution', fontsize=12, fontweight='bold')\n    axes[0, 0].set_ylabel('Number of Patients')\n    for i, v in enumerate([fracture_counts[0], fracture_counts[1]]):\n        axes[0, 0].text(i, v, f'{v}\\n({v/len(train_df)*100:.1f}%)', \n                       ha='center', va='bottom', fontweight='bold')\n    \n    # 2. Fracture distribution by vertebra\n    vertebrae_cols = [f'C{i}' for i in range(1, 8)]\n    fracture_by_vertebra = train_df[vertebrae_cols].sum()\n    \n    axes[0, 1].bar(vertebrae_cols, fracture_by_vertebra, color='steelblue')\n    axes[0, 1].set_title('Fractures by Vertebra', fontsize=12, fontweight='bold')\n    axes[0, 1].set_ylabel('Number of Fractures')\n    axes[0, 1].set_xlabel('Vertebra')\n    for i, v in enumerate(fracture_by_vertebra):\n        axes[0, 1].text(i, v, f'{int(v)}', ha='center', va='bottom')\n    \n    # 3. Percentage by vertebra\n    fracture_pct = (train_df[vertebrae_cols].sum() / len(train_df) * 100).sort_values(ascending=False)\n    axes[0, 2].barh(fracture_pct.index, fracture_pct.values, color='coral')\n    axes[0, 2].set_title('Fracture Rate by Vertebra (%)', fontsize=12, fontweight='bold')\n    axes[0, 2].set_xlabel('Percentage of Patients')\n    for i, v in enumerate(fracture_pct.values):\n        axes[0, 2].text(v, i, f'{v:.1f}%', va='center')\n    \n    # 4. Number of fractured vertebrae per patient\n    num_fractures = train_df[vertebrae_cols].sum(axis=1)\n    fracture_dist = num_fractures.value_counts().sort_index()\n    \n    axes[1, 0].bar(fracture_dist.index, fracture_dist.values, color='purple', alpha=0.7)\n    axes[1, 0].set_title('Number of Fractured Vertebrae per Patient', \n                        fontsize=12, fontweight='bold')\n    axes[1, 0].set_xlabel('Number of Fractured Vertebrae')\n    axes[1, 0].set_ylabel('Number of Patients')\n    for i, v in enumerate(fracture_dist.values):\n        axes[1, 0].text(fracture_dist.index[i], v, f'{v}', ha='center', va='bottom')\n    \n    # 5. Correlation heatmap\n    corr_matrix = train_df[['patient_overall'] + vertebrae_cols].corr()\n    sns.heatmap(corr_matrix, annot=True, fmt='.2f', cmap='coolwarm', \n                center=0, ax=axes[1, 1], cbar_kws={'label': 'Correlation'})\n    axes[1, 1].set_title('Correlation Matrix', fontsize=12, fontweight='bold')\n    \n    # 6. Class imbalance visualization\n    label_cols = ['patient_overall'] + vertebrae_cols\n    class_weights = []\n    for col in label_cols:\n        pos_rate = train_df[col].mean()\n        weight = (1 - pos_rate) / (pos_rate + 1e-6)\n        class_weights.append(weight)\n    \n    axes[1, 2].bar(range(len(label_cols)), class_weights, color='orange', alpha=0.7)\n    axes[1, 2].set_xticks(range(len(label_cols)))\n    axes[1, 2].set_xticklabels(label_cols, rotation=45)\n    axes[1, 2].set_title('Class Weights (for Loss Function)', \n                        fontsize=12, fontweight='bold')\n    axes[1, 2].set_ylabel('Weight')\n    axes[1, 2].axhline(y=1, color='red', linestyle='--', alpha=0.5, label='Balanced')\n    axes[1, 2].legend()\n    \n    plt.tight_layout()\n    plt.show()\n    \n    # Print summary statistics\n    print(f\"\\n{'='*60}\")\n    print(f\"DATASET SUMMARY\")\n    print(f\"{'='*60}\")\n    print(f\"Total patients: {len(train_df)}\")\n    print(f\"Patients with fractures: {train_df['patient_overall'].sum()} ({train_df['patient_overall'].mean()*100:.1f}%)\")\n    print(f\"Patients without fractures: {(1-train_df['patient_overall']).sum()} ({(1-train_df['patient_overall']).mean()*100:.1f}%)\")\n    print(f\"\\nMost common fractured vertebra: {fracture_by_vertebra.idxmax()} ({fracture_by_vertebra.max()} cases)\")\n    print(f\"Least common fractured vertebra: {fracture_by_vertebra.idxmin()} ({fracture_by_vertebra.min()} cases)\")\n    print(f\"{'='*60}\\n\")\n\n\n# ============================================================================\n# EXAMPLE USAGE\n# ============================================================================\n\nif __name__ == \"__main__\":\n    \n    print(\"=\"*80)\n    print(\"VISUALIZE PREPROCESSED DATA\")\n    print(\"=\"*80)\n    \n    # Example 1: Visualize from DataLoader\n    print(\"\\n1. Visualizing batch from DataLoader...\")\n    print(\"-\" * 60)\n    \n    # Assuming you have train_loader\n    # visualize_batch(train_loader, num_samples=4)\n    \n    # Example 2: Visualize single patient\n    print(\"\\n2. To visualize a single patient:\")\n    print(\"-\" * 60)\n    print(\"\"\"\n    # Get one sample from dataloader\n    volumes, labels, patient_ids = next(iter(train_loader))\n    \n    # Visualize first patient in batch\n    visualize_volume_with_labels(\n        volume=volumes[0],\n        labels=labels[0],\n        patient_id=patient_ids[0],\n        num_slices=12\n    )\n    \"\"\")\n    \n    # Example 3: Explore all slices\n    print(\"\\n3. To explore all slices of a volume:\")\n    print(\"-\" * 60)\n    print(\"\"\"\n    explore_volume_slices(\n        volume=volumes[0],\n        labels=labels[0],\n        patient_id=patient_ids[0],\n        view='axial'  # or 'sagittal' or 'coronal'\n    )\n    \"\"\")\n    \n    # Example 4: Compare multiple patients\n    print(\"\\n4. To compare multiple patients:\")\n    print(\"-\" * 60)\n    print(\"\"\"\n    # Get a batch\n    volumes, labels, patient_ids = next(iter(train_loader))\n    \n    # Compare first 3 patients\n    compare_patients(\n        volumes_list=[volumes[0], volumes[1], volumes[2]],\n        labels_list=[labels[0], labels[1], labels[2]],\n        patient_ids_list=[patient_ids[0], patient_ids[1], patient_ids[2]],\n        slice_position=0.5\n    )\n    \"\"\")\n    \n    # Example 5: Dataset statistics\n    print(\"\\n5. To visualize dataset statistics:\")\n    print(\"-\" * 60)\n    print(\"\"\"\n    import pandas as pd\n    train_df = pd.read_csv('/kaggle/input/.../train.csv')\n    visualize_dataset_statistics(train_df)\n    \"\"\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T17:06:24.498194Z","iopub.execute_input":"2025-12-30T17:06:24.498561Z","iopub.status.idle":"2025-12-30T17:06:24.666571Z","shell.execute_reply.started":"2025-12-30T17:06:24.498507Z","shell.execute_reply":"2025-12-30T17:06:24.665676Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"    # Get one sample from dataloader\n    volumes, labels, patient_ids = next(iter(train_loader))\n    \n    # Visualize first patient in batch\n    visualize_volume_with_labels(\n        volume=volumes[0],\n        labels=labels[0],\n        patient_id=patient_ids[0],\n        num_slices=12\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T17:06:24.667575Z","iopub.execute_input":"2025-12-30T17:06:24.66794Z","iopub.status.idle":"2025-12-30T17:06:52.207953Z","shell.execute_reply.started":"2025-12-30T17:06:24.667915Z","shell.execute_reply":"2025-12-30T17:06:52.207135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"    explore_volume_slices(\n        volume=volumes[0],\n        labels=labels[0],\n        patient_id=patient_ids[0],\n        view='axial'  # or 'sagittal' or 'coronal'\n    )","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T17:06:52.209194Z","iopub.execute_input":"2025-12-30T17:06:52.210012Z","iopub.status.idle":"2025-12-30T17:06:58.156537Z","shell.execute_reply.started":"2025-12-30T17:06:52.209968Z","shell.execute_reply":"2025-12-30T17:06:58.155694Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"    import pandas as pd\n    train_df = pd.read_csv('/kaggle/input/rsna-2022-cervical-spine-fracture-detection/train.csv')\n    visualize_dataset_statistics(train_df)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T17:06:58.15759Z","iopub.execute_input":"2025-12-30T17:06:58.157837Z","iopub.status.idle":"2025-12-30T17:06:59.581392Z","shell.execute_reply.started":"2025-12-30T17:06:58.157815Z","shell.execute_reply":"2025-12-30T17:06:59.580481Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nEFFICIENTNET TRAINING - Complete Training Pipeline with Advanced Metrics\nIncludes: Precision-Recall Curves, mAP@0.5, Confusion Matrix\n\"\"\"\n\nimport torch\nimport torch.nn as nn\nimport torch.optim as optim\nimport numpy as np\nfrom sklearn.metrics import (roc_auc_score, precision_recall_curve, average_precision_score,\n                              confusion_matrix, classification_report, auc)\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nimport time\nimport os\nimport gc\nimport warnings\nwarnings.filterwarnings('ignore')\n\nprint(\"=\"*80)\nprint(\"🚀 EFFICIENTNET-B0 TRAINING FOR SPINE FRACTURES\")\nprint(\"=\"*80)\n\n# ============================================================================\n# IMPORT EFFICIENTNET MODEL\n# ============================================================================\n\nimport torch.nn.functional as F\n\nclass Conv3dSame(nn.Module):\n    \"\"\"3D convolution with 'SAME' padding\"\"\"\n    def __init__(self, in_channels, out_channels, kernel_size, stride=1, dilation=1, groups=1, bias=True):\n        super().__init__()\n        self.conv = nn.Conv3d(in_channels, out_channels, kernel_size, stride, \n                              padding=0, dilation=dilation, groups=groups, bias=bias)\n        self.stride = stride\n        self.kernel_size = kernel_size if isinstance(kernel_size, tuple) else (kernel_size,) * 3\n        \n    def forward(self, x):\n        d, h, w = x.shape[2:]\n        pad_d = max((self.stride[0] if isinstance(self.stride, tuple) else self.stride) * \n                    ((d - 1) // (self.stride[0] if isinstance(self.stride, tuple) else self.stride)) + \n                    self.kernel_size[0] - d, 0)\n        pad_h = max((self.stride[1] if isinstance(self.stride, tuple) else self.stride) * \n                    ((h - 1) // (self.stride[1] if isinstance(self.stride, tuple) else self.stride)) + \n                    self.kernel_size[1] - h, 0)\n        pad_w = max((self.stride[2] if isinstance(self.stride, tuple) else self.stride) * \n                    ((w - 1) // (self.stride[2] if isinstance(self.stride, tuple) else self.stride)) + \n                    self.kernel_size[2] - w, 0)\n        \n        if pad_d > 0 or pad_h > 0 or pad_w > 0:\n            x = F.pad(x, [pad_w // 2, pad_w - pad_w // 2,\n                         pad_h // 2, pad_h - pad_h // 2,\n                         pad_d // 2, pad_d - pad_d // 2])\n        \n        return self.conv(x)\n\n\nclass SqueezeExcitation3D(nn.Module):\n    \"\"\"3D Squeeze-and-Excitation block\"\"\"\n    def __init__(self, channels, reduction=4):\n        super().__init__()\n        reduced_channels = max(1, channels // reduction)\n        self.se = nn.Sequential(\n            nn.AdaptiveAvgPool3d(1),\n            nn.Conv3d(channels, reduced_channels, 1),\n            nn.SiLU(inplace=True),\n            nn.Conv3d(reduced_channels, channels, 1),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        return x * self.se(x)\n\n\nclass MBConv3D(nn.Module):\n    \"\"\"3D Mobile Inverted Bottleneck Convolution\"\"\"\n    def __init__(self, in_channels, out_channels, kernel_size, stride, expand_ratio, se_ratio=0.25):\n        super().__init__()\n        self.use_residual = (stride == 1 and in_channels == out_channels)\n        hidden_dim = int(in_channels * expand_ratio)\n        \n        layers = []\n        if expand_ratio != 1:\n            layers.extend([\n                nn.Conv3d(in_channels, hidden_dim, 1, bias=False),\n                nn.BatchNorm3d(hidden_dim),\n                nn.SiLU(inplace=True)\n            ])\n        \n        layers.extend([\n            Conv3dSame(hidden_dim, hidden_dim, kernel_size, stride=stride, groups=hidden_dim, bias=False),\n            nn.BatchNorm3d(hidden_dim),\n            nn.SiLU(inplace=True)\n        ])\n        \n        if se_ratio > 0:\n            layers.append(SqueezeExcitation3D(hidden_dim, int(1/se_ratio)))\n        \n        layers.extend([\n            nn.Conv3d(hidden_dim, out_channels, 1, bias=False),\n            nn.BatchNorm3d(out_channels)\n        ])\n        \n        self.block = nn.Sequential(*layers)\n        self.dropout = nn.Dropout(0.2) if self.use_residual else None\n    \n    def forward(self, x):\n        if self.use_residual:\n            return x + self.dropout(self.block(x))\n        else:\n            return self.block(x)\n\n\nclass EfficientNet3D_B0(nn.Module):\n    \"\"\"3D EfficientNet-B0 for spine fracture detection\"\"\"\n    def __init__(self, num_classes=8, dropout=0.3):\n        super().__init__()\n        \n        self.stem = nn.Sequential(\n            Conv3dSame(1, 32, kernel_size=3, stride=2, bias=False),\n            nn.BatchNorm3d(32),\n            nn.SiLU(inplace=True)\n        )\n        \n        blocks_config = [\n            [1, 32, 16, 3, 1, 1],\n            [2, 16, 24, 3, 2, 6],\n            [2, 24, 40, 5, 2, 6],\n            [3, 40, 80, 3, 2, 6],\n            [3, 80, 112, 5, 1, 6],\n            [4, 112, 192, 5, 2, 6],\n            [1, 192, 320, 3, 1, 6],\n        ]\n        \n        self.blocks = nn.ModuleList()\n        for num_layers, in_ch, out_ch, kernel, stride, expand in blocks_config:\n            for i in range(num_layers):\n                self.blocks.append(\n                    MBConv3D(\n                        in_channels=in_ch if i == 0 else out_ch,\n                        out_channels=out_ch,\n                        kernel_size=kernel,\n                        stride=stride if i == 0 else 1,\n                        expand_ratio=expand,\n                        se_ratio=0.25\n                    )\n                )\n        \n        self.head = nn.Sequential(\n            nn.Conv3d(320, 1280, 1, bias=False),\n            nn.BatchNorm3d(1280),\n            nn.SiLU(inplace=True),\n            nn.AdaptiveAvgPool3d(1),\n            nn.Flatten()\n        )\n        \n        self.dropout = nn.Dropout(dropout)\n        self.fc_overall = nn.Linear(1280, 1)\n        self.fc_vertebrae = nn.Linear(1280, 7)\n        \n        self._init_weights()\n    \n    def _init_weights(self):\n        for m in self.modules():\n            if isinstance(m, nn.Conv3d):\n                nn.init.kaiming_normal_(m.weight, mode='fan_out', nonlinearity='relu')\n            elif isinstance(m, nn.BatchNorm3d):\n                nn.init.constant_(m.weight, 1)\n                nn.init.constant_(m.bias, 0)\n            elif isinstance(m, nn.Linear):\n                nn.init.xavier_uniform_(m.weight)\n                nn.init.zeros_(m.bias)\n    \n    def forward(self, x):\n        x = self.stem(x)\n        \n        for block in self.blocks:\n            x = block(x)\n        \n        features = self.head(x)\n        features = self.dropout(features)\n        \n        overall = self.fc_overall(features)\n        vertebrae = self.fc_vertebrae(features)\n        \n        return torch.cat([overall, vertebrae], dim=1)\n\n\n# ============================================================================\n# EVALUATION METRICS FUNCTIONS\n# ============================================================================\n\ndef plot_precision_recall_curves(all_labels, all_preds, class_names, save_path):\n    \"\"\"Plot Precision-Recall curves for all classes\"\"\"\n    fig, axes = plt.subplots(2, 4, figsize=(20, 10))\n    axes = axes.ravel()\n    \n    aps = []\n    \n    for i, (ax, class_name) in enumerate(zip(axes, class_names)):\n        if len(np.unique(all_labels[:, i])) > 1:\n            precision, recall, _ = precision_recall_curve(all_labels[:, i], all_preds[:, i])\n            ap = average_precision_score(all_labels[:, i], all_preds[:, i])\n            pr_auc = auc(recall, precision)\n            aps.append(ap)\n            \n            ax.plot(recall, precision, linewidth=2, label=f'AP={ap:.3f}, AUC={pr_auc:.3f}')\n            ax.fill_between(recall, precision, alpha=0.2)\n            ax.set_xlabel('Recall', fontsize=10)\n            ax.set_ylabel('Precision', fontsize=10)\n            ax.set_title(f'{class_name}', fontsize=12, fontweight='bold')\n            ax.legend(loc='best')\n            ax.grid(True, alpha=0.3)\n            ax.set_xlim([0, 1])\n            ax.set_ylim([0, 1.05])\n        else:\n            ax.text(0.5, 0.5, f'{class_name}\\nNo positive samples', \n                   ha='center', va='center', fontsize=10)\n            ax.set_xlim([0, 1])\n            ax.set_ylim([0, 1])\n            aps.append(0.0)\n    \n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.close()\n    \n    return aps\n\n\ndef calculate_map_at_threshold(all_labels, all_preds, threshold=0.5):\n    \"\"\"Calculate mAP@threshold (mean Average Precision)\"\"\"\n    aps = []\n    \n    for i in range(all_labels.shape[1]):\n        if len(np.unique(all_labels[:, i])) > 1:\n            ap = average_precision_score(all_labels[:, i], all_preds[:, i])\n            aps.append(ap)\n        else:\n            aps.append(0.0)\n    \n    return np.mean(aps), aps\n\n\ndef plot_confusion_matrices(all_labels, all_preds, class_names, save_path, threshold=0.5):\n    \"\"\"Plot confusion matrices for all classes\"\"\"\n    fig, axes = plt.subplots(2, 4, figsize=(20, 10))\n    axes = axes.ravel()\n    \n    pred_binary = (all_preds > threshold).astype(int)\n    \n    for i, (ax, class_name) in enumerate(zip(axes, class_names)):\n        cm = confusion_matrix(all_labels[:, i], pred_binary[:, i])\n        \n        sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', ax=ax,\n                   xticklabels=['Negative', 'Positive'],\n                   yticklabels=['Negative', 'Positive'],\n                   cbar_kws={'label': 'Count'})\n        \n        ax.set_title(f'{class_name}', fontsize=12, fontweight='bold')\n        ax.set_ylabel('True Label', fontsize=10)\n        ax.set_xlabel('Predicted Label', fontsize=10)\n        \n        # Calculate metrics\n        tn, fp, fn, tp = cm.ravel() if cm.size == 4 else (0, 0, 0, 0)\n        accuracy = (tp + tn) / (tp + tn + fp + fn) if (tp + tn + fp + fn) > 0 else 0\n        precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n        recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n        f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0\n        \n        metrics_text = f'Acc: {accuracy:.3f}\\nPrec: {precision:.3f}\\nRec: {recall:.3f}\\nF1: {f1:.3f}'\n        ax.text(1.5, 0.5, metrics_text, fontsize=9, \n               bbox=dict(boxstyle='round', facecolor='wheat', alpha=0.5))\n    \n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    plt.close()\n\n\ndef print_detailed_metrics(all_labels, all_preds, class_names, threshold=0.5):\n    \"\"\"Print detailed classification metrics\"\"\"\n    pred_binary = (all_preds > threshold).astype(int)\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"📊 DETAILED CLASSIFICATION METRICS\")\n    print(\"=\"*80)\n    \n    for i, class_name in enumerate(class_names):\n        print(f\"\\n{class_name}:\")\n        print(\"-\" * 60)\n        \n        if len(np.unique(all_labels[:, i])) > 1:\n            cm = confusion_matrix(all_labels[:, i], pred_binary[:, i])\n            tn, fp, fn, tp = cm.ravel() if cm.size == 4 else (0, 0, 0, 0)\n            \n            accuracy = (tp + tn) / (tp + tn + fp + fn)\n            precision = tp / (tp + fp) if (tp + fp) > 0 else 0\n            recall = tp / (tp + fn) if (tp + fn) > 0 else 0\n            specificity = tn / (tn + fp) if (tn + fp) > 0 else 0\n            f1 = 2 * precision * recall / (precision + recall) if (precision + recall) > 0 else 0\n            \n            print(f\"  Confusion Matrix: TP={tp}, TN={tn}, FP={fp}, FN={fn}\")\n            print(f\"  Accuracy:    {accuracy:.4f}\")\n            print(f\"  Precision:   {precision:.4f}\")\n            print(f\"  Recall:      {recall:.4f}\")\n            print(f\"  Specificity: {specificity:.4f}\")\n            print(f\"  F1-Score:    {f1:.4f}\")\n            \n            # Average Precision\n            ap = average_precision_score(all_labels[:, i], all_preds[:, i])\n            print(f\"  Avg Precision (AP): {ap:.4f}\")\n        else:\n            print(\"  No positive samples in validation set\")\n\n\n# ============================================================================\n# TRAINING CONFIGURATION\n# ============================================================================\n\nCONFIG = {\n    'num_epochs': 6,\n    'batch_size': 2,\n    'train_subset_ratio': 0.50,\n    'val_subset_ratio': 0.40,\n    'max_train_batches': 120,\n    'max_val_batches': 60,\n    'learning_rate': 3e-4,\n    'use_amp': True,\n    'gradient_accumulation': 4,\n    'num_workers': 2,\n    'pin_memory': False,\n    'prefetch_factor': 2,\n    'warmup_epochs': 1,\n    'weight_decay': 0.01,\n    'max_grad_norm': 1.0,\n    'save_dir': '/kaggle/working',\n    'verbose': True,\n}\n\nprint(f\"\\n⚙️  EfficientNet Training Configuration:\")\nprint(f\"  📊 Model: EfficientNet-B0 (Efficient & Accurate)\")\nprint(f\"  📈 Epochs: {CONFIG['num_epochs']}\")\nprint(f\"  🔢 Batch size: {CONFIG['batch_size']} (effective: {CONFIG['batch_size']*CONFIG['gradient_accumulation']})\")\nprint(f\"  📚 Training batches: {CONFIG['max_train_batches']}\")\nprint(f\"  ✅ Validation batches: {CONFIG['max_val_batches']}\")\nprint(f\"  🎓 Learning rate: {CONFIG['learning_rate']}\")\nprint(f\"  📊 Metrics: PR Curves, mAP@0.5, Confusion Matrix\")\n\n# ============================================================================\n# DEVICE SETUP\n# ============================================================================\n\nprint(f\"\\n🧹 Preparing GPU...\")\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n    torch.cuda.synchronize()\n    gc.collect()\n    \ntorch.backends.cudnn.benchmark = True\n\nif torch.cuda.is_available():\n    mem_free = torch.cuda.get_device_properties(0).total_memory - torch.cuda.memory_allocated(0)\n    print(f\"🖥️  GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"  Total: {torch.cuda.get_device_properties(0).total_memory / 1024**3:.1f} GB\")\n    print(f\"  Free: {mem_free / 1024**3:.1f} GB\")\n\n# ============================================================================\n# CREATE DATALOADERS\n# ============================================================================\n\nprint(f\"\\n📦 Creating data loaders...\")\n\ndef create_balanced_loader(original_loader, subset_ratio, batch_size, max_batches, shuffle=True):\n    dataset = original_loader.dataset\n    total_size = len(dataset)\n    subset_size = int(total_size * subset_ratio)\n    \n    indices = np.random.choice(total_size, subset_size, replace=False)\n    subset = torch.utils.data.Subset(dataset, indices)\n    \n    loader = torch.utils.data.DataLoader(\n        subset,\n        batch_size=batch_size,\n        shuffle=shuffle,\n        num_workers=CONFIG['num_workers'],\n        pin_memory=CONFIG['pin_memory'],\n        prefetch_factor=CONFIG['prefetch_factor'] if CONFIG['num_workers'] > 0 else None,\n        persistent_workers=True if CONFIG['num_workers'] > 0 else False,\n        drop_last=True,\n    )\n    \n    return loader\n\ntry:\n    balanced_train_loader = create_balanced_loader(\n        train_loader, CONFIG['train_subset_ratio'],\n        CONFIG['batch_size'], CONFIG['max_train_batches'], shuffle=True\n    )\n    balanced_val_loader = create_balanced_loader(\n        val_loader, CONFIG['val_subset_ratio'],\n        CONFIG['batch_size'], CONFIG['max_val_batches'], shuffle=False\n    )\n    print(f\"  ✓ Using batch_size={CONFIG['batch_size']}\")\n    \nexcept RuntimeError as e:\n    if \"out of memory\" in str(e):\n        print(f\"  ⚠️  OOM with batch_size={CONFIG['batch_size']}, reducing to 1\")\n        CONFIG['batch_size'] = 1\n        torch.cuda.empty_cache()\n        \n        balanced_train_loader = create_balanced_loader(\n            train_loader, CONFIG['train_subset_ratio'],\n            CONFIG['batch_size'], CONFIG['max_train_batches'], shuffle=True\n        )\n        balanced_val_loader = create_balanced_loader(\n            val_loader, CONFIG['val_subset_ratio'],\n            CONFIG['batch_size'], CONFIG['max_val_batches'], shuffle=False\n        )\n\nactual_train_batches = min(len(balanced_train_loader), CONFIG['max_train_batches'])\nactual_val_batches = min(len(balanced_val_loader), CONFIG['max_val_batches'])\n\nprint(f\"  ✓ Train: {actual_train_batches} batches\")\nprint(f\"  ✓ Val: {actual_val_batches} batches\")\n\n# ============================================================================\n# CLASS WEIGHTS\n# ============================================================================\n\nprint(f\"\\n⚖️  Calculating class weights...\")\nlabel_cols = ['patient_overall'] + [f'C{i}' for i in range(1, 8)]\n\nsample_df = train_df.sample(n=min(2000, len(train_df)), random_state=42)\npos_weights = []\nfor col in label_cols:\n    pos_rate = sample_df[col].mean()\n    weight = max(1.0, min(10.0, (1 - pos_rate) / (pos_rate + 1e-6)))\n    pos_weights.append(weight)\n\npos_weights_tensor = torch.tensor(pos_weights, dtype=torch.float32).to(device)\nprint(f\"  ✓ Weights computed: Overall={pos_weights[0]:.2f}, C1-C7={pos_weights[1]:.2f} avg\")\n\n# ============================================================================\n# MODEL & TRAINING SETUP\n# ============================================================================\n\nprint(f\"\\n🏗️  Loading EfficientNet-B0...\")\nmodel = EfficientNet3D_B0(num_classes=8, dropout=0.3)\nmodel = model.to(device)\n\ntorch.cuda.empty_cache()\nnum_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\nprint(f\"  ✓ Model loaded: {num_params:,} parameters ({num_params*4/1024**2:.1f} MB)\")\n\nprint(f\"\\n🎯 Setting up training components...\")\n\ncriterion = nn.BCEWithLogitsLoss(pos_weight=pos_weights_tensor)\n\noptimizer = optim.AdamW(\n    model.parameters(), \n    lr=CONFIG['learning_rate'],\n    weight_decay=CONFIG['weight_decay'],\n    betas=(0.9, 0.999)\n)\n\ntotal_steps = (actual_train_batches // CONFIG['gradient_accumulation']) * CONFIG['num_epochs']\nwarmup_steps = (actual_train_batches // CONFIG['gradient_accumulation']) * CONFIG['warmup_epochs']\n\nscheduler = optim.lr_scheduler.OneCycleLR(\n    optimizer,\n    max_lr=CONFIG['learning_rate'],\n    total_steps=total_steps,\n    pct_start=warmup_steps/total_steps,\n    anneal_strategy='cos',\n    div_factor=25.0,\n    final_div_factor=10000.0\n)\n\nscaler = torch.amp.GradScaler('cuda') if CONFIG['use_amp'] else None\n\nprint(f\"  ✓ Loss: Weighted BCEWithLogitsLoss\")\nprint(f\"  ✓ Optimizer: AdamW\")\nprint(f\"  ✓ Scheduler: OneCycleLR\")\n\n# ============================================================================\n# TRAINING FUNCTIONS\n# ============================================================================\n\ndef train_one_epoch(model, loader, criterion, optimizer, scheduler, device, scaler, \n                    max_batches, grad_accum_steps, epoch_num):\n    model.train()\n    running_loss = 0.0\n    num_batches = 0\n    \n    optimizer.zero_grad()\n    \n    from tqdm import tqdm\n    pbar = tqdm(enumerate(loader), total=min(len(loader), max_batches), \n                desc=f\"Epoch {epoch_num+1} Train\", leave=True)\n    \n    for batch_idx, batch_data in pbar:\n        if batch_idx >= max_batches:\n            break\n        \n        try:\n            if len(batch_data) == 3:\n                volumes, labels, _ = batch_data\n            else:\n                volumes, labels = batch_data[0], batch_data[1]\n            \n            volumes = volumes.to(device, non_blocking=True)\n            labels = labels.to(device, non_blocking=True)\n            \n            if CONFIG['use_amp'] and scaler:\n                with torch.amp.autocast('cuda'):\n                    outputs = model(volumes)\n                    loss = criterion(outputs, labels) / grad_accum_steps\n                \n                scaler.scale(loss).backward()\n                \n                if (batch_idx + 1) % grad_accum_steps == 0:\n                    scaler.unscale_(optimizer)\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), CONFIG['max_grad_norm'])\n                    scaler.step(optimizer)\n                    scaler.update()\n                    optimizer.zero_grad()\n                    scheduler.step()\n            else:\n                outputs = model(volumes)\n                loss = criterion(outputs, labels) / grad_accum_steps\n                loss.backward()\n                \n                if (batch_idx + 1) % grad_accum_steps == 0:\n                    torch.nn.utils.clip_grad_norm_(model.parameters(), CONFIG['max_grad_norm'])\n                    optimizer.step()\n                    optimizer.zero_grad()\n                    scheduler.step()\n            \n            running_loss += loss.item() * grad_accum_steps\n            num_batches += 1\n            \n            pbar.set_postfix({'loss': f'{loss.item()*grad_accum_steps:.4f}'})\n        \n        except Exception as e:\n            continue\n    \n    return running_loss / num_batches if num_batches > 0 else 0\n\n\ndef validate_with_metrics(model, loader, criterion, device, max_batches):\n    \"\"\"Enhanced validation with all predictions and labels for metrics\"\"\"\n    model.eval()\n    running_loss = 0.0\n    num_batches = 0\n    all_preds = []\n    all_labels = []\n    \n    from tqdm import tqdm\n    pbar = tqdm(enumerate(loader), total=min(len(loader), max_batches), \n                desc=\"Validation\", leave=True)\n    \n    with torch.no_grad():\n        for batch_idx, batch_data in pbar:\n            if batch_idx >= max_batches:\n                break\n            \n            try:\n                if len(batch_data) == 3:\n                    volumes, labels, _ = batch_data\n                else:\n                    volumes, labels = batch_data[0], batch_data[1]\n                \n                volumes = volumes.to(device, non_blocking=True)\n                labels = labels.to(device, non_blocking=True)\n                \n                with torch.amp.autocast('cuda'):\n                    outputs = model(volumes)\n                    loss = criterion(outputs, labels)\n                \n                running_loss += loss.item()\n                num_batches += 1\n                \n                preds = torch.sigmoid(outputs).cpu().numpy()\n                all_preds.append(preds)\n                all_labels.append(labels.cpu().numpy())\n                \n                pbar.set_postfix({'loss': f'{loss.item():.4f}'})\n            \n            except Exception as e:\n                continue\n    \n    if num_batches == 0 or len(all_preds) == 0:\n        return 0, 0.5, [0.5]*8, 0.5, None, None\n    \n    all_preds = np.vstack(all_preds)\n    all_labels = np.vstack(all_labels)\n    \n    # Calculate AUCs\n    aucs = []\n    for i in range(8):\n        try:\n            if len(np.unique(all_labels[:, i])) > 1:\n                auc = roc_auc_score(all_labels[:, i], all_preds[:, i])\n                aucs.append(auc)\n            else:\n                aucs.append(0.5)\n        except:\n            aucs.append(0.5)\n    \n    mean_auc = np.mean(aucs)\n    pred_binary = (all_preds > 0.5).astype(int)\n    accuracy = (pred_binary == all_labels).mean()\n    avg_loss = running_loss / num_batches\n    \n    return avg_loss, mean_auc, aucs, accuracy, all_preds, all_labels\n\n# ============================================================================\n# MAIN TRAINING LOOP\n# ============================================================================\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"🚀 STARTING EFFICIENTNET TRAINING\")\nprint(\"=\"*80)\n\nclass_names = ['Overall'] + [f'C{i}' for i in range(1, 8)]\n\nhistory = {\n    'train_loss': [], 'val_loss': [], 'val_auc': [], 'val_acc': [],\n    'learning_rate': [], 'epoch_time': [], 'map_scores': []\n}\nbest_auc = 0.0\nbest_aucs = [0.5] * 8\nbest_epoch = 0\n\ntotal_start = time.time()\n\nfor epoch in range(CONFIG['num_epochs']):\n    epoch_start = time.time()\n    \n    print(f\"\\n{'='*70}\")\n    print(f\"📅 Epoch {epoch+1}/{CONFIG['num_epochs']}\")\n    print(f\"{'='*70}\")\n    \n    train_loss = train_one_epoch(\n        model, balanced_train_loader, criterion, optimizer, scheduler, device,\n        scaler, CONFIG['max_train_batches'], CONFIG['gradient_accumulation'], epoch\n    )\n    \n    torch.cuda.empty_cache()\n    \n    val_loss, mean_auc, aucs, val_acc, all_preds, all_labels = validate_with_metrics(\n        model, balanced_val_loader, criterion, device, CONFIG['max_val_batches']\n    )\n    \n    # Calculate mAP@0.5\n    map_score, ap_scores = calculate_map_at_threshold(all_labels, all_preds, threshold=0.5)\n    \n    current_lr = optimizer.param_groups[0]['lr']\n    epoch_time = time.time() - epoch_start\n    \n    history['train_loss'].append(train_loss)\n    history['val_loss'].append(val_loss)\n    history['val_auc'].append(mean_auc)\n    history['val_acc'].append(val_acc)\n    history['learning_rate'].append(current_lr)\n    history['epoch_time'].append(epoch_time)\n    history['map_scores'].append(map_score)\n    \n    print(f\"\\n📊 Epoch {epoch+1} Results:\")\n    print(f\"   Train Loss:    {train_loss:.4f}\")\n    print(f\"   Val Loss:      {val_loss:.4f}\")\n    print(f\"   Val AUC:       {mean_auc:.4f} {'🎯 NEW BEST!' if mean_auc > best_auc else ''}\")\n    print(f\"   Val Accuracy:  {val_acc:.4f}\")\n    print(f\"   mAP@0.5:       {map_score:.4f}\")\n    print(f\"   Learning Rate: {current_lr:.6f}\")\n    print(f\"   Epoch Time:    {epoch_time:.1f}s\")\n    \n    print(f\"\\n   Individual AUCs:\")\n    for i, (name, auc_val) in enumerate(zip(class_names, aucs)):\n        print(f\"      {name:8s}: {auc_val:.4f}\")\n    \n    print(f\"\\n   Individual APs (Average Precision):\")\n    for i, (name, ap_val) in enumerate(zip(class_names, ap_scores)):\n        print(f\"      {name:8s}: {ap_val:.4f}\")\n    \n    # Save best model\n    if mean_auc > best_auc:\n        best_auc = mean_auc\n        best_aucs = aucs.copy()\n        best_epoch = epoch\n        \n        model_path = os.path.join(CONFIG['save_dir'], 'efficientnet_best_model.pth')\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_auc': mean_auc,\n            'val_loss': val_loss,\n            'aucs': aucs,\n            'config': CONFIG\n        }, model_path)\n        print(f\"\\n   💾 Model saved: {model_path}\")\n    \n    torch.cuda.empty_cache()\n    gc.collect()\n\ntotal_time = time.time() - total_start\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"✅ TRAINING COMPLETED!\")\nprint(\"=\"*80)\nprint(f\"⏱️  Total training time: {total_time/60:.1f} minutes\")\nprint(f\"🏆 Best validation AUC: {best_auc:.4f} (Epoch {best_epoch+1})\")\nprint(f\"\\n   Best Individual AUCs:\")\nfor name, auc_val in zip(class_names, best_aucs):\n    print(f\"      {name:8s}: {auc_val:.4f}\")\n\n# ============================================================================\n# FINAL EVALUATION WITH ALL METRICS\n# ============================================================================\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"📈 GENERATING COMPREHENSIVE EVALUATION METRICS\")\nprint(\"=\"*80)\n\n# Load best model for final evaluation\nprint(\"\\n📂 Loading best model for final evaluation...\")\ncheckpoint = torch.load(os.path.join(CONFIG['save_dir'], 'efficientnet_best_model.pth'))\nmodel.load_state_dict(checkpoint['model_state_dict'])\nprint(\"   ✓ Best model loaded\")\n\n# Get predictions on validation set\nprint(\"\\n🔍 Running final validation pass...\")\nval_loss, mean_auc, aucs, val_acc, all_preds, all_labels = validate_with_metrics(\n    model, balanced_val_loader, criterion, device, CONFIG['max_val_batches']\n)\n\n# Calculate mAP@0.5\nmap_score, ap_scores = calculate_map_at_threshold(all_labels, all_preds, threshold=0.5)\n\nprint(f\"\\n✅ Final Validation Metrics:\")\nprint(f\"   Mean AUC:      {mean_auc:.4f}\")\nprint(f\"   Accuracy:      {val_acc:.4f}\")\nprint(f\"   mAP@0.5:       {map_score:.4f}\")\n\n# Generate Precision-Recall Curves\nprint(\"\\n📊 Generating Precision-Recall curves...\")\npr_curve_path = os.path.join(CONFIG['save_dir'], 'precision_recall_curves.png')\naps = plot_precision_recall_curves(all_labels, all_preds, class_names, pr_curve_path)\nprint(f\"   ✓ Saved: {pr_curve_path}\")\n\n# Generate Confusion Matrices\nprint(\"\\n📊 Generating Confusion matrices...\")\ncm_path = os.path.join(CONFIG['save_dir'], 'confusion_matrices.png')\nplot_confusion_matrices(all_labels, all_preds, class_names, cm_path, threshold=0.5)\nprint(f\"   ✓ Saved: {cm_path}\")\n\n# Print detailed metrics\nprint_detailed_metrics(all_labels, all_preds, class_names, threshold=0.5)\n\n# ============================================================================\n# PLOT TRAINING HISTORY\n# ============================================================================\n\nprint(\"\\n📊 Generating training history plots...\")\n\nfig, axes = plt.subplots(2, 3, figsize=(18, 10))\n\n# Loss plot\naxes[0, 0].plot(history['train_loss'], label='Train Loss', linewidth=2, marker='o')\naxes[0, 0].plot(history['val_loss'], label='Val Loss', linewidth=2, marker='s')\naxes[0, 0].set_xlabel('Epoch', fontsize=10)\naxes[0, 0].set_ylabel('Loss', fontsize=10)\naxes[0, 0].set_title('Training & Validation Loss', fontsize=12, fontweight='bold')\naxes[0, 0].legend()\naxes[0, 0].grid(True, alpha=0.3)\n\n# AUC plot\naxes[0, 1].plot(history['val_auc'], label='Val AUC', linewidth=2, marker='o', color='green')\naxes[0, 1].axhline(y=best_auc, color='r', linestyle='--', label=f'Best: {best_auc:.4f}')\naxes[0, 1].set_xlabel('Epoch', fontsize=10)\naxes[0, 1].set_ylabel('AUC', fontsize=10)\naxes[0, 1].set_title('Validation AUC', fontsize=12, fontweight='bold')\naxes[0, 1].legend()\naxes[0, 1].grid(True, alpha=0.3)\n\n# Accuracy plot\naxes[0, 2].plot(history['val_acc'], label='Val Accuracy', linewidth=2, marker='o', color='purple')\naxes[0, 2].set_xlabel('Epoch', fontsize=10)\naxes[0, 2].set_ylabel('Accuracy', fontsize=10)\naxes[0, 2].set_title('Validation Accuracy', fontsize=12, fontweight='bold')\naxes[0, 2].legend()\naxes[0, 2].grid(True, alpha=0.3)\n\n# Learning rate plot\naxes[1, 0].plot(history['learning_rate'], linewidth=2, marker='o', color='orange')\naxes[1, 0].set_xlabel('Epoch', fontsize=10)\naxes[1, 0].set_ylabel('Learning Rate', fontsize=10)\naxes[1, 0].set_title('Learning Rate Schedule', fontsize=12, fontweight='bold')\naxes[1, 0].set_yscale('log')\naxes[1, 0].grid(True, alpha=0.3)\n\n# Epoch time plot\naxes[1, 1].bar(range(len(history['epoch_time'])), history['epoch_time'], color='teal', alpha=0.7)\naxes[1, 1].set_xlabel('Epoch', fontsize=10)\naxes[1, 1].set_ylabel('Time (seconds)', fontsize=10)\naxes[1, 1].set_title('Epoch Training Time', fontsize=12, fontweight='bold')\naxes[1, 1].grid(True, alpha=0.3, axis='y')\n\n# mAP plot\naxes[1, 2].plot(history['map_scores'], label='mAP@0.5', linewidth=2, marker='o', color='red')\naxes[1, 2].set_xlabel('Epoch', fontsize=10)\naxes[1, 2].set_ylabel('mAP@0.5', fontsize=10)\naxes[1, 2].set_title('Mean Average Precision @0.5', fontsize=12, fontweight='bold')\naxes[1, 2].legend()\naxes[1, 2].grid(True, alpha=0.3)\n\nplt.tight_layout()\nhistory_path = os.path.join(CONFIG['save_dir'], 'training_history.png')\nplt.savefig(history_path, dpi=150, bbox_inches='tight')\nplt.close()\nprint(f\"   ✓ Saved: {history_path}\")\n\n# ============================================================================\n# FINAL SUMMARY\n# ============================================================================\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"🎉 EFFICIENTNET TRAINING COMPLETE - FINAL SUMMARY\")\nprint(\"=\"*80)\n\nprint(f\"\\n📋 Training Configuration:\")\nprint(f\"   Model:              EfficientNet-B0\")\nprint(f\"   Epochs:             {CONFIG['num_epochs']}\")\nprint(f\"   Batch Size:         {CONFIG['batch_size']}\")\nprint(f\"   Gradient Accum:     {CONFIG['gradient_accumulation']}\")\nprint(f\"   Learning Rate:      {CONFIG['learning_rate']}\")\nprint(f\"   Training Batches:   {actual_train_batches}\")\nprint(f\"   Validation Batches: {actual_val_batches}\")\n\nprint(f\"\\n🏆 Best Results (Epoch {best_epoch+1}):\")\nprint(f\"   Validation AUC:     {best_auc:.4f}\")\nprint(f\"   mAP@0.5:            {map_score:.4f}\")\nprint(f\"   Accuracy:           {val_acc:.4f}\")\n\nprint(f\"\\n📊 Saved Outputs:\")\nprint(f\"   ✓ Model checkpoint:        efficientnet_best_model.pth\")\nprint(f\"   ✓ Training history:        training_history.png\")\nprint(f\"   ✓ Precision-Recall curves: precision_recall_curves.png\")\nprint(f\"   ✓ Confusion matrices:      confusion_matrices.png\")\n\nprint(f\"\\n⏱️  Performance:\")\nprint(f\"   Total Training Time: {total_time/60:.1f} minutes\")\nprint(f\"   Avg Time per Epoch:  {np.mean(history['epoch_time']):.1f} seconds\")\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"✨ All metrics generated successfully!\")\nprint(\"=\"*80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T17:06:59.582816Z","iopub.execute_input":"2025-12-30T17:06:59.583095Z","iopub.status.idle":"2025-12-30T19:22:25.905533Z","shell.execute_reply.started":"2025-12-30T17:06:59.583068Z","shell.execute_reply":"2025-12-30T19:22:25.902258Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"import torch\nimport numpy as np\nimport os\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"📈 GENERATING COMPREHENSIVE EVALUATION METRICS\")\nprint(\"=\"*80)\n\n# Load best model for final evaluation (WITH FIX)\nprint(\"\\n📂 Loading best model for final evaluation...\")\ncheckpoint = torch.load(\n    os.path.join(CONFIG['save_dir'], 'efficientnet_best_model.pth'),\n    weights_only=False  # <-- This is the fix!\n)\nmodel.load_state_dict(checkpoint['model_state_dict'])\nprint(\"   ✓ Best model loaded\")\n\n# Get predictions on validation set\nprint(\"\\n🔍 Running final validation pass...\")\nval_loss, mean_auc, aucs, val_acc, all_preds, all_labels = validate_with_metrics(\n    model, balanced_val_loader, criterion, device, CONFIG['max_val_batches']\n)\n\n# Calculate mAP@0.5\nmap_score, ap_scores = calculate_map_at_threshold(all_labels, all_preds, threshold=0.5)\n\nprint(f\"\\n✅ Final Validation Metrics:\")\nprint(f\"   Mean AUC:      {mean_auc:.4f}\")\nprint(f\"   Accuracy:      {val_acc:.4f}\")\nprint(f\"   mAP@0.5:       {map_score:.4f}\")\n\n# Generate Precision-Recall Curves\nprint(\"\\n📊 Generating Precision-Recall curves...\")\npr_curve_path = os.path.join(CONFIG['save_dir'], 'precision_recall_curves.png')\naps = plot_precision_recall_curves(all_labels, all_preds, class_names, pr_curve_path)\nprint(f\"   ✓ Saved: {pr_curve_path}\")\n\n# Generate Confusion Matrices\nprint(\"\\n📊 Generating Confusion matrices...\")\ncm_path = os.path.join(CONFIG['save_dir'], 'confusion_matrices.png')\nplot_confusion_matrices(all_labels, all_preds, class_names, cm_path, threshold=0.5)\nprint(f\"   ✓ Saved: {cm_path}\")\n\n# Print detailed metrics\nprint_detailed_metrics(all_labels, all_preds, class_names, threshold=0.5)\n\n# Plot training history\nprint(\"\\n📊 Generating training history plots...\")\n\nfig, axes = plt.subplots(2, 3, figsize=(18, 10))\n\n# Loss plot\naxes[0, 0].plot(history['train_loss'], label='Train Loss', linewidth=2, marker='o')\naxes[0, 0].plot(history['val_loss'], label='Val Loss', linewidth=2, marker='s')\naxes[0, 0].set_xlabel('Epoch', fontsize=10)\naxes[0, 0].set_ylabel('Loss', fontsize=10)\naxes[0, 0].set_title('Training & Validation Loss', fontsize=12, fontweight='bold')\naxes[0, 0].legend()\naxes[0, 0].grid(True, alpha=0.3)\n\n# AUC plot\naxes[0, 1].plot(history['val_auc'], label='Val AUC', linewidth=2, marker='o', color='green')\naxes[0, 1].axhline(y=best_auc, color='r', linestyle='--', label=f'Best: {best_auc:.4f}')\naxes[0, 1].set_xlabel('Epoch', fontsize=10)\naxes[0, 1].set_ylabel('AUC', fontsize=10)\naxes[0, 1].set_title('Validation AUC', fontsize=12, fontweight='bold')\naxes[0, 1].legend()\naxes[0, 1].grid(True, alpha=0.3)\n\n# Accuracy plot\naxes[0, 2].plot(history['val_acc'], label='Val Accuracy', linewidth=2, marker='o', color='purple')\naxes[0, 2].set_xlabel('Epoch', fontsize=10)\naxes[0, 2].set_ylabel('Accuracy', fontsize=10)\naxes[0, 2].set_title('Validation Accuracy', fontsize=12, fontweight='bold')\naxes[0, 2].legend()\naxes[0, 2].grid(True, alpha=0.3)\n\n# Learning rate plot\naxes[1, 0].plot(history['learning_rate'], linewidth=2, marker='o', color='orange')\naxes[1, 0].set_xlabel('Epoch', fontsize=10)\naxes[1, 0].set_ylabel('Learning Rate', fontsize=10)\naxes[1, 0].set_title('Learning Rate Schedule', fontsize=12, fontweight='bold')\naxes[1, 0].set_yscale('log')\naxes[1, 0].grid(True, alpha=0.3)\n\n# Epoch time plot\naxes[1, 1].bar(range(len(history['epoch_time'])), history['epoch_time'], color='teal', alpha=0.7)\naxes[1, 1].set_xlabel('Epoch', fontsize=10)\naxes[1, 1].set_ylabel('Time (seconds)', fontsize=10)\naxes[1, 1].set_title('Epoch Training Time', fontsize=12, fontweight='bold')\naxes[1, 1].grid(True, alpha=0.3, axis='y')\n\n# mAP plot\naxes[1, 2].plot(history['map_scores'], label='mAP@0.5', linewidth=2, marker='o', color='red')\naxes[1, 2].set_xlabel('Epoch', fontsize=10)\naxes[1, 2].set_ylabel('mAP@0.5', fontsize=10)\naxes[1, 2].set_title('Mean Average Precision @0.5', fontsize=12, fontweight='bold')\naxes[1, 2].legend()\naxes[1, 2].grid(True, alpha=0.3)\n\nplt.tight_layout()\nhistory_path = os.path.join(CONFIG['save_dir'], 'training_history.png')\nplt.savefig(history_path, dpi=150, bbox_inches='tight')\nplt.close()\nprint(f\"   ✓ Saved: {history_path}\")\n\n# Final summary\nprint(\"\\n\" + \"=\"*80)\nprint(\"🎉 EFFICIENTNET TRAINING COMPLETE - FINAL SUMMARY\")\nprint(\"=\"*80)\n\nprint(f\"\\n📋 Training Configuration:\")\nprint(f\"   Model:              EfficientNet-B0\")\nprint(f\"   Epochs:             {CONFIG['num_epochs']}\")\nprint(f\"   Best Epoch:         {best_epoch+1}\")\n\nprint(f\"\\n🏆 Best Results:\")\nprint(f\"   Validation AUC:     {best_auc:.4f}\")\nprint(f\"   mAP@0.5:            {map_score:.4f}\")\nprint(f\"   Accuracy:           {val_acc:.4f}\")\n\nprint(f\"\\n📊 Saved Outputs:\")\nprint(f\"   ✓ Model checkpoint:        efficientnet_best_model.pth\")\nprint(f\"   ✓ Training history:        training_history.png\")\nprint(f\"   ✓ Precision-Recall curves: precision_recall_curves.png\")\nprint(f\"   ✓ Confusion matrices:      confusion_matrices.png\")\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"✨ All metrics generated successfully!\")\nprint(\"=\"*80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T19:32:22.360219Z","iopub.execute_input":"2025-12-30T19:32:22.360889Z","iopub.status.idle":"2025-12-30T19:38:52.397145Z","shell.execute_reply.started":"2025-12-30T19:32:22.360845Z","shell.execute_reply":"2025-12-30T19:38:52.396253Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nCOMPLETE GRAD-CAM VISUALIZATION - SINGLE SCRIPT\nEverything needed: imports, model, data loading, and visualization\nJust run this entire cell after your training!\n\"\"\"\n\nimport os\nimport gc\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom matplotlib.colors import LinearSegmentedColormap\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torchvision.models.video import r3d_18, R3D_18_Weights\n\nprint(\"=\"*80)\nprint(\"🔥 COMPLETE GRAD-CAM VISUALIZATION SCRIPT\")\nprint(\"=\"*80)\n\n# ============================================================================\n# 1. MODEL DEFINITION\n# ============================================================================\n\nprint(\"\\n📦 Step 1: Defining Model Architecture...\")\n\nclass SpineFractureResNet3D(nn.Module):\n    \"\"\"3D ResNet18 for fracture detection\"\"\"\n    \n    def __init__(self, num_classes=8, pretrained=False, dropout=0.3):\n        super(SpineFractureResNet3D, self).__init__()\n        \n        if pretrained:\n            self.backbone = r3d_18(weights=R3D_18_Weights.DEFAULT)\n        else:\n            self.backbone = r3d_18(weights=None)\n        \n        self.backbone.stem[0] = nn.Conv3d(\n            1, 64, kernel_size=(3, 7, 7),\n            stride=(1, 2, 2), padding=(1, 3, 3), bias=False\n        )\n        \n        in_features = self.backbone.fc.in_features\n        self.backbone.fc = nn.Identity()\n        self.dropout = nn.Dropout(dropout)\n        self.fc_overall = nn.Linear(in_features, 1)\n        self.fc_vertebrae = nn.Linear(in_features, 7)\n        \n        nn.init.xavier_uniform_(self.fc_overall.weight)\n        nn.init.zeros_(self.fc_overall.bias)\n        nn.init.xavier_uniform_(self.fc_vertebrae.weight)\n        nn.init.zeros_(self.fc_vertebrae.bias)\n    \n    def forward(self, x):\n        features = self.backbone(x)\n        features = self.dropout(features)\n        overall = self.fc_overall(features)\n        vertebrae = self.fc_vertebrae(features)\n        output = torch.cat([overall, vertebrae], dim=1)\n        return output\n\nprint(\"  ✓ Model architecture defined\")\n\n# ============================================================================\n# 2. GRAD-CAM IMPLEMENTATION\n# ============================================================================\n\nprint(\"\\n🔍 Step 2: Setting up Grad-CAM...\")\n\nclass GradCAM3D:\n    \"\"\"3D Grad-CAM for fracture localization\"\"\"\n    \n    def __init__(self, model, target_layer):\n        self.model = model\n        self.target_layer = target_layer\n        self.gradients = None\n        self.activations = None\n        \n        self.forward_handle = target_layer.register_forward_hook(self._forward_hook)\n        self.backward_handle = target_layer.register_full_backward_hook(self._backward_hook)\n    \n    def _forward_hook(self, module, input, output):\n        self.activations = output.detach()\n    \n    def _backward_hook(self, module, grad_input, grad_output):\n        self.gradients = grad_output[0].detach()\n    \n    def generate_cam(self, input_volume, target_class=0):\n        self.model.eval()\n        output = self.model(input_volume)\n        \n        self.model.zero_grad()\n        output[0, target_class].backward()\n        \n        gradients = self.gradients[0]\n        activations = self.activations[0]\n        \n        weights = gradients.mean(dim=(1, 2, 3), keepdim=True)\n        cam = (weights * activations).sum(dim=0)\n        \n        cam = F.relu(cam)\n        cam = cam - cam.min()\n        if cam.max() > 0:\n            cam = cam / cam.max()\n        \n        return cam.cpu().numpy()\n    \n    def remove_hooks(self):\n        self.forward_handle.remove()\n        self.backward_handle.remove()\n\nprint(\"  ✓ Grad-CAM class ready\")\n\n# ============================================================================\n# 3. VISUALIZATION FUNCTIONS\n# ============================================================================\n\nprint(\"\\n🎨 Step 3: Setting up visualization functions...\")\n\ndef visualize_gradcam_comprehensive(volume, cam, predictions, labels, patient_id, save_path=None):\n    \"\"\"\n    Comprehensive Grad-CAM visualization with 12 slices\n    \"\"\"\n    # Select 12 evenly spaced slices\n    depth = volume.shape[0]\n    slice_indices = np.linspace(0, depth-1, 12, dtype=int)\n    \n    # Resize CAM if needed\n    if cam.shape != volume.shape:\n        from scipy.ndimage import zoom\n        zoom_factors = np.array(volume.shape) / np.array(cam.shape)\n        cam_resized = zoom(cam, zoom_factors, order=1)\n    else:\n        cam_resized = cam\n    \n    # Create figure with 4x3 grid\n    fig, axes = plt.subplots(3, 4, figsize=(20, 15))\n    axes = axes.flatten()\n    \n    # Custom colormap (blue to red for heatmap)\n    colors = ['darkblue', 'blue', 'cyan', 'yellow', 'orange', 'red', 'darkred']\n    cmap = LinearSegmentedColormap.from_list('fracture_heatmap', colors, N=256)\n    \n    for idx, slice_idx in enumerate(slice_indices):\n        ax = axes[idx]\n        \n        # Show CT slice in grayscale\n        ax.imshow(volume[slice_idx], cmap='gray', vmin=0, vmax=1)\n        \n        # Overlay CAM heatmap (only significant regions)\n        cam_slice = cam_resized[slice_idx]\n        masked_cam = np.ma.masked_where(cam_slice < 0.3, cam_slice)\n        im = ax.imshow(masked_cam, cmap=cmap, alpha=0.7, vmin=0, vmax=1)\n        \n        ax.set_title(f'Slice {slice_idx}/{depth}', fontsize=11, fontweight='bold')\n        ax.axis('off')\n    \n    # Add colorbar\n    cbar = plt.colorbar(im, ax=axes, orientation='horizontal', \n                        pad=0.02, fraction=0.046, aspect=40)\n    cbar.set_label('Fracture Attention (Model Focus)', fontsize=12, fontweight='bold')\n    \n    # Overall title with prediction\n    fracture_status = \"FRACTURE DETECTED\" if predictions[0] > 0.5 else \"NO FRACTURE\"\n    confidence = predictions[0] * 100\n    color = 'red' if predictions[0] > 0.5 else 'green'\n    \n    fig.suptitle(\n        f'Grad-CAM Fracture Localization: {patient_id}\\n' +\n        f'{fracture_status} (Model Confidence: {confidence:.1f}%)',\n        fontsize=18, fontweight='bold', color=color, y=0.98\n    )\n    \n    # Add detailed predictions panel\n    label_names = ['Overall'] + [f'C{i}' for i in range(1, 8)]\n    info_text = \"MODEL PREDICTIONS:\\n\" + \"=\"*35 + \"\\n\"\n    \n    for i, name in enumerate(label_names):\n        pred_prob = predictions[i]\n        gt = int(labels[i])\n        pred = int(pred_prob > 0.5)\n        match = '✓ CORRECT' if pred == gt else '✗ WRONG'\n        \n        status = \"FRACTURE\" if pred == 1 else \"Normal\"\n        info_text += f\"{name:8s}: {status:10s} ({pred_prob*100:5.1f}%)\"\n        \n        if gt is not None:\n            info_text += f\" | GT:{gt} {match}\"\n        \n        info_text += \"\\n\"\n    \n    fig.text(0.02, 0.5, info_text, fontsize=10, verticalalignment='center',\n             family='monospace',\n             bbox=dict(boxstyle='round', facecolor='lightyellow', \n                      alpha=0.9, edgecolor='black', linewidth=2))\n    \n    plt.tight_layout(rect=[0.12, 0, 1, 0.95])\n    \n    if save_path:\n        plt.savefig(save_path, dpi=150, bbox_inches='tight')\n        print(f\"  ✓ Saved visualization: {save_path}\")\n    \n    plt.show()\n    \n    return fig\n\n# ============================================================================\n# 4. LOAD MODEL AND DATA\n# ============================================================================\n\nprint(\"\\n🔧 Step 4: Loading model and preparing data...\")\n\n# Setup device\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"  Device: {device}\")\n\nif torch.cuda.is_available():\n    torch.cuda.empty_cache()\n    gc.collect()\n    print(f\"  ✓ GPU memory cleared\")\n\n# Load trained model\ncheckpoint_path = '/kaggle/working/balanced_demo_best.pth'\n\nif os.path.exists(checkpoint_path):\n    print(f\"  Loading trained model...\")\n    model = SpineFractureResNet3D(num_classes=8, pretrained=False, dropout=0.3)\n    checkpoint = torch.load(checkpoint_path, map_location=device, weights_only=False)\n    model.load_state_dict(checkpoint['model_state_dict'])\n    print(f\"  ✓ Model loaded (Epoch {checkpoint['epoch']}, AUC: {checkpoint['best_auc']:.4f})\")\nelse:\n    print(f\"  ⚠️  No checkpoint found, creating untrained model\")\n    model = SpineFractureResNet3D(num_classes=8, pretrained=False, dropout=0.3)\n\nmodel = model.to(device)\nmodel.eval()\n\nnum_params = sum(p.numel() for p in model.parameters())\nprint(f\"  ✓ Model ready ({num_params:,} parameters)\")\n\n# ============================================================================\n# 5. RUN GRAD-CAM ON VALIDATION DATA\n# ============================================================================\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"🚀 RUNNING GRAD-CAM ANALYSIS\")\nprint(\"=\"*80)\n\n# Check if validation loader exists\ntry:\n    if 'balanced_val_loader' in dir():\n        val_loader = balanced_val_loader\n        print(\"  Using: balanced_val_loader\")\n    elif 'val_loader' in dir():\n        val_loader = val_loader\n        print(\"  Using: val_loader\")\n    else:\n        raise NameError(\"No validation loader found\")\n    \n    print(f\"  ✓ Validation loader ready ({len(val_loader)} batches)\")\n    \nexcept NameError:\n    print(\"  ❌ No validation loader found!\")\n    print(\"\\n  You need to run the dataloader creation code first.\")\n    print(\"  Skipping Grad-CAM visualization...\")\n    \n    # Exit gracefully\n    print(\"\\n\" + \"=\"*80)\n    print(\"⚠️  GRAD-CAM SKIPPED - Create validation loader first\")\n    print(\"=\"*80)\n    raise SystemExit\n\n# Find interesting cases (fracture + no fracture)\nprint(\"\\n🔍 Finding interesting cases...\")\n\ncases_found = {'fracture': None, 'no_fracture': None}\nnum_checked = 0\n\nfor batch_data in val_loader:\n    if len(batch_data) == 3:\n        volumes, labels, patient_ids = batch_data\n    else:\n        volumes, labels = batch_data[0], batch_data[1]\n        patient_ids = [f\"Patient_{i}\" for i in range(len(volumes))]\n    \n    for i in range(len(volumes)):\n        has_fracture = labels[i][0].item() == 1\n        \n        if has_fracture and cases_found['fracture'] is None:\n            cases_found['fracture'] = (volumes[i:i+1], labels[i], patient_ids[i])\n            print(f\"  ✓ Found fracture case: {patient_ids[i]}\")\n        \n        if not has_fracture and cases_found['no_fracture'] is None:\n            cases_found['no_fracture'] = (volumes[i:i+1], labels[i], patient_ids[i])\n            print(f\"  ✓ Found no-fracture case: {patient_ids[i]}\")\n        \n        if cases_found['fracture'] and cases_found['no_fracture']:\n            break\n    \n    num_checked += 1\n    if cases_found['fracture'] and cases_found['no_fracture']:\n        break\n    if num_checked >= 10:  # Check max 10 batches\n        break\n\n# Process each case\nresults = []\n\nfor case_name, case_data in cases_found.items():\n    if case_data is None:\n        print(f\"\\n  ⚠️  No {case_name} case found\")\n        continue\n    \n    volume_tensor, label, patient_id = case_data\n    \n    print(f\"\\n{'='*60}\")\n    print(f\"📊 Analyzing: {patient_id} ({case_name.replace('_', ' ')})\")\n    print(f\"{'='*60}\")\n    \n    # Move to device\n    volume_tensor = volume_tensor.to(device)\n    volume_tensor.requires_grad = True\n    \n    # Get predictions\n    with torch.no_grad():\n        output = model(volume_tensor)\n        predictions = torch.sigmoid(output).cpu().numpy()[0]\n    \n    print(f\"  Model Prediction: {predictions[0]*100:.1f}% fracture probability\")\n    print(f\"  Ground Truth: {'FRACTURE' if label[0]==1 else 'NO FRACTURE'}\")\n    \n    # Generate Grad-CAM\n    print(f\"  Generating Grad-CAM heatmap...\")\n    gradcam = GradCAM3D(model, model.backbone.layer4[-1])\n    \n    # Enable gradients for backward pass\n    volume_tensor.requires_grad = True\n    cam = gradcam.generate_cam(volume_tensor, target_class=0)\n    \n    gradcam.remove_hooks()\n    \n    print(f\"  ✓ Grad-CAM complete (shape: {cam.shape})\")\n    \n    # Get volume for visualization\n    volume_np = volume_tensor[0, 0].detach().cpu().numpy()\n    \n    # Create comprehensive visualization\n    save_path = f'/kaggle/working/gradcam_{case_name}_{patient_id}.png'\n    \n    print(f\"  Creating visualization...\")\n    fig = visualize_gradcam_comprehensive(\n        volume_np, cam, predictions, label.numpy(),\n        patient_id, save_path=save_path\n    )\n    \n    results.append({\n        'case': case_name,\n        'patient_id': patient_id,\n        'prediction': predictions[0],\n        'ground_truth': label[0].item(),\n        'save_path': save_path\n    })\n    \n    print(f\"  ✓ Visualization complete!\")\n\n# ============================================================================\n# 6. SUMMARY\n# ============================================================================\n\nprint(\"\\n\" + \"=\"*80)\nprint(\"🎉 GRAD-CAM ANALYSIS COMPLETE!\")\nprint(\"=\"*80)\n\nif results:\n    print(f\"\\n📊 Summary:\")\n    for r in results:\n        pred_label = \"FRACTURE\" if r['prediction'] > 0.5 else \"NO FRACTURE\"\n        gt_label = \"FRACTURE\" if r['ground_truth'] == 1 else \"NO FRACTURE\"\n        match = \"✓\" if (r['prediction'] > 0.5) == (r['ground_truth'] == 1) else \"✗\"\n        \n        print(f\"\\n  {r['case'].upper()}:\")\n        print(f\"    Patient: {r['patient_id']}\")\n        print(f\"    Prediction: {pred_label} ({r['prediction']*100:.1f}%)\")\n        print(f\"    Ground Truth: {gt_label}\")\n        print(f\"    Match: {match}\")\n        print(f\"    Saved: {r['save_path']}\")\n\nprint(f\"\\n💡 What the heatmap shows:\")\nprint(f\"   • RED areas = High attention (model suspects fracture)\")\nprint(f\"   • BLUE areas = Low attention (model thinks normal)\")\nprint(f\"   • Intensity = Confidence level\")\n\nprint(f\"\\n🎯 Perfect for presentation!\")\nprint(f\"   These visualizations show WHERE your model detects fractures!\")\n\nprint(\"\\n\" + \"=\"*80)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T20:04:31.309268Z","iopub.execute_input":"2025-12-30T20:04:31.309816Z","iopub.status.idle":"2025-12-30T20:07:36.534605Z","shell.execute_reply.started":"2025-12-30T20:04:31.309787Z","shell.execute_reply":"2025-12-30T20:07:36.533792Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nINTERACTIVE 3D SPINE VISUALIZATION\nRotating 3D reconstruction with fracture heatmap overlay\nWorks with your current demo-trained model!\n\"\"\"\n\nimport numpy as np\nimport torch\nimport torch.nn.functional as F\nimport plotly.graph_objects as go\nfrom plotly.subplots import make_subplots\nfrom skimage import measure\nfrom scipy.ndimage import zoom\nimport warnings\nwarnings.filterwarnings('ignore')\n\nprint(\"=\"*80)\nprint(\"🌐 INTERACTIVE 3D SPINE VISUALIZATION\")\nprint(\"=\"*80)\n\n# ============================================================================\n# GRAD-CAM CLASS (embedded)\n# ============================================================================\n\nclass GradCAM3D:\n    \"\"\"3D Grad-CAM for fracture localization\"\"\"\n    \n    def __init__(self, model, target_layer):\n        self.model = model\n        self.target_layer = target_layer\n        self.gradients = None\n        self.activations = None\n        \n        self.forward_handle = target_layer.register_forward_hook(self._forward_hook)\n        self.backward_handle = target_layer.register_full_backward_hook(self._backward_hook)\n    \n    def _forward_hook(self, module, input, output):\n        self.activations = output.detach()\n    \n    def _backward_hook(self, module, grad_input, grad_output):\n        self.gradients = grad_output[0].detach()\n    \n    def generate_cam(self, input_volume, target_class=0):\n        self.model.eval()\n        output = self.model(input_volume)\n        \n        self.model.zero_grad()\n        output[0, target_class].backward()\n        \n        gradients = self.gradients[0]\n        activations = self.activations[0]\n        \n        weights = gradients.mean(dim=(1, 2, 3), keepdim=True)\n        cam = (weights * activations).sum(dim=0)\n        \n        cam = F.relu(cam)\n        cam = cam - cam.min()\n        if cam.max() > 0:\n            cam = cam / cam.max()\n        \n        return cam.cpu().numpy()\n    \n    def remove_hooks(self):\n        self.forward_handle.remove()\n        self.backward_handle.remove()\n\nprint(\"  ✓ GradCAM3D class loaded\")\n\n# ============================================================================\n# 3D RECONSTRUCTION FUNCTIONS\n# ============================================================================\n\ndef create_3d_spine_mesh(volume, threshold=0.4, downsample=0.5):\n    \"\"\"\n    Create 3D mesh from CT volume using marching cubes\n    \n    Args:\n        volume: CT volume (D, H, W)\n        threshold: Threshold for bone segmentation\n        downsample: Factor to reduce size (0.5 = half size)\n    \n    Returns:\n        verts, faces: Mesh vertices and faces\n    \"\"\"\n    print(f\"  Creating 3D mesh from volume...\")\n    \n    # Downsample for performance\n    if downsample < 1.0:\n        volume_small = zoom(volume, downsample, order=1)\n    else:\n        volume_small = volume\n    \n    print(f\"    Volume shape: {volume.shape} → {volume_small.shape}\")\n    \n    # Create binary mask for bone\n    bone_mask = volume_small > threshold\n    \n    # Apply marching cubes to get mesh\n    try:\n        verts, faces, normals, values = measure.marching_cubes(\n            bone_mask,\n            level=0,\n            spacing=(1.0, 1.0, 1.0),\n            allow_degenerate=False\n        )\n        print(f\"    ✓ Mesh created: {len(verts)} vertices, {len(faces)} faces\")\n        return verts, faces\n    except Exception as e:\n        print(f\"    ✗ Marching cubes failed: {e}\")\n        return None, None\n\n\ndef map_gradcam_to_mesh(verts, cam_volume, volume_shape):\n    \"\"\"\n    Map Grad-CAM values to mesh vertices\n    \n    Args:\n        verts: Mesh vertices (N, 3)\n        cam_volume: Grad-CAM heatmap (D, H, W)\n        volume_shape: Original volume shape\n    \n    Returns:\n        colors: Color values for each vertex\n    \"\"\"\n    print(f\"  Mapping Grad-CAM to mesh vertices...\")\n    \n    # Resize CAM to match mesh scale\n    if cam_volume.shape != volume_shape:\n        zoom_factors = np.array(volume_shape) / np.array(cam_volume.shape)\n        cam_resized = zoom(cam_volume, zoom_factors, order=1)\n    else:\n        cam_resized = cam_volume\n    \n    # Sample CAM values at vertex positions\n    colors = []\n    for vert in verts:\n        z, y, x = vert\n        \n        # Convert to array indices\n        zi = int(np.clip(z, 0, cam_resized.shape[0] - 1))\n        yi = int(np.clip(y, 0, cam_resized.shape[1] - 1))\n        xi = int(np.clip(x, 0, cam_resized.shape[2] - 1))\n        \n        cam_value = cam_resized[zi, yi, xi]\n        colors.append(cam_value)\n    \n    colors = np.array(colors)\n    print(f\"    ✓ Mapped {len(colors)} vertex colors\")\n    print(f\"    Color range: [{colors.min():.3f}, {colors.max():.3f}]\")\n    \n    return colors\n\n\ndef create_interactive_3d_visualization(volume, cam, predictions, labels, patient_id,\n                                       threshold=0.4, downsample=0.5):\n    \"\"\"\n    Create interactive 3D visualization with Plotly\n    \n    Args:\n        volume: CT volume (D, H, W)\n        cam: Grad-CAM heatmap (D, H, W)\n        predictions: Model predictions (8,)\n        labels: Ground truth labels (8,)\n        patient_id: Patient ID\n        threshold: Bone segmentation threshold\n        downsample: Downsampling factor for performance\n    \"\"\"\n    \n    print(f\"\\n{'='*60}\")\n    print(f\"🎨 Creating 3D visualization for {patient_id}\")\n    print(f\"{'='*60}\")\n    \n    # Create mesh\n    verts, faces = create_3d_spine_mesh(volume, threshold, downsample)\n    \n    if verts is None or faces is None:\n        print(\"  ✗ Could not create mesh\")\n        return None\n    \n    # Map Grad-CAM to vertices\n    colors = map_gradcam_to_mesh(verts, cam, volume.shape)\n    \n    # Determine fracture status\n    has_fracture = predictions[0] > 0.5\n    confidence = predictions[0] * 100\n    \n    fracture_text = \"FRACTURE DETECTED\" if has_fracture else \"NO FRACTURE\"\n    title_color = 'red' if has_fracture else 'green'\n    \n    print(f\"\\n  Prediction: {fracture_text} ({confidence:.1f}%)\")\n    \n    # Create Plotly figure\n    print(f\"  Creating interactive plot...\")\n    \n    fig = go.Figure(data=[\n        go.Mesh3d(\n            x=verts[:, 0],\n            y=verts[:, 1],\n            z=verts[:, 2],\n            i=faces[:, 0],\n            j=faces[:, 1],\n            k=faces[:, 2],\n            intensity=colors,\n            colorscale=[\n                [0.0, 'rgb(0, 0, 100)'],      # Dark blue (low attention)\n                [0.3, 'rgb(0, 100, 200)'],    # Blue\n                [0.5, 'rgb(0, 200, 200)'],    # Cyan\n                [0.7, 'rgb(255, 255, 0)'],    # Yellow\n                [0.85, 'rgb(255, 150, 0)'],   # Orange\n                [1.0, 'rgb(255, 0, 0)']       # Red (high attention - fracture)\n            ],\n            cmin=0,\n            cmax=1,\n            colorbar=dict(\n                title=dict(\n                    text=\"Fracture<br>Attention\",\n                    font=dict(size=14, color='white')\n                ),\n                titleside=\"right\",\n                tickmode=\"linear\",\n                tick0=0,\n                dtick=0.2,\n                tickfont=dict(size=12, color='white'),\n                len=0.7,\n                thickness=20,\n                x=1.0\n            ),\n            opacity=0.95,\n            flatshading=False,\n            lighting=dict(\n                ambient=0.5,\n                diffuse=0.8,\n                specular=0.3,\n                roughness=0.4,\n                fresnel=0.2\n            ),\n            lightposition=dict(\n                x=100,\n                y=100,\n                z=1000\n            ),\n            hovertemplate='<b>Position</b><br>' +\n                         'X: %{x:.1f}<br>' +\n                         'Y: %{y:.1f}<br>' +\n                         'Z: %{z:.1f}<br>' +\n                         '<b>Attention: %{intensity:.3f}</b><br>' +\n                         '<extra></extra>'\n        )\n    ])\n    \n    # Add annotations with predictions\n    label_names = ['Overall'] + [f'C{i}' for i in range(1, 8)]\n    annotation_text = \"<b>PREDICTIONS:</b><br>\"\n    \n    for i, name in enumerate(label_names):\n        pred_prob = predictions[i]\n        pred_status = \"FRACTURE\" if pred_prob > 0.5 else \"Normal\"\n        gt = int(labels[i]) if labels is not None else None\n        \n        annotation_text += f\"{name}: {pred_status} ({pred_prob*100:.1f}%)\"\n        \n        if gt is not None:\n            match = '✓' if (pred_prob > 0.5) == (gt == 1) else '✗'\n            annotation_text += f\" {match}\"\n        \n        annotation_text += \"<br>\"\n    \n    # Update layout with dark theme\n    fig.update_layout(\n        title=dict(\n            text=f'<b>3D Cervical Spine Reconstruction</b><br>' +\n                 f'Patient: {patient_id}<br>' +\n                 f'<span style=\"color:{title_color};\">{fracture_text}</span> ' +\n                 f'(Confidence: {confidence:.1f}%)',\n            font=dict(size=18, color='white'),\n            x=0.5,\n            xanchor='center'\n        ),\n        scene=dict(\n            xaxis=dict(\n                title='Superior ← → Inferior',\n                titlefont=dict(size=12, color='white'),\n                gridcolor='rgb(50, 50, 50)',\n                showbackground=True,\n                backgroundcolor='rgb(20, 20, 20)',\n                tickfont=dict(color='white')\n            ),\n            yaxis=dict(\n                title='Anterior ← → Posterior',\n                titlefont=dict(size=12, color='white'),\n                gridcolor='rgb(50, 50, 50)',\n                showbackground=True,\n                backgroundcolor='rgb(20, 20, 20)',\n                tickfont=dict(color='white')\n            ),\n            zaxis=dict(\n                title='Left ← → Right',\n                titlefont=dict(size=12, color='white'),\n                gridcolor='rgb(50, 50, 50)',\n                showbackground=True,\n                backgroundcolor='rgb(20, 20, 20)',\n                tickfont=dict(color='white')\n            ),\n            aspectmode='data',\n            camera=dict(\n                eye=dict(x=1.8, y=1.8, z=1.5),\n                center=dict(x=0, y=0, z=0),\n                up=dict(x=0, y=0, z=1)\n            ),\n            bgcolor='rgb(10, 10, 10)'\n        ),\n        paper_bgcolor='rgb(15, 15, 15)',\n        plot_bgcolor='rgb(15, 15, 15)',\n        font=dict(color='white'),\n        width=1200,\n        height=900,\n        annotations=[\n            dict(\n                text=annotation_text,\n                xref=\"paper\",\n                yref=\"paper\",\n                x=0.02,\n                y=0.98,\n                xanchor='left',\n                yanchor='top',\n                showarrow=False,\n                font=dict(size=11, family='monospace', color='white'),\n                bgcolor='rgba(0, 0, 0, 0.7)',\n                bordercolor='white',\n                borderwidth=2,\n                borderpad=10\n            )\n        ],\n        showlegend=False,\n        hovermode='closest'\n    )\n    \n    print(f\"  ✓ Interactive visualization ready!\")\n    \n    return fig\n\n\n# ============================================================================\n# MAIN DEMO FUNCTION\n# ============================================================================\n\ndef demo_interactive_3d(model, val_loader, device, save_html=True):\n    \"\"\"\n    Complete demo with interactive 3D visualization\n    \n    Args:\n        model: Trained model\n        val_loader: Validation DataLoader\n        device: Device\n        save_html: Whether to save HTML file\n    \"\"\"\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"🚀 RUNNING INTERACTIVE 3D VISUALIZATION DEMO\")\n    print(\"=\"*80)\n    \n    # Load model\n    model = model.to(device)\n    model.eval()\n    \n    # Find a patient with fracture\n    print(\"\\n🔍 Finding patient with fracture...\")\n    \n    selected_volume = None\n    selected_label = None\n    selected_id = None\n    \n    for batch_data in val_loader:\n        if len(batch_data) == 3:\n            volumes, labels, patient_ids = batch_data\n        else:\n            volumes, labels = batch_data[0], batch_data[1]\n            patient_ids = [f\"Patient_{i}\" for i in range(len(volumes))]\n        \n        for i in range(len(volumes)):\n            if labels[i][0].item() == 1:  # Has fracture\n                selected_volume = volumes[i:i+1]\n                selected_label = labels[i]\n                selected_id = patient_ids[i]\n                print(f\"  ✓ Found fracture case: {selected_id}\")\n                break\n        \n        if selected_volume is not None:\n            break\n    \n    if selected_volume is None:\n        print(\"  ⚠️  No fracture found, using first patient\")\n        selected_volume = volumes[0:0+1]\n        selected_label = labels[0]\n        selected_id = patient_ids[0]\n    \n    # Move to device and get predictions\n    selected_volume = selected_volume.to(device)\n    \n    with torch.no_grad():\n        output = model(selected_volume)\n        predictions = torch.sigmoid(output).cpu().numpy()[0]\n    \n    print(f\"\\n  Model Prediction: {predictions[0]*100:.1f}% fracture probability\")\n    \n    # Generate Grad-CAM\n    print(f\"\\n📊 Generating Grad-CAM...\")\n    \n    gradcam = GradCAM3D(model, model.backbone.layer4[-1])\n    selected_volume.requires_grad = True\n    cam = gradcam.generate_cam(selected_volume, target_class=0)\n    gradcam.remove_hooks()\n    \n    print(f\"  ✓ Grad-CAM generated\")\n    \n    # Get volume for visualization\n    volume_np = selected_volume[0, 0].detach().cpu().numpy()\n    \n    # Create interactive 3D visualization\n    fig = create_interactive_3d_visualization(\n        volume=volume_np,\n        cam=cam,\n        predictions=predictions,\n        labels=selected_label.numpy(),\n        patient_id=selected_id,\n        threshold=0.4,\n        downsample=0.4  # Reduce for performance\n    )\n    \n    if fig is None:\n        print(\"\\n  ✗ Visualization failed\")\n        return None\n    \n    # Save HTML\n    if save_html:\n        html_path = f'/kaggle/working/interactive_3d_{selected_id}.html'\n        fig.write_html(html_path)\n        print(f\"\\n💾 Saved interactive HTML: {html_path}\")\n        print(f\"   You can download and open this in a browser!\")\n    \n    # Display\n    print(f\"\\n🌐 Displaying interactive visualization...\")\n    print(f\"   • Rotate: Click and drag\")\n    print(f\"   • Zoom: Scroll wheel\")\n    print(f\"   • Pan: Right-click and drag\")\n    print(f\"   • Hover: See attention values\")\n    \n    fig.show()\n    \n    return fig\n\n\n# ============================================================================\n# EXAMPLE USAGE\n# ============================================================================\n\nif __name__ == \"__main__\":\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"📋 READY TO CREATE INTERACTIVE 3D VISUALIZATION\")\n    print(\"=\"*80)\n    \n    print(\"\"\"\nTo run the interactive 3D visualization:\n\n# Make sure you have:\n# 1. Trained model (model)\n# 2. Validation loader (balanced_val_loader or val_loader)\n# 3. Device (device)\n\n# Run the demo:\nfig = demo_interactive_3d(\n    model=model,\n    val_loader=balanced_val_loader,  # or val_loader\n    device=device,\n    save_html=True\n)\n\n# This will:\n# 1. Find a patient with fracture\n# 2. Create 3D mesh of the spine\n# 3. Overlay Grad-CAM heatmap\n# 4. Create interactive Plotly visualization\n# 5. Save as HTML file (downloadable)\n# 6. Display in notebook\n\n# The result is a rotating 3D spine with:\n# • Color-coded fracture attention (blue → red)\n# • Interactive rotation, zoom, pan\n# • Hover to see attention values\n# • Predictions panel overlay\n# • Dark professional theme\n\nPerfect for presentations! Show your sir a rotating 3D spine! 🚀\n    \"\"\")\n    \n    print(\"\\n\" + \"=\"*80)\n    print(\"🎯 FEATURES:\")\n    print(\"=\"*80)\n    print(\"\"\"\n✓ Interactive 3D mesh reconstruction\n✓ Grad-CAM heatmap overlay (blue = normal, red = fracture)\n✓ Smooth rotation and zoom\n✓ Hover tooltips with attention values\n✓ Predictions panel showing all vertebrae\n✓ Professional dark theme\n✓ Exportable as HTML (shareable file)\n✓ Works with demo-trained model (no full training needed!)\n\nThis is the MOST IMPRESSIVE visualization for your presentation!\n    \"\"\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T20:07:58.436364Z","iopub.execute_input":"2025-12-30T20:07:58.437003Z","iopub.status.idle":"2025-12-30T20:07:58.474782Z","shell.execute_reply.started":"2025-12-30T20:07:58.436973Z","shell.execute_reply":"2025-12-30T20:07:58.474113Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"fig = demo_interactive_3d(\n    model=model,\n    val_loader=balanced_val_loader,  # or val_loader\n    device=device,\n    save_html=True\n)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-12-30T20:08:12.379906Z","iopub.execute_input":"2025-12-30T20:08:12.38064Z","iopub.status.idle":"2025-12-30T20:08:20.32672Z","shell.execute_reply.started":"2025-12-30T20:08:12.380613Z","shell.execute_reply":"2025-12-30T20:08:20.325913Z"}},"outputs":[],"execution_count":null}]}