{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":52254,"databundleVersionId":9674523,"sourceType":"competition"},{"sourceId":11555525,"sourceType":"datasetVersion","datasetId":7245708},{"sourceId":11593143,"sourceType":"datasetVersion","datasetId":7269783}],"dockerImageVersionId":31011,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Cell 0: 设置 Matplotlib 中文字体 (Kaggle 环境) \nimport matplotlib\nimport matplotlib.pyplot as plt\nimport matplotlib.font_manager as fm\nimport os\nimport requests # 保留以备未来可能需要下载\nimport zipfile\nimport shutil\n\n# --- 中文字体设置 ---\nprint(\"开始设置 Matplotlib 中文字体 (V5 - 使用上传字体)...\")\n\nfont_file_name = \"SIMHEI.TTF\" # 从你的日志看，似乎是这个文件名\n\nfont_dataset_name = 'simhei'\nfont_file_uploaded_path = f\"/kaggle/input/{font_dataset_name}/{font_file_name}\"\n\n# 目标用户字体目录 (Kaggle中通常可写)\nuser_font_dir = os.path.join(os.path.expanduser('~'), '.local/share/fonts')\nfont_file_dest_path = os.path.join(user_font_dir, font_file_name)\n\nfont_installed_and_setup = False\n\n# 2. 检查上传的字体文件是否存在\nif not os.path.exists(font_file_uploaded_path):\n    print(f\"错误: 未在指定路径找到上传的字体文件: {font_file_uploaded_path}\")\n    print(\"请确认:\")\n    print(f\"  1. 你已经将名为 '{font_file_name}' 的字体文件上传为一个 Kaggle Dataset。\")\n    print(f\"  2. 你已经将该 Dataset ('{font_dataset_name}') 添加到了这个 Notebook 的输入中。\")\n    print(f\"  3. 代码中的 `font_dataset_name` 已设置为 '{font_dataset_name}' (如果不是，请修改)。\")\nelse:\n    print(f\"找到上传的字体文件: {font_file_uploaded_path}\")\n    try:\n        # --- 将字体复制到系统可识别的位置 ---\n        os.makedirs(user_font_dir, exist_ok=True)\n        if not os.path.exists(font_file_dest_path):\n            shutil.copy(font_file_uploaded_path, font_file_dest_path)\n            print(f\"字体已复制到: {font_file_dest_path}\")\n        else:\n            print(f\"字体已存在于目标目录: {font_file_dest_path}\")\n\n        # --- 添加字体到 Matplotlib 管理器 ---\n        # 使用字体文件的完整路径来添加\n        font_entry = fm.FontEntry(fname=font_file_dest_path, name=os.path.splitext(font_file_name)[0]) # 使用文件名作为字体名\n        existing_font_paths = [f.fname for f in fm.fontManager.ttflist]\n        if font_file_dest_path not in existing_font_paths:\n            fm.fontManager.addfont(font_file_dest_path)\n            print(f\"字体 {font_file_dest_path} 已添加到 FontManager\")\n        else:\n             print(f\"字体 {font_file_dest_path} 已在 FontManager 的列表 Ttflist 中。\")\n\n        font_installed_and_setup = True # 标记字体文件已就位\n\n    except Exception as e_copy_add:\n        print(f\"复制或添加字体时出错: {e_copy_add}\")\n\n\n# 3. 如果字体文件在目标位置，则清理缓存并设置rcParams\nif font_installed_and_setup: # 仅当字体文件已复制/存在于目标位置时执行\n    try:\n        # --- 清理缓存 (使用正确的函数 matplotlib.get_cachedir()) ---\n        try:\n            cache_dir = matplotlib.get_cachedir() # <--- 使用 matplotlib.get_cachedir()\n            cache_cleaned = False\n            if os.path.exists(cache_dir):\n                print(f\"尝试清理 Matplotlib 字体缓存目录: {cache_dir}\")\n                for file in os.listdir(cache_dir):\n                    # 匹配更通用的缓存文件名模式\n                    if file.startswith('fontlist') and file.endswith(('.json', '.cache', '.afm', '.pickle')):\n                        try:\n                            os.remove(os.path.join(cache_dir, file))\n                            print(f\"  已删除缓存文件: {file}\")\n                            cache_cleaned = True\n                        except Exception as e_rm:\n                            print(f\"  删除缓存文件 {file} 失败: {e_rm}\")\n                if cache_cleaned:\n                     print(\"Matplotlib 字体缓存文件已清理。可能需要重启 Kernel 使其完全生效。\")\n                else:\n                     print(\"  未找到需要清理的字体缓存文件。\")\n            else:\n                 print(\"未找到 Matplotlib 缓存目录。\")\n        except AttributeError:\n             # 如果连 matplotlib.get_cachedir() 都没有 (极旧版本?)，则跳过\n             print(\"警告: 无法使用 matplotlib.get_cachedir()。跳过缓存清理。\")\n        except Exception as e_cache:\n             print(f\"清理字体缓存时出错: {e_cache}\")\n\n\n        # --- 设置 Matplotlib 参数 ---\n        # 尝试从字体文件获取标准字体名 (如 'SimHei')\n        try:\n            font_prop = fm.FontProperties(fname=font_file_dest_path)\n            font_name = font_prop.get_name()\n            print(f\"从文件推断出的字体名称: {font_name}\")\n        except Exception:\n            # 如果失败，则使用文件名（不含扩展名）作为备选\n            font_name = os.path.splitext(font_file_name)[0]\n            print(f\"无法从文件获取字体名，使用文件名作为名称: {font_name}\")\n\n        # 设置 matplotlib 默认字体\n        plt.rcParams['font.family'] = 'sans-serif' # 设置通用族\n        plt.rcParams['font.sans-serif'] = [font_name, 'sans-serif'] # 将你的字体名加入列表首位\n        plt.rcParams['axes.unicode_minus'] = False # 正确显示负号\n        print(f\"Matplotlib RCParams 已设置为优先使用 '{font_name}' 显示中文。\")\n\n        # 验证一下字体是否被正确识别（可选）\n        # if font_name in fm.findSystemFonts(fontpaths=[user_font_dir]):\n        #      print(f\"验证：字体 '{font_name}' 在管理器中找到。\")\n        # else:\n        #      print(f\"警告：字体 '{font_name}' 可能未被管理器完全识别，如果绘图仍有问题请重启Kernel。\")\n\n\n    except Exception as e_setup:\n        print(f\"设置 Matplotlib 字体参数或清理缓存时出错: {e_setup}\")\n        font_installed_and_setup = False # 标记设置失败\nelse:\n     # 如果前面步骤失败，提醒用户\n     if not os.path.exists(font_file_uploaded_path):\n        print(\"错误：未找到上传的字体文件，无法继续设置。\")\n     else:\n        print(\"字体复制或添加到管理器时失败，中文可能无法正确显示。\")\n\n\n# --- 结束字体设置 ---\nprint(\"-\" * 30)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:41.714488Z","iopub.execute_input":"2025-04-28T03:48:41.714873Z","iopub.status.idle":"2025-04-28T03:48:41.73008Z","shell.execute_reply.started":"2025-04-28T03:48:41.714843Z","shell.execute_reply":"2025-04-28T03:48:41.729487Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 导入库&环境设置","metadata":{}},{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n# === 导入必要的库 ===\nimport gc\nimport os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\n# import seaborn as sns # 根据需要取消注释\nfrom tqdm.notebook import tqdm\nimport cv2\nimport tensorflow as tf\nfrom tensorflow.keras import layers, models, optimizers, callbacks\nfrom tensorflow.keras.applications import EfficientNetB0\nfrom sklearn.model_selection import train_test_split # 如果需要K折交叉验证，可能需要 KFold 或 GroupKFold\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport random\nimport glob\nimport nibabel as nib\nimport math\nfrom concurrent.futures import ThreadPoolExecutor, ProcessPoolExecutor\n# import matplotlib as mpl # 根据需要取消注释\n\n# === 配置 ===\n# --- 数据路径 ---\nDATA_DIR = '/kaggle/input/rsna-2023-abdominal-trauma-detection'\nTRAIN_IMAGES_DIR = os.path.join(DATA_DIR, 'train_images')\nSEGMENTATION_DIR = os.path.join(DATA_DIR, 'segmentations') # 确认这个路径正确\nPREPROCESSED_DIR = '/kaggle/input/rsna-uneted/preprocessed_data' # 新增：存储预处理数据的目录\n\n# --- 输出路径 ---\nOUTPUT_DIR = '/kaggle/working/'\nMODEL_OUTPUT_DIR = os.path.join(OUTPUT_DIR, 'unet_model_v2')\nPREDICTION_OUTPUT_DIR = os.path.join(OUTPUT_DIR, 'segmentation_predictions_multi_v2')\n\nos.makedirs(MODEL_OUTPUT_DIR, exist_ok=True)\nos.makedirs(PREDICTION_OUTPUT_DIR, exist_ok=True)\nos.makedirs(PREPROCESSED_DIR, exist_ok=True) # 创建预处理数据目录\n\n# --- 图像和模型参数 ---\nIMG_SIZE = (224, 224)\nTARGET_SIZE = IMG_SIZE[0]\nN_INPUT_CHANNELS = 3\n\n# 定义器官映射 (修正后，合并左右肾)\n# NII标签: 1:肝脏, 2:脾脏, 3:左肾, 4:右肾, 5:肠道\nORGAN_MAP_NII = {\n    1: 'liver',\n    2: 'spleen',\n    3: 'kidney',  # 合并标签 3 和 4\n    5: 'bowel'\n}\n# 输出通道映射 (模型输出顺序)\nORGAN_CHANNEL_MAP = {\n    'liver': 0,\n    'spleen': 1,\n    'kidney': 2,\n    'bowel': 3\n}\nNUM_ORGANS = len(ORGAN_CHANNEL_MAP) # 现在是 4\n\nprint(f\"分割目标数量: {NUM_ORGANS}\")\nprint(f\"器官 -> NII值 映射 (处理方式): {ORGAN_MAP_NII}\")\nprint(f\"器官 -> 模型输出通道 映射: {ORGAN_CHANNEL_MAP}\")\n\nOUTPUT_MODEL_FILENAME = f'unet_effb0_multi_organ_{TARGET_SIZE}px_v2.keras'\nMODEL_SAVE_PATH = os.path.join(MODEL_OUTPUT_DIR, OUTPUT_MODEL_FILENAME)\n\n# --- 训练参数 ---\nVALIDATION_SPLIT = 0.15 # 考虑使用 K-Fold 交叉验证以更好地复现论文\nRANDOM_STATE = 42\nBATCH_SIZE = 8  # 增大批量大小以提高训练速度\nEPOCHS_STAGE1 = 20 # 可以根据需要调整各阶段Epochs\nEPOCHS_STAGE2 = 15\nEPOCHS_STAGE3 = 15\nEPOCHS_STAGE4 = 20 # 最后阶段可以多训练一些\nLEARNING_RATE = 1e-4\nEARLY_STOPPING_PATIENCE = 10\nREDUCE_LR_PATIENCE = 4\nREDUCE_LR_FACTOR = 0.2\nMIN_LR = 1e-6\n\n# --- 推理参数 ---\nINFERENCE_BATCH_SIZE = 32  # 增大批量大小以提高推理速度\nPREDICTION_THRESHOLD = 0.5\n\n# --- 预处理参数 ---\nPREFETCH_BUFFER_SIZE = tf.data.AUTOTUNE\nPARALLEL_CALLS = tf.data.AUTOTUNE\nCACHE_DATASET = False\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:41.731206Z","iopub.execute_input":"2025-04-28T03:48:41.731426Z","iopub.status.idle":"2025-04-28T03:48:41.753706Z","shell.execute_reply.started":"2025-04-28T03:48:41.731409Z","shell.execute_reply":"2025-04-28T03:48:41.753129Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 随机种子&辅助函数","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# === 设置随机种子 ===\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    print(f\"随机种子设置为: {seed}\")\n\nset_seed(RANDOM_STATE)\n\n# === 辅助函数 ===\ndef load_dicom_slice(path):\n    \"\"\"加载单个DICOM切片，应用VOI LUT，归一化，并获取InstanceNumber\"\"\"\n    try:\n        dicom_file = pydicom.dcmread(path)\n        instance_number = int(dicom_file.InstanceNumber)\n        image = apply_voi_lut(dicom_file.pixel_array, dicom_file)\n\n        min_val = np.min(image)\n        max_val = np.max(image)\n        if max_val > min_val:\n            image = (image - min_val) / (max_val - min_val)\n        else:\n            image = np.zeros_like(image) # 处理全黑或全白图像\n\n        image = image.astype(np.float32)\n\n        # 处理 MONOCHROME1 (图像像素值需要反转)\n        if 'PhotometricInterpretation' in dicom_file and dicom_file.PhotometricInterpretation == \"MONOCHROME1\":\n             image = 1.0 - image\n\n        return image, instance_number\n    except Exception as e:\n        # print(f\"加载 DICOM 错误 {path}: {e}\") # 可以取消注释以调试\n        return None, None\n\ndef load_multi_organ_segmentation_mask(nii_data_array, slice_index, organ_map_nii, organ_channel_map, num_organs, target_size, nii_path_for_error_msg=\"\"):\n    \"\"\"\n    从已加载的 NII 数据数组中提取特定切片的分割掩码, 创建多通道掩码。\n    根据提供的映射处理标签 (合并左右肾到'kidney', 标签5到'bowel')。\n    返回多通道二值掩码(0或1), 形状为(target_size, target_size, num_organs)\n    \"\"\"\n    try:\n        # 使用传入的已加载数据:\n        seg_data = nii_data_array # <--- 使用传入的 3D numpy 数组\n\n        if not isinstance(seg_data, np.ndarray) or seg_data.ndim != 3: # 检查传入的是否为有效的3D数组\n             print(f\"警告: 传入的NII数据不是有效的3D NumPy数组: {nii_path_for_error_msg}, 形状或类型: {seg_data.shape if isinstance(seg_data, np.ndarray) else type(seg_data)}\")\n             return None\n\n        num_slices = seg_data.shape[2]\n        # 检查slice_index是否有效\n        if not (0 <= slice_index < num_slices):\n            # print(f\"警告: 切片索引 {slice_index} 超出范围 [0, {num_slices}) for {nii_path_for_error_msg}\")\n            return None # 或者返回全零掩码\n\n        # 提取目标切片 (注意NIfTI的轴顺序可能与期望不同, 可能需要transpose)\n        # 假设最后一维是切片维\n        mask_slice_float = seg_data[:, :, slice_index]\n        # 为了安全比较，转换为整数\n        mask_slice = np.round(mask_slice_float).astype(np.int16) # <--- 从float转为int进行比较\n\n        # 创建多通道掩码\n        multi_channel_mask = np.zeros((mask_slice.shape[0], mask_slice.shape[1], num_organs), dtype=np.float32)\n\n        # --- 基于NII值填充通道 (这部分逻辑不变) ---\n        for nii_value, organ_name in organ_map_nii.items():\n            if organ_name in organ_channel_map:\n                channel_idx = organ_channel_map[organ_name]\n                if organ_name == 'kidney':\n                    # 合并左右肾 (标签 3 和 4)\n                    binary_mask_organ = ((mask_slice == 3) | (mask_slice == 4)).astype(np.float32)\n                else:\n                    # 处理其他器官 (肝脏 1, 脾脏 2, 肠道 5)\n                    binary_mask_organ = (mask_slice == nii_value).astype(np.float32)\n\n                # 使用加法以防万一映射重叠（虽然在此配置中不应发生）\n                multi_channel_mask[:, :, channel_idx] += binary_mask_organ\n        \n        # 调整大小\n        if multi_channel_mask.shape[0] != target_size or multi_channel_mask.shape[1] != target_size:\n            # 使用 INTER_NEAREST 来保持标签的离散性\n            resized_mask = cv2.resize(\n                multi_channel_mask,\n                (target_size, target_size),\n                interpolation=cv2.INTER_NEAREST\n            )\n            # cv2.resize可能压缩单通道输出，需要重新扩展维度\n            if len(resized_mask.shape) == 2 and num_organs == 1:\n                 resized_mask = np.expand_dims(resized_mask, axis=-1)\n            elif len(resized_mask.shape) == 2 and num_organs > 1:\n                print(f\"警告: 调整大小后多通道掩码被压缩为2D: {resized_mask.shape}, 目标通道: {num_organs} from {nii_path_for_error_msg}\")\n                # 尝试基于第一个通道恢复（可能不准确）或返回None\n                return None # 更安全的选择\n            resized_mask = (resized_mask > 0.5).astype(np.float32) # 二值化确保是0或1\n        else:\n             resized_mask = (multi_channel_mask > 0.5).astype(np.float32) # 确保原始尺寸也是二值化\n\n        # 验证最终形状\n        if resized_mask.shape != (target_size, target_size, num_organs):\n             print(f\"警告: 最终掩码形状不正确: {resized_mask.shape}，预期: {(target_size, target_size, num_organs)} from {nii_path_for_error_msg}\")\n             return None\n\n        return resized_mask\n\n    except Exception as e:\n        print(f\"处理分割掩码错误 (来自预加载数据) {nii_path_for_error_msg}, 切片 {slice_index}: {e}\")\n        return None\n\ndef preprocess_image_for_unet(image, target_size):\n    \"\"\"准备单个图像切片作为U-Net输入\"\"\"\n    # 调整图像大小\n    image_resized = cv2.resize(image, (target_size, target_size), interpolation=cv2.INTER_LINEAR)\n    # 扩展到3个通道 (对于需要3通道输入的模型如EfficientNet)\n    image_rgb = np.stack([image_resized] * N_INPUT_CHANNELS, axis=-1)\n    return image_rgb.astype(np.float32)\n\ndef get_dicom_files_dict(patient_id, series_id, dicom_tags_df):\n    \"\"\"辅助函数：获取并排序某个序列的DICOM文件信息\"\"\"\n    sorted_dicom_info = []\n    patient_dir = os.path.join(TRAIN_IMAGES_DIR, str(patient_id))\n    series_folder = os.path.join(patient_dir, str(series_id))\n    dicom_files = glob.glob(os.path.join(series_folder, '*.dcm'))\n    if not dicom_files: return []\n\n    # --- 尝试使用 InstanceNumber 排序 ---\n    dicom_tuples = []\n    use_tags = False\n    # 检查 DICOM tags 是否包含必要信息\n    if dicom_tags_df is not None and all(col in dicom_tags_df.columns for col in ['PatientID', 'SeriesInstanceUID', 'InstanceNumber', 'SOPInstanceUID']):\n        try:\n            # 确保类型匹配\n            patient_id_str = str(patient_id)\n            \n            # 从tags_df获取此序列的信息\n            tags_subset = dicom_tags_df[\n                 (dicom_tags_df['PatientID'].astype(str) == patient_id_str) &\n                 (dicom_tags_df['series_id_extracted'] == str(series_id))\n            ][['InstanceNumber', 'SOPInstanceUID']].dropna()\n\n            if not tags_subset.empty:\n                sop_to_inst = dict(zip(tags_subset['SOPInstanceUID'], tags_subset['InstanceNumber'].astype(int)))\n                use_tags = True # 标记成功使用tags\n\n                for f_path in dicom_files:\n                    try:\n                        ds_sop = pydicom.dcmread(f_path, stop_before_pixels=True).SOPInstanceUID\n                        if ds_sop in sop_to_inst:\n                             dicom_tuples.append((sop_to_inst[ds_sop], f_path))\n                        else: # Fallback: read InstanceNumber directly from header\n                           ds_num = pydicom.dcmread(f_path, stop_before_pixels=True)\n                           dicom_tuples.append((int(ds_num.InstanceNumber), f_path))\n                    except: # 文件读取失败或其他异常\n                         pass # 跳过无法处理的文件\n\n        except Exception as e:\n            use_tags = False # 出错则回退\n\n    # --- 如果 Tags 排序失败或不可用，尝试直接从DICOM头读取 InstanceNumber ---\n    if not use_tags or not dicom_tuples:\n        dicom_tuples = []\n        for f_path in dicom_files:\n            try:\n                ds = pydicom.dcmread(f_path, stop_before_pixels=True)\n                dicom_tuples.append((int(ds.InstanceNumber), f_path))\n            except Exception:\n                pass # 跳过无法读取的文件\n\n    # --- 如果 DICOM 头读取也失败，则按文件名排序 ---\n    if not dicom_tuples:\n        try:\n           # 尝试按文件名中的数字排序\n           dicom_tuples = sorted([(int(os.path.splitext(os.path.basename(f))[0]), f) for f in dicom_files])\n        except ValueError:\n           # 如果文件名不是纯数字，则按字母顺序排序\n           dicom_tuples = sorted([(i, f) for i, f in enumerate(sorted(dicom_files))])\n\n    # 按 InstanceNumber (或其他排序键) 排序\n    dicom_tuples.sort(key=lambda x: x[0])\n    sorted_dicom_info = [(item[0], item[1]) for item in dicom_tuples] # 返回 (InstanceNumber, path)\n\n    return sorted_dicom_info","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:41.900532Z","iopub.execute_input":"2025-04-28T03:48:41.900742Z","iopub.status.idle":"2025-04-28T03:48:41.964338Z","shell.execute_reply.started":"2025-04-28T03:48:41.900725Z","shell.execute_reply":"2025-04-28T03:48:41.963682Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_preprocessed_samples(preprocessed_dir, num_patients=3, samples_per_patient=2):\n    \"\"\"\n    从预处理数据中可视化几个样本，检查图像和掩码是否对齐\n    \n    参数:\n        preprocessed_dir: 预处理数据目录\n        num_patients: 要检查的患者数量\n        samples_per_patient: 每个患者检查的样本数量\n    \"\"\"\n    print(f\"检查预处理数据的对齐情况...\")\n    \n    # 获取所有预处理过的患者ID\n    patient_dirs = [d for d in os.listdir(preprocessed_dir) \n                   if os.path.isdir(os.path.join(preprocessed_dir, d))]\n    \n    if not patient_dirs:\n        print(\"没有找到预处理数据目录\")\n        return\n    \n    # 随机选择几个患者\n    selected_patients = np.random.choice(patient_dirs, \n                                        min(num_patients, len(patient_dirs)), \n                                        replace=False)\n    \n    for patient_id in selected_patients:\n        patient_dir = os.path.join(preprocessed_dir, patient_id)\n        npz_files = glob.glob(os.path.join(patient_dir, \"*.npz\"))\n        \n        if not npz_files:\n            print(f\"患者 {patient_id} 没有预处理文件\")\n            continue\n        \n        print(f\"检查患者 {patient_id} 的预处理数据\")\n        \n        # 随机选择几个样本\n        selected_files = np.random.choice(npz_files, \n                                         min(samples_per_patient, len(npz_files)), \n                                         replace=False)\n        \n        for file_path in selected_files:\n            try:\n                # 加载NPZ文件\n                data = np.load(file_path)\n                image = data['image']\n                mask = data['mask']\n                \n                # 获取文件名作为切片标识\n                slice_id = os.path.splitext(os.path.basename(file_path))[0]\n                \n                # 创建彩色掩码叠加\n                colors = {\n                    'liver': [1.0, 0.0, 0.0],  # 红色\n                    'spleen': [0.0, 1.0, 0.0],  # 绿色\n                    'kidney': [0.0, 0.0, 1.0],  # 蓝色\n                    'bowel': [1.0, 1.0, 0.0]    # 黄色\n                }\n                \n                # 准备可视化\n                fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n                \n                # 显示原始图像\n                if image.shape[2] == 3:  # 如果是3通道图像\n                    axes[0].imshow(image)\n                else:  # 如果是单通道图像\n                    axes[0].imshow(image[:,:,0], cmap='gray')\n                axes[0].set_title(f\"原始图像 - 患者 {patient_id}, 切片 {slice_id}\")\n                axes[0].axis('off')\n                \n                # 显示多通道掩码 (各通道不同颜色)\n                overlay_mask = np.zeros((*mask.shape[0:2], 3))\n                for i, organ_name in enumerate(ORGAN_CHANNEL_MAP.keys()):\n                    if i < mask.shape[2]:  # 确保通道索引有效\n                        color = colors[organ_name]\n                        for c in range(3):  # 对RGB三个通道\n                            overlay_mask[:,:,c] += mask[:,:,i] * color[c]\n                \n                # 将掩码值限制在[0,1]范围内\n                overlay_mask = np.clip(overlay_mask, 0, 1)\n                \n                # 显示掩码\n                axes[1].imshow(overlay_mask)\n                axes[1].set_title(\"分割掩码\")\n                axes[1].axis('off')\n                \n                # 显示图像和掩码叠加\n                # 将单通道图像转为RGB\n                if image.shape[2] == 3:\n                    rgb_image = image\n                else:\n                    rgb_image = np.stack([image[:,:,0]] * 3, axis=-1)\n                \n                # 叠加图像\n                alpha = 0.5\n                blended = rgb_image * (1 - alpha) + overlay_mask * alpha\n                blended = np.clip(blended, 0, 1)\n                \n                axes[2].imshow(blended)\n                axes[2].set_title(\"图像+掩码叠加\")\n                axes[2].axis('off')\n                \n                # 添加图例\n                legend_elements = [plt.Rectangle((0, 0), 1, 1, fc=colors[organ], label=organ)\n                                  for organ in ORGAN_CHANNEL_MAP.keys()]\n                fig.legend(handles=legend_elements, loc='lower center', ncol=len(legend_elements))\n                \n                plt.tight_layout(rect=[0, 0.05, 1, 0.95])\n                plt.show()\n                \n            except Exception as e:\n                print(f\"可视化文件 {file_path} 时出错: {e}\")\n    \n    print(\"预处理数据对齐检查完成\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:41.965665Z","iopub.execute_input":"2025-04-28T03:48:41.965873Z","iopub.status.idle":"2025-04-28T03:48:41.984462Z","shell.execute_reply.started":"2025-04-28T03:48:41.965858Z","shell.execute_reply":"2025-04-28T03:48:41.983781Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 离线预处理函数","metadata":{}},{"cell_type":"code","source":"def preprocess_and_save_data(patient_id, image_paths, nii_path, output_dir, target_size):\n    \"\"\"\n    预处理并保存患者的图像和掩码数据\n    \n    参数:\n        patient_id: 患者ID\n        image_paths: 该患者的DICOM路径列表\n        nii_path: 该患者的NII文件路径\n        output_dir: 输出目录\n        target_size: 目标图像大小\n    \n    返回:\n        处理成功的切片数量\n    \"\"\"\n    patient_output_dir = os.path.join(output_dir, str(patient_id))\n    os.makedirs(patient_output_dir, exist_ok=True)\n    \n    # 加载NII数据\n    try:\n        nii_img = nib.load(nii_path)\n        nii_data = nii_img.get_fdata()\n    except Exception as e:\n        print(f\"无法加载患者 {patient_id} 的NII文件: {e}\")\n        return 0\n    \n    processed_count = 0\n    \n    # 处理每个切片\n    for slice_idx, dicom_path in enumerate(image_paths):\n        if slice_idx >= nii_data.shape[2]:  # 确保不超出NII数据的切片范围\n            continue\n            \n        # 加载DICOM图像\n        image, instance_number = load_dicom_slice(dicom_path)\n        if image is None:\n            continue\n            \n        # 获取掩码\n        mask = load_multi_organ_segmentation_mask(\n            nii_data, slice_idx, \n            ORGAN_MAP_NII, ORGAN_CHANNEL_MAP, NUM_ORGANS, \n            target_size, nii_path\n        )\n        if mask is None:\n            continue\n            \n        # 预处理图像\n        processed_image = preprocess_image_for_unet(image, target_size)\n        \n        # 保存处理后的数据\n        output_path = os.path.join(patient_output_dir, f\"{instance_number if instance_number else slice_idx}.npz\")\n        np.savez_compressed(\n            output_path,\n            image=processed_image,\n            mask=mask\n        )\n        \n        processed_count += 1\n    \n    return processed_count\n\ndef perform_offline_preprocessing(image_paths_dict, segmentation_map, output_dir, target_size):\n    \"\"\"\n    对所有患者数据进行离线预处理\n    \n    参数:\n        image_paths_dict: 患者ID到DICOM路径列表的映射\n        segmentation_map: 患者ID到NII路径的映射\n        output_dir: 输出目录\n        target_size: 目标图像大小\n    \n    返回:\n        预处理数据的患者ID列表\n    \"\"\"\n    print(\"开始离线预处理数据...\")\n    \n    # 获取所有需要处理的患者ID\n    patient_ids = sorted(list(set(image_paths_dict.keys()) & set(segmentation_map.keys())))\n    \n    if not patient_ids:\n        print(\"没有找到同时包含图像和分割数据的患者\")\n        return []\n    \n    print(f\"将对 {len(patient_ids)} 位患者的数据进行预处理\")\n    \n    # 使用多进程处理\n    total_processed = 0\n    preprocessed_patients = []\n    \n    with ProcessPoolExecutor(max_workers=os.cpu_count()) as executor:\n        future_to_patient = {\n            executor.submit(\n                preprocess_and_save_data,\n                patient_id,\n                image_paths_dict[patient_id],\n                segmentation_map[patient_id],\n                output_dir,\n                target_size\n            ): patient_id for patient_id in patient_ids\n        }\n        \n        for future in tqdm(concurrent.futures.as_completed(future_to_patient), total=len(patient_ids), desc=\"预处理患者数据\"):\n            patient_id = future_to_patient[future]\n            try:\n                processed_count = future.result()\n                if processed_count > 0:\n                    total_processed += processed_count\n                    preprocessed_patients.append(patient_id)\n                    print(f\"患者 {patient_id} 预处理完成: {processed_count} 个切片\")\n                else:\n                    print(f\"患者 {patient_id} 没有处理成功的切片\")\n            except Exception as e:\n                print(f\"处理患者 {patient_id} 时出错: {e}\")\n    \n    print(f\"预处理完成，共处理 {len(preprocessed_patients)}/{len(patient_ids)} 位患者的 {total_processed} 个切片\")\n    return preprocessed_patients\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:41.985244Z","iopub.execute_input":"2025-04-28T03:48:41.985492Z","iopub.status.idle":"2025-04-28T03:48:42.004805Z","shell.execute_reply.started":"2025-04-28T03:48:41.985476Z","shell.execute_reply":"2025-04-28T03:48:42.004101Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 内存高效版 - 使用生成器而不是预加载\ndef create_tf_dataset_from_preprocessed(patient_ids, preprocessed_dir, batch_size, augment=True, shuffle=True):\n    \"\"\"\n    从预处理数据创建tf.data.Dataset，使用生成器方式避免内存溢出\n    \n    参数:\n        patient_ids: 患者ID列表\n        preprocessed_dir: 预处理数据目录\n        batch_size: 批量大小\n        augment: 是否进行数据增强\n        shuffle: 是否打乱数据\n    \n    返回:\n        tf.data.Dataset对象\n    \"\"\"\n    # 收集所有预处理文件的路径\n    all_files = []\n    for patient_id in patient_ids:\n        patient_dir = os.path.join(preprocessed_dir, str(patient_id))\n        if not os.path.exists(patient_dir):\n            continue\n        \n        npz_files = glob.glob(os.path.join(patient_dir, \"*.npz\"))\n        all_files.extend(npz_files)\n    \n    if not all_files:\n        raise ValueError(f\"没有找到预处理数据文件，请先运行预处理\")\n    \n    print(f\"找到 {len(all_files)} 个预处理数据文件\")\n    \n    # 创建一个基于文件路径的数据集\n    paths_dataset = tf.data.Dataset.from_tensor_slices(all_files)\n    \n    if shuffle:\n        # 限制缓冲区大小，避免内存问题\n        buffer_size = min(len(all_files), 10000)\n        paths_dataset = paths_dataset.shuffle(buffer_size=buffer_size, reshuffle_each_iteration=True)\n    \n    # 定义加载函数\n    def load_npz_file(file_path):\n        \"\"\"加载单个NPZ文件\"\"\"\n        # 将张量转换为字符串\n        file_path_str = file_path.numpy().decode('utf-8')\n        \n        try:\n            data = np.load(file_path_str)\n            image = data['image'].astype(np.float32)\n            mask = data['mask'].astype(np.float32)\n            \n            # 确保形状正确\n            if image.shape != (TARGET_SIZE, TARGET_SIZE, N_INPUT_CHANNELS) or mask.shape != (TARGET_SIZE, TARGET_SIZE, NUM_ORGANS):\n                print(f\"警告: 文件 {file_path_str} 的形状不正确，图像: {image.shape}, 掩码: {mask.shape}\")\n                # 返回正确形状的零数组\n                return (\n                    np.zeros((TARGET_SIZE, TARGET_SIZE, N_INPUT_CHANNELS), dtype=np.float32),\n                    np.zeros((TARGET_SIZE, TARGET_SIZE, NUM_ORGANS), dtype=np.float32)\n                )\n                \n            return image, mask\n            \n        except Exception as e:\n            print(f\"加载文件 {file_path_str} 失败: {e}\")\n            # 返回零数组\n            return (\n                np.zeros((TARGET_SIZE, TARGET_SIZE, N_INPUT_CHANNELS), dtype=np.float32),\n                np.zeros((TARGET_SIZE, TARGET_SIZE, NUM_ORGANS), dtype=np.float32)\n            )\n    \n    # 使用py_function将路径映射到图像和掩码\n    def load_and_process(file_path):\n        image, mask = tf.py_function(\n            load_npz_file,\n            [file_path],\n            [tf.float32, tf.float32]\n        )\n        # 设置形状，避免形状推断问题\n        image.set_shape((TARGET_SIZE, TARGET_SIZE, N_INPUT_CHANNELS))\n        mask.set_shape((TARGET_SIZE, TARGET_SIZE, NUM_ORGANS))\n        return image, mask\n    \n    # 映射加载函数\n    dataset = paths_dataset.map(load_and_process, num_parallel_calls=PARALLEL_CALLS)\n    \n    # 过滤掉加载失败的文件（可选）\n    # dataset = dataset.filter(lambda img, mask: tf.reduce_sum(img) > 0)\n    \n    # 数据增强\n    def augment_data(image, mask):\n        # 随机水平翻转\n        if tf.random.uniform(()) > 0.5:\n            image = tf.image.flip_left_right(image)\n            mask = tf.image.flip_left_right(mask)\n        \n        # 随机亮度\n        if tf.random.uniform(()) > 0.5:\n            image = tf.image.random_brightness(image, max_delta=0.1)\n            # 确保值在[0,1]范围内\n            image = tf.clip_by_value(image, 0.0, 1.0)\n        \n        # 随机对比度\n        if tf.random.uniform(()) > 0.5:\n            image = tf.image.random_contrast(image, lower=0.9, upper=1.1)\n            image = tf.clip_by_value(image, 0.0, 1.0)\n        \n        return image, mask\n    \n    # 应用数据增强\n    if augment:\n        dataset = dataset.map(augment_data, num_parallel_calls=PARALLEL_CALLS)\n    \n    # 批处理和预取\n    dataset = dataset.batch(batch_size)\n    \n    # 对于大型数据集，最好不要缓存\n    if CACHE_DATASET:\n        print(\"警告: 对大型数据集启用缓存可能导致内存问题，考虑设置CACHE_DATASET=False\")\n        dataset = dataset.cache()\n    \n    return dataset.prefetch(PREFETCH_BUFFER_SIZE)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:42.005697Z","iopub.execute_input":"2025-04-28T03:48:42.006168Z","iopub.status.idle":"2025-04-28T03:48:42.026032Z","shell.execute_reply.started":"2025-04-28T03:48:42.006145Z","shell.execute_reply":"2025-04-28T03:48:42.025378Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 数据分析函数","metadata":{}},{"cell_type":"code","source":"# 数据分析函数 - 完整版\ndef verify_nii_labels(segmentation_map, expected_map):\n    \"\"\"验证NII文件中的标签值是否与预期的器官映射匹配\"\"\"\n    print(\"开始验证NII文件中的标签值...\")\n    # 定义原始NII文件中的预期非背景标签值\n    expected_nii_labels_set = {1, 2, 3, 4, 5} # 肝、脾、左肾、右肾、肠\n\n    found_labels = set()\n    label_counts = {}\n    label_pixels = {val: [] for val in range(6)} # 统计0-5\n\n    sample_patients = list(segmentation_map.keys())[:20]\n    print(f\"将抽样检查 {len(sample_patients)} 位患者的NII文件...\")\n\n    for patient_id in tqdm(sample_patients, desc=\"验证NII标签\"):\n        nii_path = segmentation_map.get(patient_id)\n        if not nii_path: continue\n        try:\n            nii_img = nib.load(nii_path)\n            # 使用默认的 float dtype 加载数据，避免类型错误\n            seg_data = nii_img.get_fdata() \n\n            unique_labels_in_file = np.unique(seg_data)\n            found_labels.update(unique_labels_in_file)\n\n            for label in unique_labels_in_file:\n                 # 转换为整数进行比较和字典键查找\n                 label_int = int(np.round(label)) # 四舍五入并转整数，处理可能的浮点误差\n                 if 0 <= label_int <= 5: # 只统计0-5的标签\n                    # 使用原始浮点标签值进行精确计数\n                    count = np.sum(np.round(seg_data) == label_int) # 使用四舍五入比较\n                    if count > 0:\n                       label_pixels[label_int].append(count)\n\n        except Exception as e:\n            print(f\"处理患者 {patient_id} 的NII文件时出错: {e}\")\n\n    print(\"\\n=== NII文件标签验证结果 ===\")\n    # 从找到的标签中提取整数标签\n    found_int_labels = {int(np.round(l)) for l in found_labels if l == np.round(l)}\n    print(f\"在抽样的NII文件中发现的所有整数标签值: {sorted(list(found_int_labels))}\")\n\n\n    found_numeric_labels = {int(l) for l in found_int_labels if l != 0}\n\n    missing_labels = expected_nii_labels_set - found_numeric_labels\n    extra_labels = found_numeric_labels - expected_nii_labels_set\n\n    if missing_labels:\n        print(f\"警告: 以下预期标签在抽样NII文件中未找到: {missing_labels}\")\n    else:\n        print(\"所有预期的器官标签 (1-5) 至少在部分抽样文件中存在。\")\n\n    if extra_labels:\n        print(f\"警告: NII文件中包含以下预期之外的整数标签: {extra_labels}\")\n    else:\n        print(\"未发现预期之外的整数标签。\")\n\n    # 打印像素统计\n    print(\"\\n各标签值的像素数量统计 (基于抽样文件):\")\n    name_map = {1: 'liver', 2: 'spleen', 3: 'kidney_left', 4: 'kidney_right', 5: 'bowel', 0: 'background'}\n    for label_int, counts in label_pixels.items():\n        if not counts: continue\n        organ_name = name_map.get(label_int, \"未知\")\n        avg_count = np.mean(counts)\n        min_count = np.min(counts)\n        max_count = np.max(counts)\n        print(f\"  标签 {label_int} ({organ_name}): 平均像素数 = {avg_count:.0f}, 最小值 = {min_count:.0f}, 最大值 = {max_count:.0f}, 样本数 = {len(counts)}\")\n\n    return found_labels # 返回原始找到的标签（可能包含浮点数）\n\ndef analyze_organ_distribution(segmentation_map, organ_map_nii_model):\n    \"\"\"分析各器官（按模型定义合并）在数据集中的分布情况\"\"\"\n    print(\"开始分析器官分布...\")\n\n    organ_slice_counts = {organ: 0 for organ in organ_map_nii_model.values()}\n    organ_pixel_counts = {organ: 0 for organ in organ_map_nii_model.values()}\n    total_slices = 0\n    total_patients = len(segmentation_map)\n    patients_with_organ = {organ: set() for organ in organ_map_nii_model.values()}\n\n    for patient_id, nii_path in tqdm(segmentation_map.items(), desc=\"分析器官分布\"):\n        try:\n            nii_img = nib.load(nii_path)\n            # 使用默认的 float dtype 加载数据\n            seg_data = nii_img.get_fdata() \n            # 四舍五入为整数以进行标签比较\n            seg_data_int = np.round(seg_data).astype(np.int16)\n\n            num_slices_in_scan = seg_data_int.shape[2]\n            total_slices += num_slices_in_scan\n\n            for slice_idx in range(num_slices_in_scan):\n                slice_data_int = seg_data_int[:, :, slice_idx]\n\n                # 检查每个模型定义的器官是否存在于切片中\n                for nii_value, organ_name in organ_map_nii_model.items():\n                    if organ_name == 'kidney': # 合并处理肾脏\n                        has_organ = np.any((slice_data_int == 3) | (slice_data_int == 4))\n                        pixel_count = np.sum((slice_data_int == 3) | (slice_data_int == 4))\n                    else: # 处理其他器官 (肝脏 1, 脾脏 2, 肠道 5)\n                        has_organ = np.any(slice_data_int == nii_value)\n                        pixel_count = np.sum(slice_data_int == nii_value)\n\n                    if has_organ:\n                        organ_slice_counts[organ_name] += 1\n                        organ_pixel_counts[organ_name] += pixel_count\n                        patients_with_organ[organ_name].add(patient_id)\n\n        except Exception as e:\n            print(f\"分析患者 {patient_id} 时出错: {e}\")\n\n    organ_patient_counts = {organ: len(pids) for organ, pids in patients_with_organ.items()}\n\n    print(\"\\n=== 器官分布分析结果 ===\")\n    print(f\"总患者数: {total_patients}\")\n    print(f\"总切片数 (所有NII文件): {total_slices}\")\n\n    print(\"\\n器官在患者中的分布:\")\n    for organ, count in organ_patient_counts.items():\n        percentage = count / total_patients * 100 if total_patients > 0 else 0\n        print(f\"  {organ}: {count}/{total_patients} 患者 ({percentage:.2f}%)\")\n\n    print(\"\\n器官在切片中的分布 (至少有一个像素):\")\n    for organ, count in organ_slice_counts.items():\n        percentage = count / total_slices * 100 if total_slices > 0 else 0\n        print(f\"  {organ}: {count}/{total_slices} 切片 ({percentage:.2f}%)\")\n\n    print(\"\\n器官总像素数量:\")\n    for organ, count in organ_pixel_counts.items():\n        avg_per_slice_present = count / max(organ_slice_counts[organ], 1)\n        print(f\"  {organ}: 总像素数 = {count}, 平均每(含器官)切片像素数 = {avg_per_slice_present:.2f}\")\n\n    # === 计算类别权重 (基于切片频率倒数) ===\n    class_weights_slice_inv = {}\n    if total_slices > 0: # 使用总切片数作为分母计算频率\n        max_slice_count = max(organ_slice_counts.values()) if organ_slice_counts else 1.0\n        # 另一种方法：权重与频率成反比，再归一化\n        for organ, count in organ_slice_counts.items():\n             # 频率 = 该器官出现切片数 / 总切片数\n             frequency = (count + 1e-6) / total_slices\n             # 权重与频率成反比，用最大计数的倒数比例\n             weight = max_slice_count / (count + 1e-6)\n             class_weights_slice_inv[organ] = weight\n\n        # 归一化 (例如，使最小权重为1)\n        min_weight = min(class_weights_slice_inv.values()) if class_weights_slice_inv else 1.0\n        if min_weight > 0:\n             for organ in class_weights_slice_inv:\n                  class_weights_slice_inv[organ] /= min_weight\n        else: # 如果有器官从未出现，权重可能无限大，需要处理\n             max_finite_weight = max([w for w in class_weights_slice_inv.values() if np.isfinite(w)], default=1.0)\n             for organ in class_weights_slice_inv:\n                  if not np.isfinite(class_weights_slice_inv[organ]):\n                       class_weights_slice_inv[organ] = max_finite_weight * 2 # 给一个较大的有限值\n             min_weight = min(class_weights_slice_inv.values())\n             if min_weight > 0:\n                for organ in class_weights_slice_inv:\n                     class_weights_slice_inv[organ] /= min_weight\n    else: # 如果没有有效的切片\n        class_weights_slice_inv = {organ: 1.0 for organ in organ_map_nii_model.values()}\n\n    print(\"\\n建议的类别权重 (基于切片频率倒数，归一化):\")\n    for organ, weight in class_weights_slice_inv.items():\n        print(f\"  {organ}: {weight:.2f}\")\n\n    return {\n        'organ_patient_counts': organ_patient_counts,\n        'organ_slice_counts': organ_slice_counts,\n        'organ_pixel_counts': organ_pixel_counts,\n        'class_weights': class_weights_slice_inv\n    }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:42.086001Z","iopub.execute_input":"2025-04-28T03:48:42.086183Z","iopub.status.idle":"2025-04-28T03:48:42.104119Z","shell.execute_reply.started":"2025-04-28T03:48:42.086169Z","shell.execute_reply":"2025-04-28T03:48:42.103556Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# Unet模型构建","metadata":{}},{"cell_type":"code","source":"# === U-Net模型定义 ===\ndef build_unet_multi_organ(input_shape, num_organs, dropout_rate=0.3):\n    \"\"\"构建带有EfficientNetB0编码器的2D U-Net模型用于多器官分割\"\"\"\n    print(f\"构建多器官U-Net (EfficientNetB0 编码器), 输入形状:{input_shape}, 输出通道数:{num_organs}\")\n\n    # 加载预训练的EfficientNetB0作为编码器\n    efficientnet = EfficientNetB0(include_top=False, weights='imagenet', input_shape=input_shape)\n\n    # 获取跳跃连接层 (使用名称获取更稳定)\n    try:\n        # 这些层名通常在EfficientNet B0-B7中比较稳定\n        s1 = efficientnet.get_layer('block2a_expand_activation').output # 112x112\n        s2 = efficientnet.get_layer('block3a_expand_activation').output # 56x56\n        s3 = efficientnet.get_layer('block4a_expand_activation').output # 28x28\n        s4 = efficientnet.get_layer('block6a_expand_activation').output # 14x14\n        b0 = efficientnet.output # 7x7 (瓶颈)\n        skip_connections = [s1, s2, s3, s4]\n        print(\"成功获取EfficientNetB0中间层作为跳跃连接。\")\n    except ValueError as e:\n        print(f\"错误：无法获取指定的EfficientNetB0层。请检查层名称: {e}\")\n        print(\"建议使用 model.summary() 检查实际层名并更新。\")\n        # print(efficientnet.summary()) # 打印模型结构以帮助调试\n        raise e\n\n    # === 解码器 ===\n    # 定义上采样/解码器块的滤波器数量\n    decoder_filters = [256, 128, 64, 32] # 从瓶颈向上\n    x = b0\n\n    # 解码器路径与跳跃连接\n    for i in range(len(decoder_filters)):\n        filters = decoder_filters[i]\n        # 获取对应的跳跃连接 (从深到浅)\n        skip = skip_connections[len(skip_connections) - 1 - i]\n\n        # 上采样 (Conv2DTranspose)\n        x = layers.Conv2DTranspose(filters, (2, 2), strides=2, padding='same')(x)\n\n        # 检查并调整尺寸以匹配跳跃连接 (如果需要)\n        # if x.shape[1:3] != skip.shape[1:3]:\n        #     print(f\"尺寸不匹配: 上采样后 {x.shape[1:3]}, 跳跃连接 {skip.shape[1:3]}. 调整解码器输出大小。\")\n        #     x = tf.image.resize(x, skip.shape[1:3], method='bilinear')\n        #   或者调整跳跃连接 (有时更简单):\n        #   skip = layers.Conv2D(filters, 1, padding='same', activation='relu')(skip) # 用1x1卷积调整通道数\n        #   skip = tf.image.resize(skip, x.shape[1:3], method='bilinear')\n\n        # 连接跳跃特征\n        x = layers.concatenate([x, skip], axis=-1)\n\n        # 两个卷积层 + ReLU + Dropout\n        x = layers.Conv2D(filters, 3, padding='same', activation='relu')(x)\n        x = layers.BatchNormalization()(x) # 添加BN层有助于稳定训练\n        x = layers.Dropout(dropout_rate)(x)\n        x = layers.Conv2D(filters, 3, padding='same', activation='relu')(x)\n        x = layers.BatchNormalization()(x)\n\n    # 最终上采样到原始输入大小 (224x224)\n    # 当前 x 的大小应为 112x112 (经过4次上采样)\n    # 再进行一次上采样\n    x = layers.Conv2DTranspose(16, (2, 2), strides=2, padding='same', activation='relu')(x) # 输出 224x224x16\n    x = layers.Conv2D(16, 3, padding='same', activation='relu')(x)\n    x = layers.BatchNormalization()(x)\n\n    # 输出层: 1x1 卷积，通道数为器官数，激活函数为 sigmoid (用于多标签分割)\n    outputs = layers.Conv2D(num_organs, 1, activation='sigmoid', name='multi_organ_mask')(x)\n\n    # 创建模型\n    model = models.Model(inputs=efficientnet.input, outputs=outputs, name=f\"U-Net_EffB0_{num_organs}Organ\")\n\n    return model","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:42.105281Z","iopub.execute_input":"2025-04-28T03:48:42.105541Z","iopub.status.idle":"2025-04-28T03:48:42.124219Z","shell.execute_reply.started":"2025-04-28T03:48:42.105526Z","shell.execute_reply":"2025-04-28T03:48:42.12369Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 损失函数&评价指标","metadata":{}},{"cell_type":"code","source":"# 损失函数与评估指标\nSMOOTH = 1e-6\n\n@tf.function\ndef dice_coefficient(y_true, y_pred):\n    \"\"\"计算单个通道/类别的Dice系数\"\"\"\n    y_true_f = tf.keras.backend.flatten(y_true)\n    y_pred_f = tf.keras.backend.flatten(y_pred)\n    intersection = tf.keras.backend.sum(y_true_f * y_pred_f)\n    dice = (2. * intersection + SMOOTH) / (tf.keras.backend.sum(y_true_f) + tf.keras.backend.sum(y_pred_f) + SMOOTH)\n    return dice\n\n@tf.function\ndef average_dice_coefficient(y_true, y_pred):\n    \"\"\"计算所有通道/器官的平均Dice系数\"\"\"\n    num_organs = NUM_ORGANS  # 使用全局常量\n    total_dice = tf.constant(0.0, dtype=tf.float32)\n    \n    for i in range(num_organs):\n        total_dice += dice_coefficient(y_true[..., i], y_pred[..., i])\n    \n    return total_dice / tf.cast(num_organs, tf.float32)\n\n@tf.function\ndef dice_loss(y_true, y_pred):\n    \"\"\"计算 1 - 平均Dice系数 作为损失\"\"\"\n    return 1.0 - average_dice_coefficient(y_true, y_pred)\n\ndef weighted_dice_loss(class_weights):\n    \"\"\"\n    创建带有类别权重的 Dice 损失函数。\n    Args:\n        class_weights: 一个tf.constant张量或列表，包含每个通道的权重。\n                       顺序应与 ORGAN_CHANNEL_MAP 中的通道索引一致。\n    Returns:\n        一个损失函数 (y_true, y_pred) -> loss_value\n    \"\"\"\n    class_weights = tf.constant(class_weights, dtype=tf.float32)\n    num_classes = NUM_ORGANS\n    total_weight = tf.reduce_sum(class_weights)\n    if total_weight <= 0:  # 防止除以零\n        total_weight = tf.cast(num_classes, tf.float32)\n\n    @tf.function\n    def loss(y_true, y_pred):\n        total_weighted_dice = tf.constant(0.0, dtype=tf.float32)\n        for i in range(num_classes):\n            dice = dice_coefficient(y_true[..., i], y_pred[..., i])\n            total_weighted_dice += class_weights[i] * dice\n        # 返回加权平均 Dice 的补数作为损失\n        weighted_avg_dice = total_weighted_dice / total_weight\n        return 1.0 - weighted_avg_dice\n    return loss\n\n# --- 各器官的Dice系数指标 (用于评估) ---\n@tf.function\ndef dice_liver(y_true, y_pred):\n    channel_idx = ORGAN_CHANNEL_MAP['liver']\n    return dice_coefficient(y_true[..., channel_idx], y_pred[..., channel_idx])\n\n@tf.function\ndef dice_spleen(y_true, y_pred):\n    channel_idx = ORGAN_CHANNEL_MAP['spleen']\n    return dice_coefficient(y_true[..., channel_idx], y_pred[..., channel_idx])\n\n@tf.function\ndef dice_kidney(y_true, y_pred):\n    channel_idx = ORGAN_CHANNEL_MAP['kidney']\n    return dice_coefficient(y_true[..., channel_idx], y_pred[..., channel_idx])\n\n@tf.function\ndef dice_bowel(y_true, y_pred):\n    channel_idx = ORGAN_CHANNEL_MAP['bowel']\n    return dice_coefficient(y_true[..., channel_idx], y_pred[..., channel_idx])\n\n# 用于训练的评估指标列表\nMETRICS = [\n    average_dice_coefficient,\n    dice_liver,\n    dice_spleen,\n    dice_kidney,\n    dice_bowel\n]\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:42.124827Z","iopub.execute_input":"2025-04-28T03:48:42.124984Z","iopub.status.idle":"2025-04-28T03:48:42.141892Z","shell.execute_reply.started":"2025-04-28T03:48:42.124972Z","shell.execute_reply":"2025-04-28T03:48:42.141219Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 分阶段训练函数","metadata":{}},{"cell_type":"code","source":"# 分阶段训练函数 - 完整优化版(续)\ndef train_in_stages(train_pids, val_pids, preprocessed_dir, target_size,\n                    input_channels, num_organs, batch_size, \n                    learning_rate, epochs_per_stage, class_weights_map):\n    \"\"\"分阶段训练模型 (使用预处理数据)\"\"\"\n    input_shape = (target_size, target_size, input_channels)\n\n    # === 构建模型 ===\n    unet_model = build_unet_multi_organ(input_shape, num_organs, dropout_rate=0.3)\n\n    # 提取数据分析得出的类别权重，并映射到通道索引顺序\n    # 确保权重顺序与 ORGAN_CHANNEL_MAP 对应: liver, spleen, kidney, bowel\n    data_derived_weights = [class_weights_map.get(organ, 1.0) for organ in ORGAN_CHANNEL_MAP.keys()]\n    print(f\"数据驱动的类别权重 (顺序: {list(ORGAN_CHANNEL_MAP.keys())}): {data_derived_weights}\")\n\n    # --- 定义各阶段参数 ---\n    stages = [\n        {'name': 'Stage1_Liver', 'focus': ['liver'], 'epochs': epochs_per_stage[0], 'lr_factor': 1.0,\n         'loss': lambda yt, yp: 1.0 - dice_liver(yt, yp), 'monitor': 'val_dice_liver'},\n        {'name': 'Stage2_LiverSpleen', 'focus': ['liver', 'spleen'], 'epochs': epochs_per_stage[1], 'lr_factor': 0.5,\n         'loss': weighted_dice_loss([data_derived_weights[0], data_derived_weights[1], 0.0, 0.0]), 'monitor': 'val_average_dice_coefficient'},\n        {'name': 'Stage3_LiverSpleenKidney', 'focus': ['liver', 'spleen', 'kidney'], 'epochs': epochs_per_stage[2], 'lr_factor': 0.2,\n         'loss': weighted_dice_loss([data_derived_weights[0], data_derived_weights[1], data_derived_weights[2], 0.0]), 'monitor': 'val_average_dice_coefficient'},\n        {'name': 'Stage4_AllOrgans', 'focus': None, 'epochs': epochs_per_stage[3], 'lr_factor': 0.1,\n         'loss': weighted_dice_loss(data_derived_weights), 'monitor': 'val_average_dice_coefficient'}\n    ]\n\n    stage_histories = {}\n    best_model_path_overall = MODEL_SAVE_PATH # 最终模型的保存路径\n\n    for stage_info in stages:\n        stage_name = stage_info['name']\n        focus_organs = stage_info['focus']\n        epochs = stage_info['epochs']\n        current_lr = learning_rate * stage_info['lr_factor']\n        loss_func = stage_info['loss']\n        monitor_metric = stage_info['monitor']\n        stage_model_save_path = os.path.join(MODEL_OUTPUT_DIR, f\"{stage_name}_model.keras\")\n\n        print(f\"\\n=== {stage_name} 训练 ===\")\n        print(f\"  学习率: {current_lr}\")\n        print(f\"  监控指标: {monitor_metric}\")\n\n        # 创建数据集\n        train_dataset = create_tf_dataset_from_preprocessed(\n            train_pids, preprocessed_dir, batch_size, augment=True, shuffle=True\n        )\n        val_dataset = create_tf_dataset_from_preprocessed(\n            val_pids, preprocessed_dir, batch_size, augment=False, shuffle=False\n        )\n\n        # 编译模型\n        optimizer = optimizers.Adam(learning_rate=current_lr)\n        unet_model.compile(optimizer=optimizer, loss=loss_func, metrics=METRICS)\n\n        # 设置回调\n        checkpoint_callback_stage = callbacks.ModelCheckpoint(\n            stage_model_save_path,\n            monitor=monitor_metric,\n            mode='max',\n            save_best_only=True,\n            save_weights_only=False,\n            verbose=1\n        )\n        \n        callbacks_list = [checkpoint_callback_stage]\n        if stage_name == stages[-1]['name']:\n             checkpoint_callback_final = callbacks.ModelCheckpoint(\n                 best_model_path_overall,\n                 monitor=monitor_metric,\n                 mode='max',\n                 save_best_only=True,\n                 save_weights_only=False,\n                 verbose=1\n             )\n             callbacks_list.append(checkpoint_callback_final)\n\n        early_stopping = callbacks.EarlyStopping(\n            monitor=monitor_metric,\n            mode='max',\n            patience=EARLY_STOPPING_PATIENCE,\n            verbose=1,\n            restore_best_weights=True\n        )\n        reduce_lr = callbacks.ReduceLROnPlateau(\n            monitor=monitor_metric,\n            mode='max',\n            factor=REDUCE_LR_FACTOR,\n            patience=REDUCE_LR_PATIENCE,\n            min_lr=MIN_LR,\n            verbose=1\n        )\n        callbacks_list.extend([early_stopping, reduce_lr])\n\n        # 添加TensorBoard回调\n        tensorboard_callback = callbacks.TensorBoard(\n            log_dir=os.path.join(OUTPUT_DIR, 'logs', stage_name),\n            histogram_freq=1,\n            update_freq='epoch'\n        )\n        callbacks_list.append(tensorboard_callback)\n\n        # 估算训练集和验证集大小\n        # 使用tf.data.experimental.cardinality获取数据集大小\n        train_size = tf.data.experimental.cardinality(train_dataset).numpy()\n        val_size = tf.data.experimental.cardinality(val_dataset).numpy()\n        \n        print(f\"  训练批次数: {train_size}, 验证批次数: {val_size}\")\n\n        # 训练当前阶段\n        history = unet_model.fit(\n            train_dataset,\n            validation_data=val_dataset,\n            epochs=epochs,\n            callbacks=callbacks_list,\n            verbose=1\n        )\n        stage_histories[stage_name] = history.history\n\n        # 加载本阶段保存的最佳模型权重\n        if os.path.exists(stage_model_save_path):\n             print(f\"加载阶段最佳模型: {stage_model_save_path}\")\n             custom_objects = {\n                 'dice_loss': dice_loss,\n                 'loss': loss_func,\n                 'average_dice_coefficient': average_dice_coefficient,\n                 'dice_liver': dice_liver, 'dice_spleen': dice_spleen,\n                 'dice_kidney': dice_kidney, 'dice_bowel': dice_bowel\n             }\n             try:\n                unet_model = models.load_model(stage_model_save_path, custom_objects=custom_objects, compile=(loss_func is not None))\n                if loss_func is None:\n                    unet_model.compile(optimizer=optimizer, loss=weighted_dice_loss(data_derived_weights), metrics=METRICS)\n             except Exception as e:\n                 print(f\"加载阶段模型失败: {e}。继续使用当前模型权重。\")\n                 if not getattr(unet_model, '_is_compiled', False):\n                       unet_model.compile(optimizer=optimizer, loss=loss_func, metrics=METRICS)\n\n        # 清理显存\n        tf.keras.backend.clear_session()\n        gc.collect()\n\n    # 返回所有阶段的历史记录\n    return stage_histories\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:42.142616Z","iopub.execute_input":"2025-04-28T03:48:42.142849Z","iopub.status.idle":"2025-04-28T03:48:42.166597Z","shell.execute_reply.started":"2025-04-28T03:48:42.142825Z","shell.execute_reply":"2025-04-28T03:48:42.165844Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 优化版推理函数 - 完整版\ndef predict_batch_with_model(model, batch_images):\n    \"\"\"使用模型预测一批图像\"\"\"\n    return model.predict(batch_images, verbose=0)\n\ndef predict_and_save_masks_optimized(model, patient_ids, image_paths_dict, output_dir, batch_size=32):\n    \"\"\"对多个患者的图像进行批量预测和保存 (优化版)\"\"\"\n    print(f\"开始对 {len(patient_ids)} 位患者进行分割预测...\")\n    \n    # 统计计数器\n    processed_slices = 0\n    failed_slices = 0\n    \n    # 批量处理变量\n    batch_images = []\n    batch_info = []  # 存储(patient_id, series_id, instance_number, dicom_path)\n    \n    # 对每个患者进行处理\n    for patient_id in tqdm(patient_ids, desc=\"处理患者\"):\n        patient_dir = os.path.join(TRAIN_IMAGES_DIR, str(patient_id))\n        if not os.path.isdir(patient_dir):\n            continue\n            \n        # 获取该患者的所有系列\n        series_folders = [d for d in os.listdir(patient_dir) if os.path.isdir(os.path.join(patient_dir, d))]\n        \n        for series_id in series_folders:\n            series_dir = os.path.join(patient_dir, series_id)\n            dicom_files = glob.glob(os.path.join(series_dir, \"*.dcm\"))\n            \n            if not dicom_files:\n                continue\n                \n            # 为每个DICOM文件加载和预处理图像\n            for dicom_path in dicom_files:\n                # 加载DICOM图像\n                image, instance_number = load_dicom_slice(dicom_path)\n                \n                if image is None:\n                    failed_slices += 1\n                    continue\n                    \n                # 如果没有获取到instance_number，使用文件名作为替代\n                if instance_number is None:\n                    instance_number = os.path.splitext(os.path.basename(dicom_path))[0]\n                    \n                # 预处理图像\n                processed_image = preprocess_image_for_unet(image, TARGET_SIZE)\n                \n                # 添加到批次\n                batch_images.append(processed_image)\n                batch_info.append((patient_id, series_id, instance_number, dicom_path))\n                \n                # 当批次达到指定大小或处理完所有图像时进行预测\n                if len(batch_images) >= batch_size:\n                    # 批量预测\n                    try:\n                        batch_predictions = predict_batch_with_model(model, np.array(batch_images))\n                        \n                        # 处理预测结果\n                        for i, (pid, sid, inst_num, _) in enumerate(batch_info):\n                            # 创建输出目录\n                            patient_series_dir = os.path.join(output_dir, str(pid), str(sid))\n                            os.makedirs(patient_series_dir, exist_ok=True)\n                            \n                            # 保存预测掩码\n                            output_path = os.path.join(patient_series_dir, f\"{inst_num}.npz\")\n                            pred_mask = (batch_predictions[i] > PREDICTION_THRESHOLD).astype(np.uint8)\n                            \n                            np.savez_compressed(output_path, mask=pred_mask)\n                            processed_slices += 1\n                            \n                    except Exception as e:\n                        print(f\"批量预测失败: {e}\")\n                        failed_slices += len(batch_images)\n                        \n                    # 清空批次\n                    batch_images = []\n                    batch_info = []\n    \n    # 处理剩余的批次\n    if batch_images:\n        try:\n            batch_predictions = predict_batch_with_model(model, np.array(batch_images))\n            \n            for i, (pid, sid, inst_num, _) in enumerate(batch_info):\n                patient_series_dir = os.path.join(output_dir, str(pid), str(sid))\n                os.makedirs(patient_series_dir, exist_ok=True)\n                \n                output_path = os.path.join(patient_series_dir, f\"{inst_num}.npz\")\n                pred_mask = (batch_predictions[i] > PREDICTION_THRESHOLD).astype(np.uint8)\n                \n                np.savez_compressed(output_path, mask=pred_mask)\n                processed_slices += 1\n                \n        except Exception as e:\n            print(f\"处理最后一批预测失败: {e}\")\n            failed_slices += len(batch_images)\n    \n    print(f\"预测完成。成功处理: {processed_slices} 切片，失败: {failed_slices} 切片\")\n    return processed_slices, failed_slices\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:42.168257Z","iopub.execute_input":"2025-04-28T03:48:42.168528Z","iopub.status.idle":"2025-04-28T03:48:42.19118Z","shell.execute_reply.started":"2025-04-28T03:48:42.168513Z","shell.execute_reply":"2025-04-28T03:48:42.190584Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 模型评估","metadata":{}},{"cell_type":"code","source":"# 评估函数\ndef process_patient_evaluation(patient_id, nii_path, dicom_paths, series_pred_dir, pred_files,\n                              organ_map_nii, organ_channel_map, target_size, prediction_threshold):\n    \"\"\"处理单个患者的评估，用于并行执行\"\"\"\n    local_dice_scores = {organ: [] for organ in organ_channel_map.keys()}\n    local_iou_scores = {organ: [] for organ in organ_channel_map.keys()}\n    local_processed = 0\n    \n    try:\n        # 加载整个NII文件\n        nii_img = nib.load(nii_path)\n        gt_data = nii_img.get_fdata(dtype=np.int16)\n        num_slices_nii = gt_data.shape[2]\n    except Exception as e:\n        print(f\"无法加载患者 {patient_id} 的真实掩码: {e}\")\n        return local_dice_scores, local_iou_scores, local_processed\n    \n    # 创建预测文件映射\n    instance_to_pred_path = {int(os.path.splitext(os.path.basename(f))[0]): f for f in pred_files if os.path.splitext(os.path.basename(f))[0].isdigit()}\n    # 对于非数字文件名，使用字符串作为键\n    for f in pred_files:\n        basename = os.path.splitext(os.path.basename(f))[0]\n        if not basename.isdigit():\n            instance_to_pred_path[basename] = f\n    \n    # 遍历DICOM文件对应的切片索引\n    for slice_idx, dicom_path in enumerate(dicom_paths):\n        try:\n            # 尝试从DICOM头读取InstanceNumber作为切片标识符\n            ds = pydicom.dcmread(dicom_path, stop_before_pixels=True)\n            instance_number = int(ds.InstanceNumber)\n        except Exception:\n            # 如果无法读取InstanceNumber，使用切片索引作为近似值\n            instance_number = slice_idx + 1 # DICOM InstanceNumber通常从1开始\n\n        # 检查NII和预测中是否存在对应的切片\n        if not (0 <= slice_idx < num_slices_nii): \n            continue\n            \n        pred_file_path = instance_to_pred_path.get(instance_number)\n        if not pred_file_path and str(instance_number) in instance_to_pred_path:\n            pred_file_path = instance_to_pred_path[str(instance_number)]\n        \n        if not pred_file_path: \n            continue\n\n        # 加载预测掩码\n        try:\n            pred_data = np.load(pred_file_path)\n            pred_mask_binary = pred_data['mask'] # 假设已经是二值掩码\n            if pred_mask_binary.shape[:2] != (target_size, target_size) or pred_mask_binary.shape[2] != len(organ_channel_map):\n               continue\n        except Exception:\n            continue\n\n        # 提取真实掩码切片\n        gt_slice = gt_data[:, :, slice_idx]\n\n        # 对每个器官通道进行评估\n        for organ_name, channel_idx in organ_channel_map.items():\n            # 准备真实掩码 (二值化, 调整大小)\n            if organ_name == 'kidney':\n                gt_organ_binary = ((gt_slice == 3) | (gt_slice == 4)).astype(np.float32)\n            elif organ_name == 'liver':\n                gt_organ_binary = (gt_slice == 1).astype(np.float32)\n            elif organ_name == 'spleen':\n                gt_organ_binary = (gt_slice == 2).astype(np.float32)\n            elif organ_name == 'bowel':\n                gt_organ_binary = (gt_slice == 5).astype(np.float32)\n            else:\n                continue # 跳过未定义的器官\n\n            # Resize GT mask to target size using nearest neighbor\n            if gt_organ_binary.shape != (target_size, target_size):\n                gt_organ_resized = cv2.resize(gt_organ_binary, (target_size, target_size), interpolation=cv2.INTER_NEAREST)\n            else:\n                gt_organ_resized = gt_organ_binary\n\n            # 获取预测掩码的对应通道\n            pred_organ_binary = pred_mask_binary[..., channel_idx]\n\n            # 计算 Dice 和 IoU (仅在真实掩码或预测掩码至少有一个像素时计算才有意义)\n            if np.sum(gt_organ_resized) > 0 or np.sum(pred_organ_binary) > 0:\n                dice = (2. * np.sum(gt_organ_resized * pred_organ_binary) + 1e-6) / (np.sum(gt_organ_resized) + np.sum(pred_organ_binary) + 1e-6)\n                intersection = np.sum(gt_organ_resized * pred_organ_binary)\n                union = np.sum(gt_organ_resized) + np.sum(pred_organ_binary) - intersection\n                iou = (intersection + 1e-6) / (union + 1e-6)\n                \n                local_dice_scores[organ_name].append(dice)\n                local_iou_scores[organ_name].append(iou)\n                \n        local_processed += 1\n    \n    return local_dice_scores, local_iou_scores, local_processed\n\ndef evaluate_segmentation_results(prediction_dir, segmentation_map, image_paths_dict,\n                                 organ_map_nii, organ_channel_map, num_organs, target_size,\n                                 prediction_threshold):\n    \"\"\"评估分割结果与真实掩码的匹配程度 (优化版)\"\"\"\n    print(\"开始评估分割结果...\")\n\n    # 初始化分数记录\n    dice_scores = {organ: [] for organ in organ_channel_map.keys()}\n    iou_scores = {organ: [] for organ in organ_channel_map.keys()}\n\n    # 获取有真实掩码的患者ID列表\n    patient_ids_with_gt = list(segmentation_map.keys())\n    print(f\"找到 {len(patient_ids_with_gt)} 个有真实掩码的患者用于评估\")\n\n    processed_slices = 0\n    \n    # 使用ThreadPoolExecutor并行处理多个患者\n    with ThreadPoolExecutor(max_workers=os.cpu_count()) as executor:\n        futures = []\n        \n        for patient_id in patient_ids_with_gt:\n            nii_path = segmentation_map.get(patient_id)\n            dicom_paths = image_paths_dict.get(patient_id)\n            if not nii_path or not dicom_paths: \n                continue\n\n            # 查找该患者的预测文件\n            pred_patient_dir = os.path.join(prediction_dir, patient_id)\n            if not os.path.exists(pred_patient_dir):\n                continue\n\n            # 获取该病人第一个有预测文件的series\n            series_dirs = [os.path.join(pred_patient_dir, d) for d in os.listdir(pred_patient_dir)\n                           if os.path.isdir(os.path.join(pred_patient_dir, d))]\n            if not series_dirs:\n                continue\n                \n            series_pred_dir = series_dirs[0] # 假设评估第一个找到的series\n            pred_files = glob.glob(os.path.join(series_pred_dir, \"*.npz\"))\n            if not pred_files:\n                continue\n                \n            # 提交任务到线程池\n            futures.append(executor.submit(\n                process_patient_evaluation, \n                patient_id, nii_path, dicom_paths, series_pred_dir, pred_files,\n                organ_map_nii, organ_channel_map, target_size, prediction_threshold\n            ))\n        \n        # 收集结果\n        for future in tqdm(concurrent.futures.as_completed(futures), total=len(futures), desc=\"评估患者\"):\n            try:\n                patient_dice_scores, patient_iou_scores, patient_processed = future.result()\n                \n                # 合并结果\n                for organ in organ_channel_map.keys():\n                    dice_scores[organ].extend(patient_dice_scores[organ])\n                    iou_scores[organ].extend(patient_iou_scores[organ])\n                    \n                processed_slices += patient_processed\n            except Exception as e:\n                print(f\"处理评估结果时出错: {e}\")\n\n    # --- 输出和绘制结果 ---\n    print(f\"\\n评估完成，共处理 {processed_slices} 个有效切片。\")\n    print(\"=== 分割评估结果 (平均 Dice 和 IoU) ===\")\n    all_dices = []\n    all_ious = []\n    for organ in organ_channel_map.keys():\n        mean_dice = np.mean(dice_scores[organ]) if dice_scores[organ] else 0\n        mean_iou = np.mean(iou_scores[organ]) if iou_scores[organ] else 0\n        print(f\"  {organ}: Dice={mean_dice:.4f}, IoU={mean_iou:.4f}, 样本数={len(dice_scores[organ])}\")\n        all_dices.extend(dice_scores[organ])\n        all_ious.extend(iou_scores[organ])\n\n    overall_mean_dice = np.mean(all_dices) if all_dices else 0\n    overall_mean_iou = np.mean(all_ious) if all_ious else 0\n    print(f\"\\n  总体平均: Dice={overall_mean_dice:.4f}, IoU={overall_mean_iou:.4f}, 总样本数={len(all_dices)}\")\n\n    # --- 绘制评估结果箱线图 ---\n    fig, ax = plt.subplots(1, 2, figsize=(14, 6))\n    labels = list(organ_channel_map.keys())\n\n    # Dice 分数箱线图\n    dice_data_for_plot = [dice_scores[organ] for organ in labels]\n    ax[0].boxplot(dice_data_for_plot, labels=labels, showfliers=False) # showfliers=False 隐藏异常值\n    ax[0].set_title('各器官 Dice 系数分布')\n    ax[0].set_ylabel('Dice 系数')\n    ax[0].grid(True)\n\n    # IoU 分数箱线图\n    iou_data_for_plot = [iou_scores[organ] for organ in labels]\n    ax[1].boxplot(iou_data_for_plot, labels=labels, showfliers=False)\n    ax[1].set_title('各器官 IoU 分数分布')\n    ax[1].set_ylabel('IoU 分数')\n    ax[1].grid(True)\n\n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, \"segmentation_evaluation_boxplot.png\"))\n    plt.show()\n\n    return dice_scores, iou_scores\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:42.19205Z","iopub.execute_input":"2025-04-28T03:48:42.192262Z","iopub.status.idle":"2025-04-28T03:48:42.213736Z","shell.execute_reply.started":"2025-04-28T03:48:42.192241Z","shell.execute_reply":"2025-04-28T03:48:42.213169Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 保存并预测掩码","metadata":{}},{"cell_type":"code","source":"# === 保存预测掩码 ===\ndef postprocess_and_save_masks_npz(pred_mask_prob, threshold, output_dir, patient_id, series_id, instance_number):\n    \"\"\"处理并保存预测的分割掩码为NPZ格式\"\"\"\n    try:\n        # 创建输出目录\n        patient_series_dir = os.path.join(output_dir, str(patient_id), str(series_id))\n        os.makedirs(patient_series_dir, exist_ok=True)\n\n        # 确定输出路径\n        output_path = os.path.join(patient_series_dir, f\"{instance_number}.npz\")\n\n        # 可以选择只保存二值掩码以节省空间，或同时保存概率和二值掩码\n        binary_mask = (pred_mask_prob > threshold).astype(np.uint8)\n\n        # 使用压缩格式保存\n        np.savez_compressed(\n            output_path,\n            # mask_prob=pred_mask_prob.astype(np.float16), # 可选：保存概率图 (使用float16节省空间)\n            mask=binary_mask # 必须：保存二值掩码\n        )\n        return True\n    except Exception as e:\n        print(f\"保存掩码失败 (P{patient_id} S{series_id} I{instance_number}): {e}\")\n        return False","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:42.214464Z","iopub.execute_input":"2025-04-28T03:48:42.21471Z","iopub.status.idle":"2025-04-28T03:48:42.233157Z","shell.execute_reply.started":"2025-04-28T03:48:42.214689Z","shell.execute_reply":"2025-04-28T03:48:42.232515Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 主函数&可视化","metadata":{}},{"cell_type":"code","source":"def main():\n    \"\"\"主程序流程 (优化版)\"\"\"\n    print(\"--- 1. 构建文件映射 ---\")\n    # 加载系列元数据\n    train_meta_path = os.path.join(DATA_DIR, 'train_series_meta.csv')\n    if not os.path.exists(train_meta_path):\n        raise FileNotFoundError(f\"找不到系列元数据文件: {train_meta_path}\")\n    train_series_meta = pd.read_csv(train_meta_path)\n\n    # 创建 series_id -> patient_id 映射\n    series_to_patient = dict(zip(\n        train_series_meta['series_id'].astype(str),\n        train_series_meta['patient_id'].astype(str)\n    ))\n\n    # 创建 NII 分割文件映射 (series_id -> nii_path)\n    series_to_nii = {}\n    segmentation_files = glob.glob(os.path.join(SEGMENTATION_DIR, \"*.nii\"))\n    print(f\"发现 {len(segmentation_files)} 个 NII 文件。\")\n    for fpath in segmentation_files:\n        series_id = os.path.splitext(os.path.basename(fpath))[0]\n        if series_id in series_to_patient: # 确保这个系列在元数据中\n            series_to_nii[series_id] = fpath\n\n    print(f\"成功映射 {len(series_to_nii)} 个 NII 文件到 series_id。\")\n\n    # 加载 DICOM tags (如果可用)\n    dicom_tags_df = None\n    dicom_tags_path = os.path.join(DATA_DIR, 'train_dicom_tags.parquet')\n    if os.path.exists(dicom_tags_path):\n        print(\"加载 DICOM tags...\")\n        try:\n            dicom_tags_df = pd.read_parquet(dicom_tags_path)\n            # 预处理 tags DataFrame\n            if 'PatientID' in dicom_tags_df.columns:\n                dicom_tags_df['PatientID'] = dicom_tags_df['PatientID'].astype(str)\n            # 尝试提取 series_id\n            if 'SeriesInstanceUID' in dicom_tags_df.columns:\n                 dicom_tags_df['series_id_extracted'] = dicom_tags_df['SeriesInstanceUID'].str.split('.').str[-2]\n                 dicom_tags_df = dicom_tags_df.dropna(subset=['series_id_extracted'])\n            print(\"DICOM tags 加载完成。\")\n        except Exception as e:\n            print(f\"加载 DICOM tags 失败: {e}. 将不使用 tags 进行排序。\")\n            dicom_tags_df = None\n    else:\n        print(\"未找到 DICOM tags 文件，将仅依赖DICOM头或文件名排序。\")\n\n    # --- 关联图像和分割 ---\n    image_paths_dict = {}  # {patient_id: [sorted_list_of_dicom_paths]}\n    segmentation_map = {}  # {patient_id: nii_path}\n    \n    patient_ids_in_train_images = os.listdir(TRAIN_IMAGES_DIR)\n    valid_patients = []  # 存储有图像和对应NII文件的患者ID\n\n    print(\"开始关联患者图像和分割文件...\")\n    for patient_id in tqdm(patient_ids_in_train_images, desc=\"处理患者\"):\n        if not os.path.isdir(os.path.join(TRAIN_IMAGES_DIR, patient_id)):\n            continue\n\n        patient_series_ids = [s for s, p in series_to_patient.items() if p == patient_id]\n        if not patient_series_ids: \n            continue\n\n        # 查找该患者是否有系列同时存在于图像目录和NII分割中\n        found_valid_series = False\n        for series_id in patient_series_ids:\n            series_img_path = os.path.join(TRAIN_IMAGES_DIR, patient_id, series_id)\n            nii_path = series_to_nii.get(series_id)\n\n            if os.path.isdir(series_img_path) and nii_path:\n                # 获取并排序DICOM文件\n                dicom_info = get_dicom_files_dict(patient_id, series_id, dicom_tags_df)\n                if dicom_info:  # 确保系列中有有效的DICOM文件\n                    image_paths_dict[patient_id] = [item[1] for item in dicom_info]  # 只存路径\n                    segmentation_map[patient_id] = nii_path\n                    valid_patients.append(patient_id)\n                    found_valid_series = True\n                    break  # 每个患者只使用一个有效的 series 和 NII\n\n    final_patient_ids = sorted(list(set(valid_patients)))  # 去重并排序\n    print(f\"成功映射了 {len(final_patient_ids)} 位患者的图像和分割文件。\")\n    if not final_patient_ids:\n        raise SystemExit(\"错误：未能找到任何包含有效图像序列及对应NII文件的患者。\")\n\n    print(\"\\n--- 2. 验证NII文件标签 ---\")\n    # 使用修正后的器官映射进行验证\n    verify_nii_labels(segmentation_map, {1: 'liver', 2: 'spleen', 3: 'kidney_left', 4: 'kidney_right', 5: 'bowel'})\n\n    print(\"\\n--- 3. 分析器官分布 ---\")\n    # 使用模型将要使用的器官映射进行分析 (合并肾脏)\n    distribution_results = analyze_organ_distribution(segmentation_map, ORGAN_MAP_NII)\n    calculated_class_weights = distribution_results['class_weights']  # 保存计算出的权重\n\n    print(\"\\n--- 4. 划分训练集和验证集 ---\")\n    train_pids, val_pids = train_test_split(final_patient_ids, test_size=VALIDATION_SPLIT, random_state=RANDOM_STATE)\n    print(f\"训练集患者数: {len(train_pids)}\")\n    print(f\"验证集患者数: {len(val_pids)}\")\n\n    print(\"\\n--- 5. 使用已有预处理数据 ---\")\n    # 检查预处理数据目录\n    preprocessed_patients = [d for d in os.listdir(PREPROCESSED_DIR) \n                          if os.path.isdir(os.path.join(PREPROCESSED_DIR, d))]\n    print(f\"找到 {len(preprocessed_patients)} 个已预处理的患者数据\")\n\n    # 添加这一行，在训练前检查预处理数据的对齐情况\n    print(\"\\n--- 5b. 检查预处理数据对齐 ---\")\n    visualize_preprocessed_samples(PREPROCESSED_DIR, num_patients=3, samples_per_patient=2)\n    \n    # 确认预处理的患者包含了训练和验证集\n    train_pids_preprocessed = [pid for pid in train_pids if pid in preprocessed_patients]\n    val_pids_preprocessed = [pid for pid in val_pids if pid in preprocessed_patients]\n\n    \n    print(f\"预处理后的训练集患者数: {len(train_pids_preprocessed)}/{len(train_pids)}\")\n    print(f\"预处理后的验证集患者数: {len(val_pids_preprocessed)}/{len(val_pids)}\")\n    \n    if len(train_pids_preprocessed) == 0 or len(val_pids_preprocessed) == 0:\n        raise SystemExit(\"错误：预处理后训练集或验证集为空。\")\n\n    # --- 准备训练参数 ---\n    epochs_config = [EPOCHS_STAGE1, EPOCHS_STAGE2, EPOCHS_STAGE3, EPOCHS_STAGE4]\n\n    print(\"\\n--- 6. 开始分阶段训练模型 ---\")\n    stage_histories = train_in_stages(\n        train_pids_preprocessed, val_pids_preprocessed, PREPROCESSED_DIR,\n        TARGET_SIZE, N_INPUT_CHANNELS, NUM_ORGANS,\n        BATCH_SIZE, LEARNING_RATE, epochs_config, calculated_class_weights\n    )\n\n    # --- 7. 绘制训练历史曲线 ---\n    print(\"\\n--- 7. 绘制训练历史曲线 ---\")\n    plt.figure(figsize=(18, 12))\n    num_stages = len(stage_histories)\n    colors = plt.cm.viridis(np.linspace(0, 1, num_stages))\n\n    # 绘制各阶段的平均Dice系数\n    plt.subplot(2, 2, 1)\n    for i, (stage_name, history) in enumerate(stage_histories.items()):\n        epochs = range(1, len(history['average_dice_coefficient']) + 1)\n        plt.plot(epochs, history['average_dice_coefficient'], label=f'{stage_name} Train Avg Dice', color=colors[i], linestyle='--')\n        if 'val_average_dice_coefficient' in history:\n             plt.plot(epochs, history['val_average_dice_coefficient'], label=f'{stage_name} Val Avg Dice', color=colors[i])\n    plt.title('平均Dice系数 (所有阶段)')\n    plt.xlabel('Epochs')\n    plt.ylabel('Dice系数')\n    plt.legend()\n    plt.grid(True)\n\n    # 绘制各阶段的损失\n    plt.subplot(2, 2, 2)\n    for i, (stage_name, history) in enumerate(stage_histories.items()):\n         epochs = range(1, len(history['loss']) + 1)\n         plt.plot(epochs, history['loss'], label=f'{stage_name} Train Loss', color=colors[i], linestyle='--')\n         if 'val_loss' in history:\n              plt.plot(epochs, history['val_loss'], label=f'{stage_name} Val Loss', color=colors[i])\n    plt.title('损失 (所有阶段)')\n    plt.xlabel('Epochs')\n    plt.ylabel('Loss')\n    plt.legend()\n    plt.grid(True)\n\n    # 绘制最终阶段的各器官验证Dice系数\n    plt.subplot(2, 2, 3)\n    final_stage_name = list(stage_histories.keys())[-1]\n    final_history = stage_histories[final_stage_name]\n    epochs = range(1, len(final_history['val_dice_liver']) + 1) # 假设所有指标长度相同\n    plt.plot(epochs, final_history['val_dice_liver'], label='肝脏 (Val)')\n    plt.plot(epochs, final_history['val_dice_spleen'], label='脾脏 (Val)')\n    plt.plot(epochs, final_history['val_dice_kidney'], label='肾脏 (Val)')\n    plt.plot(epochs, final_history['val_dice_bowel'], label='肠道 (Val)')\n    plt.title(f'各器官 Dice 系数 ({final_stage_name} - 验证集)')\n    plt.xlabel('Epochs')\n    plt.ylabel('Dice系数')\n    plt.legend()\n    plt.grid(True)\n\n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, \"training_history_all_stages_v2.png\"))\n    plt.show()\n\n    # --- 8. 加载最终最佳模型进行推理 ---\n    print(\"\\n--- 8. 加载最终最佳模型进行推理 ---\")\n    # 确保最终模型路径存在\n    if not os.path.exists(MODEL_SAVE_PATH):\n         print(f\"警告: 最终最佳模型 {MODEL_SAVE_PATH} 未找到。可能需要检查训练过程。跳过推理和评估。\")\n         return # 或者加载最后一个阶段的模型\n\n    custom_objects = {\n        'dice_loss': dice_loss, # 基础loss\n        'average_dice_coefficient': average_dice_coefficient,\n        'dice_liver': dice_liver, 'dice_spleen': dice_spleen,\n        'dice_kidney': dice_kidney, 'dice_bowel': dice_bowel\n    }\n    \n    try:\n         # 先尝试不编译加载\n         best_model = models.load_model(MODEL_SAVE_PATH, custom_objects=custom_objects, compile=False)\n         print(f\"成功加载模型: {MODEL_SAVE_PATH}\")\n    except Exception as e:\n         print(f\"加载最终模型失败: {e}. 尝试提供损失函数...\")\n         try:\n             final_loss_func = weighted_dice_loss(calculated_class_weights) # 假设这是最终阶段损失\n             custom_objects['loss'] = final_loss_func # 添加到自定义对象\n             best_model = models.load_model(MODEL_SAVE_PATH, custom_objects=custom_objects)\n             print(f\"成功加载并编译模型: {MODEL_SAVE_PATH}\")\n         except Exception as e2:\n              print(f\"再次加载模型失败: {e2}. 跳过推理和评估。\")\n              return\n\n    # --- 9. 进行推理 ---\n    print(\"\\n--- 9. 开始处理病人并生成分割掩码 (推理) ---\")\n    # 对所有在 train_images 中的患者进行推理\n    all_patient_ids_inference = sorted([pid for pid in os.listdir(TRAIN_IMAGES_DIR)\n                                       if os.path.isdir(os.path.join(TRAIN_IMAGES_DIR, pid))])\n    print(f\"将为 {len(all_patient_ids_inference)} 个患者生成分割掩码。\")\n\n    # 使用优化的批量推理函数\n    processed_slices_total_inference, failed_saves_inference = predict_and_save_masks_optimized(\n        best_model, all_patient_ids_inference, image_paths_dict, \n        PREDICTION_OUTPUT_DIR, INFERENCE_BATCH_SIZE\n    )\n\n    print(\"-\" * 30)\n    print(\"所有病人处理完毕 (推理)。\")\n    print(f\"总共成功生成的分割掩码文件数量: {processed_slices_total_inference}\")\n    print(f\"保存失败的掩码文件数量: {failed_saves_inference}\")\n    print(f\"分割掩码已保存在: {PREDICTION_OUTPUT_DIR}\")\n\n    # --- 10. 评估分割结果 ---\n    print(\"\\n--- 10. 评估分割结果 ---\")\n    # 使用验证集子集进行评估，因为只有这些有GT NII文件\n    evaluate_segmentation_results(\n        PREDICTION_OUTPUT_DIR,\n        segmentation_map, # 包含 patient_id -> nii_path 的映射\n        image_paths_dict, # 包含 patient_id -> [dicom_paths] 的映射\n        ORGAN_MAP_NII, ORGAN_CHANNEL_MAP, NUM_ORGANS, TARGET_SIZE,\n        PREDICTION_THRESHOLD\n    )\n    print(\"-\" * 30)\n    print(\"分割评估完成。\")\n\n    # --- 11. 可视化一些分割结果 ---\n    print(\"\\n--- 11. 可视化一些分割结果 ---\")\n    # 从验证集中随机选择几个患者进行可视化\n    if val_pids:\n        vis_patient_ids = np.random.choice(val_pids, min(5, len(val_pids)), replace=False)\n    else:\n        vis_patient_ids = np.random.choice(final_patient_ids, min(5, len(final_patient_ids)), replace=False)\n\n    for patient_id in vis_patient_ids:\n        if patient_id not in image_paths_dict or patient_id not in segmentation_map:\n            continue\n\n        dicom_paths = image_paths_dict[patient_id]\n        nii_path = segmentation_map[patient_id]\n\n        # 随机选择一个切片索引进行可视化\n        if not dicom_paths: continue\n        slice_vis_idx = np.random.randint(0, len(dicom_paths))\n        dicom_path_vis = dicom_paths[slice_vis_idx]\n\n        # 加载DICOM图像\n        image_vis, inst_num_vis = load_dicom_slice(dicom_path_vis)\n        if image_vis is None: continue\n\n        # 加载真实掩码\n        try:\n            nii_img_vis = nib.load(nii_path)\n            gt_data_vis = nii_img_vis.get_fdata(dtype=np.int16)\n            if 0 <= slice_vis_idx < gt_data_vis.shape[2]:\n                gt_slice_vis = gt_data_vis[:, :, slice_vis_idx]\n            else:\n                 print(f\"可视化时切片索引 {slice_vis_idx} 无效 for NII {nii_path}\")\n                 continue\n        except Exception as e:\n            print(f\"加载真实掩码进行可视化时出错: {e}\")\n            continue\n\n        # 查找对应的预测掩码\n        series_id_vis = dicom_path_vis.split('/')[-2] # 从路径提取series_id\n        instance_number_vis = inst_num_vis if inst_num_vis is not None else slice_vis_idx + 1\n\n        pred_vis_path = os.path.join(PREDICTION_OUTPUT_DIR, patient_id, series_id_vis, f\"{instance_number_vis}.npz\")\n\n        if not os.path.exists(pred_vis_path):\n            print(f\"未找到预测文件进行可视化: {pred_vis_path}\")\n            # 尝试模型实时预测\n            try:\n                 processed_image_vis = preprocess_image_for_unet(image_vis, TARGET_SIZE)\n                 prediction_prob = best_model.predict(np.expand_dims(processed_image_vis, axis=0), verbose=0)[0]\n                 pred_mask_binary = (prediction_prob > PREDICTION_THRESHOLD).astype(np.uint8)\n            except Exception as live_pred_e:\n                 print(f\"实时预测失败: {live_pred_e}\")\n                 continue\n        else:\n             # 加载保存的预测二值掩码\n             try:\n                 pred_data_vis = np.load(pred_vis_path)\n                 pred_mask_binary = pred_data_vis['mask'] # 假设二值掩码保存在'mask'\n             except Exception as e:\n                  print(f\"加载预测掩码进行可视化时出错: {e}\")\n                  continue\n\n        # --- 创建彩色叠加图 ---\n        colors = { # BGR 颜色\n            'liver': [0, 0, 255],  # 红色\n            'spleen': [0, 255, 0],  # 绿色\n            'kidney': [255, 0, 0],  # 蓝色\n            'bowel': [0, 255, 255]   # 黄色\n        }\n\n        # 调整原始图像用于显示 (灰度图转BGR)\n        display_image = cv2.cvtColor((image_vis * 255).astype(np.uint8), cv2.COLOR_GRAY2BGR)\n        display_image_resized = cv2.resize(display_image, (TARGET_SIZE, TARGET_SIZE))\n\n        # 创建真实掩码的彩色叠加版本\n        gt_overlay = np.zeros_like(display_image_resized, dtype=np.uint8)\n        for nii_val, org_name in ORGAN_MAP_NII.items():\n             if org_name in colors:\n                 if org_name == 'kidney':\n                     mask_ch = ((gt_slice_vis == 3) | (gt_slice_vis == 4)).astype(np.uint8)\n                 else:\n                     mask_ch = (gt_slice_vis == nii_val).astype(np.uint8)\n\n                 mask_ch_resized = cv2.resize(mask_ch, (TARGET_SIZE, TARGET_SIZE), interpolation=cv2.INTER_NEAREST)\n                 gt_overlay[mask_ch_resized > 0] = colors[org_name]\n\n        # 创建预测掩码的彩色叠加版本\n        pred_overlay = np.zeros_like(display_image_resized, dtype=np.uint8)\n        for organ_name, channel_idx in ORGAN_CHANNEL_MAP.items():\n             if organ_name in colors:\n                  mask_ch_pred = pred_mask_binary[:, :, channel_idx] # 已经是 (target_size, target_size)\n                  pred_overlay[mask_ch_pred > 0] = colors[organ_name]\n\n        # 混合图像和掩码\n        alpha = 0.4\n        gt_blended = cv2.addWeighted(display_image_resized, 1 - alpha, gt_overlay, alpha, 0)\n        pred_blended = cv2.addWeighted(display_image_resized, 1 - alpha, pred_overlay, alpha, 0)\n\n        # --- 显示结果 ---\n        fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n\n        axes[0].imshow(cv2.cvtColor(display_image_resized, cv2.COLOR_BGR2RGB))\n        axes[0].set_title(f\"患者 {patient_id} - 切片 {instance_number_vis}\")\n        axes[0].axis('off')\n\n        axes[1].imshow(cv2.cvtColor(gt_blended, cv2.COLOR_BGR2RGB))\n        axes[1].set_title(\"真实分割掩码 (叠加)\")\n        axes[1].axis('off')\n\n        axes[2].imshow(cv2.cvtColor(pred_blended, cv2.COLOR_BGR2RGB))\n        axes[2].set_title(\"预测分割掩码 (叠加)\")\n        axes[2].axis('off')\n\n        # 添加图例\n        legend_elements = [plt.Rectangle((0, 0), 1, 1, fc=[c/255. for c in colors[org][::-1]], label=org)\n                           for org in ORGAN_CHANNEL_MAP.keys() if org in colors]\n        fig.legend(handles=legend_elements, loc='lower center', ncol=len(legend_elements))\n\n        plt.tight_layout(rect=[0, 0.05, 1, 0.97]) # 调整布局以防图例重叠\n        plt.savefig(os.path.join(OUTPUT_DIR, f\"segmentation_vis_{patient_id}_{instance_number_vis}.png\"))\n        plt.show()\n\n    print(\"-\" * 30)\n    print(\"程序执行完成。\")\n    print(\"-\" * 30)\n\nif __name__ == \"__main__\":\n    main()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-28T03:48:42.270296Z","iopub.execute_input":"2025-04-28T03:48:42.270485Z","execution_failed":"2025-04-28T04:12:06.697Z"}},"outputs":[],"execution_count":null}]}