{"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":"none","dataSources":[{"sourceId":52254,"databundleVersionId":9674523,"sourceType":"competition"},{"sourceId":11555525,"sourceType":"datasetVersion","datasetId":7245708},{"sourceId":11779768,"sourceType":"datasetVersion","datasetId":7269783},{"sourceId":11846236,"sourceType":"datasetVersion","datasetId":7424084},{"sourceId":11864742,"sourceType":"datasetVersion","datasetId":7443333},{"sourceId":392249,"sourceType":"modelInstanceVersion","modelInstanceId":322963,"modelId":343654},{"sourceId":395733,"sourceType":"modelInstanceVersion","modelInstanceId":324962,"modelId":345797}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"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-05-17T15:25:54.922278Z","iopub.execute_input":"2025-05-17T15:25:54.923142Z","iopub.status.idle":"2025-05-17T15:25:54.941622Z","shell.execute_reply.started":"2025-05-17T15:25:54.923106Z","shell.execute_reply":"2025-05-17T15:25:54.940867Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 导入库&环境设置","metadata":{}},{"cell_type":"code","source":"# -*- coding: utf-8 -*-\n# === 导入必要的库 ===\nimport concurrent.futures\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 # 根据需要取消注释\nimport pickle\nimport datetime\nfrom tensorflow.keras import mixed_precision\nimport shutil\n\ntry:\n    policy = mixed_precision.Policy('mixed_float16')\n    mixed_precision.set_global_policy(policy)\n    print(\"全局混合精度策略已成功设置为 'mixed_float16'。\")\n    print(f\"计算数据类型: {policy.compute_dtype}\") # 应该是 float16\n    print(f\"变量数据类型: {policy.variable_dtype}\") # 应该是 float32\nexcept Exception as e:\n    print(f\"设置混合精度策略失败: {e}\")\n    print(\"将继续使用默认精度 (float32)。\")\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)\n\n# --- 图像和模型参数 ---\nIMG_SIZE = (224, 224)\nTARGET_SIZE = IMG_SIZE[0]\nTARGET_SIZE_INT = IMG_SIZE[0]\nN_INPUT_CHANNELS = 3\nBEST_NIFTI_ORIENTATION_TRANSFORM = ['rot90']\nUSE_REVERSE_NIFTI_MAPPING = True # 设置为True来启用反向映射\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) \nprint(f\"最终最佳模型将保存到 (可写路径): {MODEL_SAVE_PATH}\")\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-05-17T15:25:54.954033Z","iopub.execute_input":"2025-05-17T15:25:54.954356Z","iopub.status.idle":"2025-05-17T15:25:54.967634Z","shell.execute_reply.started":"2025-05-17T15:25:54.954329Z","shell.execute_reply":"2025-05-17T15:25:54.966806Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 随机种子&辅助函数","metadata":{}},{"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\n\n    \n# 在您的 Cell 3 (随机种子&辅助函数) 中修改或添加\n\ndef apply_orientation_transform(mask_slice, transform_ops=None):\n    \"\"\"\n    应用方向变换到掩码切片。\n    transform_ops: 一个包含操作字符串的列表，例如 ['rot90', 'fliplr']\n                   可能的op: 'rot90', 'rot90_2', 'rot90_3', 'fliplr', 'flipud'\n    \"\"\"\n    if transform_ops is None:\n        return mask_slice\n\n    transformed_mask = mask_slice.copy()\n    for op in transform_ops:\n        if op == 'rot90':\n            transformed_mask = np.rot90(transformed_mask)\n        elif op == 'rot90_2': # 旋转180度\n            transformed_mask = np.rot90(transformed_mask, k=2)\n        elif op == 'rot90_3': # 旋转270度\n            transformed_mask = np.rot90(transformed_mask, k=3)\n        elif op == 'fliplr': # 左右翻转\n            transformed_mask = np.fliplr(transformed_mask)\n        elif op == 'flipud': # 上下翻转\n            transformed_mask = np.flipud(transformed_mask)\n        else:\n            print(f\"警告: 未知的方向变换操作 '{op}'\")\n    return transformed_mask\n\ndef load_multi_organ_segmentation_mask(\n    nii_data_array,\n    slice_index,\n    organ_map_nii,\n    organ_channel_map,\n    num_organs,\n    target_size, # 目标尺寸暂时保留，但我们先在原始尺寸上做方向调整\n    nii_path_for_error_msg=\"\",\n    orientation_transform_ops=None # 新增参数，例如 ['rot90'] 或 ['rot90', 'fliplr']\n):\n    \"\"\"\n    从已加载的 NII 数据数组中提取特定切片的分割掩码, 创建多通道掩码。\n    根据提供的映射处理标签 (合并左右肾到'kidney', 标签5到'bowel')。\n    返回多通道二值掩码(0或1), 形状为(target_size, target_size, num_organs)\n    \"\"\"\n    try:\n        seg_data = nii_data_array\n\n        if not isinstance(seg_data, np.ndarray) or seg_data.ndim != 3:\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_nii = seg_data.shape[2]\n        if not (0 <= slice_index < num_slices_nii):\n            # print(f\"警告: 切片索引 {slice_index} 超出范围 [0, {num_slices_nii-1}) for {nii_path_for_error_msg}\")\n            return None\n\n        mask_slice_float = seg_data[:, :, slice_index]\n        # --- 关键修改点：应用方向变换 ---\n        if orientation_transform_ops:\n            print(f\"对掩码切片应用方向变换: {orientation_transform_ops}\")\n            mask_slice_float = apply_orientation_transform(mask_slice_float, orientation_transform_ops)\n        # --- 方向变换结束 ---\n        \n        mask_slice_int = np.round(mask_slice_float).astype(np.int16)\n\n        # 创建多通道掩码 (在变换后的掩码尺寸上创建)\n        # 注意：这里 multi_channel_mask 的尺寸是变换后的原始掩码尺寸，还未resize\n        multi_channel_mask = np.zeros((mask_slice_int.shape[0], mask_slice_int.shape[1], num_organs), dtype=np.float32)\n\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                    binary_mask_organ = ((mask_slice_int == 3) | (mask_slice_int == 4)).astype(np.float32)\n                else:\n                    binary_mask_organ = (mask_slice_int == nii_value).astype(np.float32)\n                \n                # 确保 binary_mask_organ 和 multi_channel_mask 的前两维匹配\n                if binary_mask_organ.shape == multi_channel_mask.shape[:2]:\n                    multi_channel_mask[:, :, channel_idx] += binary_mask_organ\n                else:\n                    # 如果因为旋转导致尺寸不一致（应该不会，因为旋转保持尺寸），这里需要处理或报错\n                    print(f\"警告: 器官 {organ_name} 的二值掩码形状 {binary_mask_organ.shape} 与多通道掩码基底形状 {multi_channel_mask.shape[:2]} 不匹配。\")\n                    # 尝试resize binary_mask_organ 到 multi_channel_mask 的尺寸\n                    resized_binary_mask_organ = cv2.resize(binary_mask_organ, (multi_channel_mask.shape[1], multi_channel_mask.shape[0]), interpolation=cv2.INTER_NEAREST)\n                    multi_channel_mask[:, :, channel_idx] += resized_binary_mask_organ\n\n\n        # --- Resize 操作 ---\n        # 现在 multi_channel_mask 是在（可能旋转/翻转过的）原始切片分辨率下的\n        # 将其resize到目标尺寸\n        if multi_channel_mask.shape[0] != target_size or multi_channel_mask.shape[1] != target_size:\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                # 这种情况不应该发生，如果发生了，说明resize逻辑有问题，或者输入掩码有问题\n                print(f\"警告: Resize后多通道掩码被意外压缩为2D: {resized_mask.shape}, 目标通道: {num_organs} from {nii_path_for_error_msg}\")\n                # 为避免错误，返回None或全零掩码\n                return np.zeros((target_size, target_size, num_organs), dtype=np.float32) # 返回全零\n            resized_mask = (resized_mask > 0.5).astype(np.float32) # 二值化确保是0或1\n        else:\n            # 如果原始尺寸（经过变换后）已经是目标尺寸，则直接二值化\n            resized_mask = (multi_channel_mask > 0.5).astype(np.float32)\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             # 可以选择返回None或一个全零的掩码\n             return np.zeros((target_size, target_size, num_organs), dtype=np.float32) # 返回全零\n\n        return resized_mask\n\n    except Exception as e:\n        print(f\"处理分割掩码错误 (来自预加载数据) {nii_path_for_error_msg}, 切片 {slice_index}, 变换 {orientation_transform_ops}: {e}\")\n        import traceback\n        traceback.print_exc() # 打印详细的错误堆栈\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-05-17T15:25:55.024841Z","iopub.execute_input":"2025-05-17T15:25:55.025153Z","iopub.status.idle":"2025-05-17T15:25:55.052461Z","shell.execute_reply.started":"2025-05-17T15:25:55.025127Z","shell.execute_reply":"2025-05-17T15:25:55.051664Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 建议在 In[4] visualize_preprocessed_samples 函数的上方或一个新的Cell中添加此函数\n\ndef debug_dicom_nifti_alignment(\n    dicom_path,\n    nii_data_array, # 预加载的整个3D NII数据\n    nii_slice_idx,  # 要从NII数据中提取的切片索引\n    organ_map_nii,\n    organ_channel_map,\n    num_organs,\n    target_size, # 这是最终模型期望的尺寸\n    orientation_transform_ops_list=None # 一个包含多种变换操作列表的列表, e.g., [None, ['rot90'], ['fliplr']]\n):\n    \"\"\"\n    调试单个DICOM图像和其对应的NIFTI掩码（应用不同方向变换）的对齐情况。\n    \"\"\"\n    print(f\"调试对齐: DICOM='{os.path.basename(dicom_path)}', NII切片索引={nii_slice_idx}\")\n\n    # 1. 加载和预处理DICOM图像\n    dicom_image_raw, instance_number = load_dicom_slice(dicom_path)\n    if dicom_image_raw is None:\n        print(f\"无法加载DICOM图像: {dicom_path}\")\n        return\n    \n    # 将DICOM图像调整到target_size以进行比较（注意：U-Net输入是3通道的）\n    # 为了可视化，我们先用原始单通道灰度图\n    dicom_display = cv2.resize(dicom_image_raw, (target_size, target_size), interpolation=cv2.INTER_LINEAR)\n\n    if orientation_transform_ops_list is None:\n        orientation_transform_ops_list = [None] # 默认只显示原始（无变换）\n\n    num_transforms = len(orientation_transform_ops_list)\n    \n    # 为每个变换创建一个图\n    for i, current_ops in enumerate(orientation_transform_ops_list):\n        print(f\"\\n尝试变换: {current_ops}\")\n        \n        # 2. 加载和处理NIFTI掩码（应用当前方向变换）\n        # 注意：这里调用修改后的 load_multi_organ_segmentation_mask\n        # 它会在内部处理方向，然后resize到target_size\n        nifti_mask_multichannel = load_multi_organ_segmentation_mask(\n            nii_data_array,\n            nii_slice_idx,\n            organ_map_nii,\n            organ_channel_map,\n            num_organs,\n            target_size, # 确保掩码也被resize到同样大小\n            nii_path_for_error_msg=f\"debug_patient_slice_{nii_slice_idx}\",\n            orientation_transform_ops=current_ops\n        )\n\n        if nifti_mask_multichannel is None:\n            print(f\"无法为变换 {current_ops} 加载NIFTI掩码。\")\n            # 可以选择画一个空白掩码图\n            fig_title_suffix = f\"(变换: {current_ops}) - 掩码加载失败\"\n            nifti_mask_multichannel_display = np.zeros((target_size, target_size, 3), dtype=np.uint8) # 用于显示的空白彩色图\n            blended_display = (dicom_display * 255).astype(np.uint8)\n            if len(blended_display.shape) == 2: # 如果是单通道灰度图，转为BGR\n                blended_display = cv2.cvtColor(blended_display, cv2.COLOR_GRAY2BGR)\n\n        else:\n            fig_title_suffix = f\"(变换: {current_ops if current_ops else '无'})\"\n            # 创建彩色叠加图进行可视化 (与您 visualize_preprocessed_samples 中的逻辑类似)\n            colors = {\n                'liver': [255, 0, 0],  # 红色 (注意这里用BGR顺序，因为OpenCV常用BGR)\n                'spleen': [0, 255, 0],  # 绿色\n                'kidney': [0, 0, 255],  # 蓝色\n                'bowel': [255, 255, 0]    # 黄色\n            }\n            \n            # 将单通道DICOM显示图像转换为BGR，以便与彩色掩码叠加\n            dicom_display_bgr = (dicom_display * 255).astype(np.uint8)\n            if len(dicom_display_bgr.shape) == 2:\n                dicom_display_bgr = cv2.cvtColor(dicom_display_bgr, cv2.COLOR_GRAY2BGR)\n\n            # 创建掩码的彩色叠加版本\n            overlay_mask_display = np.zeros_like(dicom_display_bgr, dtype=np.uint8) # BGR\n            for organ_name_map, channel_idx_map in organ_channel_map.items():\n                if organ_name_map in colors:\n                    color_bgr = colors[organ_name_map] # 直接使用BGR\n                    # nifti_mask_multichannel 是 (target_size, target_size, num_organs)\n                    organ_mask_slice = nifti_mask_multichannel[:, :, channel_idx_map]\n                    # 将单通道二值掩码应用颜色，并叠加到 overlay_mask_display\n                    for c in range(3): # B, G, R\n                        overlay_mask_display[organ_mask_slice > 0, c] = color_bgr[c]\n            \n            # 混合图像和掩码\n            alpha = 0.4\n            blended_display = cv2.addWeighted(dicom_display_bgr, 1 - alpha, overlay_mask_display, alpha, 0)\n\n\n        # 3. 可视化\n        plt.figure(figsize=(8, 8))\n        plt.imshow(cv2.cvtColor(blended_display, cv2.COLOR_BGR2RGB)) # Matplotlib期望RGB\n        plt.title(f\"DICOM与NIFTI掩码叠加 {fig_title_suffix}\\nDICOM: {os.path.basename(dicom_path)}, NII切片: {nii_slice_idx}\")\n        plt.axis('off')\n        \n        # 添加图例 (可选，但推荐)\n        legend_elements = [plt.Rectangle((0, 0), 1, 1, color=[c/255. for c in colors[org][::-1]], label=org) #转RGB给matplotlib\n                           for org in organ_channel_map.keys() if org in colors]\n        plt.legend(handles=legend_elements, bbox_to_anchor=(1.05, 1), loc='upper left')\n        plt.tight_layout(rect=[0, 0, 0.85, 1]) # 为图例留出空间\n        plt.show()\n\n\ndef 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-05-17T15:25:55.094794Z","iopub.execute_input":"2025-05-17T15:25:55.095461Z","iopub.status.idle":"2025-05-17T15:25:55.125009Z","shell.execute_reply.started":"2025-05-17T15:25:55.095396Z","shell.execute_reply":"2025-05-17T15:25:55.123907Z"}},"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-05-17T15:25:55.126501Z","iopub.execute_input":"2025-05-17T15:25:55.126816Z","iopub.status.idle":"2025-05-17T15:25:55.153543Z","shell.execute_reply.started":"2025-05-17T15:25:55.12679Z","shell.execute_reply":"2025-05-17T15:25:55.152566Z"}},"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-05-17T15:25:55.154474Z","iopub.execute_input":"2025-05-17T15:25:55.154766Z","iopub.status.idle":"2025-05-17T15:25:55.175994Z","shell.execute_reply.started":"2025-05-17T15:25:55.154741Z","shell.execute_reply":"2025-05-17T15:25:55.174921Z"}},"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-05-17T15:25:55.177953Z","iopub.execute_input":"2025-05-17T15:25:55.178285Z","iopub.status.idle":"2025-05-17T15:25:55.207861Z","shell.execute_reply.started":"2025-05-17T15:25:55.178257Z","shell.execute_reply":"2025-05-17T15:25:55.207033Z"}},"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-05-17T15:25:55.208751Z","iopub.execute_input":"2025-05-17T15:25:55.209622Z","iopub.status.idle":"2025-05-17T15:25:55.229613Z","shell.execute_reply.started":"2025-05-17T15:25:55.209592Z","shell.execute_reply":"2025-05-17T15:25:55.228723Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 损失函数&评价指标","metadata":{}},{"cell_type":"code","source":"# 在 Cell 9 (\"损失函数&评价指标\")\n\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 dice_loss_single_channel(y_true, y_pred):\n    \"\"\"计算单个通道的 1 - Dice系数 作为损失\"\"\"\n    return 1.0 - dice_coefficient(y_true, y_pred)\n\n# --- 平均Dice系数 ---\n@tf.function\ndef average_dice_coefficient(y_true, y_pred):\n    \"\"\"计算所有类别Dice系数的平均值\"\"\"\n    total_dice = 0.0\n    # 假设 NUM_ORGANS 是全局定义的，代表总的器官通道数 (例如 4)\n    # 并且 dice_coefficient 函数也已定义\n    for i in range(NUM_ORGANS):\n        dice_ch = dice_coefficient(y_true[..., i], y_pred[..., i])\n        total_dice += dice_ch\n    return total_dice / tf.cast(NUM_ORGANS, tf.float32)\n\n# --- Focal Loss (基于Binary Crossentropy) ---\n@tf.function\ndef focal_loss_bce(y_true, y_pred, gamma=2.0, alpha=0.25):\n    \"\"\"\n    Binary Focal Loss.\n    FL(pt) = -alpha_t * (1 - pt)**gamma * log(pt)\n    pt is the probability of the true class.\n    \"\"\"\n    y_pred = tf.clip_by_value(y_pred, SMOOTH, 1.0 - SMOOTH) # 避免log(0)\n    \n    # Calculate Chross Entropy\n    cross_entropy = -y_true * tf.math.log(y_pred) - (1.0 - y_true) * tf.math.log(1.0 - y_pred)\n    \n    # Calculate P_t\n    p_t = (y_true * y_pred) + ((1.0 - y_true) * (1.0 - y_pred))\n    \n    # Calculate Focal Loss\n    focal_term = (1.0 - p_t) ** gamma\n    \n    # Weighted Focal Loss\n    loss = alpha * focal_term * cross_entropy # 使用固定的 alpha (可调整)\n    \n    return tf.reduce_mean(loss) # 对batch和像素取平均\n\n\n# --- 新的 Focal Dice Loss ---\ndef create_focal_dice_loss(gamma_focal=2.0, alpha_focal=0.25, lambda_focal=0.5, lambda_dice=0.5, class_weights=None):\n    \"\"\"\n    创建结合 Focal Loss (基于BCE) 和 Dice Loss 的损失函数，支持类别权重。\n    Args:\n        gamma_focal: Focal loss的gamma参数.\n        alpha_focal: Focal loss的alpha参数 (单个值，或每个类别的列表/数组).\n        lambda_focal: Focal loss的权重.\n        lambda_dice: Dice loss的权重.\n        class_weights: 每个器官通道的权重列表/数组，用于加权Dice Loss和Focal Loss (如果alpha_focal是单个值).\n                       顺序应与 ORGAN_CHANNEL_MAP 中的通道索引一致。\n    \"\"\"\n    _class_weights = tf.constant(class_weights if class_weights is not None else [1.0] * NUM_ORGANS, dtype=tf.float32)\n    _alpha_focal = alpha_focal # 可以是单个值或列表/数组\n\n    @tf.function\n    def focal_dice_loss_fn(y_true, y_pred):\n        total_loss = tf.constant(0.0, dtype=tf.float32)\n        \n        for i in range(NUM_ORGANS):\n            y_true_ch = y_true[..., i]\n            y_pred_ch = y_pred[..., i]\n            \n            # Dice Loss for this channel\n            dice_l = dice_loss_single_channel(y_true_ch, y_pred_ch)\n            \n            # Focal Loss (BCE based) for this channel\n            # 如果 alpha_focal 是列表，则按通道取值\n            current_alpha = _alpha_focal[i] if isinstance(_alpha_focal, (list, tuple, tf.Tensor, np.ndarray)) and len(_alpha_focal) == NUM_ORGANS else _alpha_focal\n            focal_l = focal_loss_bce(y_true_ch, y_pred_ch, gamma=gamma_focal, alpha=current_alpha)\n            \n            # 结合 Focal Loss 和 Dice Loss，并应用类别权重\n            channel_loss = (lambda_focal * focal_l + lambda_dice * dice_l) * _class_weights[i]\n            total_loss += channel_loss\n            \n        return total_loss / tf.reduce_sum(_class_weights) # 加权平均或总和，这里用加权平均\n        # 或者 return total_loss / tf.cast(NUM_ORGANS, tf.float32) 如果不希望权重影响总损失的尺度\n\n    return focal_dice_loss_fn\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]","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-17T15:25:55.230551Z","iopub.execute_input":"2025-05-17T15:25:55.231323Z","iopub.status.idle":"2025-05-17T15:25:55.25158Z","shell.execute_reply.started":"2025-05-17T15:25:55.231294Z","shell.execute_reply":"2025-05-17T15:25:55.250654Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 分阶段训练函数","metadata":{}},{"cell_type":"code","source":"# 在 Cell 10 (\"分阶段训练函数\")\n\ndef train_in_stages(train_pids, val_pids, preprocessed_dir, target_size_int_param,\n                      input_channels, num_organs_param, batch_size, \n                      initial_learning_rate, epochs_per_stage, class_weights_map): # class_weights_map 从 analyze_organ_distribution 获取\n    \n    input_shape = (target_size_int_param, target_size_int_param, input_channels)\n    print(f\"构建模型，输入形状: {input_shape}, 器官数: {num_organs_param}\")\n    unet_model = build_unet_multi_organ(input_shape, num_organs_param, dropout_rate=0.3) # 使用您 Cell 8 的定义\n\n    # 准备类别权重，确保顺序与 ORGAN_CHANNEL_MAP 一致\n    # 这里的 class_weights_map 应该是类似 {'liver': w1, 'spleen': w2, ...} 的字典\n    # 我们需要将其转换为一个列表，顺序与模型输出通道对应\n    ordered_class_weights = [1.0] * num_organs_param # 初始化为1.0\n    if class_weights_map: # 确保 class_weights_map 不是 None\n        for organ, idx in ORGAN_CHANNEL_MAP.items(): # ORGAN_CHANNEL_MAP 是全局的\n            if organ in class_weights_map:\n                ordered_class_weights[idx] = class_weights_map[organ]\n    print(f\"训练中使用的有序类别权重: {ordered_class_weights}\")\n\n    # ***** 定义新的损失函数实例 *****\n    # 您可以调整 FocalDiceLoss 的超参数\n    # 例如，给罕见或难分的器官更高的权重（通过 class_weights）\n    # lambda_focal 和 lambda_dice 控制两部分损失的贡献，论文没有指明，可以设为0.5, 0.5开始\n    final_stage_loss = create_focal_dice_loss(\n        gamma_focal=2.0, \n        alpha_focal=0.25, # 或者可以是一个列表，为每个通道设置不同的alpha\n        lambda_focal=0.5, \n        lambda_dice=0.5, \n        class_weights=ordered_class_weights\n    )\n    # ********************************\n\n    # --- 定义各阶段参数 ---\n    # 对于前几个阶段，如果只想关注特定器官，可以创建只针对那些器官的损失或权重\n    # 例如，Stage1_Liver 可以继续使用 1.0 - dice_liver\n    # 或者也使用 FocalDiceLoss，但权重只给 liver\n    \n    # 为每个阶段动态创建损失函数\n    stage_losses = []\n    for i in range(len(epochs_per_stage)):\n        current_weights = [0.0] * num_organs_param\n        if i == 0: # Stage 1: Liver\n            current_weights[ORGAN_CHANNEL_MAP['liver']] = ordered_class_weights[ORGAN_CHANNEL_MAP['liver']] \\\n                                                            if 'liver' in ORGAN_CHANNEL_MAP and ordered_class_weights else 1.0\n            stage_losses.append(create_focal_dice_loss(class_weights=current_weights, lambda_focal=0.5, lambda_dice=0.5)) # 可以调整lambda\n        elif i == 1: # Stage 2: Liver, Spleen\n            if 'liver' in ORGAN_CHANNEL_MAP: current_weights[ORGAN_CHANNEL_MAP['liver']] = ordered_class_weights[ORGAN_CHANNEL_MAP['liver']] if ordered_class_weights else 1.0\n            if 'spleen' in ORGAN_CHANNEL_MAP: current_weights[ORGAN_CHANNEL_MAP['spleen']] = ordered_class_weights[ORGAN_CHANNEL_MAP['spleen']] if ordered_class_weights else 1.0\n            stage_losses.append(create_focal_dice_loss(class_weights=current_weights, lambda_focal=0.5, lambda_dice=0.5))\n        elif i == 2: # Stage 3: Liver, Spleen, Kidney\n            if 'liver' in ORGAN_CHANNEL_MAP: current_weights[ORGAN_CHANNEL_MAP['liver']] = ordered_class_weights[ORGAN_CHANNEL_MAP['liver']] if ordered_class_weights else 1.0\n            if 'spleen' in ORGAN_CHANNEL_MAP: current_weights[ORGAN_CHANNEL_MAP['spleen']] = ordered_class_weights[ORGAN_CHANNEL_MAP['spleen']] if ordered_class_weights else 1.0\n            if 'kidney' in ORGAN_CHANNEL_MAP: current_weights[ORGAN_CHANNEL_MAP['kidney']] = ordered_class_weights[ORGAN_CHANNEL_MAP['kidney']] if ordered_class_weights else 1.0\n            stage_losses.append(create_focal_dice_loss(class_weights=current_weights, lambda_focal=0.5, lambda_dice=0.5))\n        elif i == 3: # Stage 4: All Organs\n            stage_losses.append(final_stage_loss) # 使用为所有器官配置的FocalDiceLoss\n\n    stages = [\n        {'name': 'Stage1_LiverFocus', 'epochs': epochs_per_stage[0], 'lr_factor': 1.0, \n         'loss': stage_losses[0], 'monitor': 'val_dice_liver'}, # 仍然监控val_dice_liver\n        {'name': 'Stage2_LiverSpleenFocus', 'epochs': epochs_per_stage[1], 'lr_factor': 0.5,\n         'loss': stage_losses[1], 'monitor': 'val_average_dice_coefficient'}, # 监控平均Dice\n        {'name': 'Stage3_LiverSpleenKidneyFocus', 'epochs': epochs_per_stage[2], 'lr_factor': 0.2,\n         'loss': stage_losses[2], 'monitor': 'val_average_dice_coefficient'},\n        {'name': 'Stage4_AllOrgans', 'epochs': epochs_per_stage[3], 'lr_factor': 0.1,\n         'loss': stage_losses[3], 'monitor': 'val_average_dice_coefficient'} # 最终监控平均Dice\n    ]\n\n    stage_histories = {}\n    # MODEL_SAVE_PATH 现在应该指向 /kaggle/working/...\n    # best_model_path_overall = MODEL_SAVE_PATH # 已在 Cell 2 全局定义并修正\n\n    for i_stage_loop, stage_info in enumerate(stages): # 使用新的索引名避免与外部i_stage冲突\n        stage_name = stage_info['name']\n        epochs = stage_info['epochs']\n        current_lr = initial_learning_rate * stage_info['lr_factor']\n        loss_func_for_stage = stage_info['loss'] # 这是已经创建好的损失函数实例\n        monitor_metric = stage_info['monitor']\n        \n        # 确保MODEL_OUTPUT_DIR是全局定义的 /kaggle/working/unet_model_v2\n        stage_model_save_path = os.path.join(MODEL_OUTPUT_DIR, f\"{stage_name}_best_model.keras\")\n\n        print(f\"\\n=== {stage_name} 训练 ===\")\n        # ... (打印学习率、监控指标等信息不变) ...\n        print(f\"  阶段模型将保存到: {stage_model_save_path}\")\n        if stage_name == stages[-1]['name']: # 检查是否是最后一个阶段\n            print(f\"  最终最佳模型将保存到: {MODEL_SAVE_PATH}\") # 使用全局 MODEL_SAVE_PATH\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        optimizer = optimizers.Adam(learning_rate=current_lr)\n        \n        # 在这里编译模型，使用当前阶段的损失函数\n        unet_model.compile(optimizer=optimizer, loss=loss_func_for_stage, metrics=METRICS)\n        print(f\"模型已为阶段 {stage_name} 编译，损失函数: {loss_func_for_stage.__name__ if hasattr(loss_func_for_stage, '__name__') else str(loss_func_for_stage)}\")\n\n\n        # 回调函数 (与您之前的版本类似，确保路径正确)\n        callbacks_list_stage = [\n            callbacks.ModelCheckpoint(\n                stage_model_save_path, # 保存到阶段特定的路径\n                monitor=monitor_metric, mode='max', save_best_only=True,\n                save_weights_only=False, verbose=1\n            ),\n            callbacks.EarlyStopping(\n                monitor=monitor_metric, mode='max', patience=EARLY_STOPPING_PATIENCE,\n                verbose=1, restore_best_weights=True\n            ),\n            callbacks.ReduceLROnPlateau(\n                monitor=monitor_metric, mode='max', factor=REDUCE_LR_FACTOR,\n                patience=REDUCE_LR_PATIENCE, min_lr=MIN_LR, verbose=1\n            ),\n            callbacks.TensorBoard(\n                log_dir=os.path.join(OUTPUT_DIR, 'logs', stage_name),\n                histogram_freq=1, update_freq='epoch'\n            )\n        ]\n        \n        if stage_name == stages[-1]['name']: # 如果是最后一个阶段\n            checkpoint_callback_final = callbacks.ModelCheckpoint(\n                MODEL_SAVE_PATH, # 全局定义的最终模型保存路径\n                monitor=monitor_metric, mode='max', save_best_only=True,\n                save_weights_only=False, verbose=1, save_freq='epoch'\n            )\n            callbacks_list_stage.append(checkpoint_callback_final)\n        \n        print(f\"开始训练阶段: {stage_name}，共 {epochs} 个 Epochs\")\n        history = unet_model.fit(\n            train_dataset,\n            validation_data=val_dataset,\n            epochs=epochs,\n            callbacks=callbacks_list_stage,\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_name}' 完成。从检查点 '{stage_model_save_path}' 加载此阶段的最佳模型...\")\n            custom_objects_for_load = { # 只需要自定义指标\n                'average_dice_coefficient': average_dice_coefficient, # 如果在METRICS中使用了\n                'dice_liver': dice_liver, 'dice_spleen': dice_spleen,\n                'dice_kidney': dice_kidney, 'dice_bowel': dice_bowel\n            }\n            try:\n                # 加载模型时不编译，因为下一阶段会重新编译\n                unet_model = models.load_model(stage_model_save_path, custom_objects=custom_objects_for_load, compile=False)\n                print(f\"模型已从 {stage_model_save_path} 成功加载结构和权重。\")\n            except Exception as e_load:\n                print(f\"警告: 从阶段检查点 '{stage_model_save_path}' 加载模型失败: {e_load}\")\n                print(\"将继续使用内存中当前的模型（可能已由EarlyStopping恢复了最佳权重）。\")\n        else:\n            print(f\"警告: 阶段检查点文件 '{stage_model_save_path}' 未找到。将使用内存中当前阶段训练后的模型。\")\n        \n        gc.collect()\n\n    print(\"\\n所有训练阶段完成。\")\n    return stage_histories","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-17T15:25:55.253354Z","iopub.execute_input":"2025-05-17T15:25:55.253736Z","iopub.status.idle":"2025-05-17T15:25:55.274955Z","shell.execute_reply.started":"2025-05-17T15:25:55.253713Z","shell.execute_reply":"2025-05-17T15:25:55.274137Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 模型评估","metadata":{}},{"cell_type":"code","source":"# 评估函数\n# 在 Cell 11 (\"模型评估\") 中，修改 process_patient_evaluation 函数\n\n# (确保 apply_orientation_transform 函数在此作用域内可用，或者 BEST_NIFTI_ORIENTATION_TRANSFORM 和 USE_REVERSE_NIFTI_MAPPING 是全局的)\n\ndef process_patient_evaluation(patient_id, nii_path, dicom_paths, series_pred_dir, pred_files,\n                                 organ_map_nii_local, organ_channel_map_local, target_size_local_int, # 使用局部变量名和整数尺寸\n                                 prediction_threshold):\n    local_dice_scores = {organ: [] for organ in organ_channel_map_local.keys()}\n    local_iou_scores = {organ: [] for organ in organ_channel_map_local.keys()}\n    local_processed = 0\n\n    try:\n        nii_img = nib.load(nii_path)\n        gt_data_float = nii_img.get_fdata(dtype=np.float32) # 直接加载为 float32\n        nii_total_slices = gt_data_float.shape[2]\n    except Exception as e:\n        print(f\"评估时无法加载患者 {patient_id} 的真实掩码 '{nii_path}': {e}\")\n        return local_dice_scores, local_iou_scores, local_processed\n\n    # ... (instance_to_pred_path 的逻辑不变) ...\n    instance_to_pred_path = {\n        int(os.path.splitext(os.path.basename(f))[0]): os.path.join(series_pred_dir, f) # 确保路径完整\n        for f in os.listdir(series_pred_dir) # 直接 listdir series_pred_dir\n        if os.path.splitext(os.path.basename(f))[0].isdigit() and f.endswith(\".npz\")\n    }\n    # 为非数字文件名添加 (如果您的预测文件名可能不是纯数字)\n    for f in os.listdir(series_pred_dir):\n        if f.endswith(\".npz\"):\n            basename = os.path.splitext(os.path.basename(f))[0]\n            if not basename.isdigit():\n                instance_to_pred_path[basename] = os.path.join(series_pred_dir, f)\n\n\n    # dicom_paths 是一个 (instance_number, path) 的元组列表\n    # 我们需要的是DICOM在其原始排序列表中的索引 (dicom_list_idx)\n    for dicom_list_idx, (instance_number, dicom_path) in enumerate(dicom_paths):\n        # instance_number 已经是整数了，来自 get_dicom_files_dict\n        if instance_number is None: # 以防万一\n            print(f\"警告: 患者 {patient_id} 的 DICOM {dicom_path} 缺少InstanceNumber，使用列表索引。\")\n            # 如果instance_number可能为None，需要一个备用方案来匹配预测文件，\n            # 或者在get_dicom_files_dict中确保instance_number总是一个有效值或唯一标识符\n            id_for_pred = f\"idx{dicom_list_idx}\" # 假设预测文件名可能是基于索引的\n        else:\n            id_for_pred = instance_number\n\n\n        # ***** 核心修改：获取正确的NIFTI切片索引 *****\n        nii_slice_idx_to_use = -1\n        if USE_REVERSE_NIFTI_MAPPING: # 使用全局配置\n            nii_slice_idx_to_use = nii_total_slices - 1 - dicom_list_idx\n        else:\n            nii_slice_idx_to_use = dicom_list_idx\n        # *********************************************\n\n        if not (0 <= nii_slice_idx_to_use < nii_total_slices):\n            # print(f\"评估患者 {patient_id}: NIFTI索引 {nii_slice_idx_to_use} (来自DICOM列表索引 {dicom_list_idx}) 超出范围。\")\n            continue\n            \n        pred_file_path = instance_to_pred_path.get(id_for_pred)\n        if not pred_file_path and str(id_for_pred) in instance_to_pred_path: # 再尝试字符串形式的键\n            pred_file_path = instance_to_pred_path[str(id_for_pred)]\n            \n        if not pred_file_path:\n            # print(f\"评估患者 {patient_id}: 未找到 InstanceNumber/ID {id_for_pred} 对应的预测文件。\")\n            continue\n            \n        try:\n            pred_data = np.load(pred_file_path)\n            # 假设 'mask' 是 (H, W, NumChannels) 并且是概率值或二值化后的 (0或1)\n            pred_mask_from_npz = pred_data['mask'] \n            # 如果保存的是概率，在这里应用阈值；如果已经是二值，确保是float32\n            pred_mask_binary = (pred_mask_from_npz > prediction_threshold).astype(np.float32)\n\n            if pred_mask_binary.shape[:2] != (target_size_local_int, target_size_local_int) or \\\n               pred_mask_binary.shape[2] != len(organ_channel_map_local):\n                print(f\"评估患者 {patient_id}: 预测掩码 {os.path.basename(pred_file_path)} 形状 {pred_mask_binary.shape} 不正确。预期 ({target_size_local_int},{target_size_local_int},{len(organ_channel_map_local)})\")\n                continue\n        except Exception as e_load_pred:\n            print(f\"评估患者 {patient_id}: 加载或处理预测掩码 {os.path.basename(pred_file_path)} 失败: {e_load_pred}\")\n            continue\n\n        gt_slice_raw = gt_data_float[:, :, nii_slice_idx_to_use]\n\n        # ***** 核心修改：对真实掩码应用方向变换 *****\n        gt_slice_oriented = apply_orientation_transform(gt_slice_raw, BEST_NIFTI_ORIENTATION_TRANSFORM) # 使用全局变量\n        # *******************************************\n        \n        gt_slice_int = np.round(gt_slice_oriented).astype(np.int16)\n\n        for organ_name, channel_idx in organ_channel_map_local.items():\n            if organ_name == 'kidney':\n                gt_organ_binary_raw = ((gt_slice_int == 3) | (gt_slice_int == 4)).astype(np.float32)\n            elif organ_name == 'liver':\n                gt_organ_binary_raw = (gt_slice_int == 1).astype(np.float32)\n            elif organ_name == 'spleen':\n                gt_organ_binary_raw = (gt_slice_int == 2).astype(np.float32)\n            elif organ_name == 'bowel':\n                gt_organ_binary_raw = (gt_slice_int == 5).astype(np.float32)\n            else:\n                continue\n\n            if gt_organ_binary_raw.shape != (target_size_local_int, target_size_local_int):\n                gt_organ_resized = cv2.resize(gt_organ_binary_raw, \n                                              (target_size_local_int, target_size_local_int), \n                                              interpolation=cv2.INTER_NEAREST)\n            else:\n                gt_organ_resized = gt_organ_binary_raw\n            \n            gt_organ_resized = (gt_organ_resized > 0.5).astype(np.float32)\n            pred_organ_binary_channel = pred_mask_binary[..., channel_idx]\n\n            if np.sum(gt_organ_resized) > 0 or np.sum(pred_organ_binary_channel) > 0:\n                dice = (2. * np.sum(gt_organ_resized * pred_organ_binary_channel) + 1e-6) / \\\n                       (np.sum(gt_organ_resized) + np.sum(pred_organ_binary_channel) + 1e-6)\n                intersection = np.sum(gt_organ_resized * pred_organ_binary_channel)\n                union = np.sum(gt_organ_resized) + np.sum(pred_organ_binary_channel) - intersection\n                iou = (intersection + 1e-6) / (union + 1e-6)\n                local_dice_scores[organ_name].append(dice)\n                local_iou_scores[organ_name].append(iou)\n        local_processed += 1\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-05-17T15:25:55.28574Z","iopub.execute_input":"2025-05-17T15:25:55.286415Z","iopub.status.idle":"2025-05-17T15:25:55.313332Z","shell.execute_reply.started":"2025-05-17T15:25:55.28639Z","shell.execute_reply":"2025-05-17T15:25:55.312645Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 保存并预测掩码","metadata":{}},{"cell_type":"code","source":"def predict_and_save_masks_optimized(model, patient_ids_to_predict, image_paths_dict_all_patients, output_dir, batch_size=32):\n    processed_slices_total = 0\n    failed_saves = 0\n\n    for patient_id in tqdm(patient_ids_to_predict, desc=\"为患者预测掩码\"):\n        if patient_id not in image_paths_dict_all_patients:\n            print(f\"警告: 患者ID {patient_id} 在 image_paths_dict 中未找到。跳过。\")\n            continue\n            \n        dicom_items = image_paths_dict_all_patients[patient_id]\n        if not dicom_items:\n            print(f\"警告: 患者 {patient_id} 没有DICOM项目。跳过。\")\n            continue\n            \n        try:\n            first_dicom_path = dicom_items[0][1]\n            series_id = os.path.basename(os.path.dirname(first_dicom_path))\n        except Exception:\n            series_id = \"unknown_series\"\n            \n        patient_output_dir = os.path.join(output_dir, str(patient_id), str(series_id))\n        os.makedirs(patient_output_dir, exist_ok=True)\n        \n        slices_to_process = []\n        instance_numbers_for_saving = []\n        \n        for inst_num_from_dict, actual_dicom_path_str in dicom_items:\n            image, _ = load_dicom_slice(actual_dicom_path_str) # load_dicom_slice 也返回 instance_num，但我们用字典中的\n            \n            if image is None:\n                continue\n            \n            instance_number_to_use = inst_num_from_dict\n            \n            processed_image = preprocess_image_for_unet(image, TARGET_SIZE)\n            slices_to_process.append(processed_image)\n            instance_numbers_for_saving.append(instance_number_to_use)\n            \n        if not slices_to_process:\n            print(f\"患者 {patient_id} 序列 {series_id} 没有需要处理的切片。跳过。\")\n            continue\n            \n        try:\n            for i in range(0, len(slices_to_process), batch_size):\n                batch_images = np.array(slices_to_process[i:i+batch_size])\n                batch_predictions_prob = model.predict(batch_images, verbose=0)\n                \n                for j, pred_mask_prob in enumerate(batch_predictions_prob):\n                    current_item_idx = i + j\n                    instance_number = instance_numbers_for_saving[current_item_idx]\n                    output_filename = f\"{instance_number}.npz\"\n                    output_path = os.path.join(patient_output_dir, output_filename)\n                    \n                    pred_mask_binary = (pred_mask_prob > PREDICTION_THRESHOLD).astype(np.float32)\n                    \n                    try:\n                        np.savez_compressed(output_path, mask=pred_mask_binary)\n                        processed_slices_total += 1\n                    except Exception as e_save:\n                        print(f\"保存掩码 {output_path} 失败: {e_save}\")\n                        failed_saves += 1\n                        \n                del batch_images\n                del batch_predictions_prob\n                gc.collect()\n                \n        except Exception as e_predict:\n            print(f\"患者 {patient_id} 序列 {series_id} 的批量预测过程中出错: {e_predict}\")\n    \n    print(f\"预测完成。总共处理的切片数: {processed_slices_total}，保存失败次数: {failed_saves}\")\n    return processed_slices_total, failed_saves","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-17T15:25:55.314929Z","iopub.execute_input":"2025-05-17T15:25:55.31528Z","iopub.status.idle":"2025-05-17T15:25:55.3331Z","shell.execute_reply.started":"2025-05-17T15:25:55.315258Z","shell.execute_reply":"2025-05-17T15:25:55.332385Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 主函数&可视化","metadata":{}},{"cell_type":"code","source":"def load_processed_items_from_log(log_path):\n    \"\"\"从日志文件中加载已成功处理的 (PatientID, SeriesID) 集合。\"\"\"\n    processed_items = set()\n    if not os.path.exists(log_path):\n        print(f\"日志文件 {log_path} 未找到。将从头开始处理。\")\n        return processed_items\n    \n    print(f\"尝试从现有日志加载已处理项: {log_path}\")\n    try:\n        with open(log_path, 'r', encoding='utf-8') as f_log_read:\n            header_skipped = False\n            for line_number, line in enumerate(f_log_read):\n                line = line.strip()\n                if not line or line.startswith(\"---\"): # 跳过空行和分隔符\n                    continue\n                \n                # 跳过表头行\n                if not header_skipped and line.lower().startswith(\"patientid,seriesid,status\"):\n                    header_skipped = True\n                    continue\n                \n                parts = line.split(',')\n                # 期望至少有 PatientID, SeriesID, Status 三个部分\n                if len(parts) >= 3:\n                    patient_id_log = parts[0].strip()\n                    series_id_log = parts[1].strip()\n                    status_log = parts[2].strip()\n                    \n                    # 只将状态为 \"Finished\" 的项视为已成功处理并可跳过\n                    if status_log == \"Finished\":\n                        processed_items.add((patient_id_log, series_id_log))\n                # else:\n                #     if line: # 如果行不为空但格式不正确，可以选择打印警告\n                #         print(f\"警告: 日志文件 {log_path} 中行 #{line_number+1} 格式不正确: '{line}'\")\n    except Exception as e_read_log:\n        print(f\"警告: 读取或解析日志文件 {log_path} 时发生错误: {e_read_log}。将认为没有项目被预先处理。\")\n        return set() # 发生错误时返回空集合，避免状态不一致\n    return processed_items","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-17T15:25:55.334014Z","iopub.execute_input":"2025-05-17T15:25:55.33429Z","iopub.status.idle":"2025-05-17T15:25:55.352741Z","shell.execute_reply.started":"2025-05-17T15:25:55.334262Z","shell.execute_reply":"2025-05-17T15:25:55.351906Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main_generate_unet_predictions_for_all():\n    readonly_log_path = '/kaggle/input/processed-log/prediction_processed_log.txt'\n    working_log_path = os.path.join(OUTPUT_DIR, 'prediction_processed_log.txt')\n    \n    # 复制文件（如果存在）\n    if os.path.exists(readonly_log_path):\n        print(f\"正在将日志文件从 {readonly_log_path} 复制到 {working_log_path}\")\n        shutil.copy2(readonly_log_path, working_log_path)\n        print(\"复制成功！\")\n    else:\n        print(f\"在 {readonly_log_path} 未找到日志文件，将在 {working_log_path} 创建一个新文件\")\n        # 创建空文件\n        with open(working_log_path, 'w') as f:\n            pass\n\n    calculated_class_weights = {\n        'liver': 2.51,\n        'spleen': 4.19,\n        'kidney': 3.82,\n        'bowel': 1.00\n    }\n    print(f\"使用的类别权重: {calculated_class_weights}\") # 打印出来确认一下\n    \n    print(\"--- 开始为所有患者的所有序列生成U-Net .npz预测文件 ---\")\n\n    # 更新日志文件路径为可写版本\n    log_file_path_processed_patients = working_log_path\n    print(f\"预测处理日志将/已保存到: {log_file_path_processed_patients}\")\n\n    # 在开始时加载已处理的项\n    already_processed_items = load_processed_items_from_log(log_file_path_processed_patients)\n    if already_processed_items:\n        print(f\"从日志中加载了 {len(already_processed_items)} 个先前已成功处理的 病人/序列 组合。\")\n\n    # 定义缓存文件路径\n    cache_file_path = '/kaggle/input/series-cache/other/default/1/all_patient_series_dicom_items_cache.pkl'\n    all_patient_series_dicom_items = {}\n\n    # --- 1. 构建或加载所有患者所有序列的DICOM文件映射 ---\n    if os.path.exists(cache_file_path):\n        print(f\"发现已缓存的DICOM文件映射，正在从 '{cache_file_path}' 加载...\")\n        try:\n            with open(cache_file_path, 'rb') as f:\n                all_patient_series_dicom_items = pickle.load(f)\n            print(f\"成功从缓存加载了 {len(all_patient_series_dicom_items)} 位患者的DICOM映射信息。\")\n        except Exception as e:\n            print(f\"从缓存文件 '{cache_file_path}' 加载失败: {e}. 将重新进行映射。\")\n            all_patient_series_dicom_items = {} # 重置以确保执行映射\n    \n    if not all_patient_series_dicom_items: # 如果缓存不存在或加载失败，则执行映射\n        print(\"--- 1a. 开始构建所有患者所有序列的DICOM文件映射 (缓存未找到或加载失败) ---\")\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            try:\n                dicom_tags_df = pd.read_parquet(dicom_tags_path)\n                if 'PatientID' in dicom_tags_df.columns:\n                    dicom_tags_df['PatientID'] = dicom_tags_df['PatientID'].astype(str)\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}\")\n        else:\n            print(\"未找到 DICOM tags 文件。\")\n\n        if not os.path.exists(TRAIN_IMAGES_DIR):\n            print(f\"错误: TRAIN_IMAGES_DIR '{TRAIN_IMAGES_DIR}' 不存在。\")\n            return\n            \n        all_pids_in_images_dir = sorted([\n            pid for pid in os.listdir(TRAIN_IMAGES_DIR) \n            if os.path.isdir(os.path.join(TRAIN_IMAGES_DIR, pid))\n        ])\n\n        if not all_pids_in_images_dir:\n            print(f\"错误: 在 '{TRAIN_IMAGES_DIR}' 中没有找到患者文件夹。\")\n            return\n\n        for patient_id in tqdm(all_pids_in_images_dir, desc=\"映射所有患者DICOM文件\"):\n            patient_dir_path = os.path.join(TRAIN_IMAGES_DIR, patient_id)\n            all_patient_series_dicom_items[patient_id] = {}\n            \n            series_ids_for_patient = sorted([\n                sid for sid in os.listdir(patient_dir_path)\n                if os.path.isdir(os.path.join(patient_dir_path, sid))\n            ])\n            \n            for series_id in series_ids_for_patient:\n                dicom_info_list = get_dicom_files_dict(patient_id, series_id, dicom_tags_df)\n                if dicom_info_list:\n                    all_patient_series_dicom_items[patient_id][series_id] = dicom_info_list\n        \n        print(f\"为 {len(all_patient_series_dicom_items)} 位患者的所有有效序列映射了DICOM文件。\")\n        \n        # 保存新构建的映射到缓存文件\n        working_cache_path = os.path.join(OUTPUT_DIR, 'all_patient_series_dicom_items_cache.pkl')\n        print(f\"正在将新的DICOM文件映射保存到缓存: '{working_cache_path}'\")\n        try:\n            with open(working_cache_path, 'wb') as f:\n                pickle.dump(all_patient_series_dicom_items, f)\n            print(\"DICOM文件映射已成功保存到缓存。\")\n        except Exception as e:\n            print(f\"保存DICOM文件映射到缓存失败: {e}\")\n\n    print(\"\\n--- 2. 加载最终最佳模型进行推理 ---\")\n    USER_PRETRAINED_MODEL_PATH = \"/kaggle/input/unet_effb0/keras/default/1/unet_effb0_multi_organ_224px_v2.keras\"\n    ordered_class_weights_for_load = [calculated_class_weights.get(organ, 1.0) for organ in ORGAN_CHANNEL_MAP.keys()]\n    final_loss_for_loading = create_focal_dice_loss( # 使用与训练时相同的参数\n        gamma_focal=2.0, alpha_focal=0.25, \n        lambda_focal=0.5, lambda_dice=0.5, \n        class_weights=ordered_class_weights_for_load\n    )\n    \n    custom_objects = {\n        'focal_dice_loss_fn': final_loss_for_loading, # 使用创建函数返回的实际损失函数\n        # 或者，如果损失函数被命名，使用那个名字，并在全局定义它\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        # TensorFlow有时可以直接反序列化函数对象，但更可靠的是传递名称或Loss类\n        # 如果上面的 focal_dice_loss_fn 不能直接被识别，\n        # 你可能需要将 create_focal_dice_loss 返回的函数在全局命名，\n        # 或者将 FocalDiceLoss 实现为一个 tf.keras.losses.Loss 的子类。\n        # 鉴于之前的错误，我们先尝试加载时不编译或只带指标编译。\n        best_model = models.load_model(USER_PRETRAINED_MODEL_PATH, custom_objects=custom_objects, compile=False) # 尝试 compile=False\n        # 如果需要评估，后续再用优化器和损失函数编译一次\n        best_model.compile(optimizer=optimizers.Adam(learning_rate=LEARNING_RATE*0.01), # 用一个小的学习率重新编译\n                           loss=final_loss_for_loading, \n                           metrics=METRICS)\n        print(f\"成功加载并重新编译最终模型: {USER_PRETRAINED_MODEL_PATH}\")\n    except Exception as e:\n        print(f\"加载最终模型 {USER_PRETRAINED_MODEL_PATH} 失败: {e}\")\n        print(\"确保 custom_objects 中的损失函数名称与保存时一致，或尝试仅加载权重。\")\n        # 如果完全失败，后续的推理和评估将无法进行\n        return\n\n    print(\"\\n--- 3. 开始为所有映射的患者和序列生成U-Net的.npz预测文件 ---\")\n    os.makedirs(PREDICTION_OUTPUT_DIR, exist_ok=True)\n    # print(f\"U-Net预测的.npz文件将保存在: {PREDICTION_OUTPUT_DIR}\") # 这行可以移到最后\n\n    total_slices_predicted_overall = 0\n    total_failed_saves_overall = 0\n    skipped_count = 0\n\n    with open(log_file_path_processed_patients, 'a', encoding='utf-8') as log_file:\n        # 只有在文件为空（第一次创建）或者特定条件下才写入表头和开始标记\n        # 简单起见，每次运行都在日志中追加一个新的运行段落\n        log_file.write(f\"\\n--- Prediction Run Appended/Started: {datetime.datetime.now().isoformat()} ---\\n\")\n        log_file.write(f\"Model Used: {USER_PRETRAINED_MODEL_PATH}\\n\")\n        # 考虑是否每次都写表头，如果文件已存在且有内容，可能不需要重复写\n        # if log_file.tell() == 0: # 如果是文件开头\n        #     log_file.write(\"PatientID,SeriesID,Status,SlicesPredicted,SavesFailed,ErrorMessage\\n\")\n\n        for patient_id, series_dict in tqdm(all_patient_series_dicom_items.items(), desc=\"整体U-Net预测进度\"):\n            for series_id, dicom_items_list_for_series in series_dict.items():\n                current_processing_key = (str(patient_id), str(series_id)) # 使用字符串确保一致性\n\n                # 检查是否已处理并跳过\n                if current_processing_key in already_processed_items:\n                    skip_log_message = f\"{patient_id},{series_id},Skipped (previously processed),0,0,Previously Finished\\n\"\n                    log_file.write(skip_log_message)\n                    skipped_count += 1\n                    continue # 跳到下一个序列\n\n                if not dicom_items_list_for_series:\n                    log_message = f\"{patient_id},{series_id},Skipped - No DICOM items,0,0,\\n\"\n                    log_file.write(log_message)\n                    continue\n                \n                current_patient_data_for_series = {patient_id: dicom_items_list_for_series}\n                processed_slices = 0\n                failed_saves = 0\n                error_message = \"\"\n                status = \"Attempted\" # 初始状态\n\n                try:\n                    # 在实际处理前，可以先记录一个\"Attempted\"状态\n                    # log_file.write(f\"{patient_id},{series_id},{status},0,0,Starting processing\\n\")\n                    # log_file.flush()\n\n                    processed_slices, failed_saves = predict_and_save_masks_optimized(\n                        model=best_model,\n                        patient_ids_to_predict=[patient_id],\n                        image_paths_dict_all_patients=current_patient_data_for_series,\n                        output_dir=PREDICTION_OUTPUT_DIR,\n                        batch_size=INFERENCE_BATCH_SIZE\n                    )\n                    total_slices_predicted_overall += processed_slices\n                    total_failed_saves_overall += failed_saves\n                    status = \"Finished\" # 只有成功完成才标记为 Finished\n                except Exception as e_predict_series:\n                    error_message = str(e_predict_series).replace(',',';').replace('\\n', ' ') # 清理错误信息\n                    status = \"Error\"\n                    print(f\"  对患者 {patient_id} 序列 {series_id} 的预测过程中出错: {error_message}\")\n                \n                log_message = f\"{patient_id},{series_id},{status},{processed_slices},{failed_saves},{error_message}\\n\"\n                log_file.write(log_message)\n                log_file.flush()\n\n        log_file.write(f\"--- Prediction Run Ended: {datetime.datetime.now().isoformat()} ---\\n\")\n        log_file.write(f\"Items Skipped (previously processed): {skipped_count}\\n\")\n        log_file.write(f\"Newly Processed Slices: {total_slices_predicted_overall}\\n\") # 注意这是本次运行新处理的\n        log_file.write(f\"Newly Failed Saves: {total_failed_saves_overall}\\n\\n\")\n            \n    print(f\"\\nU-Net预测完成。\")\n    print(f\"本次运行跳过了 {skipped_count} 个已处理的 病人/序列 组合。\")\n    print(f\"本次运行新处理的切片数: {total_slices_predicted_overall}\")\n    print(f\"本次运行新发生的保存失败次数: {total_failed_saves_overall}\")\n    print(f\"所有预测的.npz文件已（或之前已）保存在: {PREDICTION_OUTPUT_DIR}\")\n    print(f\"预测处理日志已更新/创建于: {log_file_path_processed_patients}\")\n    print(\"--- U-Net .npz 文件生成流程结束 ---\")\n\n\nif __name__ == \"__main__\":\n    main_generate_unet_predictions_for_all()\n\n\ndef 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] = 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#     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    USER_PRETRAINED_MODEL_PATH = \"/kaggle/input/unet_effb0/keras/default/1/unet_effb0_multi_organ_224px_v2.keras\"\n    ordered_class_weights_for_load = [calculated_class_weights.get(organ, 1.0) for organ in ORGAN_CHANNEL_MAP.keys()]\n    final_loss_for_loading = create_focal_dice_loss( # 使用与训练时相同的参数\n        gamma_focal=2.0, alpha_focal=0.25, \n        lambda_focal=0.5, lambda_dice=0.5, \n        class_weights=ordered_class_weights_for_load\n    )\n    \n    custom_objects = {\n        'focal_dice_loss_fn': final_loss_for_loading, # 使用创建函数返回的实际损失函数\n        # 或者，如果损失函数被命名，使用那个名字，并在全局定义它\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        # TensorFlow有时可以直接反序列化函数对象，但更可靠的是传递名称或Loss类\n        # 如果上面的 focal_dice_loss_fn 不能直接被识别，\n        # 你可能需要将 create_focal_dice_loss 返回的函数在全局命名，\n        # 或者将 FocalDiceLoss 实现为一个 tf.keras.losses.Loss 的子类。\n        # 鉴于之前的错误，我们先尝试加载时不编译或只带指标编译。\n        best_model = models.load_model(USER_PRETRAINED_MODEL_PATH, custom_objects=custom_objects, compile=False) # 尝试 compile=False\n        # 如果需要评估，后续再用优化器和损失函数编译一次\n        best_model.compile(optimizer=optimizers.Adam(learning_rate=LEARNING_RATE*0.01), # 用一个小的学习率重新编译\n                           loss=final_loss_for_loading, \n                           metrics=METRICS)\n        print(f\"成功加载并重新编译最终模型: {USER_PRETRAINED_MODEL_PATH}\")\n    except Exception as e:\n        print(f\"加载最终模型 {USER_PRETRAINED_MODEL_PATH} 失败: {e}\")\n        print(\"确保 custom_objects 中的损失函数名称与保存时一致，或尝试仅加载权重。\")\n        # 如果完全失败，后续的推理和评估将无法进行\n        return \n    \n    # --- 9. 使用修改后的方法进行预测和评估 ---\n    print(\"\\n--- 9. 使用修改后的方法进行预测和评估 ---\")\n    \n    # 获取已存在的预测结果患者列表\n    previous_results_dir = '/kaggle/input/rsna-uneted-output/merged_predictions'\n    if os.path.exists(previous_results_dir):\n        existing_patients = set(os.listdir(previous_results_dir))\n        print(f\"已有预测结果中包含 {len(existing_patients)} 个患者\")\n    else:\n        existing_patients = set()\n        print(\"未找到已有预测结果\")\n    \n    # 直接调用新的评估函数，它会处理未处理的患者并评估所有结果\n    dice_scores, iou_scores = modified_evaluation_approach(best_model, image_paths_dict, segmentation_map)\n\n\n# --- 10. 可视化一些分割结果 (修改版，可视化更多样本并从指定路径加载预测) ---\n    print(\"\\n--- 10. 可视化更多分割结果 (使用指定路径的预测) ---\")\n\n    # 用户指定的预测结果路径\n    USER_SPECIFIED_PREDICTION_DIR = \"/kaggle/input/val-unetd/segmentation_predictions_multi_v2\"\n    print(f\"将从以下路径加载预测掩码: {USER_SPECIFIED_PREDICTION_DIR}\")\n    if not os.path.exists(USER_SPECIFIED_PREDICTION_DIR):\n        print(f\"警告: 指定的预测路径 {USER_SPECIFIED_PREDICTION_DIR} 不存在！将无法加载已保存的预测。\")\n        # 如果路径不存在，后续的预测加载会失败，可视化中将只显示原始图和真实掩码（如果存在）\n\n    # 定义可视化时使用的颜色 (BGR格式)\n    colors_for_visualization = {\n        'liver': [0, 0, 255],  # 红色\n        'spleen': [0, 255, 0], # 绿色\n        'kidney': [255, 0, 0], # 蓝色\n        'bowel': [0, 255, 255]   # 黄色 (BGR中的Cyan对应RGB中的Yellow)\n    }\n\n    # 确定用于选择可视化患者的ID列表\n    vis_patient_ids_options = []\n    if 'val_pids' in locals() and val_pids and len(val_pids) > 0: # locals() 检查变量是否存在\n        vis_patient_ids_options = [pid for pid in val_pids if pid in image_paths_dict and pid in segmentation_map]\n        print(f\"从验证集选择患者进行可视化 (共 {len(vis_patient_ids_options)} 位候选)。\")\n    elif 'final_patient_ids' in locals() and final_patient_ids and len(final_patient_ids) > 0:\n        vis_patient_ids_options = [pid for pid in final_patient_ids if pid in image_paths_dict and pid in segmentation_map]\n        print(f\"从所有有效患者中选择进行可视化 (共 {len(vis_patient_ids_options)} 位候选)。\")\n    else:\n        candidate_pids = list(image_paths_dict.keys())\n        vis_patient_ids_options = [pid for pid in candidate_pids if pid in segmentation_map]\n        print(f\"从 image_paths_dict 和 segmentation_map 的交集中选择患者进行可视化 (共 {len(vis_patient_ids_options)} 位候选)。\")\n\n    if not vis_patient_ids_options:\n        print(\"错误：没有可供选择的患者ID进行可视化。请检查 image_paths_dict 和 segmentation_map 是否已正确填充。\")\n    else:\n        num_patients_to_visualize = min(10, len(vis_patient_ids_options)) # 可视化最多10个患者\n        if num_patients_to_visualize == 0 :\n             print(\"没有符合条件的患者可供可视化。\")\n        else:\n            actual_num_to_sample = min(num_patients_to_visualize, len(vis_patient_ids_options))\n            vis_patient_ids_selected = np.random.choice(vis_patient_ids_options, actual_num_to_sample, replace=False)\n            print(f\"将可视化来自以下随机选择的 {len(vis_patient_ids_selected)} 位患者的随机切片: {vis_patient_ids_selected.tolist()}\") # .tolist() 以获得更好的打印输出\n\n            for patient_id_vis in vis_patient_ids_selected:\n                print(f\"\\n--- 正在处理可视化: 患者 {patient_id_vis} ---\")\n                \n                dicom_items_current_patient = image_paths_dict.get(patient_id_vis)\n                nii_path_current_patient = segmentation_map.get(patient_id_vis)\n\n                if not dicom_items_current_patient:\n                    print(f\"  跳过患者 {patient_id_vis}: 在 image_paths_dict 中找不到DICOM信息。\")\n                    continue\n                if not nii_path_current_patient:\n                    print(f\"  跳过患者 {patient_id_vis}: 在 segmentation_map 中找不到NIFTI路径。\")\n                    continue\n                \n                dicom_list_idx_vis = np.random.randint(0, len(dicom_items_current_patient))\n                inst_num_vis, selected_dicom_path_vis = dicom_items_current_patient[dicom_list_idx_vis]\n\n                print(f\"  可视化切片信息: DICOM列表索引 {dicom_list_idx_vis}, InstanceNumber {inst_num_vis}\")\n\n                image_vis_raw, _ = load_dicom_slice(selected_dicom_path_vis)\n                if image_vis_raw is None:\n                    print(f\"  无法加载DICOM图像: {selected_dicom_path_vis}\")\n                    continue\n                \n                gt_slice_int_oriented_raw_res = None\n                nii_slice_idx_for_gt = -1\n                try:\n                    nii_img_vis = nib.load(nii_path_current_patient)\n                    gt_data_vis_float = nii_img_vis.get_fdata(dtype=np.float32)\n                    nii_total_slices_vis = gt_data_vis_float.shape[2]\n\n                    if USE_REVERSE_NIFTI_MAPPING:\n                        nii_slice_idx_for_gt = nii_total_slices_vis - 1 - dicom_list_idx_vis\n                    else:\n                        nii_slice_idx_for_gt = dicom_list_idx_vis\n                    \n                    if not (0 <= nii_slice_idx_for_gt < nii_total_slices_vis):\n                        print(f\"  警告: 为DICOM索引 {dicom_list_idx_vis} 计算的NIFTI索引 {nii_slice_idx_for_gt} 超出范围 ({nii_total_slices_vis}片)。真实掩码将不可用。\")\n                    else:\n                        print(f\"  对应的NIFTI切片索引: {nii_slice_idx_for_gt}\")\n                        gt_slice_raw_from_nii = gt_data_vis_float[:, :, nii_slice_idx_for_gt]\n                        gt_slice_oriented_raw_res = apply_orientation_transform(gt_slice_raw_from_nii, BEST_NIFTI_ORIENTATION_TRANSFORM)\n                        gt_slice_int_oriented_raw_res = np.round(gt_slice_oriented_raw_res).astype(np.int16)\n                except Exception as e_gt:\n                    print(f\"  加载或处理真实NIFTI掩码时出错: {e_gt}\")\n                \n                series_id_vis = os.path.basename(os.path.dirname(selected_dicom_path_vis))\n                pred_mask_npz_path = os.path.join(USER_SPECIFIED_PREDICTION_DIR, str(patient_id_vis), str(series_id_vis), f\"{inst_num_vis}.npz\")\n                pred_mask_binary_vis = None\n\n                if os.path.exists(pred_mask_npz_path):\n                    print(f\"  找到预测文件: {pred_mask_npz_path}\")\n                    try:\n                        pred_data_vis = np.load(pred_mask_npz_path)\n                        pred_mask_from_npz = pred_data_vis['mask']\n                        pred_mask_binary_vis = (pred_mask_from_npz > PREDICTION_THRESHOLD).astype(np.uint8)\n                        if pred_mask_binary_vis.shape[:2] != (TARGET_SIZE, TARGET_SIZE) or pred_mask_binary_vis.ndim != 3 or pred_mask_binary_vis.shape[2] != NUM_ORGANS:\n                             print(f\"  警告: 预测掩码 {os.path.basename(pred_mask_npz_path)} 形状 {pred_mask_binary_vis.shape} 不正确。预期: ({TARGET_SIZE},{TARGET_SIZE},{NUM_ORGANS})。将忽略此预测。\")\n                             pred_mask_binary_vis = None\n                    except Exception as e_load_pred:\n                        print(f\"  加载或处理预测掩码 {pred_mask_npz_path} 失败: {e_load_pred}\")\n                else:\n                    print(f\"  未在指定路径找到预测文件: {pred_mask_npz_path}。\")\n\n                display_image_resized = cv2.resize(image_vis_raw, (TARGET_SIZE, TARGET_SIZE), interpolation=cv2.INTER_LINEAR)\n                if display_image_resized.dtype == np.float32 or display_image_resized.dtype == np.float64:\n                    display_image_for_bgr = (display_image_resized * 255).astype(np.uint8)\n                else:\n                    display_image_for_bgr = display_image_resized\n                display_image_rgb_for_plot = cv2.cvtColor(display_image_for_bgr, cv2.COLOR_GRAY2BGR) # Renamed for clarity\n\n                gt_blended_vis = display_image_rgb_for_plot.copy()\n                if gt_slice_int_oriented_raw_res is not None:\n                    gt_overlay_vis = np.zeros_like(display_image_rgb_for_plot, dtype=np.uint8)\n                    for nii_val_map_key, org_name_map_val in ORGAN_MAP_NII.items(): # Corrected iteration\n                         if org_name_map_val in colors_for_visualization:\n                             color_bgr_val = colors_for_visualization[org_name_map_val]\n                             current_gt_mask_channel_raw = np.zeros(gt_slice_int_oriented_raw_res.shape[:2], dtype=np.uint8)\n                             if org_name_map_val == 'kidney': # kidney is special as it combines NII 3 and 4\n                                 current_gt_mask_channel_raw = ((gt_slice_int_oriented_raw_res == 3) | (gt_slice_int_oriented_raw_res == 4)).astype(np.uint8)\n                             else: # For liver, spleen, bowel, use the nii_val_map_key directly\n                                 current_gt_mask_channel_raw = (gt_slice_int_oriented_raw_res == nii_val_map_key).astype(np.uint8)\n                             \n                             if np.any(current_gt_mask_channel_raw):\n                                 current_gt_mask_channel_resized = cv2.resize(current_gt_mask_channel_raw, (TARGET_SIZE, TARGET_SIZE), interpolation=cv2.INTER_NEAREST)\n                                 for c_idx_rgb in range(3):\n                                    gt_overlay_vis[current_gt_mask_channel_resized > 0, c_idx_rgb] = color_bgr_val[c_idx_rgb]\n                    alpha_blend = 0.4\n                    gt_blended_vis = cv2.addWeighted(display_image_rgb_for_plot, 1 - alpha_blend, gt_overlay_vis, alpha_blend, 0)\n\n                pred_blended_vis = display_image_rgb_for_plot.copy()\n                if pred_mask_binary_vis is not None:\n                    pred_overlay_vis = np.zeros_like(display_image_rgb_for_plot, dtype=np.uint8)\n                    for org_name_pred, channel_idx_pred in ORGAN_CHANNEL_MAP.items():\n                        if org_name_pred in colors_for_visualization and channel_idx_pred < pred_mask_binary_vis.shape[2]:\n                            mask_ch_pred_display = pred_mask_binary_vis[:, :, channel_idx_pred]\n                            color_bgr_pred_val = colors_for_visualization[org_name_pred]\n                            for c_idx_rgb_pred in range(3):\n                                pred_overlay_vis[mask_ch_pred_display > 0, c_idx_rgb_pred] = color_bgr_pred_val[c_idx_rgb_pred]\n                    pred_blended_vis = cv2.addWeighted(display_image_rgb_for_plot, 1 - alpha_blend, pred_overlay_vis, alpha_blend, 0)\n                else:\n                    print(f\"  患者 {patient_id_vis} 切片 {inst_num_vis} 无有效预测掩码用于叠加。\")\n\n                fig, axes = plt.subplots(1, 3, figsize=(20, 7))\n                title_info = f\"患者 {patient_id_vis} - DICOM列表索引 {dicom_list_idx_vis} (Inst: {inst_num_vis})\"\n                if nii_slice_idx_for_gt != -1:\n                    title_info += f\" - NII索引 {nii_slice_idx_for_gt}\"\n                fig.suptitle(title_info, fontsize=14)\n\n                axes[0].imshow(cv2.cvtColor(display_image_rgb_for_plot, cv2.COLOR_BGR2RGB))\n                axes[0].set_title(\"原始DICOM (调整大小)\", fontsize=10)\n                axes[0].axis('off')\n\n                axes[1].imshow(cv2.cvtColor(gt_blended_vis, cv2.COLOR_BGR2RGB))\n                axes[1].set_title(\"真实掩码叠加\", fontsize=10)\n                axes[1].axis('off')\n\n                axes[2].imshow(cv2.cvtColor(pred_blended_vis, cv2.COLOR_BGR2RGB))\n                axes[2].set_title(\"预测掩码叠加 (来自指定路径)\", fontsize=10)\n                axes[2].axis('off')\n                \n                legend_elements = [plt.Rectangle((0, 0), 1, 1, color=[c/255. for c in colors_for_visualization[org][::-1]], label=org)\n                                   for org in ORGAN_CHANNEL_MAP.keys() if org in colors_for_visualization]\n                fig.legend(handles=legend_elements, loc='lower center', ncol=len(ORGAN_CHANNEL_MAP.keys()), bbox_to_anchor=(0.5, 0.01), fontsize=8)\n                \n                plt.tight_layout(rect=[0, 0.05, 1, 0.93])\n                \n                vis_output_filename = f\"vis_pred_gt_{patient_id_vis}_dcm_idx{dicom_list_idx_vis}_inst{inst_num_vis}_nii_idx{nii_slice_idx_for_gt}.png\"\n                vis_output_path = os.path.join(OUTPUT_DIR, vis_output_filename)\n                try:\n                    plt.savefig(vis_output_path)\n                    print(f\"  可视化图像已保存到: {vis_output_path}\")\n                except Exception as e_save_fig:\n                    print(f\"  保存可视化图像失败: {e_save_fig}\")\n                plt.show()\n                \n                gc.collect()\n    \n    print(\"-\" * 30)\n    print(\"可视化部分执行完毕。\")\n    print(\"-\" * 30)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-17T15:25:55.363008Z","iopub.execute_input":"2025-05-17T15:25:55.363621Z","execution_failed":"2025-05-17T15:28:01.18Z"}},"outputs":[],"execution_count":null}]}