{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.11.11","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"none","dataSources":[{"sourceId":52254,"databundleVersionId":9674523,"sourceType":"competition"},{"sourceId":11555525,"sourceType":"datasetVersion","datasetId":7245708},{"sourceId":11885494,"sourceType":"datasetVersion","datasetId":7470201}],"dockerImageVersionId":31040,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"code","source":"import 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":{"_uuid":"8f2839f25d086af736a60e9eeb907d3b93b6e0e5","_cell_guid":"b1076dfc-b9ad-4769-8c92-a6c4dae69d19","trusted":true,"execution":{"iopub.status.busy":"2025-05-22T14:49:27.75794Z","iopub.execute_input":"2025-05-22T14:49:27.758175Z","iopub.status.idle":"2025-05-22T14:49:28.048701Z","shell.execute_reply.started":"2025-05-22T14:49:27.758152Z","shell.execute_reply":"2025-05-22T14:49:28.048002Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 环境设置","metadata":{}},{"cell_type":"code","source":"# 导入必要的库\nimport os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport pydicom\nfrom pydicom.pixel_data_handlers.util import apply_voi_lut\nimport cv2\nimport torch\nimport torch.nn as nn\nimport torch.nn.functional as F\nfrom torch.utils.data import Dataset, DataLoader, Subset\nfrom torch.optim.lr_scheduler import ReduceLROnPlateau\nfrom torchvision import models\nimport random\nimport glob\nfrom tqdm.notebook import tqdm\nimport pickle\nfrom sklearn.model_selection import KFold\nimport warnings\nwarnings.filterwarnings('ignore')\n\n# 设置随机种子确保可复现性\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    torch.manual_seed(seed)\n    torch.cuda.manual_seed(seed)\n    torch.backends.cudnn.deterministic = True\n    torch.backends.cudnn.benchmark = False\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\nset_seed(42)\n\n# 检查GPU可用性\ndevice = torch.device('cuda' if torch.cuda.is_available() else 'cpu')\nprint(f\"使用设备: {device}\")\n\n# 设置路径\nDATA_DIR = '/kaggle/input/rsna-2023-abdominal-trauma-detection'\nTRAIN_IMAGES_DIR = os.path.join(DATA_DIR, 'train_images')\nTRAIN_CSV_PATH = os.path.join(DATA_DIR, 'train_2024.csv')\nTRAIN_META_PATH = os.path.join(DATA_DIR, 'train_series_meta.csv')\nSEGMENTATION_DIR = '/kaggle/input/unet-cache/segmentation_predictions_multi_v2'  # 您生成的分割掩码路径\n\n# 创建输出目录\nOUTPUT_DIR = '/kaggle/working/model_output'\nos.makedirs(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 = None  # 不需要旋转\nUSE_REVERSE_NIFTI_MAPPING = False  # 使用正向映射\n\n# 训练参数\nBATCH_SIZE = 4\nNUM_EPOCHS = 10\nLEARNING_RATE = 1e-4\nNUM_SLICES = 16  # 论文中使用32个切片\nNUM_FOLDS = 5\n\nprint(f\"图像尺寸设置为: {IMG_SIZE}\")\nprint(f\"NIFTI方向变换: {BEST_NIFTI_ORIENTATION_TRANSFORM}\")\nprint(f\"使用反向映射: {USE_REVERSE_NIFTI_MAPPING}\")\nprint(\"环境和参数设置完成\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T14:49:28.049358Z","iopub.execute_input":"2025-05-22T14:49:28.049831Z","iopub.status.idle":"2025-05-22T14:49:37.919033Z","shell.execute_reply.started":"2025-05-22T14:49:28.049811Z","shell.execute_reply":"2025-05-22T14:49:37.918426Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 数据预处理和辅助函数","metadata":{}},{"cell_type":"code","source":"# 加载DICOM切片\ndef load_dicom_slice(path):\n    \"\"\"加载单个DICOM切片，应用VOI LUT，归一化\"\"\"\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        # 转换为Hounsfield单位\n        intercept = dicom_file.RescaleIntercept\n        slope = dicom_file.RescaleSlope\n        image = image * slope + intercept\n\n        # 窗口化处理\n        window_center = 50  # 腹窗\n        window_width = 350\n        img_min = window_center - window_width // 2\n        img_max = window_center + window_width // 2\n        image = image.copy()\n        image[image < img_min] = img_min\n        image[image > img_max] = img_max\n\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        # 处理 MONOCHROME1 (图像像素值需要反转)\n        if 'PhotometricInterpretation' in dicom_file and dicom_file.PhotometricInterpretation == \"MONOCHROME1\":\n            image = 1.0 - image\n\n        return image, instance_number, dicom_file\n    except Exception as e:\n        # print(f\"加载 DICOM 错误 {path}: {e}\")\n        return None, None, None\n\n# 获取DICOM文件列表\ndef get_dicom_files_dict(patient_id):\n    \"\"\"获取患者的DICOM文件列表\"\"\"\n    dicom_info = {}\n    patient_dir = os.path.join(TRAIN_IMAGES_DIR, str(patient_id))\n    \n    if not os.path.exists(patient_dir):\n        return dicom_info\n        \n    for series_id in os.listdir(patient_dir):\n        series_dir = os.path.join(patient_dir, series_id)\n        if not os.path.isdir(series_dir):\n            continue\n            \n        dicom_files = glob.glob(os.path.join(series_dir, '*.dcm'))\n        dicom_tuples = []\n        \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:\n                pass\n                \n        # 按InstanceNumber排序\n        dicom_tuples.sort(key=lambda x: x[0])\n        dicom_info[series_id] = dicom_tuples\n        \n    return dicom_info\n\n# 从分割掩码文件加载掩码\ndef load_segmentation_mask(mask_path, target_size=TARGET_SIZE_INT):\n    \"\"\"加载分割掩码\"\"\"\n    try:\n        mask_data = np.load(mask_path)\n        mask = mask_data['mask']\n        \n        # 调整大小\n        if mask.shape[0] != target_size or mask.shape[1] != target_size:\n            resized_mask = np.zeros((target_size, target_size, mask.shape[2]), dtype=mask.dtype)\n            for c in range(mask.shape[2]):\n                resized_mask[:, :, c] = cv2.resize(mask[:, :, c], (target_size, target_size), \n                                                  interpolation=cv2.INTER_NEAREST)\n            mask = resized_mask\n            \n        return mask\n    except Exception as e:\n        # print(f\"加载掩码错误 {mask_path}: {e}\")\n        return None\n\n# 增强外渗特征\ndef enhance_extravasation_features(image, dicom_file, aortic_hu=None):\n    \"\"\"基于CT值特性增强外渗特征\"\"\"\n    if dicom_file is None:\n        return np.zeros_like(image)\n        \n    # 获取原始HU值\n    pixel_array = dicom_file.pixel_array\n    intercept = dicom_file.RescaleIntercept\n    slope = dicom_file.RescaleSlope\n    hu_image = pixel_array * slope + intercept\n    \n    # 设置阈值\n    if aortic_hu is not None and aortic_hu > 100:  # 确保是造影检查\n        lower_threshold = aortic_hu * 0.6\n        upper_threshold = aortic_hu * 1.1\n    else:\n        # 默认值\n        lower_threshold = 100\n        upper_threshold = 300\n    \n    # 创建掩码\n    extravasation_mask = np.zeros_like(hu_image)\n    extravasation_mask[(hu_image >= lower_threshold) & (hu_image <= upper_threshold)] = 1\n    \n    # 形态学操作去除噪声\n    kernel = np.ones((3, 3), np.uint8)\n    extravasation_mask = cv2.morphologyEx(extravasation_mask, cv2.MORPH_OPEN, kernel)\n    extravasation_mask = cv2.morphologyEx(extravasation_mask, cv2.MORPH_CLOSE, kernel)\n    \n    # 调整大小以匹配图像\n    if extravasation_mask.shape != image.shape:\n        extravasation_mask = cv2.resize(extravasation_mask, (image.shape[1], image.shape[0]), \n                                       interpolation=cv2.INTER_NEAREST)\n    \n    return extravasation_mask\n\n# 选择代表性切片\ndef select_representative_slices(slices_data, num_slices=NUM_SLICES, organs_of_interest=['liver', 'spleen', 'kidney', 'bowel']):\n    \"\"\"选择代表性切片，确保每个器官至少出现在一定数量的切片中\"\"\"\n    if len(slices_data) <= num_slices:\n        return slices_data\n        \n    # 计算每个切片中各器官的存在情况\n    organ_presence = []\n    for slice_data in slices_data:\n        mask = slice_data.get('mask')\n        presence = {organ: False for organ in organs_of_interest}\n        \n        if mask is not None:\n            for i, organ in enumerate(organs_of_interest):\n                if i < mask.shape[2] and np.any(mask[:, :, i] > 0.5):\n                    presence[organ] = True\n                    \n        organ_presence.append(presence)\n    \n    # 分数计算：每个切片覆盖的器官数量\n    slice_scores = []\n    for i, presence in enumerate(organ_presence):\n        score = sum(presence.values())\n        slice_scores.append((i, score))\n    \n    # 按分数降序排序\n    slice_scores.sort(key=lambda x: x[1], reverse=True)\n    \n    # 选择分数最高的切片\n    selected_indices = [score[0] for score in slice_scores[:num_slices]]\n    selected_indices.sort()  # 按原始顺序排序\n    \n    return [slices_data[i] for i in selected_indices]\n\n# 图像预处理和调整大小\ndef preprocess_image(image, target_size=TARGET_SIZE_INT):\n    \"\"\"预处理图像，调整大小\"\"\"\n    if image.shape[0] != target_size or image.shape[1] != target_size:\n        image = cv2.resize(image, (target_size, target_size), interpolation=cv2.INTER_LINEAR)\n    return image","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T14:49:37.920603Z","iopub.execute_input":"2025-05-22T14:49:37.92095Z","iopub.status.idle":"2025-05-22T14:49:37.937425Z","shell.execute_reply.started":"2025-05-22T14:49:37.920931Z","shell.execute_reply":"2025-05-22T14:49:37.936744Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 数据集类定义","metadata":{}},{"cell_type":"code","source":"class AbdominalTraumaDataset(Dataset):\n    def __init__(self, patient_ids, meta_df, transform=None, num_slices=NUM_SLICES, mode='train'):\n        self.patient_ids = patient_ids\n        self.meta_df = meta_df\n        self.transform = transform\n        self.num_slices = num_slices\n        self.mode = mode\n        self.labels_df = pd.read_csv(TRAIN_CSV_PATH)\n\n        self.aortic_hu_map = {}\n        for _, row in meta_df.iterrows():\n            patient_id = str(row['patient_id'])\n            series_id = str(row['series_id'])\n            aortic_hu = row['aortic_hu']\n            if patient_id not in self.aortic_hu_map:\n                self.aortic_hu_map[patient_id] = {}\n            self.aortic_hu_map[patient_id][series_id] = aortic_hu\n\n    def __len__(self):\n        return len(self.patient_ids)\n\n    def __getitem__(self, idx):\n        patient_id = str(self.patient_ids[idx]) # 确保 patient_id 是字符串\n        dicom_files_dict = get_dicom_files_dict(patient_id)\n\n        if not dicom_files_dict:\n            return self._create_empty_sample()\n\n        series_id, dicom_tuples = next(iter(dicom_files_dict.items()))\n\n        aortic_hu = 200.0 # 默认值\n        if patient_id in self.aortic_hu_map and series_id in self.aortic_hu_map[patient_id]:\n            aortic_hu = self.aortic_hu_map[patient_id][series_id]\n            if pd.isna(aortic_hu): # 处理可能的 NaN 值\n                 aortic_hu = 200.0\n\n        slices_data = []\n        for dicom_idx, (instance_number, dicom_path) in enumerate(dicom_tuples):\n            raw_dicom_image, _, dicom_file_obj = load_dicom_slice(dicom_path)\n\n            if raw_dicom_image is None:\n                continue\n\n            current_extravasation_features_orig_size = enhance_extravasation_features(raw_dicom_image, dicom_file_obj, aortic_hu)\n            processed_main_image = preprocess_image(raw_dicom_image.copy(), target_size=TARGET_SIZE_INT)\n\n            mask_path = os.path.join(SEGMENTATION_DIR, patient_id, series_id, f\"{instance_number}.npz\")\n            mask = None\n            if os.path.exists(mask_path):\n                mask = load_segmentation_mask(mask_path, target_size=TARGET_SIZE_INT)\n\n            if current_extravasation_features_orig_size.shape[0] != TARGET_SIZE_INT or \\\n               current_extravasation_features_orig_size.shape[1] != TARGET_SIZE_INT:\n                resized_extravasation_features = cv2.resize(\n                    current_extravasation_features_orig_size,\n                    (TARGET_SIZE_INT, TARGET_SIZE_INT),\n                    interpolation=cv2.INTER_NEAREST\n                )\n            else:\n                resized_extravasation_features = current_extravasation_features_orig_size\n\n            slices_data.append({\n                'instance_number': instance_number,\n                'image': processed_main_image,\n                'mask': mask,\n                'extravasation_features': resized_extravasation_features\n            })\n\n        if not slices_data: # 如果所有切片都无法加载\n            return self._create_empty_sample()\n\n        if len(slices_data) > self.num_slices:\n            slices_data = select_representative_slices(slices_data, self.num_slices)\n\n        while len(slices_data) < self.num_slices:\n            if slices_data: # 确保 slices_data 不为空\n                slices_data.append(slices_data[-1])\n            else: # 理论上不应该执行到这里，因为前面有 if not slices_data: return ...\n                return self._create_empty_sample()\n\n\n        images_list = []\n        masks_list = []\n        extravasation_features_list = []\n\n        for slice_data in slices_data:\n            img_np = slice_data['image'] # 已经是 numpy array (H, W)\n            if self.transform:\n                # ToPILImage 需要 (H, W) 或 (H, W, C)\n                # 如果 img_np 是 (H,W)，ToPILImage 会将其视为 'L' 模式\n                # ToTensor 会将 PIL 'L' 模式图像转换为 [1, H, W]\n                img_tensor = self.transform(img_np)\n            else:\n                img_tensor = torch.tensor(img_np, dtype=torch.float32).unsqueeze(0) # [1, H, W]\n            images_list.append(img_tensor)\n\n            if slice_data['mask'] is not None:\n                masks_list.append(slice_data['mask']) # (H, W, 4)\n            else:\n                empty_mask = np.zeros((TARGET_SIZE_INT, TARGET_SIZE_INT, 4), dtype=np.float32)\n                masks_list.append(empty_mask)\n\n            extravasation_features_list.append(slice_data['extravasation_features']) # (H, W)\n\n        images = torch.stack(images_list).float() # [num_slices, 1, H, W]\n        # masks_list 是 [(H,W,4), (H,W,4), ...]\n        # np.stack(masks_list) -> [num_slices, H, W, 4]\n        masks = torch.tensor(np.stack(masks_list), dtype=torch.float32)\n        # extravasation_features_list 是 [(H,W), (H,W), ...]\n        # np.stack(extravasation_features_list) -> [num_slices, H, W]\n        # .unsqueeze(1) -> [num_slices, 1, H, W]\n        extravasation_features_tensor = torch.tensor(np.stack(extravasation_features_list), dtype=torch.float32).unsqueeze(1)\n        aortic_hu_tensor = torch.tensor([float(aortic_hu)], dtype=torch.float32)\n\n\n        if self.mode == 'train' or self.mode == 'val':\n            # patient_id 在这里应该是 int 类型以便于在 labels_df 中查找\n            label_row = self.labels_df[self.labels_df['patient_id'] == int(patient_id)].iloc[0]\n            bowel_label = int(label_row['bowel_healthy'] == 0)\n            extravasation_label = int(label_row['extravasation_healthy'] == 0)\n            kidney_label = int(label_row['kidney_healthy'] == 0) * (1 + int(label_row['kidney_low'] == 0))\n            liver_label = int(label_row['liver_healthy'] == 0) * (1 + int(label_row['liver_low'] == 0))\n            spleen_label = int(label_row['spleen_healthy'] == 0) * (1 + int(label_row['spleen_low'] == 0))\n        else: # test mode\n            bowel_label = 0\n            extravasation_label = 0\n            kidney_label = 0\n            liver_label = 0\n            spleen_label = 0\n\n        return {\n            'patient_id': patient_id, # 返回原始的 patient_id (str)\n            'images': images,\n            'masks': masks,\n            'extravasation_features': extravasation_features_tensor,\n            'aortic_hu': aortic_hu_tensor,\n            'bowel': torch.tensor([bowel_label], dtype=torch.float32),\n            'extravasation': torch.tensor([extravasation_label], dtype=torch.float32),\n            'kidney': torch.tensor(kidney_label, dtype=torch.long),\n            'liver': torch.tensor(liver_label, dtype=torch.long),\n            'spleen': torch.tensor(spleen_label, dtype=torch.long)\n        }\n\n    def _create_empty_sample(self):\n        empty_images = torch.zeros((self.num_slices, 1, TARGET_SIZE_INT, TARGET_SIZE_INT), dtype=torch.float32)\n        empty_masks = torch.zeros((self.num_slices, TARGET_SIZE_INT, TARGET_SIZE_INT, 4), dtype=torch.float32)\n        empty_extravasation = torch.zeros((self.num_slices, 1, TARGET_SIZE_INT, TARGET_SIZE_INT), dtype=torch.float32)\n        empty_aortic_hu = torch.tensor([0.0], dtype=torch.float32)\n\n        return {\n            'patient_id': '0', # 虚拟 patient_id\n            'images': empty_images,\n            'masks': empty_masks,\n            'extravasation_features': empty_extravasation,\n            'aortic_hu': empty_aortic_hu,\n            'bowel': torch.tensor([0], dtype=torch.float32),\n            'extravasation': torch.tensor([0], dtype=torch.float32),\n            'kidney': torch.tensor(0, dtype=torch.long),\n            'liver': torch.tensor(0, dtype=torch.long),\n            'spleen': torch.tensor(0, dtype=torch.long)\n        }\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T14:49:37.938146Z","iopub.execute_input":"2025-05-22T14:49:37.938396Z","iopub.status.idle":"2025-05-22T14:49:37.969145Z","shell.execute_reply.started":"2025-05-22T14:49:37.938378Z","shell.execute_reply":"2025-05-22T14:49:37.968419Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 数据增强和数据加载器","metadata":{}},{"cell_type":"code","source":"from torchvision.transforms import Compose, RandomHorizontalFlip, RandomRotation, RandomAffine, ColorJitter, GaussianBlur, ToTensor, ToPILImage\nimport PIL.Image as Image\n\n# 定义数据增强\ndef get_transforms(mode='train'):\n    if mode == 'train':\n        return Compose([\n            ToPILImage(),  # 先将NumPy数组转换为PIL Image\n            RandomHorizontalFlip(p=0.5),\n            RandomRotation(degrees=10),\n            RandomAffine(degrees=0, translate=(0.05, 0.05), scale=(0.95, 1.05)),\n            ColorJitter(brightness=0.2, contrast=0.2),\n            GaussianBlur(kernel_size=3, sigma=(0.1, 0.5)),\n            ToTensor(),  # 转换回Tensor\n        ])\n    else:\n        return ToTensor()  # 验证集只需要转换为Tensor\n\n# 准备数据集\ndef prepare_data(meta_df, num_folds=NUM_FOLDS, fold_idx=0):\n    # 获取所有患者ID\n    all_patients = meta_df['patient_id'].unique()\n    \n    # 创建KFold分割\n    kf = KFold(n_splits=num_folds, shuffle=True, random_state=42)\n    folds = list(kf.split(all_patients))\n    \n    train_indices, val_indices = folds[fold_idx]\n    train_patients = [str(all_patients[i]) for i in train_indices]\n    val_patients = [str(all_patients[i]) for i in val_indices]\n    \n    # 创建数据集\n    train_dataset = AbdominalTraumaDataset(\n        train_patients, meta_df, transform=get_transforms('train'), mode='train'\n    )\n    val_dataset = AbdominalTraumaDataset(\n        val_patients, meta_df, transform=get_transforms('val'), mode='train'\n    )\n    \n    # 创建数据加载器 - 减少worker数量以避免潜在的多进程问题\n    train_loader = DataLoader(\n        train_dataset, batch_size=BATCH_SIZE, shuffle=True, num_workers=4, pin_memory=True, drop_last=True\n    )\n    val_loader = DataLoader(\n        val_dataset, batch_size=BATCH_SIZE, shuffle=False, num_workers=4, pin_memory=True\n    )\n    \n    return train_loader, val_loader\n\n# 加载元数据\ndef load_meta_data():\n    meta_df = pd.read_csv(TRAIN_META_PATH)\n    # 确保数据类型正确\n    meta_df['patient_id'] = meta_df['patient_id'].astype(str)\n    meta_df['series_id'] = meta_df['series_id'].astype(str)\n    return meta_df\n\nprint(\"数据加载和增强函数定义完成\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T14:49:37.969977Z","iopub.execute_input":"2025-05-22T14:49:37.970243Z","iopub.status.idle":"2025-05-22T14:49:37.989896Z","shell.execute_reply.started":"2025-05-22T14:49:37.970226Z","shell.execute_reply":"2025-05-22T14:49:37.989333Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 模型定义","metadata":{}},{"cell_type":"code","source":"class AbdominalTraumaModel(nn.Module):\n    \"\"\"腹部外伤2.5D分类模型，实现论文中描述的架构\"\"\"\n    \n    def __init__(self, num_slices=NUM_SLICES):\n        super().__init__()\n        \n        # EfficientNetB1作为特征提取器\n        self.efficientnet = models.efficientnet_b1(pretrained=True)\n        self.feature_dim = 1280  # EfficientNetB1的特征维度\n        \n        # 掩码处理分支\n        self.mask_conv = nn.Conv2d(4, 16, kernel_size=3, padding=1)\n        self.mask_pool = nn.AdaptiveAvgPool2d(1)\n        self.mask_fc = nn.Linear(16, 64)\n        \n        # 外渗特征处理分支\n        self.extra_conv = nn.Conv2d(1, 8, kernel_size=3, padding=1)\n        self.extra_pool = nn.AdaptiveAvgPool2d(1)\n        self.extra_fc = nn.Linear(8, 32)\n        \n        # aortic_hu处理\n        self.aortic_encoder = nn.Sequential(\n            nn.Linear(1, 16),\n            nn.ReLU(),\n            nn.Linear(16, 32),\n            nn.ReLU()\n        )\n        \n        # LSTM处理时序特征\n        self.lstm = nn.LSTM(\n            input_size=self.feature_dim,\n            hidden_size=512,\n            num_layers=2,\n            batch_first=True,\n            bidirectional=True\n        )\n        \n        # Neck结构\n        self.neck = nn.Sequential(\n            nn.Linear(512*2, 512),  # 双向LSTM，所以是512*2\n            nn.BatchNorm1d(512),\n            nn.ReLU(),\n            nn.Dropout(0.3),\n            nn.Linear(512, 256),\n            nn.BatchNorm1d(256),\n            nn.ReLU(),\n            nn.Dropout(0.2)\n        )\n        \n        # 多标签分类头\n        combined_features = 256 + 64 + 32 + 32  # Neck + 掩码特征 + 外渗特征 + aortic_hu特征\n        self.bowel_head = nn.Linear(combined_features, 1)\n        self.extravasation_head = nn.Linear(combined_features, 1)\n        self.kidney_head = nn.Linear(combined_features, 3)\n        self.liver_head = nn.Linear(combined_features, 3)\n        self.spleen_head = nn.Linear(combined_features, 3)\n    \n    def forward(self, images, masks, extravasation_features, aortic_hu):\n        \"\"\"\n        参数:\n            images: [B, num_slices, C, H, W] - 批次中每个患者的切片，C可能是1或3\n            masks: [B, num_slices, H, W, 4] - 对应的分割掩码\n            extravasation_features: [B, num_slices, 1, H, W] - 外渗特征\n            aortic_hu: [B, 1] - 主动脉HU值\n        返回:\n            各器官的损伤预测结果\n        \"\"\"\n        batch_size = images.size(0)\n        seq_len = images.size(1)  # 应该是num_slices\n        \n        # 检查输入通道数，确保是3通道\n        if images.size(2) == 1:  # 如果是单通道\n            # 重塑以便批量处理所有切片\n            images_reshaped = images.view(batch_size * seq_len, 1, images.size(3), images.size(4))\n            # 扩展为3通道\n            images_reshaped = images_reshaped.repeat(1, 3, 1, 1)\n        else:  # 已经是3通道\n            images_reshaped = images.view(batch_size * seq_len, images.size(2), images.size(3), images.size(4))\n        \n        # 处理掩码\n        masks_reshaped = masks.view(batch_size * seq_len, masks.size(2), masks.size(3), masks.size(4))\n        masks_reshaped = masks_reshaped.permute(0, 3, 1, 2)  # [B*seq_len, 4, H, W]\n        \n        # 处理外渗特征\n        extra_reshaped = extravasation_features.view(batch_size * seq_len, 1, extravasation_features.size(3), extravasation_features.size(4))\n        \n        # 特征提取\n        features = self.efficientnet.features(images_reshaped)  # [B*seq_len, 1280, h, w]\n        \n        # 掩码特征处理\n        mask_features = self.mask_conv(masks_reshaped)  # [B*seq_len, 16, h, w]\n        mask_features = self.mask_pool(mask_features).squeeze(-1).squeeze(-1)  # [B*seq_len, 16]\n        mask_features = self.mask_fc(mask_features)  # [B*seq_len, 64]\n        \n        # 外渗特征处理\n        extra_features = self.extra_conv(extra_reshaped)  # [B*seq_len, 8, h, w]\n        extra_features = self.extra_pool(extra_features).squeeze(-1).squeeze(-1)  # [B*seq_len, 8]\n        extra_features = self.extra_fc(extra_features)  # [B*seq_len, 32]\n        \n        # 全局平均池化EfficientNet特征\n        pooled_features = F.adaptive_avg_pool2d(features, (1, 1)).squeeze(-1).squeeze(-1)  # [B*seq_len, 1280]\n        \n        # 重塑回序列形式\n        sequence_features = pooled_features.view(batch_size, seq_len, -1)  # [B, seq_len, 1280]\n        \n        # LSTM处理序列特征\n        lstm_out, _ = self.lstm(sequence_features)  # [B, seq_len, 512*2]\n        \n        # 取最后一个时间步的输出\n        final_lstm_out = lstm_out[:, -1, :]  # [B, 512*2]\n        \n        # Neck结构处理\n        neck_out = self.neck(final_lstm_out)  # [B, 256]\n        \n        # 重塑掩码特征和外渗特征以匹配批次大小\n        mask_features = mask_features.view(batch_size, seq_len, -1)  # [B, seq_len, 64]\n        mask_features_avg = torch.mean(mask_features, dim=1)  # [B, 64]\n        \n        extra_features = extra_features.view(batch_size, seq_len, -1)  # [B, seq_len, 32]\n        extra_features_avg = torch.mean(extra_features, dim=1)  # [B, 32]\n        \n        # 处理aortic_hu\n        aortic_features = self.aortic_encoder(aortic_hu)  # [B, 32]\n        \n        # 连接特征\n        combined_features = torch.cat([neck_out, mask_features_avg, extra_features_avg, aortic_features], dim=1)  # [B, 256+64+32+32]\n        \n        # 多标签分类\n        bowel_out = self.bowel_head(combined_features)\n        extravasation_out = self.extravasation_head(combined_features)\n        kidney_out = self.kidney_head(combined_features)\n        liver_out = self.liver_head(combined_features)\n        spleen_out = self.spleen_head(combined_features)\n        \n        return {\n            'bowel': bowel_out,\n            'extravasation': extravasation_out,\n            'kidney': kidney_out,\n            'liver': liver_out,\n            'spleen': spleen_out\n        }\n\n# 初始化模型\nmodel = AbdominalTraumaModel().to(device)\nprint(\"模型定义完成\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T14:49:37.990783Z","iopub.execute_input":"2025-05-22T14:49:37.99095Z","iopub.status.idle":"2025-05-22T14:49:42.187928Z","shell.execute_reply.started":"2025-05-22T14:49:37.990936Z","shell.execute_reply":"2025-05-22T14:49:42.187193Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 损失函数和优化器","metadata":{}},{"cell_type":"code","source":"# 定义损失函数\nbce_bowel = nn.BCEWithLogitsLoss(pos_weight=torch.tensor([2.0]).to(device))\nbce_extravasation = nn.BCEWithLogitsLoss(pos_weight=torch.tensor([4.0]).to(device))\nce_loss = nn.CrossEntropyLoss(label_smoothing=0.05, weight=torch.tensor([1.0, 2.0, 4.0]).to(device))\n\n# 定义优化器\noptimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE)\nscheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n    optimizer, mode='min', patience=3, factor=0.5, verbose=True\n)\n\n# 指标计算\nclass MetricsCalculator:\n    def __init__(self, mode='binary'):\n        self.probabilities = []\n        self.predictions = []\n        self.targets = []\n        self.mode = mode\n    \n    def update(self, logits, target):\n        if self.mode == 'binary':\n            probabilities = torch.sigmoid(logits)\n            predicted = (probabilities > 0.5)\n        else:\n            probabilities = F.softmax(logits, dim=1)\n            predicted = torch.argmax(probabilities, dim=1)\n            \n        self.probabilities.extend(probabilities.detach().cpu().numpy())\n        self.predictions.extend(predicted.detach().cpu().numpy())\n        self.targets.extend(target.detach().cpu().numpy())\n    \n    def reset(self):\n        self.probabilities = []\n        self.predictions = []\n        self.targets = []\n    \n    def compute_accuracy(self):\n        if not self.predictions:\n            return 0.0\n        return np.mean(np.array(self.predictions) == np.array(self.targets))\n    \n    def compute_auc(self):\n        if not self.probabilities or len(np.unique(self.targets)) < 2:\n            return 0.0\n        try:\n            if self.mode == 'multi':\n                from sklearn.metrics import roc_auc_score\n                return roc_auc_score(self.targets, self.probabilities, multi_class='ovo')\n            else:\n                from sklearn.metrics import roc_auc_score\n                return roc_auc_score(self.targets, self.probabilities)\n        except:\n            return 0.0\n\n# 初始化指标计算器\ntrain_metrics = {\n    'bowel': MetricsCalculator('binary'),\n    'extravasation': MetricsCalculator('binary'),\n    'kidney': MetricsCalculator('multi'),\n    'liver': MetricsCalculator('multi'),\n    'spleen': MetricsCalculator('multi')\n}\n\nval_metrics = {\n    'bowel': MetricsCalculator('binary'),\n    'extravasation': MetricsCalculator('binary'),\n    'kidney': MetricsCalculator('multi'),\n    'liver': MetricsCalculator('multi'),\n    'spleen': MetricsCalculator('multi')\n}\n\nprint(\"损失函数和优化器设置完成\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T14:49:42.188716Z","iopub.execute_input":"2025-05-22T14:49:42.188946Z","iopub.status.idle":"2025-05-22T14:49:42.201678Z","shell.execute_reply.started":"2025-05-22T14:49:42.188929Z","shell.execute_reply":"2025-05-22T14:49:42.201067Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 训练函数","metadata":{}},{"cell_type":"code","source":"def train_epoch(model, train_loader, optimizer, metrics):\n    \"\"\"训练单个epoch\"\"\"\n    model.train()\n    total_loss = 0\n    \n    # 重置指标\n    for metric in metrics.values():\n        metric.reset()\n    \n    # 进度条\n    progress_bar = tqdm(train_loader, desc=\"训练\")\n    \n    for batch_idx, batch in enumerate(progress_bar):\n        # 获取数据\n        images = batch['images'].to(device)\n        masks = batch['masks'].to(device)\n        extravasation_features = batch['extravasation_features'].to(device)\n        aortic_hu = batch['aortic_hu'].to(device)\n        \n        # 获取标签\n        bowel = batch['bowel'].to(device)\n        extravasation = batch['extravasation'].to(device)\n        kidney = batch['kidney'].to(device)\n        liver = batch['liver'].to(device)\n        spleen = batch['spleen'].to(device)\n        \n        # 前向传播\n        optimizer.zero_grad()\n        outputs = model(images, masks, extravasation_features, aortic_hu)\n        \n        # 计算损失\n        bowel_loss = bce_bowel(outputs['bowel'], bowel)\n        extravasation_loss = bce_extravasation(outputs['extravasation'], extravasation)\n        kidney_loss = ce_loss(outputs['kidney'], kidney)\n        liver_loss = ce_loss(outputs['liver'], liver)\n        spleen_loss = ce_loss(outputs['spleen'], spleen)\n        \n        # 总损失\n        loss = bowel_loss + extravasation_loss + kidney_loss + liver_loss + spleen_loss\n        \n        # 反向传播\n        loss.backward()\n        optimizer.step()\n        \n        # 更新总损失\n        total_loss += loss.item()\n        \n        # 更新指标\n        metrics['bowel'].update(outputs['bowel'], bowel)\n        metrics['extravasation'].update(outputs['extravasation'], extravasation)\n        metrics['kidney'].update(outputs['kidney'], kidney)\n        metrics['liver'].update(outputs['liver'], liver)\n        metrics['spleen'].update(outputs['spleen'], spleen)\n        \n        # 更新进度条\n        progress_bar.set_postfix(loss=loss.item())\n    \n    # 计算平均损失\n    avg_loss = total_loss / len(train_loader)\n    \n    # 计算指标\n    metrics_results = {\n        'loss': avg_loss,\n        'bowel_acc': metrics['bowel'].compute_accuracy(),\n        'extravasation_acc': metrics['extravasation'].compute_accuracy(),\n        'kidney_acc': metrics['kidney'].compute_accuracy(),\n        'liver_acc': metrics['liver'].compute_accuracy(),\n        'spleen_acc': metrics['spleen'].compute_accuracy(),\n        'bowel_auc': metrics['bowel'].compute_auc(),\n        'extravasation_auc': metrics['extravasation'].compute_auc(),\n        'kidney_auc': metrics['kidney'].compute_auc(),\n        'liver_auc': metrics['liver'].compute_auc(),\n        'spleen_auc': metrics['spleen'].compute_auc()\n    }\n    \n    return metrics_results\n\ndef validate(model, val_loader, metrics):\n    \"\"\"验证模型\"\"\"\n    model.eval()\n    total_loss = 0\n    \n    # 重置指标\n    for metric in metrics.values():\n        metric.reset()\n    \n    # 进度条\n    progress_bar = tqdm(val_loader, desc=\"验证\")\n    \n    with torch.no_grad():\n        for batch_idx, batch in enumerate(progress_bar):\n            # 获取数据\n            images = batch['images'].to(device)\n            masks = batch['masks'].to(device)\n            extravasation_features = batch['extravasation_features'].to(device)\n            aortic_hu = batch['aortic_hu'].to(device)\n            \n            # 获取标签\n            bowel = batch['bowel'].to(device)\n            extravasation = batch['extravasation'].to(device)\n            kidney = batch['kidney'].to(device)\n            liver = batch['liver'].to(device)\n            spleen = batch['spleen'].to(device)\n            \n            # 前向传播\n            outputs = model(images, masks, extravasation_features, aortic_hu)\n            \n            # 计算损失\n            bowel_loss = bce_bowel(outputs['bowel'], bowel)\n            extravasation_loss = bce_extravasation(outputs['extravasation'], extravasation)\n            kidney_loss = ce_loss(outputs['kidney'], kidney)\n            liver_loss = ce_loss(outputs['liver'], liver)\n            spleen_loss = ce_loss(outputs['spleen'], spleen)\n            \n            # 总损失\n            loss = bowel_loss + extravasation_loss + kidney_loss + liver_loss + spleen_loss\n            \n            # 更新总损失\n            total_loss += loss.item()\n            \n            # 更新指标\n            metrics['bowel'].update(outputs['bowel'], bowel)\n            metrics['extravasation'].update(outputs['extravasation'], extravasation)\n            metrics['kidney'].update(outputs['kidney'], kidney)\n            metrics['liver'].update(outputs['liver'], liver)\n            metrics['spleen'].update(outputs['spleen'], spleen)\n\n            # 打印每步信息\n            if (batch_idx + 1) % 10 == 0:  # 每10步打印一次，可以根据需要调整\n                print(f\"步骤 [{batch_idx+1}/{len(train_loader)}] 损失: {loss.item():.4f}\")\n            \n            # 更新进度条\n            progress_bar.set_postfix(loss=loss.item())\n    \n    # 计算平均损失\n    avg_loss = total_loss / len(val_loader)\n    \n    # 计算指标\n    metrics_results = {\n        'loss': avg_loss,\n        'bowel_acc': metrics['bowel'].compute_accuracy(),\n        'extravasation_acc': metrics['extravasation'].compute_accuracy(),\n        'kidney_acc': metrics['kidney'].compute_accuracy(),\n        'liver_acc': metrics['liver'].compute_accuracy(),\n        'spleen_acc': metrics['spleen'].compute_accuracy(),\n        'bowel_auc': metrics['bowel'].compute_auc(),\n        'extravasation_auc': metrics['extravasation'].compute_auc(),\n        'kidney_auc': metrics['kidney'].compute_auc(),\n        'liver_auc': metrics['liver'].compute_auc(),\n        'spleen_auc': metrics['spleen'].compute_auc()\n    }\n    \n    return metrics_results\n\nprint(\"训练和验证函数定义完成\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T14:49:42.203385Z","iopub.execute_input":"2025-05-22T14:49:42.203562Z","iopub.status.idle":"2025-05-22T14:49:42.22654Z","shell.execute_reply.started":"2025-05-22T14:49:42.203548Z","shell.execute_reply":"2025-05-22T14:49:42.225858Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 主训练循环","metadata":{}},{"cell_type":"code","source":"# 加载元数据\nmeta_df = load_meta_data()\n\n# 记录训练历史\nhistory = {\n    'train_loss': [],\n    'val_loss': [],\n    'train_metrics': [],\n    'val_metrics': []\n}\n\n# 最佳验证损失\nbest_val_loss = float('inf')\n\n# 使用多折交叉验证\nfor fold in range(NUM_FOLDS):\n    print(f\"\\n===== 开始训练折 {fold+1}/{NUM_FOLDS} =====\")\n    \n    # 准备数据\n    train_loader, val_loader = prepare_data(meta_df, NUM_FOLDS, fold)\n    \n    # 重新初始化模型\n    if fold > 0:\n        model = AbdominalTraumaModel().to(device)\n        optimizer = torch.optim.Adam(model.parameters(), lr=LEARNING_RATE)\n        scheduler = torch.optim.lr_scheduler.ReduceLROnPlateau(\n            optimizer, mode='min', patience=3, factor=0.5, verbose=True\n        )\n    \n    # 训练循环\n    for epoch in range(NUM_EPOCHS):\n        print(f\"\\nEpoch {epoch+1}/{NUM_EPOCHS}\")\n        \n        # 训练\n        train_metrics_results = train_epoch(model, train_loader, optimizer, train_metrics)\n        history['train_loss'].append(train_metrics_results['loss'])\n        history['train_metrics'].append(train_metrics_results)\n        \n        # 验证\n        val_metrics_results = validate(model, val_loader, val_metrics)\n        history['val_loss'].append(val_metrics_results['loss'])\n        history['val_metrics'].append(val_metrics_results)\n        \n        # 打印结果\n        print(f\"训练损失: {train_metrics_results['loss']:.4f}, 验证损失: {val_metrics_results['loss']:.4f}\")\n        print(f\"训练 AUC - 肠道: {train_metrics_results['bowel_auc']:.4f}, 外渗: {train_metrics_results['extravasation_auc']:.4f}, \"\n              f\"肾脏: {train_metrics_results['kidney_auc']:.4f}, 肝脏: {train_metrics_results['liver_auc']:.4f}, \"\n              f\"脾脏: {train_metrics_results['spleen_auc']:.4f}\")\n        print(f\"验证 AUC - 肠道: {val_metrics_results['bowel_auc']:.4f}, 外渗: {val_metrics_results['extravasation_auc']:.4f}, \"\n              f\"肾脏: {val_metrics_results['kidney_auc']:.4f}, 肝脏: {val_metrics_results['liver_auc']:.4f}, \"\n              f\"脾脏: {val_metrics_results['spleen_auc']:.4f}\")\n        \n        # 更新学习率\n        scheduler.step(val_metrics_results['loss'])\n        \n        # 保存最佳模型\n        if val_metrics_results['loss'] < best_val_loss:\n            best_val_loss = val_metrics_results['loss']\n            torch.save(model.state_dict(), os.path.join(OUTPUT_DIR, f'best_model_fold{fold}.pt'))\n            print(f\"保存最佳模型，验证损失: {best_val_loss:.4f}\")\n    \n    # 保存最终模型\n    torch.save(model.state_dict(), os.path.join(OUTPUT_DIR, f'final_model_fold{fold}.pt'))\n    print(f\"折 {fold+1} 训练完成，保存最终模型\")\n\nprint(\"\\n训练完成！\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-05-22T14:49:42.227297Z","iopub.execute_input":"2025-05-22T14:49:42.227526Z","execution_failed":"2025-05-22T15:01:06.622Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 可视化训练结果","metadata":{}},{"cell_type":"code","source":"# 绘制损失曲线\nplt.figure(figsize=(12, 5))\nplt.subplot(1, 2, 1)\nplt.plot(history['train_loss'], label='训练损失')\nplt.plot(history['val_loss'], label='验证损失')\nplt.title('训练和验证损失')\nplt.xlabel('Epoch')\nplt.ylabel('损失')\nplt.legend()\n\n# 绘制AUC曲线\nplt.subplot(1, 2, 2)\ntrain_bowel_auc = [m['bowel_auc'] for m in history['train_metrics']]\ntrain_extravasation_auc = [m['extravasation_auc'] for m in history['train_metrics']]\ntrain_liver_auc = [m['liver_auc'] for m in history['train_metrics']]\ntrain_kidney_auc = [m['kidney_auc'] for m in history['train_metrics']]\ntrain_spleen_auc = [m['spleen_auc'] for m in history['train_metrics']]\n\nval_bowel_auc = [m['bowel_auc'] for m in history['val_metrics']]\nval_extravasation_auc = [m['extravasation_auc'] for m in history['val_metrics']]\nval_liver_auc = [m['liver_auc'] for m in history['val_metrics']]\nval_kidney_auc = [m['kidney_auc'] for m in history['val_metrics']]\nval_spleen_auc = [m['spleen_auc'] for m in history['val_metrics']]\n\nplt.plot(val_bowel_auc, label='肠道')\nplt.plot(val_extravasation_auc, label='外渗')\nplt.plot(val_liver_auc, label='肝脏')\nplt.plot(val_kidney_auc, label='肾脏')\nplt.plot(val_spleen_auc, label='脾脏')\nplt.title('验证集各器官AUC')\nplt.xlabel('Epoch')\nplt.ylabel('AUC')\nplt.legend()\n\nplt.tight_layout()\nplt.savefig(os.path.join(OUTPUT_DIR, 'training_history.png'))\nplt.show()\n\n# 打印最终结果表格\nfrom tabulate import tabulate\n\nfinal_metrics = history['val_metrics'][-1]\nmetrics_table = [\n    ['器官', 'AUC', '准确率'],\n    ['肠道', f\"{final_metrics['bowel_auc']:.4f}\", f\"{final_metrics['bowel_acc']:.4f}\"],\n    ['外渗', f\"{final_metrics['extravasation_auc']:.4f}\", f\"{final_metrics['extravasation_acc']:.4f}\"],\n    ['肝脏', f\"{final_metrics['liver_auc']:.4f}\", f\"{final_metrics['liver_acc']:.4f}\"],\n    ['肾脏', f\"{final_metrics['kidney_auc']:.4f}\", f\"{final_metrics['kidney_acc']:.4f}\"],\n    ['脾脏', f\"{final_metrics['spleen_auc']:.4f}\", f\"{final_metrics['spleen_acc']:.4f}\"]\n]\n\nprint(\"\\n最终验证结果:\")\nprint(tabulate(metrics_table, headers='firstrow', tablefmt='grid'))","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-22T15:01:06.622Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 可视化训练结果","metadata":{}},{"cell_type":"code","source":"def load_best_model():\n    \"\"\"加载最佳模型\"\"\"\n    best_model = AbdominalTraumaModel().to(device)\n    best_model_path = os.path.join(OUTPUT_DIR, 'best_model_fold0.pt')\n    \n    if os.path.exists(best_model_path):\n        best_model.load_state_dict(torch.load(best_model_path))\n        print(f\"加载最佳模型: {best_model_path}\")\n    else:\n        print(f\"未找到最佳模型，使用当前模型\")\n        best_model = model\n        \n    return best_model\n\ndef visualize_predictions(model, val_loader, num_samples=3):\n    \"\"\"可视化模型预测结果\"\"\"\n    model.eval()\n    \n    samples_visualized = 0\n    \n    with torch.no_grad():\n        for batch in val_loader:\n            if samples_visualized >= num_samples:\n                break\n                \n            # 获取数据\n            patient_id = batch['patient_id'][0]\n            images = batch['images'].to(device)\n            masks = batch['masks'].to(device)\n            extravasation_features = batch['extravasation_features'].to(device)\n            aortic_hu = batch['aortic_hu'].to(device)\n            \n            # 获取标签\n            bowel_gt = batch['bowel'].cpu().numpy()[0][0]\n            extravasation_gt = batch['extravasation'].cpu().numpy()[0][0]\n            kidney_gt = batch['kidney'].cpu().numpy()[0]\n            liver_gt = batch['liver'].cpu().numpy()[0]\n            spleen_gt = batch['spleen'].cpu().numpy()[0]\n            \n            # 前向传播\n            outputs = model(images, masks, extravasation_features, aortic_hu)\n            \n            # 获取预测结果\n            bowel_pred = torch.sigmoid(outputs['bowel']).cpu().numpy()[0][0] > 0.5\n            extravasation_pred = torch.sigmoid(outputs['extravasation']).cpu().numpy()[0][0] > 0.5\n            kidney_pred = torch.argmax(outputs['kidney'], dim=1).cpu().numpy()[0]\n            liver_pred = torch.argmax(outputs['liver'], dim=1).cpu().numpy()[0]\n            spleen_pred = torch.argmax(outputs['spleen'], dim=1).cpu().numpy()[0]\n            \n            # 可视化\n            plt.figure(figsize=(20, 10))\n            \n            # 选择4个代表性切片进行显示\n            slice_indices = np.linspace(0, images.size(1) - 1, 4, dtype=int)\n            \n            for i, idx in enumerate(slice_indices):\n                plt.subplot(2, 4, i + 1)\n                plt.imshow(images[0, idx].cpu().numpy(), cmap='gray')\n                plt.title(f'切片 {idx}')\n                plt.axis('off')\n                \n                # 显示掩码叠加\n                plt.subplot(2, 4, i + 5)\n                \n                # 创建RGB掩码\n                mask_rgb = np.zeros((images.shape[2], images.shape[3], 3))\n                \n                # 肝脏 - 红色\n                if masks[0, idx, :, :, 0].sum() > 0:\n                    mask_rgb[:, :, 0] += masks[0, idx, :, :, 0].cpu().numpy() * 0.5\n                \n                # 脾脏 - 绿色\n                if masks[0, idx, :, :, 1].sum() > 0:\n                    mask_rgb[:, :, 1] += masks[0, idx, :, :, 1].cpu().numpy() * 0.5\n                \n                # 肾脏 - 蓝色\n                if masks[0, idx, :, :, 2].sum() > 0:\n                    mask_rgb[:, :, 2] += masks[0, idx, :, :, 2].cpu().numpy() * 0.5\n                \n                # 肠道 - 黄色\n                if masks[0, idx, :, :, 3].sum() > 0:\n                    mask_rgb[:, :, 0] += masks[0, idx, :, :, 3].cpu().numpy() * 0.5\n                    mask_rgb[:, :, 1] += masks[0, idx, :, :, 3].cpu().numpy() * 0.5\n                \n                # 外渗 - 紫色\n                if extravasation_features[0, idx, 0].sum() > 0:\n                    mask_rgb[:, :, 0] += extravasation_features[0, idx, 0].cpu().numpy() * 0.5\n                    mask_rgb[:, :, 2] += extravasation_features[0, idx, 0].cpu().numpy() * 0.5\n                \n                # 叠加到原图\n                img_rgb = np.stack([images[0, idx].cpu().numpy()] * 3, axis=2)\n                overlay = img_rgb * 0.7 + mask_rgb * 0.3\n                \n                plt.imshow(np.clip(overlay, 0, 1))\n                plt.title(f'掩码叠加')\n                plt.axis('off')\n            \n            # 显示预测结果\n            plt.suptitle(f'患者 {patient_id} - aortic_hu: {aortic_hu.item():.1f}\\n'\n                         f'肠道: 真实={bowel_gt}, 预测={bowel_pred} | '\n                         f'外渗: 真实={extravasation_gt}, 预测={extravasation_pred} | '\n                         f'肾脏: 真实={kidney_gt}, 预测={kidney_pred} | '\n                         f'肝脏: 真实={liver_gt}, 预测={liver_pred} | '\n                         f'脾脏: 真实={spleen_gt}, 预测={spleen_pred}', \n                         fontsize=16)\n            \n            plt.tight_layout()\n            plt.savefig(os.path.join(OUTPUT_DIR, f'prediction_{patient_id}.png'))\n            plt.show()\n            \n            samples_visualized += 1\n\n# 加载最佳模型\nbest_model = load_best_model()\n\n# 准备验证数据\n_, val_loader = prepare_data(meta_df, NUM_FOLDS, 0)\n\n# 可视化预测结果\nvisualize_predictions(best_model, val_loader, num_samples=3)","metadata":{"trusted":true,"execution":{"execution_failed":"2025-05-22T15:01:06.622Z"}},"outputs":[],"execution_count":null}]}