{"metadata":{"kernelspec":{"display_name":"Python 3","language":"python","name":"python3"},"language_info":{"codemirror_mode":{"name":"ipython","version":3},"file_extension":".py","mimetype":"text/x-python","name":"python","nbconvert_exporter":"python","pygments_lexer":"ipython3","version":"3.11.13"},"kaggle":{"accelerator":"nvidiaTeslaT4","dataSources":[{"sourceId":99552,"databundleVersionId":13762876,"sourceType":"competition"},{"sourceId":12637336,"sourceType":"datasetVersion","datasetId":7981664}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"## References\n- https://www.kaggle.com/code/dennisfong/dicom-pngs-for-rsna-intracranial-aneurysm/notebook\n\n## Notebooks\n- My Train Notebook: here\n\n## Training Overview - 2D U-Net Segmentation\n### What to Train:\n- 2D U-Net Architecture for Medical Image Segmentation\n- Multi-class Segmentation (13 anatomical locations + Background)\n\n### What to Train With:\n- 2D slices from 3D volumes (224×224)\n- Segmentation masks from NII files\n- Dice Loss + BCE Loss combination\n\n### Key Features:\n- **True Patient Separation**: DICOM StudyInstanceUID-based cross-validation for patient-level separation\n- **3D to 2D Conversion**: Convert NII segmentation masks to 2D slices\n- **Smart Data Matching**: Match segmentation files with corresponding image series\n- **CLAHE Contrast Adaptation**: Modality-specific enhancement for CTA/MRA/MRI variations\n- **Strong Augmentation**: 15° rotation, elastic transforms, noise simulation for scanner robustness\n- **Robust Percentile Normalization**: Outlier-resistant preprocessing using 1st-99th percentile clipping\n- **Medical Metadata Integration**: Patient age and sex features for enhanced segmentation\n- **Modality-specific Windowing**: Optimized intensity windows (CTA: 50/350, MRA: 600/1200, MRI: 40/80)\n- **Mixed Precision Training**: GPU-optimized training with gradient accumulation\n- **LRU Caching**: Performance optimization for frequently accessed DICOM data\n\n### Improvements:\n- **Patient Leakage Prevention**: DICOM metadata extraction ensures no patient overlap between train/validation\n- **Segmentation-focused**: Direct pixel-level prediction instead of classification\n\n### Evaluation Metric:\n- **Dice Score**: Primary metric for segmentation quality\n- **IoU (Intersection over Union)**: Secondary metric\n- **Pixel Accuracy**: Overall accuracy metric\n\n### Expected Performance:\n- **Improved Segmentation**: Direct pixel-level prediction for better localization\n- **Better Generalization**: Strategic sampling and robust preprocessing for real-world variation\n- **Reduced CV/LB Gap**: From ~0.44 gap to healthy 0.10-0.15 range through proper validation","metadata":{}},{"cell_type":"code","source":"# Environment setup and library imports\nimport os\nimport glob\nimport random\nimport warnings\nimport numpy as np\nimport pandas as pd\nimport cv2\nimport functools\nfrom pathlib import Path\nfrom tqdm import tqdm\nimport matplotlib.pyplot as plt\nimport seaborn as sns\nfrom typing import List, Tuple, Optional\nimport gc\n\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader\nfrom torch.optim import AdamW\nfrom torch.optim.lr_scheduler import CosineAnnealingLR\n\nimport timm\nimport albumentations as A\nfrom albumentations.pytorch import ToTensorV2\nfrom sklearn.model_selection import StratifiedKFold, GroupKFold\nfrom sklearn.metrics import roc_auc_score\nimport pydicom\nimport nibabel as nib  # For NII file handling\nfrom scipy.ndimage import zoom\n\nwarnings.filterwarnings('ignore')\n\ndef set_seed(seed=42):\n    \"\"\"Set all random seeds for reproducibility\"\"\"\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.cuda.manual_seed_all(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = True\n\nset_seed(42)\n\n# Device configuration\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"Using device: {device}\")\nif torch.cuda.is_available():\n    print(f\"GPU: {torch.cuda.get_device_name(0)}\")\n    print(f\"Memory: {torch.cuda.get_device_properties(0).total_memory / 1024**3:.1f} GB\")\n    print(f\"CUDA version: {torch.version.cuda}\")\n    torch.cuda.empty_cache()\nelse:\n    raise RuntimeError(\"CUDA is not available! This code requires GPU.\")","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:23:09.832807Z","iopub.execute_input":"2025-09-23T14:23:09.833038Z","iopub.status.idle":"2025-09-23T14:23:30.216882Z","shell.execute_reply.started":"2025-09-23T14:23:09.83302Z","shell.execute_reply":"2025-09-23T14:23:30.215987Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class Config:\n\n    TRAIN_CSV_PATH = \"/kaggle/input/rsna-intracranial-aneurysm-detection/train.csv\"\n    \n\n    # Alternative paths to try if the main path doesn't exist\n    SEGMENTATION_DIR_ALTERNATIVES = [\n        \"/kaggle/input/rsna-segmentation-masks\",\n        \"/kaggle/input/rsna-intracranial-aneurysm-detection/segmentations\",\n        \"/kaggle/input/rsna-intracranial-aneurysm-detection/segmentation\",\n        \"./segmentation\"  # Local path\n    ]\n    DICOM_SERIES_DIR = \"/kaggle/input/rsna-intracranial-aneurysm-detection/series\"\n    \n    # Model parameters for 2D U-Net segmentation\n    IMAGE_SIZE = 224\n    NUM_CLASSES = 14  # Binary segmentation: 0=background, 1=foreground\n    BATCH_SIZE = 16  # Adjusted for 2D processing\n    NUM_EPOCHS = 50\n    LEARNING_RATE = 1e-3\n    \n    # Model configuration\n    MODEL_TYPE = \"unet\"  # Changed from EfficientNet to U-Net\n    USE_METADATA = True\n    USE_WINDOWING = True\n    USE_CLAHE = True\n    USE_STRONG_AUGMENTATION = True\n    \n    # Segmentation specific settings\n    USE_DICE_LOSS = True\n    DICE_WEIGHT = 0.5\n    BCE_WEIGHT = 0.5\n    USE_FOCAL_LOSS = False\n    \n    # GPU optimization settings\n    NUM_WORKERS = 2\n    PIN_MEMORY = True\n    PREFETCH_FACTOR = 2\n    PERSISTENT_WORKERS = True\n    \n    # Training parameters with robust cross-validation\n    NUM_FOLDS = 5\n    FOLD = 0\n    ACCUMULATION_STEPS = 4\n    EARLY_STOPPING_PATIENCE = 5\n    USE_GROUP_CV = True\n    \n    # Data loading optimization\n    CACHE_SIZE = 100\n    \n    # Output\n    OUTPUT_DIR = \"/kaggle/working\"\n    MODEL_NAME = \"2d_unet_segmentation\"\n\nconfig = Config()\n\nprint(\"=== Configuration Summary - 2D U-Net Segmentation ===\")\nprint(f\"Model Type: {config.MODEL_TYPE}\")\nprint(f\"Image Size: {config.IMAGE_SIZE}\")\nprint(f\"Number of Classes: {config.NUM_CLASSES}\")\nprint(f\"Batch Size: {config.BATCH_SIZE}\")\nprint(f\"Accumulation Steps: {config.ACCUMULATION_STEPS}\")\nprint(f\"Effective Batch Size: {config.BATCH_SIZE * config.ACCUMULATION_STEPS}\")\nprint(f\"CLAHE Enabled: {config.USE_CLAHE}\")\nprint(f\"Strong Augmentation: {config.USE_STRONG_AUGMENTATION}\")\nprint(f\"Group Cross-Validation: {config.USE_GROUP_CV}\")\nprint(f\"Dice Loss Weight: {config.DICE_WEIGHT}\")\nprint(f\"BCE Loss Weight: {config.BCE_WEIGHT}\")","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:23:30.218542Z","iopub.execute_input":"2025-09-23T14:23:30.218778Z","iopub.status.idle":"2025-09-23T14:23:30.229037Z","shell.execute_reply.started":"2025-09-23T14:23:30.21876Z","shell.execute_reply":"2025-09-23T14:23:30.228463Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Load data\nprint(\"Loading data...\")\ntrain_df = pd.read_csv(config.TRAIN_CSV_PATH)\n\nprint(f\"Train data shape: {train_df.shape}\")\n\n# Define target columns for segmentation (13 anatomical locations + background)\nTARGET_COLS = [ \n    'Other Posterior Circulation',\n    'Basilar Tip',\n    'Right Posterior Communicating Artery',\n    'Left Posterior Communicating Artery',\n    'Right Anterior Cerebral Artery', \n    'Left Anterior Cerebral Artery',\n    'Anterior Communicating Artery', \n    'Right Middle Cerebral Artery', \n    'Left Middle Cerebral Artery',\n    'Right Supraclinoid Internal Carotid Artery',\n    'Left Supraclinoid Internal Carotid Artery',\n    'Right Infraclinoid Internal Carotid Artery',\n    'Left Infraclinoid Internal Carotid Artery',\n]\n\n# Class mapping for segmentation (0 = background, 1-13 = anatomical locations)\nCLASS_MAPPING = {col: idx + 1 for idx, col in enumerate(TARGET_COLS)}\nCLASS_MAPPING['Background'] = 0\n\nprint(f\"Target columns: {len(TARGET_COLS)}\")\nprint(f\"Class mapping: {CLASS_MAPPING}\")","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:23:30.22962Z","iopub.execute_input":"2025-09-23T14:23:30.229815Z","iopub.status.idle":"2025-09-23T14:23:30.302035Z","shell.execute_reply.started":"2025-09-23T14:23:30.2298Z","shell.execute_reply":"2025-09-23T14:23:30.301219Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_windowing_params(modality: str) -> Tuple[float, float]:\n    \"\"\"Get optimal windowing parameters for different modalities\"\"\"\n    windows = {\n        'CT': (40, 80),\n        'CTA': (50, 350), \n        'MRA': (600, 1200),\n        'MRI': (40, 80),\n        'MR': (40, 80)\n    }\n    return windows.get(modality, (40, 80))\n\ndef apply_dicom_windowing(img: np.ndarray, window_center: float, window_width: float) -> np.ndarray:\n    \"\"\"Apply DICOM windowing to normalize image intensities\"\"\"\n    img_min = window_center - window_width // 2\n    img_max = window_center + window_width // 2\n    img = np.clip(img, img_min, img_max)\n    img = (img - img_min) / (img_max - img_min + 1e-7)\n    return (img * 255).astype(np.uint8)\n\ndef apply_clahe_normalization(img: np.ndarray, modality: str) -> np.ndarray:\n    \"\"\"Apply CLAHE with modality-specific optimization\"\"\"\n    if not config.USE_CLAHE:\n        return img\n        \n    if modality in ['CTA', 'MRA']:\n        # Vascular imaging: stronger contrast improvement\n        clahe = cv2.createCLAHE(clipLimit=3.0, tileGridSize=(8, 8))\n        img_clahe = clahe.apply(img.astype(np.uint8))\n        img_clahe = cv2.convertScaleAbs(img_clahe, alpha=1.1, beta=5)\n    elif modality in ['MRI', 'MR']:\n        # MRI: gentler improvement with gamma correction\n        clahe = cv2.createCLAHE(clipLimit=2.0, tileGridSize=(8, 8))\n        img_clahe = clahe.apply(img.astype(np.uint8))\n        img_clahe = np.power(img_clahe / 255.0, 0.9) * 255\n        img_clahe = img_clahe.astype(np.uint8)\n    else:\n        # CT: standard CLAHE\n        clahe = cv2.createCLAHE(clipLimit=2.5, tileGridSize=(8, 8))\n        img_clahe = clahe.apply(img.astype(np.uint8))\n    \n    return img_clahe\n\ndef robust_normalization(volume: np.ndarray) -> np.ndarray:\n    \"\"\"Apply robust normalization using percentiles\"\"\"\n    p1, p99 = np.percentile(volume.flatten(), [1, 99])\n    volume_norm = np.clip(volume, p1, p99)\n    \n    if p99 > p1:\n        volume_norm = (volume_norm - p1) / (p99 - p1 + 1e-7)\n    else:\n        volume_norm = np.zeros_like(volume_norm)\n        \n    return (volume_norm * 255).astype(np.uint8)","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:23:30.302775Z","iopub.execute_input":"2025-09-23T14:23:30.303051Z","iopub.status.idle":"2025-09-23T14:23:30.311982Z","shell.execute_reply.started":"2025-09-23T14:23:30.303032Z","shell.execute_reply":"2025-09-23T14:23:30.311221Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 修改后的NII文件匹配函数，支持_cowseg后缀和文件夹验证\ndef find_matching_segmentation_file_enhanced(series_uid: str, segmentation_dir: str, dicom_series_dir: str) -> Optional[str]:\n    \"\"\"Find matching segmentation NII file for a given series UID by matching with DICOM folders\n    \n    Enhanced version that handles _cowseg suffix and verifies DICOM folder existence\n    \"\"\"\n    \n    # First, check if the corresponding DICOM series folder exists\n    dicom_series_path = os.path.join(dicom_series_dir, series_uid)\n    if not os.path.exists(dicom_series_path):\n        return None\n    \n    # Get list of DICOM files in the series folder\n    try:\n        dicom_files = [f for f in os.listdir(dicom_series_path) if f.endswith('.dcm')]\n        if len(dicom_files) == 0:\n            return None\n    except Exception as e:\n        return None\n    \n    # Enhanced matching: Handle _cowseg suffix and verify DICOM folder exists\n    try:\n        for root, dirs, files in os.walk(segmentation_dir):\n            for file in files:\n                if file.endswith(('.nii', '.nii.gz')):\n                    # Remove extensions to get the base name\n                    file_base = file.replace('.nii.gz', '').replace('.nii', '')\n                    \n                    # Handle _cowseg suffix: remove it if present\n                    if file_base.endswith('_cowseg'):\n                        file_base = file_base[:-7]  # Remove '_cowseg' (7 characters)\n                    else:\n                        continue\n                    # Check if the base name matches series_uid\n                    if file_base == series_uid:\n                        # Double-check that the corresponding DICOM folder exists\n                        dicom_folder_path = os.path.join(dicom_series_dir, file_base)\n                        if os.path.exists(dicom_folder_path):\n                            # Verify the folder contains DICOM files\n                            try:\n                                dcm_files = [f for f in os.listdir(dicom_folder_path) if f.endswith('.dcm')]\n                                if len(dcm_files) > 0:\n                                    return os.path.join(root, file)\n                            except Exception:\n                                continue\n                            \n    except Exception as e:\n        pass\n    \n    return None\n\n# 更新validate_dicom_nii_matching函数以使用新的匹配逻辑\ndef validate_dicom_nii_matching_enhanced(segmentation_dir: str, dicom_series_dir: str) -> dict:\n    \"\"\"Validate the matching between DICOM series and NII files with enhanced matching logic\"\"\"\n    print(\"Validating DICOM-NII matching with enhanced logic...\")\n    \n    # Get all DICOM series folders\n    dicom_series_folders = []\n    try:\n        for item in os.listdir(dicom_series_dir):\n            item_path = os.path.join(dicom_series_dir, item)\n            if os.path.isdir(item_path):\n                # Check if it contains DICOM files\n                dcm_files = [f for f in os.listdir(item_path) if f.endswith('.dcm')]\n                if len(dcm_files) > 0:\n                    dicom_series_folders.append(item)\n    except Exception as e:\n        print(f\"Error reading DICOM series directory: {e}\")\n        return {}\n    \n    # Get all NII files\n    nii_files = list_available_nii_files(segmentation_dir)\n    \n    print(f\"Found {len(dicom_series_folders)} DICOM series folders\")\n    print(f\"Found {len(nii_files)} NII files\")\n    \n    # Enhanced matching: Handle _cowseg suffix and verify DICOM folder exists\n    matches = []\n    unmatched_dicom = []\n    unmatched_nii = []\n    \n    for series_uid in dicom_series_folders:\n        matched_nii = None\n        for nii_file in nii_files:\n            nii_basename = os.path.basename(nii_file)\n            # Remove extensions to get the base name\n            nii_base = nii_basename.replace('.nii.gz', '').replace('.nii', '')\n            \n            # Handle _cowseg suffix: remove it if present\n            if nii_base.endswith('_cowseg'):\n                nii_base = nii_base[:-7]  # Remove '_cowseg' (7 characters)\n            \n            # Check if the base name matches series_uid\n            if series_uid == nii_base:\n                # Double-check that the corresponding DICOM folder exists\n                dicom_folder_path = os.path.join(dicom_series_dir, nii_base)\n                if os.path.exists(dicom_folder_path):\n                    # Verify the folder contains DICOM files\n                    try:\n                        dcm_files = [f for f in os.listdir(dicom_folder_path) if f.endswith('.dcm')]\n                        if len(dcm_files) > 0:\n                            matched_nii = nii_file\n                            break\n                    except Exception:\n                        continue\n        \n        if matched_nii:\n            matches.append((series_uid, matched_nii))\n        else:\n            unmatched_dicom.append(series_uid)\n    \n    # Find unmatched NII files\n    matched_nii_files = [match[1] for match in matches]\n    for nii_file in nii_files:\n        if nii_file not in matched_nii_files:\n            unmatched_nii.append(nii_file)\n    \n    print(f\"\\\\nEnhanced Matching Results:\")\n    print(f\"- Successful matches: {len(matches)}\")\n    print(f\"- Unmatched DICOM series: {len(unmatched_dicom)}\")\n    print(f\"- Unmatched NII files: {len(unmatched_nii)}\")\n    \n    if len(matches) > 0:\n        print(f\"\\\\nSample matches:\")\n        for i, (series_uid, nii_file) in enumerate(matches[:5]):\n            print(f\"  {i+1}. {series_uid} <-> {os.path.basename(nii_file)}\")\n    \n    if len(unmatched_dicom) > 0:\n        print(f\"\\\\nSample unmatched DICOM series:\")\n        for i, series_uid in enumerate(unmatched_dicom[:5]):\n            print(f\"  {i+1}. {series_uid}\")\n    \n    if len(unmatched_nii) > 0:\n        print(f\"\\\\nSample unmatched NII files:\")\n        for i, nii_file in enumerate(unmatched_nii[:5]):\n            print(f\"  {i+1}. {os.path.basename(nii_file)}\")\n    \n    return {\n        'matches': matches,\n        'unmatched_dicom': unmatched_dicom,\n        'unmatched_nii': unmatched_nii,\n        'total_dicom': len(dicom_series_folders),\n        'total_nii': len(nii_files)\n    }\n\nprint(\"✅ 已创建增强版的NII文件匹配函数\")\nprint(\"   - find_matching_segmentation_file_enhanced: 支持_cowseg后缀和文件夹验证\")\nprint(\"   - validate_dicom_nii_matching_enhanced: 增强版验证函数\")\n","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:23:30.313683Z","iopub.execute_input":"2025-09-23T14:23:30.314286Z","iopub.status.idle":"2025-09-23T14:23:30.336732Z","shell.execute_reply.started":"2025-09-23T14:23:30.314262Z","shell.execute_reply":"2025-09-23T14:23:30.336079Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 添加二值化处理函数\ndef apply_binary_thresholding(mask: np.ndarray, threshold: int = 127) -> np.ndarray:\n    \"\"\"\n    对分割掩码进行二值化处理\n    大于threshold的像素值设为1，小于等于threshold的像素值设为0\n    \n    Args:\n        mask: 输入的分割掩码\n        threshold: 二值化阈值，默认为127\n    \n    Returns:\n        二值化后的掩码 (0或1)\n    \"\"\"\n    if mask.max() <= 1:\n        # 如果已经是二值化的，直接返回\n        return mask.astype(np.uint8)\n    \n    # 应用二值化阈值\n    binary_mask = (mask > threshold).astype(np.uint8)\n    return binary_mask\n\ndef load_nii_segmentation(nii_path: str) -> np.ndarray:\n    \"\"\"Load NII segmentation file and return 3D array\"\"\"\n    try:\n        nii_img = nib.load(nii_path)\n        segmentation = nii_img.get_fdata()\n        return segmentation.astype(np.uint8)\n    except Exception as e:\n        print(f\"Error loading NII file {nii_path}: {e}\")\n        return None\nimport os\nimport time\nimport glob\nimport numpy as np\nimport cv2\nfrom typing import List\n\ndef convert_3d_to_2d_slices_intelligent(volume_3d: np.ndarray, target_size: int = 224, \n                                       series_uid: str = None, dicom_series_dir: str = None) -> Tuple[List[np.ndarray], List[str]]:\n    \"\"\"使用OpenCV进行空间对齐的智能3D到2D转换（向量化批处理版本）\n    \n    Args:\n        volume_3d: 3D分割体积\n        target_size: 2D切片的目标尺寸\n        series_uid: 序列UID，用于检查对应的DICOM文件\n        dicom_series_dir: 包含DICOM序列的目录\n    \n    Returns:\n        (slices, dicom_filenames) 二元组：\n        - slices: 2D切片列表（顺序与排序后的DICOM文件名一致）\n        - dicom_filenames: 对应的DICOM基文件名（已排序）\n    \"\"\"\n    if volume_3d is None:\n        return [], []\n    \n    # 如果没有提供series_uid或dicom_series_dir，直接忽略\n    if not series_uid or not dicom_series_dir:\n        print(\"未提供series_uid或dicom_series_dir，忽略此NII文件\")\n        return [], []\n    \n    # 检查对应的DICOM文件夹是否存在\n    dicom_folder_path = os.path.join(dicom_series_dir, series_uid)\n    if not os.path.exists(dicom_folder_path):\n        print(f\"未找到DICOM文件夹: {dicom_folder_path}，忽略此NII文件\")\n        return [], []\n    \n    # 记录总处理时间\n    total_start_time = time.time()\n    \n    try:\n        # 快速扫描目录获取.dcm文件\n        print(\"扫描DICOM文件...\")\n        dicom_files = glob.glob(os.path.join(dicom_folder_path, \"*.dcm\"))\n        dicom_files.extend(glob.glob(os.path.join(dicom_folder_path, \"*.DCM\")))\n        \n        if len(dicom_files) == 0:\n            print(\"DICOM文件夹中没有找到有效的DICOM文件，忽略此NII文件\")\n            return [], []\n        \n        # 排序以建立稳定映射\n        dicom_files = sorted(dicom_files)\n        dicom_basenames = [os.path.basename(p) for p in dicom_files]\n        \n        # 检查NII的Z维与DICOM文件数量是否匹配（假定NII为 (W, H, Z)）\n        nii_shape = volume_3d.shape\n        print(f\"NII体积形状: {nii_shape}, DICOM文件数量: {len(dicom_files)}\")\n        \n        if len(nii_shape) < 3 or nii_shape[2] != len(dicom_files):\n            print(f\"NII的Z维 ({nii_shape[2] if len(nii_shape)>=3 else 'N/A'}) 与DICOM文件数量 ({len(dicom_files)}) 不匹配\")\n            return [], []\n        \n        # 调整NII的维度顺序以匹配DICOM（(W,H,Z) -> (Z,H,W)）\n        print(\"调整NII维度顺序以匹配DICOM文件...\")\n        aligned_seg_vol = np.transpose(volume_3d, (2, 1, 0))  # 将NII从(W, H, Z)调整为(Z, H, W)\n        \n        print(f\"调整后的NII体积形状: {aligned_seg_vol.shape} (Z, H, W)\")\n        \n        # 向量化处理：一次性处理所有切片\n        z_slices = aligned_seg_vol.shape[0]\n        \n        # 创建一个存储所有切片的列表\n        all_slices = []\n        \n        for i in range(z_slices):\n            slice_2d = aligned_seg_vol[i]\n            all_slices.append(slice_2d)\n        \n        # 检查是否需要调整大小\n        if aligned_seg_vol.shape[1:] != (target_size, target_size):\n            start_time = time.time()\n            print(f\"调整到目标尺寸: {aligned_seg_vol.shape[1:]} -> ({target_size}, {target_size})\")\n            resized_slices = []\n            \n            # 批量调整大小\n            for i in range(0, len(all_slices), 32):  # 每批处理32个切片\n                batch_end = min(i + 32, len(all_slices))\n                batch_slices = all_slices[i:batch_end]\n                \n                # 批量调整大小\n                batch_resized = np.array([\n                    cv2.resize(slice_2d, (target_size, target_size), \n                             interpolation=cv2.INTER_NEAREST)\n                    for slice_2d in batch_slices\n                ])\n                resized_slices.append(batch_resized)\n            \n            # 合并所有批次\n            all_slices = np.vstack(resized_slices)\n            resize_time = time.time() - start_time\n            print(f\"图像resize耗时: {resize_time:.3f}秒\")\n        \n        # 转换为列表格式\n        slices = [all_slices[i] for i in range(len(all_slices))]\n        \n        total_time = time.time() - total_start_time\n        print(f\"=== 总处理耗时: {total_time:.3f}秒 ===\")\n        print(f\"成功生成了{len(slices)}个对齐的2D切片\")\n        return slices, dicom_basenames\n        \n    except Exception as e:\n        total_time = time.time() - total_start_time\n        print(f\"=== 处理失败，总耗时: {total_time:.3f}秒 ===\")\n        print(f\"OpenCV处理过程中出错: {e}，忽略此NII文件\")\n        return [], []\n\n\n\n\n\ndef process_2d_slice_intelligent(slice_2d: np.ndarray, target_size: int) -> np.ndarray:\n    \"\"\"处理单个2D切片：调整大小并应用二值化\"\"\"\n    # 调整到目标大小\n    if slice_2d.shape != (target_size, target_size):\n        slice_2d = cv2.resize(slice_2d, (target_size, target_size), \n                            interpolation=cv2.INTER_NEAREST)  # 对分割掩码使用最近邻插值\n    \n    # 应用二值化\n    slice_2d = apply_binary_thresholding(slice_2d, threshold=0)\n    \n    return slice_2d\n\ndef analyze_3d_volume_structure(volume_3d: np.ndarray, series_uid: str = None, dicom_series_dir: str = None) -> dict:\n    \"\"\"使用SimpleITK分析3D体积结构\n    \n    Args:\n        volume_3d: 3D分割体积\n        series_uid: 序列UID\n        dicom_series_dir: DICOM序列目录\n    \n    Returns:\n        包含分析结果的字典\n    \"\"\"\n    if volume_3d is None:\n        return {}\n    \n    analysis = {\n        'volume_shape': volume_3d.shape,\n        'dicom_file_count': 0,\n        'alignment_successful': False,\n        'aligned_shape': None,\n        'confidence': 'none'\n    }\n    \n    if series_uid and dicom_series_dir:\n        dicom_folder_path = os.path.join(dicom_series_dir, series_uid)\n        if os.path.exists(dicom_folder_path):\n            try:\n                # 使用SimpleITK读取DICOM序列\n                reader = sitk.ImageSeriesReader()\n                dicom_names = reader.GetGDCMSeriesFileNames(dicom_folder_path)\n                \n                if len(dicom_names) > 0:\n                    reader.SetFileNames(dicom_names)\n                    dicom_image_sitk = reader.Execute()\n                    dicom_vol = sitk.GetArrayFromImage(dicom_image_sitk)\n                    \n                    analysis['dicom_file_count'] = len(dicom_names)\n                    analysis['dicom_shape'] = dicom_vol.shape\n                    \n                    # 尝试对齐\n                    seg_image_sitk = sitk.GetImageFromArray(volume_3d)\n                    resampler = sitk.ResampleImageFilter()\n                    resampler.SetReferenceImage(dicom_image_sitk)\n                    resampler.SetInterpolator(sitk.sitkNearestNeighbor)\n                    resampled_seg_sitk = resampler.Execute(seg_image_sitk)\n                    aligned_seg_vol = sitk.GetArrayFromImage(resampled_seg_sitk)\n                    \n                    analysis['aligned_shape'] = aligned_seg_vol.shape\n                    \n                    # 检查对齐是否成功\n                    if aligned_seg_vol.shape == dicom_vol.shape:\n                        analysis['alignment_successful'] = True\n                        analysis['confidence'] = 'high'\n                    else:\n                        analysis['alignment_successful'] = False\n                        analysis['confidence'] = 'low'\n                        \n            except Exception as e:\n                analysis['error'] = str(e)\n                analysis['confidence'] = 'none'\n        else:\n            analysis['confidence'] = 'none'\n    else:\n        analysis['confidence'] = 'none'\n    \n    return analysis\n\nprint(\"✅ 已创建智能3D到2D转换函数\")\nprint(\"   - convert_3d_to_2d_slices_intelligent: 根据NII形状和DICOM文件数量智能选择切分轴\")\nprint(\"   - process_2d_slice_intelligent: 处理单个2D切片\")\nprint(\"   - analyze_3d_volume_structure: 分析3D体积结构并提供切分建议\")\n\ndef create_segmentation_mask_from_labels(labels: np.ndarray, class_mapping: dict) -> np.ndarray:\n    \"\"\"Create binary segmentation mask from labels (0=background, 1=foreground)\"\"\"\n    mask = np.zeros((config.IMAGE_SIZE, config.IMAGE_SIZE), dtype=np.uint8)\n    \n    # Check if any anatomical location is present\n    has_aneurysm = any(labels == 1)\n    \n    if has_aneurysm:\n        # Create a simple foreground region in the center\n        # In practice, you'd use actual segmentation data\n        center_y, center_x = config.IMAGE_SIZE // 2, config.IMAGE_SIZE // 2\n        y_start = max(0, center_y - 15)\n        y_end = min(config.IMAGE_SIZE, center_y + 15)\n        x_start = max(0, center_x - 15)\n        x_end = min(config.IMAGE_SIZE, center_x + 15)\n        mask[y_start:y_end, x_start:x_end] = 1  # Binary: 1 = foreground\n    \n    return mask\n\ndef list_available_nii_files(segmentation_dir: str) -> List[str]:\n    \"\"\"List all available NII files in the segmentation directory for debugging\"\"\"\n    nii_files = []\n    try:\n        for root, dirs, files in os.walk(segmentation_dir):\n            for file in files:\n                if file.endswith(('.nii', '.nii.gz')):\n                    nii_files.append(os.path.join(root, file))\n    except Exception as e:\n        print(f\"Error listing NII files: {e}\")\n    \n    return nii_files\n\ndef find_available_segmentation_dir(alternative_paths: List[str]) -> Optional[str]:\n    \"\"\"Find the first available segmentation directory from a list of alternatives\"\"\"\n    for path in alternative_paths:\n        if os.path.exists(path):\n            nii_files = list_available_nii_files(path)\n            if len(nii_files) > 0:\n                print(f\"Found segmentation directory: {path} with {len(nii_files)} NII files\")\n                return path\n            else:\n                print(f\"Directory exists but no NII files found: {path}\")\n        else:\n            print(f\"Directory does not exist: {path}\")\n    \n    print(\"No valid segmentation directory found!\")\n    return None\n","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:23:30.337715Z","iopub.execute_input":"2025-09-23T14:23:30.337938Z","iopub.status.idle":"2025-09-23T14:23:30.369882Z","shell.execute_reply.started":"2025-09-23T14:23:30.337902Z","shell.execute_reply":"2025-09-23T14:23:30.369221Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def extract_dicom_patient_info(series_uid: str) -> Tuple[str, str]:\n    \"\"\"Extract StudyInstanceUID and PatientID from DICOM metadata\"\"\"\n    try:\n        dicom_dir = f\"/kaggle/input/rsna-intracranial-aneurysm-detection/series/{series_uid}\"\n        if os.path.exists(dicom_dir):\n            dcm_files = [f for f in os.listdir(dicom_dir) if f.endswith('.dcm')]\n            if dcm_files:\n                ds = pydicom.dcmread(\n                    os.path.join(dicom_dir, dcm_files[0]), \n                    stop_before_pixels=True, \n                    force=True\n                )\n                study_uid = getattr(ds, 'StudyInstanceUID', None)\n                patient_id = getattr(ds, 'PatientID', None)\n                return study_uid or f\"fallback_{series_uid[:32]}\", patient_id\n    except Exception:\n        pass\n    \n    # Fallback: use longer prefix from series UID\n    return f\"fallback_{series_uid[:32]}\", f\"fallback_{series_uid[:32]}\"\n\n@functools.lru_cache(maxsize=5000)\ndef get_patient_group_cached(series_uid: str) -> str:\n    \"\"\"Get patient group with caching for performance\"\"\"\n    study_uid, patient_id = extract_dicom_patient_info(series_uid)\n    # Use StudyInstanceUID as primary identifier\n    return study_uid if study_uid and not study_uid.startswith('fallback_') else patient_id\n\n# 更新数据映射函数以使用智能3D到2D转换\ndef create_segmentation_data_mapping_intelligent():\n    \"\"\"进一步优化版分割数据映射 - 直接遍历NII文件匹配训练数据\"\"\"\n    segmentation_mapping = {}\n    found_nii_count = 0\n    placeholder_count = 0\n    skipped_no_positive = 0\n    skipped_no_dicom = 0\n    \n    print(\"创建进一步优化版分割数据映射 - 直接遍历NII文件...\")\n    \n    # 尝试找到可用的分割目录\n    segmentation_dir = find_available_segmentation_dir(config.SEGMENTATION_DIR_ALTERNATIVES)\n    if segmentation_dir is None:\n        print(\"警告: 未找到分割目录。将仅使用占位符掩码。\")\n        segmentation_dir = config.SEGMENTATION_DIR  # 用于占位符创建的默认路径\n    \n    print(f\"使用分割目录: {segmentation_dir}\")\n    \n    # 一次性获取所有NII文件\n    all_nii_files = list_available_nii_files(segmentation_dir)\n    print(f\"找到 {len(all_nii_files)} 个NII文件\")\n    \n    # 创建训练数据的快速查找字典（不过滤阳性，全部纳入）\n    train_data_dict = {}\n    for _, row in train_df.iterrows():\n        series_uid = row['SeriesInstanceUID']\n        train_data_dict[series_uid] = row\n    \n    print(f\"训练数据中共有 {len(train_data_dict)} 个序列\")\n    \n    # 直接遍历NII文件进行处理\n    print(f\"\\\\n直接遍历NII文件进行匹配...\")\n    for nii_file in tqdm(all_nii_files, desc=\"处理NII文件\"):\n        # 从NII文件名提取series_uid\n        if '_cowseg' in nii_file:\n            series_uid = nii_file.split('/')[-1].replace('.nii','').replace('_cowseg','')\n        else:\n            continue\n        # if os.path.exist'/kaggle/input/rsna-intracranial-aneurysm-detection/series'+series_uid)\n        # if series_uid is None:\n        #     print(f\"⚠️  无法解析NII文件名: {os.path.basename(nii_file)}\")\n        #     continue\n        \n        # 检查这个series_uid是否在训练数据中（不过滤阳性）\n        if series_uid not in train_data_dict:\n            skipped_no_positive += 1\n            continue\n        \n        # # 验证对应的DICOM文件夹是否存在\n        # if not validate_dicom_folder_exists(series_uid, config.DICOM_SERIES_DIR):\n        #     skipped_no_dicom += 1\n        #     print(f\"⚠️  {series_uid}: 未找到对应的DICOM文件夹\")\n        #     continue\n        \n        # 获取训练数据行\n        train_row = train_data_dict[series_uid]\n        \n        # 加载分割数据\n        segmentation_3d = load_nii_segmentation(nii_file)\n        if segmentation_3d is not None:\n            # 分析3D体积结构\n            # analysis = analyze_3d_volume_structure(segmentation_3d, series_uid, config.DICOM_SERIES_DIR)\n            \n            # 使用智能转换转换为2D切片（返回切片与对应DICOM文件名）\n            slices_2d, dicom_filenames = convert_3d_to_2d_slices_intelligent(\n                segmentation_3d, config.IMAGE_SIZE, series_uid, config.DICOM_SERIES_DIR\n            )\n            \n            if len(slices_2d) > 0:  # 确保成功生成了切片\n                # 构建文件名 -> 切片 的映射\n                slice_map = {fname: slices_2d[i] for i, fname in enumerate(dicom_filenames)}\n                \n                # 过滤全零掩码切片（仅保留含前景像素的切片）\n                nonzero_indices = [i for i, s in enumerate(slices_2d) if (np.asarray(s) > 0).any()]\n                if len(nonzero_indices) != len(slices_2d):\n                    print(f\"{series_uid}: 过滤掉 {len(slices_2d)-len(nonzero_indices)} 个全零掩码切片\")\n                \n                if len(nonzero_indices) > 0:\n                    slices_2d = [slices_2d[i] for i in nonzero_indices]\n                    dicom_filenames = [dicom_filenames[i] for i in nonzero_indices]\n                    slice_map = {dicom_filenames[i]: slices_2d[i] for i in range(len(slices_2d))}\n                else:\n                    # 若全部为零，置空（该序列将不会贡献样本）\n                    slices_2d = []\n                    dicom_filenames = []\n                    slice_map = {}\n                \n                segmentation_mapping[series_uid] = {\n                    'segmentation_file': nii_file,\n                    'slices_2d': slices_2d,\n                    'slice_map': slice_map,  # 新增：文件名到切片的映射（已过滤零掩码）\n                    'dicom_filenames': dicom_filenames,  # 新增：已排序的文件名列表（与切片同步）\n                    'labels': train_row[TARGET_COLS].values,\n                    'volume_analysis': None\n                }\n                found_nii_count += 1\n                \n                if found_nii_count <= 5:  # 打印前5个匹配用于调试\n                    print(f\"✅ 处理NII文件: {os.path.basename(nii_file)}\")\n                    print(f\"   对应序列: {series_uid}\")\n                    print(f\"   生成了 {len(slices_2d)} 个2D切片\")\n            else:\n                print(f\"⚠️  {series_uid}: NII文件加载成功但未生成2D切片\")\n        else:\n            print(f\"❌ 加载NII文件失败: {nii_file}\")\n    \n    # 为有阳性标签但没有NII文件的序列创建占位符\n    print(f\"\\\\n为没有NII文件的序列创建占位符...\")\n    for series_uid, train_row in train_data_dict.items():\n        if series_uid not in segmentation_mapping:\n            placeholder_mask = create_segmentation_mask_from_labels(\n                train_row[TARGET_COLS].values, CLASS_MAPPING\n            )\n            segmentation_mapping[series_uid] = {\n                'segmentation_file': None,\n                'slices_2d': [placeholder_mask],  # 单个切片占位符\n                'labels': train_row[TARGET_COLS].values,\n                'volume_analysis': None\n            }\n            placeholder_count += 1\n    \n    print(f\"\\\\n进一步优化版分割映射摘要:\")\n    print(f\"- 总NII文件数: {len(all_nii_files)}\")\n    print(f\"- 成功处理NII文件: {found_nii_count}\")\n    print(f\"- 跳过(不在训练数据中): {skipped_no_positive}\")\n    print(f\"- 跳过(无DICOM文件夹): {skipped_no_dicom}\")\n    print(f\"- 使用占位符: {placeholder_count}\")\n    print(f\"- 最终映射序列数: {len(segmentation_mapping)}\")\n    print(f\"- NII文件利用率: {found_nii_count}/{len(all_nii_files)} = {found_nii_count/len(all_nii_files)*100:.1f}%\")\n    print(f\"- 训练数据覆盖率: {len(segmentation_mapping)}/{len(train_data_dict)} = {len(segmentation_mapping)/len(train_data_dict)*100:.1f}%\")\n    \n    return segmentation_mapping\n\nprint(\"✅ 已创建智能数据映射函数\")\nprint(\"   - create_segmentation_data_mapping_intelligent: 使用智能3D到2D转换\")\n\n# Create segmentation data mapping\nsegmentation_data_dict = create_segmentation_data_mapping_intelligent()\nprint(f\"Created segmentation mapping for {len(segmentation_data_dict)} series\")\n\n# Filter data to only include series with segmentation data\nvalid_series = list(segmentation_data_dict.keys())\ntrain_df_filtered = train_df[train_df['SeriesInstanceUID'].isin(valid_series)].copy()\nprint(f\"Filtered train data shape: {train_df_filtered.shape}\")\n\n# Check distribution\npositive_labels_count = train_df_filtered[TARGET_COLS].sum().sum()\nprint(f\"Total positive labels: {positive_labels_count}\")\nprint(f\"Average positive labels per series: {positive_labels_count / len(train_df_filtered):.2f}\")","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:23:30.370675Z","iopub.execute_input":"2025-09-23T14:23:30.37089Z","iopub.status.idle":"2025-09-23T14:26:52.049782Z","shell.execute_reply.started":"2025-09-23T14:23:30.370867Z","shell.execute_reply":"2025-09-23T14:26:52.0489Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2D灰度图像二分类数据增强\nprint(\"=== 配置2D灰度图像二分类数据增强 ===\")\n\nif config.USE_STRONG_AUGMENTATION:\n    print(\"使用强数据增强 - 针对2D灰度医学图像二分类优化\")\n    train_transform = A.Compose([\n        # === 几何变换 (医学图像安全) ===\n        # 旋转 - 医学图像通常可以小幅旋转\n        A.Rotate(limit=15, p=0.7, border_mode=cv2.BORDER_CONSTANT, value=0),\n        \n        # 翻转 - 对于脑血管图像，水平翻转是安全的\n        A.HorizontalFlip(p=0.5),\n        \n        # 缩放和平移 - 模拟不同的扫描视野\n        A.ShiftScaleRotate(\n            shift_limit=0.1, \n            scale_limit=0.15, \n            rotate_limit=10, \n            border_mode=cv2.BORDER_CONSTANT,\n            value=0,\n            p=0.6\n        ),\n        \n        \n        # === 图像质量变化 (模拟不同扫描仪/协议) ===\n        # 亮度和对比度调整\n        A.RandomBrightnessContrast(\n            brightness_limit=0.2, \n            contrast_limit=0.2, \n            p=0.6\n        ),\n        \n        # CLAHE - 自适应直方图均衡化\n        A.CLAHE(\n            clip_limit=2.0, \n            tile_grid_size=(8, 8), \n            p=0.4\n        ),\n        \n        # Gamma校正 - 模拟不同的显示特性\n        A.RandomGamma(\n            gamma_limit=(85, 115), \n            p=0.4\n        ),\n        \n        # === 噪声模拟 (扫描仪差异) ===\n        # 高斯噪声 - 模拟热噪声\n        A.GaussNoise(\n            var_limit=(5, 25), \n            mean=0,\n            p=0.3\n        ),\n        \n        # 模糊 - 模拟运动伪影或低分辨率\n        A.OneOf([\n            A.Blur(blur_limit=3, p=1.0),\n            A.MotionBlur(blur_limit=3, p=1.0),\n            A.GaussianBlur(blur_limit=3, p=1.0),\n        ], p=0.2),\n        \n \n        \n        # # 像素级dropout - 模拟噪声像素\n        # A.PixelDropout(\n        #     dropout_prob=0.01,\n        #     p=0.1\n        # ),\n        \n        # === 归一化和张量转换 ===\n        # 针对灰度图像的归一化 (ImageNet单通道统计)\n        # A.Normalize(\n        #     mean=[0.485],  # 灰度图像单通道\n        #     std=[0.229]\n        # ),\n        ToTensorV2()\n    ])\nelse:\n    print(\"使用标准数据增强 - 2D灰度医学图像二分类\")\n    train_transform = A.Compose([\n        # 基础几何变换\n        A.Rotate(limit=10, p=0.5, border_mode=cv2.BORDER_CONSTANT, value=0),\n        A.HorizontalFlip(p=0.5),\n        \n        # 基础图像质量调整\n        A.RandomBrightnessContrast(\n            brightness_limit=0.15, \n            contrast_limit=0.15, \n            p=0.4\n        ),\n        \n        # 轻微噪声\n        A.GaussNoise(\n            var_limit=(5, 20), \n            p=0.2\n        ),\n        \n        # 归一化\n        # A.Normalize(mean=[0.485], std=[0.229]),\n        ToTensorV2()\n    ])\n\n# 验证集变换 - 仅归一化\nval_transform = A.Compose([\n    A.Normalize(mean=[0.485], std=[0.229]),  # 灰度图像归一化\n    ToTensorV2()\n])\n\nprint(f\"训练变换: {'强增强' if config.USE_STRONG_AUGMENTATION else '标准增强'}\")\nprint(f\"验证变换: 仅归一化\")\nprint(\"✅ 2D灰度图像二分类数据增强配置完成\")\n","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:26:52.05068Z","iopub.execute_input":"2025-09-23T14:26:52.050904Z","iopub.status.idle":"2025-09-23T14:26:52.068359Z","shell.execute_reply.started":"2025-09-23T14:26:52.050885Z","shell.execute_reply":"2025-09-23T14:26:52.067641Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 2D灰度图像二分类专用数据增强工具函数\nprint(\"=== 二分类任务专用数据增强工具 ===\")\n\ndef get_classification_transforms_2d_grayscale(image_size=224, is_training=True, augmentation_level='strong'):\n    \"\"\"\n    为2D灰度医学图像二分类任务创建专用的数据增强管道\n    \n    Args:\n        image_size (int): 目标图像尺寸\n        is_training (bool): 是否为训练模式\n        augmentation_level (str): 增强级别 ('none', 'light', 'medium', 'strong')\n    \n    Returns:\n        albumentations.Compose: 数据增强管道\n    \"\"\"\n    \n    if not is_training:\n        # 验证/测试时只进行归一化\n        return A.Compose([\n            A.Normalize(mean=[0.485], std=[0.229]),\n            ToTensorV2()\n        ])\n    \n    # 训练时的数据增强\n    transforms_list = []\n    \n    if augmentation_level == 'none':\n        # 无增强，仅归一化\n        transforms_list = [\n            A.Normalize(mean=[0.485], std=[0.229]),\n            ToTensorV2()\n        ]\n    \n    elif augmentation_level == 'light':\n        # 轻度增强 - 仅基础变换\n        transforms_list = [\n            # 几何变换\n            A.HorizontalFlip(p=0.5),\n            A.Rotate(limit=10, p=0.3, border_mode=cv2.BORDER_CONSTANT, value=0),\n            \n            # 图像质量\n            A.RandomBrightnessContrast(brightness_limit=0.1, contrast_limit=0.1, p=0.3),\n            \n            # 归一化\n            A.Normalize(mean=[0.485], std=[0.229]),\n            ToTensorV2()\n        ]\n    \n    elif augmentation_level == 'medium':\n        # 中等增强\n        transforms_list = [\n            # 几何变换\n            A.HorizontalFlip(p=0.5),\n            A.Rotate(limit=15, p=0.5, border_mode=cv2.BORDER_CONSTANT, value=0),\n            A.ShiftScaleRotate(\n                shift_limit=0.1, scale_limit=0.1, rotate_limit=10,\n                border_mode=cv2.BORDER_CONSTANT, value=0, p=0.4\n            ),\n            \n            # 图像质量\n            A.RandomBrightnessContrast(brightness_limit=0.15, contrast_limit=0.15, p=0.5),\n            A.RandomGamma(gamma_limit=(90, 110), p=0.3),\n            \n            # 噪声\n            A.GaussNoise(var_limit=(5, 20), p=0.2),\n            \n            # 归一化\n            A.Normalize(mean=[0.485], std=[0.229]),\n            ToTensorV2()\n        ]\n    \n    elif augmentation_level == 'strong':\n        # 强增强 - 完整的增强管道\n        transforms_list = [\n            # === 几何变换 ===\n            A.HorizontalFlip(p=0.5),\n            A.Rotate(limit=20, p=0.7, border_mode=cv2.BORDER_CONSTANT, value=0),\n            A.ShiftScaleRotate(\n                shift_limit=0.15, scale_limit=0.2, rotate_limit=15,\n                border_mode=cv2.BORDER_CONSTANT, value=0, p=0.6\n            ),\n        \n            \n            # === 图像质量变换 ===\n            A.RandomBrightnessContrast(brightness_limit=0.2, contrast_limit=0.2, p=0.6),\n            A.RandomGamma(gamma_limit=(80, 120), p=0.4),\n            A.CLAHE(clip_limit=2.0, tile_grid_size=(8, 8), p=0.4),\n            \n            # === 噪声和模糊 ===\n            A.OneOf([\n                A.Blur(blur_limit=3, p=1.0),\n                A.MotionBlur(blur_limit=3, p=1.0),\n                A.GaussianBlur(blur_limit=3, p=1.0),\n            ], p=0.2),\n            \n\n            # 归一化\n            A.Normalize(mean=[0.485], std=[0.229]),\n            ToTensorV2()\n        ]\n    \n    return A.Compose(transforms_list)\n\ndef get_test_time_augmentation_2d_classification():\n    \"\"\"\n    为2D灰度图像二分类创建测试时增强(TTA)变换\n    \n    Returns:\n        list: TTA变换列表\n    \"\"\"\n    tta_transforms = [\n        # 原始图像\n        A.Compose([\n            A.Normalize(mean=[0.485], std=[0.229]),\n            ToTensorV2()\n        ]),\n        \n        # 水平翻转\n        A.Compose([\n            A.HorizontalFlip(p=1.0),\n            A.Normalize(mean=[0.485], std=[0.229]),\n            ToTensorV2()\n        ]),\n        \n        # 垂直翻转\n        A.Compose([\n            A.VerticalFlip(p=1.0),\n            A.Normalize(mean=[0.485], std=[0.229]),\n            ToTensorV2()\n        ]),\n        \n        # 90度旋转\n        A.Compose([\n            A.RandomRotate90(p=1.0),\n            A.Normalize(mean=[0.485], std=[0.229]),\n            ToTensorV2()\n        ]),\n        \n        # 水平+垂直翻转\n        A.Compose([\n            A.HorizontalFlip(p=1.0),\n            A.VerticalFlip(p=1.0),\n            A.Normalize(mean=[0.485], std=[0.229]),\n            ToTensorV2()\n        ])\n    ]\n    \n    return tta_transforms\n\ndef visualize_augmentations_2d_classification(image, transforms, num_samples=4):\n    \"\"\"\n    可视化2D灰度图像二分类的数据增强效果\n    \n    Args:\n        image (np.ndarray): 输入图像 (H, W)\n        transforms (albumentations.Compose): 数据增强管道\n        num_samples (int): 生成的样本数量\n    \"\"\"\n    import matplotlib.pyplot as plt\n    \n    fig, axes = plt.subplots(1, num_samples + 1, figsize=(20, 4))\n    \n    # 显示原始图像\n    axes[0].imshow(image, cmap='gray')\n    axes[0].set_title('原始图像')\n    axes[0].axis('off')\n    \n    # 生成增强样本\n    for i in range(num_samples):\n        # 应用增强（需要移除最后的归一化和ToTensor步骤用于可视化）\n        aug_transforms = A.Compose(transforms.transforms[:-2])  # 移除Normalize和ToTensorV2\n        augmented = aug_transforms(image=image)\n        aug_image = augmented['image']\n        \n        # 显示增强后的图像\n        axes[i + 1].imshow(aug_image, cmap='gray')\n        axes[i + 1].set_title(f'增强样本 {i + 1}')\n        axes[i + 1].axis('off')\n    \n    plt.tight_layout()\n    plt.show()\n\n# 应用新的数据增强配置\naugmentation_level = 'strong' if config.USE_STRONG_AUGMENTATION else 'medium'\n\ntrain_transform = get_classification_transforms_2d_grayscale(\n    image_size=config.IMAGE_SIZE,\n    is_training=True,\n    augmentation_level=augmentation_level\n)\n\nval_transform = get_classification_transforms_2d_grayscale(\n    image_size=config.IMAGE_SIZE,\n    is_training=False\n)\n\nprint(f\"✅ 已配置 {augmentation_level} 级别的2D灰度图像二分类数据增强\")\nprint(f\"   - 训练变换: {len(train_transform.transforms)} 个步骤\")\nprint(f\"   - 验证变换: {len(val_transform.transforms)} 个步骤\")\nprint(f\"   - 目标图像尺寸: {config.IMAGE_SIZE}x{config.IMAGE_SIZE}\")\nprint(\"   - 支持测试时增强(TTA)和可视化功能\")\n","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:26:52.069148Z","iopub.execute_input":"2025-09-23T14:26:52.069407Z","iopub.status.idle":"2025-09-23T14:26:52.091208Z","shell.execute_reply.started":"2025-09-23T14:26:52.069358Z","shell.execute_reply":"2025-09-23T14:26:52.090509Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 测试和可视化2D灰度图像二分类数据增强\nprint(\"=== 测试2D灰度图像二分类数据增强效果 ===\")\n\n# 创建测试用的虚拟图像\ndef create_test_sample():\n    \"\"\"创建测试用的2D灰度图像\"\"\"\n    # 创建一个模拟的脑血管图像\n    test_image = np.zeros((224, 224), dtype=np.uint8)\n    \n    # 添加一些结构（模拟血管）\n    # 主血管\n    cv2.line(test_image, (50, 50), (170, 170), 200, 3)\n    cv2.line(test_image, (170, 50), (50, 170), 180, 2)\n    \n    # 分支血管\n    cv2.line(test_image, (100, 30), (100, 100), 160, 2)\n    cv2.line(test_image, (30, 100), (100, 100), 160, 2)\n    \n    # 添加一些噪声和纹理\n    noise = np.random.normal(0, 20, (224, 224))\n    test_image = np.clip(test_image.astype(np.float32) + noise, 0, 255).astype(np.uint8)\n    \n    return test_image\n\n# 测试不同增强级别\naugmentation_levels = ['none', 'light', 'medium', 'strong']\n\nfor level in augmentation_levels:\n    print(f\"\\n--- {level.upper()} 级别数据增强 ---\")\n    \n    # 创建对应级别的变换\n    transform = get_classification_transforms_2d_grayscale(\n        image_size=config.IMAGE_SIZE,\n        is_training=True,\n        augmentation_level=level\n    )\n    \n    print(f\"变换步骤数: {len(transform.transforms)}\")\n    print(\"变换列表:\")\n    for i, t in enumerate(transform.transforms):\n        print(f\"  {i+1}. {type(t).__name__}\")\n\n# 可视化数据增强效果（如果需要）\nif hasattr(config, 'VISUALIZE_AUGMENTATION') and config.VISUALIZE_AUGMENTATION:\n    print(\"\\n=== 可视化数据增强效果 ===\")\n    \n    # 创建测试样本\n    test_image = create_test_sample()\n    \n    # 可视化强增强效果\n    strong_transform = get_classification_transforms_2d_grayscale(\n        image_size=config.IMAGE_SIZE,\n        is_training=True,\n        augmentation_level='strong'\n    )\n    \n    print(\"生成数据增强可视化...\")\n    # 注意：这里需要在有matplotlib的环境中运行\n    # visualize_augmentations_2d_classification(test_image, strong_transform, num_samples=4)\n\n# 测试数据增强的性能\nprint(\"\\n=== 数据增强性能测试 ===\")\n\nimport time\n\n# 创建测试数据\ntest_image = create_test_sample()\n\n# 测试不同级别的增强速度\nfor level in ['light', 'medium', 'strong']:\n    transform = get_classification_transforms_2d_grayscale(\n        image_size=config.IMAGE_SIZE,\n        is_training=True,\n        augmentation_level=level\n    )\n    \n    # 性能测试\n    start_time = time.time()\n    for i in range(100):\n        # 移除最后两个步骤（Normalize和ToTensorV2）进行速度测试\n        test_transform = A.Compose(transform.transforms[:-2])\n        augmented = test_transform(image=test_image)\n    \n    end_time = time.time()\n    avg_time = (end_time - start_time) / 100 * 1000  # 转换为毫秒\n    \n    print(f\"{level.upper()} 级别: {avg_time:.2f}ms/样本 ({len(transform.transforms)} 个变换)\")\n\n# 验证数据增强的正确性\nprint(\"\\n=== 数据增强正确性验证 ===\")\n\n# 测试图像变换\ntest_transform = get_classification_transforms_2d_grayscale(\n    image_size=config.IMAGE_SIZE,\n    is_training=True,\n    augmentation_level='medium'\n)\n\n# 移除归一化步骤进行测试\ntest_transform_no_norm = A.Compose(test_transform.transforms[:-2])\n\nfor i in range(3):\n    augmented = test_transform_no_norm(image=test_image)\n    aug_image = augmented['image']\n    \n    print(f\"测试 {i+1}:\")\n    print(f\"  原始图像形状: {test_image.shape}, 数据类型: {test_image.dtype}\")\n    print(f\"  增强图像形状: {aug_image.shape}, 数据类型: {aug_image.dtype}\")\n    print(f\"  像素值范围: {aug_image.min():.2f} - {aug_image.max():.2f}\")\n    \n    # 验证图像形状和数据类型\n    assert aug_image.shape == test_image.shape, \"图像形状不匹配\"\n    assert aug_image.dtype == test_image.dtype, \"数据类型不匹配\"\n\nprint(\"\\n✅ 数据增强正确性验证通过\")\nprint(\"✅ 2D灰度图像二分类数据增强配置和测试完成\")\n\n# 最终配置确认\nprint(f\"\\n=== 最终配置确认 ===\")\nprint(f\"当前增强级别: {'strong' if config.USE_STRONG_AUGMENTATION else 'medium'}\")\nprint(f\"训练变换步骤: {len(train_transform.transforms)}\")\nprint(f\"验证变换步骤: {len(val_transform.transforms)}\")\nprint(f\"图像尺寸: {config.IMAGE_SIZE}x{config.IMAGE_SIZE}\")\nprint(f\"输入通道数: 1 (灰度图像)\")\nprint(f\"任务类型: 二分类 (0/1)\")\nprint(f\"支持TTA: 是\")\n# print(f\"支持可视化: 是\")\"\n","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:26:52.092075Z","iopub.execute_input":"2025-09-23T14:26:52.092305Z","iopub.status.idle":"2025-09-23T14:26:52.684069Z","shell.execute_reply.started":"2025-09-23T14:26:52.09229Z","shell.execute_reply":"2025-09-23T14:26:52.683222Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_robust_cv_split(train_df, n_splits=5):\n    \"\"\"Create robust cross-validation split with true patient separation from DICOM\"\"\"\n    \n    print(\"Creating patient-separated cross-validation split...\")\n    print(\"Extracting true patient IDs from DICOM metadata...\")\n    print(\"This will take a few minutes but ensures proper patient separation.\")\n    \n    # Extract true patient groups from DICOM metadata\n    patient_groups = []\n    for series_uid in tqdm(train_df['SeriesInstanceUID'], desc=\"Reading DICOM patient info\"):\n        patient_group = get_patient_group_cached(series_uid)\n        patient_groups.append(patient_group)\n    \n    # Add patient groups to dataframe\n    train_df = train_df.copy()\n    train_df['patient_id'] = patient_groups\n    \n    n_groups = train_df['patient_id'].nunique()\n    print(f\"True patient groups found: {n_groups}\")\n    \n    # Check if we have enough patient groups\n    if n_groups < n_splits:\n        print(f\"Not enough patient groups ({n_groups}) for {n_splits}-fold CV.\")\n        print(\"Falling back to StratifiedKFold...\")\n        skf = StratifiedKFold(n_splits=n_splits, shuffle=True, random_state=42)\n        return list(skf.split(train_df, train_df['Aneurysm Present']))\n    \n    # Create stratification key combining modality and aneurysm presence\n    train_df['stratify_key'] = (\n        train_df['Modality'].astype(str) + '_' + \n        train_df['Aneurysm Present'].astype(str)\n    )\n    \n    print(f\"Stratification keys: {train_df['stratify_key'].unique()}\")\n    \n    # Use GroupKFold to ensure patient-level separation\n    group_kfold = GroupKFold(n_splits=n_splits)\n    \n    splits = []\n    for fold_idx, (train_idx, val_idx) in enumerate(group_kfold.split(\n        train_df, \n        groups=train_df['patient_id']\n    )):\n        # Validate patient separation\n        train_fold = train_df.iloc[train_idx]\n        val_fold = train_df.iloc[val_idx]\n        \n        # Check for patient overlap (should be 0)\n        train_patients = set(train_fold['patient_id'])\n        val_patients = set(val_fold['patient_id'])\n        overlap = train_patients.intersection(val_patients)\n        \n        train_dist = train_fold['Aneurysm Present'].value_counts(normalize=True)\n        val_dist = val_fold['Aneurysm Present'].value_counts(normalize=True)\n        \n        print(f\"Fold {fold_idx}:\")\n        print(f\"  Train: {len(train_fold)} samples ({len(train_patients)} patients)\")\n        print(f\"  Val: {len(val_fold)} samples ({len(val_patients)} patients)\")\n        print(f\"  Patient overlap: {len(overlap)} (should be 0!)\")\n        print(f\"  Aneurysm Present - Train: {train_dist.get(1, 0):.3f}, Val: {val_dist.get(1, 0):.3f}\")\n        \n        if len(overlap) > 0:\n            print(f\"  WARNING: Found {len(overlap)} overlapping patients!\")\n        \n        splits.append((train_idx, val_idx))\n    \n    return splits\n\n# Create robust train/validation split\ncv_splits = create_robust_cv_split(train_df_filtered, config.NUM_FOLDS)\ntrain_indices, val_indices = cv_splits[config.FOLD]\n\ntrain_fold_df = train_df_filtered.iloc[train_indices]\nval_fold_df = train_df_filtered.iloc[val_indices]\n\nprint(f\"\\nRobust CV Fold {config.FOLD} Summary:\")\nprint(f\"Train fold size: {len(train_fold_df)}\")\nprint(f\"Validation fold size: {len(val_fold_df)}\")\n\n# Check distributions\nprint(f\"Train Aneurysm Present: {train_fold_df['Aneurysm Present'].value_counts().to_dict()}\")\nprint(f\"Val Aneurysm Present: {val_fold_df['Aneurysm Present'].value_counts().to_dict()}\")\n\n# Check modality distribution\nprint(f\"Train Modality distribution: {train_fold_df['Modality'].value_counts().to_dict()}\")\nprint(f\"Val Modality distribution: {val_fold_df['Modality'].value_counts().to_dict()}\")","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:26:52.684884Z","iopub.execute_input":"2025-09-23T14:26:52.68517Z","iopub.status.idle":"2025-09-23T14:31:40.882262Z","shell.execute_reply.started":"2025-09-23T14:26:52.685146Z","shell.execute_reply":"2025-09-23T14:31:40.881441Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class SegmentationDataset(Dataset):\n    \"\"\"Dataset for 2D U-Net segmentation training with foreground pixel filtering\"\"\"\n    def __init__(self, df, segmentation_data_dict, series_mapping_df=None, \n                 transform=None, is_training=True, min_foreground_ratio=0.01):\n        self.df = df.reset_index(drop=True)\n        self.segmentation_data_dict = segmentation_data_dict\n        self.series_mapping_df = series_mapping_df\n        self.transform = transform\n        self.is_training = is_training\n        self.min_foreground_ratio = min_foreground_ratio  # 最小前景像素比例\n        \n        # Create list of (series_uid, slice_idx) pairs for all available slices\n        # 在初始化时就过滤掉前景像素太少的切片\n        self.samples = []\n        self._filter_samples()\n        \n        # Simple LRU cache for recently accessed data\n        self._cache = {}\n        self._cache_keys = []\n        self._max_cache_size = config.CACHE_SIZE\n        \n        print(f\"数据集初始化完成: {len(self.samples)} 个有效样本 (最小前景像素比例: {self.min_foreground_ratio})\")\n    \n    def _filter_samples(self):\n        \"\"\"过滤掉前景像素比例太低的切片\"\"\"\n        total_slices = 0\n        filtered_slices = 0\n        \n        for series_uid in self.df['SeriesInstanceUID'].unique():\n            if series_uid in self.segmentation_data_dict:\n                slices_2d = self.segmentation_data_dict[series_uid]['slices_2d']\n                num_slices = len(slices_2d)\n                total_slices += num_slices\n                \n                for slice_idx in range(num_slices):\n                    # 检查前景像素比例\n                    mask = slices_2d[slice_idx]\n                    foreground_ratio = np.sum(mask > 0) / mask.size\n                    \n                    # 只保留前景像素比例足够的切片\n                    if foreground_ratio >= self.min_foreground_ratio:\n                        self.samples.append((series_uid, slice_idx))\n                        filtered_slices += 1\n        \n        print(f\"切片过滤统计: 总切片 {total_slices}, 保留切片 {filtered_slices}, 过滤率 {(total_slices-filtered_slices)/total_slices*100:.1f}%\")\n    \n    def __len__(self):\n        return len(self.samples)\n    \n    def __getitem__(self, idx):\n        # Check cache first\n        if idx in self._cache:\n            return self._cache[idx]\n        \n        series_uid, slice_idx = self.samples[idx]\n        row = self.df[self.df['SeriesInstanceUID'] == series_uid].iloc[0]\n        \n        # Get segmentation mask\n        mask = self.segmentation_data_dict[series_uid]['slices_2d'][slice_idx]\n        \n        # 再次检查前景像素比例（双重保险）\n        foreground_ratio = np.sum(mask > 0) / mask.size\n        if foreground_ratio < self.min_foreground_ratio:\n            # 如果前景像素太少，返回一个随机的有效样本\n            return self.__getitem__(np.random.randint(0, len(self.samples)))\n        \n        # Load corresponding image slice\n        image = self._load_image_slice(series_uid, slice_idx, row)\n        \n        # Extract metadata\n        metadata = self._extract_metadata(row)\n        \n        # Apply transforms\n        if self.transform:\n            transformed = self.transform(image=image, mask=mask)\n            image = transformed['image']\n            mask = transformed['mask']\n        \n        # Convert mask to tensor\n        mask_tensor = mask.long()\n        \n        result = (image, mask_tensor, metadata)\n        \n        # Update cache\n        self._update_cache(idx, result)\n        \n        return result\n    \n    def _update_cache(self, idx, data):\n        \"\"\"Update LRU cache\"\"\"\n        if len(self._cache) >= self._max_cache_size:\n            # Remove oldest entry\n            oldest_idx = self._cache_keys.pop(0)\n            del self._cache[oldest_idx]\n        \n        self._cache[idx] = data\n        self._cache_keys.append(idx)\n    \n    def _extract_metadata(self, row):\n        \"\"\"Extract metadata from row\"\"\"\n        metadata = {\n            'age': row.get('Patient Age', 0),\n            'sex': 1 if row.get('Patient Sex', 'M') == 'M' else 0,\n            'modality': row.get('Modality', 'CTA')\n        }\n        return metadata\n    \n    def _load_image_slice(self, series_uid: str, slice_idx: int, row) -> np.ndarray:\n        \"\"\"Load real DICOM image slice\"\"\"\n        try:\n            # 构建DICOM文件路径\n            dicom_dir = os.path.join(config.DICOM_SERIES_DIR, series_uid)\n            \n            if not os.path.exists(dicom_dir):\n                print(f\"Warning: DICOM directory not found: {dicom_dir}\")\n                return self._create_fallback_image()\n            \n            # 获取DICOM文件列表\n            dcm_files = [f for f in os.listdir(dicom_dir) if f.endswith('.dcm')]\n            if not dcm_files:\n                print(f\"Warning: No DICOM files found in {dicom_dir}\")\n                return self._create_fallback_image()\n            \n            # 按文件名排序确保顺序一致\n            dcm_files.sort()\n            \n            # 检查slice_idx是否在有效范围内\n            if slice_idx >= len(dcm_files):\n                print(f\"Warning: slice_idx {slice_idx} out of range for {series_uid}\")\n                return self._create_fallback_image()\n            \n            # 加载对应的DICOM文件\n            dcm_path = os.path.join(dicom_dir, dcm_files[slice_idx])\n            \n            try:\n                # 使用pydicom读取DICOM文件\n                ds = pydicom.dcmread(dcm_path, stop_before_pixels=False, force=True)\n                img = ds.pixel_array.astype(np.float32)\n                \n                # 关键检查：跳过3D DICOM文件\n                if len(img.shape) != 2:\n                    # print(f\"跳过3D DICOM文件: {dcm_path}, 形状: {img.shape}\")\n                    return self._create_fallback_image()\n                \n                # 检查图像尺寸有效性\n                if img.shape[0] == 0 or img.shape[1] == 0:\n                    print(f\"无效图像尺寸: {img.shape}\")\n                    return self._create_fallback_image()\n                \n                # 获取模态信息\n                modality = row.get('Modality', 'CTA')\n                \n                # 应用窗宽窗位\n                if config.USE_WINDOWING:\n                    window_center, window_width = get_windowing_params(modality)\n                    img = apply_dicom_windowing(img, window_center, window_width)\n                else:\n                    # 使用鲁棒归一化\n                    img = robust_normalization(img)\n                \n                # 应用CLAHE增强\n                if config.USE_CLAHE:\n                    img = apply_clahe_normalization(img, modality)\n                \n                # 调整到目标尺寸\n                if img.shape != (config.IMAGE_SIZE, config.IMAGE_SIZE):\n                    img = cv2.resize(img, (config.IMAGE_SIZE, config.IMAGE_SIZE), \n                                   interpolation=cv2.INTER_AREA)\n                \n                # 最终检查：确保输出是2D\n                if len(img.shape) != 2:\n                    print(f\"输出图像不是2D: {img.shape}\")\n                    return self._create_fallback_image()\n                \n                return img.astype(np.uint8)\n                \n            except Exception as e:\n                print(f\"Error reading DICOM file {dcm_path}: {e}\")\n                return self._create_fallback_image()\n                \n        except Exception as e:\n            print(f\"Error loading image slice {slice_idx} for {series_uid}: {e}\")\n            return self._create_fallback_image()\n    \n    def _create_fallback_image(self) -> np.ndarray:\n        \"\"\"创建备用图像（当无法加载真实DICOM时）\"\"\"\n        # 创建一个简单的测试图像而不是随机噪声\n        img = np.zeros((config.IMAGE_SIZE, config.IMAGE_SIZE), dtype=np.uint8)\n        \n        # 添加一些简单的几何形状作为测试\n        center = config.IMAGE_SIZE // 2\n        cv2.circle(img, (center, center), 30, 128, -1)\n        cv2.rectangle(img, (center-20, center-20), (center+20, center+20), 200, 2)\n        \n        return img\n","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:31:40.883196Z","iopub.execute_input":"2025-09-23T14:31:40.88345Z","iopub.status.idle":"2025-09-23T14:31:40.901794Z","shell.execute_reply.started":"2025-09-23T14:31:40.88343Z","shell.execute_reply":"2025-09-23T14:31:40.901262Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# Create segmentation datasets\nprint(\"Creating segmentation datasets...\")\ntrain_dataset = SegmentationDataset(\n    train_fold_df, \n    segmentation_data_dict, \n    series_mapping_df=None,  # 不再依赖series_mapping_df\n    transform=train_transform,\n    is_training=True,\n    min_foreground_ratio=0.0\n)\n\nval_dataset = SegmentationDataset(\n    val_fold_df,\n    segmentation_data_dict,\n    series_mapping_df=None,  # 不再依赖series_mapping_df\n    transform=val_transform,\n    is_training=False,\n    min_foreground_ratio=0.0\n)\n\n# Create optimized data loaders\nprint(\"Creating optimized data loaders for segmentation...\")\ntrain_loader = DataLoader(\n    train_dataset,\n    batch_size=config.BATCH_SIZE,\n    shuffle=True,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=config.PIN_MEMORY,\n    drop_last=True,\n    prefetch_factor=config.PREFETCH_FACTOR,\n    persistent_workers=config.PERSISTENT_WORKERS\n)\n\nval_loader = DataLoader(\n    val_dataset,\n    batch_size=config.BATCH_SIZE,\n    shuffle=False,\n    num_workers=config.NUM_WORKERS,\n    pin_memory=config.PIN_MEMORY,\n    prefetch_factor=config.PREFETCH_FACTOR,\n    persistent_workers=config.PERSISTENT_WORKERS\n)\n\nprint(f\"Train batches: {len(train_loader)}\")\nprint(f\"Validation batches: {len(val_loader)}\")\n\n# Test segmentation data loading speed and analyze label pixels\nprint(\"Testing segmentation data loading speed...\")\nimport time\n\nstart_time = time.time()\nlabel_stats = {\n    'unique_values': set(),\n    'value_counts': {},\n    'total_pixels': 0,\n    'non_zero_pixels': 0,\n    'class_distribution': {i: 0 for i in range(config.NUM_CLASSES)}\n}\n\nfor i, batch in enumerate(train_loader):\n    if i >= 5:  # Test first 5 batches\n        break\n    images, masks, metadata = batch\n    print(f\"Batch {i+1}: Images shape: {images.shape}, Masks shape: {masks.shape}, Device: {images.device}\")\n    \n    # 分析标签像素分布\n    for mask in masks:\n        # 转换为numpy进行分析\n        mask_np = mask.numpy()\n        \n        # 收集唯一值\n        unique_vals = np.unique(mask_np)\n        label_stats['unique_values'].update(unique_vals)\n        \n        # 计算每个类别的像素数量\n        for val in unique_vals:\n            count = np.sum(mask_np == val)\n            if val not in label_stats['value_counts']:\n                label_stats['value_counts'][val] = 0\n            label_stats['value_counts'][val] += count\n            \n            # 统计类别分布\n            if val < config.NUM_CLASSES:\n                label_stats['class_distribution'][val] += count\n        \n        # 统计总像素和非零像素\n        total_pixels = mask_np.size\n        non_zero_pixels = np.sum(mask_np > 0)\n        label_stats['total_pixels'] += total_pixels\n        label_stats['non_zero_pixels'] += non_zero_pixels\n\nelapsed = time.time() - start_time\nprint(f\"Loaded 5 batches in {elapsed:.2f} seconds ({elapsed/5:.2f}s per batch)\")\n\n# 打印标签像素分析结果\nprint(\"\\n\" + \"=\"*60)\nprint(\"分割标签像素分析结果\")\nprint(\"=\"*60)\n\nprint(f\"发现的唯一标签值: {sorted(label_stats['unique_values'])}\")\nprint(f\"预期类别数: {config.NUM_CLASSES}\")\n\nprint(f\"\\n标签值像素统计:\")\nfor val in sorted(label_stats['value_counts'].keys()):\n    count = label_stats['value_counts'][val]\n    percentage = (count / label_stats['total_pixels']) * 100 if label_stats['total_pixels'] > 0 else 0\n    class_name = list(CLASS_MAPPING.keys())[val] if val < len(CLASS_MAPPING) else f\"未知类别{val}\"\n    print(f\"  标签 {val} ({class_name}): {count:,} 像素 ({percentage:.2f}%)\")\n\nprint(f\"\\n类别分布:\")\nfor class_id in range(config.NUM_CLASSES):\n    count = label_stats['class_distribution'][class_id]\n    percentage = (count / label_stats['total_pixels']) * 100 if label_stats['total_pixels'] > 0 else 0\n    class_name = list(CLASS_MAPPING.keys())[class_id] if class_id < len(CLASS_MAPPING) else f\"类别{class_id}\"\n    print(f\"  {class_name}: {count:,} 像素 ({percentage:.2f}%)\")\n\nprint(f\"\\n总体统计:\")\nprint(f\"  总像素数: {label_stats['total_pixels']:,}\")\nprint(f\"  非零像素数: {label_stats['non_zero_pixels']:,}\")\nprint(f\"  背景像素数: {label_stats['total_pixels'] - label_stats['non_zero_pixels']:,}\")\nprint(f\"  非零像素比例: {(label_stats['non_zero_pixels'] / label_stats['total_pixels']) * 100:.2f}%\")\n\n# 检查数据平衡性\nprint(f\"\\n数据平衡性分析:\")\nbackground_pixels = label_stats['class_distribution'][0]\nforeground_pixels = label_stats['non_zero_pixels'] - background_pixels\nif foreground_pixels > 0:\n    imbalance_ratio = background_pixels / foreground_pixels\n    print(f\"  背景/前景像素比例: {imbalance_ratio:.2f}:1\")\n    print(f\"  数据不平衡程度: {'严重' if imbalance_ratio > 100 else '中等' if imbalance_ratio > 10 else '轻微'}\")\n\nprint(\"\\n\" + \"=\"*60)\n\n# 检查图片和标签的压缩情况\nprint(\"\\n检查图片和标签的压缩情况:\")\nprint(\"=\"*60)\n\nif len(train_dataset) > 0:\n    sample_image, sample_mask, sample_metadata = train_dataset[0]\n    \n    # 检查图像\n    if isinstance(sample_image, torch.Tensor):\n        image_np = sample_image.squeeze().numpy()\n        print(f\"图像形状: {sample_image.shape}\")\n        print(f\"图像数据类型: {sample_image.dtype}\")\n        print(f\"图像值范围: {sample_image.min().item():.3f} - {sample_image.max().item():.3f}\")\n        print(f\"图像是否归一化: {'是' if sample_image.min() >= 0 and sample_image.max() <= 1 else '否'}\")\n    else:\n        image_np = sample_image\n        print(f\"图像形状: {sample_image.shape}\")\n        print(f\"图像数据类型: {sample_image.dtype}\")\n        print(f\"图像值范围: {sample_image.min():.3f} - {sample_image.max():.3f}\")\n    \n    # 检查标签\n    if isinstance(sample_mask, torch.Tensor):\n        mask_np = sample_mask.numpy()\n        print(f\"标签形状: {sample_mask.shape}\")\n        print(f\"标签数据类型: {sample_mask.dtype}\")\n        print(f\"标签值范围: {sample_mask.min().item()} - {sample_mask.max().item()}\")\n    else:\n        mask_np = sample_mask\n        print(f\"标签形状: {sample_mask.shape}\")\n        print(f\"标签数据类型: {sample_mask.dtype}\")\n        print(f\"标签值范围: {sample_mask.min()} - {sample_mask.max()}\")\n    \n    # 检查标签的完整性\n    unique_vals = np.unique(mask_np)\n    print(f\"标签唯一值: {unique_vals}\")\n    print(f\"预期类别数: {config.NUM_CLASSES}\")\n    \n    # 检查是否有压缩或数据丢失\n    if len(unique_vals) > config.NUM_CLASSES:\n        print(f\"⚠️ 警告: 发现 {len(unique_vals)} 个唯一值，超过预期的 {config.NUM_CLASSES} 个类别\")\n    \n    # 检查标签值是否连续\n    expected_vals = set(range(config.NUM_CLASSES))\n    actual_vals = set(unique_vals)\n    missing_vals = expected_vals - actual_vals\n    extra_vals = actual_vals - expected_vals\n    \n    if missing_vals:\n        print(f\"⚠️ 警告: 缺失的标签值: {missing_vals}\")\n    if extra_vals:\n        print(f\"⚠️ 警告: 额外的标签值: {extra_vals}\")\n    \n    # 检查数据压缩\n    print(f\"\\n数据压缩检查:\")\n    print(f\"图像内存占用: {sample_image.element_size() * sample_image.nelement() / 1024:.2f} KB\")\n    print(f\"标签内存占用: {sample_mask.element_size() * sample_mask.nelement() / 1024:.2f} KB\")\n    \n    # 检查标签分布\n    for val in unique_vals:\n        count = np.sum(mask_np == val)\n        percentage = (count / mask_np.size) * 100\n        class_name = list(CLASS_MAPPING.keys())[val] if val < len(CLASS_MAPPING) else f\"未知类别{val}\"\n        print(f\"  标签 {val} ({class_name}): {count:,} 像素 ({percentage:.2f}%)\")\n\nprint(\"\\n\" + \"=\"*60)\n","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:31:40.902652Z","iopub.execute_input":"2025-09-23T14:31:40.902864Z","iopub.status.idle":"2025-09-23T14:31:42.792601Z","shell.execute_reply.started":"2025-09-23T14:31:40.902848Z","shell.execute_reply":"2025-09-23T14:31:42.791321Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNet2D(nn.Module):\n    \"\"\"2D U-Net architecture for medical image segmentation\"\"\"\n    def __init__(self, in_channels=1, num_classes=14, base_features=64):\n        super(UNet2D, self).__init__()\n        self.num_classes = num_classes\n        \n        # Encoder (Contracting Path)\n        self.enc1 = self._conv_block(in_channels, base_features)\n        self.enc2 = self._conv_block(base_features, base_features * 2)\n        self.enc3 = self._conv_block(base_features * 2, base_features * 4)\n        self.enc4 = self._conv_block(base_features * 4, base_features * 8)\n        \n        # Bottleneck\n        self.bottleneck = self._conv_block(base_features * 8, base_features * 16)\n        \n        # Decoder (Expanding Path)\n        self.upconv4 = nn.ConvTranspose2d(base_features * 16, base_features * 8, kernel_size=2, stride=2)\n        self.dec4 = self._conv_block(base_features * 16, base_features * 8)\n        \n        self.upconv3 = nn.ConvTranspose2d(base_features * 8, base_features * 4, kernel_size=2, stride=2)\n        self.dec3 = self._conv_block(base_features * 8, base_features * 4)\n        \n        self.upconv2 = nn.ConvTranspose2d(base_features * 4, base_features * 2, kernel_size=2, stride=2)\n        self.dec2 = self._conv_block(base_features * 4, base_features * 2)\n        \n        self.upconv1 = nn.ConvTranspose2d(base_features * 2, base_features, kernel_size=2, stride=2)\n        self.dec1 = self._conv_block(base_features * 2, base_features)\n        \n        # Final classification layer\n        self.final_conv = nn.Conv2d(base_features, num_classes, kernel_size=1)\n        \n        # Max pooling\n        self.pool = nn.MaxPool2d(kernel_size=2, stride=2)\n        \n        # Dropout for regularization\n        self.dropout = nn.Dropout2d(0.2)\n        \n    def _conv_block(self, in_channels, out_channels):\n        \"\"\"Convolutional block with two conv layers\"\"\"\n        return nn.Sequential(\n            nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True),\n            nn.Conv2d(out_channels, out_channels, kernel_size=3, padding=1),\n            nn.BatchNorm2d(out_channels),\n            nn.ReLU(inplace=True)\n        )\n    \n    def forward(self, x):\n        # Encoder\n        enc1 = self.enc1(x)\n        enc2 = self.enc2(self.pool(enc1))\n        enc3 = self.enc3(self.pool(enc2))\n        enc4 = self.enc4(self.pool(enc3))\n        \n        # Bottleneck\n        bottleneck = self.bottleneck(self.pool(enc4))\n        bottleneck = self.dropout(bottleneck)\n        \n        # Decoder with skip connections\n        dec4 = self.upconv4(bottleneck)\n        dec4 = torch.cat((dec4, enc4), dim=1)\n        dec4 = self.dec4(dec4)\n        \n        dec3 = self.upconv3(dec4)\n        dec3 = torch.cat((dec3, enc3), dim=1)\n        dec3 = self.dec3(dec3)\n        \n        dec2 = self.upconv2(dec3)\n        dec2 = torch.cat((dec2, enc2), dim=1)\n        dec2 = self.dec2(dec2)\n        \n        dec1 = self.upconv1(dec2)\n        dec1 = torch.cat((dec1, enc1), dim=1)\n        dec1 = self.dec1(dec1)\n        \n        # Final classification\n        output = self.final_conv(dec1)\n        \n        return output\n\n# Initialize 2D U-Net model\nprint(\"Initializing 2D U-Net model...\")\nmodel = UNet2D(\n    in_channels=1,  # Grayscale input\n    num_classes=config.NUM_CLASSES,\n    base_features=64\n)\n\nmodel = model.to(device)\n\n# Count parameters\ntotal_params = sum(p.numel() for p in model.parameters())\ntrainable_params = sum(p.numel() for p in model.parameters() if p.requires_grad)\n\nprint(f\"Total parameters: {total_params:,}\")\nprint(f\"Trainable parameters: {trainable_params:,}\")\nprint(f\"Model device: {next(model.parameters()).device}\")\n","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:31:42.797185Z","iopub.execute_input":"2025-09-23T14:31:42.797485Z","iopub.status.idle":"2025-09-23T14:31:43.249041Z","shell.execute_reply.started":"2025-09-23T14:31:42.797449Z","shell.execute_reply":"2025-09-23T14:31:43.248127Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DiceLoss(nn.Module):\n    \"\"\"Dice Loss for segmentation\"\"\"\n    def __init__(self, smooth=1e-5):\n        super(DiceLoss, self).__init__()\n        self.smooth = smooth\n    \n    def forward(self, inputs, targets):\n        # Convert to probabilities\n        inputs = torch.softmax(inputs, dim=1)\n        \n        # Flatten tensors\n        inputs = inputs.contiguous().view(-1)\n        targets = targets.contiguous().view(-1)\n        \n        # Calculate Dice coefficient\n        intersection = (inputs * targets).sum()\n        dice = (2. * intersection + self.smooth) / (inputs.sum() + targets.sum() + self.smooth)\n        \n        return 1 - dice\n\nclass CombinedSegmentationLoss(nn.Module):\n    \"\"\"Combined Dice Loss and Cross Entropy Loss for segmentation\"\"\"\n    def __init__(self, dice_weight=0.5, ce_weight=0.5, class_weights=None):\n        super(CombinedSegmentationLoss, self).__init__()\n        self.dice_weight = dice_weight\n        self.ce_weight = ce_weight\n        \n        self.dice_loss = DiceLoss()\n        self.ce_loss = nn.CrossEntropyLoss(weight=class_weights)\n    \n    def forward(self, inputs, targets):\n        # Cross Entropy Loss\n        ce_loss = self.ce_loss(inputs, targets)\n        \n        # Dice Loss (convert targets to one-hot for dice calculation)\n        targets_one_hot = F.one_hot(targets, num_classes=config.NUM_CLASSES).permute(0, 3, 1, 2).float()\n        dice_loss = self.dice_loss(inputs, targets_one_hot)\n        \n        # Combined loss\n        total_loss = self.ce_weight * ce_loss + self.dice_weight * dice_loss\n        \n        return total_loss\n\ndef calculate_dice_score(pred, target, num_classes, smooth=1e-5):\n    \"\"\"Calculate Dice score for each class\"\"\"\n    pred = torch.softmax(pred, dim=1)\n    pred = torch.argmax(pred, dim=1)\n    \n    dice_scores = []\n    for i in range(num_classes):\n        pred_i = (pred == i).float()\n        target_i = (target == i).float()\n        \n        intersection = (pred_i * target_i).sum()\n        dice = (2. * intersection + smooth) / (pred_i.sum() + target_i.sum() + smooth)\n        dice_scores.append(dice.item())\n    \n    return dice_scores\n\ndef calculate_iou(pred, target, num_classes, smooth=1e-5):\n    \"\"\"Calculate IoU for each class\"\"\"\n    pred = torch.softmax(pred, dim=1)\n    pred = torch.argmax(pred, dim=1)\n    \n    iou_scores = []\n    for i in range(num_classes):\n        pred_i = (pred == i).float()\n        target_i = (target == i).float()\n        \n        intersection = (pred_i * target_i).sum()\n        union = pred_i.sum() + target_i.sum() - intersection\n        iou = (intersection + smooth) / (union + smooth)\n        iou_scores.append(iou.item())\n    \n    return iou_scores\n\n# Training setup\ncriterion = CombinedSegmentationLoss(\n    dice_weight=config.DICE_WEIGHT,\n    ce_weight=config.BCE_WEIGHT\n)\noptimizer = AdamW(model.parameters(), lr=config.LEARNING_RATE, weight_decay=1e-4)\nscheduler = CosineAnnealingLR(optimizer, T_max=config.NUM_EPOCHS, eta_min=1e-6)\n\n# Mixed precision training\nscaler = torch.cuda.amp.GradScaler()\n\nprint(\"Training setup complete\")\nprint(f\"Using loss function: {type(criterion).__name__}\")\nprint(f\"Dice weight: {config.DICE_WEIGHT}, CE weight: {config.BCE_WEIGHT}\")\n","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:31:43.249948Z","iopub.execute_input":"2025-09-23T14:31:43.250198Z","iopub.status.idle":"2025-09-23T14:31:43.26895Z","shell.execute_reply.started":"2025-09-23T14:31:43.250178Z","shell.execute_reply":"2025-09-23T14:31:43.268296Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 训练和验证函数定义\ndef train_epoch_segmentation(model, train_loader, criterion, optimizer, scaler, device, accumulation_steps):\n    \"\"\"训练一个epoch的分割模型\"\"\"\n    model.train()\n    total_loss = 0.0\n    total_dice = 0.0\n    total_iou = 0.0\n    num_batches = 0\n    \n    optimizer.zero_grad()\n    \n    # 使用tqdm显示训练进度\n    pbar = tqdm(train_loader, desc=\"训练中\", leave=False)\n    \n    for batch_idx, (images, masks, metadata) in enumerate(pbar):\n        images = images.to(device, non_blocking=True).half()\n        masks = masks.to(device, non_blocking=True)\n        \n        # 前向传播\n        with torch.cuda.amp.autocast():\n            outputs = model(images)\n            loss = criterion(outputs, masks)\n            \n            # 计算Dice和IoU\n            dice_scores = calculate_dice_score(outputs, masks, config.NUM_CLASSES)\n            iou_scores = calculate_iou(outputs, masks, config.NUM_CLASSES)\n            \n            # 平均Dice和IoU\n            avg_dice = np.mean(dice_scores)\n            avg_iou = np.mean(iou_scores)\n        \n        # 反向传播（梯度累积）\n        loss = loss / accumulation_steps\n        scaler.scale(loss).backward()\n        \n        if (batch_idx + 1) % accumulation_steps == 0:\n            scaler.step(optimizer)\n            scaler.update()\n            optimizer.zero_grad()\n        \n        total_loss += loss.item() * accumulation_steps\n        total_dice += avg_dice\n        total_iou += avg_iou\n        num_batches += 1\n        \n        # 更新进度条显示\n        pbar.set_postfix({\n            'Loss': f'{loss.item() * accumulation_steps:.4f}',\n            'Dice': f'{avg_dice:.4f}',\n            'IoU': f'{avg_iou:.4f}'\n        })\n    \n    avg_loss = total_loss / num_batches\n    avg_dice = total_dice / num_batches\n    avg_iou = total_iou / num_batches\n    \n    return avg_loss, avg_dice, avg_iou\n\ndef validate_epoch_segmentation(model, val_loader, criterion, device):\n    \"\"\"验证一个epoch的分割模型\"\"\"\n    model.eval()\n    total_loss = 0.0\n    total_dice = 0.0\n    total_iou = 0.0\n    num_batches = 0\n    \n    all_dice_per_class = [[] for _ in range(config.NUM_CLASSES)]\n    all_iou_per_class = [[] for _ in range(config.NUM_CLASSES)]\n    \n    with torch.no_grad():\n        # 使用tqdm显示验证进度\n        pbar = tqdm(val_loader, desc=\"验证中\", leave=False)\n        \n        for images, masks, metadata in pbar:\n            images = images.to(device, non_blocking=True).half()\n            masks = masks.to(device, non_blocking=True)\n            \n            # 前向传播\n            with torch.cuda.amp.autocast():\n                outputs = model(images)\n                loss = criterion(outputs, masks)\n                \n                # 计算Dice和IoU\n                dice_scores = calculate_dice_score(outputs, masks, config.NUM_CLASSES)\n                iou_scores = calculate_iou(outputs, masks, config.NUM_CLASSES)\n                \n                # 平均Dice和IoU\n                avg_dice = np.mean(dice_scores)\n                avg_iou = np.mean(iou_scores)\n            \n            total_loss += loss.item()\n            total_dice += avg_dice\n            total_iou += avg_iou\n            num_batches += 1\n            \n            # 收集每个类别的分数\n            for i in range(config.NUM_CLASSES):\n                all_dice_per_class[i].append(dice_scores[i])\n                all_iou_per_class[i].append(iou_scores[i])\n            \n            # 更新进度条显示\n            pbar.set_postfix({\n                'Loss': f'{loss.item():.4f}',\n                'Dice': f'{avg_dice:.4f}',\n                'IoU': f'{avg_iou:.4f}'\n            })\n    \n    avg_loss = total_loss / num_batches\n    avg_dice = total_dice / num_batches\n    avg_iou = total_iou / num_batches\n    \n    # 计算每个类别的平均分数\n    dice_per_class = [np.mean(scores) for scores in all_dice_per_class]\n    iou_per_class = [np.mean(scores) for scores in all_iou_per_class]\n    \n    return avg_loss, avg_dice, avg_iou, dice_per_class, iou_per_class\n\ndef check_gpu_utilization():\n    \"\"\"检查GPU利用率\"\"\"\n    if torch.cuda.is_available():\n        gpu_memory = torch.cuda.memory_allocated() / 1024**3\n        gpu_memory_max = torch.cuda.max_memory_allocated() / 1024**3\n        return f\"GPU内存: {gpu_memory:.2f}GB / {gpu_memory_max:.2f}GB\"\n    return \"GPU不可用\"\n\nprint(\"✅ 训练和验证函数已定义\")\nprint(\"   - train_epoch_segmentation: 训练一个epoch\")\nprint(\"   - validate_epoch_segmentation: 验证一个epoch\")\nprint(\"   - check_gpu_utilization: 检查GPU利用率\")\n","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:31:43.269933Z","iopub.execute_input":"2025-09-23T14:31:43.270547Z","iopub.status.idle":"2025-09-23T14:31:43.295937Z","shell.execute_reply.started":"2025-09-23T14:31:43.270522Z","shell.execute_reply":"2025-09-23T14:31:43.29509Z"},"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 5-Fold Cross Validation Training Loop\nprint(\"=== 开始5折交叉验证训练 ===\")\nprint(f\"总折数: {config.NUM_FOLDS}\")\nprint(f\"将依次训练所有 {config.NUM_FOLDS} 个fold\")\n\n# 存储所有fold的结果\nall_fold_results = []\n\n# 循环训练每个fold\nfor fold_idx in range(config.NUM_FOLDS):\n    print(f\"\\n{'='*60}\")\n    print(f\"开始训练 Fold {fold_idx + 1}/{config.NUM_FOLDS}\")\n    print(f\"{'='*60}\")\n    \n    # 获取当前fold的数据划分\n    train_indices, val_indices = cv_splits[fold_idx]\n    \n    train_fold_df = train_df_filtered.iloc[train_indices]\n    val_fold_df = train_df_filtered.iloc[val_indices]\n    \n    print(f\"Fold {fold_idx} 数据划分:\")\n    print(f\"  训练集大小: {len(train_fold_df)}\")\n    print(f\"  验证集大小: {len(val_fold_df)}\")\n    \n    # 检查数据分布\n    train_dist = train_fold_df['Aneurysm Present'].value_counts(normalize=True)\n    val_dist = val_fold_df['Aneurysm Present'].value_counts(normalize=True)\n    print(f\"  训练集动脉瘤阳性比例: {train_dist.get(1, 0):.3f}\")\n    print(f\"  验证集动脉瘤阳性比例: {val_dist.get(1, 0):.3f}\")\n    \n    # 检查模态分布\n    train_modality = train_fold_df['Modality'].value_counts().to_dict()\n    val_modality = val_fold_df['Modality'].value_counts().to_dict()\n    print(f\"  训练集模态分布: {train_modality}\")\n    print(f\"  验证集模态分布: {val_modality}\")\n    \n    # 创建当前fold的数据集\n    print(f\"\\n创建 Fold {fold_idx} 的数据集...\")\n    train_dataset = SegmentationDataset(\n        train_fold_df, \n        segmentation_data_dict, \n        series_mapping_df=None,  # 不再依赖系列映射\n        transform=train_transform,\n        is_training=True,\n        min_foreground_ratio=0.0\n    )\n    \n    val_dataset = SegmentationDataset(\n        val_fold_df,\n        segmentation_data_dict,\n        series_mapping_df=None,  # 不再依赖系列映射\n        transform=val_transform,\n        is_training=False,\n        min_foreground_ratio=0.0\n    )\n    \n    # 创建数据加载器\n    train_loader = DataLoader(\n        train_dataset,\n        batch_size=config.BATCH_SIZE,\n        shuffle=True,\n        num_workers=config.NUM_WORKERS,\n        pin_memory=config.PIN_MEMORY,\n        drop_last=True,\n        prefetch_factor=config.PREFETCH_FACTOR,\n        persistent_workers=config.PERSISTENT_WORKERS\n    )\n    \n    val_loader = DataLoader(\n        val_dataset,\n        batch_size=config.BATCH_SIZE,\n        shuffle=False,\n        num_workers=config.NUM_WORKERS,\n        pin_memory=config.PIN_MEMORY,\n        prefetch_factor=config.PREFETCH_FACTOR,\n        persistent_workers=config.PERSISTENT_WORKERS\n    )\n    \n    print(f\"Fold {fold_idx} 数据加载器:\")\n    print(f\"  训练批次数: {len(train_loader)}\")\n    print(f\"  验证批次数: {len(val_loader)}\")\n    \n    # 初始化模型（每个fold使用新的模型实例）\n    print(f\"\\n初始化 Fold {fold_idx} 的模型...\")\n    model = UNet2D(\n        in_channels=1,  # 灰度输入\n        num_classes=config.NUM_CLASSES,\n        base_features=64\n    )\n    model = model.to(device)\n    \n    # 训练设置\n    criterion = CombinedSegmentationLoss(\n        dice_weight=config.DICE_WEIGHT,\n        ce_weight=config.BCE_WEIGHT\n    )\n    optimizer = AdamW(model.parameters(), lr=config.LEARNING_RATE, weight_decay=1e-4)\n    scheduler = CosineAnnealingLR(optimizer, T_max=config.NUM_EPOCHS, eta_min=1e-6)\n    scaler = torch.cuda.amp.GradScaler()\n    \n    # 训练循环变量\n    best_dice = 0.0\n    best_epoch = 0\n    patience_counter = 0\n    train_losses = []\n    val_losses = []\n    val_dice_scores = []\n    val_iou_scores = []\n    \n    print(f\"\\n开始 Fold {fold_idx} 的训练...\")\n    print(f\"批次大小: {config.BATCH_SIZE}, 工作进程: {config.NUM_WORKERS}\")\n    print(f\"图像大小: {config.IMAGE_SIZE}\")\n    print(f\"类别数: {config.NUM_CLASSES}\")\n    print(f\"CLAHE启用: {config.USE_CLAHE}\")\n    print(f\"强数据增强: {config.USE_STRONG_AUGMENTATION}\")\n    print(f\"真实患者分离: {config.USE_GROUP_CV}\")\n    \n    # 训练循环\n    for epoch in range(config.NUM_EPOCHS):\n        print(f\"\\nFold {fold_idx} - Epoch {epoch+1}/{config.NUM_EPOCHS}\")\n        print(\"-\" * 50)\n        \n        # 训练\n        train_loss, train_dice, train_iou = train_epoch_segmentation(\n            model, train_loader, criterion, optimizer, scaler, device, config.ACCUMULATION_STEPS\n        )\n        \n        # 验证\n        val_loss, val_dice, val_iou, val_dice_per_class, val_iou_per_class = validate_epoch_segmentation(\n            model, val_loader, criterion, device\n        )\n        \n        # 学习率调度\n        scheduler.step()\n        \n        # 记录指标\n        train_losses.append(train_loss)\n        val_losses.append(val_loss)\n        val_dice_scores.append(val_dice)\n        val_iou_scores.append(val_iou)\n        \n        print(f\"训练损失: {train_loss:.6f}, 训练Dice: {train_dice:.6f}, 训练IoU: {train_iou:.6f}\")\n        print(f\"验证损失: {val_loss:.6f}, 验证Dice: {val_dice:.6f}, 验证IoU: {val_iou:.6f}\")\n        print(f\"学习率: {optimizer.param_groups[0]['lr']:.8f}\")\n        \n        # 打印前5个类别的Dice分数\n        print(\"各类别Dice分数 (前5个类别):\")\n        for i in range(min(5, len(val_dice_per_class))):\n            class_name = list(CLASS_MAPPING.keys())[i] if i < len(CLASS_MAPPING) else f\"类别 {i}\"\n            print(f\"  {class_name}: {val_dice_per_class[i]:.4f}\")\n        \n        # GPU利用率\n        gpu_util = check_gpu_utilization()\n        \n        # 早停和模型保存\n        if val_dice > best_dice:\n            best_dice = val_dice\n            best_epoch = epoch + 1\n            patience_counter = 0\n            \n            # 保存模型（包含fold信息）\n            model_path = os.path.join(config.OUTPUT_DIR, f\"{config.MODEL_NAME}_fold{fold_idx}_best.pth\")\n            torch.save({\n                'epoch': epoch + 1,\n                'fold': fold_idx,\n                'model_state_dict': model.state_dict(),\n                'optimizer_state_dict': optimizer.state_dict(),\n                'scheduler_state_dict': scheduler.state_dict(),\n                'best_dice': best_dice,\n                'val_loss': val_loss,\n                'val_dice': val_dice,\n                'val_iou': val_iou,\n                'val_dice_per_class': val_dice_per_class,\n                'val_iou_per_class': val_iou_per_class,\n                'config': config,\n                'model_config': {\n                    'model_type': config.MODEL_TYPE,\n                    'num_classes': config.NUM_CLASSES,\n                    'use_clahe': config.USE_CLAHE,\n                    'use_strong_augmentation': config.USE_STRONG_AUGMENTATION,\n                    'use_group_cv': config.USE_GROUP_CV,\n                    'dice_weight': config.DICE_WEIGHT,\n                    'ce_weight': config.BCE_WEIGHT\n                }\n            }, model_path)\n            \n            print(f\"新的最佳模型已保存! Dice: {best_dice:.6f}\")\n        else:\n            patience_counter += 1\n            print(f\"无改善. 耐心计数: {patience_counter}/{config.EARLY_STOPPING_PATIENCE}\")\n            \n            if patience_counter >= config.EARLY_STOPPING_PATIENCE:\n                print(f\"早停触发于 epoch {epoch + 1}\")\n                break\n        \n        # 内存清理\n        torch.cuda.empty_cache()\n    \n    # 记录当前fold的结果\n    fold_result = {\n        'fold': fold_idx,\n        'best_dice': best_dice,\n        'best_epoch': best_epoch,\n        'final_val_loss': val_loss,\n        'final_val_dice': val_dice,\n        'final_val_iou': val_iou,\n        'train_losses': train_losses,\n        'val_losses': val_losses,\n        'val_dice_scores': val_dice_scores,\n        'val_iou_scores': val_iou_scores,\n        'val_dice_per_class': val_dice_per_class,\n        'val_iou_per_class': val_iou_per_class\n    }\n    all_fold_results.append(fold_result)\n    \n    print(f\"\\nFold {fold_idx} 训练完成!\")\n    print(f\"最佳Dice分数: {best_dice:.6f} (Epoch {best_epoch})\")\n    print(f\"最终验证Dice: {val_dice:.6f}\")\n    print(f\"最终验证IoU: {val_iou:.6f}\")\n    \n    # 清理当前fold的模型和数据集\n    del model, train_dataset, val_dataset, train_loader, val_loader\n    torch.cuda.empty_cache()\n    gc.collect()\n\nprint(f\"\\n{'='*70}\")\nprint(\"所有5个fold训练完成!\")\nprint(f\"{'='*70}\")\n\n# 计算和显示所有fold的统计结果\nall_best_dices = [result['best_dice'] for result in all_fold_results]\nall_final_dices = [result['final_val_dice'] for result in all_fold_results]\nall_final_ious = [result['final_val_iou'] for result in all_fold_results]\n\nprint(f\"\\n=== 5折交叉验证结果汇总 ===\")\nprint(f\"最佳Dice分数:\")\nfor i, dice in enumerate(all_best_dices):\n    print(f\"  Fold {i}: {dice:.6f}\")\n\nprint(f\"\\n最终验证Dice分数:\")\nfor i, dice in enumerate(all_final_dices):\n    print(f\"  Fold {i}: {dice:.6f}\")\n\nprint(f\"\\n最终验证IoU分数:\")\nfor i, iou in enumerate(all_final_ious):\n    print(f\"  Fold {i}: {iou:.6f}\")\n\nprint(f\"\\n=== 统计摘要 ===\")\nprint(f\"最佳Dice分数 - 平均: {np.mean(all_best_dices):.6f}, 标准差: {np.std(all_best_dices):.6f}\")\nprint(f\"最终验证Dice - 平均: {np.mean(all_final_dices):.6f}, 标准差: {np.std(all_final_dices):.6f}\")\nprint(f\"最终验证IoU - 平均: {np.mean(all_final_ious):.6f}, 标准差: {np.std(all_final_ious):.6f}\")\n\nprint(f\"\\n=== 模型文件保存位置 ===\")\nfor i in range(config.NUM_FOLDS):\n    model_path = os.path.join(config.OUTPUT_DIR, f\"{config.MODEL_NAME}_fold{i}_best.pth\")\n    if os.path.exists(model_path):\n        file_size = os.path.getsize(model_path) / (1024*1024)\n        print(f\"  Fold {i}: {model_path} ({file_size:.1f} MB)\")\n\nprint(f\"\\n所有5个fold的模型已保存，可用于集成推理!\")\n","metadata":{"execution":{"iopub.status.busy":"2025-09-23T14:31:43.297099Z","iopub.execute_input":"2025-09-23T14:31:43.297597Z","execution_failed":"2025-09-23T14:43:04.084Z"},"trusted":true},"outputs":[],"execution_count":null}]}