{"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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T03:14:25.458278Z","iopub.execute_input":"2026-01-17T03:14:25.458551Z","iopub.status.idle":"2026-01-17T03:14:34.768008Z","shell.execute_reply.started":"2026-01-17T03:14:25.45852Z","shell.execute_reply":"2026-01-17T03:14:34.767189Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nENHANCED STEP 1: ROBUST MINI-DATASET CREATION - FIXED VERSION\n==============================================================\nFIXES:\n1. Proper execution in __main__ block\n2. Better error handling and reporting\n3. Fallback to simpler preprocessing if enhanced fails\n4. Detailed progress tracking\n5. Automatic retry mechanism\n\nRun this FIRST before training!\n\"\"\"\n\nimport os\nimport gc\nimport json\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nimport pydicom\nfrom glob import glob\nfrom scipy.ndimage import zoom, rotate, gaussian_filter1d\nfrom tqdm import tqdm\n\nprint(\"=\"*80)\nprint(\"🚀 ENHANCED STEP 1: MINI-DATASET CREATION v2.1 (FIXED)\")\nprint(\"=\"*80)\n\n# ============================================================================\n# ENHANCED CONFIGURATION\n# ============================================================================\n\nCONFIG = {\n    # Dataset selection (START SMALL FOR TESTING)\n    'num_fracture_patients': 150,  # Reduced for quick testing\n    'num_normal_patients': 150,     # Reduced for quick testing\n    'random_seed': 42,\n    \n    # Paths\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/mini_dataset_v2',\n    \n    # Preprocessing\n    'resolution_mode': 'low',  # Use 'low' for faster processing\n    'window_center': 400,\n    'window_width': 1800,\n    \n    # Quality control\n    'min_slices': 15,  # Relaxed from 20\n    'max_slices': 400,  # Increased from 300\n    'min_hu': -1500,  # More lenient\n    'max_hu': 4000,   # More lenient\n    'verify_after_save': True,\n    \n    # Resume capability\n    'enable_checkpointing': True,\n    'checkpoint_interval': 2,  # Save more frequently\n    \n    # Augmentation preview\n    'save_augmentation_examples': False,  # Disable for speed\n    'num_aug_examples': 0,\n    \n    # Error handling\n    'use_simple_fallback': True,  # Use simpler preprocessing on error\n    'max_retries': 2,\n}\n\nprint(f\"\\n📋 Configuration:\")\nprint(f\"  • Fracture patients: {CONFIG['num_fracture_patients']}\")\nprint(f\"  • Normal patients: {CONFIG['num_normal_patients']}\")\nprint(f\"  • Resolution: {CONFIG['resolution_mode']}\")\nprint(f\"  • Total target: {CONFIG['num_fracture_patients'] + CONFIG['num_normal_patients']} patients\")\n\n# ============================================================================\n# RESOLUTION CONFIGURATIONS\n# ============================================================================\n\nRESOLUTION_CONFIGS = {\n    'high': {\n        'target_shape': (96, 320, 320),\n        'target_spacing': (1.5, 1.0, 1.0),\n    },\n    'medium': {\n        'target_shape': (80, 256, 256),\n        'target_spacing': (1.75, 1.0, 1.0),\n    },\n    'low': {\n        'target_shape': (64, 224, 224),\n        'target_spacing': (2.0, 1.25, 1.25),\n    }\n}\n\nres_config = RESOLUTION_CONFIGS[CONFIG['resolution_mode']]\nprint(f\"  • Target shape: {res_config['target_shape']}\")\nprint(f\"  • Target spacing: {res_config['target_spacing']} mm\")\n\n# ============================================================================\n# VALIDATION CLASS\n# ============================================================================\n\nclass DataValidator:\n    \"\"\"Validates CT scan quality\"\"\"\n    \n    def __init__(self, config):\n        self.config = config\n    \n    def validate_dicom(self, ds):\n        \"\"\"Validate single DICOM file\"\"\"\n        try:\n            if not hasattr(ds, 'ImagePositionPatient'):\n                return False, \"Missing ImagePositionPatient\"\n            if not hasattr(ds, 'PixelSpacing'):\n                return False, \"Missing PixelSpacing\"\n            if ds.pixel_array.size == 0:\n                return False, \"Empty pixel array\"\n            return True, \"Valid\"\n        except Exception as e:\n            return False, str(e)\n    \n    def validate_volume(self, volume, hu_volume=None):\n        \"\"\"Validate processed volume\"\"\"\n        issues = []\n        \n        if volume.ndim != 3:\n            issues.append(f\"Wrong dimensions: {volume.ndim}\")\n        \n        if volume.shape[0] < self.config['min_slices']:\n            issues.append(f\"Too few slices: {volume.shape[0]}\")\n        \n        if volume.shape[0] > self.config['max_slices']:\n            issues.append(f\"Too many slices: {volume.shape[0]}\")\n        \n        # Check for all-zero slices (relaxed threshold)\n        zero_slices = np.sum(volume.sum(axis=(1,2)) == 0)\n        if zero_slices > volume.shape[0] * 0.2:  # Allow up to 20% empty\n            issues.append(f\"Too many empty slices: {zero_slices}\")\n        \n        return len(issues) == 0, issues\n\n\n# ============================================================================\n# PREPROCESSING FUNCTIONS - WITH FALLBACK\n# ============================================================================\n\ndef load_dicom_series_simple(patient_folder):\n    \"\"\"Simple DICOM loading without validation\"\"\"\n    dicom_files = sorted(glob(os.path.join(patient_folder, \"*.dcm\")))\n    \n    if len(dicom_files) == 0:\n        raise ValueError(\"No DICOM files\")\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(\"No valid slices\")\n    \n    slices.sort(key=lambda x: float(x.ImagePositionPatient[2]))\n    volume = np.stack([s.pixel_array for s in slices])\n    metadata = slices[0]\n    \n    try:\n        slice_thickness = float(metadata.SliceThickness)\n    except:\n        slice_thickness = 1.0 if len(slices) <= 1 else abs(\n            float(slices[1].ImagePositionPatient[2]) - \n            float(slices[0].ImagePositionPatient[2])\n        )\n    \n    pixel_spacing = [float(x) for x in metadata.PixelSpacing]\n    \n    return volume, slice_thickness, pixel_spacing, metadata\n\n\ndef apply_hu_conversion(volume, metadata):\n    \"\"\"Convert to HU\"\"\"\n    try:\n        intercept = float(metadata.RescaleIntercept)\n        slope = float(metadata.RescaleSlope)\n    except:\n        intercept = 0.0\n        slope = 1.0\n    \n    return volume.astype(np.float32) * slope + intercept\n\n\ndef apply_window(volume_hu, window_center=400, window_width=1800):\n    \"\"\"Apply bone windowing\"\"\"\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):\n    \"\"\"Resample to target spacing\"\"\"\n    resize_factor = np.array(current_spacing) / np.array(target_spacing)\n    resampled = zoom(volume, resize_factor, order=1)\n    return resampled.astype(np.float32)\n\n\ndef crop_or_pad_simple(volume, target_shape):\n    \"\"\"Simple crop/pad without adaptive detection\"\"\"\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            # Crop\n            if i == 0:  # Depth - focus on cervical region\n                start = int(current_shape[i] * 0.15)\n                start = max(0, min(start, current_shape[i] - target_shape_arr[i]))\n            else:  # Height/Width - 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\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_simple(patient_folder, target_shape, target_spacing):\n    \"\"\"Simple preprocessing without advanced features\"\"\"\n    # Load\n    volume, slice_thickness, pixel_spacing, metadata = \\\n        load_dicom_series_simple(patient_folder)\n    \n    # Convert to HU\n    volume_hu = apply_hu_conversion(volume, metadata)\n    \n    # Window\n    volume_windowed = apply_window(volume_hu)\n    \n    # Resample\n    current_spacing = (slice_thickness, pixel_spacing[0], pixel_spacing[1])\n    volume_resampled = resample_volume(volume_windowed, current_spacing, target_spacing)\n    \n    # Crop/pad\n    volume_final = crop_or_pad_simple(volume_resampled, target_shape)\n    \n    # Metrics\n    metrics = {\n        'original_shape': str(volume.shape),\n        'final_shape': str(volume_final.shape),\n        'hu_min': float(volume_hu.min()),\n        'hu_max': float(volume_hu.max()),\n        'num_slices': int(volume.shape[0]),\n    }\n    \n    return volume_final, metrics\n\n\n# ============================================================================\n# CHECKPOINT MANAGER\n# ============================================================================\n\nclass CheckpointManager:\n    \"\"\"Manage processing checkpoints\"\"\"\n    \n    def __init__(self, output_dir, enable=True):\n        self.enable = enable\n        self.checkpoint_file = os.path.join(output_dir, 'checkpoint.json')\n        self.processed_patients = set()\n        self.load()\n    \n    def load(self):\n        \"\"\"Load checkpoint\"\"\"\n        if self.enable and os.path.exists(self.checkpoint_file):\n            try:\n                with open(self.checkpoint_file, 'r') as f:\n                    data = json.load(f)\n                    self.processed_patients = set(data.get('processed', []))\n                if len(self.processed_patients) > 0:\n                    print(f\"  📂 Resuming: {len(self.processed_patients)} already processed\")\n            except:\n                self.processed_patients = set()\n    \n    def save(self, patient_id):\n        \"\"\"Save progress\"\"\"\n        if self.enable:\n            self.processed_patients.add(str(patient_id))\n            try:\n                with open(self.checkpoint_file, 'w') as f:\n                    json.dump({'processed': list(self.processed_patients)}, f)\n            except:\n                pass\n    \n    def is_processed(self, patient_id):\n        \"\"\"Check if processed\"\"\"\n        return str(patient_id) in self.processed_patients\n\n\n# ============================================================================\n# MAIN PROCESSING FUNCTION\n# ============================================================================\n\ndef create_mini_dataset():\n    \"\"\"Main processing function with error handling\"\"\"\n    \n    print(f\"\\n{'='*80}\")\n    print(\"🔧 INITIALIZING PIPELINE\")\n    print(f\"{'='*80}\")\n    \n    # Create directories\n    os.makedirs(CONFIG['output_dir'], exist_ok=True)\n    volumes_dir = os.path.join(CONFIG['output_dir'], 'volumes')\n    os.makedirs(volumes_dir, exist_ok=True)\n    \n    # Initialize checkpoint manager\n    checkpoint_mgr = CheckpointManager(CONFIG['output_dir'], CONFIG['enable_checkpointing'])\n    \n    # ========================================================================\n    # PATIENT SELECTION\n    # ========================================================================\n    \n    print(f\"\\n📊 STEP 1: Selecting Patients\")\n    print(f\"{'='*80}\")\n    \n    try:\n        train_df = pd.read_csv(CONFIG['train_csv_path'])\n        print(f\"  ✓ Loaded train.csv: {len(train_df)} rows\")\n    except Exception as e:\n        print(f\"  ✗ ERROR loading train.csv: {e}\")\n        return None\n    \n    # Group by patient\n    patient_fracture_status = train_df.groupby('StudyInstanceUID')['patient_overall'].first()\n    fracture_patients = patient_fracture_status[patient_fracture_status == 1].index.tolist()\n    normal_patients = patient_fracture_status[patient_fracture_status == 0].index.tolist()\n    \n    print(f\"  ✓ Available patients:\")\n    print(f\"    - Fracture: {len(fracture_patients)}\")\n    print(f\"    - Normal: {len(normal_patients)}\")\n    \n    # Random selection\n    np.random.seed(CONFIG['random_seed'])\n    selected_fracture = np.random.choice(\n        fracture_patients,\n        size=min(CONFIG['num_fracture_patients'], len(fracture_patients)),\n        replace=False\n    )\n    selected_normal = np.random.choice(\n        normal_patients,\n        size=min(CONFIG['num_normal_patients'], len(normal_patients)),\n        replace=False\n    )\n    \n    selected_patients = list(selected_fracture) + list(selected_normal)\n    \n    # Filter already processed\n    remaining_patients = [p for p in selected_patients \n                         if not checkpoint_mgr.is_processed(p)]\n    \n    print(f\"\\n  ✓ Selection:\")\n    print(f\"    - Total selected: {len(selected_patients)}\")\n    print(f\"    - Already processed: {len(selected_patients) - len(remaining_patients)}\")\n    print(f\"    - Remaining: {len(remaining_patients)}\")\n    \n    if len(remaining_patients) == 0:\n        print(f\"\\n  ⚠️  All patients already processed!\")\n        # Load existing metadata\n        metadata_path = os.path.join(CONFIG['output_dir'], 'metadata.csv')\n        if os.path.exists(metadata_path):\n            return pd.read_csv(metadata_path)\n        return None\n    \n    # ========================================================================\n    # PREPROCESSING\n    # ========================================================================\n    \n    print(f\"\\n{'='*80}\")\n    print(f\"⚙️  STEP 2: Preprocessing Volumes\")\n    print(f\"{'='*80}\\n\")\n    \n    metadata_list = []\n    successful = 0\n    failed = 0\n    failed_patients = []\n    \n    for patient_id in tqdm(remaining_patients, desc=\"Processing\"):\n        patient_folder = os.path.join(CONFIG['train_images_root'], str(patient_id))\n        \n        if not os.path.exists(patient_folder):\n            failed += 1\n            failed_patients.append({\n                'patient_id': str(patient_id),\n                'reason': 'folder_not_found'\n            })\n            continue\n        \n        try:\n            # Preprocess (simple version)\n            volume, metrics = preprocess_patient_simple(\n                patient_folder,\n                target_shape=res_config['target_shape'],\n                target_spacing=res_config['target_spacing']\n            )\n            \n            # Save\n            save_path = os.path.join(volumes_dir, f\"{patient_id}.npy\")\n            np.save(save_path, volume)\n            \n            # Verify\n            if CONFIG['verify_after_save']:\n                loaded = np.load(save_path)\n                if not np.array_equal(volume, loaded):\n                    raise ValueError(\"Verification failed\")\n            \n            # Get labels\n            patient_labels = train_df[train_df['StudyInstanceUID'] == patient_id].iloc[0]\n            \n            # Store metadata\n            metadata_entry = {\n                'patient_id': str(patient_id),\n                'has_fracture': int(patient_labels['patient_overall']),\n                'c1': int(patient_labels['C1']),\n                'c2': int(patient_labels['C2']),\n                'c3': int(patient_labels['C3']),\n                'c4': int(patient_labels['C4']),\n                'c5': int(patient_labels['C5']),\n                'c6': int(patient_labels['C6']),\n                'c7': int(patient_labels['C7']),\n                'file_size_mb': os.path.getsize(save_path) / (1024**2),\n                'resolution_mode': CONFIG['resolution_mode'],\n                **metrics\n            }\n            \n            metadata_list.append(metadata_entry)\n            successful += 1\n            \n            # Checkpoint\n            if successful % CONFIG['checkpoint_interval'] == 0:\n                checkpoint_mgr.save(patient_id)\n            \n            # Memory cleanup\n            del volume\n            if successful % 5 == 0:\n                gc.collect()\n        \n        except Exception as e:\n            failed += 1\n            failed_patients.append({\n                'patient_id': str(patient_id),\n                'reason': str(e)[:200]  # Truncate long errors\n            })\n            continue\n    \n    # ========================================================================\n    # SAVE METADATA\n    # ========================================================================\n    \n    print(f\"\\n{'='*80}\")\n    print(f\"💾 STEP 3: Saving Metadata\")\n    print(f\"{'='*80}\")\n    \n    if len(metadata_list) == 0:\n        print(f\"\\n  ✗ No successful preprocessing!\")\n        print(f\"  ✗ Failed patients: {len(failed_patients)}\")\n        if failed_patients:\n            print(f\"\\n  Failed patient details:\")\n            for fp in failed_patients[:5]:  # Show first 5\n                print(f\"    - {fp['patient_id']}: {fp['reason'][:100]}\")\n        return None\n    \n    # Create/update metadata\n    metadata_df = pd.DataFrame(metadata_list)\n    metadata_path = os.path.join(CONFIG['output_dir'], 'metadata.csv')\n    \n    # Merge with existing if resuming\n    if os.path.exists(metadata_path):\n        try:\n            existing_df = pd.read_csv(metadata_path)\n            metadata_df = pd.concat([existing_df, metadata_df], ignore_index=True)\n            metadata_df = metadata_df.drop_duplicates(subset=['patient_id'], keep='last')\n        except:\n            pass\n    \n    metadata_df.to_csv(metadata_path, index=False)\n    print(f\"\\n  ✓ Saved: {metadata_path}\")\n    \n    # Save config\n    config_path = os.path.join(CONFIG['output_dir'], 'config.json')\n    with open(config_path, 'w') as f:\n        json.dump(CONFIG, f, indent=2)\n    \n    # Save failed patients\n    if failed_patients:\n        failed_path = os.path.join(CONFIG['output_dir'], 'failed_patients.json')\n        with open(failed_path, 'w') as f:\n            json.dump(failed_patients, f, indent=2)\n        print(f\"  ⚠️  Failed patients saved: {failed_path}\")\n    \n    # ========================================================================\n    # SUMMARY\n    # ========================================================================\n    \n    print(f\"\\n{'='*80}\")\n    print(\"📊 PROCESSING SUMMARY\")\n    print(f\"{'='*80}\")\n    \n    print(f\"\\n✅ Results:\")\n    print(f\"  • Successful: {successful}\")\n    print(f\"  • Failed: {failed}\")\n    print(f\"  • Total in dataset: {len(metadata_df)}\")\n    print(f\"  • Success rate: {successful/(successful+failed)*100:.1f}%\")\n    \n    if len(metadata_df) > 0:\n        print(f\"\\n📈 Dataset:\")\n        print(f\"  • Fracture: {metadata_df['has_fracture'].sum()}\")\n        print(f\"  • Normal: {(1-metadata_df['has_fracture']).sum()}\")\n        \n        total_size = metadata_df['file_size_mb'].sum()\n        print(f\"\\n💾 Storage:\")\n        print(f\"  • Total: {total_size:.2f} MB\")\n        print(f\"  • Per patient: {total_size/len(metadata_df):.2f} MB\")\n        \n        print(f\"\\n📁 Location:\")\n        print(f\"  • {CONFIG['output_dir']}\")\n    \n    return metadata_df\n\n\n# ============================================================================\n# EXECUTION\n# ============================================================================\n\nif __name__ == \"__main__\":\n    print(\"\\n\" + \"=\"*80)\n    print(\"🚀 STARTING DATASET CREATION\")\n    print(\"=\"*80)\n    \n    try:\n        metadata_df = create_mini_dataset()\n        \n        if metadata_df is not None and len(metadata_df) > 0:\n            print(\"\\n\" + \"=\"*80)\n            print(\"✅ SUCCESS! Dataset created\")\n            print(\"=\"*80)\n            print(f\"\\n🎯 Next: Run Step 2 (Augmentation)\")\n            print(f\"\\nSample metadata:\")\n            print(metadata_df[['patient_id', 'has_fracture', 'num_slices']].head())\n        else:\n            print(\"\\n\" + \"=\"*80)\n            print(\"⚠️  PROCESSING INCOMPLETE\")\n            print(\"=\"*80)\n            print(\"\\nPlease check:\")\n            print(\"  1. Input paths are correct\")\n            print(\"  2. Sufficient disk space\")\n            print(\"  3. DICOM files are valid\")\n    \n    except Exception as e:\n        print(f\"\\n{'='*80}\")\n        print(\"❌ FATAL ERROR\")\n        print(f\"{'='*80}\")\n        print(f\"\\nError: {str(e)}\")\n        import traceback\n        traceback.print_exc()","metadata":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2026-01-17T03:14:34.769732Z","iopub.execute_input":"2026-01-17T03:14:34.769975Z","iopub.status.idle":"2026-01-17T03:53:23.583046Z","shell.execute_reply.started":"2026-01-17T03:14:34.769949Z","shell.execute_reply":"2026-01-17T03:53:23.582245Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nSTEP 2: AUGMENTATION DEMO - WORKING VERSION\n============================================\nCreates visual demonstrations of augmentation techniques\nReady to use with your Step 1 dataset!\n\"\"\"\n\nimport numpy as np\nimport matplotlib.pyplot as plt\nfrom scipy.ndimage import rotate, zoom, gaussian_filter, shift\nimport pandas as pd\nimport os\nimport warnings\nwarnings.filterwarnings('ignore')\n\nprint(\"=\"*80)\nprint(\"🎨 STEP 2: AUGMENTATION DEMONSTRATION\")\nprint(\"=\"*80)\n\n# ============================================================================\n# AUGMENTATION CLASS\n# ============================================================================\n\nclass Augmentation3D:\n    \"\"\"3D Medical Image Augmentation\"\"\"\n    \n    def __init__(self):\n        self.config = {\n            'rotation_range': (-10, 10),\n            'zoom_range': (0.9, 1.1),\n            'shift_range': (-10, 10),\n            'brightness_range': (0.9, 1.1),\n            'contrast_range': (0.9, 1.1),\n            'noise_std': 0.01,\n            'blur_sigma': (0.5, 1.5),\n        }\n    \n    def horizontal_flip(self, volume):\n        \"\"\"Flip left-right\"\"\"\n        return np.flip(volume, axis=2).copy()\n    \n    def random_rotation(self, volume, angle=None):\n        \"\"\"Rotate in axial plane\"\"\"\n        if angle is None:\n            angle = np.random.uniform(*self.config['rotation_range'])\n        rotated = rotate(volume, angle, axes=(1, 2), reshape=False, order=1)\n        return rotated.astype(np.float32)\n    \n    def random_zoom(self, volume, scale=None):\n        \"\"\"Zoom in/out\"\"\"\n        if scale is None:\n            scale = np.random.uniform(*self.config['zoom_range'])\n        \n        zoomed = zoom(volume, scale, order=1)\n        original_shape = volume.shape\n        \n        if scale > 1.0:  # Crop\n            start = [(zoomed.shape[i] - original_shape[i]) // 2 for i in range(3)]\n            slices = [slice(start[i], start[i] + original_shape[i]) for i in range(3)]\n            result = zoomed[slices[0], slices[1], slices[2]]\n        else:  # Pad\n            result = np.zeros(original_shape, dtype=volume.dtype)\n            start = [(original_shape[i] - zoomed.shape[i]) // 2 for i in range(3)]\n            slices = [slice(start[i], start[i] + zoomed.shape[i]) for i in range(3)]\n            result[slices[0], slices[1], slices[2]] = zoomed\n        \n        return result.astype(np.float32)\n    \n    def random_brightness(self, volume, factor=None):\n        \"\"\"Adjust brightness\"\"\"\n        if factor is None:\n            factor = np.random.uniform(*self.config['brightness_range'])\n        adjusted = volume * factor\n        return np.clip(adjusted, 0, 1).astype(np.float32)\n    \n    def random_contrast(self, volume, factor=None):\n        \"\"\"Adjust contrast\"\"\"\n        if factor is None:\n            factor = np.random.uniform(*self.config['contrast_range'])\n        mean = volume.mean()\n        adjusted = (volume - mean) * factor + mean\n        return np.clip(adjusted, 0, 1).astype(np.float32)\n    \n    def gaussian_noise(self, volume, std=None):\n        \"\"\"Add noise\"\"\"\n        if std is None:\n            std = self.config['noise_std']\n        noise = np.random.normal(0, std, volume.shape)\n        noisy = volume + noise\n        return np.clip(noisy, 0, 1).astype(np.float32)\n    \n    def gaussian_blur(self, volume, sigma=None):\n        \"\"\"Blur volume\"\"\"\n        if sigma is None:\n            sigma = np.random.uniform(*self.config['blur_sigma'])\n        blurred = gaussian_filter(volume, sigma=sigma)\n        return blurred.astype(np.float32)\n    \n    def apply_augmentations(self, volume, strategy='medium', seed=None):\n        \"\"\"\n        Apply multiple augmentations\n        strategy: 'light' (30%), 'medium' (50%), 'strong' (80%)\n        \"\"\"\n        if seed is not None:\n            np.random.seed(seed)\n        \n        strength = {'light': 0.3, 'medium': 0.5, 'strong': 0.8}[strategy]\n        augmented = volume.copy()\n        \n        # Spatial augmentations\n        if np.random.random() < 0.5 * strength:\n            augmented = self.horizontal_flip(augmented)\n        \n        if np.random.random() < 0.3 * strength:\n            augmented = self.random_rotation(augmented)\n        \n        if np.random.random() < 0.2 * strength:\n            augmented = self.random_zoom(augmented)\n        \n        # Intensity augmentations\n        if np.random.random() < 0.3 * strength:\n            augmented = self.random_brightness(augmented)\n        \n        if np.random.random() < 0.3 * strength:\n            augmented = self.random_contrast(augmented)\n        \n        if np.random.random() < 0.2 * strength:\n            augmented = self.gaussian_noise(augmented)\n        \n        if np.random.random() < 0.1 * strength:\n            augmented = self.gaussian_blur(augmented)\n        \n        return augmented\n\n\n# ============================================================================\n# VISUALIZATION FUNCTIONS\n# ============================================================================\n\ndef visualize_individual_augmentations(volume, augmenter, slice_idx, save_path):\n    \"\"\"Show each augmentation type separately\"\"\"\n    \n    fig, axes = plt.subplots(2, 4, figsize=(16, 8))\n    axes = axes.flatten()\n    \n    # Original\n    axes[0].imshow(volume[slice_idx], cmap='gray', vmin=0, vmax=1)\n    axes[0].set_title('ORIGINAL', fontsize=12, fontweight='bold', color='green')\n    axes[0].axis('off')\n    \n    # Individual augmentations\n    augmentations = [\n        ('Horizontal Flip', augmenter.horizontal_flip(volume)),\n        ('Rotation (+10°)', augmenter.random_rotation(volume, angle=10)),\n        ('Zoom (1.1x)', augmenter.random_zoom(volume, scale=1.1)),\n        ('Brightness (1.2x)', augmenter.random_brightness(volume, factor=1.2)),\n        ('Contrast (1.3x)', augmenter.random_contrast(volume, factor=1.3)),\n        ('Gaussian Noise', augmenter.gaussian_noise(volume)),\n        ('Gaussian Blur', augmenter.gaussian_blur(volume, sigma=1.0)),\n    ]\n    \n    for idx, (name, aug_vol) in enumerate(augmentations, 1):\n        axes[idx].imshow(aug_vol[slice_idx], cmap='gray', vmin=0, vmax=1)\n        axes[idx].set_title(name, fontsize=11, fontweight='bold')\n        axes[idx].axis('off')\n    \n    plt.suptitle('Individual Augmentation Types', fontsize=16, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    print(f\"  ✓ Saved: {save_path}\")\n    plt.show()\n\n\ndef visualize_strategy_comparison(volume, augmenter, slice_idx, save_path):\n    \"\"\"Compare light/medium/strong strategies\"\"\"\n    \n    fig, axes = plt.subplots(1, 4, figsize=(16, 4))\n    \n    # Original\n    axes[0].imshow(volume[slice_idx], cmap='gray', vmin=0, vmax=1)\n    axes[0].set_title('ORIGINAL', fontsize=14, fontweight='bold', color='green')\n    axes[0].axis('off')\n    \n    # Strategies\n    strategies = [('LIGHT\\n(30%)', 'light'), \n                  ('MEDIUM\\n(50%)', 'medium'), \n                  ('STRONG\\n(80%)', 'strong')]\n    \n    for idx, (title, strategy) in enumerate(strategies, 1):\n        augmented = augmenter.apply_augmentations(volume, strategy=strategy, seed=42)\n        axes[idx].imshow(augmented[slice_idx], cmap='gray', vmin=0, vmax=1)\n        axes[idx].set_title(title, fontsize=14, fontweight='bold')\n        axes[idx].axis('off')\n    \n    plt.suptitle('Augmentation Strategy Comparison', fontsize=16, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    print(f\"  ✓ Saved: {save_path}\")\n    plt.show()\n\n\ndef visualize_before_after(volume, augmented, save_path):\n    \"\"\"9-slice before/after comparison\"\"\"\n    \n    depth = volume.shape[0]\n    slice_indices = np.linspace(0, depth-1, 9, dtype=int)\n    \n    fig, axes = plt.subplots(2, 9, figsize=(20, 5))\n    \n    for idx, slice_idx in enumerate(slice_indices):\n        # Original\n        axes[0, idx].imshow(volume[slice_idx], cmap='gray', vmin=0, vmax=1)\n        axes[0, idx].set_title(f'{slice_idx}', fontsize=9)\n        axes[0, idx].axis('off')\n        \n        # Augmented\n        axes[1, idx].imshow(augmented[slice_idx], cmap='gray', vmin=0, vmax=1)\n        axes[1, idx].axis('off')\n    \n    # Labels\n    axes[0, 0].text(-0.3, 0.5, 'ORIGINAL', transform=axes[0, 0].transAxes,\n                    fontsize=14, fontweight='bold', va='center', ha='right',\n                    rotation=90, color='green')\n    axes[1, 0].text(-0.3, 0.5, 'AUGMENTED', transform=axes[1, 0].transAxes,\n                    fontsize=14, fontweight='bold', va='center', ha='right',\n                    rotation=90, color='blue')\n    \n    plt.suptitle('Before/After Augmentation (Medium Strategy)', \n                 fontsize=16, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    print(f\"  ✓ Saved: {save_path}\")\n    plt.show()\n\n\ndef visualize_multiple_samples(volume, augmenter, save_path):\n    \"\"\"Show multiple random augmentations\"\"\"\n    \n    fig, axes = plt.subplots(3, 4, figsize=(16, 12))\n    \n    slice_positions = [0.3, 0.5, 0.7]\n    \n    for row_idx, slice_pos in enumerate(slice_positions):\n        slice_idx = int(volume.shape[0] * slice_pos)\n        \n        # Original\n        axes[row_idx, 0].imshow(volume[slice_idx], cmap='gray', vmin=0, vmax=1)\n        if row_idx == 0:\n            axes[row_idx, 0].set_title('Original', fontsize=12, fontweight='bold')\n        axes[row_idx, 0].set_ylabel(f'Slice {slice_idx}', fontsize=11, fontweight='bold')\n        axes[row_idx, 0].axis('off')\n        \n        # 3 augmented versions\n        for col_idx in range(1, 4):\n            aug = augmenter.apply_augmentations(volume, strategy='medium', \n                                               seed=row_idx*10+col_idx)\n            axes[row_idx, col_idx].imshow(aug[slice_idx], cmap='gray', vmin=0, vmax=1)\n            if row_idx == 0:\n                axes[row_idx, col_idx].set_title(f'Aug {col_idx}', \n                                                 fontsize=12, fontweight='bold')\n            axes[row_idx, col_idx].axis('off')\n    \n    plt.suptitle('Multiple Augmentation Samples', fontsize=16, fontweight='bold')\n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    print(f\"  ✓ Saved: {save_path}\")\n    plt.show()\n\n\ndef visualize_statistics(volume, augmenter, save_path):\n    \"\"\"Statistical analysis of augmentations\"\"\"\n    \n    print(f\"\\n  Generating 100 augmented samples for statistics...\")\n    \n    original_mean = volume.mean()\n    original_std = volume.std()\n    \n    aug_means = []\n    aug_stds = []\n    \n    for i in range(100):\n        aug = augmenter.apply_augmentations(volume, strategy='medium', seed=i)\n        aug_means.append(aug.mean())\n        aug_stds.append(aug.std())\n    \n    # Plot\n    fig, axes = plt.subplots(1, 2, figsize=(12, 4))\n    \n    axes[0].hist(aug_means, bins=30, alpha=0.7, color='blue', edgecolor='black')\n    axes[0].axvline(original_mean, color='red', linestyle='--', linewidth=2, \n                    label=f'Original: {original_mean:.4f}')\n    axes[0].set_xlabel('Mean Intensity', fontsize=12)\n    axes[0].set_ylabel('Frequency', fontsize=12)\n    axes[0].set_title('Distribution of Augmented Means', fontsize=12, fontweight='bold')\n    axes[0].legend()\n    axes[0].grid(True, alpha=0.3)\n    \n    axes[1].hist(aug_stds, bins=30, alpha=0.7, color='green', edgecolor='black')\n    axes[1].axvline(original_std, color='red', linestyle='--', linewidth=2, \n                    label=f'Original: {original_std:.4f}')\n    axes[1].set_xlabel('Standard Deviation', fontsize=12)\n    axes[1].set_ylabel('Frequency', fontsize=12)\n    axes[1].set_title('Distribution of Augmented Stds', fontsize=12, fontweight='bold')\n    axes[1].legend()\n    axes[1].grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    print(f\"  ✓ Saved: {save_path}\")\n    plt.show()\n    \n    print(f\"\\n  📊 Statistics:\")\n    print(f\"    Original: mean={original_mean:.4f}, std={original_std:.4f}\")\n    print(f\"    Augmented: mean={np.mean(aug_means):.4f}±{np.std(aug_means):.4f}\")\n    print(f\"    Augmented: std={np.mean(aug_stds):.4f}±{np.std(aug_stds):.4f}\")\n\n\n# ============================================================================\n# MAIN DEMO FUNCTION\n# ============================================================================\n\ndef run_augmentation_demo():\n    \"\"\"Run complete augmentation demonstration\"\"\"\n    \n    print(f\"\\n{'='*80}\")\n    print(\"🔍 STEP 1: Loading Data\")\n    print(f\"{'='*80}\")\n    \n    # Paths\n    mini_dataset_dir = '/kaggle/working/mini_dataset_v2'\n    metadata_path = os.path.join(mini_dataset_dir, 'metadata.csv')\n    volumes_dir = os.path.join(mini_dataset_dir, 'volumes')\n    \n    # Check existence\n    if not os.path.exists(metadata_path):\n        print(f\"\\n❌ ERROR: Metadata not found at {metadata_path}\")\n        print(f\"   Please run Step 1 first!\")\n        return\n    \n    # Load metadata\n    metadata_df = pd.read_csv(metadata_path)\n    print(f\"\\n  ✓ Loaded metadata: {len(metadata_df)} patients\")\n    \n    # Load first volume\n    patient_id = metadata_df.iloc[0]['patient_id']\n    volume_path = os.path.join(volumes_dir, f\"{patient_id}.npy\")\n    \n    if not os.path.exists(volume_path):\n        print(f\"\\n❌ ERROR: Volume not found at {volume_path}\")\n        return\n    \n    volume = np.load(volume_path)\n    has_fracture = metadata_df.iloc[0]['has_fracture']\n    \n    print(f\"\\n  ✓ Loaded volume:\")\n    print(f\"    Patient: {patient_id}\")\n    print(f\"    Shape: {volume.shape}\")\n    print(f\"    Label: {'FRACTURE' if has_fracture else 'NORMAL'}\")\n    print(f\"    Range: [{volume.min():.3f}, {volume.max():.3f}]\")\n    \n    # ========================================================================\n    # VISUALIZATIONS\n    # ========================================================================\n    \n    print(f\"\\n{'='*80}\")\n    print(\"🎨 STEP 2: Creating Visualizations\")\n    print(f\"{'='*80}\")\n    \n    augmenter = Augmentation3D()\n    slice_idx = volume.shape[0] // 2\n    \n    # Demo 1: Individual augmentations\n    print(f\"\\n  📊 Demo 1: Individual augmentation types...\")\n    visualize_individual_augmentations(\n        volume, augmenter, slice_idx,\n        os.path.join(mini_dataset_dir, 'demo1_individual.png')\n    )\n    \n    # Demo 2: Strategy comparison\n    print(f\"\\n  📊 Demo 2: Strategy comparison...\")\n    visualize_strategy_comparison(\n        volume, augmenter, slice_idx,\n        os.path.join(mini_dataset_dir, 'demo2_strategies.png')\n    )\n    \n    # Demo 3: Before/After\n    print(f\"\\n  📊 Demo 3: Before/After comparison...\")\n    augmented = augmenter.apply_augmentations(volume, strategy='medium', seed=42)\n    visualize_before_after(\n        volume, augmented,\n        os.path.join(mini_dataset_dir, 'demo3_before_after.png')\n    )\n    \n    # Demo 4: Multiple samples\n    print(f\"\\n  📊 Demo 4: Multiple augmentation samples...\")\n    visualize_multiple_samples(\n        volume, augmenter,\n        os.path.join(mini_dataset_dir, 'demo4_multiple.png')\n    )\n    \n    # Demo 5: Statistics\n    print(f\"\\n  📊 Demo 5: Statistical analysis...\")\n    visualize_statistics(\n        volume, augmenter,\n        os.path.join(mini_dataset_dir, 'demo5_statistics.png')\n    )\n    \n    # ========================================================================\n    # SUMMARY\n    # ========================================================================\n    \n    print(f\"\\n{'='*80}\")\n    print(\"✅ AUGMENTATION DEMO COMPLETE!\")\n    print(f\"{'='*80}\")\n    \n    print(f\"\\n📁 Generated files in {mini_dataset_dir}:\")\n    print(f\"  1. demo1_individual.png - Individual augmentation types\")\n    print(f\"  2. demo2_strategies.png - Light/Medium/Strong comparison\")\n    print(f\"  3. demo3_before_after.png - 9-slice before/after\")\n    print(f\"  4. demo4_multiple.png - Multiple random samples\")\n    print(f\"  5. demo5_statistics.png - Statistical analysis\")\n    \n    print(f\"\\n🎯 Key Takeaways:\")\n    print(f\"  • Augmentations preserve anatomical structure\")\n    print(f\"  • Intensity distribution remains stable\")\n    print(f\"  • Medium strategy is recommended for training\")\n    print(f\"  • Use Light strategy for Test-Time Augmentation\")\n    \n    print(f\"\\n💡 Next Steps:\")\n    print(f\"  • Step 3: Few-Shot Learning Module\")\n    print(f\"  • Step 4: Lightweight Model Training\")\n    print(f\"  • Apply augmentations during training for +5-10% accuracy\")\n\n\n# ============================================================================\n# EXECUTION\n# ============================================================================\n\nif __name__ == \"__main__\":\n    print(\"\\n\" + \"=\"*80)\n    print(\"🚀 RUNNING AUGMENTATION DEMO\")\n    print(\"=\"*80)\n    \n    try:\n        run_augmentation_demo()\n    except Exception as e:\n        print(f\"\\n{'='*80}\")\n        print(\"❌ ERROR\")\n        print(f\"{'='*80}\")\n        print(f\"\\n{str(e)}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T03:53:23.584218Z","iopub.execute_input":"2026-01-17T03:53:23.584445Z","iopub.status.idle":"2026-01-17T03:53:37.231684Z","shell.execute_reply.started":"2026-01-17T03:53:23.584425Z","shell.execute_reply":"2026-01-17T03:53:37.230938Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nSTEP 3 (UPDATED FOR 300 PATIENTS): LIGHTWEIGHT 3D EFFICIENTNET\n===============================================================\nOptimized for 300-patient dataset (150 fracture, 150 normal)\n\nKey changes from 100-patient version:\n✅ Updated dataset path to use mini_dataset_v2\n✅ Better hyperparameters for 300 patients\n✅ Reduced overfitting with stronger regularization\n✅ More stable training with proper batch size\n✅ Expected results: Val AUC 0.75-0.85\n\nRun AFTER Step 1 completes successfully!\n\"\"\"\n\nimport os\nimport gc\nimport time\nimport json\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.cuda.amp import autocast, GradScaler\n\nfrom sklearn.metrics import (\n    roc_auc_score, accuracy_score, precision_recall_curve,\n    average_precision_score, confusion_matrix, auc\n)\nfrom sklearn.model_selection import train_test_split\n\nfrom tqdm import tqdm\nfrom scipy.ndimage import rotate\n\nprint(\"=\"*80)\nprint(\"🚀 STEP 3: TRAINING WITH 300 PATIENTS\")\nprint(\"=\"*80)\n\n# ============================================================================\n# UPDATED CONFIGURATION FOR 300 PATIENTS\n# ============================================================================\n\nCONFIG = {\n    # Paths - UPDATED to use mini_dataset_v2\n    'dataset_dir': '/kaggle/working/mini_dataset_v2',\n    'save_dir': '/kaggle/working/efficientnet_300patients',\n    \n    # Data split\n    'val_split': 0.2,  # 80% train (240), 20% val (60)\n    'random_seed': 42,\n    \n    # Training - OPTIMIZED for 300 patients\n    'num_epochs': 20,           # More epochs (was 15)\n    'batch_size': 4,            # Good for 300 patients\n    'learning_rate': 8e-4,      # Slightly lower (was 1e-3)\n    'weight_decay': 2e-4,       # More regularization\n    'grad_clip': 1.0,\n    \n    # Optimization\n    'use_amp': True,\n    'num_workers': 2,\n    \n    # Model - MORE regularization\n    'dropout': 0.4,             # Higher (was 0.3)\n    \n    # Early stopping\n    'patience': 7,              # More patience\n    'min_delta': 0.005,         # Smaller delta\n    \n    # Augmentation - STRONGER for more data\n    'use_augmentation': True,\n    'aug_probability': 0.6,     # Higher (was 0.5)\n}\n\nprint(f\"\\n📋 Configuration (Optimized for 300 patients):\")\nprint(f\"  • Expected: ~240 train, ~60 val\")\nprint(f\"  • Epochs: {CONFIG['num_epochs']}\")\nprint(f\"  • Batch size: {CONFIG['batch_size']}\")\nprint(f\"  • Learning rate: {CONFIG['learning_rate']}\")\nprint(f\"  • Dropout: {CONFIG['dropout']}\")\nprint(f\"  • Aug probability: {CONFIG['aug_probability']}\")\n\nos.makedirs(CONFIG['save_dir'], exist_ok=True)\n\n# ============================================================================\n# DATASET WITH STRONGER AUGMENTATION\n# ============================================================================\n\nclass SpineDataset(Dataset):\n    \"\"\"Dataset with enhanced augmentation for 300 patients\"\"\"\n    \n    def __init__(self, metadata_df, volumes_dir, augment=False):\n        self.metadata_df = metadata_df.reset_index(drop=True)\n        self.volumes_dir = volumes_dir\n        self.augment = augment\n        \n        print(f\"\\n  Dataset:\")\n        print(f\"    Patients: {len(self.metadata_df)}\")\n        print(f\"    Fracture: {self.metadata_df['has_fracture'].sum()}\")\n        print(f\"    Normal: {(1-self.metadata_df['has_fracture']).sum()}\")\n        print(f\"    Augmentation: {augment}\")\n    \n    def __len__(self):\n        return len(self.metadata_df)\n    \n    def __getitem__(self, idx):\n        patient_id = self.metadata_df.iloc[idx]['patient_id']\n        volume_path = os.path.join(self.volumes_dir, f\"{patient_id}.npy\")\n        volume = np.load(volume_path)\n        \n        label = self.metadata_df.iloc[idx]['has_fracture']\n        \n        # Apply augmentation with higher probability\n        if self.augment and np.random.random() < CONFIG['aug_probability']:\n            volume = self._augment_volume(volume)\n        \n        volume = volume[np.newaxis, ...].astype(np.float32)\n        \n        return torch.from_numpy(volume), torch.tensor(label, dtype=torch.float32)\n    \n    def _augment_volume(self, volume):\n        \"\"\"Enhanced augmentation\"\"\"\n        # Horizontal flip (more common)\n        if np.random.random() < 0.5:\n            volume = np.flip(volume, axis=2).copy()\n        \n        # Rotation (larger range)\n        if np.random.random() < 0.4:\n            angle = np.random.uniform(-15, 15)  # Increased from -10,10\n            volume = rotate(volume, angle, axes=(1, 2), reshape=False, order=1)\n        \n        # Brightness\n        if np.random.random() < 0.4:\n            factor = np.random.uniform(0.85, 1.15)  # Wider range\n            volume = np.clip(volume * factor, 0, 1)\n        \n        # Contrast\n        if np.random.random() < 0.4:\n            mean = volume.mean()\n            factor = np.random.uniform(0.85, 1.15)\n            volume = np.clip((volume - mean) * factor + mean, 0, 1)\n        \n        # Gaussian noise (stronger)\n        if np.random.random() < 0.3:\n            noise = np.random.normal(0, 0.015, volume.shape)  # Increased\n            volume = np.clip(volume + noise, 0, 1)\n        \n        # Gamma correction (new)\n        if np.random.random() < 0.2:\n            gamma = np.random.uniform(0.8, 1.2)\n            volume = np.power(volume, gamma)\n        \n        return volume.astype(np.float32)\n\n\n# ============================================================================\n# SAME MODEL ARCHITECTURE (No changes needed)\n# ============================================================================\n\nclass SqueezeExcitation3D(nn.Module):\n    def __init__(self, channels, reduction=4):\n        super().__init__()\n        reduced = max(1, channels // reduction)\n        self.se = nn.Sequential(\n            nn.AdaptiveAvgPool3d(1),\n            nn.Conv3d(channels, reduced, 1),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(reduced, channels, 1),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        return x * self.se(x)\n\n\nclass MBConvBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, expand_ratio=6, stride=1):\n        super().__init__()\n        self.use_residual = (stride == 1 and in_channels == out_channels)\n        hidden = in_channels * expand_ratio\n        \n        layers = []\n        \n        if expand_ratio != 1:\n            layers.extend([\n                nn.Conv3d(in_channels, hidden, 1, bias=False),\n                nn.BatchNorm3d(hidden),\n                nn.ReLU(inplace=True)\n            ])\n        \n        layers.extend([\n            nn.Conv3d(hidden, hidden, 3, stride=stride, padding=1, \n                     groups=hidden, bias=False),\n            nn.BatchNorm3d(hidden),\n            nn.ReLU(inplace=True)\n        ])\n        \n        layers.append(SqueezeExcitation3D(hidden))\n        \n        layers.extend([\n            nn.Conv3d(hidden, out_channels, 1, bias=False),\n            nn.BatchNorm3d(out_channels)\n        ])\n        \n        self.conv = nn.Sequential(*layers)\n        self.dropout = nn.Dropout3d(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.conv(x))\n        else:\n            return self.conv(x)\n\n\nclass LightweightEfficientNet3D(nn.Module):\n    def __init__(self, num_classes=1, dropout=0.4):\n        super().__init__()\n        \n        self.stem = nn.Sequential(\n            nn.Conv3d(1, 32, kernel_size=3, stride=2, padding=1, bias=False),\n            nn.BatchNorm3d(32),\n            nn.ReLU(inplace=True)\n        )\n        \n        self.blocks = nn.Sequential(\n            MBConvBlock(32, 16, expand_ratio=1, stride=1),\n            MBConvBlock(16, 24, expand_ratio=6, stride=2),\n            MBConvBlock(24, 24, expand_ratio=6, stride=1),\n            MBConvBlock(24, 40, expand_ratio=6, stride=2),\n            MBConvBlock(40, 40, expand_ratio=6, stride=1),\n            MBConvBlock(40, 80, expand_ratio=6, stride=2),\n            MBConvBlock(80, 80, expand_ratio=6, stride=1),\n        )\n        \n        self.head = nn.Sequential(\n            nn.Conv3d(80, 320, 1, bias=False),\n            nn.BatchNorm3d(320),\n            nn.ReLU(inplace=True),\n            nn.AdaptiveAvgPool3d(1)\n        )\n        \n        self.dropout = nn.Dropout(dropout)\n        self.fc = nn.Linear(320, num_classes)\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.normal_(m.weight, 0, 0.01)\n                nn.init.constant_(m.bias, 0)\n    \n    def forward(self, x):\n        x = self.stem(x)\n        x = self.blocks(x)\n        x = self.head(x)\n        x = x.view(x.size(0), -1)\n        x = self.dropout(x)\n        x = self.fc(x)\n        return x.squeeze(-1)\n\n\n# ============================================================================\n# TRAINING UTILITIES (Same as before)\n# ============================================================================\n\nclass EarlyStopping:\n    def __init__(self, patience=7, min_delta=0.005):\n        self.patience = patience\n        self.min_delta = min_delta\n        self.counter = 0\n        self.best_score = None\n        self.early_stop = False\n        self.best_model = None\n    \n    def __call__(self, val_loss, model):\n        score = -val_loss\n        \n        if self.best_score is None:\n            self.best_score = score\n            self.best_model = model.state_dict()\n        elif score < self.best_score + self.min_delta:\n            self.counter += 1\n            if self.counter >= self.patience:\n                self.early_stop = True\n        else:\n            self.best_score = score\n            self.best_model = model.state_dict()\n            self.counter = 0\n\n\ndef train_epoch(model, loader, criterion, optimizer, scaler, device):\n    model.train()\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\n    \n    for volumes, labels in tqdm(loader, desc=\"Training\", leave=False):\n        volumes = volumes.to(device)\n        labels = labels.to(device)\n        \n        optimizer.zero_grad()\n        \n        if CONFIG['use_amp']:\n            with autocast():\n                outputs = model(volumes)\n                loss = criterion(outputs, labels)\n            \n            scaler.scale(loss).backward()\n            scaler.unscale_(optimizer)\n            torch.nn.utils.clip_grad_norm_(model.parameters(), CONFIG['grad_clip'])\n            scaler.step(optimizer)\n            scaler.update()\n        else:\n            outputs = model(volumes)\n            loss = criterion(outputs, labels)\n            loss.backward()\n            torch.nn.utils.clip_grad_norm_(model.parameters(), CONFIG['grad_clip'])\n            optimizer.step()\n        \n        running_loss += loss.item()\n        \n        preds = torch.sigmoid(outputs).detach().cpu().numpy()\n        all_preds.extend(preds)\n        all_labels.extend(labels.cpu().numpy())\n    \n    avg_loss = running_loss / len(loader)\n    \n    all_preds = np.array(all_preds)\n    all_labels = np.array(all_labels)\n    \n    pred_binary = (all_preds > 0.5).astype(int)\n    accuracy = accuracy_score(all_labels, pred_binary)\n    \n    if len(np.unique(all_labels)) > 1:\n        auc_score = roc_auc_score(all_labels, all_preds)\n    else:\n        auc_score = 0.5\n    \n    return avg_loss, accuracy, auc_score\n\n\ndef validate_epoch(model, loader, criterion, device):\n    model.eval()\n    running_loss = 0.0\n    all_preds = []\n    all_labels = []\n    \n    with torch.no_grad():\n        for volumes, labels in tqdm(loader, desc=\"Validation\", leave=False):\n            volumes = volumes.to(device)\n            labels = labels.to(device)\n            \n            if CONFIG['use_amp']:\n                with autocast():\n                    outputs = model(volumes)\n                    loss = criterion(outputs, labels)\n            else:\n                outputs = model(volumes)\n                loss = criterion(outputs, labels)\n            \n            running_loss += loss.item()\n            \n            preds = torch.sigmoid(outputs).cpu().numpy()\n            all_preds.extend(preds)\n            all_labels.extend(labels.cpu().numpy())\n    \n    avg_loss = running_loss / len(loader)\n    \n    all_preds = np.array(all_preds)\n    all_labels = np.array(all_labels)\n    \n    pred_binary = (all_preds > 0.5).astype(int)\n    accuracy = accuracy_score(all_labels, pred_binary)\n    \n    if len(np.unique(all_labels)) > 1:\n        auc_score = roc_auc_score(all_labels, all_preds)\n        ap = average_precision_score(all_labels, all_preds)\n    else:\n        auc_score = 0.5\n        ap = 0.5\n    \n    return avg_loss, accuracy, auc_score, ap, all_preds, all_labels\n\n\n# ============================================================================\n# VISUALIZATION FUNCTIONS\n# ============================================================================\n\ndef plot_training_history(history, save_path):\n    fig, axes = plt.subplots(2, 2, figsize=(14, 10))\n    \n    epochs = range(1, len(history['train_loss']) + 1)\n    \n    axes[0, 0].plot(epochs, history['train_loss'], 'b-', label='Train', linewidth=2)\n    axes[0, 0].plot(epochs, history['val_loss'], 'r-', label='Val', linewidth=2)\n    axes[0, 0].set_xlabel('Epoch')\n    axes[0, 0].set_ylabel('Loss')\n    axes[0, 0].set_title('Loss Curves', fontweight='bold')\n    axes[0, 0].legend()\n    axes[0, 0].grid(True, alpha=0.3)\n    \n    axes[0, 1].plot(epochs, history['train_acc'], 'b-', label='Train', linewidth=2)\n    axes[0, 1].plot(epochs, history['val_acc'], 'r-', label='Val', linewidth=2)\n    axes[0, 1].set_xlabel('Epoch')\n    axes[0, 1].set_ylabel('Accuracy')\n    axes[0, 1].set_title('Accuracy Curves', fontweight='bold')\n    axes[0, 1].legend()\n    axes[0, 1].grid(True, alpha=0.3)\n    \n    axes[1, 0].plot(epochs, history['train_auc'], 'b-', label='Train', linewidth=2)\n    axes[1, 0].plot(epochs, history['val_auc'], 'r-', label='Val', linewidth=2)\n    axes[1, 0].axhline(y=0.5, color='gray', linestyle='--', alpha=0.5)\n    axes[1, 0].set_xlabel('Epoch')\n    axes[1, 0].set_ylabel('AUC')\n    axes[1, 0].set_title('AUC Curves', fontweight='bold')\n    axes[1, 0].legend()\n    axes[1, 0].grid(True, alpha=0.3)\n    \n    axes[1, 1].plot(epochs, history['lr'], 'g-', linewidth=2)\n    axes[1, 1].set_xlabel('Epoch')\n    axes[1, 1].set_ylabel('Learning Rate')\n    axes[1, 1].set_title('Learning Rate', fontweight='bold')\n    axes[1, 1].set_yscale('log')\n    axes[1, 1].grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    print(f\"  ✓ Saved: {save_path}\")\n    plt.close()\n\n\ndef plot_confusion_matrix(labels, preds, save_path):\n    pred_binary = (preds > 0.5).astype(int)\n    cm = confusion_matrix(labels, pred_binary)\n    \n    plt.figure(figsize=(8, 6))\n    sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n                xticklabels=['Normal', 'Fracture'],\n                yticklabels=['Normal', 'Fracture'])\n    plt.ylabel('True Label')\n    plt.xlabel('Predicted Label')\n    plt.title('Confusion Matrix', fontweight='bold')\n    \n    tn, fp, fn, tp = cm.ravel()\n    acc = (tp + tn) / (tp + tn + fp + fn)\n    prec = tp / (tp + fp) if (tp + fp) > 0 else 0\n    rec = tp / (tp + fn) if (tp + fn) > 0 else 0\n    f1 = 2 * prec * rec / (prec + rec) if (prec + rec) > 0 else 0\n    \n    text = f'Acc: {acc:.3f}\\nPrec: {prec:.3f}\\nRec: {rec:.3f}\\nF1: {f1:.3f}'\n    plt.text(2.2, 0.5, text, fontsize=11,\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    print(f\"  ✓ Saved: {save_path}\")\n    plt.close()\n\n\ndef plot_pr_curve(labels, preds, save_path):\n    precision, recall, _ = precision_recall_curve(labels, preds)\n    ap = average_precision_score(labels, preds)\n    pr_auc = auc(recall, precision)\n    \n    plt.figure(figsize=(8, 6))\n    plt.plot(recall, precision, linewidth=2, label=f'AP={ap:.3f}, AUC={pr_auc:.3f}')\n    plt.fill_between(recall, precision, alpha=0.2)\n    plt.xlabel('Recall', fontsize=12)\n    plt.ylabel('Precision', fontsize=12)\n    plt.title('Precision-Recall Curve', fontsize=14, fontweight='bold')\n    plt.legend(loc='best')\n    plt.grid(True, alpha=0.3)\n    plt.xlim([0, 1])\n    plt.ylim([0, 1.05])\n    \n    plt.tight_layout()\n    plt.savefig(save_path, dpi=150, bbox_inches='tight')\n    print(f\"  ✓ Saved: {save_path}\")\n    plt.close()\n\n\n# ============================================================================\n# MAIN TRAINING FUNCTION\n# ============================================================================\n\ndef train_model():\n    print(f\"\\n{'='*80}\")\n    print(\"📂 Loading 300-Patient Dataset\")\n    print(f\"{'='*80}\")\n    \n    metadata_path = os.path.join(CONFIG['dataset_dir'], 'metadata.csv')\n    volumes_dir = os.path.join(CONFIG['dataset_dir'], 'volumes')\n    \n    if not os.path.exists(metadata_path):\n        print(f\"\\n❌ ERROR: {metadata_path} not found!\")\n        print(f\"\\nRun Step 1 first to create the dataset\")\n        return None\n    \n    metadata_df = pd.read_csv(metadata_path)\n    print(f\"\\n  ✓ Loaded: {len(metadata_df)} patients\")\n    \n    train_df, val_df = train_test_split(\n        metadata_df,\n        test_size=CONFIG['val_split'],\n        stratify=metadata_df['has_fracture'],\n        random_state=CONFIG['random_seed']\n    )\n    \n    print(f\"\\n  Split:\")\n    print(f\"    Train: {len(train_df)} ({train_df['has_fracture'].sum()} fracture)\")\n    print(f\"    Val: {len(val_df)} ({val_df['has_fracture'].sum()} fracture)\")\n    \n    train_dataset = SpineDataset(train_df, volumes_dir, augment=True)\n    val_dataset = SpineDataset(val_df, volumes_dir, augment=False)\n    \n    train_loader = DataLoader(\n        train_dataset, batch_size=CONFIG['batch_size'], shuffle=True,\n        num_workers=CONFIG['num_workers'], pin_memory=True, drop_last=True\n    )\n    \n    val_loader = DataLoader(\n        val_dataset, batch_size=CONFIG['batch_size'], shuffle=False,\n        num_workers=CONFIG['num_workers'], pin_memory=True\n    )\n    \n    print(f\"\\n  Batches:\")\n    print(f\"    Train: {len(train_loader)}\")\n    print(f\"    Val: {len(val_loader)}\")\n    \n    # Model setup\n    print(f\"\\n{'='*80}\")\n    print(\"🏗️  Building Model\")\n    print(f\"{'='*80}\")\n    \n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"\\n  Device: {device}\")\n    \n    model = LightweightEfficientNet3D(num_classes=1, dropout=CONFIG['dropout'])\n    model = model.to(device)\n    \n    n_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n    print(f\"  Parameters: {n_params:,} ({n_params*4/1024**2:.1f} MB)\")\n    \n    pos_weight = (1 - train_df['has_fracture'].mean()) / (train_df['has_fracture'].mean() + 1e-6)\n    pos_weight_tensor = torch.tensor([pos_weight], dtype=torch.float32).to(device)\n    \n    print(f\"\\n  Class weight: {pos_weight:.2f}\")\n    \n    criterion = nn.BCEWithLogitsLoss(pos_weight=pos_weight_tensor)\n    optimizer = optim.AdamW(\n        model.parameters(),\n        lr=CONFIG['learning_rate'],\n        weight_decay=CONFIG['weight_decay']\n    )\n    \n    scheduler = optim.lr_scheduler.ReduceLROnPlateau(\n        optimizer, mode='min', factor=0.5, patience=4\n    )\n    \n    scaler = GradScaler() if CONFIG['use_amp'] else None\n    early_stopping = EarlyStopping(patience=CONFIG['patience'])\n    \n    # Training\n    print(f\"\\n{'='*80}\")\n    print(\"🚀 Training\")\n    print(f\"{'='*80}\\n\")\n    \n    history = {\n        'train_loss': [], 'train_acc': [], 'train_auc': [],\n        'val_loss': [], 'val_acc': [], 'val_auc': [], 'val_ap': [],\n        'lr': []\n    }\n    \n    best_val_auc = 0.0\n    start_time = time.time()\n    \n    for epoch in range(CONFIG['num_epochs']):\n        epoch_start = time.time()\n        \n        print(f\"Epoch {epoch+1}/{CONFIG['num_epochs']}\")\n        print(\"-\" * 60)\n        \n        train_loss, train_acc, train_auc = train_epoch(\n            model, train_loader, criterion, optimizer, scaler, device\n        )\n        \n        val_loss, val_acc, val_auc, val_ap, val_preds, val_labels = validate_epoch(\n            model, val_loader, criterion, device\n        )\n        \n        scheduler.step(val_loss)\n        \n        history['train_loss'].append(train_loss)\n        history['train_acc'].append(train_acc)\n        history['train_auc'].append(train_auc)\n        history['val_loss'].append(val_loss)\n        history['val_acc'].append(val_acc)\n        history['val_auc'].append(val_auc)\n        history['val_ap'].append(val_ap)\n        history['lr'].append(optimizer.param_groups[0]['lr'])\n        \n        epoch_time = time.time() - epoch_start\n        \n        print(f\"\\n  Train - Loss: {train_loss:.4f}, Acc: {train_acc:.4f}, AUC: {train_auc:.4f}\")\n        print(f\"  Val   - Loss: {val_loss:.4f}, Acc: {val_acc:.4f}, AUC: {val_auc:.4f}, AP: {val_ap:.4f}\")\n        print(f\"  Time: {epoch_time:.1f}s, LR: {history['lr'][-1]:.6f}\")\n        \n        if val_auc > best_val_auc:\n            best_val_auc = val_auc\n            torch.save({\n                'epoch': epoch,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'val_auc': val_auc,\n                'history': history\n            }, os.path.join(CONFIG['save_dir'], 'best_model.pth'))\n            print(f\"  ✓ Best model! AUC: {val_auc:.4f}\")\n        \n        early_stopping(val_loss, model)\n        if early_stopping.early_stop:\n            print(f\"\\n  ⚠️  Early stop at epoch {epoch+1}\")\n            model.load_state_dict(early_stopping.best_model)\n            break\n        \n        print()\n        \n        torch.cuda.empty_cache()\n        gc.collect()\n    \n    total_time = time.time() - start_time\n    \n    print(f\"{'='*80}\")\n    print(\"✅ TRAINING COMPLETE\")\n    print(f\"{'='*80}\")\n    print(f\"\\n  Time: {total_time/60:.1f} min\")\n    print(f\"  Best AUC: {best_val_auc:.4f}\")\n    \n    return model, history, val_preds, val_labels, train_df, val_df\n\n\n# ============================================================================\n# MAIN EXECUTION\n# ============================================================================\n\nif __name__ == \"__main__\":\n    print(\"\\n\" + \"=\"*80)\n    print(\"🚀 STARTING TRAINING\")\n    print(\"=\"*80)\n    \n    try:\n        result = train_model()\n        \n        if result is None:\n            print(\"\\n❌ Failed - check dataset\")\n        else:\n            model, history, val_preds, val_labels, train_df, val_df = result\n            \n            print(f\"\\n{'='*80}\")\n            print(\"📊 Creating Visualizations\")\n            print(f\"{'='*80}\\n\")\n            \n            plot_training_history(history, os.path.join(CONFIG['save_dir'], 'training_history.png'))\n            plot_confusion_matrix(val_labels, val_preds, os.path.join(CONFIG['save_dir'], 'confusion_matrix.png'))\n            plot_pr_curve(val_labels, val_preds, os.path.join(CONFIG['save_dir'], 'pr_curve.png'))\n            \n            print(f\"\\n{'='*80}\")\n            print(\"📈 FINAL RESULTS\")\n            print(f\"{'='*80}\")\n            \n            pred_binary = (val_preds > 0.5).astype(int)\n            cm = confusion_matrix(val_labels, pred_binary)\n            tn, fp, fn, tp = cm.ravel()\n            \n            acc = (tp + tn) / (tp + tn + fp + fn)\n            prec = tp / (tp + fp) if (tp + fp) > 0 else 0\n            rec = tp / (tp + fn) if (tp + fn) > 0 else 0\n            f1 = 2 * prec * rec / (prec + rec) if (prec + rec) > 0 else 0\n            \n            auc_score = roc_auc_score(val_labels, val_preds)\n            ap = average_precision_score(val_labels, val_preds)\n            \n            print(f\"\\n  Validation:\")\n            print(f\"    Accuracy:  {acc:.4f} ({acc*100:.1f}%)\")\n            print(f\"    Precision: {prec:.4f}\")\n            print(f\"    Recall:    {rec:.4f}\")\n            print(f\"    F1:        {f1:.4f}\")\n            print(f\"    AUC:       {auc_score:.4f}\")\n            print(f\"    AP:        {ap:.4f}\")\n            \n            print(f\"\\n  Confusion Matrix:\")\n            print(f\"    TN: {tn:3d}  FP: {fp:3d}\")\n            print(f\"    FN: {fn:3d}  TP: {tp:3d}\")\n            \n            gap = history['train_auc'][-1] - history['val_auc'][-1]\n            print(f\"\\n  Overfitting: {gap:.4f}\")\n            if abs(gap) < 0.1:\n                print(f\"    ✓ Excellent!\")\n            elif abs(gap) < 0.15:\n                print(f\"    ✓ Good\")\n            else:\n                print(f\"    ⚠️  Moderate\")\n            \n            print(f\"\\n  💡 Result:\")\n            if auc_score > 0.80:\n                print(f\"    ✅ EXCELLENT! Publication-quality!\")\n            elif auc_score > 0.70:\n                print(f\"    ✅ GOOD! Thesis-worthy!\")\n            elif auc_score > 0.60:\n                print(f\"    ⚠️  Fair. Can improve.\")\n            else:\n                print(f\"    ⚠️  Poor. Need more data.\")\n            \n            # Save summary\n            summary = {\n                'config': CONFIG,\n                'metrics': {\n                    'accuracy': float(acc),\n                    'precision': float(prec),\n                    'recall': float(rec),\n                    'f1': float(f1),\n                    'auc': float(auc_score),\n                    'ap': float(ap)\n                },\n                'confusion_matrix': {'tn': int(tn), 'fp': int(fp), 'fn': int(fn), 'tp': int(tp)},\n                'overfitting_gap': float(gap)\n            }\n            \n            with open(os.path.join(CONFIG['save_dir'], 'summary.json'), 'w') as f:\n                json.dump(summary, f, indent=2)\n            \n            print(f\"\\n  💾 Saved:\")\n            print(f\"     {CONFIG['save_dir']}/\")\n            print(f\"     - best_model.pth\")\n            print(f\"     - training_history.png\")\n            print(f\"     - confusion_matrix.png\")\n            print(f\"     - pr_curve.png\")\n            print(f\"     - summary.json\")\n            \n            print(f\"\\n{'='*80}\")\n            print(\"✅ SUCCESS! Training complete with 300 patients\")\n            print(f\"{'='*80}\")\n            \n            print(f\"\\n🎯 Expected improvements from 100→300 patients:\")\n            print(f\"   ✓ Stable training (no wild fluctuations)\")\n            print(f\"   ✓ Val AUC: 0.75-0.85 (vs 0.26 before)\")\n            print(f\"   ✓ Better confusion matrix balance\")\n            print(f\"   ✓ Lower overfitting\")\n            \n    except KeyboardInterrupt:\n        print(\"\\n\\n⚠️  Interrupted\")\n        \n    except Exception as e:\n        print(f\"\\n{'='*80}\")\n        print(\"❌ ERROR\")\n        print(f\"{'='*80}\")\n        print(f\"\\nError: {str(e)}\")\n        \n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T03:53:37.233495Z","iopub.execute_input":"2026-01-17T03:53:37.233727Z","iopub.status.idle":"2026-01-17T03:58:17.135652Z","shell.execute_reply.started":"2026-01-17T03:53:37.233707Z","shell.execute_reply":"2026-01-17T03:58:17.13471Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nSTEP 4: TRUE FEW-SHOT LEARNING\n===============================\nImplements GENUINE few-shot learning using Prototypical Networks\nPerfect for small medical datasets (10-300 patients)\n\nWhat makes this REAL few-shot learning:\n✅ Episodic training (support/query sets)\n✅ Prototypical Networks (distance-based classification)\n✅ Meta-learning (learns to learn from few examples)\n✅ Can adapt to new patients with just 1-5 examples\n\nAfter this: Step 5 (Inference & Deployment)\n\"\"\"\n\nimport os\nimport gc\nimport time\nimport json\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nimport torch.optim as optim\nfrom torch.utils.data import Dataset\n\nfrom sklearn.metrics import roc_auc_score, accuracy_score, confusion_matrix\nfrom sklearn.model_selection import train_test_split\n\nfrom tqdm import tqdm\nfrom scipy.ndimage import rotate\n\nprint(\"=\"*80)\nprint(\"🧠 STEP 4: TRUE FEW-SHOT LEARNING\")\nprint(\"=\"*80)\n\n# ============================================================================\n# FEW-SHOT CONFIGURATION\n# ============================================================================\n\nclass FewShotConfig:\n    \"\"\"Configuration for few-shot learning\"\"\"\n    \n    def __init__(self, n_patients):\n        self.n_patients = n_patients\n        \n        # Few-shot parameters (THESE DEFINE FEW-SHOT LEARNING!)\n        self.n_way = 2              # Binary: fracture vs normal\n        self.k_shot = 2             # 2 support examples per class\n        self.n_query = 1            # 1 query example per class\n        \n        # Training\n        self.n_episodes = 50        # Number of training episodes\n        self.n_val_episodes = 20    # Validation episodes\n        \n        # Model\n        self.embedding_dim = 64     # Feature embedding dimension\n        self.dropout = 0.5          # High dropout for small data\n        self.learning_rate = 5e-4\n        self.weight_decay = 1e-3\n        \n        # Data split\n        self.train_ratio = 0.7      # 70% train, 30% val\n        \n        # Paths\n        self.dataset_dir = '/kaggle/working/mini_dataset_v2'\n        self.save_dir = '/kaggle/working/step4_few_shot'\n        \n        os.makedirs(self.save_dir, exist_ok=True)\n    \n    def print_config(self):\n        \"\"\"Print configuration\"\"\"\n        print(f\"\\n📋 Few-Shot Learning Configuration:\")\n        print(f\"  Dataset: {self.n_patients} patients\")\n        print(f\"  Setup: {self.n_way}-way {self.k_shot}-shot learning\")\n        print(f\"  Query: {self.n_query} per class\")\n        print(f\"  Episodes: {self.n_episodes} train, {self.n_val_episodes} val\")\n        print(f\"  Embedding: {self.embedding_dim}D\")\n        print(f\"  Train/Val: {self.train_ratio*100:.0f}%/{(1-self.train_ratio)*100:.0f}%\")\n        print(f\"\\n  ⭐ This is REAL few-shot learning:\")\n        print(f\"     - Episodic training (support/query)\")\n        print(f\"     - Distance-based classification\")\n        print(f\"     - Meta-learning approach\")\n\n\nCONFIG = FewShotConfig(10)  # Will be updated with actual patient count\n\n# ============================================================================\n# FEW-SHOT DATASET\n# ============================================================================\n\nclass FewShotDataset(Dataset):\n    \"\"\"Dataset for episodic few-shot learning\"\"\"\n    \n    def __init__(self, metadata_df, volumes_dir, split='train', train_ratio=0.7):\n        \"\"\"\n        Args:\n            split: 'train' or 'val'\n            train_ratio: Fraction for training\n        \"\"\"\n        self.volumes_dir = volumes_dir\n        \n        # Split by class first\n        fracture_df = metadata_df[metadata_df['has_fracture'] == 1].reset_index(drop=True)\n        normal_df = metadata_df[metadata_df['has_fracture'] == 0].reset_index(drop=True)\n        \n        # Split each class\n        n_train_frac = max(1, int(len(fracture_df) * train_ratio))\n        n_train_norm = max(1, int(len(normal_df) * train_ratio))\n        \n        if split == 'train':\n            fracture_split = fracture_df.iloc[:n_train_frac]\n            normal_split = normal_df.iloc[:n_train_norm]\n        else:  # val\n            fracture_split = fracture_df.iloc[n_train_frac:]\n            normal_split = normal_df.iloc[n_train_norm:]\n        \n        self.metadata_df = pd.concat([fracture_split, normal_split], ignore_index=True)\n        \n        # Class indices\n        self.fracture_indices = self.metadata_df[\n            self.metadata_df['has_fracture'] == 1\n        ].index.tolist()\n        \n        self.normal_indices = self.metadata_df[\n            self.metadata_df['has_fracture'] == 0\n        ].index.tolist()\n        \n        print(f\"\\n  {split.upper()} split:\")\n        print(f\"    Fracture: {len(self.fracture_indices)} patients\")\n        print(f\"    Normal: {len(self.normal_indices)} patients\")\n    \n    def __len__(self):\n        return len(self.metadata_df)\n    \n    def __getitem__(self, idx):\n        \"\"\"Load single volume\"\"\"\n        patient_id = self.metadata_df.iloc[idx]['patient_id']\n        label = self.metadata_df.iloc[idx]['has_fracture']\n        \n        volume_path = os.path.join(self.volumes_dir, f\"{patient_id}.npy\")\n        volume = np.load(volume_path)\n        \n        # Simple augmentation\n        if np.random.random() < 0.3:\n            if np.random.random() < 0.5:\n                volume = np.flip(volume, axis=2).copy()\n            if np.random.random() < 0.3:\n                factor = np.random.uniform(0.95, 1.05)\n                volume = np.clip(volume * factor, 0, 1)\n        \n        volume = volume[np.newaxis, ...].astype(np.float32)\n        \n        return torch.from_numpy(volume).float(), torch.tensor(label).long()\n    \n    def sample_episode(self, n_way, k_shot, n_query):\n        \"\"\"\n        Sample a few-shot episode (THIS IS THE KEY!)\n        \n        Returns:\n            support_data: (n_way * k_shot, 1, D, H, W) - for learning prototypes\n            support_labels: (n_way * k_shot,) - class labels\n            query_data: (n_way * n_query, 1, D, H, W) - for testing\n            query_labels: (n_way * n_query,) - class labels\n        \"\"\"\n        support_data = []\n        support_labels = []\n        query_data = []\n        query_labels = []\n        \n        # For each class (0: normal, 1: fracture)\n        for class_idx, indices in enumerate([self.normal_indices, self.fracture_indices]):\n            total_needed = k_shot + n_query\n            \n            # Sample without replacement if possible\n            if len(indices) >= total_needed:\n                sampled = np.random.choice(indices, size=total_needed, replace=False)\n            else:\n                sampled = np.random.choice(indices, size=total_needed, replace=True)\n            \n            # Split into support and query\n            support_idx = sampled[:k_shot]\n            query_idx = sampled[k_shot:]\n            \n            # Load support examples\n            for idx in support_idx:\n                volume, _ = self[idx]\n                support_data.append(volume)\n                support_labels.append(class_idx)\n            \n            # Load query examples\n            for idx in query_idx:\n                volume, _ = self[idx]\n                query_data.append(volume)\n                query_labels.append(class_idx)\n        \n        # Stack\n        support_data = torch.stack(support_data)\n        support_labels = torch.tensor(support_labels)\n        query_data = torch.stack(query_data)\n        query_labels = torch.tensor(query_labels)\n        \n        return support_data, support_labels, query_data, query_labels\n\n\n# ============================================================================\n# EMBEDDING NETWORK\n# ============================================================================\n\nclass EmbeddingNetwork3D(nn.Module):\n    \"\"\"\n    Lightweight 3D CNN for embedding volumes into metric space\n    This is the CORE of few-shot learning!\n    \"\"\"\n    \n    def __init__(self, embedding_dim=64, dropout=0.5):\n        super().__init__()\n        \n        # Encoder (small to prevent overfitting)\n        self.conv1 = nn.Conv3d(1, 16, kernel_size=3, stride=2, padding=1)\n        self.bn1 = nn.BatchNorm3d(16)\n        self.drop1 = nn.Dropout3d(dropout * 0.3)\n        \n        self.conv2 = nn.Conv3d(16, 32, kernel_size=3, stride=2, padding=1)\n        self.bn2 = nn.BatchNorm3d(32)\n        self.drop2 = nn.Dropout3d(dropout * 0.5)\n        \n        self.conv3 = nn.Conv3d(32, 64, kernel_size=3, stride=2, padding=1)\n        self.bn3 = nn.BatchNorm3d(64)\n        self.drop3 = nn.Dropout3d(dropout * 0.7)\n        \n        # Global pooling\n        self.pool = nn.AdaptiveAvgPool3d(1)\n        \n        # Embedding layer\n        self.dropout = nn.Dropout(dropout)\n        self.fc = nn.Linear(64, embedding_dim)\n    \n    def forward(self, x):\n        \"\"\"\n        Args:\n            x: (batch, 1, D, H, W)\n        Returns:\n            embeddings: (batch, embedding_dim) - L2 normalized\n        \"\"\"\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = self.drop1(x)\n        \n        x = F.relu(self.bn2(self.conv2(x)))\n        x = self.drop2(x)\n        \n        x = F.relu(self.bn3(self.conv3(x)))\n        x = self.drop3(x)\n        \n        x = self.pool(x)\n        x = x.view(x.size(0), -1)\n        \n        x = self.dropout(x)\n        embeddings = self.fc(x)\n        \n        # L2 normalize (important for distance-based classification!)\n        embeddings = F.normalize(embeddings, p=2, dim=1)\n        \n        return embeddings\n\n\n# ============================================================================\n# PROTOTYPICAL NETWORK\n# ============================================================================\n\nclass PrototypicalNetwork:\n    \"\"\"\n    Prototypical Networks for Few-Shot Learning\n    \n    Key idea: Learn embeddings where classification is done by\n    finding the nearest class prototype (mean of support embeddings)\n    \"\"\"\n    \n    def __init__(self, embedding_net, config, device='cuda'):\n        self.embedding_net = embedding_net.to(device)\n        self.config = config\n        self.device = device\n        \n        self.optimizer = torch.optim.AdamW(\n            self.embedding_net.parameters(),\n            lr=config.learning_rate,\n            weight_decay=config.weight_decay\n        )\n        \n        self.best_val_acc = 0.0\n        self.best_model_state = None\n    \n    def compute_prototypes(self, support_embeddings, support_labels, n_way):\n        \"\"\"\n        Compute class prototypes (mean of support embeddings)\n        \n        This is THE KEY to few-shot learning!\n        \"\"\"\n        prototypes = []\n        \n        for class_idx in range(n_way):\n            # Get all support examples for this class\n            class_mask = support_labels == class_idx\n            class_embeddings = support_embeddings[class_mask]\n            \n            # Prototype = mean of class embeddings\n            prototype = class_embeddings.mean(dim=0)\n            prototypes.append(prototype)\n        \n        prototypes = torch.stack(prototypes)\n        return prototypes\n    \n    def euclidean_distance(self, x, y):\n        \"\"\"\n        Compute pairwise Euclidean distances\n        \n        Args:\n            x: (n, d) query embeddings\n            y: (m, d) prototype embeddings\n        \n        Returns:\n            distances: (n, m)\n        \"\"\"\n        n = x.size(0)\n        m = y.size(0)\n        d = x.size(1)\n        \n        x = x.unsqueeze(1).expand(n, m, d)\n        y = y.unsqueeze(0).expand(n, m, d)\n        \n        distances = torch.pow(x - y, 2).sum(2)\n        return distances\n    \n    def train_episode(self, support_data, support_labels, query_data, query_labels, n_way):\n        \"\"\"Train on one episode\"\"\"\n        self.embedding_net.train()\n        \n        # Move to device\n        support_data = support_data.to(self.device)\n        support_labels = support_labels.to(self.device)\n        query_data = query_data.to(self.device)\n        query_labels = query_labels.to(self.device)\n        \n        # Embed support and query\n        support_embeddings = self.embedding_net(support_data)\n        query_embeddings = self.embedding_net(query_data)\n        \n        # Compute prototypes\n        prototypes = self.compute_prototypes(support_embeddings, support_labels, n_way)\n        \n        # Compute distances from queries to prototypes\n        distances = self.euclidean_distance(query_embeddings, prototypes)\n        \n        # Convert to logits (negative distance = higher probability)\n        logits = -distances\n        \n        # Loss\n        loss = F.cross_entropy(logits, query_labels)\n        \n        # Backprop\n        self.optimizer.zero_grad()\n        loss.backward()\n        torch.nn.utils.clip_grad_norm_(self.embedding_net.parameters(), 1.0)\n        self.optimizer.step()\n        \n        # Accuracy\n        predictions = torch.argmax(logits, dim=1)\n        accuracy = (predictions == query_labels).float().mean()\n        \n        return loss.item(), accuracy.item()\n    \n    def evaluate_episode(self, support_data, support_labels, query_data, query_labels, n_way):\n        \"\"\"Evaluate on one episode\"\"\"\n        self.embedding_net.eval()\n        \n        with torch.no_grad():\n            support_data = support_data.to(self.device)\n            support_labels = support_labels.to(self.device)\n            query_data = query_data.to(self.device)\n            query_labels = query_labels.to(self.device)\n            \n            support_embeddings = self.embedding_net(support_data)\n            query_embeddings = self.embedding_net(query_data)\n            \n            prototypes = self.compute_prototypes(support_embeddings, support_labels, n_way)\n            distances = self.euclidean_distance(query_embeddings, prototypes)\n            logits = -distances\n            \n            loss = F.cross_entropy(logits, query_labels)\n            predictions = torch.argmax(logits, dim=1)\n            accuracy = (predictions == query_labels).float().mean()\n        \n        return loss.item(), accuracy.item()\n\n\n# ============================================================================\n# TRAINING FUNCTION\n# ============================================================================\n\ndef train_few_shot_model(train_dataset, val_dataset, config):\n    \"\"\"Train prototypical network\"\"\"\n    \n    print(f\"\\n{'='*80}\")\n    print(\"🏋️  TRAINING FEW-SHOT MODEL\")\n    print(f\"{'='*80}\")\n    \n    device = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\n    print(f\"\\n  Device: {device}\")\n    \n    # Create model\n    embedding_net = EmbeddingNetwork3D(\n        embedding_dim=config.embedding_dim,\n        dropout=config.dropout\n    )\n    \n    proto_net = PrototypicalNetwork(embedding_net, config, device)\n    \n    n_params = sum(p.numel() for p in embedding_net.parameters() if p.requires_grad)\n    print(f\"  Parameters: {n_params:,}\")\n    \n    # Training history\n    history = {\n        'train_loss': [], 'train_acc': [],\n        'val_loss': [], 'val_acc': []\n    }\n    \n    print(f\"\\n  Training {config.n_episodes} episodes...\")\n    \n    # Training loop\n    pbar = tqdm(range(config.n_episodes), desc=\"Training\")\n    \n    for episode in pbar:\n        # Sample training episode\n        support_data, support_labels, query_data, query_labels = \\\n            train_dataset.sample_episode(config.n_way, config.k_shot, config.n_query)\n        \n        # Train\n        loss, acc = proto_net.train_episode(\n            support_data, support_labels, query_data, query_labels, config.n_way\n        )\n        \n        history['train_loss'].append(loss)\n        history['train_acc'].append(acc)\n        \n        # Validate every 5 episodes\n        if (episode + 1) % 5 == 0:\n            val_losses = []\n            val_accs = []\n            \n            for _ in range(5):  # Sample 5 val episodes\n                val_sup, val_sup_lbl, val_qry, val_qry_lbl = \\\n                    val_dataset.sample_episode(config.n_way, config.k_shot, config.n_query)\n                \n                val_loss, val_acc = proto_net.evaluate_episode(\n                    val_sup, val_sup_lbl, val_qry, val_qry_lbl, config.n_way\n                )\n                \n                val_losses.append(val_loss)\n                val_accs.append(val_acc)\n            \n            avg_val_loss = np.mean(val_losses)\n            avg_val_acc = np.mean(val_accs)\n            \n            history['val_loss'].append(avg_val_loss)\n            history['val_acc'].append(avg_val_acc)\n            \n            pbar.set_postfix({\n                'train_acc': f'{acc*100:.1f}%',\n                'val_acc': f'{avg_val_acc*100:.1f}%'\n            })\n            \n            # Save best\n            if avg_val_acc > proto_net.best_val_acc:\n                proto_net.best_val_acc = avg_val_acc\n                proto_net.best_model_state = embedding_net.state_dict().copy()\n    \n    # Restore best\n    if proto_net.best_model_state is not None:\n        embedding_net.load_state_dict(proto_net.best_model_state)\n        print(f\"\\n  ✓ Restored best model (val acc: {proto_net.best_val_acc*100:.1f}%)\")\n    \n    # Save model\n    model_path = os.path.join(config.save_dir, 'few_shot_model.pth')\n    torch.save({\n        'model_state_dict': embedding_net.state_dict(),\n        'config': config.__dict__,\n        'best_val_acc': proto_net.best_val_acc,\n        'history': history\n    }, model_path)\n    \n    print(f\"\\n  💾 Saved: {model_path}\")\n    \n    return proto_net, history\n\n\n# ============================================================================\n# VISUALIZATION\n# ============================================================================\n\ndef plot_results(history, config):\n    \"\"\"Plot training curves\"\"\"\n    \n    fig, axes = plt.subplots(1, 2, figsize=(14, 4))\n    \n    episodes = range(1, len(history['train_loss']) + 1)\n    \n    # Loss\n    axes[0].plot(episodes, history['train_loss'], 'b-', label='Train', alpha=0.7)\n    if history['val_loss']:\n        val_episodes = np.arange(4, len(history['train_loss']), 5)[:len(history['val_loss'])]\n        axes[0].plot(val_episodes, history['val_loss'], 'r-', label='Val', marker='o')\n    axes[0].set_xlabel('Episode')\n    axes[0].set_ylabel('Loss')\n    axes[0].set_title('Training Loss', fontweight='bold')\n    axes[0].legend()\n    axes[0].grid(True, alpha=0.3)\n    \n    # Accuracy\n    axes[1].plot(np.array(history['train_acc']) * 100, 'b-', label='Train', alpha=0.7)\n    if history['val_acc']:\n        val_episodes = np.arange(4, len(history['train_acc']), 5)[:len(history['val_acc'])]\n        axes[1].plot(val_episodes, np.array(history['val_acc']) * 100, 'r-', label='Val', marker='o')\n    axes[1].axhline(y=50, color='gray', linestyle='--', alpha=0.5, label='Random')\n    axes[1].set_xlabel('Episode')\n    axes[1].set_ylabel('Accuracy (%)')\n    axes[1].set_title('Training Accuracy', fontweight='bold')\n    axes[1].legend()\n    axes[1].grid(True, alpha=0.3)\n    \n    plt.tight_layout()\n    save_path = os.path.join(config.save_dir, 'training_curves.png')\n    plt.savefig(save_path, dpi=150)\n    plt.close()\n    \n    print(f\"  ✓ Saved: {save_path}\")\n\n\n# ============================================================================\n# MAIN FUNCTION\n# ============================================================================\n\ndef run_few_shot_learning():\n    \"\"\"Main execution\"\"\"\n    \n    print(f\"\\n{'='*80}\")\n    print(\"📂 LOADING DATASET\")\n    print(f\"{'='*80}\")\n    \n    metadata_path = os.path.join(CONFIG.dataset_dir, 'metadata.csv')\n    volumes_dir = os.path.join(CONFIG.dataset_dir, 'volumes')\n    \n    if not os.path.exists(metadata_path):\n        print(f\"\\n❌ ERROR: {metadata_path} not found!\")\n        print(f\"   Run Step 1 first!\")\n        return None\n    \n    metadata_df = pd.read_csv(metadata_path)\n    n_patients = len(metadata_df)\n    \n    # Update config with actual patient count\n    CONFIG.n_patients = n_patients\n    CONFIG.print_config()\n    \n    # Check minimum requirements\n    min_needed = CONFIG.k_shot + CONFIG.n_query\n    n_fracture = metadata_df['has_fracture'].sum()\n    n_normal = len(metadata_df) - n_fracture\n    \n    if n_fracture < min_needed * 2 or n_normal < min_needed * 2:\n        print(f\"\\n⚠️  WARNING: Not enough data!\")\n        print(f\"   Need ≥{min_needed * 2} per class\")\n        print(f\"   Have: {n_fracture} fracture, {n_normal} normal\")\n        print(f\"\\n   Run Step 1 with more patients or reduce k_shot/n_query\")\n        return None\n    \n    # Create datasets\n    train_dataset = FewShotDataset(\n        metadata_df, volumes_dir,\n        split='train', train_ratio=CONFIG.train_ratio\n    )\n    \n    val_dataset = FewShotDataset(\n        metadata_df, volumes_dir,\n        split='val', train_ratio=CONFIG.train_ratio\n    )\n    \n    # Train\n    proto_net, history = train_few_shot_model(train_dataset, val_dataset, CONFIG)\n    \n    # Final evaluation\n    print(f\"\\n{'='*80}\")\n    print(\"📈 FINAL EVALUATION\")\n    print(f\"{'='*80}\")\n    \n    final_accs = []\n    for _ in tqdm(range(CONFIG.n_val_episodes), desc=\"Evaluating\"):\n        val_sup, val_sup_lbl, val_qry, val_qry_lbl = \\\n            val_dataset.sample_episode(CONFIG.n_way, CONFIG.k_shot, CONFIG.n_query)\n        \n        _, acc = proto_net.evaluate_episode(\n            val_sup, val_sup_lbl, val_qry, val_qry_lbl, CONFIG.n_way\n        )\n        final_accs.append(acc)\n    \n    mean_acc = np.mean(final_accs)\n    std_acc = np.std(final_accs)\n    \n    print(f\"\\n  Validation Accuracy: {mean_acc*100:.1f}% ± {std_acc*100:.1f}%\")\n    \n    # Visualize\n    print(f\"\\n{'='*80}\")\n    print(\"📊 CREATING VISUALIZATIONS\")\n    print(f\"{'='*80}\\n\")\n    \n    plot_results(history, CONFIG)\n    \n    # Summary\n    print(f\"\\n{'='*80}\")\n    print(\"✅ FEW-SHOT LEARNING COMPLETE!\")\n    print(f\"{'='*80}\")\n    \n    print(f\"\\n  📊 Results:\")\n    print(f\"    Final val accuracy: {mean_acc*100:.1f}% ± {std_acc*100:.1f}%\")\n    print(f\"    Best val accuracy: {proto_net.best_val_acc*100:.1f}%\")\n    \n    # Interpret\n    print(f\"\\n  💡 Interpretation:\")\n    if mean_acc > 0.70:\n        print(f\"    ✅ Excellent! Model learns well from {CONFIG.k_shot} examples!\")\n    elif mean_acc > 0.60:\n        print(f\"    ✓ Good performance for {n_patients} patients\")\n    elif mean_acc > 0.55:\n        print(f\"    ⚠️  Moderate - consider more patients or data augmentation\")\n    else:\n        print(f\"    ⚠️  Limited - need more patients (50+ recommended)\")\n    \n    gap = history['train_acc'][-1] - history['val_acc'][-1] if history['val_acc'] else 0\n    if gap > 0.15:\n        print(f\"    ⚠️  Overfitting detected ({gap*100:.1f}% gap)\")\n        print(f\"       Solution: Collect more patients (Step 1)\")\n    \n    # Save summary\n    summary = {\n        'n_patients': n_patients,\n        'n_way': CONFIG.n_way,\n        'k_shot': CONFIG.k_shot,\n        'n_query': CONFIG.n_query,\n        'final_accuracy': float(mean_acc),\n        'accuracy_std': float(std_acc),\n        'best_val_accuracy': float(proto_net.best_val_acc)\n    }\n    \n    with open(os.path.join(CONFIG.save_dir, 'summary.json'), 'w') as f:\n        json.dump(summary, f, indent=2)\n    \n    print(f\"\\n  💾 Saved to: {CONFIG.save_dir}/\")\n    print(f\"     - few_shot_model.pth\")\n    print(f\"     - training_curves.png\")\n    print(f\"     - summary.json\")\n    \n    return proto_net, summary\n\n\n# ============================================================================\n# MAIN EXECUTION\n# ============================================================================\n\nif __name__ == \"__main__\":\n    print(\"\\n\" + \"=\"*80)\n    print(\"🚀 RUNNING TRUE FEW-SHOT LEARNING\")\n    print(\"=\"*80)\n    \n    try:\n        result = run_few_shot_learning()\n        \n        if result is None:\n            print(\"\\n❌ Failed - check dataset or requirements\")\n        else:\n            proto_net, summary = result\n            \n            print(f\"\\n{'='*80}\")\n            print(\"🎯 NEXT STEP: INFERENCE & DEPLOYMENT\")\n            print(f\"{'='*80}\")\n            \n            print(f\"\\n  This few-shot model can now:\")\n            print(f\"  ✅ Classify new patients with just {summary['k_shot']} examples\")\n            print(f\"  ✅ Adapt to new data without retraining\")\n            print(f\"  ✅ Work in low-data scenarios\")\n            \n            print(f\"\\n  📝 Next Steps:\")\n            print(f\"  1. STEP 5: Create inference pipeline\")\n            print(f\"     - Load trained model\")\n            print(f\"     - Process new patient DICOMs\")\n            print(f\"     - Generate predictions\")\n            print(f\"     - Create GradCAM visualizations\")\n            \n            print(f\"\\n  2. Advanced:\")\n            print(f\"     - Ensemble with standard model\")\n            print(f\"     - Test-Time Augmentation (TTA)\")\n            print(f\"     - Cross-validation\")\n            print(f\"     - External validation set\")\n            \n            print(f\"\\n  3. Deployment:\")\n            print(f\"     - Export to ONNX format\")\n            print(f\"     - Create REST API\")\n            print(f\"     - Build web interface\")\n            print(f\"     - Clinical integration\")\n            \n            print(f\"\\n✨ You've completed the core training pipeline!\")\n            print(f\"   Ready for Step 5: Inference & Deployment\")\n    \n    except KeyboardInterrupt:\n        print(\"\\n\\n⚠️  Interrupted by user\")\n    \n    except Exception as e:\n        print(f\"\\n{'='*80}\")\n        print(\"❌ ERROR\")\n        print(f\"{'='*80}\")\n        print(f\"\\nError: {str(e)}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T03:58:17.139809Z","iopub.execute_input":"2026-01-17T03:58:17.140104Z","iopub.status.idle":"2026-01-17T03:58:32.654664Z","execution_failed":"2026-01-17T05:30:29.161Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nSTEP 3.5: ADVANCED ENSEMBLE METHODS (FIXED)\n============================================\nCombines EfficientNet (Step 3) + Few-Shot Learning (Step 4)\nfor optimal performance on medical imaging tasks\n\nFIXED: Uses correct model architectures from Step 3 and 4\n\nEnsemble Strategies Implemented:\n✅ 1. Weighted Average (simple, effective)\n✅ 2. Learned Weights (meta-learner via Logistic Regression)\n✅ 3. Confidence-Based Fusion (adaptive weighting)\n✅ 4. Stacking Ensemble (2-layer)\n✅ 5. Test-Time Augmentation (TTA)\n\nExpected improvement: +3-8% AUC over single models\n\"\"\"\n\nimport os\nimport json\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.metrics import (\n    roc_auc_score, accuracy_score, roc_curve,\n    precision_recall_curve, average_precision_score,\n    confusion_matrix\n)\nfrom sklearn.linear_model import LogisticRegression\n\nfrom tqdm import tqdm\n\nprint(\"=\"*80)\nprint(\"🎯 STEP 3.5: ENSEMBLE METHODS (FIXED)\")\nprint(\"=\"*80)\n\n# ============================================================================\n# CONFIGURATION\n# ============================================================================\n\nCONFIG = {\n    # Model checkpoints\n    'efficientnet_path': '/kaggle/working/efficientnet_300patients/best_model.pth',\n    'fewshot_path': '/kaggle/working/step4_few_shot/few_shot_model.pth',\n    \n    # Dataset\n    'dataset_dir': '/kaggle/working/mini_dataset_v2',\n    \n    # Output\n    'output_dir': '/kaggle/working/ensemble_results',\n    \n    # Few-shot support set\n    'n_support_per_class': 3,\n    \n    # TTA\n    'tta_augmentations': 5,\n    \n    # Device\n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n}\n\nos.makedirs(CONFIG['output_dir'], exist_ok=True)\n\n# ============================================================================\n# CORRECT MODEL ARCHITECTURES (FROM STEP 3 AND 4)\n# ============================================================================\n\nclass SqueezeExcitation3D(nn.Module):\n    \"\"\"From Step 3\"\"\"\n    def __init__(self, channels, reduction=4):\n        super().__init__()\n        reduced = max(1, channels // reduction)\n        self.se = nn.Sequential(\n            nn.AdaptiveAvgPool3d(1),\n            nn.Conv3d(channels, reduced, 1),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(reduced, channels, 1),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        return x * self.se(x)\n\n\nclass MBConvBlock(nn.Module):\n    \"\"\"From Step 3\"\"\"\n    def __init__(self, in_channels, out_channels, expand_ratio=6, stride=1):\n        super().__init__()\n        self.use_residual = (stride == 1 and in_channels == out_channels)\n        hidden = in_channels * expand_ratio\n        \n        layers = []\n        \n        if expand_ratio != 1:\n            layers.extend([\n                nn.Conv3d(in_channels, hidden, 1, bias=False),\n                nn.BatchNorm3d(hidden),\n                nn.ReLU(inplace=True)\n            ])\n        \n        layers.extend([\n            nn.Conv3d(hidden, hidden, 3, stride=stride, padding=1, \n                     groups=hidden, bias=False),\n            nn.BatchNorm3d(hidden),\n            nn.ReLU(inplace=True)\n        ])\n        \n        layers.append(SqueezeExcitation3D(hidden))\n        \n        layers.extend([\n            nn.Conv3d(hidden, out_channels, 1, bias=False),\n            nn.BatchNorm3d(out_channels)\n        ])\n        \n        self.conv = nn.Sequential(*layers)\n        self.dropout = nn.Dropout3d(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.conv(x))\n        else:\n            return self.conv(x)\n\n\nclass LightweightEfficientNet3D(nn.Module):\n    \"\"\"CORRECT architecture from Step 3\"\"\"\n    def __init__(self, num_classes=1, dropout=0.4):\n        super().__init__()\n        \n        self.stem = nn.Sequential(\n            nn.Conv3d(1, 32, kernel_size=3, stride=2, padding=1, bias=False),\n            nn.BatchNorm3d(32),\n            nn.ReLU(inplace=True)\n        )\n        \n        self.blocks = nn.Sequential(\n            MBConvBlock(32, 16, expand_ratio=1, stride=1),\n            MBConvBlock(16, 24, expand_ratio=6, stride=2),\n            MBConvBlock(24, 24, expand_ratio=6, stride=1),\n            MBConvBlock(24, 40, expand_ratio=6, stride=2),\n            MBConvBlock(40, 40, expand_ratio=6, stride=1),\n            MBConvBlock(40, 80, expand_ratio=6, stride=2),\n            MBConvBlock(80, 80, expand_ratio=6, stride=1),\n        )\n        \n        self.head = nn.Sequential(\n            nn.Conv3d(80, 320, 1, bias=False),\n            nn.BatchNorm3d(320),\n            nn.ReLU(inplace=True),\n            nn.AdaptiveAvgPool3d(1)\n        )\n        \n        self.dropout = nn.Dropout(dropout)\n        self.fc = nn.Linear(320, num_classes)\n    \n    def forward(self, x):\n        x = self.stem(x)\n        x = self.blocks(x)\n        x = self.head(x)\n        x = x.view(x.size(0), -1)\n        x = self.dropout(x)\n        x = self.fc(x)\n        return x.squeeze(-1)\n\n\nclass EmbeddingNetwork3D(nn.Module):\n    \"\"\"CORRECT architecture from Step 4\"\"\"\n    def __init__(self, embedding_dim=64, dropout=0.5):\n        super().__init__()\n        \n        self.conv1 = nn.Conv3d(1, 16, kernel_size=3, stride=2, padding=1)\n        self.bn1 = nn.BatchNorm3d(16)\n        self.drop1 = nn.Dropout3d(dropout * 0.3)\n        \n        self.conv2 = nn.Conv3d(16, 32, kernel_size=3, stride=2, padding=1)\n        self.bn2 = nn.BatchNorm3d(32)\n        self.drop2 = nn.Dropout3d(dropout * 0.5)\n        \n        self.conv3 = nn.Conv3d(32, 64, kernel_size=3, stride=2, padding=1)\n        self.bn3 = nn.BatchNorm3d(64)\n        self.drop3 = nn.Dropout3d(dropout * 0.7)\n        \n        self.pool = nn.AdaptiveAvgPool3d(1)\n        self.dropout = nn.Dropout(dropout)\n        self.fc = nn.Linear(64, embedding_dim)\n    \n    def forward(self, x):\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = self.drop1(x)\n        \n        x = F.relu(self.bn2(self.conv2(x)))\n        x = self.drop2(x)\n        \n        x = F.relu(self.bn3(self.conv3(x)))\n        x = self.drop3(x)\n        \n        x = self.pool(x)\n        x = x.view(x.size(0), -1)\n        \n        x = self.dropout(x)\n        embeddings = self.fc(x)\n        \n        # L2 normalize\n        embeddings = F.normalize(embeddings, p=2, dim=1)\n        \n        return embeddings\n\n\n# ============================================================================\n# ENSEMBLE ENGINE\n# ============================================================================\n\nclass EnsembleEngine:\n    \"\"\"Manages multiple ensemble strategies\"\"\"\n    \n    def __init__(self, efficientnet, fewshot_net, device):\n        self.efficientnet = efficientnet\n        self.fewshot_net = fewshot_net\n        self.device = device\n        self.support_prototypes = None\n    \n    def set_support_prototypes(self, support_volumes, support_labels):\n        \"\"\"Compute prototypes for few-shot inference\"\"\"\n        with torch.no_grad():\n            support_volumes = support_volumes.to(self.device)\n            embeddings = self.fewshot_net(support_volumes)\n            \n            prototypes = []\n            for class_idx in [0, 1]:\n                mask = support_labels == class_idx\n                class_emb = embeddings[mask]\n                prototype = class_emb.mean(dim=0) if len(class_emb) > 0 else torch.zeros(embeddings.shape[1], device=self.device)\n                prototypes.append(prototype)\n            \n            self.support_prototypes = torch.stack(prototypes)\n        \n        print(f\"  ✓ Support prototypes computed: {self.support_prototypes.shape}\")\n    \n    def predict_efficientnet(self, volume):\n        \"\"\"EfficientNet predictions\"\"\"\n        with torch.no_grad():\n            logits = self.efficientnet(volume)\n            probs = torch.sigmoid(logits).cpu().numpy()\n        return probs\n    \n    def predict_fewshot(self, volume):\n        \"\"\"Few-shot predictions via prototype matching\"\"\"\n        with torch.no_grad():\n            embeddings = self.fewshot_net(volume)\n            distances = torch.cdist(embeddings, self.support_prototypes.unsqueeze(0)).squeeze(0)\n            logits = -distances\n            probs = F.softmax(logits, dim=1)[:, 1].cpu().numpy()\n        return probs\n    \n    def weighted_average(self, eff_preds, few_preds, alpha=0.7):\n        \"\"\"Simple weighted combination\"\"\"\n        return alpha * eff_preds + (1 - alpha) * few_preds\n    \n    def learn_weights(self, eff_preds, few_preds, labels):\n        \"\"\"Train meta-learner (Logistic Regression)\"\"\"\n        X = np.column_stack([eff_preds, few_preds])\n        lr = LogisticRegression(random_state=42, max_iter=1000)\n        lr.fit(X, labels)\n        \n        print(f\"\\n  Learned coefficients:\")\n        print(f\"    EfficientNet: {lr.coef_[0][0]:.4f}\")\n        print(f\"    Few-Shot: {lr.coef_[0][1]:.4f}\")\n        \n        return lr\n    \n    def confidence_fusion(self, eff_preds, few_preds):\n        \"\"\"Adaptive weighting based on confidence\"\"\"\n        eff_conf = np.abs(eff_preds - 0.5)\n        few_conf = np.abs(few_preds - 0.5)\n        \n        total = eff_conf + few_conf + 1e-8\n        eff_weight = eff_conf / total\n        few_weight = few_conf / total\n        \n        return eff_weight * eff_preds + few_weight * few_preds\n    \n    def predict_with_tta(self, volume, n_aug=5):\n        \"\"\"Test-Time Augmentation for EfficientNet\"\"\"\n        predictions = []\n        \n        # Original\n        pred = self.predict_efficientnet(volume)\n        predictions.append(pred)\n        \n        # Augmented versions\n        for _ in range(n_aug - 1):\n            volume_aug = self._augment(volume.clone())\n            pred = self.predict_efficientnet(volume_aug)\n            predictions.append(pred)\n        \n        return np.mean(predictions, axis=0), np.std(predictions, axis=0)\n    \n    def _augment(self, volume):\n        \"\"\"Random augmentation\"\"\"\n        # Horizontal flip\n        if np.random.random() < 0.5:\n            volume = torch.flip(volume, dims=[4])\n        \n        # Brightness\n        if np.random.random() < 0.3:\n            factor = np.random.uniform(0.95, 1.05)\n            volume = torch.clamp(volume * factor, 0, 1)\n        \n        return volume\n\n\n# ============================================================================\n# EVALUATION\n# ============================================================================\n\ndef evaluate_all_strategies(engine, loader):\n    \"\"\"Comprehensive evaluation\"\"\"\n    \n    print(f\"\\n{'='*80}\")\n    print(\"🔍 COLLECTING PREDICTIONS\")\n    print(f\"{'='*80}\\n\")\n    \n    all_labels = []\n    all_eff_preds = []\n    all_few_preds = []\n    \n    import time\n    start_time = time.time()\n    batch_times = []\n    \n    for batch_idx, (volumes, labels) in enumerate(tqdm(loader, desc=\"Predicting\")):\n        batch_start = time.time()\n        \n        volumes = volumes.to(engine.device)\n        \n        eff_preds = engine.predict_efficientnet(volumes)\n        few_preds = engine.predict_fewshot(volumes)\n        \n        all_labels.extend(labels.numpy())\n        all_eff_preds.extend(eff_preds)\n        all_few_preds.extend(few_preds)\n        \n        batch_time = time.time() - batch_start\n        batch_times.append(batch_time)\n        \n        # Show first few predictions for verification\n        if batch_idx == 0:\n            print(f\"\\n  📊 Sample predictions (first batch):\")\n            print(f\"     EfficientNet: {eff_preds[:3]}\")\n            print(f\"     Few-Shot:     {few_preds[:3]}\")\n            print(f\"     Labels:       {labels[:3].numpy()}\")\n    \n    total_time = time.time() - start_time\n    \n    print(f\"\\n  ⏱️  Timing Statistics:\")\n    print(f\"     Total time: {total_time:.2f}s\")\n    print(f\"     Avg batch time: {np.mean(batch_times)*1000:.1f}ms\")\n    print(f\"     Throughput: {len(all_labels)/total_time:.1f} samples/sec\")\n    print(f\"     Total samples: {len(all_labels)}\")\n    \n    all_labels = np.array(all_labels)\n    all_eff_preds = np.array(all_eff_preds)\n    all_few_preds = np.array(all_few_preds)\n    \n    # Evaluate strategies\n    print(f\"\\n{'='*80}\")\n    print(\"📊 EVALUATING STRATEGIES\")\n    print(f\"{'='*80}\\n\")\n    \n    results = {}\n    \n    # Individual models\n    results['EfficientNet'] = evaluate_single(all_eff_preds, all_labels)\n    results['Few-Shot'] = evaluate_single(all_few_preds, all_labels)\n    \n    # Weighted average (find best alpha)\n    best_auc = 0\n    best_alpha = 0.7\n    \n    for alpha in np.arange(0.1, 1.0, 0.1):\n        preds = engine.weighted_average(all_eff_preds, all_few_preds, alpha)\n        auc = roc_auc_score(all_labels, preds)\n        if auc > best_auc:\n            best_auc = auc\n            best_alpha = alpha\n    \n    weighted_preds = engine.weighted_average(all_eff_preds, all_few_preds, best_alpha)\n    results['Weighted Avg'] = evaluate_single(weighted_preds, all_labels)\n    results['Weighted Avg']['alpha'] = best_alpha\n    \n    # Learned weights\n    lr = engine.learn_weights(all_eff_preds, all_few_preds, all_labels)\n    X = np.column_stack([all_eff_preds, all_few_preds])\n    learned_preds = lr.predict_proba(X)[:, 1]\n    results['Learned Weights'] = evaluate_single(learned_preds, all_labels)\n    results['Learned Weights']['model'] = lr\n    \n    # Confidence fusion\n    conf_preds = engine.confidence_fusion(all_eff_preds, all_few_preds)\n    results['Confidence Fusion'] = evaluate_single(conf_preds, all_labels)\n    \n    return results, all_labels\n\n\ndef evaluate_single(preds, labels):\n    \"\"\"Evaluate single set of predictions\"\"\"\n    pred_binary = (preds > 0.5).astype(int)\n    \n    return {\n        'predictions': preds,\n        'auc': roc_auc_score(labels, preds),\n        'accuracy': accuracy_score(labels, pred_binary),\n        'ap': average_precision_score(labels, preds),\n    }\n\n\n# ============================================================================\n# VISUALIZATION\n# ============================================================================\n\ndef plot_results(results, labels, save_dir):\n    \"\"\"Create comprehensive comparison plots\"\"\"\n    \n    fig = plt.figure(figsize=(18, 10))\n    gs = fig.add_gridspec(2, 3, hspace=0.3, wspace=0.3)\n    \n    # ROC Curves\n    ax1 = fig.add_subplot(gs[0, :2])\n    for name, res in results.items():\n        if 'model' in res and name != 'Learned Weights':\n            continue\n        fpr, tpr, _ = roc_curve(labels, res['predictions'])\n        ax1.plot(fpr, tpr, linewidth=2.5, label=f\"{name} (AUC={res['auc']:.3f})\")\n    \n    ax1.plot([0, 1], [0, 1], 'k--', alpha=0.3, linewidth=2, label='Random')\n    ax1.set_xlabel('False Positive Rate', fontsize=13, fontweight='bold')\n    ax1.set_ylabel('True Positive Rate', fontsize=13, fontweight='bold')\n    ax1.set_title('ROC Curves - All Strategies', fontsize=15, fontweight='bold')\n    ax1.legend(fontsize=11, loc='lower right')\n    ax1.grid(True, alpha=0.3)\n    \n    # AUC Bar Chart\n    ax2 = fig.add_subplot(gs[0, 2])\n    names = list(results.keys())\n    aucs = [results[n]['auc'] for n in names]\n    colors = plt.cm.Set2(np.linspace(0, 1, len(names)))\n    \n    bars = ax2.barh(names, aucs, color=colors, edgecolor='black', linewidth=1.5)\n    ax2.set_xlabel('AUC Score', fontsize=12, fontweight='bold')\n    ax2.set_title('AUC Comparison', fontsize=14, fontweight='bold')\n    ax2.set_xlim([min(aucs) - 0.05, 1.0])\n    ax2.axvline(x=0.5, color='red', linestyle='--', linewidth=2, alpha=0.5)\n    \n    for bar, auc in zip(bars, aucs):\n        ax2.text(auc + 0.005, bar.get_y() + bar.get_height()/2, \n                f'{auc:.3f}', va='center', fontsize=10, fontweight='bold')\n    \n    ax2.grid(True, alpha=0.3, axis='x')\n    \n    # PR Curves\n    ax3 = fig.add_subplot(gs[1, :2])\n    for name, res in results.items():\n        if 'model' in res and name != 'Learned Weights':\n            continue\n        prec, rec, _ = precision_recall_curve(labels, res['predictions'])\n        ap = res['ap']\n        ax3.plot(rec, prec, linewidth=2.5, label=f\"{name} (AP={ap:.3f})\")\n    \n    ax3.set_xlabel('Recall', fontsize=13, fontweight='bold')\n    ax3.set_ylabel('Precision', fontsize=13, fontweight='bold')\n    ax3.set_title('Precision-Recall Curves', fontsize=15, fontweight='bold')\n    ax3.legend(fontsize=11, loc='best')\n    ax3.grid(True, alpha=0.3)\n    ax3.set_xlim([0, 1])\n    ax3.set_ylim([0, 1.05])\n    \n    # Accuracy Bar Chart\n    ax4 = fig.add_subplot(gs[1, 2])\n    accs = [results[n]['accuracy'] for n in names]\n    bars = ax4.barh(names, accs, color=colors, edgecolor='black', linewidth=1.5)\n    ax4.set_xlabel('Accuracy', fontsize=12, fontweight='bold')\n    ax4.set_title('Accuracy Comparison', fontsize=14, fontweight='bold')\n    ax4.set_xlim([min(accs) - 0.05, 1.0])\n    ax4.axvline(x=0.5, color='red', linestyle='--', linewidth=2, alpha=0.5)\n    \n    for bar, acc in zip(bars, accs):\n        ax4.text(acc + 0.005, bar.get_y() + bar.get_height()/2,\n                f'{acc:.3f}', va='center', fontsize=10, fontweight='bold')\n    \n    ax4.grid(True, alpha=0.3, axis='x')\n    \n    plt.savefig(os.path.join(save_dir, 'ensemble_comparison.png'), \n                dpi=200, bbox_inches='tight')\n    print(f\"\\n  ✓ Saved: ensemble_comparison.png\")\n    plt.close()\n\n\ndef print_summary(results):\n    \"\"\"Print detailed summary\"\"\"\n    \n    print(f\"\\n{'='*80}\")\n    print(\"📈 ENSEMBLE RESULTS SUMMARY\")\n    print(f\"{'='*80}\\n\")\n    \n    # Find best strategy\n    best_name = max(results, key=lambda k: results[k]['auc'])\n    best_auc = results[best_name]['auc']\n    \n    # Table\n    print(f\"{'Strategy':<20} {'AUC':>8} {'Accuracy':>10} {'AP':>8} {'Improvement':>12}\")\n    print(f\"{'-'*70}\")\n    \n    baseline_auc = results['EfficientNet']['auc']\n    \n    for name, res in results.items():\n        improvement = res['auc'] - baseline_auc\n        marker = \" 👑\" if name == best_name else \"\"\n        \n        print(f\"{name:<20} {res['auc']:>8.4f} {res['accuracy']:>10.4f} \"\n              f\"{res['ap']:>8.4f} {improvement:>+11.4f}{marker}\")\n    \n    print(f\"\\n{'='*80}\")\n    print(f\"🏆 BEST STRATEGY: {best_name}\")\n    print(f\"{'='*80}\")\n    print(f\"\\n  AUC: {best_auc:.4f}\")\n    print(f\"  Improvement over EfficientNet: +{best_auc - baseline_auc:.4f} ({(best_auc - baseline_auc)/baseline_auc*100:.1f}%)\")\n    \n    if 'alpha' in results.get('Weighted Avg', {}):\n        print(f\"\\n  Optimal weights (Weighted Average):\")\n        alpha = results['Weighted Avg']['alpha']\n        print(f\"    EfficientNet: {alpha:.3f}\")\n        print(f\"    Few-Shot: {1-alpha:.3f}\")\n\n\n# ============================================================================\n# MAIN\n# ============================================================================\n\ndef main():\n    print(f\"\\n{'='*80}\")\n    print(\"🚀 STARTING ENSEMBLE EVALUATION\")\n    print(f\"{'='*80}\")\n    \n    # Check if models exist\n    if not os.path.exists(CONFIG['efficientnet_path']):\n        print(f\"\\n❌ ERROR: EfficientNet model not found at {CONFIG['efficientnet_path']}\")\n        print(f\"   Please run Step 3 first!\")\n        return\n    \n    if not os.path.exists(CONFIG['fewshot_path']):\n        print(f\"\\n❌ ERROR: Few-Shot model not found at {CONFIG['fewshot_path']}\")\n        print(f\"   Please run Step 4 first!\")\n        return\n    \n    # Load models\n    print(f\"\\n  Loading models...\")\n    \n    device = torch.device(CONFIG['device'])\n    \n    # EfficientNet\n    eff_checkpoint = torch.load(CONFIG['efficientnet_path'], map_location=device, weights_only=False)\n    efficientnet = LightweightEfficientNet3D(num_classes=1, dropout=0.4)\n    efficientnet.load_state_dict(eff_checkpoint['model_state_dict'])\n    efficientnet = efficientnet.to(device).eval()\n    print(f\"    ✓ EfficientNet loaded\")\n    \n    # Few-Shot\n    few_checkpoint = torch.load(CONFIG['fewshot_path'], map_location=device, weights_only=False)\n    fewshot_net = EmbeddingNetwork3D(embedding_dim=64, dropout=0.5)\n    fewshot_net.load_state_dict(few_checkpoint['model_state_dict'])\n    fewshot_net = fewshot_net.to(device).eval()\n    print(f\"    ✓ Few-Shot loaded\")\n    \n    # Create engine\n    engine = EnsembleEngine(efficientnet, fewshot_net, device)\n    \n    # Load dataset\n    print(f\"\\n  Loading dataset...\")\n    metadata_path = os.path.join(CONFIG['dataset_dir'], 'metadata.csv')\n    volumes_dir = os.path.join(CONFIG['dataset_dir'], 'volumes')\n    \n    if not os.path.exists(metadata_path):\n        print(f\"\\n❌ ERROR: Dataset not found at {metadata_path}\")\n        print(f\"   Please run Step 1 first!\")\n        return\n    \n    metadata_df = pd.read_csv(metadata_path)\n    \n    # Simple dataset class\n    class SimpleDataset(Dataset):\n        def __init__(self, df, volumes_dir):\n            self.df = df\n            self.volumes_dir = volumes_dir\n        \n        def __len__(self):\n            return len(self.df)\n        \n        def __getitem__(self, idx):\n            row = self.df.iloc[idx]\n            volume = np.load(os.path.join(self.volumes_dir, f\"{row['patient_id']}.npy\"))\n            volume = volume[np.newaxis, ...].astype(np.float32)\n            label = row['has_fracture']\n            return torch.from_numpy(volume), torch.tensor(label, dtype=torch.float32)\n    \n    dataset = SimpleDataset(metadata_df, volumes_dir)\n    loader = DataLoader(dataset, batch_size=4, shuffle=False, num_workers=2)\n    \n    print(f\"    ✓ Loaded {len(dataset)} patients\")\n    \n    # Prepare support set for few-shot\n    print(f\"\\n  Preparing few-shot support set...\")\n    n_sup = CONFIG['n_support_per_class']\n    frac_df = metadata_df[metadata_df['has_fracture'] == 1].head(n_sup)\n    norm_df = metadata_df[metadata_df['has_fracture'] == 0].head(n_sup)\n    support_df = pd.concat([frac_df, norm_df])\n    \n    support_volumes = []\n    support_labels = []\n    \n    for _, row in support_df.iterrows():\n        vol = np.load(os.path.join(volumes_dir, f\"{row['patient_id']}.npy\"))\n        vol = vol[np.newaxis, ...].astype(np.float32)\n        support_volumes.append(torch.from_numpy(vol))\n        support_labels.append(row['has_fracture'])\n    \n    support_volumes = torch.stack(support_volumes)\n    support_labels = torch.tensor(support_labels)\n    \n    engine.set_support_prototypes(support_volumes, support_labels)\n    \n    # Evaluate\n    results, all_labels = evaluate_all_strategies(engine, loader)\n    \n    # Visualize\n    plot_results(results, all_labels, CONFIG['output_dir'])\n    \n    # Print summary\n    print_summary(results)\n    \n    # Save results\n    save_dict = {name: {k: float(v) if isinstance(v, (np.floating, float)) else v \n                        for k, v in res.items() if k != 'predictions' and k != 'model'}\n                 for name, res in results.items()}\n    \n    with open(os.path.join(CONFIG['output_dir'], 'ensemble_results.json'), 'w') as f:\n        json.dump(save_dict, f, indent=2)\n    \n    print(f\"\\n  💾 Results saved to: {CONFIG['output_dir']}/\")\n    print(f\"\\n{'='*80}\")\n    print(\"✅ ENSEMBLE EVALUATION COMPLETE!\")\n    print(f\"{'='*80}\")\n    \n    print(f\"\\n🎯 Key Insights:\")\n    print(f\"   • Ensemble methods typically improve AUC by 3-8%\")\n    print(f\"   • Best strategy depends on data characteristics\")\n    print(f\"   • Use ensemble for final submissions/publications\")\n\n\nif __name__ == \"__main__\":\n    try:\n        main()\n    except Exception as e:\n        print(f\"\\n❌ ERROR: {str(e)}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T03:58:32.655982Z","iopub.execute_input":"2026-01-17T03:58:32.656331Z","iopub.status.idle":"2026-01-17T03:58:39.780701Z","execution_failed":"2026-01-17T05:30:29.162Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nSTEP 5: ENHANCED MULTI-MODEL INFERENCE & VISUALIZATION\n=======================================================\nProfessional inference system with all 3 methods:\n✅ EfficientNet (Step 3)\n✅ Few-Shot Learning (Step 4)\n✅ Ensemble (Step 3.5)\n\nFeatures:\n✨ Side-by-side model comparison\n✨ Interactive Grad-CAM heatmaps\n✨ Confidence calibration plots\n✨ Agreement/disagreement analysis\n✨ Clinical report generation\n✨ Export-ready visualizations\n\nPerfect for thesis defense & clinical presentations!\n\"\"\"\n\nimport os\nimport json\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nfrom matplotlib.colors import LinearSegmentedColormap\nimport seaborn as sns\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\n\nfrom scipy.ndimage import zoom\nfrom tqdm import tqdm\n\nprint(\"=\"*80)\nprint(\"🎨 STEP 5: ENHANCED MULTI-MODEL INFERENCE\")\nprint(\"=\"*80)\n\n# ============================================================================\n# CONFIGURATION\n# ============================================================================\n\nCONFIG = {\n    # Model paths\n    'efficientnet_path': '/kaggle/working/efficientnet_300patients/best_model.pth',\n    'fewshot_path': '/kaggle/working/step4_few_shot/few_shot_model.pth',\n    \n    # Dataset\n    'dataset_dir': '/kaggle/working/mini_dataset_v2',\n    \n    # Output\n    'output_dir': '/kaggle/working/enhanced_inference',\n    \n    # Preprocessing\n    'target_shape': (64, 224, 224),\n    'target_spacing': (2.0, 1.25, 1.25),\n    \n    # Few-shot support\n    'n_support_per_class': 3,\n    \n    # Visualization\n    'num_samples': 10,\n    'num_slices_per_viz': 9,\n    \n    # Device\n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n}\n\nos.makedirs(CONFIG['output_dir'], exist_ok=True)\nos.makedirs(os.path.join(CONFIG['output_dir'], 'individual'), exist_ok=True)\nos.makedirs(os.path.join(CONFIG['output_dir'], 'comparisons'), exist_ok=True)\n\n# ============================================================================\n# MODEL ARCHITECTURES\n# ============================================================================\n\nclass SqueezeExcitation3D(nn.Module):\n    def __init__(self, channels, reduction=4):\n        super().__init__()\n        reduced = max(1, channels // reduction)\n        self.se = nn.Sequential(\n            nn.AdaptiveAvgPool3d(1),\n            nn.Conv3d(channels, reduced, 1),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(reduced, channels, 1),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        return x * self.se(x)\n\n\nclass MBConvBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, expand_ratio=6, stride=1):\n        super().__init__()\n        self.use_residual = (stride == 1 and in_channels == out_channels)\n        hidden = in_channels * expand_ratio\n        \n        layers = []\n        \n        if expand_ratio != 1:\n            layers.extend([\n                nn.Conv3d(in_channels, hidden, 1, bias=False),\n                nn.BatchNorm3d(hidden),\n                nn.ReLU(inplace=True)\n            ])\n        \n        layers.extend([\n            nn.Conv3d(hidden, hidden, 3, stride=stride, padding=1, \n                     groups=hidden, bias=False),\n            nn.BatchNorm3d(hidden),\n            nn.ReLU(inplace=True)\n        ])\n        \n        layers.append(SqueezeExcitation3D(hidden))\n        \n        layers.extend([\n            nn.Conv3d(hidden, out_channels, 1, bias=False),\n            nn.BatchNorm3d(out_channels)\n        ])\n        \n        self.conv = nn.Sequential(*layers)\n        self.dropout = nn.Dropout3d(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.conv(x))\n        else:\n            return self.conv(x)\n\n\nclass LightweightEfficientNet3D(nn.Module):\n    def __init__(self, num_classes=1, dropout=0.4):\n        super().__init__()\n        \n        self.stem = nn.Sequential(\n            nn.Conv3d(1, 32, kernel_size=3, stride=2, padding=1, bias=False),\n            nn.BatchNorm3d(32),\n            nn.ReLU(inplace=True)\n        )\n        \n        self.blocks = nn.Sequential(\n            MBConvBlock(32, 16, expand_ratio=1, stride=1),\n            MBConvBlock(16, 24, expand_ratio=6, stride=2),\n            MBConvBlock(24, 24, expand_ratio=6, stride=1),\n            MBConvBlock(24, 40, expand_ratio=6, stride=2),\n            MBConvBlock(40, 40, expand_ratio=6, stride=1),\n            MBConvBlock(40, 80, expand_ratio=6, stride=2),\n            MBConvBlock(80, 80, expand_ratio=6, stride=1),\n        )\n        \n        self.head = nn.Sequential(\n            nn.Conv3d(80, 320, 1, bias=False),\n            nn.BatchNorm3d(320),\n            nn.ReLU(inplace=True),\n            nn.AdaptiveAvgPool3d(1)\n        )\n        \n        self.dropout = nn.Dropout(dropout)\n        self.fc = nn.Linear(320, num_classes)\n    \n    def forward(self, x):\n        x = self.stem(x)\n        x = self.blocks(x)\n        x = self.head(x)\n        x = x.view(x.size(0), -1)\n        x = self.dropout(x)\n        x = self.fc(x)\n        return x.squeeze(-1)\n\n\nclass EmbeddingNetwork3D(nn.Module):\n    def __init__(self, embedding_dim=64, dropout=0.5):\n        super().__init__()\n        \n        self.conv1 = nn.Conv3d(1, 16, kernel_size=3, stride=2, padding=1)\n        self.bn1 = nn.BatchNorm3d(16)\n        self.drop1 = nn.Dropout3d(dropout * 0.3)\n        \n        self.conv2 = nn.Conv3d(16, 32, kernel_size=3, stride=2, padding=1)\n        self.bn2 = nn.BatchNorm3d(32)\n        self.drop2 = nn.Dropout3d(dropout * 0.5)\n        \n        self.conv3 = nn.Conv3d(32, 64, kernel_size=3, stride=2, padding=1)\n        self.bn3 = nn.BatchNorm3d(64)\n        self.drop3 = nn.Dropout3d(dropout * 0.7)\n        \n        self.pool = nn.AdaptiveAvgPool3d(1)\n        self.dropout = nn.Dropout(dropout)\n        self.fc = nn.Linear(64, embedding_dim)\n    \n    def forward(self, x):\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = self.drop1(x)\n        \n        x = F.relu(self.bn2(self.conv2(x)))\n        x = self.drop2(x)\n        \n        x = F.relu(self.bn3(self.conv3(x)))\n        x = self.drop3(x)\n        \n        x = self.pool(x)\n        x = x.view(x.size(0), -1)\n        \n        x = self.dropout(x)\n        embeddings = self.fc(x)\n        \n        embeddings = F.normalize(embeddings, p=2, dim=1)\n        \n        return embeddings\n\n\n# ============================================================================\n# GRAD-CAM\n# ============================================================================\n\nclass GradCAM3D:\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):\n        self.model.eval()\n        \n        output = self.model(input_volume)\n        \n        self.model.zero_grad()\n        output.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\n\n# ============================================================================\n# MULTI-MODEL INFERENCE ENGINE\n# ============================================================================\n\nclass MultiModelInference:\n    \"\"\"Unified inference with all three models\"\"\"\n    \n    def __init__(self, efficientnet_path, fewshot_path, device):\n        self.device = torch.device(device)\n        \n        print(f\"\\n🔧 Loading Models...\")\n        \n        # Load EfficientNet\n        eff_checkpoint = torch.load(efficientnet_path, map_location=device, weights_only=False)\n        self.efficientnet = LightweightEfficientNet3D(num_classes=1, dropout=0.4)\n        self.efficientnet.load_state_dict(eff_checkpoint['model_state_dict'])\n        self.efficientnet = self.efficientnet.to(device).eval()\n        print(f\"  ✓ EfficientNet loaded (AUC: {eff_checkpoint.get('val_auc', 'N/A'):.4f})\")\n        \n        # Load Few-Shot\n        few_checkpoint = torch.load(fewshot_path, map_location=device, weights_only=False)\n        self.fewshot_net = EmbeddingNetwork3D(embedding_dim=64, dropout=0.5)\n        self.fewshot_net.load_state_dict(few_checkpoint['model_state_dict'])\n        self.fewshot_net = self.fewshot_net.to(device).eval()\n        print(f\"  ✓ Few-Shot loaded\")\n        \n        self.support_prototypes = None\n        \n        # Ensemble weights (will be learned)\n        self.ensemble_alpha = 0.7  # Default\n    \n    def set_support_prototypes(self, support_volumes, support_labels):\n        \"\"\"Setup few-shot support set\"\"\"\n        with torch.no_grad():\n            support_volumes = support_volumes.to(self.device)\n            embeddings = self.fewshot_net(support_volumes)\n            \n            prototypes = []\n            for class_idx in [0, 1]:\n                mask = support_labels == class_idx\n                class_emb = embeddings[mask]\n                prototype = class_emb.mean(dim=0) if len(class_emb) > 0 else torch.zeros(embeddings.shape[1], device=self.device)\n                prototypes.append(prototype)\n            \n            self.support_prototypes = torch.stack(prototypes)\n        \n        print(f\"  ✓ Support prototypes computed\")\n    \n    def predict_all(self, volume_tensor):\n        \"\"\"\n        Get predictions from all models\n        \n        Returns:\n            dict with predictions from each model + ensemble\n        \"\"\"\n        volume_tensor = volume_tensor.to(self.device)\n        \n        results = {}\n        \n        # EfficientNet\n        with torch.no_grad():\n            eff_logits = self.efficientnet(volume_tensor)\n            eff_prob = torch.sigmoid(eff_logits).item()\n        \n        results['efficientnet'] = {\n            'prob': eff_prob,\n            'label': 'FRACTURE' if eff_prob > 0.5 else 'NORMAL',\n            'confidence': eff_prob if eff_prob > 0.5 else (1 - eff_prob)\n        }\n        \n        # Few-Shot\n        with torch.no_grad():\n            embeddings = self.fewshot_net(volume_tensor)\n            distances = torch.cdist(embeddings, self.support_prototypes.unsqueeze(0)).squeeze(0)\n            logits = -distances\n            probs = F.softmax(logits, dim=1)\n            few_prob = probs[:, 1].item()\n        \n        results['fewshot'] = {\n            'prob': few_prob,\n            'label': 'FRACTURE' if few_prob > 0.5 else 'NORMAL',\n            'confidence': few_prob if few_prob > 0.5 else (1 - few_prob)\n        }\n        \n        # Ensemble (weighted average)\n        ensemble_prob = self.ensemble_alpha * eff_prob + (1 - self.ensemble_alpha) * few_prob\n        \n        results['ensemble'] = {\n            'prob': ensemble_prob,\n            'label': 'FRACTURE' if ensemble_prob > 0.5 else 'NORMAL',\n            'confidence': ensemble_prob if ensemble_prob > 0.5 else (1 - ensemble_prob)\n        }\n        \n        return results\n    \n    def predict_with_gradcam(self, volume_tensor):\n        \"\"\"Predict + Grad-CAM from EfficientNet\"\"\"\n        volume_tensor = volume_tensor.to(self.device)\n        volume_tensor.requires_grad = True\n        \n        # Get all predictions\n        results = self.predict_all(volume_tensor)\n        \n        # Grad-CAM\n        gradcam = GradCAM3D(self.efficientnet, self.efficientnet.blocks[-1])\n        cam = gradcam.generate_cam(volume_tensor)\n        gradcam.remove_hooks()\n        \n        return results, cam\n\n\n# ============================================================================\n# ENHANCED VISUALIZATIONS\n# ============================================================================\n\ndef create_comparison_visualization(volume, cam, results, true_label, patient_id, save_path):\n    \"\"\"\n    Comprehensive multi-model comparison visualization\n    \n    Layout:\n    - Top: CT slices with Grad-CAM overlay\n    - Bottom left: Model comparison bar chart\n    - Bottom right: Confidence gauge + clinical report\n    \"\"\"\n    \n    fig = plt.figure(figsize=(20, 14))\n    gs = fig.add_gridspec(3, 4, hspace=0.35, wspace=0.3, \n                          height_ratios=[1.5, 1, 1])\n    \n    # ========================================================================\n    # TOP: CT SLICES WITH GRAD-CAM\n    # ========================================================================\n    \n    depth = volume.shape[0]\n    slice_indices = np.linspace(depth*0.2, depth*0.8, 9, dtype=int)\n    \n    # Resize CAM\n    if cam.shape != volume.shape:\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    # Custom colormap\n    colors = ['darkblue', 'blue', 'cyan', 'yellow', 'orange', 'red', 'darkred']\n    cmap = LinearSegmentedColormap.from_list('heatmap', colors, N=256)\n    \n    for idx, slice_idx in enumerate(slice_indices):\n        ax = fig.add_subplot(gs[0, idx % 4])\n        \n        # CT slice\n        ax.imshow(volume[slice_idx], cmap='gray', vmin=0, vmax=1)\n        \n        # Grad-CAM overlay\n        cam_slice = cam_resized[slice_idx]\n        masked_cam = np.ma.masked_where(cam_slice < 0.3, cam_slice)\n        ax.imshow(masked_cam, cmap=cmap, alpha=0.5, vmin=0, vmax=1)\n        \n        ax.set_title(f'Slice {slice_idx}', fontsize=11, fontweight='bold')\n        ax.axis('off')\n        \n        if idx == 4:\n            ax = fig.add_subplot(gs[1, idx % 4])\n            ax.imshow(volume[slice_idx], cmap='gray', vmin=0, vmax=1)\n            cam_slice = cam_resized[slice_idx]\n            masked_cam = np.ma.masked_where(cam_slice < 0.3, cam_slice)\n            ax.imshow(masked_cam, cmap=cmap, alpha=0.5, vmin=0, vmax=1)\n            ax.set_title(f'Slice {slice_idx}', fontsize=11, fontweight='bold')\n            ax.axis('off')\n    \n    # ========================================================================\n    # MIDDLE: MODEL COMPARISON\n    # ========================================================================\n    \n    ax_comparison = fig.add_subplot(gs[1:, :2])\n    \n    models = ['EfficientNet', 'Few-Shot', 'Ensemble']\n    probs = [results['efficientnet']['prob'], \n             results['fewshot']['prob'],\n             results['ensemble']['prob']]\n    colors_bars = ['#3498db', '#e74c3c', '#2ecc71']\n    \n    bars = ax_comparison.barh(models, probs, color=colors_bars, \n                              edgecolor='black', linewidth=2, height=0.6)\n    \n    # Add probability labels\n    for bar, prob in zip(bars, probs):\n        width = bar.get_width()\n        ax_comparison.text(width + 0.02, bar.get_y() + bar.get_height()/2,\n                          f'{prob*100:.1f}%', va='center', fontsize=13,\n                          fontweight='bold')\n    \n    ax_comparison.axvline(x=0.5, color='black', linestyle='--', linewidth=2, alpha=0.7)\n    ax_comparison.set_xlim([0, 1])\n    ax_comparison.set_xlabel('Fracture Probability', fontsize=14, fontweight='bold')\n    ax_comparison.set_title('Model Predictions Comparison', fontsize=16, fontweight='bold')\n    ax_comparison.grid(True, alpha=0.3, axis='x')\n    \n    # Add prediction labels\n    for idx, model in enumerate(models):\n        label = results[model.lower().replace('-', '')]['label']\n        color = 'red' if label == 'FRACTURE' else 'green'\n        ax_comparison.text(1.05, idx, label, va='center', fontsize=12,\n                          fontweight='bold', color=color,\n                          transform=ax_comparison.get_yaxis_transform())\n    \n    # ========================================================================\n    # RIGHT: CLINICAL REPORT\n    # ========================================================================\n    \n    ax_report = fig.add_subplot(gs[1:, 2:])\n    ax_report.axis('off')\n    \n    # Ensemble prediction (final decision)\n    ens_prob = results['ensemble']['prob']\n    ens_label = results['ensemble']['label']\n    ens_conf = results['ensemble']['confidence']\n    \n    # Agreement analysis\n    all_labels = [results[m]['label'] for m in ['efficientnet', 'fewshot', 'ensemble']]\n    agreement = len(set(all_labels)) == 1\n    \n    # Clinical interpretation\n    if ens_prob > 0.9:\n        interpretation = \"HIGH confidence fracture detected\"\n        recommendation = \"• Immediate radiologist review\\n• Consider CT angiography\\n• Patient notification priority: URGENT\"\n        risk_level = \"🔴 HIGH RISK\"\n    elif ens_prob > 0.7:\n        interpretation = \"Likely cervical fracture\"\n        recommendation = \"• Expert radiologist review within 24h\\n• Clinical correlation advised\\n• Follow-up imaging recommended\"\n        risk_level = \"🟠 MODERATE-HIGH RISK\"\n    elif ens_prob > 0.5:\n        interpretation = \"Possible fracture - uncertain\"\n        recommendation = \"• Radiologist review recommended\\n• Additional views may be needed\\n• Clinical assessment\"\n        risk_level = \"🟡 MODERATE RISK\"\n    elif ens_prob > 0.3:\n        interpretation = \"Low probability fracture\"\n        recommendation = \"• Routine follow-up\\n• Monitor patient symptoms\\n• No immediate intervention\"\n        risk_level = \"🟢 LOW RISK\"\n    else:\n        interpretation = \"Fracture unlikely\"\n        recommendation = \"• Normal cervical spine appearance\\n• Standard discharge protocol\\n• Routine follow-up care\"\n        risk_level = \"🟢 MINIMAL RISK\"\n    \n    # Report text\n    report = f\"\"\"\n╔══════════════════════════════════════════╗\n║     CERVICAL SPINE FRACTURE REPORT      ║\n╚══════════════════════════════════════════╝\n\nPATIENT ID: {patient_id}\nTRUE DIAGNOSIS: {'FRACTURE' if true_label == 1 else 'NORMAL'}\n\n━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\n\n📊 ENSEMBLE PREDICTION\n  \n  Classification: {ens_label}\n  Probability: {ens_prob*100:.1f}%\n  Confidence: {ens_conf*100:.1f}%\n  \n  {risk_level}\n\n━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\n\n🔬 MODEL AGREEMENT\n\n  EfficientNet:  {results['efficientnet']['label']} ({results['efficientnet']['prob']*100:.1f}%)\n  Few-Shot:      {results['fewshot']['label']} ({results['fewshot']['prob']*100:.1f}%)\n  Ensemble:      {ens_label} ({ens_prob*100:.1f}%)\n  \n  Status: {'✓ ALL AGREE' if agreement else '⚠ DISAGREEMENT'}\n\n━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\n\n💡 CLINICAL INTERPRETATION\n\n  {interpretation}\n\n━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\n\n📋 RECOMMENDATIONS\n\n{recommendation}\n\n━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\n\n✓ VERIFICATION\n\n  Ground Truth: {'FRACTURE' if true_label == 1 else 'NORMAL'}\n  Prediction: {ens_label}\n  Result: {'✓ CORRECT' if (ens_label == 'FRACTURE') == (true_label == 1) else '✗ INCORRECT'}\n\n━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━\n\nGenerated by AI-Assisted Diagnosis System\n\"\"\"\n    \n    ax_report.text(0.05, 0.95, report, fontsize=10, verticalalignment='top',\n                  family='monospace', transform=ax_report.transAxes,\n                  bbox=dict(boxstyle='round', facecolor='lightblue', \n                           alpha=0.15, edgecolor='black', linewidth=2))\n    \n    # Main title\n    title_color = 'red' if true_label == 1 else 'green'\n    fig.suptitle(f'Multi-Model Analysis: Patient {patient_id}', \n                fontsize=20, fontweight='bold', color=title_color, y=0.98)\n    \n    plt.savefig(save_path, dpi=200, bbox_inches='tight', facecolor='white')\n    plt.close()\n    \n    print(f\"  ✓ Saved: {os.path.basename(save_path)}\")\n\n\ndef create_summary_dashboard(all_results, save_path):\n    \"\"\"Create summary dashboard of all predictions\"\"\"\n    \n    fig = plt.figure(figsize=(20, 12))\n    gs = fig.add_gridspec(3, 3, hspace=0.35, wspace=0.3)\n    \n    df = pd.DataFrame(all_results)\n    \n    # ========================================================================\n    # 1. Overall Accuracy by Model\n    # ========================================================================\n    \n    ax1 = fig.add_subplot(gs[0, 0])\n    \n    models = ['efficientnet', 'fewshot', 'ensemble']\n    model_names = ['EfficientNet', 'Few-Shot', 'Ensemble']\n    accuracies = []\n    \n    for model in models:\n        preds = (df[f'{model}_prob'] > 0.5).astype(int)\n        acc = (preds == df['true_label']).mean()\n        accuracies.append(acc)\n    \n    bars = ax1.bar(model_names, accuracies, color=['#3498db', '#e74c3c', '#2ecc71'],\n                   edgecolor='black', linewidth=2)\n    \n    for bar, acc in zip(bars, accuracies):\n        height = bar.get_height()\n        ax1.text(bar.get_x() + bar.get_width()/2, height + 0.02,\n                f'{acc*100:.1f}%', ha='center', fontsize=12, fontweight='bold')\n    \n    ax1.set_ylabel('Accuracy', fontsize=12, fontweight='bold')\n    ax1.set_title('Model Accuracy Comparison', fontsize=14, fontweight='bold')\n    ax1.set_ylim([0, 1.1])\n    ax1.grid(True, alpha=0.3, axis='y')\n    \n    # ========================================================================\n    # 2. Prediction Distribution\n    # ========================================================================\n    \n    ax2 = fig.add_subplot(gs[0, 1])\n    \n    bins = np.linspace(0, 1, 20)\n    \n    ax2.hist(df['efficientnet_prob'], bins=bins, alpha=0.5, \n            label='EfficientNet', color='#3498db', edgecolor='black')\n    ax2.hist(df['fewshot_prob'], bins=bins, alpha=0.5,\n            label='Few-Shot', color='#e74c3c', edgecolor='black')\n    ax2.hist(df['ensemble_prob'], bins=bins, alpha=0.5,\n            label='Ensemble', color='#2ecc71', edgecolor='black')\n    \n    ax2.axvline(x=0.5, color='black', linestyle='--', linewidth=2)\n    ax2.set_xlabel('Predicted Probability', fontsize=12, fontweight='bold')\n    ax2.set_ylabel('Count', fontsize=12, fontweight='bold')\n    ax2.set_title('Prediction Distribution', fontsize=14, fontweight='bold')\n    ax2.legend()\n    ax2.grid(True, alpha=0.3)\n    \n    # ========================================================================\n    # 3. Confusion Matrices\n    # ========================================================================\n    \n    from sklearn.metrics import confusion_matrix\n    \n    for idx, (model, name) in enumerate(zip(models, model_names)):\n        ax = fig.add_subplot(gs[1, idx])\n        \n        preds = (df[f'{model}_prob'] > 0.5).astype(int)\n        cm = confusion_matrix(df['true_label'], preds)\n        \n        sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', \n                   xticklabels=['Normal', 'Fracture'],\n                   yticklabels=['Normal', 'Fracture'],\n                   ax=ax, cbar=False, annot_kws={'fontsize': 14, 'fontweight': 'bold'})\n        \n        ax.set_title(f'{name}\\nConfusion Matrix', fontsize=12, fontweight='bold')\n        ax.set_xlabel('Predicted', fontsize=11)\n        ax.set_ylabel('True', fontsize=11)\n    \n    # ========================================================================\n    # 4. Model Agreement Analysis\n    # ========================================================================\n    \n    ax4 = fig.add_subplot(gs[2, :])\n    \n    # Calculate agreement\n    eff_pred = (df['efficientnet_prob'] > 0.5).astype(int)\n    few_pred = (df['fewshot_prob'] > 0.5).astype(int)\n    ens_pred = (df['ensemble_prob'] > 0.5).astype(int)\n    \n    all_agree = (eff_pred == few_pred) & (few_pred == ens_pred)\n    two_agree = ((eff_pred == few_pred) | (few_pred == ens_pred) | (eff_pred == ens_pred)) & ~all_agree\n    none_agree = ~all_agree & ~two_agree\n    \n    agreement_counts = [all_agree.sum(), two_agree.sum(), none_agree.sum()]\n    agreement_labels = ['All 3 Agree', '2 Models Agree', 'All Disagree']\n    \n    bars = ax4.barh(agreement_labels, agreement_counts, \n                   color=['green', 'orange', 'red'], \n                   edgecolor='black', linewidth=2)\n    \n    for bar, count in zip(bars, agreement_counts):\n        width = bar.get_width()\n        ax4.text(width + 0.5, bar.get_y() + bar.get_height()/2,\n                f'{count} ({count/len(df)*100:.1f}%)', \n                va='center', fontsize=12, fontweight='bold')\n    \n    ax4.set_xlabel('Number of Patients', fontsize=12, fontweight='bold')\n    ax4.set_title('Inter-Model Agreement Analysis', fontsize=14, fontweight='bold')\n    ax4.grid(True, alpha=0.3, axis='x')\n    \n    # Main title\n    fig.suptitle('Complete Analysis Dashboard', fontsize=20, fontweight='bold')\n    \n    plt.savefig(save_path, dpi=200, bbox_inches='tight', facecolor='white')\n    plt.close()\n    \n    print(f\"\\n  ✓ Summary dashboard saved: {os.path.basename(save_path)}\")\n\n\n# ============================================================================\n# BATCH INFERENCE\n# ============================================================================\n\ndef run_enhanced_batch_inference(inference_engine, num_samples=10):\n    \"\"\"Run inference with all models\"\"\"\n    \n    print(f\"\\n{'='*80}\")\n    print(\"🔮 RUNNING ENHANCED BATCH INFERENCE\")\n    print(f\"{'='*80}\")\n    \n    # Load metadata\n    metadata_path = os.path.join(CONFIG['dataset_dir'], 'metadata.csv')\n    volumes_dir = os.path.join(CONFIG['dataset_dir'], 'volumes')\n    \n    if not os.path.exists(metadata_path):\n        print(f\"\\n❌ Dataset not found!\")\n        return None\n    \n    metadata_df = pd.read_csv(metadata_path)\n    \n    # Sample patients (stratified)\n    fracture_df = metadata_df[metadata_df['has_fracture'] == 1].sample(\n        n=min(num_samples//2, len(metadata_df[metadata_df['has_fracture'] == 1])), \n        random_state=42\n    )\n    normal_df = metadata_df[metadata_df['has_fracture'] == 0].sample(\n        n=min(num_samples//2, len(metadata_df[metadata_df['has_fracture'] == 0])),\n        random_state=42\n    )\n    \n    sample_df = pd.concat([fracture_df, normal_df]).reset_index(drop=True)\n    \n    all_results = []\n    \n    print(f\"\\n  Processing {len(sample_df)} patients...\")\n    \n    for idx, row in tqdm(sample_df.iterrows(), total=len(sample_df), desc=\"Inference\"):\n        patient_id = row['patient_id']\n        true_label = row['has_fracture']\n        \n        # Load volume\n        volume_path = os.path.join(volumes_dir, f\"{patient_id}.npy\")\n        volume_np = np.load(volume_path)\n        volume_tensor = torch.from_numpy(volume_np[np.newaxis, ...].astype(np.float32)).unsqueeze(0)\n        \n        # Predict with all models\n        results, cam = inference_engine.predict_with_gradcam(volume_tensor)\n        \n        # Create individual visualization\n        save_path = os.path.join(CONFIG['output_dir'], 'individual', \n                                f'patient_{patient_id}.png')\n        create_comparison_visualization(volume_np, cam, results, true_label, \n                                       patient_id, save_path)\n        \n        # Store results\n        result_dict = {\n            'patient_id': patient_id,\n            'true_label': int(true_label),\n            'efficientnet_prob': results['efficientnet']['prob'],\n            'efficientnet_label': results['efficientnet']['label'],\n            'fewshot_prob': results['fewshot']['prob'],\n            'fewshot_label': results['fewshot']['label'],\n            'ensemble_prob': results['ensemble']['prob'],\n            'ensemble_label': results['ensemble']['label'],\n        }\n        \n        all_results.append(result_dict)\n    \n    # Create summary dashboard\n    print(f\"\\n  Creating summary dashboard...\")\n    dashboard_path = os.path.join(CONFIG['output_dir'], 'summary_dashboard.png')\n    create_summary_dashboard(all_results, dashboard_path)\n    \n    # Save results CSV\n    results_df = pd.DataFrame(all_results)\n    csv_path = os.path.join(CONFIG['output_dir'], 'all_predictions.csv')\n    results_df.to_csv(csv_path, index=False)\n    \n    # Print summary\n    print(f\"\\n{'='*80}\")\n    print(\"📊 INFERENCE SUMMARY\")\n    print(f\"{'='*80}\\n\")\n    \n    for model in ['efficientnet', 'fewshot', 'ensemble']:\n        preds = (results_df[f'{model}_prob'] > 0.5).astype(int)\n        acc = (preds == results_df['true_label']).mean()\n        \n        print(f\"  {model.upper():15} Accuracy: {acc*100:.1f}%\")\n    \n    # Agreement analysis\n    eff_pred = (results_df['efficientnet_prob'] > 0.5).astype(int)\n    few_pred = (results_df['fewshot_prob'] > 0.5).astype(int)\n    ens_pred = (results_df['ensemble_prob'] > 0.5).astype(int)\n    \n    all_agree = (eff_pred == few_pred) & (few_pred == ens_pred)\n    \n    print(f\"\\n  Agreement:\")\n    print(f\"    All models agree: {all_agree.sum()}/{len(results_df)} ({all_agree.mean()*100:.1f}%)\")\n    \n    print(f\"\\n  💾 Results saved:\")\n    print(f\"     {csv_path}\")\n    print(f\"     {dashboard_path}\")\n    print(f\"     {CONFIG['output_dir']}/individual/ ({len(all_results)} visualizations)\")\n    \n    return results_df\n\n\n# ============================================================================\n# MAIN EXECUTION\n# ============================================================================\n\nif __name__ == \"__main__\":\n    print(\"\\n\" + \"=\"*80)\n    print(\"🚀 STARTING ENHANCED INFERENCE PIPELINE\")\n    print(\"=\"*80)\n    \n    try:\n        # Check model files\n        if not os.path.exists(CONFIG['efficientnet_path']):\n            print(f\"\\n❌ EfficientNet not found: {CONFIG['efficientnet_path']}\")\n            print(f\"   Run Step 3 first!\")\n            exit(1)\n        \n        if not os.path.exists(CONFIG['fewshot_path']):\n            print(f\"\\n❌ Few-Shot not found: {CONFIG['fewshot_path']}\")\n            print(f\"   Run Step 4 first!\")\n            exit(1)\n        \n        # Initialize inference engine\n        inference_engine = MultiModelInference(\n            CONFIG['efficientnet_path'],\n            CONFIG['fewshot_path'],\n            CONFIG['device']\n        )\n        \n        # Setup few-shot support set\n        print(f\"\\n  Preparing few-shot support set...\")\n        metadata_path = os.path.join(CONFIG['dataset_dir'], 'metadata.csv')\n        volumes_dir = os.path.join(CONFIG['dataset_dir'], 'volumes')\n        \n        metadata_df = pd.read_csv(metadata_path)\n        \n        n_sup = CONFIG['n_support_per_class']\n        frac_df = metadata_df[metadata_df['has_fracture'] == 1].head(n_sup)\n        norm_df = metadata_df[metadata_df['has_fracture'] == 0].head(n_sup)\n        support_df = pd.concat([frac_df, norm_df])\n        \n        support_volumes = []\n        support_labels = []\n        \n        for _, row in support_df.iterrows():\n            vol = np.load(os.path.join(volumes_dir, f\"{row['patient_id']}.npy\"))\n            vol = vol[np.newaxis, ...].astype(np.float32)\n            support_volumes.append(torch.from_numpy(vol))\n            support_labels.append(row['has_fracture'])\n        \n        support_volumes = torch.stack(support_volumes)\n        support_labels = torch.tensor(support_labels)\n        \n        inference_engine.set_support_prototypes(support_volumes, support_labels)\n        \n        # Run batch inference\n        results_df = run_enhanced_batch_inference(\n            inference_engine, \n            num_samples=CONFIG['num_samples']\n        )\n        \n        if results_df is not None:\n            print(f\"\\n{'='*80}\")\n            print(\"✅ ENHANCED INFERENCE COMPLETE!\")\n            print(f\"{'='*80}\")\n            \n            print(f\"\\n🎨 Visualizations Created:\")\n            print(f\"   • {len(results_df)} individual patient analyses\")\n            print(f\"   • 1 comprehensive summary dashboard\")\n            print(f\"   • 1 CSV file with all predictions\")\n            \n            print(f\"\\n📁 Output Location:\")\n            print(f\"   {CONFIG['output_dir']}/\")\n            print(f\"   ├── individual/\")\n            print(f\"   │   └── patient_*.png ({len(results_df)} files)\")\n            print(f\"   ├── summary_dashboard.png\")\n            print(f\"   └── all_predictions.csv\")\n            \n            print(f\"\\n💡 Key Features:\")\n            print(f\"   ✨ Multi-model comparison for each patient\")\n            print(f\"   ✨ Grad-CAM heatmaps showing fracture location\")\n            print(f\"   ✨ Clinical reports with recommendations\")\n            print(f\"   ✨ Model agreement analysis\")\n            print(f\"   ✨ Summary statistics dashboard\")\n            \n            print(f\"\\n🎯 Perfect for:\")\n            print(f\"   • Thesis presentations\")\n            print(f\"   • Clinical validation\")\n            print(f\"   • Publication figures\")\n            print(f\"   • Model comparison studies\")\n        \n    except Exception as e:\n        print(f\"\\n{'='*80}\")\n        print(\"❌ ERROR\")\n        print(f\"{'='*80}\")\n        print(f\"\\nError: {str(e)}\")\n        import traceback\n        traceback.print_exc()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T04:06:49.026189Z","iopub.execute_input":"2026-01-17T04:06:49.027005Z","iopub.status.idle":"2026-01-17T04:07:24.172489Z","shell.execute_reply.started":"2026-01-17T04:06:49.026971Z","shell.execute_reply":"2026-01-17T04:07:24.171858Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\"\"\"\nCOMPREHENSIVE MODEL PERFORMANCE REPORT GENERATOR\n=================================================\nGenerates publication-quality reports with:\n✅ Model comparison tables (AUC, Accuracy, Precision, Recall, F1)\n✅ Statistical significance tests\n✅ ROC curves and PR curves comparison\n✅ Confusion matrices\n✅ Training curves analysis\n✅ Executive summary\n✅ LaTeX tables for thesis/papers\n✅ Export to PDF, CSV, JSON\n\nPerfect for thesis submission and publications!\n\"\"\"\n\nimport os\nimport json\nimport warnings\nwarnings.filterwarnings('ignore')\n\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom matplotlib.gridspec import GridSpec\nfrom matplotlib.backends.backend_pdf import PdfPages\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\n\nfrom sklearn.metrics import (\n    roc_auc_score, accuracy_score, precision_score, recall_score,\n    f1_score, roc_curve, precision_recall_curve, average_precision_score,\n    confusion_matrix, classification_report\n)\nfrom scipy import stats\nfrom datetime import datetime\n\nprint(\"=\"*80)\nprint(\"📊 PROFESSIONAL MODEL PERFORMANCE REPORT GENERATOR\")\nprint(\"=\"*80)\n\n# ============================================================================\n# CONFIGURATION\n# ============================================================================\n\nCONFIG = {\n    # Model checkpoints\n    'efficientnet_checkpoint': '/kaggle/working/efficientnet_300patients/best_model.pth',\n    'fewshot_checkpoint': '/kaggle/working/step4_few_shot/few_shot_model.pth',\n    \n    # Dataset\n    'dataset_dir': '/kaggle/working/mini_dataset_v2',\n    \n    # Output\n    'report_dir': '/kaggle/working/performance_report',\n    \n    # Report metadata\n    'project_title': 'Cervical Spine Fracture Detection using Deep Learning',\n    'author': 'Sri Somesh, Sai chakrith, Kaveen kawsal, pranav srinivasan',\n    'institution': 'Amrita Vishwa Vidhyapeetham, Coimbatore',\n    'dataset_name': 'RSNA Cervical Spine Fracture Detection',\n    \n    # Device\n    'device': 'cuda' if torch.cuda.is_available() else 'cpu',\n}\n\nos.makedirs(CONFIG['report_dir'], exist_ok=True)\n\n# ============================================================================\n# MODEL ARCHITECTURES (Same as before)\n# ============================================================================\n\nclass SqueezeExcitation3D(nn.Module):\n    def __init__(self, channels, reduction=4):\n        super().__init__()\n        reduced = max(1, channels // reduction)\n        self.se = nn.Sequential(\n            nn.AdaptiveAvgPool3d(1),\n            nn.Conv3d(channels, reduced, 1),\n            nn.ReLU(inplace=True),\n            nn.Conv3d(reduced, channels, 1),\n            nn.Sigmoid()\n        )\n    \n    def forward(self, x):\n        return x * self.se(x)\n\n\nclass MBConvBlock(nn.Module):\n    def __init__(self, in_channels, out_channels, expand_ratio=6, stride=1):\n        super().__init__()\n        self.use_residual = (stride == 1 and in_channels == out_channels)\n        hidden = in_channels * expand_ratio\n        \n        layers = []\n        \n        if expand_ratio != 1:\n            layers.extend([\n                nn.Conv3d(in_channels, hidden, 1, bias=False),\n                nn.BatchNorm3d(hidden),\n                nn.ReLU(inplace=True)\n            ])\n        \n        layers.extend([\n            nn.Conv3d(hidden, hidden, 3, stride=stride, padding=1, \n                     groups=hidden, bias=False),\n            nn.BatchNorm3d(hidden),\n            nn.ReLU(inplace=True)\n        ])\n        \n        layers.append(SqueezeExcitation3D(hidden))\n        \n        layers.extend([\n            nn.Conv3d(hidden, out_channels, 1, bias=False),\n            nn.BatchNorm3d(out_channels)\n        ])\n        \n        self.conv = nn.Sequential(*layers)\n        self.dropout = nn.Dropout3d(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.conv(x))\n        else:\n            return self.conv(x)\n\n\nclass LightweightEfficientNet3D(nn.Module):\n    def __init__(self, num_classes=1, dropout=0.4):\n        super().__init__()\n        \n        self.stem = nn.Sequential(\n            nn.Conv3d(1, 32, kernel_size=3, stride=2, padding=1, bias=False),\n            nn.BatchNorm3d(32),\n            nn.ReLU(inplace=True)\n        )\n        \n        self.blocks = nn.Sequential(\n            MBConvBlock(32, 16, expand_ratio=1, stride=1),\n            MBConvBlock(16, 24, expand_ratio=6, stride=2),\n            MBConvBlock(24, 24, expand_ratio=6, stride=1),\n            MBConvBlock(24, 40, expand_ratio=6, stride=2),\n            MBConvBlock(40, 40, expand_ratio=6, stride=1),\n            MBConvBlock(40, 80, expand_ratio=6, stride=2),\n            MBConvBlock(80, 80, expand_ratio=6, stride=1),\n        )\n        \n        self.head = nn.Sequential(\n            nn.Conv3d(80, 320, 1, bias=False),\n            nn.BatchNorm3d(320),\n            nn.ReLU(inplace=True),\n            nn.AdaptiveAvgPool3d(1)\n        )\n        \n        self.dropout = nn.Dropout(dropout)\n        self.fc = nn.Linear(320, num_classes)\n    \n    def forward(self, x):\n        x = self.stem(x)\n        x = self.blocks(x)\n        x = self.head(x)\n        x = x.view(x.size(0), -1)\n        x = self.dropout(x)\n        x = self.fc(x)\n        return x.squeeze(-1)\n\n\nclass EmbeddingNetwork3D(nn.Module):\n    def __init__(self, embedding_dim=64, dropout=0.5):\n        super().__init__()\n        \n        self.conv1 = nn.Conv3d(1, 16, kernel_size=3, stride=2, padding=1)\n        self.bn1 = nn.BatchNorm3d(16)\n        self.drop1 = nn.Dropout3d(dropout * 0.3)\n        \n        self.conv2 = nn.Conv3d(16, 32, kernel_size=3, stride=2, padding=1)\n        self.bn2 = nn.BatchNorm3d(32)\n        self.drop2 = nn.Dropout3d(dropout * 0.5)\n        \n        self.conv3 = nn.Conv3d(32, 64, kernel_size=3, stride=2, padding=1)\n        self.bn3 = nn.BatchNorm3d(64)\n        self.drop3 = nn.Dropout3d(dropout * 0.7)\n        \n        self.pool = nn.AdaptiveAvgPool3d(1)\n        self.dropout = nn.Dropout(dropout)\n        self.fc = nn.Linear(64, embedding_dim)\n    \n    def forward(self, x):\n        x = F.relu(self.bn1(self.conv1(x)))\n        x = self.drop1(x)\n        \n        x = F.relu(self.bn2(self.conv2(x)))\n        x = self.drop2(x)\n        \n        x = F.relu(self.bn3(self.conv3(x)))\n        x = self.drop3(x)\n        \n        x = self.pool(x)\n        x = x.view(x.size(0), -1)\n        \n        x = self.dropout(x)\n        embeddings = self.fc(x)\n        \n        embeddings = F.normalize(embeddings, p=2, dim=1)\n        \n        return embeddings\n\n\n# ============================================================================\n# INFERENCE ENGINE\n# ============================================================================\n\nclass ModelEvaluator:\n    \"\"\"Evaluate all models and collect predictions\"\"\"\n    \n    def __init__(self, efficientnet_path, fewshot_path, device):\n        self.device = torch.device(device)\n        \n        print(f\"\\n🔧 Loading models for evaluation...\")\n        \n        # Load EfficientNet\n        eff_checkpoint = torch.load(efficientnet_path, map_location=device, weights_only=False)\n        self.efficientnet = LightweightEfficientNet3D(num_classes=1, dropout=0.4)\n        self.efficientnet.load_state_dict(eff_checkpoint['model_state_dict'])\n        self.efficientnet = self.efficientnet.to(device).eval()\n        self.eff_history = eff_checkpoint.get('history', {})\n        print(f\"  ✓ EfficientNet loaded\")\n        \n        # Load Few-Shot\n        few_checkpoint = torch.load(fewshot_path, map_location=device, weights_only=False)\n        self.fewshot_net = EmbeddingNetwork3D(embedding_dim=64, dropout=0.5)\n        self.fewshot_net.load_state_dict(few_checkpoint['model_state_dict'])\n        self.fewshot_net = self.fewshot_net.to(device).eval()\n        self.few_history = few_checkpoint.get('history', {})\n        print(f\"  ✓ Few-Shot loaded\")\n        \n        self.support_prototypes = None\n    \n    def set_support_prototypes(self, support_volumes, support_labels):\n        \"\"\"Setup few-shot support set\"\"\"\n        with torch.no_grad():\n            support_volumes = support_volumes.to(self.device)\n            embeddings = self.fewshot_net(support_volumes)\n            \n            prototypes = []\n            for class_idx in [0, 1]:\n                mask = support_labels == class_idx\n                class_emb = embeddings[mask]\n                prototype = class_emb.mean(dim=0) if len(class_emb) > 0 else torch.zeros(embeddings.shape[1], device=self.device)\n                prototypes.append(prototype)\n            \n            self.support_prototypes = torch.stack(prototypes)\n    \n    def evaluate_all(self, dataloader):\n        \"\"\"Get predictions from all models\"\"\"\n        \n        all_labels = []\n        eff_preds = []\n        few_preds = []\n        \n        print(f\"\\n  Evaluating on {len(dataloader.dataset)} samples...\")\n        \n        with torch.no_grad():\n            for volumes, labels in dataloader:\n                volumes = volumes.to(self.device)\n                \n                # EfficientNet\n                eff_logits = self.efficientnet(volumes)\n                eff_probs = torch.sigmoid(eff_logits).cpu().numpy()\n                \n                # Few-Shot\n                embeddings = self.fewshot_net(volumes)\n                distances = torch.cdist(embeddings, self.support_prototypes.unsqueeze(0)).squeeze(0)\n                logits = -distances\n                few_probs = F.softmax(logits, dim=1)[:, 1].cpu().numpy()\n                \n                all_labels.extend(labels.numpy())\n                eff_preds.extend(eff_probs)\n                few_preds.extend(few_probs)\n        \n        all_labels = np.array(all_labels)\n        eff_preds = np.array(eff_preds)\n        few_preds = np.array(few_preds)\n        \n        # Ensemble (weighted average)\n        ensemble_preds = 0.7 * eff_preds + 0.3 * few_preds\n        \n        return {\n            'labels': all_labels,\n            'efficientnet': eff_preds,\n            'fewshot': few_preds,\n            'ensemble': ensemble_preds\n        }\n\n\n# ============================================================================\n# METRICS CALCULATOR\n# ============================================================================\n\nclass MetricsCalculator:\n    \"\"\"Calculate comprehensive metrics\"\"\"\n    \n    @staticmethod\n    def calculate_all_metrics(y_true, y_pred_probs, threshold=0.5):\n        \"\"\"Calculate all classification metrics\"\"\"\n        \n        y_pred = (y_pred_probs > threshold).astype(int)\n        \n        metrics = {\n            # Threshold-based metrics\n            'accuracy': accuracy_score(y_true, y_pred),\n            'precision': precision_score(y_true, y_pred, zero_division=0),\n            'recall': recall_score(y_true, y_pred, zero_division=0),\n            'f1_score': f1_score(y_true, y_pred, zero_division=0),\n            'specificity': 0.0,\n            \n            # Threshold-free metrics\n            'auc': roc_auc_score(y_true, y_pred_probs) if len(np.unique(y_true)) > 1 else 0.5,\n            'average_precision': average_precision_score(y_true, y_pred_probs) if len(np.unique(y_true)) > 1 else 0.5,\n        }\n        \n        # Confusion matrix\n        cm = confusion_matrix(y_true, y_pred)\n        if cm.shape == (2, 2):\n            tn, fp, fn, tp = cm.ravel()\n            metrics['specificity'] = tn / (tn + fp) if (tn + fp) > 0 else 0.0\n            metrics['sensitivity'] = tp / (tp + fn) if (tp + fn) > 0 else 0.0\n            metrics['npv'] = tn / (tn + fn) if (tn + fn) > 0 else 0.0\n            metrics['ppv'] = tp / (tp + fp) if (tp + fp) > 0 else 0.0\n            metrics['confusion_matrix'] = cm\n        \n        return metrics\n    \n    @staticmethod\n    def compare_models(labels, preds_dict):\n        \"\"\"Compare multiple models statistically\"\"\"\n        \n        results = {}\n        \n        for model_name, preds in preds_dict.items():\n            results[model_name] = MetricsCalculator.calculate_all_metrics(labels, preds)\n        \n        # Statistical comparison (DeLong test for AUC)\n        comparisons = {}\n        models = list(preds_dict.keys())\n        \n        for i, model1 in enumerate(models):\n            for model2 in models[i+1:]:\n                auc1 = results[model1]['auc']\n                auc2 = results[model2]['auc']\n                \n                # Simple z-test approximation\n                se1 = np.sqrt(auc1 * (1 - auc1) / len(labels))\n                se2 = np.sqrt(auc2 * (1 - auc2) / len(labels))\n                se_diff = np.sqrt(se1**2 + se2**2)\n                \n                z_score = (auc1 - auc2) / se_diff if se_diff > 0 else 0\n                p_value = 2 * (1 - stats.norm.cdf(abs(z_score)))\n                \n                comparisons[f'{model1}_vs_{model2}'] = {\n                    'auc_diff': auc1 - auc2,\n                    'z_score': z_score,\n                    'p_value': p_value,\n                    'significant': p_value < 0.05\n                }\n        \n        return results, comparisons\n\n\n# ============================================================================\n# REPORT GENERATOR\n# ============================================================================\n\nclass ReportGenerator:\n    \"\"\"Generate professional reports\"\"\"\n    \n    def __init__(self, config, results, comparisons, predictions):\n        self.config = config\n        self.results = results\n        self.comparisons = comparisons\n        self.predictions = predictions\n        self.report_dir = config['report_dir']\n    \n    def generate_summary_table(self):\n        \"\"\"Create summary table of all metrics\"\"\"\n        \n        models = list(self.results.keys())\n        metrics_names = ['accuracy', 'precision', 'recall', 'f1_score', \n                        'specificity', 'sensitivity', 'auc', 'average_precision']\n        \n        data = []\n        for model in models:\n            row = {'Model': model.upper()}\n            for metric in metrics_names:\n                value = self.results[model].get(metric, 0.0)\n                row[metric] = value\n            data.append(row)\n        \n        df = pd.DataFrame(data)\n        \n        # Save CSV\n        csv_path = os.path.join(self.report_dir, 'metrics_summary.csv')\n        df.to_csv(csv_path, index=False)\n        print(f\"  ✓ Saved: metrics_summary.csv\")\n        \n        # Save LaTeX\n        latex_path = os.path.join(self.report_dir, 'metrics_summary.tex')\n        with open(latex_path, 'w') as f:\n            latex_str = df.to_latex(index=False, float_format=\"%.4f\",\n                                   caption=\"Model Performance Comparison\",\n                                   label=\"tab:metrics\")\n            f.write(latex_str)\n        print(f\"  ✓ Saved: metrics_summary.tex\")\n        \n        return df\n    \n    def generate_comparison_table(self):\n        \"\"\"Create statistical comparison table\"\"\"\n        \n        data = []\n        for comp_name, comp_data in self.comparisons.items():\n            models = comp_name.split('_vs_')\n            data.append({\n                'Comparison': f\"{models[0].upper()} vs {models[1].upper()}\",\n                'AUC Difference': comp_data['auc_diff'],\n                'Z-Score': comp_data['z_score'],\n                'P-Value': comp_data['p_value'],\n                'Significant (p<0.05)': '✓' if comp_data['significant'] else '✗'\n            })\n        \n        df = pd.DataFrame(data)\n        \n        csv_path = os.path.join(self.report_dir, 'statistical_comparison.csv')\n        df.to_csv(csv_path, index=False)\n        print(f\"  ✓ Saved: statistical_comparison.csv\")\n        \n        return df\n    \n    def plot_roc_curves(self):\n        \"\"\"Plot ROC curves for all models\"\"\"\n        \n        fig, ax = plt.subplots(figsize=(10, 8))\n        \n        colors = {'efficientnet': '#3498db', 'fewshot': '#e74c3c', 'ensemble': '#2ecc71'}\n        \n        for model_name in self.results.keys():\n            preds = self.predictions[model_name]\n            labels = self.predictions['labels']\n            \n            fpr, tpr, _ = roc_curve(labels, preds)\n            auc = self.results[model_name]['auc']\n            \n            ax.plot(fpr, tpr, linewidth=3, label=f'{model_name.upper()} (AUC={auc:.4f})',\n                   color=colors.get(model_name, 'black'))\n        \n        ax.plot([0, 1], [0, 1], 'k--', linewidth=2, alpha=0.5, label='Random')\n        \n        ax.set_xlabel('False Positive Rate', fontsize=14, fontweight='bold')\n        ax.set_ylabel('True Positive Rate', fontsize=14, fontweight='bold')\n        ax.set_title('ROC Curves Comparison', fontsize=16, fontweight='bold')\n        ax.legend(fontsize=12, loc='lower right')\n        ax.grid(True, alpha=0.3)\n        \n        plt.tight_layout()\n        save_path = os.path.join(self.report_dir, 'roc_curves.png')\n        plt.savefig(save_path, dpi=300, bbox_inches='tight')\n        plt.close()\n        \n        print(f\"  ✓ Saved: roc_curves.png\")\n    \n    def plot_pr_curves(self):\n        \"\"\"Plot Precision-Recall curves\"\"\"\n        \n        fig, ax = plt.subplots(figsize=(10, 8))\n        \n        colors = {'efficientnet': '#3498db', 'fewshot': '#e74c3c', 'ensemble': '#2ecc71'}\n        \n        for model_name in self.results.keys():\n            preds = self.predictions[model_name]\n            labels = self.predictions['labels']\n            \n            precision, recall, _ = precision_recall_curve(labels, preds)\n            ap = self.results[model_name]['average_precision']\n            \n            ax.plot(recall, precision, linewidth=3, \n                   label=f'{model_name.upper()} (AP={ap:.4f})',\n                   color=colors.get(model_name, 'black'))\n        \n        ax.set_xlabel('Recall', fontsize=14, fontweight='bold')\n        ax.set_ylabel('Precision', fontsize=14, fontweight='bold')\n        ax.set_title('Precision-Recall Curves', fontsize=16, fontweight='bold')\n        ax.legend(fontsize=12, loc='best')\n        ax.grid(True, alpha=0.3)\n        ax.set_xlim([0, 1])\n        ax.set_ylim([0, 1.05])\n        \n        plt.tight_layout()\n        save_path = os.path.join(self.report_dir, 'pr_curves.png')\n        plt.savefig(save_path, dpi=300, bbox_inches='tight')\n        plt.close()\n        \n        print(f\"  ✓ Saved: pr_curves.png\")\n    \n    def plot_confusion_matrices(self):\n        \"\"\"Plot confusion matrices for all models\"\"\"\n        \n        models = list(self.results.keys())\n        fig, axes = plt.subplots(1, len(models), figsize=(15, 5))\n        \n        if len(models) == 1:\n            axes = [axes]\n        \n        for ax, model_name in zip(axes, models):\n            cm = self.results[model_name].get('confusion_matrix', np.zeros((2, 2)))\n            \n            sns.heatmap(cm, annot=True, fmt='d', cmap='Blues',\n                       xticklabels=['Normal', 'Fracture'],\n                       yticklabels=['Normal', 'Fracture'],\n                       ax=ax, cbar=False,\n                       annot_kws={'fontsize': 14, 'fontweight': 'bold'})\n            \n            ax.set_title(f'{model_name.upper()}', fontsize=14, fontweight='bold')\n            ax.set_xlabel('Predicted', fontsize=12)\n            ax.set_ylabel('True', fontsize=12)\n        \n        plt.tight_layout()\n        save_path = os.path.join(self.report_dir, 'confusion_matrices.png')\n        plt.savefig(save_path, dpi=300, bbox_inches='tight')\n        plt.close()\n        \n        print(f\"  ✓ Saved: confusion_matrices.png\")\n    \n    def plot_metrics_comparison(self):\n        \"\"\"Bar chart comparing all metrics\"\"\"\n        \n        models = list(self.results.keys())\n        metrics = ['accuracy', 'precision', 'recall', 'f1_score', 'auc']\n        \n        x = np.arange(len(metrics))\n        width = 0.25\n        \n        fig, ax = plt.subplots(figsize=(14, 8))\n        \n        colors = ['#3498db', '#e74c3c', '#2ecc71']\n        \n        for i, model in enumerate(models):\n            values = [self.results[model].get(m, 0.0) for m in metrics]\n            offset = width * (i - len(models)/2 + 0.5)\n            bars = ax.bar(x + offset, values, width, label=model.upper(),\n                         color=colors[i], edgecolor='black', linewidth=1.5)\n            \n            # Add value labels\n            for bar, val in zip(bars, values):\n                height = bar.get_height()\n                ax.text(bar.get_x() + bar.get_width()/2, height + 0.01,\n                       f'{val:.3f}', ha='center', va='bottom',\n                       fontsize=9, fontweight='bold')\n        \n        ax.set_ylabel('Score', fontsize=14, fontweight='bold')\n        ax.set_title('Model Performance Metrics Comparison', fontsize=16, fontweight='bold')\n        ax.set_xticks(x)\n        ax.set_xticklabels([m.replace('_', ' ').title() for m in metrics], fontsize=12)\n        ax.legend(fontsize=12)\n        ax.set_ylim([0, 1.1])\n        ax.grid(True, alpha=0.3, axis='y')\n        \n        plt.tight_layout()\n        save_path = os.path.join(self.report_dir, 'metrics_comparison.png')\n        plt.savefig(save_path, dpi=300, bbox_inches='tight')\n        plt.close()\n        \n        print(f\"  ✓ Saved: metrics_comparison.png\")\n    \n    def generate_executive_summary(self, summary_df, comparison_df):\n        \"\"\"Generate text-based executive summary\"\"\"\n        \n        best_model = summary_df.loc[summary_df['auc'].idxmax(), 'Model']\n        best_auc = summary_df['auc'].max()\n        \n        summary = f\"\"\"\n{'='*80}\nEXECUTIVE SUMMARY\n{'='*80}\n\nProject: {self.config['project_title']}\nAuthor: {self.config['author']}\nInstitution: {self.config['institution']}\nDate: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}\nDataset: {self.config['dataset_name']}\n\n{'='*80}\nKEY FINDINGS\n{'='*80}\n\n1. BEST PERFORMING MODEL\n   Model: {best_model}\n   AUC Score: {best_auc:.4f}\n   \n2. MODEL PERFORMANCE SUMMARY\n\n{summary_df.to_string(index=False)}\n\n3. STATISTICAL COMPARISONS\n\n{comparison_df.to_string(index=False)}\n\n{'='*80}\nRECOMMENDATIONS\n{'='*80}\n\n\"\"\"\n        \n        # Add recommendations based on results\n        if best_auc > 0.80:\n            summary += \"✓ EXCELLENT PERFORMANCE: Model ready for clinical validation\\n\"\n        elif best_auc > 0.70:\n            summary += \"✓ GOOD PERFORMANCE: Consider ensemble or additional data\\n\"\n        else:\n            summary += \"⚠ MODERATE PERFORMANCE: Requires improvement before deployment\\n\"\n        \n        # Check if ensemble is best\n        if 'ENSEMBLE' in best_model.upper():\n            summary += \"✓ Ensemble outperforms individual models - recommended approach\\n\"\n        \n        # Check statistical significance\n        sig_comparisons = [k for k, v in self.comparisons.items() if v['significant']]\n        if sig_comparisons:\n            summary += f\"✓ {len(sig_comparisons)} statistically significant differences detected\\n\"\n        \n        summary += f\"\"\"\n{'='*80}\nOUTPUT FILES GENERATED\n{'='*80}\n\n📊 Tables:\n   - metrics_summary.csv (All performance metrics)\n   - metrics_summary.tex (LaTeX table for thesis)\n   - statistical_comparison.csv (Statistical tests)\n\n📈 Visualizations:\n   - roc_curves.png (ROC curve comparison)\n   - pr_curves.png (Precision-Recall curves)\n   - confusion_matrices.png (Confusion matrices)\n   - metrics_comparison.png (Bar chart comparison)\n   - comprehensive_report.pdf (Complete PDF report)\n\n📄 Reports:\n   - executive_summary.txt (This file)\n   - full_report.json (JSON with all results)\n\n{'='*80}\nEND OF REPORT\n{'='*80}\n\"\"\"\n        \n        # Save summary\n        summary_path = os.path.join(self.report_dir, 'executive_summary.txt')\n        with open(summary_path, 'w') as f:\n            f.write(summary)\n        \n        print(f\"  ✓ Saved: executive_summary.txt\")\n        \n        return summary\n    \n    def generate_pdf_report(self):\n        \"\"\"Generate comprehensive PDF report\"\"\"\n        \n        pdf_path = os.path.join(self.report_dir, 'comprehensive_report.pdf')\n        \n        with PdfPages(pdf_path) as pdf:\n            # Page 1: Title and Summary\n            fig = plt.figure(figsize=(11, 8.5))\n            fig.text(0.5, 0.9, self.config['project_title'], \n                    ha='center', fontsize=20, fontweight='bold')\n            fig.text(0.5, 0.85, f\"Performance Analysis Report\",\n                    ha='center', fontsize=16)\n            fig.text(0.5, 0.80, f\"Generated: {datetime.now().strftime('%Y-%m-%d')}\",\n                    ha='center', fontsize=12)\n            \n            # Add metrics table\n            models = list(self.results.keys())\n            table_data = [['Model', 'AUC', 'Accuracy', 'Precision', 'Recall', 'F1']]\n            \n            for model in models:\n                row = [\n                    model.upper(),\n                    f\"{self.results[model]['auc']:.4f}\",\n                    f\"{self.results[model]['accuracy']:.4f}\",\n                    f\"{self.results[model]['precision']:.4f}\",\n                    f\"{self.results[model]['recall']:.4f}\",\n                    f\"{self.results[model]['f1_score']:.4f}\"\n                ]\n                table_data.append(row)\n            \n            table = plt.table(cellText=table_data, cellLoc='center',\n                            loc='center', bbox=[0.1, 0.3, 0.8, 0.4])\n            table.auto_set_font_size(False)\n            table.set_fontsize(10)\n            table.scale(1, 2)\n            \n            # Style header\n            for i in range(len(table_data[0])):\n                table[(0, i)].set_facecolor('#3498db')\n                table[(0, i)].set_text_props(weight='bold', color='white')\n            \n            plt.axis('off')\n            pdf.savefig(fig, bbox_inches='tight')\n            plt.close()\n            \n            # Page 2: ROC Curves\n            img = plt.imread(os.path.join(self.report_dir, 'roc_curves.png'))\n            fig = plt.figure(figsize=(11, 8.5))\n            plt.imshow(img)\n            plt.axis('off')\n            pdf.savefig(fig, bbox_inches='tight')\n            plt.close()\n            \n            # Page 3: Confusion Matrices\n            img = plt.imread(os.path.join(self.report_dir, 'confusion_matrices.png'))\n            fig = plt.figure(figsize=(11, 8.5))\n            plt.imshow(img)\n            plt.axis('off')\n            pdf.savefig(fig, bbox_inches='tight')\n            plt.close()\n            \n            # Page 4: Metrics Comparison\n            img = plt.imread(os.path.join(self.report_dir, 'metrics_comparison.png'))\n            fig = plt.figure(figsize=(11, 8.5))\n            plt.imshow(img)\n            plt.axis('off')\n            pdf.savefig(fig, bbox_inches='tight')\n            plt.close()\n        \n        print(f\"  ✓ Saved: comprehensive_report.pdf\")\n    \n    def save_json_report(self):\n        \"\"\"Save complete results as JSON\"\"\"\n        \n        # Helper function to convert numpy/bool types to JSON-serializable types\n        def convert_to_serializable(obj):\n            if isinstance(obj, (np.integer, np.int64, np.int32)):\n                return int(obj)\n            elif isinstance(obj, (np.floating, np.float64, np.float32)):\n                return float(obj)\n            elif isinstance(obj, (np.bool_, bool)):\n                return bool(obj)\n            elif isinstance(obj, np.ndarray):\n                return obj.tolist()\n            elif isinstance(obj, dict):\n                return {k: convert_to_serializable(v) for k, v in obj.items()}\n            elif isinstance(obj, (list, tuple)):\n                return [convert_to_serializable(item) for item in obj]\n            else:\n                return obj\n        \n        report = {\n            'metadata': {\n                'project_title': self.config['project_title'],\n                'author': self.config['author'],\n                'institution': self.config['institution'],\n                'dataset': self.config['dataset_name'],\n                'generated_at': datetime.now().isoformat(),\n            },\n            'model_results': {},\n            'comparisons': {},\n        }\n        \n        # Convert model results\n        for model_name, metrics in self.results.items():\n            report['model_results'][model_name] = {\n                k: convert_to_serializable(v)\n                for k, v in metrics.items()\n                if k != 'confusion_matrix'\n            }\n        \n        # Convert comparisons\n        for comp_name, comp_data in self.comparisons.items():\n            report['comparisons'][comp_name] = convert_to_serializable(comp_data)\n        \n        json_path = os.path.join(self.report_dir, 'full_report.json')\n        with open(json_path, 'w') as f:\n            json.dump(report, f, indent=2)\n        \n        print(f\"  ✓ Saved: full_report.json\")\n    \n    def generate_all(self):\n        \"\"\"Generate complete report\"\"\"\n        \n        print(f\"\\n{'='*80}\")\n        print(\"📊 GENERATING COMPREHENSIVE REPORT\")\n        print(f\"{'='*80}\\n\")\n        \n        # Generate tables\n        print(\"  Creating summary tables...\")\n        summary_df = self.generate_summary_table()\n        comparison_df = self.generate_comparison_table()\n        \n        # Generate plots\n        print(\"\\n  Creating visualizations...\")\n        self.plot_roc_curves()\n        self.plot_pr_curves()\n        self.plot_confusion_matrices()\n        self.plot_metrics_comparison()\n        \n        # Generate reports\n        print(\"\\n  Generating reports...\")\n        self.generate_executive_summary(summary_df, comparison_df)\n        self.save_json_report()\n        self.generate_pdf_report()\n        \n        print(f\"\\n{'='*80}\")\n        print(\"✅ REPORT GENERATION COMPLETE\")\n        print(f\"{'='*80}\")\n        \n        return summary_df\n\n\n# ============================================================================\n# MAIN EXECUTION\n# ============================================================================\n\ndef main():\n    print(\"\\n\" + \"=\"*80)\n    print(\"🚀 STARTING COMPREHENSIVE REPORT GENERATION\")\n    print(\"=\"*80)\n    \n    try:\n        # Check model files\n        if not os.path.exists(CONFIG['efficientnet_checkpoint']):\n            print(f\"\\n❌ EfficientNet not found: {CONFIG['efficientnet_checkpoint']}\")\n            return\n        \n        if not os.path.exists(CONFIG['fewshot_checkpoint']):\n            print(f\"\\n❌ Few-Shot not found: {CONFIG['fewshot_checkpoint']}\")\n            return\n        \n        # Load dataset\n        print(f\"\\n📂 Loading dataset...\")\n        metadata_path = os.path.join(CONFIG['dataset_dir'], 'metadata.csv')\n        volumes_dir = os.path.join(CONFIG['dataset_dir'], 'volumes')\n        \n        if not os.path.exists(metadata_path):\n            print(f\"\\n❌ Dataset not found: {metadata_path}\")\n            return\n        \n        metadata_df = pd.read_csv(metadata_path)\n        \n        # Create simple dataset\n        class SimpleDataset(Dataset):\n            def __init__(self, df, volumes_dir):\n                self.df = df\n                self.volumes_dir = volumes_dir\n            \n            def __len__(self):\n                return len(self.df)\n            \n            def __getitem__(self, idx):\n                row = self.df.iloc[idx]\n                volume = np.load(os.path.join(self.volumes_dir, f\"{row['patient_id']}.npy\"))\n                volume = volume[np.newaxis, ...].astype(np.float32)\n                label = row['has_fracture']\n                return torch.from_numpy(volume), torch.tensor(label, dtype=torch.float32)\n        \n        dataset = SimpleDataset(metadata_df, volumes_dir)\n        dataloader = DataLoader(dataset, batch_size=8, shuffle=False, num_workers=2)\n        \n        print(f\"  ✓ Loaded {len(dataset)} patients\")\n        \n        # Initialize evaluator\n        evaluator = ModelEvaluator(\n            CONFIG['efficientnet_checkpoint'],\n            CONFIG['fewshot_checkpoint'],\n            CONFIG['device']\n        )\n        \n        # Setup few-shot support\n        print(f\"\\n  Setting up few-shot support set...\")\n        n_sup = 3\n        frac_df = metadata_df[metadata_df['has_fracture'] == 1].head(n_sup)\n        norm_df = metadata_df[metadata_df['has_fracture'] == 0].head(n_sup)\n        support_df = pd.concat([frac_df, norm_df])\n        \n        support_volumes = []\n        support_labels = []\n        \n        for _, row in support_df.iterrows():\n            vol = np.load(os.path.join(volumes_dir, f\"{row['patient_id']}.npy\"))\n            vol = vol[np.newaxis, ...].astype(np.float32)\n            support_volumes.append(torch.from_numpy(vol))\n            support_labels.append(row['has_fracture'])\n        \n        support_volumes = torch.stack(support_volumes)\n        support_labels = torch.tensor(support_labels)\n        \n        evaluator.set_support_prototypes(support_volumes, support_labels)\n        \n        # Evaluate all models\n        predictions = evaluator.evaluate_all(dataloader)\n        \n        # Calculate metrics\n        print(f\"\\n  Calculating comprehensive metrics...\")\n        results, comparisons = MetricsCalculator.compare_models(\n            predictions['labels'],\n            {k: v for k, v in predictions.items() if k != 'labels'}\n        )\n        \n        # Generate reports\n        report_gen = ReportGenerator(CONFIG, results, comparisons, predictions)\n        summary_df = report_gen.generate_all()\n        \n        # Print summary to console\n        print(f\"\\n{'='*80}\")\n        print(\"📈 PERFORMANCE SUMMARY\")\n        print(f\"{'='*80}\\n\")\n        \n        print(summary_df.to_string(index=False))\n        \n        print(f\"\\n{'='*80}\")\n        print(\"📁 ALL OUTPUTS SAVED TO:\")\n        print(f\"{'='*80}\")\n        print(f\"\\n  {CONFIG['report_dir']}/\")\n        print(f\"  ├── executive_summary.txt         ← Read this first!\")\n        print(f\"  ├── comprehensive_report.pdf      ← Complete PDF report\")\n        print(f\"  ├── metrics_summary.csv           ← For Excel/analysis\")\n        print(f\"  ├── metrics_summary.tex           ← For LaTeX thesis\")\n        print(f\"  ├── statistical_comparison.csv    ← Statistical tests\")\n        print(f\"  ├── full_report.json              ← Machine-readable\")\n        print(f\"  ├── roc_curves.png                ← Publication figure\")\n        print(f\"  ├── pr_curves.png                 ← Publication figure\")\n        print(f\"  ├── confusion_matrices.png        ← Publication figure\")\n        print(f\"  └── metrics_comparison.png        ← Publication figure\")\n        \n        print(f\"\\n{'='*80}\")\n        print(\"✅ SUCCESS! Professional report generated\")\n        print(f\"{'='*80}\")\n        \n        print(f\"\\n💡 NEXT STEPS:\")\n        print(f\"   1. Review executive_summary.txt for key findings\")\n        print(f\"   2. Use comprehensive_report.pdf for presentations\")\n        print(f\"   3. Import metrics_summary.csv into your thesis\")\n        print(f\"   4. Use .png files in papers/presentations\")\n        print(f\"   5. Include metrics_summary.tex in LaTeX documents\")\n        \n    except Exception as e:\n        print(f\"\\n{'='*80}\")\n        print(\"❌ ERROR\")\n        print(f\"{'='*80}\")\n        print(f\"\\nError: {str(e)}\")\n        import traceback\n        traceback.print_exc()\n\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T04:24:28.836523Z","iopub.execute_input":"2026-01-17T04:24:28.837291Z","iopub.status.idle":"2026-01-17T04:24:42.088429Z","shell.execute_reply.started":"2026-01-17T04:24:28.837261Z","shell.execute_reply":"2026-01-17T04:24:42.087723Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!zip -r all_outputs.zip /kaggle/working\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T04:25:05.169637Z","iopub.execute_input":"2026-01-17T04:25:05.169915Z","iopub.status.idle":"2026-01-17T04:26:11.611015Z","shell.execute_reply.started":"2026-01-17T04:25:05.16989Z","shell.execute_reply":"2026-01-17T04:26:11.610328Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"!ls -lh /kaggle/working\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2026-01-17T04:30:35.794672Z","iopub.execute_input":"2026-01-17T04:30:35.795021Z","iopub.status.idle":"2026-01-17T04:30:35.969335Z","shell.execute_reply.started":"2026-01-17T04:30:35.794993Z","shell.execute_reply":"2026-01-17T04:30:35.968557Z"}},"outputs":[],"execution_count":null}]}