{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":52254,"databundleVersionId":9674523,"sourceType":"competition"},{"sourceId":11555525,"sourceType":"datasetVersion","datasetId":7245708},{"sourceId":11679855,"sourceType":"datasetVersion","datasetId":7330592},{"sourceId":11703537,"sourceType":"datasetVersion","datasetId":7341250},{"sourceId":11779768,"sourceType":"datasetVersion","datasetId":7269783}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"# Cell 0: 设置 Matplotlib 中文字体 (Kaggle 环境) \nimport matplotlib\nimport matplotlib.pyplot as plt\nimport matplotlib.font_manager as fm\nimport os\nimport requests # 保留以备未来可能需要下载\nimport zipfile\nimport shutil\n\n# --- 中文字体设置 ---\nprint(\"开始设置 Matplotlib 中文字体 (V5 - 使用上传字体)...\")\n\nfont_file_name = \"SIMHEI.TTF\" # 从你的日志看，似乎是这个文件名\n\nfont_dataset_name = 'simhei'\nfont_file_uploaded_path = f\"/kaggle/input/{font_dataset_name}/{font_file_name}\"\n\n# 目标用户字体目录 (Kaggle中通常可写)\nuser_font_dir = os.path.join(os.path.expanduser('~'), '.local/share/fonts')\nfont_file_dest_path = os.path.join(user_font_dir, font_file_name)\n\nfont_installed_and_setup = False\n\n# 2. 检查上传的字体文件是否存在\nif not os.path.exists(font_file_uploaded_path):\n    print(f\"错误: 未在指定路径找到上传的字体文件: {font_file_uploaded_path}\")\n    print(\"请确认:\")\n    print(f\"  1. 你已经将名为 '{font_file_name}' 的字体文件上传为一个 Kaggle Dataset。\")\n    print(f\"  2. 你已经将该 Dataset ('{font_dataset_name}') 添加到了这个 Notebook 的输入中。\")\n    print(f\"  3. 代码中的 `font_dataset_name` 已设置为 '{font_dataset_name}' (如果不是，请修改)。\")\nelse:\n    print(f\"找到上传的字体文件: {font_file_uploaded_path}\")\n    try:\n        # --- 将字体复制到系统可识别的位置 ---\n        os.makedirs(user_font_dir, exist_ok=True)\n        if not os.path.exists(font_file_dest_path):\n            shutil.copy(font_file_uploaded_path, font_file_dest_path)\n            print(f\"字体已复制到: {font_file_dest_path}\")\n        else:\n            print(f\"字体已存在于目标目录: {font_file_dest_path}\")\n\n        # --- 添加字体到 Matplotlib 管理器 ---\n        # 使用字体文件的完整路径来添加\n        font_entry = fm.FontEntry(fname=font_file_dest_path, name=os.path.splitext(font_file_name)[0]) # 使用文件名作为字体名\n        existing_font_paths = [f.fname for f in fm.fontManager.ttflist]\n        if font_file_dest_path not in existing_font_paths:\n            fm.fontManager.addfont(font_file_dest_path)\n            print(f\"字体 {font_file_dest_path} 已添加到 FontManager\")\n        else:\n             print(f\"字体 {font_file_dest_path} 已在 FontManager 的列表 Ttflist 中。\")\n\n        font_installed_and_setup = True # 标记字体文件已就位\n\n    except Exception as e_copy_add:\n        print(f\"复制或添加字体时出错: {e_copy_add}\")\n\n\n# 3. 如果字体文件在目标位置，则清理缓存并设置rcParams\nif font_installed_and_setup: # 仅当字体文件已复制/存在于目标位置时执行\n    try:\n        # --- 清理缓存 (使用正确的函数 matplotlib.get_cachedir()) ---\n        try:\n            cache_dir = matplotlib.get_cachedir() # <--- 使用 matplotlib.get_cachedir()\n            cache_cleaned = False\n            if os.path.exists(cache_dir):\n                print(f\"尝试清理 Matplotlib 字体缓存目录: {cache_dir}\")\n                for file in os.listdir(cache_dir):\n                    # 匹配更通用的缓存文件名模式\n                    if file.startswith('fontlist') and file.endswith(('.json', '.cache', '.afm', '.pickle')):\n                        try:\n                            os.remove(os.path.join(cache_dir, file))\n                            print(f\"  已删除缓存文件: {file}\")\n                            cache_cleaned = True\n                        except Exception as e_rm:\n                            print(f\"  删除缓存文件 {file} 失败: {e_rm}\")\n                if cache_cleaned:\n                     print(\"Matplotlib 字体缓存文件已清理。可能需要重启 Kernel 使其完全生效。\")\n                else:\n                     print(\"  未找到需要清理的字体缓存文件。\")\n            else:\n                 print(\"未找到 Matplotlib 缓存目录。\")\n        except AttributeError:\n             # 如果连 matplotlib.get_cachedir() 都没有 (极旧版本?)，则跳过\n             print(\"警告: 无法使用 matplotlib.get_cachedir()。跳过缓存清理。\")\n        except Exception as e_cache:\n             print(f\"清理字体缓存时出错: {e_cache}\")\n\n\n        # --- 设置 Matplotlib 参数 ---\n        # 尝试从字体文件获取标准字体名 (如 'SimHei')\n        try:\n            font_prop = fm.FontProperties(fname=font_file_dest_path)\n            font_name = font_prop.get_name()\n            print(f\"从文件推断出的字体名称: {font_name}\")\n        except Exception:\n            # 如果失败，则使用文件名（不含扩展名）作为备选\n            font_name = os.path.splitext(font_file_name)[0]\n            print(f\"无法从文件获取字体名，使用文件名作为名称: {font_name}\")\n\n        # 设置 matplotlib 默认字体\n        plt.rcParams['font.family'] = 'sans-serif' # 设置通用族\n        plt.rcParams['font.sans-serif'] = [font_name, 'sans-serif'] # 将你的字体名加入列表首位\n        plt.rcParams['axes.unicode_minus'] = False # 正确显示负号\n        print(f\"Matplotlib RCParams 已设置为优先使用 '{font_name}' 显示中文。\")\n\n        # 验证一下字体是否被正确识别（可选）\n        # if font_name in fm.findSystemFonts(fontpaths=[user_font_dir]):\n        #      print(f\"验证：字体 '{font_name}' 在管理器中找到。\")\n        # else:\n        #      print(f\"警告：字体 '{font_name}' 可能未被管理器完全识别，如果绘图仍有问题请重启Kernel。\")\n\n\n    except Exception as e_setup:\n        print(f\"设置 Matplotlib 字体参数或清理缓存时出错: {e_setup}\")\n        font_installed_and_setup = False # 标记设置失败\nelse:\n     # 如果前面步骤失败，提醒用户\n     if not os.path.exists(font_file_uploaded_path):\n        print(\"错误：未找到上传的字体文件，无法继续设置。\")\n     else:\n        print(\"字体复制或添加到管理器时失败，中文可能无法正确显示。\")\n\n\n# --- 结束字体设置 ---\nprint(\"-\" * 30)","metadata":{"trusted":true},"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 # 根据需要取消注释\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},"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},"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},"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},"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},"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},"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},"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# --- 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},"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},"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},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 保存并预测掩码","metadata":{}},{"cell_type":"code","source":"def predict_and_save_masks_with_resume(model, patient_ids, image_paths_dict, output_dir, batch_size=32, \n                                      chunk_size=50, save_zips=True):\n    \"\"\"\n    支持断点续传的掩码预测和保存函数，从已解压的目录读取已有结果\n    \n    参数:\n    - chunk_size: 每个输出ZIP文件中包含的患者数量\n    - save_zips: 是否将结果保存为ZIP文件\n    \"\"\"\n    import shutil\n    import time\n    from datetime import datetime\n    \n    processed_slices_total = 0\n    failed_saves = 0\n    current_chunk_patients = 0\n    current_chunk_num = 1\n    processed_patients = []\n    \n    # 检查磁盘空间\n    def check_disk_space(path='/kaggle/working'):\n        import os\n        stat = os.statvfs(path)\n        free_bytes = stat.f_frsize * stat.f_bavail\n        free_gb = free_bytes / (1024 ** 3)\n        return free_gb\n    \n    # 已有预测结果目录\n    previous_results_dir = '/kaggle/input/rsna-uneted-output/segmentation_predictions_multi_v2'\n    has_previous_results = os.path.exists(previous_results_dir)\n    \n    # 检查是否有已完成患者记录\n    completed_patients_file = '/kaggle/working/completed_patients.txt'\n    completed_patients = set()\n    \n    # 从已有预测结果目录中识别已处理的患者\n    if has_previous_results:\n        print(f\"发现已有预测结果目录: {previous_results_dir}\")\n        \n        # 获取目录中的所有患者ID\n        try:\n            patient_dirs = [d for d in os.listdir(previous_results_dir) \n                          if os.path.isdir(os.path.join(previous_results_dir, d))]\n            \n            # 更新已完成患者集合\n            for patient_id in patient_dirs:\n                completed_patients.add(patient_id)\n            \n            print(f\"从已有结果目录中识别出 {len(completed_patients)} 个已处理的患者\")\n            \n            # 保存已完成患者列表\n            with open(completed_patients_file, 'w') as f:\n                for patient_id in sorted(completed_patients):\n                    f.write(f\"{patient_id}\\n\")\n        except Exception as e:\n            print(f\"读取已有结果目录失败: {e}\")\n            has_previous_results = False\n    \n    # 创建时间戳，用于命名输出文件\n    timestamp = datetime.now().strftime(\"%Y%m%d_%H%M%S\")\n    \n    # 处理每个患者\n    for i, patient_id in enumerate(tqdm(patient_ids, desc=\"处理患者\")):\n        # 跳过已完成的患者\n        if patient_id in completed_patients:\n            print(f\"跳过已完成的患者: {patient_id}\")\n            continue\n            \n        if patient_id not in image_paths_dict:\n            continue\n            \n        dicom_paths = image_paths_dict[patient_id]\n        if not dicom_paths:\n            continue\n        \n        # 检查剩余磁盘空间\n        free_space_gb = check_disk_space()\n        if free_space_gb < 1.0:  # 如果剩余空间小于1GB\n            print(f\"警告: 磁盘空间不足 ({free_space_gb:.2f}GB)，准备保存当前进度并清理...\")\n            \n            if save_zips and processed_patients:\n                # 保存当前批次的结果\n                chunk_zip_path = f'/kaggle/working/predictions_chunk_{current_chunk_num}_{timestamp}.zip'\n                print(f\"保存当前批次结果到: {chunk_zip_path}\")\n                \n                # 创建ZIP文件\n                import zipfile\n                with zipfile.ZipFile(chunk_zip_path, 'w', zipfile.ZIP_DEFLATED) as zipf:\n                    for pid in processed_patients:\n                        patient_dir = os.path.join(output_dir, pid)\n                        if os.path.exists(patient_dir):\n                            # 添加患者目录到ZIP\n                            for root, _, files in os.walk(patient_dir):\n                                for file in files:\n                                    file_path = os.path.join(root, file)\n                                    arc_name = os.path.relpath(file_path, output_dir)\n                                    zipf.write(file_path, arc_name)\n                \n                # 更新已完成患者记录\n                with open(completed_patients_file, 'a') as f:\n                    for pid in processed_patients:\n                        f.write(f\"{pid}\\n\")\n                        completed_patients.add(pid)\n                \n                # 清理已保存的患者目录\n                for pid in processed_patients:\n                    patient_dir = os.path.join(output_dir, pid)\n                    if os.path.exists(patient_dir):\n                        shutil.rmtree(patient_dir)\n                \n                # 重置处理状态\n                processed_patients = []\n                current_chunk_patients = 0\n                current_chunk_num += 1\n                \n                print(f\"已保存并清理空间，继续处理...\")\n            else:\n                print(\"警告: 磁盘空间不足，但没有启用ZIP保存或没有已处理患者，继续尝试...\")\n            \n        # 获取该患者的series_id (假设所有图像来自同一个系列)\n        series_id = dicom_paths[0].split('/')[-2]\n        \n        # 创建输出目录\n        patient_output_dir = os.path.join(output_dir, patient_id, series_id)\n        os.makedirs(patient_output_dir, exist_ok=True)\n        \n        # 如果有已有结果，从已有目录中复制该患者的掩码\n        if has_previous_results:\n            prev_patient_dir = os.path.join(previous_results_dir, patient_id)\n            if os.path.exists(prev_patient_dir):\n                prev_series_dirs = [d for d in os.listdir(prev_patient_dir) \n                                  if os.path.isdir(os.path.join(prev_patient_dir, d))]\n                \n                if prev_series_dirs:\n                    prev_series_dir = os.path.join(prev_patient_dir, prev_series_dirs[0])  # 使用第一个系列\n                    mask_files = glob.glob(os.path.join(prev_series_dir, \"*.npz\"))\n                    \n                    if mask_files:\n                        print(f\"在已有结果目录中找到患者 {patient_id} 的 {len(mask_files)} 个掩码文件\")\n                        \n                        # 如果掩码文件数量与DICOM文件数量相同，则认为该患者已完全处理\n                        if len(mask_files) >= len(dicom_paths):\n                            print(f\"患者 {patient_id} 在已有结果中已完全处理，标记为完成并跳过\")\n                            completed_patients.add(patient_id)\n                            with open(completed_patients_file, 'a') as f:\n                                f.write(f\"{patient_id}\\n\")\n                            continue\n                        \n                        # 复制已有的掩码文件\n                        for mask_file in mask_files:\n                            mask_filename = os.path.basename(mask_file)\n                            dest_path = os.path.join(patient_output_dir, mask_filename)\n                            \n                            if not os.path.exists(dest_path):\n                                try:\n                                    shutil.copy2(mask_file, dest_path)\n                                    processed_slices_total += 1\n                                except Exception as e:\n                                    print(f\"复制掩码文件失败 {mask_file} -> {dest_path}: {e}\")\n        \n        # 收集需要处理的切片\n        slices_to_process = []\n        instance_numbers = []\n        \n        # 检查当前输出目录中已有的掩码\n        existing_masks = glob.glob(os.path.join(patient_output_dir, \"*.npz\"))\n        existing_instances = {os.path.splitext(os.path.basename(m))[0] for m in existing_masks}\n        \n        for dicom_path in dicom_paths:\n            # 加载DICOM并获取实例号\n            image, instance_num = load_dicom_slice(dicom_path)\n            if image is None:\n                continue\n                \n            # 使用实例号作为文件名\n            instance_number = instance_num if instance_num is not None else dicom_paths.index(dicom_path) + 1\n            instance_number_str = str(instance_number)\n            \n            # 检查该切片是否已处理\n            if instance_number_str in existing_instances:\n                continue\n            \n            # 预处理图像\n            processed_image = preprocess_image_for_unet(image, TARGET_SIZE)\n            slices_to_process.append(processed_image)\n            instance_numbers.append(instance_number)\n            \n        if not slices_to_process:\n            print(f\"患者 {patient_id} 的所有切片已处理完成\")\n            # 标记该患者为已完成\n            completed_patients.add(patient_id)\n            with open(completed_patients_file, 'a') as f:\n                f.write(f\"{patient_id}\\n\")\n            continue\n            \n        print(f\"处理患者 {patient_id}: {len(slices_to_process)}/{len(dicom_paths)} 个新切片\")\n        \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 = model.predict(batch_images, verbose=0)\n                \n                # 保存预测结果\n                for j, pred in enumerate(batch_predictions):\n                    idx = i + j\n                    if idx >= len(instance_numbers):\n                        break\n                        \n                    instance_number = instance_numbers[idx]\n                    output_path = os.path.join(patient_output_dir, f\"{instance_number}.npz\")\n                    \n                    # 转换为二值掩码\n                    pred_mask_binary = (pred > PREDICTION_THRESHOLD).astype(np.uint8)\n                    \n                    try:\n                        np.savez_compressed(output_path, mask=pred_mask_binary)\n                        processed_slices_total += 1\n                    except Exception as e:\n                        print(f\"保存掩码失败 {output_path}: {e}\")\n                        failed_saves += 1\n                        \n                # 释放内存\n                del batch_images\n                del batch_predictions\n                gc.collect()\n                \n            # 该患者处理完毕，添加到处理列表\n            processed_patients.append(patient_id)\n            current_chunk_patients += 1\n            \n            # 如果达到块大小，保存并清理\n            if save_zips and current_chunk_patients >= chunk_size:\n                import zipfile\n                chunk_zip_path = f'/kaggle/working/predictions_chunk_{current_chunk_num}_{timestamp}.zip'\n                print(f\"保存当前批次结果到: {chunk_zip_path}\")\n                \n                # 创建ZIP文件\n                with zipfile.ZipFile(chunk_zip_path, 'w', zipfile.ZIP_DEFLATED) as zipf:\n                    for pid in processed_patients:\n                        patient_dir = os.path.join(output_dir, pid)\n                        if os.path.exists(patient_dir):\n                            # 添加患者目录到ZIP\n                            for root, _, files in os.walk(patient_dir):\n                                for file in files:\n                                    file_path = os.path.join(root, file)\n                                    arc_name = os.path.relpath(file_path, output_dir)\n                                    zipf.write(file_path, arc_name)\n                \n                # 更新已完成患者记录\n                with open(completed_patients_file, 'a') as f:\n                    for pid in processed_patients:\n                        f.write(f\"{pid}\\n\")\n                        completed_patients.add(pid)\n                \n                # 清理已保存的患者目录\n                for pid in processed_patients:\n                    patient_dir = os.path.join(output_dir, pid)\n                    if os.path.exists(patient_dir):\n                        shutil.rmtree(patient_dir)\n                \n                # 重置处理状态\n                processed_patients = []\n                current_chunk_patients = 0\n                current_chunk_num += 1\n                \n        except Exception as e:\n            print(f\"批量预测失败: {e}\")\n            # 尝试保存当前进度\n            if processed_patients:\n                try:\n                    import zipfile\n                    chunk_zip_path = f'/kaggle/working/predictions_partial_{current_chunk_num}_{timestamp}.zip'\n                    print(f\"保存当前进度到: {chunk_zip_path}\")\n                    \n                    with zipfile.ZipFile(chunk_zip_path, 'w', zipfile.ZIP_DEFLATED) as zipf:\n                        for pid in processed_patients:\n                            patient_dir = os.path.join(output_dir, pid)\n                            if os.path.exists(patient_dir):\n                                for root, _, files in os.walk(patient_dir):\n                                    for file in files:\n                                        file_path = os.path.join(root, file)\n                                        arc_name = os.path.relpath(file_path, output_dir)\n                                        zipf.write(file_path, arc_name)\n                    \n                    # 更新已完成患者记录\n                    with open(completed_patients_file, 'a') as f:\n                        for pid in processed_patients:\n                            if pid != patient_id:  # 不包括当前失败的患者\n                                f.write(f\"{pid}\\n\")\n                                completed_patients.add(pid)\n                except Exception as save_err:\n                    print(f\"保存进度失败: {save_err}\")\n    \n    # 保存最后一批结果\n    if save_zips and processed_patients:\n        import zipfile\n        final_zip_path = f'/kaggle/working/predictions_final_{timestamp}.zip'\n        print(f\"保存最终批次结果到: {final_zip_path}\")\n        \n        with zipfile.ZipFile(final_zip_path, 'w', zipfile.ZIP_DEFLATED) as zipf:\n            for pid in processed_patients:\n                patient_dir = os.path.join(output_dir, pid)\n                if os.path.exists(patient_dir):\n                    for root, _, files in os.walk(patient_dir):\n                        for file in files:\n                            file_path = os.path.join(root, file)\n                            arc_name = os.path.relpath(file_path, output_dir)\n                            zipf.write(file_path, arc_name)\n        \n        # 更新已完成患者记录\n        with open(completed_patients_file, 'a') as f:\n            for pid in processed_patients:\n                f.write(f\"{pid}\\n\")\n    \n    return processed_slices_total, failed_saves\n\n\ndef prepare_for_evaluation(prediction_dir, previous_results_dir=None, zip_files=None):\n    \"\"\"\n    准备评估环境，合并所有预测结果\n    \n    参数:\n    - prediction_dir: 当前预测结果目录\n    - previous_results_dir: 已有预测结果目录\n    - zip_files: 包含预测结果的ZIP文件列表\n    \n    返回:\n    - merged_dir: 合并后的预测结果目录\n    \"\"\"\n    import shutil\n    import tempfile\n    import zipfile\n    \n    print(\"准备评估环境，合并所有预测结果...\")\n    \n    # 创建合并目录\n    merged_dir = os.path.join('/kaggle/working', 'merged_predictions')\n    os.makedirs(merged_dir, exist_ok=True)\n    \n    # 1. 首先复制当前预测目录中的结果\n    if os.path.exists(prediction_dir):\n        print(f\"复制当前预测结果从 {prediction_dir}\")\n        for patient_id in os.listdir(prediction_dir):\n            patient_src = os.path.join(prediction_dir, patient_id)\n            patient_dst = os.path.join(merged_dir, patient_id)\n            \n            if os.path.isdir(patient_src) and not os.path.exists(patient_dst):\n                shutil.copytree(patient_src, patient_dst)\n    \n    # 2. 复制已有预测结果目录中的内容\n    if previous_results_dir and os.path.exists(previous_results_dir):\n        print(f\"复制已有预测结果从 {previous_results_dir}\")\n        for patient_id in os.listdir(previous_results_dir):\n            patient_src = os.path.join(previous_results_dir, patient_id)\n            patient_dst = os.path.join(merged_dir, patient_id)\n            \n            # 如果目标中已有该患者，则跳过\n            if os.path.isdir(patient_src) and not os.path.exists(patient_dst):\n                shutil.copytree(patient_src, patient_dst)\n    \n    # 3. 处理ZIP文件中的内容\n    if zip_files:\n        for zip_path in zip_files:\n            if not os.path.exists(zip_path):\n                continue\n                \n            print(f\"从ZIP文件中提取预测结果: {zip_path}\")\n            try:\n                with zipfile.ZipFile(zip_path, 'r') as zipf:\n                    # 获取ZIP中的患者目录\n                    all_files = zipf.namelist()\n                    patient_dirs = set()\n                    \n                    for file_path in all_files:\n                        parts = file_path.split('/')\n                        if len(parts) >= 1:\n                            patient_dirs.add(parts[0])\n                    \n                    # 提取每个患者的文件\n                    for patient_id in patient_dirs:\n                        # 如果合并目录中已有该患者，则跳过\n                        patient_dst = os.path.join(merged_dir, patient_id)\n                        if os.path.exists(patient_dst):\n                            continue\n                            \n                        # 创建临时目录提取文件\n                        with tempfile.TemporaryDirectory() as temp_dir:\n                            # 提取该患者的所有文件\n                            for file_info in zipf.infolist():\n                                if file_info.filename.startswith(f\"{patient_id}/\"):\n                                    zipf.extract(file_info, temp_dir)\n                            \n                            # 复制到合并目录\n                            patient_temp = os.path.join(temp_dir, patient_id)\n                            if os.path.exists(patient_temp):\n                                shutil.copytree(patient_temp, patient_dst)\n            except Exception as e:\n                print(f\"处理ZIP文件 {zip_path} 失败: {e}\")\n    \n    # 统计合并后的患者数量\n    patient_count = len([d for d in os.listdir(merged_dir) if os.path.isdir(os.path.join(merged_dir, d))])\n    print(f\"合并完成，共有 {patient_count} 个患者的预测结果\")\n    \n    return merged_dir\n\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def predict_and_save_masks_optimized(model, patient_ids, image_paths_dict, output_dir, batch_size=32):\n    \"\"\"\n    优化版的掩码预测和保存函数，不包含恢复逻辑，专注于处理新患者\n    \n    参数:\n    - model: 训练好的模型\n    - patient_ids: 需要处理的患者ID列表\n    - image_paths_dict: 患者ID到DICOM路径列表的映射\n    - output_dir: 输出目录\n    - batch_size: 批处理大小\n    \n    返回:\n    - processed_slices_total: 处理的切片总数\n    - failed_saves: 保存失败的次数\n    \"\"\"\n    processed_slices_total = 0\n    failed_saves = 0\n    \n    for patient_id in tqdm(patient_ids, desc=\"处理患者\"):\n        if patient_id not in image_paths_dict:\n            continue\n            \n        dicom_paths = image_paths_dict[patient_id]\n        if not dicom_paths:\n            continue\n            \n        # 获取该患者的series_id (假设所有图像来自同一个系列)\n        series_id = dicom_paths[0].split('/')[-2]\n        \n        # 创建输出目录\n        patient_output_dir = os.path.join(output_dir, patient_id, series_id)\n        os.makedirs(patient_output_dir, exist_ok=True)\n        \n        # 收集需要处理的切片\n        slices_to_process = []\n        instance_numbers = []\n        \n        for dicom_path in dicom_paths:\n            # 加载DICOM并获取实例号\n            image, instance_num = load_dicom_slice(dicom_path)\n            if image is None:\n                continue\n                \n            # 使用实例号作为文件名\n            instance_number = instance_num if instance_num is not None else dicom_paths.index(dicom_path) + 1\n            \n            # 预处理图像\n            processed_image = preprocess_image_for_unet(image, TARGET_SIZE)\n            slices_to_process.append(processed_image)\n            instance_numbers.append(instance_number)\n            \n        if not slices_to_process:\n            continue\n            \n        print(f\"处理患者 {patient_id}: {len(slices_to_process)} 个切片\")\n        \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 = model.predict(batch_images, verbose=0)\n                \n                # 保存预测结果\n                for j, pred in enumerate(batch_predictions):\n                    idx = i + j\n                    if idx >= len(instance_numbers):\n                        break\n                        \n                    instance_number = instance_numbers[idx]\n                    output_path = os.path.join(patient_output_dir, f\"{instance_number}.npz\")\n                    \n                    # 转换为二值掩码\n                    pred_mask_binary = (pred > PREDICTION_THRESHOLD).astype(np.uint8)\n                    \n                    try:\n                        np.savez_compressed(output_path, mask=pred_mask_binary)\n                        processed_slices_total += 1\n                    except Exception as e:\n                        print(f\"保存掩码失败 {output_path}: {e}\")\n                        failed_saves += 1\n                        \n                # 释放内存\n                del batch_images\n                del batch_predictions\n                gc.collect()\n                \n        except Exception as e:\n            print(f\"批量预测失败: {e}\")\n    \n    return processed_slices_total, failed_saves\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def modified_evaluation_approach(best_model, image_paths_dict, segmentation_map):\n    \"\"\"不合并结果，直接评估已有预测和新生成的预测\"\"\"\n    \n    # 1. 识别已有预测结果中的患者\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    # 2. 确定需要处理的患者\n    all_patient_ids = sorted([pid for pid in os.listdir(TRAIN_IMAGES_DIR) \n                             if os.path.isdir(os.path.join(TRAIN_IMAGES_DIR, pid))])\n    \n    patients_to_process = [pid for pid in all_patient_ids if pid not in existing_patients]\n    print(f\"需要处理的患者数量: {len(patients_to_process)}\")\n    \n    # 3. 处理剩余患者\n    if patients_to_process:\n        print(\"\\n--- 处理剩余患者 ---\")\n        processed_slices, failed_saves = predict_and_save_masks_optimized(\n            best_model, patients_to_process, image_paths_dict, \n            PREDICTION_OUTPUT_DIR, INFERENCE_BATCH_SIZE\n        )\n        print(f\"新处理的切片数量: {processed_slices}\")\n        print(f\"保存失败的切片数量: {failed_saves}\")\n    \n    # 4. 直接评估结果 (不复制或合并文件)\n    print(\"\\n--- 评估分割结果 ---\")\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    processed_patients = 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            if patient_id in existing_patients:\n                pred_patient_dir = os.path.join(previous_results_dir, patient_id)\n            else:\n                pred_patient_dir = os.path.join(PREDICTION_OUTPUT_DIR, patient_id)\n                \n            if not os.path.exists(pred_patient_dir):\n                continue\n                \n            # 获取该患者的系列目录\n            series_dirs = [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            # 使用第一个系列\n            series_pred_dir = os.path.join(pred_patient_dir, series_dirs[0])\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_patients += 1\n            except Exception as e:\n                print(f\"处理评估结果时出错: {e}\")\n    \n    # 输出评估结果\n    print(f\"\\n评估完成，成功评估 {processed_patients} 个患者。\")\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    try:\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    except Exception as e:\n        print(f\"绘图失败: {e}\")\n    \n    return dice_scores, iou_scores\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 主函数&可视化","metadata":{}},{"cell_type":"code","source":"def main():\n    \"\"\"主程序流程 (优化版)\"\"\"\n    print(\"--- 1. 构建文件映射 ---\")\n    # 加载系列元数据\n    train_meta_path = os.path.join(DATA_DIR, 'train_series_meta.csv')\n    if not os.path.exists(train_meta_path):\n        raise FileNotFoundError(f\"找不到系列元数据文件: {train_meta_path}\")\n    train_series_meta = pd.read_csv(train_meta_path)\n\n    # 创建 series_id -> patient_id 映射\n    series_to_patient = dict(zip(\n        train_series_meta['series_id'].astype(str),\n        train_series_meta['patient_id'].astype(str)\n    ))\n\n    # 创建 NII 分割文件映射 (series_id -> nii_path)\n    series_to_nii = {}\n    segmentation_files = glob.glob(os.path.join(SEGMENTATION_DIR, \"*.nii\"))\n    print(f\"发现 {len(segmentation_files)} 个 NII 文件。\")\n    for fpath in segmentation_files:\n        series_id = os.path.splitext(os.path.basename(fpath))[0]\n        if series_id in series_to_patient: # 确保这个系列在元数据中\n            series_to_nii[series_id] = fpath\n\n    print(f\"成功映射 {len(series_to_nii)} 个 NII 文件到 series_id。\")\n\n    # 加载 DICOM tags (如果可用)\n    dicom_tags_df = None\n    dicom_tags_path = os.path.join(DATA_DIR, 'train_dicom_tags.parquet')\n    if os.path.exists(dicom_tags_path):\n        print(\"加载 DICOM tags...\")\n        try:\n            dicom_tags_df = pd.read_parquet(dicom_tags_path)\n            # 预处理 tags DataFrame\n            if 'PatientID' in dicom_tags_df.columns:\n                dicom_tags_df['PatientID'] = dicom_tags_df['PatientID'].astype(str)\n            # 尝试提取 series_id\n            if 'SeriesInstanceUID' in dicom_tags_df.columns:\n                 dicom_tags_df['series_id_extracted'] = dicom_tags_df['SeriesInstanceUID'].str.split('.').str[-2]\n                 dicom_tags_df = dicom_tags_df.dropna(subset=['series_id_extracted'])\n            print(\"DICOM tags 加载完成。\")\n        except Exception as e:\n            print(f\"加载 DICOM tags 失败: {e}. 将不使用 tags 进行排序。\")\n            dicom_tags_df = None\n    else:\n        print(\"未找到 DICOM tags 文件，将仅依赖DICOM头或文件名排序。\")\n\n    # --- 关联图像和分割 ---\n    image_paths_dict = {}  # {patient_id: [sorted_list_of_dicom_paths]}\n    segmentation_map = {}  # {patient_id: nii_path}\n    \n    patient_ids_in_train_images = os.listdir(TRAIN_IMAGES_DIR)\n    valid_patients = []  # 存储有图像和对应NII文件的患者ID\n\n    print(\"开始关联患者图像和分割文件...\")\n    for patient_id in tqdm(patient_ids_in_train_images, desc=\"处理患者\"):\n        if not os.path.isdir(os.path.join(TRAIN_IMAGES_DIR, patient_id)):\n            continue\n\n        patient_series_ids = [s for s, p in series_to_patient.items() if p == patient_id]\n        if not patient_series_ids: \n            continue\n\n        # 查找该患者是否有系列同时存在于图像目录和NII分割中\n        found_valid_series = False\n        for series_id in patient_series_ids:\n            series_img_path = os.path.join(TRAIN_IMAGES_DIR, patient_id, series_id)\n            nii_path = series_to_nii.get(series_id)\n\n            if os.path.isdir(series_img_path) and nii_path:\n                # 获取并排序DICOM文件\n                dicom_info = get_dicom_files_dict(patient_id, series_id, dicom_tags_df)\n                if dicom_info:  # 确保系列中有有效的DICOM文件\n                    image_paths_dict[patient_id] = [item[1] for item in dicom_info]  # 只存路径\n                    segmentation_map[patient_id] = nii_path\n                    valid_patients.append(patient_id)\n                    found_valid_series = True\n                    break  # 每个患者只使用一个有效的 series 和 NII\n\n    final_patient_ids = sorted(list(set(valid_patients)))  # 去重并排序\n    print(f\"成功映射了 {len(final_patient_ids)} 位患者的图像和分割文件。\")\n    if not final_patient_ids:\n        raise SystemExit(\"错误：未能找到任何包含有效图像序列及对应NII文件的患者。\")\n        \n    print(\"\\n--- 2. 验证NII文件标签 ---\")\n    # 使用修正后的器官映射进行验证\n    verify_nii_labels(segmentation_map, {1: 'liver', 2: 'spleen', 3: 'kidney_left', 4: 'kidney_right', 5: 'bowel'})\n\n    print(\"\\n--- 3. 分析器官分布 ---\")\n    # 使用模型将要使用的器官映射进行分析 (合并肾脏)\n    distribution_results = analyze_organ_distribution(segmentation_map, ORGAN_MAP_NII)\n    calculated_class_weights = distribution_results['class_weights']  # 保存计算出的权重\n\n    print(\"\\n--- 4. 划分训练集和验证集 ---\")\n    train_pids, val_pids = train_test_split(final_patient_ids, test_size=VALIDATION_SPLIT, random_state=RANDOM_STATE)\n    print(f\"训练集患者数: {len(train_pids)}\")\n    print(f\"验证集患者数: {len(val_pids)}\")\n\n    print(\"\\n--- 5. 使用已有预处理数据 ---\")\n    # 检查预处理数据目录\n    preprocessed_patients = [d for d in os.listdir(PREPROCESSED_DIR) \n                          if os.path.isdir(os.path.join(PREPROCESSED_DIR, d))]\n    print(f\"找到 {len(preprocessed_patients)} 个已预处理的患者数据\")\n    \n    # 添加这一行，在训练前检查预处理数据的对齐情况\n    print(\"\\n--- 5b. 检查预处理数据对齐 ---\")\n    visualize_preprocessed_samples(PREPROCESSED_DIR, num_patients=3, samples_per_patient=2)\n        \n    # 确认预处理的患者包含了训练和验证集\n    train_pids_preprocessed = [pid for pid in train_pids if pid in preprocessed_patients]\n    val_pids_preprocessed = [pid for pid in val_pids if pid in preprocessed_patients]\n    \n    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    plotplt(epochs, final_history['val_dice_liver'], label='肝脏 (Val)')\n    plt.plot(epochs, final_history['val_dice_spleen'], label='脾脏 (Val)')\n    plt.plot(epochs, final_historyval['_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    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(MODEL_SAVE_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\"成功加载并重新编译最终模型: {MODEL_SAVE_PATH}\")\n    except Exception as e:\n        print(f\"加载最终模型 {MODEL_SAVE_PATH} 失败: {e}\")\n        print(\"确保 custom_objects 中的损失函数名称与保存时一致，或尝试仅加载权重。\")\n        # 如果完全失败，后续的推理和评估将无法进行\n        return \n\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    # 定义可视化时使用的颜色 (BGR格式，因为OpenCV常用，但显示时会转RGB)\n    # 您可以根据 ORGAN_CHANNEL_MAP 的顺序调整或确保颜色字典覆盖所有需要的器官\n    colors_for_visualization = { \n        'liver': [0, 0, 255],  # 红色 for liver\n        'spleen': [0, 255, 0], # 绿色 for spleen\n        'kidney': [255, 0, 0], # 蓝色 for kidney\n        'bowel': [0, 255, 255]   # 黄色 for bowel (Cyan in BGR for Yellow in RGB)\n    }\n\n    # 从验证集或所有有效患者中随机选择几个进行可视化\n    if val_pids and len(val_pids) > 0: # 优先使用验证集\n        vis_patient_ids_options = val_pids\n    elif 'final_patient_ids' in locals() and len(final_patient_ids) > 0:\n        vis_patient_ids_options = final_patient_ids\n    else:\n        vis_patient_ids_options = list(image_paths_dict.keys()) # 如果都没有，则从所有有路径的患者中选\n\n    if not vis_patient_ids_options:\n        print(\"没有可供可视化的患者ID。\")\n    else:\n        num_patients_to_visualize = min(3, len(vis_patient_ids_options)) # 最多可视化3个患者\n        vis_patient_ids = np.random.choice(vis_patient_ids_options, num_patients_to_visualize, replace=False)\n        print(f\"将可视化以下患者的随机切片: {vis_patient_ids}\")\n\n        for patient_id_vis in vis_patient_ids: # 使用新的变量名避免冲突\n            if patient_id_vis not in image_paths_dict or patient_id_vis not in segmentation_map:\n                print(f\"跳过患者 {patient_id_vis}: 缺少图像路径或分割图信息。\")\n                continue\n\n            dicom_paths_current_patient = image_paths_dict[patient_id_vis] # 这是 (instance_num, path) 的列表\n            nii_path_current_patient = segmentation_map[patient_id_vis]\n\n            if not dicom_paths_current_patient:\n                print(f\"患者 {patient_id_vis} 的DICOM路径列表为空。\")\n                continue\n            \n            # 随机选择一个DICOM切片进行可视化\n            # dicom_list_idx 是选中的DICOM在其排序列表中的索引\n            dicom_list_idx_vis = np.random.randint(0, len(dicom_paths_current_patient))\n            # 从 dicom_info_list 获取 instance_number 和 path\n            # 假设 get_dicom_files_dict 返回的是 [(instance_number, path), ...]\n            # 而 image_paths_dict[patient_id] 存储的是 [path, ...]\n            # 需要统一：假设 image_paths_dict[patient_id] 就是 get_dicom_files_dict 返回的完整元组列表\n            # 如果不是，需要调整 image_paths_dict 的构建方式，或者在这里重新获取instance_number\n            \n            # 假设 image_paths_dict[patient_id] 是路径列表，我们需要重新获取 InstanceNumber\n            # 或者，如果您的 image_paths_dict 存储的是 (instance, path) 元组，则可以直接用\n            # 为了代码的鲁棒性，我们重新读取一下InstanceNumber\n            \n            selected_dicom_path_vis = \"\"\n            # 检查 image_paths_dict[patient_id_vis] 中元素的类型\n            if isinstance(dicom_paths_current_patient[dicom_list_idx_vis], tuple): # 如果是 (inst, path)\n                inst_num_vis, selected_dicom_path_vis = dicom_paths_current_patient[dicom_list_idx_vis]\n            else: # 如果只是路径列表\n                selected_dicom_path_vis = dicom_paths_current_patient[dicom_list_idx_vis]\n                # 尝试从文件名或DICOM头读取InstanceNumber\n                try:\n                    temp_ds = pydicom.dcmread(selected_dicom_path_vis, stop_before_pixels=True)\n                    inst_num_vis = int(temp_ds.InstanceNumber)\n                except:\n                    inst_num_vis = dicom_list_idx_vis + 1 # 备用方案\n                    print(f\"警告: 无法从 {selected_dicom_path_vis} 读取InstanceNumber，使用索引 {dicom_list_idx_vis} 作为替代ID。\")\n\n\n            print(f\"\\n可视化患者 {patient_id_vis}, DICOM列表索引 {dicom_list_idx_vis}, InstanceNumber {inst_num_vis}\")\n\n            image_vis_raw, _ = load_dicom_slice(selected_dicom_path_vis) # load_dicom_slice返回归一化图像和instance_number\n            if image_vis_raw is None: \n                print(f\"  无法加载DICOM: {selected_dicom_path_vis}\")\n                continue\n\n            # --- 加载并处理真实NIFTI掩码 ---\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                nii_slice_idx_for_gt = -1\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                    continue\n\n                gt_slice_raw_from_nii = gt_data_vis_float[:, :, nii_slice_idx_for_gt]\n                \n                # 应用方向变换 (在原始分辨率下进行)\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                \n            except Exception as e:\n                print(f\"  加载或处理真实NIFTI掩码时出错 (患者 {patient_id_vis}, NII: {nii_path_current_patient}): {e}\")\n                continue\n            \n            # --- 获取预测掩码 ---\n            # 假设您的 image_paths_dict[patient_id] 中的路径是 /kaggle/input/.../patient_id/series_id/instance.dcm 格式\n            series_id_vis = os.path.basename(os.path.dirname(selected_dicom_path_vis))\n            \n            pred_mask_npz_path = os.path.join(PREDICTION_OUTPUT_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                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 != (TARGET_SIZE, TARGET_SIZE, NUM_ORGANS):\n                         print(f\"  警告: 预测掩码 {pred_mask_npz_path} 形状 {pred_mask_binary_vis.shape} 不正确，应为 {(TARGET_SIZE, TARGET_SIZE, NUM_ORGANS)}\")\n                         pred_mask_binary_vis = None # 设为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                if 'best_model' in locals() and best_model is not None:\n                    try:\n                        processed_image_for_pred = preprocess_image_for_unet(image_vis_raw, TARGET_SIZE)\n                        prediction_prob_live = best_model.predict(np.expand_dims(processed_image_for_pred, axis=0), verbose=0)[0]\n                        pred_mask_binary_vis = (prediction_prob_live > PREDICTION_THRESHOLD).astype(np.uint8)\n                    except Exception as live_pred_e:\n                        print(f\"    实时预测失败: {live_pred_e}\")\n                else:\n                    print(\"    best_model 未定义或为None，无法进行实时预测。\")\n\n            # --- 创建彩色叠加图 ---\n            display_image_resized = cv2.resize(image_vis_raw, (TARGET_SIZE, TARGET_SIZE), interpolation=cv2.INTER_LINEAR)\n            display_image_rgb = cv2.cvtColor((display_image_resized * 255).astype(np.uint8), cv2.COLOR_GRAY2BGR)\n\n            # 真实掩码叠加\n            gt_overlay_vis = np.zeros_like(display_image_rgb, dtype=np.uint8)\n            for nii_val_map, org_name_map in ORGAN_MAP_NII.items(): # 使用全局变量\n                 if org_name_map in colors_for_visualization: # 使用新定义的颜色字典\n                     color_bgr_val = colors_for_visualization[org_name_map]\n                     current_gt_mask_channel_raw = np.zeros(gt_slice_int_oriented_raw_res.shape, dtype=np.uint8)\n                     if org_name_map == 'kidney':\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:\n                         current_gt_mask_channel_raw = (gt_slice_int_oriented_raw_res == nii_val_map).astype(np.uint8)\n                     \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            \n            alpha_blend = 0.4 # 透明度\n            gt_blended_vis = cv2.addWeighted(display_image_rgb, 1 - alpha_blend, gt_overlay_vis, alpha_blend, 0)\n\n            # 预测掩码叠加\n            pred_blended_vis = display_image_rgb.copy() # 默认为原始图像\n            if pred_mask_binary_vis is not None:\n                pred_overlay_vis = np.zeros_like(display_image_rgb, 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, 1 - alpha_blend, pred_overlay_vis, alpha_blend, 0)\n            else:\n                print(f\"  患者 {patient_id_vis} 切片 {inst_num_vis} 无有效预测掩码用于叠加。\")\n\n\n            # --- 显示 ---\n            fig, axes = plt.subplots(1, 3, figsize=(18, 6))\n            fig.suptitle(f\"患者 {patient_id_vis} - DICOM索引 {dicom_list_idx_vis} (Inst: {inst_num_vis}) - NII索引 {nii_slice_idx_for_gt}\", fontsize=16)\n\n            axes[0].imshow(cv2.cvtColor(display_image_rgb, cv2.COLOR_BGR2RGB))\n            axes[0].set_title(\"原始DICOM图像 (调整大小后)\")\n            axes[0].axis('off')\n\n            axes[1].imshow(cv2.cvtColor(gt_blended_vis, cv2.COLOR_BGR2RGB))\n            axes[1].set_title(\"真实掩码叠加\")\n            axes[1].axis('off')\n\n            axes[2].imshow(cv2.cvtColor(pred_blended_vis, cv2.COLOR_BGR2RGB))\n            axes[2].set_title(\"预测掩码叠加\")\n            axes[2].axis('off')\n            \n            # 添加图例\n            legend_elements = [plt.Rectangle((0, 0), 1, 1, color=[c/255. for c in colors_for_visualization[org][::-1]], label=org) # Matplotlib 用 RGB 0-1\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.02))\n            \n            plt.tight_layout(rect=[0, 0.03, 1, 0.95]) # 调整以防标题和图例重叠\n            plt.savefig(os.path.join(OUTPUT_DIR, f\"vis_pred_gt_{patient_id_vis}_dcm{dicom_list_idx_vis}_nii{nii_slice_idx_for_gt}.png\"))\n            plt.show()\n            \n            gc.collect() # 清理内存\n            \n    print(\"-\" * 30)\n    print(\"可视化部分执行完毕。\")\n    print(\"-\" * 30)\n\nif __name__ == \"__main__\":\n    main()","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}