{"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"}],"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":false}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"# 简介\nRSNA 腹部创伤检测 AI 挑战赛旨在解决医疗保健中的一个关键问题：利用计算机断层扫描（CT）对腹部创伤进行快速准确的诊断。创伤是全球死亡的主要原因之一，CT 扫描已成为评估疑似腹部损伤患者的必要手段，因为它们能够提供详细的横断面图像。然而，解释 CT 扫描以诊断腹部创伤可能很复杂且耗时，尤其是在存在多处创伤或微妙的活跃出血区域时。\n\n本项目旨在复现Shen等人发表在《World Journal of Emergency Surgery》(2024)上题为\"The application of deep learning in abdominal trauma diagnosis by CT imaging\"的研究论文。该论文提出了一种基于深度学习的方法，用于从CT扫描图像中自动检测和诊断腹部创伤，包括肝脏、脾脏、肾脏和肠道损伤以及腹部渗出。","metadata":{},"attachments":{"d601fc64-bb38-4fe2-9f3d-a649e8c3c8b5.png":{"image/png":"iVBORw0KGgoAAAANSUhEUgAAALgAAAB+CAIAAAAY4Ew5AAAMwklEQVR4Ae2c/0sbyRvH709RCLRQ6IHggT8I/aHgD4K/9RdBlhKla44YEsQSKhIIoVJrLfnYCxZTqJZPyjV+qoKXepWmcvkEWq3U2lqpyklaCRdBEmxAojDH7uzszn6JGbtd75J9ipjdyTNfnvfzcuaZ2TQ/IPgHCjAo8AODDZiAAghAAQiYFABQmGQCIwAFGGBSAEBhkgmMABRggEkBAIVJJjACUIABJgUAFCaZwAhAAQaYFABQmGQCIwAFGGBSAEBhkgmMABRggEkBAIVJJjACUIABJgUAFCaZwAhAAQaYFABQmGQCIwAFGGBSAEBhkgmMABRggEkBAIVJJjACUIABJgUAFCaZwAhAAQaYFABQmGQCIwAFGGBSAEBhkgmMABRggEkBAIVJJjACUIABJgUAFCaZwAhAAQaYFABQmGQCIytAOcyk5xPJD4UjkLd2FDgNKEeHm+nJ237/db+Lv+a/7r/zMPkhd6jV4k2gqa7eIfy0/5rTvgn31aoAKyi5Z4HL58XwYwjk3+dbBp7RPHy+3yqbeV9Uqywwbq0CTKBsRq+ck8kwuLhw9YnCSm7W1SjaXA6u6KYbbfdwXy0KMICS/bXdQSYJR1N7cDIxN5+Ym38Y7GhUyl2Jr5TLXwu5AkBCCVL9l5VBycU6pJyjvsHzrEC7XHjmvUgmGM8z6Z3DQk76ty+zopSR91SvBRoy0sHhvmQDyBFJ/snXyqBkxtoIKG33dzRj/fBYyG2Fn/8kpdXnhY9MP63jGck86SE8kaaIjVjeOvaZarfw5pG39Ufa4EJzx81EljKByzNXoDIo1IziaPTNV9z0mgLlKDf9M9k0adhytN3fOnN5oEOiQGVQEJ2j1Dvqfmxp949Pp7fKrQhGoKRDzZeaVT9NF+X8pr5pYEkaTiYqz14XLvvGp+fmE49utjeR2eUn/wujRYr4Aq8WKsAACkLyRkazcFy83BF6oj1YMwJF60BhwYt3RnX1jtYxMlF8necJPa2/kEKEUGHeQ1ai9piyvdI2CvdWKsAECkLocGd+oK1BA4p02+SdphKIyqBsjbcSIBp9SSU9TvpJ+65p9czxJkC6vjYrZ8hWygJtaxVgBUWq9zX3Pjl5+9qVRvInLoWWWhQqgFJIen4iS8mlO++pY34qa75wUbVOXVK6UxJkrSdwb6kCpwSFGsvhVvL2VSXxlHcuJ4Ly+X4boUSXnFKgEBtNPlvvqANQqBCc5WVFUHLvxeO1xNz8/7d0s/5RekCeWnxJPO4TQNlUdtpNngVlzcEVKVCuhEin+HBP+f2HfhBnKZd9+6oIyoqCws/zWlJOAwqdwDbf+mAg+YKX5CgdjyFnNRDonyyqCApSZoj6pquPPhzKWcVR4UWghYTWUWHpoRLYurbxTbkR2ndq19N8k35OVJjuaWpu6fAEx6ffAUG0ZGd3XRkU9O5OM50rnG8QT0TogxBHnaPjMdn4KGDJ+QSdwNY7zjVozlQuNQfT2OPNX+RzFEdj+82Hc/OJ/417Wi4QHBs8C9pJ7eyksndPDKCUP0eR4udouk4lHAagKGtKmSyV5DfoaOsxlSATPqRajX5qL23vsJ2990ygiOcoydvXWqjjVDF4joZW3+SbfdWwTYEitHSYmQuon/U4zl3qCM19hslEJfTZ3rCCIo9KfqibUx4Oy29+zwv5ibPhs+Xv2RO0xaDAqUFhaBNMalABAKUGg2qFSwCKFarWYJsASg0G1QqXABQrVK3BNgGUGgyqFS4BKFaoWoNtAig1GFQrXAJQrFC1BtsEUGowqFa4BKBYoWoNtgmg1GBQrXAJQLFC1RpsE0CpwaBa4RKAYoWqNdgmgFKDQbXCJQDFClVrsE0ApQaDaoVL3wTKcam4X8jvF4ol4yGVDjLrqVQ6lUovbeTL2CCESpk1wSaVWv24Vzo2bkpfWtrbWCWN679qUG+PSsW8OFr2LhAqFf+UxpZOre0cGPlARMCNk9/lJDEYV3UVfQsoO5N9nJPnnHzkld7Z4vqD/k7xXWzDdXnDyT2tXWk7EfJKBtjYc2uBfO2O1li+P9iYHVTX6vIOzm4bhVGuU1wcEobKOftnv8iFJ13k380MenAV5XfP4Mx6Xl3ry4yfdlO+dgUmlrT/CVJdsyrvTg9KJi4LpAdlfVwIZKd/LPE2k9/PrM+N+Xmec3ojr4qUPHsLQSEG3aHY8p+FfHZjMSKy5R5dPqCsNJfHmdkbQq2ewVj6Yza/n91MxQbdQslwkm5cVa2YGu2WQsgEym7ilmjv7g1OJvC8NTcZ9LkF1NxqlDEovlHJTDD+LX431NMl+rtUdkiq8VXPzSlBOc7E/TzXFYpG+g1mlO1Yr5Pnbszs0uvIJ7GwNyZ/q1dpcUTQfeglreXuE6HB7vG1ctLhWt13U3QthBv3z+waVjtIhV08545Gh9lmlC8zA1081xWIfVR1go4Ly2FxJrtBdYRB0XVdfCWiGXyumYAMB1hFhacDZfdpgHPyA0+zu1MGoKyPuzlnX3xb435h9fFYNBJflZTDa0Ff7JPa7Hgt6uK5rtFlGjLKJL8Uj0bGFrSNb8d8POccXaYsyWVx9Z5XmMyWisthFlCKaYEnPpjQLZQIoePtWC/POd3Rt6T5MqAgtDYhODK2Sgxr4/U0oOw9D3aJEwZCRqCIMfMpM0cZgVYiXTzniq7r3hbDqQNIZ6YqOE6FnTxn2Om7aLeT7763UkKICZTiy2GnOLATSZ1Ikf86Ww6UE4akGnqV3bCDgmeC/riYchqAsv886OS58Gt0vLf8INTrEv46O939wQepHJ1tYn2HXtJlWDPc5gkJh1ZasiIMPCXBky3wBOCSkh4mUESwuHsrchsVLoxBKe08ESbd3v9qp74Krf3r32YFBWeFvZOS/wagYOHuxsWU093t6+vx9XULmSzP9UbX5SwVm4VfGyiTGuWcvH9KF3WN6f5KPDIWHQr0CCy6/Q/W1AmFYI0zHpk5JlBesfUuDwY74gqEI2NR6WdUSHv5vuGpDf2Q5HpVesEGysHriFvICtfJtGwAirwbco8sZsl8UcouDgtpIF4CBI1wwmsICmOocITwXsYVCE+t5MiopBhkxJyUSidZQMHJshZTDKWCwlj0dzJV0MOQ98bC4uUdiKgn0SpFQz1sFlCUrFCuawCKJByV7mFrnKU6RxbxXxk2MwKllBR2Q9pQyV3qLkp/bSTu6jYjKJu4IexcEtSpCQsoCGP6RH2Yo6dBHjl+S7XrKRWzhkPSDb0KCxhAwVlh+DU9nZYFpWKWmp0ZcPIc9ecui4bXC/3ZjGxgdCGdp8m18omQsC+bUsWbCZRPsR5x004mQ11v2OAkUHAV7ZB0DVVlQUVQxD9QIS0Vcg7lR8pVhZKJN6Ln5bN9MU7uiXdYILx71G+DVcnyKbRUZTbilsrJd3qooZJUSSy8tVAuBZL25yOLcjqlHgRGsIdkachgRiEVVEMihVX+WhmUhUGV6BIrelBQQTxv1Z2jHG9MCEcdoQXyNSoiN/ywtBQR/fZ+EzZNhhtdwSQrDuMWvaDgmvlZYQohC9bKBE0zucY5dQVQEMKPJrrVc6c0PukcxTshf/lceVBOu4YSCf7VrxVBMR69wdKDUAkfSt6I7yjTd3FzUtgucsMppQwfp9IH9seF9F3hmDw4qxx2Fd/Go5HY8l/SADYfiElxOJWnU9d8SsiyDU75VMNmWnoQQjhnd/K94yuqXuQnU/RpcjlQpCFRSKnGUq033xMUhIrSUTffPyzuFIb9+CmJ9iHO7pRIj2Qm7iq1Z//SIiIczOB/B68j+FmdK4Abjw4FOoUHK9qMRB8KVlAQQpnn+PkRx3sHhoR9bzjYh3tRbfIRkpaeLreyHAsTmBc/EDWelvQjq56S7wsKQqi4OTXSi49PhE2ju3dI99xVUEdr5tduKaXcSHWYVsqmIwHpbEbckXb6QhMpZRIqJ/spQBE+/aDthXP1DU+tqeYYGRR6Y+zkuS53z8BIXDkeKDei6iv/RlAqOlrKn/SBFVJd+lxLmU+KlIqGHwRB5NMwFn+nG3YhbzwG4oFtXq0CxTYC2sVRAMUukTbpJ4BiUkC7VAdQ7BJpk34CKCYFtEt1AMUukTbpJ4BiUkC7VAdQ7BJpk34CKCYFtEt1AMUukTbpJ4BiUkC7VAdQ7BJpk34CKCYFtEt1AMUukTbpJ4BiUkC7VAdQ7BJpk34CKCYFtEt1AMUukTbpJ4BiUkC7VAdQ7BJpk34CKCYFtEt1AMUukTbpJ4BiUkC7VAdQ7BJpk34CKCYFtEt1AMUukTbpJ4BiUkC7VAdQ7BJpk34CKCYFtEt1AMUukTbpJ4BiUkC7VAdQ7BJpk34CKCYFtEt1AMUukTbpJ4BiUkC7VP8b0fc8kN6TPXUAAAAASUVORK5CYII="}}},{"cell_type":"markdown","source":"# 环境设置","metadata":{}},{"cell_type":"code","source":"# 导入必要的库\nimport gc\nimport os\nimport numpy as np\nimport pandas as pd\nimport matplotlib.pyplot as plt\nimport 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, EfficientNetB1\nfrom sklearn.model_selection import KFold\nfrom sklearn.metrics import roc_auc_score, accuracy_score, precision_score, recall_score, confusion_matrix\nimport pydicom\nimport random\nimport glob\nimport pydicom\nimport nibabel as nib\nfrom skimage.transform import resize\nfrom concurrent.futures import ThreadPoolExecutor\nimport multiprocessing\n\nimport matplotlib as mpl\nimport matplotlib.font_manager as fm","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:41:41.610326Z","iopub.execute_input":"2025-04-22T15:41:41.610543Z","iopub.status.idle":"2025-04-22T15:42:09.620301Z","shell.execute_reply.started":"2025-04-22T15:41:41.610517Z","shell.execute_reply":"2025-04-22T15:42:09.61947Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 检查GPU是否可用\nprint(\"TensorFlow版本:\", tf.__version__)\nprint(\"GPU是否可用:\", tf.config.list_physical_devices('GPU'))\n\n# 限制TensorFlow最多占用显存而非一把吃掉\ngpus = tf.config.experimental.list_physical_devices('GPU')\nif gpus:\n    try:\n        for gpu in gpus:\n            tf.config.experimental.set_memory_growth(gpu, True)\n    except Exception as e:\n        print(e)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:09.621708Z","iopub.execute_input":"2025-04-22T15:42:09.622211Z","iopub.status.idle":"2025-04-22T15:42:10.633095Z","shell.execute_reply.started":"2025-04-22T15:42:09.622193Z","shell.execute_reply":"2025-04-22T15:42:10.632144Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 设置随机种子以确保结果可复现\ndef set_seed(seed=42):\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n    os.environ['PYTHONHASHSEED'] = str(seed)\n\nset_seed()","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:10.633875Z","iopub.execute_input":"2025-04-22T15:42:10.634094Z","iopub.status.idle":"2025-04-22T15:42:10.662774Z","shell.execute_reply.started":"2025-04-22T15:42:10.634075Z","shell.execute_reply":"2025-04-22T15:42:10.662028Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 数据获取与预览","metadata":{}},{"cell_type":"code","source":"# 定义数据路径\nDATA_DIR = '/kaggle/input/rsna-2023-abdominal-trauma-detection'\nTRAIN_CSV = os.path.join(DATA_DIR, 'train_2024.csv')\nTRAIN_IMAGES = os.path.join(DATA_DIR, 'train_images')\nTRAIN_META = os.path.join(DATA_DIR, 'train_series_meta.csv')\nOUTPUT_DIR = '/kaggle/working'\nCACHE_DIR = os.path.join(OUTPUT_DIR, 'cache')\nSEGMENTATION_DIR = \"/kaggle/input/rsna-2023-abdominal-trauma-detection/segmentations\"\nTRAIN_DICOM_TAGS = '/kaggle/input/rsna-2023-abdominal-trauma-detection/train_dicom_tags.parquet' \n\n# 图像尺寸统一为论文中的224×224\nIMG_SIZE = (224, 224)\nNUM_ORGANS = 5  # 肝脏, 脾脏, 肾脏, 肠道, 外渗\n\n# 创建输出目录（如果不存在）\nos.makedirs(OUTPUT_DIR, exist_ok=True)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:10.663538Z","iopub.execute_input":"2025-04-22T15:42:10.663782Z","iopub.status.idle":"2025-04-22T15:42:10.678987Z","shell.execute_reply.started":"2025-04-22T15:42:10.663754Z","shell.execute_reply":"2025-04-22T15:42:10.678315Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def load_and_process_data():\n    print(\"Loading training data...\")\n    train_df = pd.read_csv(TRAIN_CSV)\n    train_df['patient_id'] = train_df['patient_id'].astype(str)\n    organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n    organ_status_map = {}\n    for _, row in train_df.iterrows():\n        pid = row['patient_id']\n        statuses = {}\n        for organ in organs:\n            status_columns = [col for col in train_df.columns if col.startswith(f'{organ}_')]\n            organ_status = {col.split('_')[1]: row[col] for col in status_columns}\n            statuses[organ] = organ_status\n        organ_status_map[pid] = statuses\n\n    train_series_meta = pd.read_csv(TRAIN_META)\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    segmentation_map = {}\n    for fpath in glob.glob(os.path.join(SEGMENTATION_DIR, \"*.nii\")):\n        name = os.path.splitext(os.path.basename(fpath))[0]\n        if name in series_to_patient:\n            segmentation_map[series_to_patient[name]] = fpath\n\n    del train_series_meta, series_to_patient\n    gc.collect()\n\n    print(f\"Preprocessing complete. Processed {len(train_df)} records.\")\n    return train_df, organ_status_map, segmentation_map","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:10.680642Z","iopub.execute_input":"2025-04-22T15:42:10.680845Z","iopub.status.idle":"2025-04-22T15:42:10.695605Z","shell.execute_reply.started":"2025-04-22T15:42:10.680829Z","shell.execute_reply":"2025-04-22T15:42:10.695056Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_data(train_df, organ_status_map=None):\n    \"\"\"Visualize organ status distributions and missing values\"\"\"\n    organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n\n    # 1. Organ status pie charts\n    plt.figure(figsize=(20, 15))\n    for i, organ in enumerate(organs):\n        plt.subplot(2, 3, i+1)\n        status_data = {}\n        total = len(train_df)\n        for label in ['healthy', 'injury', 'low', 'high']:\n            col = f\"{organ}_{label}\"\n            if col in train_df.columns:\n                count = train_df[col].sum()\n                status_data[label.capitalize() if label == 'healthy' else ('Low-grade Injury' if label=='low' else 'High-grade Injury' if label=='high' else 'Injury')] = count\n        if not status_data:\n            continue\n        sizes = list(status_data.values())\n        labels = list(status_data.keys())\n        sizes_pct = [s/total*100 for s in sizes]\n        labels_pct = [f\"{l}: {p:.1f}%\" for l, p in zip(labels, sizes_pct)]\n        colors = plt.cm.Paired(np.arange(len(sizes)) / len(sizes))\n        explode = [0.1 if 'Injury' in l else 0 for l in labels]\n        plt.pie(sizes, explode=explode, labels=labels_pct, colors=colors,\n                autopct='%1.1f%%', shadow=True, startangle=90)\n        plt.axis('equal')\n        plt.title(f\"{organ.capitalize()} Status Distribution\", fontsize=15)\n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'organ_status_distribution.png'), dpi=300)\n    plt.show()\n\n    # 2. Injury percentage bar chart\n    percentages = []\n    for organ in organs:\n        inj = 0\n        for label in ['injury', 'low', 'high']:\n            col = f\"{organ}_{label}\"\n            if col in train_df.columns:\n                inj += train_df[col].sum()\n        percentages.append(inj/len(train_df)*100)\n    plt.figure(figsize=(12, 8))\n    bars = plt.bar(organs, percentages, color=plt.cm.viridis(np.linspace(0, 1, len(organs))))\n    for bar, pct in zip(bars, percentages):\n        plt.text(bar.get_x()+bar.get_width()/2, bar.get_height()+0.5,\n                 f\"{pct:.1f}%\", ha='center', va='bottom', fontsize=12)\n    plt.xlabel('Organ', fontsize=14)\n    plt.ylabel('Injury Percentage (%)', fontsize=14)\n    plt.title('Injury Percentage by Organ', fontsize=16)\n    plt.ylim(0, max(percentages)*1.2)\n    plt.grid(axis='y', linestyle='--', alpha=0.7)\n    plt.savefig(os.path.join(OUTPUT_DIR, 'organ_injury_percentages.png'), dpi=300)\n    plt.show()\n\n    # 3. Severity distribution pie charts\n    severity_organs = [o for o in organs if any(f\"{o}_{lab}\" in train_df.columns for lab in ['low','high'])]\n    if severity_organs:\n        plt.figure(figsize=(14, 10))\n        for i, organ in enumerate(severity_organs):\n            plt.subplot(2, 3, i+1)\n            sizes = []\n            labels = []\n            for lab, name in [('low', 'Low-grade Injury'), ('high', 'High-grade Injury')]:\n                col = f\"{organ}_{lab}\"\n                if col in train_df.columns:\n                    cnt = train_df[col].sum()\n                    sizes.append(cnt)\n                    labels.append(f\"{name}: {cnt/ (sum(sizes)) * 100:.1f}%\")\n            if not sizes or sum(sizes)==0:\n                continue\n            plt.pie(sizes, labels=labels, colors=['#ff9999','#ff3333'],\n                    autopct='%1.1f%%', shadow=True, startangle=90)\n            plt.axis('equal')\n            plt.title(f\"{organ.capitalize()} Injury Severity Distribution\", fontsize=15)\n        plt.tight_layout()\n        plt.savefig(os.path.join(OUTPUT_DIR, 'injury_severity_distribution.png'), dpi=300)\n        plt.show()\n\n    # 4. Missing values\n    missing = train_df.isnull().sum()\n    if missing.sum() > 0:\n        plt.figure(figsize=(10, 6))\n        cols = missing[missing>0].index\n        vals = missing[missing>0].values\n        plt.bar(cols, vals, color='crimson')\n        plt.xlabel('Column')\n        plt.ylabel('Missing Value Count')\n        plt.title('Missing Values in Dataset')\n        plt.xticks(rotation=45)\n        plt.tight_layout()\n        plt.savefig(os.path.join(OUTPUT_DIR, 'missing_values.png'), dpi=300)\n        plt.show()\n    else:\n        print(\"No missing values in the dataset\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:10.696272Z","iopub.execute_input":"2025-04-22T15:42:10.696427Z","iopub.status.idle":"2025-04-22T15:42:10.711664Z","shell.execute_reply.started":"2025-04-22T15:42:10.696414Z","shell.execute_reply":"2025-04-22T15:42:10.711123Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 加载并处理数据\ntrain_df, organ_status_map, segmentation_map = load_and_process_data()\ngc.collect()\nvisualize_data(train_df)\n# 示例：打印部分 segmentation_map 信息\nprint(\"部分 segmentation_map 信息：\")\nfor pid, seg_file in list(segmentation_map.items())[:5]:\n    print(f\"Patient ID: {pid} -> Segmentation file: {seg_file}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:10.712457Z","iopub.execute_input":"2025-04-22T15:42:10.712722Z","iopub.status.idle":"2025-04-22T15:42:15.774879Z","shell.execute_reply.started":"2025-04-22T15:42:10.7127Z","shell.execute_reply":"2025-04-22T15:42:15.774111Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"可以看到：\n* 数据集总共包含3147名患者\n* 数据集质量良好，没有缺失值\n* 所有器官都存在明显的类别不平衡问题，健康样本占绝大多数\n  \n各器官的健康与损伤比例：\n* 肠道(Bowel)：损伤比例最低，仅2.3%（71例）\n* 渗出(Extravasation)：6.8%（215例）的病例出现渗出\n* 肾脏(Kidney)：6.9%的损伤率，包括4.48%（141例）低度损伤和2.41%（76例）高度损伤\n* 肝脏(Liver)：10.8%的损伤率，包括8.67%（273例）低度损伤和2.13%（67例）高度损伤\n* 脾脏(Spleen)：损伤比例最高，达11.8%，包括6.67%（210例）低度损伤和5.15%（162例）高度损伤\n  \n损伤严重程度分布：\n* 肾脏：在损伤病例中，65%为低度损伤，35%为高度损伤\n* 肝脏：损伤以低度为主，占80.3%，高度损伤占19.7%\n* 脾脏：损伤程度分布最均衡，56.5%为低度损伤，43.5%为高度损伤","metadata":{}},{"cell_type":"markdown","source":"# 数据分割","metadata":{}},{"cell_type":"code","source":"def segment_slice(image, model, img_size=(224, 224)):\n    \"\"\"\n    使用训练好的U-Net模型对单张切片进行分割\n    \n    Args:\n        image: 单张CT切片图像 (numpy array)，形状应为 (height, width) 或 (height, width, channels)\n               其中 channels 可以为 1 或 3\n        model: 训练好的U-Net模型 (tensorflow.keras.Model)\n        img_size: 图像大小（tuple），用于处理不同尺寸的图像\n\n    Returns:\n        pred_mask: 分割后的掩码 (numpy array)，形状为 (height, width)\n    \"\"\"\n    try:\n        # 1. 确保图像是正确的形状\n        if len(image.shape) == 2:  # 如果是单通道的2D图像，添加通道维度\n            image = np.expand_dims(image, axis=-1)\n        elif len(image.shape) == 3 and image.shape[-1] > 3: # 如果通道数大于3 报错\n            raise ValueError(\"Image has more than 3 channels. It should have 1 or 3\")\n        \n        # 打印图像统计信息\n        print(f\"Input image stats: shape={image.shape}, max={np.max(image):.4f}, min={np.min(image):.4f}, mean={np.mean(image):.4f}\")\n        \n        # 2. 确保图像是3通道的\n        if image.shape[-1] == 1:\n            image = np.repeat(image, 3, axis=-1)\n\n        # 3. 缩放到模型所需的尺寸\n        if image.shape[0] != img_size[0] or image.shape[1] != img_size[1]:\n            image = cv2.resize(image, img_size)  # 使用 cv2 缩放图像\n\n        # 4. 转换为 float32\n        image = image.astype(np.float32)\n\n        # 5. 归一化像素值到 [0, 1] 范围 (如果需要)\n        max_val = np.max(image)\n        if max_val != 0:\n            image = image / max_val\n\n        # 6. 添加批次维度\n        input_image = np.expand_dims(image, axis=0)\n\n        # 7. 使用U-Net模型进行分割\n        pred_mask = model.predict(input_image, verbose=0)[0]\n        \n        # 打印预测掩码统计信息\n        print(f\"Pred mask stats: shape={pred_mask.shape}, max={np.max(pred_mask):.4f}, min={np.min(pred_mask):.4f}, mean={np.mean(pred_mask):.4f}\")\n\n        # 8. 从模型输出中提取单通道掩码\n        pred_mask = pred_mask[:, :, 0]\n\n        # 9. 确保输出掩码形状正确\n        assert len(pred_mask.shape) == 2, f\"Mask should be 2D, but got shape {pred_mask.shape}\"\n        return pred_mask\n        \n    except Exception as e:\n        print(f\"Error in segment_slice: {str(e)}\")\n        import traceback\n        print(traceback.format_exc())\n        # 返回一个空掩码\n        return np.zeros(img_size, dtype=np.float32)","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:15.775888Z","iopub.execute_input":"2025-04-22T15:42:15.776135Z","iopub.status.idle":"2025-04-22T15:42:15.784577Z","shell.execute_reply.started":"2025-04-22T15:42:15.776105Z","shell.execute_reply":"2025-04-22T15:42:15.78388Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 数据增强","metadata":{}},{"cell_type":"code","source":"def augment_data(image, mask=None):\n    \"\"\"\n    对图像和掩码（可选）应用数据增强。\n    \"\"\"\n    # 确保图像形状正确\n    if len(image.shape) == 3 and image.shape[-1] == 1:\n        image = image[:, :, 0]  # 如果是单通道的3D图像，转为2D\n    \n    original_shape = image.shape\n    \n    # 随机水平翻转\n    if random.random() > 0.5:\n        image = np.fliplr(image)\n        if mask is not None:\n            mask = np.fliplr(mask)\n    \n    # 随机垂直翻转\n    if random.random() > 0.5:\n        image = np.flipud(image)\n        if mask is not None:\n            mask = np.flipud(mask)\n    \n    # 随机旋转\n    angle = random.uniform(-15, 15)\n    M = cv2.getRotationMatrix2D((IMG_SIZE[0]/2, IMG_SIZE[1]/2), angle, 1)\n    image = cv2.warpAffine(image, M, IMG_SIZE)\n    if mask is not None:\n        mask = cv2.warpAffine(mask, M, IMG_SIZE, flags=cv2.INTER_NEAREST)\n    \n    # 模糊\n    if random.random() > 0.5:\n        image = cv2.GaussianBlur(image, (5, 5), 0)\n    \n    # 高斯噪声\n    if random.random() > 0.5:\n        if len(original_shape) == 2:  # 2D图像\n            row, col = image.shape\n            mean = 0\n            var = 0.1\n            sigma = var**0.5\n            gauss = np.random.normal(mean, sigma, (row, col))\n            image = image + gauss\n        elif len(original_shape) == 3:  # 3D图像\n            row, col, ch = image.shape\n            mean = 0\n            var = 0.1\n            sigma = var**0.5\n            gauss = np.random.normal(mean, sigma, (row, col, ch))\n            image = image + gauss\n    \n    # 确保值范围在[0,1]\n    image = np.clip(image, 0, 1)\n    \n    if mask is not None:\n        return image, mask\n    else:\n        return image\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:15.785348Z","iopub.execute_input":"2025-04-22T15:42:15.785796Z","iopub.status.idle":"2025-04-22T15:42:15.840809Z","shell.execute_reply.started":"2025-04-22T15:42:15.785778Z","shell.execute_reply":"2025-04-22T15:42:15.840096Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 构建U-net模型","metadata":{}},{"cell_type":"code","source":"def build_unet(input_shape):\n    \"\"\"\n    构建基于EfficientNetB0的2D U-Net模型，用于单张CT切片分割。\n    论文中使用的架构。\n    input_shape应为(height, width, channels)\n    \"\"\"\n    # 加载预训练的EfficientNetB0作为编码器\n    efficientnet = EfficientNetB0(include_top=False, weights='imagenet', input_shape=input_shape)\n    \n    # 获取编码器的中间层特征图\n    c1 = efficientnet.get_layer('block2b_add').output  # 64x64\n    c2 = efficientnet.get_layer('block3b_add').output  # 32x32\n    c3 = efficientnet.get_layer('block5c_add').output  # 16x16\n    c4 = efficientnet.get_layer('block6d_add').output  # 8x8\n    b0 = efficientnet.output  # 瓶颈: 8x8\n    \n    # 解码器部分 - 对称上采样\n    # 步骤1: 8x8 -> 16x16\n    up1 = layers.Conv2DTranspose(512, (3,3), strides=2, padding='same', activation='relu')(b0)\n    merge1 = layers.concatenate([c3, up1], axis=3)\n    conv1 = layers.Conv2D(512, 3, activation='relu', padding='same')(merge1)\n    conv1 = layers.Conv2D(512, 3, activation='relu', padding='same')(conv1)\n    \n    # 步骤2: 16x16 -> 32x32\n    up2 = layers.Conv2DTranspose(256, (3,3), strides=2, padding='same', activation='relu')(conv1)\n    merge2 = layers.concatenate([c2, up2], axis=3)\n    conv2 = layers.Conv2D(256, 3, activation='relu', padding='same')(merge2)\n    conv2 = layers.Conv2D(256, 3, activation='relu', padding='same')(conv2)\n    \n    # 步骤3: 32x32 -> 64x64\n    up3 = layers.Conv2DTranspose(128, (3,3), strides=2, padding='same', activation='relu')(conv2)\n    merge3 = layers.concatenate([c1, up3], axis=3)\n    conv3 = layers.Conv2D(128, 3, activation='relu', padding='same')(merge3)\n    conv3 = layers.Conv2D(128, 3, activation='relu', padding='same')(conv3)\n    \n    # 步骤4: 64x64 -> 128x128\n    up4 = layers.Conv2DTranspose(64, (3,3), strides=2, padding='same', activation='relu')(conv3)\n    \n    # 步骤5: 128x128 -> 256x256 (如果需要)\n    up5 = layers.Conv2DTranspose(32, (3,3), strides=2, padding='same', activation='relu')(up4)\n    \n    # 处理尺寸不匹配问题\n    # 检查输出尺寸是否需要裁剪\n    if up5.shape[1] != input_shape[0] or up5.shape[2] != input_shape[1]:\n        # 计算需要裁剪的边缘大小\n        crop_height = int((up5.shape[1] - input_shape[0]) / 2) if up5.shape[1] > input_shape[0] else 0\n        crop_width = int((up5.shape[2] - input_shape[1]) / 2) if up5.shape[2] > input_shape[1] else 0\n        \n        if crop_height > 0 or crop_width > 0:\n            up5 = layers.Cropping2D(cropping=((crop_height, crop_height), \n                                             (crop_width, crop_width)))(up5)\n    \n    # 输出层 - 单通道分割掩码\n    outputs = layers.Conv2D(1, 1, activation='sigmoid')(up5)\n    \n    # 创建模型\n    model = models.Model(inputs=efficientnet.input, outputs=outputs)\n    \n    return model\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:15.84157Z","iopub.execute_input":"2025-04-22T15:42:15.842082Z","iopub.status.idle":"2025-04-22T15:42:15.861364Z","shell.execute_reply.started":"2025-04-22T15:42:15.842065Z","shell.execute_reply":"2025-04-22T15:42:15.860844Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 构建2.5D模型","metadata":{}},{"cell_type":"code","source":"def build_25d_classifier(input_shape, num_organs=5):\n    \"\"\"\n    构建2.5D分类模型，使用32张连续切片作为输入。\n    input_shape应为单张图像的形状(height, width, channels)\n    \"\"\"\n    # 检查并修正输入通道数\n    if input_shape[-1] == 1:\n        # 如果是单通道，修改为3通道以适应EfficientNet要求\n        efficientnet_input_shape = (input_shape[0], input_shape[1], 3)\n        print(f\"注意: 将输入形状从 {input_shape} 修改为 {efficientnet_input_shape} 以适应EfficientNet\")\n    else:\n        efficientnet_input_shape = input_shape\n    \n    # 输入层 - 32张连续切片\n    input_ct = layers.Input(shape=(32,) + input_shape, name='ct_input')\n    \n    # 如果是单通道输入，需要转换为三通道\n    if input_shape[-1] == 1:\n        # 将单通道转换为三通道\n        x = layers.TimeDistributed(layers.Conv2D(3, kernel_size=1, padding='same'))(input_ct)\n    else:\n        x = input_ct\n    \n    # 使用TimeDistributed包装EfficientNetB1，处理每张切片\n    # 加载预训练的EfficientNetB1，但移除顶层\n    base_model = EfficientNetB1(\n        include_top=False, \n        weights='imagenet', \n        input_shape=efficientnet_input_shape,  # 使用修正后的形状\n        pooling='avg'\n    )\n    base_model.trainable = False  # 冻结基础模型权重\n    \n    # 使用TimeDistributed应用到每张切片\n    x = layers.TimeDistributed(base_model)(x)\n    \n    # 双向LSTM提取时序特征\n    x = layers.Bidirectional(layers.LSTM(512, return_sequences=True))(x)\n    x = layers.Bidirectional(layers.LSTM(256, return_sequences=False))(x)\n    \n    # 全连接层\n    x = layers.Dense(512, activation='relu')(x)\n    x = layers.BatchNormalization()(x)\n    x = layers.Dropout(0.5)(x)\n    x = layers.Dense(256, activation='relu')(x)\n    x = layers.Dropout(0.3)(x)\n    \n    # 输出层 - 每个器官4个类别(healthy, injury, low, high)\n    output = layers.Dense(num_organs * 4, activation='sigmoid')(x)\n    output = layers.Reshape((num_organs, 4))(output)\n    \n    model = models.Model(inputs=input_ct, outputs=output)\n    return model\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:15.861953Z","iopub.execute_input":"2025-04-22T15:42:15.862184Z","iopub.status.idle":"2025-04-22T15:42:15.878103Z","shell.execute_reply.started":"2025-04-22T15:42:15.862162Z","shell.execute_reply":"2025-04-22T15:42:15.87748Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 多线程与缓存优化","metadata":{}},{"cell_type":"code","source":"def load_single_slice(args):\n    \"\"\"\n    加载单个切片的函数，用于并行处理\n    \n    Args:\n        args: 包含image_file, image_id, img_size, segmentation_data的元组\n        \n    Returns:\n        切片数据字典或None（如果处理失败）\n    \"\"\"\n    image_file, image_id, img_size, segmentation_data = args\n    \n    try:\n        dicom = pydicom.dcmread(image_file)\n        image = dicom.pixel_array.astype(np.float32)\n        \n        # 打印图像统计信息\n        print(f\"Image (ID: {image_id}) max: {np.max(image):.4f}, min: {np.min(image):.4f}, mean: {np.mean(image):.4f}\")\n        \n        image = cv2.resize(image, img_size)\n        max_val = np.max(image)\n        image = image / max_val if max_val != 0 else image\n        \n        mask = np.zeros(img_size, dtype=np.float32)\n        if segmentation_data is not None and image_id <= segmentation_data.shape[2]:\n            mask_slice = segmentation_data[:, :, image_id-1]\n            mask = cv2.resize(mask_slice, img_size, interpolation=cv2.INTER_NEAREST)\n            \n            # 打印掩码统计信息\n            print(f\"Mask (ID: {image_id}) max: {np.max(mask):.4f}, min: {np.min(mask):.4f}, mean: {np.mean(mask):.4f}\")\n            \n        return {\n            'image': image,\n            'mask': mask,\n            'instance_number': int(image_id)\n        }\n    except Exception as e:\n        print(f\"Error loading slice {image_id}: {str(e)}\")\n        return None","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:15.878821Z","iopub.execute_input":"2025-04-22T15:42:15.879138Z","iopub.status.idle":"2025-04-22T15:42:15.900841Z","shell.execute_reply.started":"2025-04-22T15:42:15.879117Z","shell.execute_reply":"2025-04-22T15:42:15.900135Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 添加缓存功能 - 预处理和保存患者数据\ndef preprocess_patient_data(patient_id, meta_df, dicom_tags_df, segmentation_map, \n                           unet_model, output_dir, img_size=(224, 224), max_slices=32):\n    \"\"\"\n    预处理单个患者的数据并保存到磁盘\n    \n    Args:\n        patient_id: 患者ID\n        meta_df: 元数据DataFrame\n        dicom_tags_df: DICOM标签DataFrame\n        segmentation_map: 分割映射\n        unet_model: U-Net模型用于分割\n        output_dir: 输出目录\n        img_size: 图像大小\n        max_slices: 每个患者保存的最大切片数（保持32张）\n    \n    Returns:\n        成功处理返回True，否则返回False\n    \"\"\"\n    output_file = os.path.join(output_dir, f\"{patient_id}.npz\")\n    if os.path.exists(output_file):\n        return True\n        \n    try:\n        # 获取患者的系列ID\n        subset = meta_df[meta_df['patient_id'] == patient_id]\n        if subset.empty:\n            return False\n        series_id = str(subset['series_id'].iloc[0])\n        \n        # 获取图像ID\n        image_ids = dicom_tags_df[\n            (dicom_tags_df['PatientID'] == patient_id) &\n            (dicom_tags_df['series_id'] == series_id)\n        ]['InstanceNumber'].values\n        \n        if len(image_ids) == 0:\n            return False\n            \n        image_ids = sorted(image_ids)\n        \n        # 获取分割数据\n        segmentation_file = segmentation_map.get(patient_id)\n        segmentation_data = None\n        if segmentation_file and os.path.exists(segmentation_file):\n            try:\n                segmentation_data = nib.load(segmentation_file).get_fdata().astype(np.float32)\n            except Exception:\n                pass\n        \n        # 准备并行加载的参数\n        image_files = [os.path.join(TRAIN_IMAGES, str(patient_id), str(series_id), f'{image_id}.dcm') \n                      for image_id in image_ids]\n        \n        # 使用并行处理加载切片\n        args_list = [(image_file, image_id, img_size, segmentation_data) \n                    for image_file, image_id in zip(image_files, image_ids)]\n        \n        all_slices = []\n        with ThreadPoolExecutor(max_workers=8) as executor:\n            results = list(executor.map(load_single_slice, args_list))\n            all_slices = [r for r in results if r is not None]\n        \n        # 如果没有切片，跳过\n        if not all_slices:\n            return False\n        \n        # 选择中间的max_slices张切片\n        if len(all_slices) > max_slices:\n            middle_index = len(all_slices) // 2\n            start_index = max(0, middle_index - max_slices // 2)\n            all_slices = all_slices[start_index:start_index + max_slices]\n        \n        # 如果切片数量不足，则跳过\n        if len(all_slices) < max_slices:\n            return False\n        \n        # 对每个切片应用分割\n        segmented_slices = []\n        for slice_data in all_slices:\n            image = slice_data['image']\n            if len(image.shape) == 2:\n                image_for_seg = np.stack([image, image, image], axis=-1)\n            else:\n                image_for_seg = image\n            \n            segmented_slice = segment_slice(image_for_seg, unet_model, img_size)\n            segmented_slices.append(np.expand_dims(segmented_slice, axis=-1))  # 添加通道维度\n        \n        segmented_images = np.array(segmented_slices)\n        \n        # 保存到磁盘\n        np.savez_compressed(output_file, segmented_images=segmented_images)\n        \n        return True\n        \n    except Exception as e:\n        print(f\"处理患者 {patient_id} 时出错: {str(e)}\")\n        return False\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:15.901511Z","iopub.execute_input":"2025-04-22T15:42:15.901718Z","iopub.status.idle":"2025-04-22T15:42:15.917946Z","shell.execute_reply.started":"2025-04-22T15:42:15.901703Z","shell.execute_reply":"2025-04-22T15:42:15.917347Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def preprocess_patients_parallel(patient_ids, meta_df, dicom_tags_df, segmentation_map, \n                                organ_status_map, unet_model, output_dir, \n                                img_size=(224, 224), max_slices=32, batch_size=10):\n    \"\"\"\n    并行预处理多个患者的数据\n    \n    Args:\n        patient_ids: 患者ID列表\n        meta_df: 元数据DataFrame\n        dicom_tags_df: DICOM标签DataFrame\n        segmentation_map: 分割映射\n        organ_status_map: 器官状态映射\n        unet_model: U-Net模型用于分割\n        output_dir: 输出目录\n        img_size: 图像大小\n        max_slices: 每个患者保存的最大切片数（保持32张）\n        batch_size: 每批处理的患者数\n    \"\"\"\n    os.makedirs(output_dir, exist_ok=True)\n    \n    print(f\"开始为 {len(patient_ids)} 个患者处理标签数据...\")\n    \n    # 先为每个患者保存标签数据\n    for patient_id in tqdm(patient_ids, desc=\"处理患者标签\"):\n        labels_file = os.path.join(output_dir, f\"{patient_id}_labels.npz\")\n        if not os.path.exists(labels_file):\n            try:\n                # 获取器官状态标签\n                organ_status = organ_status_map.get(patient_id, {})\n                labels = []\n                organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n                for organ in organs:\n                    status = organ_status.get(organ, {})\n                    organ_label = [\n                        status.get('healthy', 0),\n                        status.get('injury', 0),\n                        status.get('low', 0),\n                        status.get('high', 0)\n                    ]\n                    labels.append(organ_label)\n                \n                labels = np.array(labels, dtype=np.float32)\n                \n                # 保存标签到磁盘\n                np.savez_compressed(labels_file, labels=labels)\n            except Exception as e:\n                print(f\"处理患者 {patient_id} 标签时出错: {str(e)}\")\n    \n    print(\"标签数据处理完成，开始处理图像数据...\")\n    \n    # 分批处理患者数据\n    total_batches = (len(patient_ids) + batch_size - 1) // batch_size\n    processed_count = 0\n    skipped_count = 0\n    \n    for i in range(0, len(patient_ids), batch_size):\n        batch = patient_ids[i:i+batch_size]\n        print(f\"处理批次 {i//batch_size + 1}/{total_batches}，共 {len(batch)} 个患者\")\n        \n        # 使用多线程并行处理每个患者\n        with ThreadPoolExecutor(max_workers=4) as executor:\n            futures = []\n            for patient_id in batch:\n                # 检查是否已经处理过\n                output_file = os.path.join(output_dir, f\"{patient_id}.npz\")\n                if os.path.exists(output_file):\n                    skipped_count += 1\n                    continue\n                \n                future = executor.submit(\n                    preprocess_patient_data,\n                    patient_id, meta_df, dicom_tags_df, segmentation_map,\n                    unet_model, output_dir, img_size, max_slices\n                )\n                futures.append(future)\n            \n            # 等待所有任务完成\n            for future in futures:\n                if future.result():\n                    processed_count += 1\n        \n        # 每批处理完成后强制垃圾回收\n        gc.collect()\n        \n        # 打印进度\n        print(f\"批次 {i//batch_size + 1}/{total_batches} 完成，总计处理: {processed_count}，跳过: {skipped_count}\")\n    \n    print(f\"预处理完成! 成功处理: {processed_count}，跳过: {skipped_count}\")\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:15.920496Z","iopub.execute_input":"2025-04-22T15:42:15.920871Z","iopub.status.idle":"2025-04-22T15:42:15.936851Z","shell.execute_reply.started":"2025-04-22T15:42:15.920855Z","shell.execute_reply":"2025-04-22T15:42:15.936362Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 加载患者数据 ","metadata":{}},{"cell_type":"code","source":"def load_patient_data_for_tf(patient_id, cache_dir):\n    \"\"\"加载单个患者的数据，用于tf.data API\"\"\"\n    try:\n        # 加载缓存的图像数据\n        image_file = os.path.join(cache_dir, f\"{patient_id}.npz\")\n        image_data = np.load(image_file, allow_pickle=True)\n        segmented_images = image_data['segmented_images']\n        \n        # 加载缓存的标签数据\n        label_file = os.path.join(cache_dir, f\"{patient_id}_labels.npz\")\n        label_data = np.load(label_file, allow_pickle=True)\n        labels = label_data['labels']\n        \n        # 确保切片数量是32\n        if segmented_images.shape[0] != 32:\n            if segmented_images.shape[0] > 32:\n                middle = segmented_images.shape[0] // 2\n                start = middle - 16\n                segmented_images = segmented_images[start:start+32]\n            else:\n                # 返回空数据，后续会被过滤掉\n                return np.zeros((32, 224, 224, 1), dtype=np.float32), np.zeros((5, 4), dtype=np.float32)\n        \n        return segmented_images.astype(np.float32), labels.astype(np.float32)\n    except Exception as e:\n        print(f\"加载患者 {patient_id} 的缓存数据时出错: {str(e)}\")\n        # 返回空数据，后续会被过滤掉\n        return np.zeros((32, 224, 224, 1), dtype=np.float32), np.zeros((5, 4), dtype=np.float32)\n\ndef create_organ_balanced_dataset(patient_ids, cache_dir, batch_size=4, shuffle=True):\n    \"\"\"创建按器官平衡的数据集，确保每个批次包含各种器官的损伤样本\"\"\"\n    print(\"创建按器官平衡的数据集...\")\n    \n    # 按器官损伤情况分类患者\n    organ_patients = {\n        'bowel': {'injured': [], 'healthy': []},\n        'extravasation': {'injured': [], 'healthy': []},\n        'kidney': {'injured': [], 'healthy': []},\n        'liver': {'injured': [], 'healthy': []},\n        'spleen': {'injured': [], 'healthy': []},\n    }\n    \n    # 分类每个患者\n    for pid in tqdm(patient_ids, desc=\"按器官分类患者\"):\n        label_file = os.path.join(cache_dir, f\"{pid}_labels.npz\")\n        if os.path.exists(label_file):\n            try:\n                labels = np.load(label_file)['labels']\n                \n                # 检查每个器官是否有损伤\n                for i, organ in enumerate(['bowel', 'extravasation', 'kidney', 'liver', 'spleen']):\n                    # 如果有损伤标签(injury, low, high中任一个)\n                    if np.any(labels[i, 1:] > 0):\n                        organ_patients[organ]['injured'].append(pid)\n                    else:\n                        organ_patients[organ]['healthy'].append(pid)\n                        \n            except Exception as e:\n                print(f\"处理患者 {pid} 标签时出错: {str(e)}\")\n                continue\n    \n    # 打印每类患者数量\n    for organ, status_dict in organ_patients.items():\n        print(f\"{organ}: {len(status_dict['injured'])} 损伤患者, {len(status_dict['healthy'])} 健康患者\")\n    \n    # 计算每个批次要包含的每种类型的样本数\n    organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n    \n    def data_generator():\n        while True:\n            batch_patients = []\n            \n            # 确保每个批次至少包含每个器官的1个损伤样本\n            for organ in organs:\n                if organ_patients[organ]['injured']:\n                    # 随机选择该器官的1个损伤患者\n                    injured_pid = random.choice(organ_patients[organ]['injured'])\n                    if injured_pid not in batch_patients:\n                        batch_patients.append(injured_pid)\n            \n            # 填充剩余位置，确保批次大小正确\n            remaining_slots = batch_size - len(batch_patients)\n            if remaining_slots > 0:\n                # 创建所有健康患者的列表\n                all_healthy = []\n                for organ in organs:\n                    all_healthy.extend([p for p in organ_patients[organ]['healthy'] if p not in batch_patients])\n                \n                # 去重\n                all_healthy = list(set(all_healthy))\n                \n                if all_healthy:\n                    # 随机选择健康患者填充批次\n                    selected_healthy = random.sample(all_healthy, min(remaining_slots, len(all_healthy)))\n                    batch_patients.extend(selected_healthy)\n            \n            # 如果批次大小不足，随机复制已有患者\n            while len(batch_patients) < batch_size:\n                batch_patients.append(random.choice(batch_patients))\n            \n            # 如果批次太大，随机删减\n            if len(batch_patients) > batch_size:\n                batch_patients = random.sample(batch_patients, batch_size)\n            \n            # 打乱批次中的患者顺序\n            random.shuffle(batch_patients)\n            \n            # 加载数据\n            batch_images = []\n            batch_labels = []\n            \n            for pid in batch_patients:\n                try:\n                    images, labels = load_patient_data_for_tf(pid, cache_dir)\n                    batch_images.append(images)\n                    batch_labels.append(labels)\n                except Exception as e:\n                    print(f\"加载患者 {pid} 数据时出错: {str(e)}\")\n                    continue\n            \n            if batch_images:  # 确保至少有一个有效样本\n                yield np.array(batch_images), np.array(batch_labels)\n    \n    # 创建tf.data.Dataset\n    output_signature = (\n        tf.TensorSpec(shape=(None, 32, 224, 224, 1), dtype=tf.float32),\n        tf.TensorSpec(shape=(None, 5, 4), dtype=tf.float32)\n    )\n    \n    dataset = tf.data.Dataset.from_generator(\n        data_generator,\n        output_signature=output_signature\n    )\n    \n    # 添加预取以提高性能\n    dataset = dataset.prefetch(tf.data.AUTOTUNE)\n    \n    # 计算每个epoch的步数\n    total_injured = sum(len(status_dict['injured']) for organ, status_dict in organ_patients.items())\n    steps_per_epoch = max(1, total_injured // batch_size)\n    \n    return dataset, steps_per_epoch\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:15.937365Z","iopub.execute_input":"2025-04-22T15:42:15.937536Z","iopub.status.idle":"2025-04-22T15:42:15.953858Z","shell.execute_reply.started":"2025-04-22T15:42:15.937518Z","shell.execute_reply":"2025-04-22T15:42:15.953178Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def create_tf_dataset(patient_ids, cache_dir, batch_size=4, shuffle=True):\n    \"\"\"创建tf.data.Dataset数据集，替代原有的Sequence生成器\"\"\"\n    print(\"创建tf.data数据集...\")\n    \n    # 过滤有效患者ID\n    valid_patient_ids = []\n    for pid in tqdm(patient_ids, desc=\"过滤有效患者ID\"):\n        image_file = os.path.join(cache_dir, f\"{pid}.npz\")\n        label_file = os.path.join(cache_dir, f\"{pid}_labels.npz\")\n        if os.path.exists(image_file) and os.path.exists(label_file):\n            valid_patient_ids.append(pid)\n    \n    print(f\"找到 {len(valid_patient_ids)} 个有效的缓存患者数据\")\n    \n    # 创建一个tf.data.Dataset，包含所有有效患者ID\n    patient_ds = tf.data.Dataset.from_tensor_slices(valid_patient_ids)\n    \n    # 如果需要洗牌\n    if shuffle:\n        patient_ds = patient_ds.shuffle(buffer_size=min(len(valid_patient_ids), 100), \n                                    reshuffle_each_iteration=True)\n    \n    # 加载患者数据\n    def load_patient(patient_id):\n        # 将字符串张量转换为Python字符串\n        pid = patient_id.numpy().decode('utf-8')\n        images, labels = load_patient_data_for_tf(pid, cache_dir)\n        return images, labels\n    \n    # 使用tf.py_function包装Python函数，使其可以在tf.data管道中使用\n    def load_patient_wrapper(patient_id):\n        images, labels = tf.py_function(\n            func=load_patient,\n            inp=[patient_id],\n            Tout=[tf.float32, tf.float32]\n        )\n        # 设置形状信息，因为py_function不会保留\n        images.set_shape((32, 224, 224, 1))\n        labels.set_shape((5, 4))\n        return images, labels\n\n    # 关键：先repeat再map，确保数据永远不会耗尽\n    patient_ds = patient_ds.repeat()\n    \n    # 映射到加载函数\n    dataset = patient_ds.map(\n        load_patient_wrapper,\n        num_parallel_calls=tf.data.AUTOTUNE\n    )\n    \n    # 过滤无效数据\n    def is_valid(images, labels):\n        return tf.math.reduce_all(tf.math.is_finite(images))\n    \n    dataset = dataset.filter(is_valid)\n    \n    # 添加repeat()方法，使数据集可以在多个epoch中重复使用\n    dataset = dataset.repeat()  # 无限重复\n    \n    # 批处理\n    dataset = dataset.batch(batch_size)\n    \n    # 预取下一批数据\n    dataset = dataset.prefetch(tf.data.AUTOTUNE)\n    \n    # 启用性能优化\n    options = tf.data.Options()\n    options.experimental_optimization.apply_default_optimizations = True\n    \n    # 应用选项\n    dataset = dataset.with_options(options)\n    \n    return dataset, len(valid_patient_ids)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:15.954756Z","iopub.execute_input":"2025-04-22T15:42:15.954999Z","iopub.status.idle":"2025-04-22T15:42:15.971427Z","shell.execute_reply.started":"2025-04-22T15:42:15.954976Z","shell.execute_reply":"2025-04-22T15:42:15.970756Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 可视化模块","metadata":{}},{"cell_type":"code","source":"# 可视化类别分布\ndef visualize_class_distribution(dataset, steps, title=\"Class Distribution\"):\n    \"\"\"可视化数据集的类别分布\"\"\"\n    # 收集数据\n    organ_counts = {\n        'bowel': {'healthy': 0, 'injury': 0},\n        'extravasation': {'healthy': 0, 'injury': 0},\n        'kidney': {'healthy': 0, 'injury': 0},\n        'liver': {'healthy': 0, 'injury': 0},\n        'spleen': {'healthy': 0, 'injury': 0}\n    }\n    \n    total_samples = 0\n    organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n    \n    for x_batch, y_batch in tqdm(dataset.take(steps), total=steps, desc=\"Collecting data\"):\n        batch_size = x_batch.shape[0]\n        total_samples += batch_size\n        \n        y_batch_np = y_batch.numpy()\n        \n        # 统计每个器官的健康和损伤样本\n        for i, organ in enumerate(organs):\n            healthy_count = np.sum(y_batch_np[:, i, 0])\n            injury_count = np.sum(y_batch_np[:, i, 1])\n            \n            organ_counts[organ]['healthy'] += healthy_count\n            organ_counts[organ]['injury'] += injury_count\n    \n    # 可视化\n    plt.figure(figsize=(14, 8))\n    \n    # 绘制条形图\n    x = np.arange(len(organs))\n    width = 0.35\n    \n    healthy_percentages = [organ_counts[organ]['healthy'] / total_samples * 100 for organ in organs]\n    injury_percentages = [organ_counts[organ]['injury'] / total_samples * 100 for organ in organs]\n    \n    plt.bar(x - width/2, healthy_percentages, width, label='Healthy')\n    plt.bar(x + width/2, injury_percentages, width, label='Injured')\n    \n    plt.xlabel('Organ')\n    plt.ylabel('Percentage (%)')\n    plt.title(title)\n    plt.xticks(x, organs)\n    plt.legend()\n    \n    # 添加数值标签\n    for i, v in enumerate(healthy_percentages):\n        plt.text(i - width/2, v + 1, f'{v:.1f}%', ha='center')\n    \n    for i, v in enumerate(injury_percentages):\n        plt.text(i + width/2, v + 1, f'{v:.1f}%', ha='center')\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, f'{title.replace(\" \", \"_\")}.png'), dpi=300)\n    plt.show()\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:15.972137Z","iopub.execute_input":"2025-04-22T15:42:15.972317Z","iopub.status.idle":"2025-04-22T15:42:15.992039Z","shell.execute_reply.started":"2025-04-22T15:42:15.972304Z","shell.execute_reply":"2025-04-22T15:42:15.991497Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 可视化训练历史记录的函数\ndef plot_training_history(history):\n    \"\"\"\n    可视化模型训练历史\n    \n    Args:\n        history: 模型训练历史对象\n    \"\"\"\n    # 绘制损失曲线\n    plt.figure(figsize=(12, 5))\n    \n    plt.subplot(1, 2, 1)\n    plt.plot(history.history['loss'], label='训练损失')\n    plt.plot(history.history['val_loss'], label='验证损失')\n    plt.title('模型损失')\n    plt.ylabel('损失')\n    plt.xlabel('Epoch')\n    plt.legend()\n    \n    plt.subplot(1, 2, 2)\n    plt.plot(history.history['accuracy'], label='训练准确率')\n    plt.plot(history.history['val_accuracy'], label='验证准确率')\n    plt.title('模型准确率')\n    plt.ylabel('准确率')\n    plt.xlabel('Epoch')\n    plt.legend()\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'training_history.png'), dpi=300)\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:15.992608Z","iopub.execute_input":"2025-04-22T15:42:15.992832Z","iopub.status.idle":"2025-04-22T15:42:16.010231Z","shell.execute_reply.started":"2025-04-22T15:42:15.992818Z","shell.execute_reply":"2025-04-22T15:42:16.009665Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_segmentation_results(unet_model, patient_id, meta_df, dicom_tags_df, segmentation_map, num_samples=5):\n    \"\"\"\n    可视化U-Net分割模型的结果\n    Args:\n        unet_model: 训练好的U-Net模型\n        patient_id: 患者ID\n        meta_df: 元数据DataFrame\n        dicom_tags_df: DICOM标签DataFrame\n        segmentation_map: 分割映射\n        num_samples: 要显示的样本数量 (增加到5)\n    \"\"\"\n    try:\n        # 获取患者的系列ID\n        subset = meta_df[meta_df['patient_id'] == patient_id]\n        if subset.empty:\n            print(f\"未找到患者 {patient_id} 的元数据\")\n            return\n\n        series_id = str(subset['series_id'].iloc[0])\n\n        # 获取图像ID\n        image_ids = dicom_tags_df[\n            (dicom_tags_df['PatientID'] == patient_id) &\n            (dicom_tags_df['series_id'] == series_id)\n        ]['InstanceNumber'].values\n\n        if len(image_ids) == 0:\n            print(f\"未找到患者 {patient_id} 的图像ID\")\n            return\n\n        image_ids = sorted(image_ids)\n\n        # 获取分割数据\n        segmentation_file = segmentation_map.get(patient_id)\n        segmentation_data = None\n        if segmentation_file and os.path.exists(segmentation_file):\n            try:\n                segmentation_data = nib.load(segmentation_file).get_fdata().astype(np.float32)\n            except Exception as e:\n                print(f\"加载分割数据时出错: {str(e)}\")\n\n        # 随机选择num_samples个切片\n        if len(image_ids) <= num_samples:\n            selected_ids = image_ids\n        else:\n            selected_ids = random.sample(list(image_ids), num_samples)\n\n        plt.figure(figsize=(15, 5 * num_samples)) # 调整图形大小\n\n        for i, image_id in enumerate(selected_ids):\n            image_file = os.path.join(TRAIN_IMAGES, str(patient_id), str(series_id), f'{image_id}.dcm')\n\n            # 加载图像\n            dicom = pydicom.dcmread(image_file)\n            image = dicom.pixel_array.astype(np.float32)\n            image = cv2.resize(image, IMG_SIZE)\n            max_val = np.max(image)\n            image = image / max_val if max_val != 0 else image\n\n            # 加载真实分割掩码\n            ground_truth = np.zeros(IMG_SIZE, dtype=np.float32)\n            if segmentation_data is not None and image_id <= segmentation_data.shape[2]:\n                mask_slice = segmentation_data[:, :, image_id-1]\n                ground_truth = cv2.resize(mask_slice, IMG_SIZE, interpolation=cv2.INTER_NEAREST)\n\n            # 使用U-Net模型进行分割\n            if len(image.shape) == 2:\n                image_for_seg = np.stack([image, image, image], axis=-1)\n            else:\n                image_for_seg = image\n\n            predicted_mask = segment_slice(image_for_seg, unet_model, IMG_SIZE)\n\n            # 可视化结果\n            plt.subplot(num_samples, 3, i * 3 + 1)  #  3 列\n            plt.imshow(image, cmap='gray')\n            plt.title(translate(f'原始图像 (ID: {image_id})'))\n            plt.axis('off')\n\n            plt.subplot(num_samples, 3, i * 3 + 2)\n            plt.imshow(ground_truth, cmap='jet', alpha=0.7)\n            plt.title(translate('真实分割掩码'))\n            plt.axis('off')\n\n            plt.subplot(num_samples, 3, i * 3 + 3)\n            plt.imshow(image, cmap='gray')\n            plt.imshow(predicted_mask, cmap='jet', alpha=0.7)\n            plt.title(translate('预测分割掩码'))\n            plt.axis('off')\n\n        plt.tight_layout()\n        plt.savefig(os.path.join(OUTPUT_DIR, f'segmentation_results_{patient_id}.png'), dpi=300)\n        plt.show()\n\n    except Exception as e:\n        print(f\"可视化分割结果时出错: {str(e)}\")\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:16.010861Z","iopub.execute_input":"2025-04-22T15:42:16.011052Z","iopub.status.idle":"2025-04-22T15:42:16.027502Z","shell.execute_reply.started":"2025-04-22T15:42:16.011038Z","shell.execute_reply":"2025-04-22T15:42:16.026842Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def visualize_classification_results(classifier_model, unet_model, patient_id, meta_df, dicom_tags_df, \n                                    segmentation_map, organ_status_map):\n    \"\"\"\n    可视化分类模型的结果\n    \"\"\"\n    try:\n        # 获取患者的系列ID\n        subset = meta_df[meta_df['patient_id'] == patient_id]\n        if subset.empty:\n            print(f\"未找到患者 {patient_id} 的元数据\")\n            return\n        \n        series_id = str(subset['series_id'].iloc[0])\n        \n        # 获取图像ID\n        image_ids = dicom_tags_df[\n            (dicom_tags_df['PatientID'] == patient_id) &\n            (dicom_tags_df['series_id'] == series_id)\n        ]['InstanceNumber'].values\n        \n        if len(image_ids) == 0:\n            print(f\"未找到患者 {patient_id} 的图像ID\")\n            return\n            \n        image_ids = sorted(image_ids)\n        \n        # 获取分割数据\n        segmentation_file = segmentation_map.get(patient_id)\n        segmentation_data = None\n        if segmentation_file and os.path.exists(segmentation_file):\n            try:\n                segmentation_data = nib.load(segmentation_file).get_fdata().astype(np.float32)\n            except Exception as e:\n                print(f\"加载分割数据时出错: {str(e)}\")\n        \n        # 加载所有切片\n        all_slices = []\n        for image_id in image_ids:\n            image_file = os.path.join(TRAIN_IMAGES, str(patient_id), str(series_id), f'{image_id}.dcm')\n            try:\n                dicom = pydicom.dcmread(image_file)\n                image = dicom.pixel_array.astype(np.float32)\n                image = cv2.resize(image, IMG_SIZE)\n                max_val = np.max(image)\n                image = image / max_val if max_val != 0 else image\n                \n                mask = np.zeros(IMG_SIZE, dtype=np.float32)\n                if segmentation_data is not None and image_id <= segmentation_data.shape[2]:\n                    mask_slice = segmentation_data[:, :, image_id-1]\n                    mask = cv2.resize(mask_slice, IMG_SIZE, interpolation=cv2.INTER_NEAREST)\n                \n                all_slices.append({\n                    'image': image,\n                    'mask': mask,\n                    'instance_number': int(image_id)\n                })\n            except Exception as e:\n                print(f\"加载切片 {image_id} 时出错: {str(e)}\")\n                continue\n        \n        # 如果切片数量不足32，则退出\n        if len(all_slices) < 32:\n            print(f\"患者 {patient_id} 的切片数量不足32，无法进行分类\")\n            return\n        \n        # 选择中间32张切片\n        middle_index = len(all_slices) // 2\n        start_index = max(0, middle_index - 16)\n        sequence = all_slices[start_index:start_index + 32]\n        \n        # 处理切片\n        processed_sequence = []\n        for slice_data in sequence:\n            image = slice_data['image']\n            if len(image.shape) == 2:\n                image = np.stack([image, image, image], axis=-1)\n            \n            segmented_slice = segment_slice(image, unet_model, IMG_SIZE)\n            processed_sequence.append(np.expand_dims(segmented_slice, axis=-1))\n        \n        # 准备输入数据\n        X = np.expand_dims(np.array(processed_sequence), axis=0)\n        \n        # 使用分类模型进行预测\n        y_pred = classifier_model.predict(X)[0]\n        \n        # 获取真实标签\n        organ_status = organ_status_map.get(patient_id, {})\n        y_true = []\n        organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n        for organ in organs:\n            status = organ_status.get(organ, {})\n            organ_label = [\n                status.get('healthy', 0),\n                status.get('injury', 0),\n                status.get('low', 0),\n                status.get('high', 0)\n            ]\n            y_true.append(organ_label)\n        \n        y_true = np.array(y_true)\n        \n        # 可视化结果 - 修改部分\n        plt.figure(figsize=(15, 10))\n        \n        # 显示中间的一张切片\n        middle_slice = processed_sequence[16][:, :, 0]\n        plt.subplot(2, 3, 1)\n        plt.imshow(middle_slice, cmap='gray')\n        plt.title('代表性切片')\n        plt.axis('off')\n        \n        # 显示每个器官的预测结果\n        colors = ['blue', 'green', 'red', 'purple', 'orange']\n        \n        plt.subplot(2, 3, 2)\n        for i, organ in enumerate(organs):\n            plt.bar(i, y_pred[i, 1], color=colors[i])\n        plt.axhline(y=0.5, color='r', linestyle='--', alpha=0.3)\n        plt.xticks(range(len(organs)), organs, rotation=45)\n        plt.ylim(0, 1)\n        plt.title('损伤预测概率')\n        \n        # 显示真实标签和预测标签的比较\n        plt.subplot(2, 3, 3)\n        for i, organ in enumerate(organs):\n            true_label = np.argmax(y_true[i])\n            pred_label = np.argmax(y_pred[i])\n            \n            plt.scatter(i-0.1, true_label, color='blue', label='真实' if i == 0 else '')\n            plt.scatter(i+0.1, pred_label, color='red', label='预测' if i == 0 else '')\n        \n        plt.xticks(range(len(organs)), organs, rotation=45)\n        plt.yticks(range(4), ['健康', '损伤', '低度', '高度'])\n        plt.title('真实标签 vs 预测标签')\n        plt.legend()\n        \n        # 显示详细的器官预测结果 - 只显示前3个器官\n        for i in range(min(3, len(organs))):\n            plt.subplot(2, 3, 4 + i)  # 最大索引为6\n            \n            bar_width = 0.35\n            index = np.arange(4)\n            \n            plt.bar(index, y_true[i], bar_width, label='真实', color='blue', alpha=0.7)\n            plt.bar(index + bar_width, y_pred[i], bar_width, label='预测', color='red', alpha=0.7)\n            \n            plt.xlabel(organs[i].capitalize())\n            plt.xticks(index + bar_width/2, ['健康', '损伤', '低度', '高度'])\n            plt.ylim(0, 1)\n            \n            if i == 0:\n                plt.legend()\n        \n        plt.tight_layout()\n        plt.savefig(os.path.join(OUTPUT_DIR, f'classification_results_{patient_id}.png'), dpi=300)\n        plt.show()\n        \n    except Exception as e:\n        print(f\"可视化分类结果时出错: {str(e)}\")\n        import traceback\n        print(traceback.format_exc())\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:16.028101Z","iopub.execute_input":"2025-04-22T15:42:16.028362Z","iopub.status.idle":"2025-04-22T15:42:16.049083Z","shell.execute_reply.started":"2025-04-22T15:42:16.028346Z","shell.execute_reply":"2025-04-22T15:42:16.048504Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 绘制ROC曲线\ndef plot_roc_curves(evaluation_results):\n    \"\"\"绘制ROC曲线\"\"\"\n    plt.figure(figsize=(15, 10))\n    organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n    \n    for i, result in enumerate([r for r in evaluation_results if r['organ'] in organs]):\n        plt.subplot(2, 3, i+1)\n        organ = result['organ']\n        \n        # 获取ROC曲线数据\n        fpr = result['fpr']\n        tpr = result['tpr']\n        auc_value = result['auc']\n        \n        # 绘制ROC曲线\n        plt.plot(fpr, tpr, lw=2, label=f'ROC (AUC = {auc_value:.3f})')\n        plt.plot([0, 1], [0, 1], 'k--', lw=2)\n        plt.xlim([0.0, 1.0])\n        plt.ylim([0.0, 1.05])\n        plt.xlabel('False Positive Rate')\n        plt.ylabel('True Positive Rate')\n        plt.title(f'{organ.capitalize()} ROC Curve')\n        plt.legend(loc=\"lower right\")\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'roc_curves.png'), dpi=300)\n    plt.show()\n\n# 绘制混淆矩阵\ndef plot_confusion_matrices(evaluation_results):\n    \"\"\"绘制混淆矩阵热图\"\"\"\n    plt.figure(figsize=(15, 10))\n    organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n    \n    for i, result in enumerate([r for r in evaluation_results if r['organ'] in organs]):\n        plt.subplot(2, 3, i+1)\n        organ = result['organ']\n        cm = result['confusion_matrix']\n        \n        # 绘制热图\n        sns.heatmap(cm, annot=True, fmt='d', cmap='Blues', cbar=False,\n                   xticklabels=['Healthy', 'Injured'], \n                   yticklabels=['Healthy', 'Injured'])\n        \n        plt.xlabel('Predicted Label')\n        plt.ylabel('True Label')\n        plt.title(f'{organ.capitalize()} Confusion Matrix')\n    \n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'confusion_matrices.png'), dpi=300)\n    plt.show()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:16.049674Z","iopub.execute_input":"2025-04-22T15:42:16.049874Z","iopub.status.idle":"2025-04-22T15:42:16.067414Z","shell.execute_reply.started":"2025-04-22T15:42:16.049854Z","shell.execute_reply":"2025-04-22T15:42:16.066747Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 显示更多信息","metadata":{}},{"cell_type":"code","source":"import time\nimport sys\n\nclass ProgressLogger:\n    def __init__(self):\n        self.start_time = time.time()\n        self.last_time = self.start_time\n        \n    def log(self, message):\n        current_time = time.time()\n        elapsed = current_time - self.last_time\n        total_elapsed = current_time - self.start_time\n        \n        time_str = time.strftime(\"%Y-%m-%d %H:%M:%S\", time.localtime())\n        elapsed_str = time.strftime(\"%H:%M:%S\", time.gmtime(total_elapsed))\n        \n        print(f\"[{time_str}] [{elapsed_str}] (+{elapsed:.2f}s) {message}\")\n        sys.stdout.flush()  # 确保立即显示\n        \n        self.last_time = current_time\n\n# 创建全局日志记录器\nlogger = ProgressLogger()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:16.068145Z","iopub.execute_input":"2025-04-22T15:42:16.068848Z","iopub.status.idle":"2025-04-22T15:42:16.086836Z","shell.execute_reply.started":"2025-04-22T15:42:16.068827Z","shell.execute_reply":"2025-04-22T15:42:16.086209Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 回调函数","metadata":{}},{"cell_type":"code","source":"class DetailedTensorBoard(tf.keras.callbacks.TensorBoard):\n    def __init__(self, log_dir=\"logs\", **kwargs):\n        super().__init__(log_dir=log_dir, **kwargs)\n        self.start_time = None\n        self.epoch_start_time = None\n        self.batch_times = []\n        \n    def on_train_begin(self, logs=None):\n        super().on_train_begin(logs)\n        self.start_time = time.time()\n        print(\"训练开始，时间：\", time.strftime(\"%Y-%m-%d %H:%M:%S\", time.localtime()))\n        \n    def on_epoch_begin(self, epoch, logs=None):\n        super().on_epoch_begin(epoch, logs)\n        self.epoch_start_time = time.time()\n        self.batch_times = []\n        print(f\"\\n===== 开始 Epoch {epoch+1} =====\")\n        \n    def on_train_batch_end(self, batch, logs=None):\n        super().on_train_batch_end(batch, logs)\n        if batch % 10 == 0:  # 每10个批次打印一次\n            current_time = time.time()\n            \n            # 获取steps_per_epoch，如果为None则显示为'?'\n            steps_per_epoch = self.params.get('steps')\n            steps_display = steps_per_epoch if steps_per_epoch is not None else '?'\n            \n            # 添加当前批次的时间\n            if self.batch_times:\n                last_time = self.batch_times[-1]\n                batch_time = current_time - last_time\n                self.batch_times.append(current_time)\n            else:\n                batch_time = current_time - self.epoch_start_time\n                self.batch_times.append(current_time)\n            \n            # 计算平均批次时间和ETA\n            if len(self.batch_times) > 1:\n                # 计算相邻时间点的差值来获取批次时间\n                recent_times = [self.batch_times[i] - self.batch_times[i-1] for i in range(1, min(11, len(self.batch_times)))]\n                avg_time = sum(recent_times) / len(recent_times)\n                \n                # 安全检查：确保steps_per_epoch不是None，才进行ETA计算\n                if steps_per_epoch is not None:\n                    eta = avg_time * (steps_per_epoch - batch)\n                    eta_str = time.strftime(\"%H:%M:%S\", time.gmtime(eta))\n                    print(f\"Batch {batch}/{steps_display}: loss={logs.get('loss', 0):.4f}, accuracy={logs.get('accuracy', 0):.4f}, ETA: {eta_str}\")\n                else:\n                    # 如果steps_per_epoch是None，只显示当前进度，不显示ETA\n                    print(f\"Batch {batch}/{steps_display}: loss={logs.get('loss', 0):.4f}, accuracy={logs.get('accuracy', 0):.4f}\")\n            else:\n                print(f\"Batch {batch}/{steps_display}: loss={logs.get('loss', 0):.4f}, accuracy={logs.get('accuracy', 0):.4f}\")\n        \n    def on_epoch_end(self, epoch, logs=None):\n        epoch_time = time.time() - self.epoch_start_time\n        print(f\"\\n===== Epoch {epoch+1} 完成 =====\")\n        print(f\"训练损失: {logs.get('loss', 0):.4f}, 训练准确率: {logs.get('accuracy', 0):.4f}\")\n        print(f\"验证损失: {logs.get('val_loss', 0):.4f}, 验证准确率: {logs.get('val_accuracy', 0):.4f}\")\n        print(f\"Epoch 耗时: {epoch_time:.2f}秒\")\n        print(f\"已训练时间: {(time.time() - self.start_time):.2f}秒\")\n        super().on_epoch_end(epoch, logs)\n        \n    def on_train_end(self, logs=None):\n        total_time = time.time() - self.start_time\n        print(\"\\n===== 训练完成 =====\")\n        print(f\"总训练时间: {total_time:.2f}秒 ({total_time/60:.2f}分钟)\")\n        super().on_train_end(logs)\n\nclass ClassBalanceMonitor(tf.keras.callbacks.Callback):\n    def __init__(self, dataset, steps):\n        super().__init__()\n        self.dataset = dataset\n        self.steps = steps\n        self.organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n        \n    def on_epoch_begin(self, epoch, logs=None):\n        # 检查训练批次中的类别分布\n        print(\"\\n检查训练数据类别分布:\")\n        \n        # 初始化计数器\n        healthy_counts = np.zeros(5)\n        injury_counts = np.zeros(5)\n        low_counts = np.zeros(5)\n        high_counts = np.zeros(5)\n        total_samples = 0\n        \n        # 只分析前5个批次，以免花费太多时间\n        for x_batch, y_batch in self.dataset.take(5):  \n            batch_size = x_batch.shape[0]\n            total_samples += batch_size\n            y_batch_np = y_batch.numpy()\n            \n            # 统计每个器官的各种损伤样本\n            for i, organ in enumerate(self.organs):\n                healthy_counts[i] += np.sum(y_batch_np[:, i, 0])\n                injury_counts[i] += np.sum(y_batch_np[:, i, 1])\n                low_counts[i] += np.sum(y_batch_np[:, i, 2])\n                high_counts[i] += np.sum(y_batch_np[:, i, 3])\n        \n        # 打印详细统计信息\n        print(f\"分析了 {total_samples} 个样本:\")\n        for i, organ in enumerate(self.organs):\n            print(f\"  {organ}:\")\n            if total_samples > 0:\n                healthy_pct = (healthy_counts[i] / total_samples) * 100\n                injury_pct = (injury_counts[i] / total_samples) * 100\n                low_pct = (low_counts[i] / total_samples) * 100\n                high_pct = (high_counts[i] / total_samples) * 100\n                \n                print(f\"    健康: {healthy_counts[i]:.0f} ({healthy_pct:.2f}%)\")\n                print(f\"    损伤: {injury_counts[i]:.0f} ({injury_pct:.2f}%)\")\n                print(f\"    低度: {low_counts[i]:.0f} ({low_pct:.2f}%)\")\n                print(f\"    高度: {high_counts[i]:.0f} ({high_pct:.2f}%)\")","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:16.087451Z","iopub.execute_input":"2025-04-22T15:42:16.087653Z","iopub.status.idle":"2025-04-22T15:42:16.105095Z","shell.execute_reply.started":"2025-04-22T15:42:16.087617Z","shell.execute_reply":"2025-04-22T15:42:16.104502Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 损失函数&训练","metadata":{}},{"cell_type":"code","source":"def weighted_cross_entropy(organ_weights=None):\n    \"\"\"\n    返回一个加权交叉熵损失函数，为不同器官设置不同权重\n    \"\"\"\n    # 如果没有提供权重，使用基于类别不平衡的默认权重\n    if organ_weights is None:\n        # 基于论文中的数据分布计算权重\n        # 肠道(Bowel)：损伤比例2.3%，权重约43\n        # 渗出(Extravasation)：6.8%，权重约15\n        # 肾脏(Kidney)：6.9%，权重约14.5\n        # 肝脏(Liver)：10.8%，权重约9.3\n        # 脾脏(Spleen)：11.8%，权重约8.5\n        organ_weights = [43.0, 15.0, 14.5, 9.3, 8.5]\n    \n    def loss(y_true, y_pred):\n        \"\"\"计算加权交叉熵损失\"\"\"\n        organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n        total_loss = 0.0\n        \n        for i, (organ, weight) in enumerate(zip(organs, organ_weights)):\n            # 提取每个器官的真实标签和预测值\n            y_true_organ = y_true[:, i, :]\n            y_pred_organ = y_pred[:, i, :]\n            \n            # 计算交叉熵损失\n            epsilon = 1e-7  # 防止log(0)\n            \n            # 单独处理每个类别\n            healthy_true = y_true_organ[:, 0]\n            injury_true = y_true_organ[:, 1]\n            \n            healthy_pred = y_pred_organ[:, 0]\n            injury_pred = y_pred_organ[:, 1]\n            \n            # 基础交叉熵\n            healthy_loss = -healthy_true * tf.math.log(healthy_pred + epsilon)\n            injury_loss = -injury_true * tf.math.log(injury_pred + epsilon)\n            \n            # 应用权重 - 损伤类别权重更高\n            weighted_loss = healthy_loss + weight * injury_loss\n            \n            # 添加到总损失\n            total_loss += tf.reduce_mean(weighted_loss)\n        \n        return total_loss / len(organs)\n    \n    return loss\n\n\n\n# 评估指标计算\ndef calculate_metrics(y_true, y_pred, threshold=0.5):\n    \"\"\"计算与论文一致的评估指标\"\"\"\n    # 将预测概率转换为二进制标签\n    y_pred_binary = (y_pred > threshold).astype(int)\n    \n    # 计算指标\n    accuracy = accuracy_score(y_true, y_pred_binary)\n    \n    # 处理可能的除零错误\n    if np.sum(y_true) > 0:  # 如果有正样本\n        precision = precision_score(y_true, y_pred_binary, zero_division=0)\n        recall = recall_score(y_true, y_pred_binary, zero_division=0)\n    else:\n        precision = 0\n        recall = 0\n    \n    # 计算混淆矩阵\n    cm = confusion_matrix(y_true, y_pred_binary)\n    \n    # 从混淆矩阵中计算TP, TN, FP, FN\n    if cm.shape == (2, 2):\n        tn, fp, fn, tp = cm.ravel()\n        # 计算特异性\n        specificity = tn / (tn + fp) if (tn + fp) > 0 else 0\n        # 计算PPV和NPV\n        ppv = precision\n        npv = tn / (tn + fn) if (tn + fn) > 0 else 0\n    else:\n        # 处理只有一个类别的情况\n        if np.all(y_true == 0):  # 全是负样本\n            tn = np.sum(y_pred_binary == 0)\n            fp = np.sum(y_pred_binary == 1)\n            fn = 0\n            tp = 0\n            specificity = tn / (tn + fp) if (tn + fp) > 0 else 1\n            ppv = 0\n            npv = tn / tn if tn > 0 else 1\n        else:  # 全是正样本\n            tn = 0\n            fp = 0\n            fn = np.sum(y_pred_binary == 0)\n            tp = np.sum(y_pred_binary == 1)\n            specificity = 0\n            ppv = tp / tp if tp > 0 else 1\n            npv = 0\n    \n    return accuracy, precision, recall, specificity, ppv, npv, cm\n\n# 按器官评估模型\ndef evaluate_model_by_organ(model, val_dataset, validation_steps):\n    \"\"\"按器官单独评估模型性能\"\"\"\n    organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n    \n    # 收集验证集预测\n    all_true = []\n    all_pred = []\n    \n    for x_batch, y_batch in tqdm(val_dataset.take(validation_steps), total=validation_steps, desc=\"Evaluating model\"):\n        y_pred = model.predict(x_batch, verbose=0)\n        all_true.append(y_batch.numpy())\n        all_pred.append(y_pred)\n    \n    # 合并批次结果\n    y_true = np.vstack(all_true)\n    y_pred = np.vstack(all_pred)\n    \n    # 按器官评估\n    results = []\n    for i, organ in enumerate(organs):\n        # 提取当前器官的标签和预测\n        organ_true = y_true[:, i, 1]  # 只考虑injury标签\n        organ_pred = y_pred[:, i, 1]\n        \n        # 计算评估指标\n        accuracy, precision, recall, specificity, ppv, npv, cm = calculate_metrics(organ_true, organ_pred)\n        \n        # 计算ROC曲线和AUC\n        try:\n            fpr, tpr, _ = roc_curve(organ_true, organ_pred)\n            auc_value = auc(fpr, tpr)\n        except:\n            fpr, tpr = [0, 1], [0, 1]\n            auc_value = 0.5\n        \n        # 记录结果\n        results.append({\n            'organ': organ,\n            'accuracy': accuracy,\n            'precision': precision,\n            'recall': recall,\n            'specificity': specificity,\n            'ppv': ppv,\n            'npv': npv,\n            'auc': auc_value,\n            'confusion_matrix': cm,\n            'fpr': fpr,\n            'tpr': tpr\n        })\n        \n        # 打印结果\n        print(f\"\\nOrgan: {organ}\")\n        print(f\"  Accuracy: {accuracy:.4f}\")\n        print(f\"  Precision: {precision:.4f}\")\n        print(f\"  Recall: {recall:.4f}\")\n        print(f\"  Specificity: {specificity:.4f}\")\n        print(f\"  PPV: {ppv:.4f}\")\n        print(f\"  NPV: {npv:.4f}\")\n        print(f\"  AUC: {auc_value:.4f}\")\n        print(f\"  Confusion Matrix: \\n{cm}\")\n    \n    return results\n\n\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:16.105691Z","iopub.execute_input":"2025-04-22T15:42:16.105873Z","iopub.status.idle":"2025-04-22T15:42:16.122701Z","shell.execute_reply.started":"2025-04-22T15:42:16.105859Z","shell.execute_reply":"2025-04-22T15:42:16.122122Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def get_lr_schedule():\n    \"\"\"创建余弦退火学习率调度器\"\"\"\n    initial_learning_rate = 3e-4\n    \n    def cosine_decay_with_warmup(epoch):\n        # 前2个epoch为预热期，逐渐增加学习率\n        if epoch < 2:\n            return initial_learning_rate * ((epoch + 1) / 2)\n        \n        # 之后使用余弦衰减\n        decay_epochs = 15 - 2  # 总epochs减去预热epochs\n        epoch_in_decay_range = epoch - 2  # 调整epoch计数\n        \n        cosine_decay = 0.5 * (1 + np.cos(np.pi * epoch_in_decay_range / decay_epochs))\n        return initial_learning_rate * cosine_decay\n    \n    return tf.keras.callbacks.LearningRateScheduler(cosine_decay_with_warmup)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:16.123303Z","iopub.execute_input":"2025-04-22T15:42:16.123495Z","iopub.status.idle":"2025-04-22T15:42:16.14017Z","shell.execute_reply.started":"2025-04-22T15:42:16.123481Z","shell.execute_reply":"2025-04-22T15:42:16.139672Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 5k-fold交叉验证实现\ndef train_with_kfold_cv(train_df, organ_status_map, segmentation_map, unet_model, k=5):\n    \"\"\"\n    使用k-fold交叉验证训练和评估模型\n    \n    Args:\n        train_df: 训练数据DataFrame\n        organ_status_map: 器官状态映射\n        segmentation_map: 分割映射\n        unet_model: U-Net模型\n        k: 交叉验证折数\n    \n    Returns:\n        训练好的模型列表和评估结果\n    \"\"\"\n    logger.log(f\"Starting {k}-fold cross-validation...\")\n    \n    # 获取所有患者ID\n    patient_ids = np.array(train_df['patient_id'].unique())\n    \n    # 创建k-fold拆分器\n    kfold = KFold(n_splits=k, shuffle=True, random_state=42)\n    \n    # 加载元数据\n    meta_df = pd.read_csv(TRAIN_META)\n    meta_df['patient_id'] = meta_df['patient_id'].astype(str)\n    \n    # 加载DICOM标签\n    logger.log(\"Loading DICOM tags...\")\n    dicom_tags_df = pd.read_parquet(TRAIN_DICOM_TAGS)\n    temp = dicom_tags_df['SeriesInstanceUID'].str.split('.', expand=True)\n    dicom_tags_df['series_id'] = temp[8]\n    dicom_tags_df['series_id'] = dicom_tags_df['series_id'].astype(str)\n    dicom_tags_df['PatientID'] = dicom_tags_df['PatientID'].astype(str)\n    gc.collect()\n    \n    # 存储每个fold的模型和结果\n    fold_models = []\n    fold_results = []\n    \n    # 定义批量大小\n    batch_size = 4\n    \n    # 遍历每个fold\n    for fold, (train_idx, val_idx) in enumerate(kfold.split(patient_ids)):\n        logger.log(f\"\\n{'='*50}\\nFold {fold+1}/{k}\\n{'='*50}\")\n        \n        # 获取当前fold的训练和验证患者ID\n        train_patient_ids = patient_ids[train_idx]\n        val_patient_ids = patient_ids[val_idx]\n        \n        logger.log(f\"Train set: {len(train_patient_ids)} patients, Validation set: {len(val_patient_ids)} patients\")\n        \n        # 为当前fold创建缓存目录\n        fold_cache_dir = os.path.join(CACHE_DIR, f\"fold_{fold+1}\")\n        os.makedirs(fold_cache_dir, exist_ok=True)\n        \n        # 预处理患者数据\n        logger.log(f\"Preprocessing patients for fold {fold+1}...\")\n        preprocess_patients_parallel(\n            train_patient_ids, meta_df, dicom_tags_df, segmentation_map, \n            organ_status_map, unet_model, fold_cache_dir, \n            img_size=IMG_SIZE, max_slices=32, batch_size=10\n        )\n        \n        preprocess_patients_parallel(\n            val_patient_ids, meta_df, dicom_tags_df, segmentation_map, \n            organ_status_map, unet_model, fold_cache_dir, \n            img_size=IMG_SIZE, max_slices=32, batch_size=10\n        )\n        \n        # 创建数据集\n        logger.log(f\"Creating datasets for fold {fold+1}...\")\n        train_dataset, train_steps = create_organ_balanced_dataset(\n            train_patient_ids, fold_cache_dir, batch_size=batch_size, shuffle=True\n        )\n        \n        val_dataset, num_val_samples = create_tf_dataset(\n            val_patient_ids, fold_cache_dir, batch_size=batch_size, shuffle=False\n        )\n        \n        val_steps = num_val_samples // batch_size\n        \n        # 构建分类器模型\n        logger.log(f\"Building classifier model for fold {fold+1}...\")\n        input_shape = (IMG_SIZE[0], IMG_SIZE[1], 1)  # mask为1通道\n        loss_fn = weighted_cross_entropy()  # 使用加权交叉熵损失函数\n        optimizer = tf.keras.optimizers.Adam(learning_rate=3e-4)\n        \n        classifier_model = build_25d_classifier(input_shape=input_shape, num_organs=NUM_ORGANS)\n        classifier_model.compile(optimizer=optimizer, loss=loss_fn, metrics=['accuracy'])\n        \n        # 设置回调函数\n        logger.log(f\"Setting up callbacks for fold {fold+1}...\")\n        model_checkpoint = tf.keras.callbacks.ModelCheckpoint(\n            filepath=os.path.join(OUTPUT_DIR, f'classifier_fold_{fold+1}_best.keras'),\n            monitor='val_loss', save_best_only=True\n        )\n        \n        early_stopping = tf.keras.callbacks.EarlyStopping(\n            monitor='val_loss', patience=5, restore_best_weights=True\n        )\n        \n        reduce_lr = tf.keras.callbacks.ReduceLROnPlateau(\n            monitor='val_loss', factor=0.2, patience=2, min_lr=1e-6\n        )\n        \n        # 使用我们定义的学习率调度器\n        lr_scheduler = get_lr_schedule()\n        \n        tensorboard_callback = DetailedTensorBoard(\n            log_dir=f\"./logs/classifier_fold_{fold+1}\", histogram_freq=1\n        )\n        \n        callbacks = [model_checkpoint, early_stopping, reduce_lr, lr_scheduler, tensorboard_callback]\n        \n        # 训练模型\n        logger.log(f\"Training classifier for fold {fold+1}...\")\n        history = classifier_model.fit(\n            train_dataset,\n            validation_data=val_dataset,\n            epochs=15,\n            steps_per_epoch=train_steps,\n            validation_steps=val_steps,\n            callbacks=callbacks,\n            verbose=1\n        )\n        \n        # 保存模型\n        classifier_model.save(os.path.join(OUTPUT_DIR, f'classifier_fold_{fold+1}.keras'))\n        \n        # 评估模型\n        logger.log(f\"Evaluating classifier for fold {fold+1}...\")\n        evaluation_results = evaluate_model_by_organ(classifier_model, val_dataset, val_steps)\n        \n        # 保存评估结果\n        results_df = pd.DataFrame([\n            {\n                'fold': fold + 1,\n                'organ': result['organ'],\n                'accuracy': result['accuracy'],\n                'precision': result['precision'],\n                'recall': result['recall'],\n                'specificity': result['specificity'],\n                'ppv': result['ppv'],\n                'npv': result['npv'],\n                'auc': result['auc']\n            }\n            for result in evaluation_results\n        ])\n        \n        results_df.to_csv(os.path.join(OUTPUT_DIR, f'evaluation_results_fold_{fold+1}.csv'), index=False)\n        \n        # 绘制ROC曲线\n        plot_roc_curves(evaluation_results)\n        plt.savefig(os.path.join(OUTPUT_DIR, f'roc_curves_fold_{fold+1}.png'), dpi=300)\n        \n        # 绘制混淆矩阵\n        plot_confusion_matrices(evaluation_results)\n        plt.savefig(os.path.join(OUTPUT_DIR, f'confusion_matrices_fold_{fold+1}.png'), dpi=300)\n        \n        # 存储模型和结果\n        fold_models.append(classifier_model)\n        fold_results.append(evaluation_results)\n        \n        # 清理内存\n        gc.collect()\n    \n    # 计算所有fold的平均性能\n    logger.log(\"Calculating average performance across all folds...\")\n    \n    # 提取所有fold的指标\n    all_metrics = []\n    for fold, results in enumerate(fold_results):\n        for result in results:\n            all_metrics.append({\n                'fold': fold + 1,\n                'organ': result['organ'],\n                'accuracy': result['accuracy'],\n                'precision': result['precision'],\n                'recall': result['recall'],\n                'specificity': result['specificity'],\n                'ppv': result['ppv'],\n                'npv': result['npv'],\n                'auc': result['auc']\n            })\n    \n    # 创建DataFrame并计算平均值\n    all_metrics_df = pd.DataFrame(all_metrics)\n    avg_metrics = all_metrics_df.groupby('organ').mean().reset_index()\n    avg_metrics['fold'] = 'Average'\n    \n    # 保存平均指标\n    avg_metrics.to_csv(os.path.join(OUTPUT_DIR, 'average_metrics.csv'), index=False)\n    \n    # 打印平均性能\n    print(\"\\nAverage Performance Across All Folds:\")\n    for _, row in avg_metrics.iterrows():\n        organ = row['organ']\n        print(f\"\\nOrgan: {organ}\")\n        print(f\"  Accuracy: {row['accuracy']:.4f}\")\n        print(f\"  AUC: {row['auc']:.4f}\")\n        print(f\"  Precision: {row['precision']:.4f}\")\n        print(f\"  Recall (Sensitivity): {row['recall']:.4f}\")\n        print(f\"  Specificity: {row['specificity']:.4f}\")\n        print(f\"  PPV: {row['ppv']:.4f}\")\n        print(f\"  NPV: {row['npv']:.4f}\")\n    \n    return fold_models, fold_results\n","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 主函数 - 修改为使用外部缓存和预训练模型\ndef main():\n    logger = ProgressLogger()\n    logger.log(\"Starting main function\")\n    \n    # 重新定义缓存路径，使用已有的缓存\n    ORIGINAL_CACHE_DIR = os.path.join(OUTPUT_DIR, 'cache')  # 原始定义\n    EXTERNAL_CACHE_DIR = '/kaggle/input/rsna-cache/cache'  # 外部已有缓存\n    \n    # 检查外部缓存是否存在\n    if os.path.exists(EXTERNAL_CACHE_DIR):\n        logger.log(f\"Found external cache: {EXTERNAL_CACHE_DIR}\")\n        CACHE_DIR = EXTERNAL_CACHE_DIR  # 使用外部缓存\n    else:\n        logger.log(f\"External cache not found, will use original cache path: {ORIGINAL_CACHE_DIR}\")\n        CACHE_DIR = ORIGINAL_CACHE_DIR\n    \n    # 1. 数据加载与预处理\n    logger.log(\"Loading training data...\")\n    train_df, organ_status_map, segmentation_map = load_and_process_data()\n    \n    # 设置5k-fold交叉验证\n    kfold = KFold(n_splits=5, shuffle=True, random_state=42)\n    patient_ids = train_df['patient_id'].unique()\n    \n    meta_df = pd.read_csv(TRAIN_META)\n    meta_df['patient_id'] = meta_df['patient_id'].astype(str)\n    \n    logger.log(\"Loading DICOM tags...\")\n    dicom_tags_df = pd.read_parquet(TRAIN_DICOM_TAGS)\n    temp = dicom_tags_df['SeriesInstanceUID'].str.split('.', expand=True)\n    dicom_tags_df['series_id'] = temp[8]\n    dicom_tags_df['series_id'] = dicom_tags_df['series_id'].astype(str)\n    dicom_tags_df['PatientID'] = dicom_tags_df['PatientID'].astype(str)\n    gc.collect()\n    logger.log(\"Data loading complete\")\n\n    # 定义全局批量大小\n    batch_size = 4\n\n    # 2. U-Net模型训练或加载\n    unet_input_shape = (IMG_SIZE[0], IMG_SIZE[1], 3)\n    \n    # 检查是否有预训练U-Net模型 - 按优先级顺序尝试不同路径\n    unet_model_paths = [\n        '/kaggle/input/unet-model/unet_model.keras',   # 外部上传模型\n        '/kaggle/input/unet_model/pytorch/default/1/unet_model.keras',  # 另一个可能的路径\n        'unet_best.keras',  # 最佳模型\n        'unet_model.keras'  # 标准模型\n    ]\n    \n    unet_model = None\n    for model_path in unet_model_paths:\n        if os.path.exists(model_path):\n            logger.log(f\"Loading pre-trained U-Net model from: {model_path}\")\n            try:\n                unet_model = tf.keras.models.load_model(model_path)\n                logger.log(\"U-Net model loaded successfully!\")\n\n                # 在这里调用可视化函数\n                sample_patient_id = patient_ids[0]  # 使用第一个训练集患者\n                visualize_segmentation_results(unet_model, sample_patient_id, meta_df, dicom_tags_df, segmentation_map, num_samples=3)\n                \n                break\n            except Exception as e:\n                logger.log(f\"Failed to load model {model_path}: {str(e)}\")\n                continue\n    \n    if unet_model is None:\n        logger.log(\"No pre-trained model found, will train a new one...\")\n        unet_model = build_unet(unet_input_shape)\n        unet_model.compile(\n            optimizer=tf.keras.optimizers.Adam(learning_rate=1e-4),\n            loss='binary_crossentropy',\n            metrics=['accuracy']\n        )\n        \n        # 为U-Net训练创建tf.data.Dataset\n        logger.log(\"Creating U-Net training dataset...\")\n        \n        # 定义加载切片的函数\n        def load_slice_for_unet(patient_id, meta_df, dicom_tags_df, segmentation_map):\n            try:\n                # 将字符串张量转换为Python字符串\n                pid = patient_id.numpy().decode('utf-8')\n                \n                # 获取患者元数据\n                subset = meta_df[meta_df['patient_id'] == pid]\n                if subset.empty:\n                    return np.zeros((224, 224, 3), dtype=np.float32), np.zeros((224, 224, 1), dtype=np.float32)\n                \n                series_id = str(subset['series_id'].iloc[0])\n                \n                # 获取图像ID列表\n                image_ids = dicom_tags_df[\n                    (dicom_tags_df['PatientID'] == pid) &\n                    (dicom_tags_df['series_id'] == series_id)\n                ]['InstanceNumber'].values\n                \n                if len(image_ids) == 0:\n                    return np.zeros((224, 224, 3), dtype=np.float32), np.zeros((224, 224, 1), dtype=np.float32)\n                \n                # 随机选择一个切片\n                image_id = np.random.choice(image_ids)\n                image_file = os.path.join(TRAIN_IMAGES, str(pid), str(series_id), f'{image_id}.dcm')\n                \n                # 加载图像\n                dicom = pydicom.dcmread(image_file)\n                image = dicom.pixel_array.astype(np.float32)\n                image = cv2.resize(image, IMG_SIZE)\n                max_val = np.max(image)\n                image = image / max_val if max_val != 0 else image\n                \n                # 确保图像是3通道的\n                if len(image.shape) == 2:\n                    image = np.stack([image, image, image], axis=-1)\n                \n                # 加载分割掩码\n                mask = np.zeros(IMG_SIZE, dtype=np.float32)\n                segmentation_file = segmentation_map.get(pid)\n                if segmentation_file and os.path.exists(segmentation_file):\n                    try:\n                        segmentation_data = nib.load(segmentation_file).get_fdata().astype(np.float32)\n                        if image_id <= segmentation_data.shape[2]:\n                            mask_slice = segmentation_data[:, :, image_id-1]\n                            mask = cv2.resize(mask_slice, IMG_SIZE, interpolation=cv2.INTER_NEAREST)\n                    except:\n                        pass\n                \n                # 数据增强\n                image, mask = augment_data(image, mask)\n                \n                # 确保掩码是单通道的\n                if len(mask.shape) == 2:\n                    mask = np.expand_dims(mask, axis=-1)\n                \n                return image, mask\n            except:\n                return np.zeros((224, 224, 3), dtype=np.float32), np.zeros((224, 224, 1), dtype=np.float32)\n        \n        # 创建tf.data.Dataset\n        def create_unet_dataset(patient_ids, batch_size=4, shuffle=True):\n            # 创建一个包含患者ID的数据集\n            patient_ds = tf.data.Dataset.from_tensor_slices(patient_ids)\n            \n            # 如果需要洗牌\n            if shuffle:\n                patient_ds = patient_ds.shuffle(buffer_size=len(patient_ids), reshuffle_each_iteration=True)\n            \n            # 加载切片\n            def load_slice_wrapper(patient_id):\n                image, mask = tf.py_function(\n                    func=lambda x: load_slice_for_unet(x, meta_df, dicom_tags_df, segmentation_map),\n                    inp=[patient_id],\n                    Tout=[tf.float32, tf.float32]\n                )\n                # 设置形状信息\n                image.set_shape((224, 224, 3))\n                mask.set_shape((224, 224, 1))\n                return image, mask\n            \n            # 映射到加载函数\n            dataset = patient_ds.map(\n                load_slice_wrapper,\n                num_parallel_calls=tf.data.AUTOTUNE\n            )\n            \n            # 过滤无效数据\n            def is_valid(image, mask):\n                return tf.math.reduce_all(tf.math.is_finite(image))\n            \n            dataset = dataset.filter(is_valid)\n            \n            # 批处理\n            dataset = dataset.batch(batch_size)\n            \n            # 预取下一批数据\n            dataset = dataset.prefetch(tf.data.AUTOTUNE)\n            \n            return dataset\n        \n        # 从train_df中获取所有患者ID并划分训练和验证集\n        from sklearn.model_selection import train_test_split\n        train_patient_ids, val_patient_ids = train_test_split(patient_ids, test_size=0.2, random_state=42)\n        logger.log(f\"Dataset split complete: {len(train_patient_ids)} training samples, {len(val_patient_ids)} validation samples\")\n        \n        # 创建训练和验证数据集\n        logger.log(\"Creating U-Net training and validation datasets...\")\n        unet_train_dataset = create_unet_dataset(train_patient_ids, batch_size=batch_size, shuffle=True)\n        unet_val_dataset = create_unet_dataset(val_patient_ids, batch_size=batch_size, shuffle=False)\n        \n        # 计算steps_per_epoch和validation_steps\n        unet_steps_per_epoch = len(train_patient_ids) // batch_size\n        unet_validation_steps = len(val_patient_ids) // batch_size\n        \n        gc.collect()\n        \n        # 定义回调\n        logger.log(\"Setting up U-Net training callbacks...\")\n        unet_checkpoint = tf.keras.callbacks.ModelCheckpoint(\n            'unet_best.keras', monitor='val_loss', save_best_only=True\n        )\n        \n        early_stopping = tf.keras.callbacks.EarlyStopping(\n            monitor='val_loss', patience=3, restore_best_weights=True\n        )\n        \n        tensorboard_callback = DetailedTensorBoard(\n            log_dir=\"./logs/unet\", histogram_freq=1\n        )\n        \n        # 使用tf.data API进行训练\n        logger.log(\"Starting U-Net model training...\")\n        history_unet = unet_model.fit(\n            unet_train_dataset,\n            validation_data=unet_val_dataset,\n            epochs=5,\n            steps_per_epoch=unet_steps_per_epoch,  # 明确指定每个epoch的步数\n            validation_steps=unet_validation_steps,  # 明确指定验证步数\n            callbacks=[unet_checkpoint, early_stopping, tensorboard_callback],\n            verbose=1\n        )\n        \n        logger.log(\"Saving U-Net model...\")\n        unet_model.save('unet_model.keras')\n        \n        # 可视化训练历史记录\n        logger.log(\"Visualizing U-Net training history...\")\n        plt.figure(figsize=(12, 5))\n        \n        plt.subplot(1, 2, 1)\n        plt.plot(history_unet.history['loss'], label='Training Loss')\n        plt.plot(history_unet.history['val_loss'], label='Validation Loss')\n        plt.title('Model Loss')\n        plt.ylabel('Loss')\n        plt.xlabel('Epoch')\n        plt.legend()\n        \n        plt.subplot(1, 2, 2)\n        plt.plot(history_unet.history['accuracy'], label='Training Accuracy')\n        plt.plot(history_unet.history['val_accuracy'], label='Validation Accuracy')\n        plt.title('Model Accuracy')\n        plt.ylabel('Accuracy')\n        plt.xlabel('Epoch')\n        plt.legend()\n        \n        plt.tight_layout()\n        plt.savefig(os.path.join(OUTPUT_DIR, 'unet_training_history.png'), dpi=300)\n        plt.show()\n        \n        del unet_train_dataset, unet_val_dataset\n        gc.collect()\n    \n    # 3. 检查缓存文件并决定是否需要预处理\n    # 检查cache文件是否完整\n    def check_cache_files(patient_ids, cache_dir):\n        valid_cached_ids = []\n        missing_ids = []\n        for pid in tqdm(patient_ids, desc=\"Checking cache files\"):\n            image_file = os.path.join(cache_dir, f\"{pid}.npz\")\n            label_file = os.path.join(cache_dir, f\"{pid}_labels.npz\")\n            if os.path.exists(image_file) and os.path.exists(label_file):\n                valid_cached_ids.append(pid)\n            else:\n                missing_ids.append(pid)\n        return valid_cached_ids, missing_ids\n\n    # 实现5k-fold交叉验证\n    logger.log(\"Implementing 5-fold cross-validation...\")\n    fold_results = []\n    \n    for fold, (train_idx, val_idx) in enumerate(kfold.split(patient_ids)):\n        logger.log(f\"Starting fold {fold+1}/5\")\n        \n        # 获取当前fold的训练和验证患者ID\n        train_patient_ids = patient_ids[train_idx]\n        val_patient_ids = patient_ids[val_idx]\n        \n        # 为当前fold创建缓存目录\n        fold_cache_dir = os.path.join(CACHE_DIR, f\"fold_{fold+1}\")\n        os.makedirs(fold_cache_dir, exist_ok=True)\n        \n        # 检查当前fold的cache文件\n        logger.log(f\"Checking cache files for fold {fold+1}...\")\n        cached_train_ids, missing_train_ids = check_cache_files(train_patient_ids, fold_cache_dir)\n        cached_val_ids, missing_val_ids = check_cache_files(val_patient_ids, fold_cache_dir)\n\n        logger.log(f\"Found {len(cached_train_ids)}/{len(train_patient_ids)} cached training samples\")\n        logger.log(f\"Found {len(cached_val_ids)}/{len(val_patient_ids)} cached validation samples\")\n        \n        # 如果有缺失的缓存文件，并且我们可以写入缓存目录\n        if (len(missing_train_ids) > 0 or len(missing_val_ids) > 0) and os.access(fold_cache_dir, os.W_OK):\n            logger.log(f\"Need to generate cache for {len(missing_train_ids) + len(missing_val_ids)} samples\")\n            \n            # 确保缓存目录存在\n            os.makedirs(fold_cache_dir, exist_ok=True)\n            \n            # 只为缺失的样本生成缓存\n            if len(missing_train_ids) > 0:\n                logger.log(f\"Preprocessing {len(missing_train_ids)} training samples...\")\n                preprocess_patients_parallel(\n                    missing_train_ids, meta_df, dicom_tags_df, segmentation_map, \n                    organ_status_map, unet_model, fold_cache_dir, \n                    img_size=IMG_SIZE, max_slices=32, batch_size=10\n                )\n            \n            if len(missing_val_ids) > 0:\n                logger.log(f\"Preprocessing {len(missing_val_ids)} validation samples...\")\n                preprocess_patients_parallel(\n                    missing_val_ids, meta_df, dicom_tags_df, segmentation_map, \n                    organ_status_map, unet_model, fold_cache_dir, \n                    img_size=IMG_SIZE, max_slices=32, batch_size=10\n                )\n        elif len(missing_train_ids) > 0 or len(missing_val_ids) > 0:\n            logger.log(f\"Warning: Cache incomplete but {fold_cache_dir} is not writable\")\n            logger.log(\"Will use available cache files, but model performance may be affected\")\n            \n            # 更新训练和验证ID列表，只使用有缓存的ID\n            train_patient_ids = cached_train_ids\n            val_patient_ids = cached_val_ids\n            \n            if len(train_patient_ids) == 0 or len(val_patient_ids) == 0:\n                logger.log(f\"Skipping fold {fold+1} due to insufficient cached data\")\n                continue\n        else:\n            logger.log(\"All data already cached, no preprocessing needed\")\n        \n        # 4. 创建数据集\n        # 使用平衡采样创建训练数据集\n        logger.log(f\"Creating training and validation datasets for fold {fold+1}...\")\n        train_dataset, train_steps = create_organ_balanced_dataset(train_patient_ids, fold_cache_dir, batch_size=batch_size, shuffle=True)\n        \n        # 验证集使用标准数据集，保持原始分布\n        val_dataset, num_val_samples = create_tf_dataset(val_patient_ids, fold_cache_dir, batch_size=batch_size, shuffle=False)\n        \n        # 计算分类器训练的steps_per_epoch和validation_steps\n        classifier_steps_per_epoch = train_steps\n        classifier_validation_steps = num_val_samples // batch_size\n        \n        logger.log(f\"Classifier training steps: {classifier_steps_per_epoch} steps/epoch, validation steps: {classifier_validation_steps} steps/epoch\")\n        \n        # 5. 构建和训练分类器模型\n        logger.log(f\"Building classifier model for fold {fold+1}...\")\n        NUM_ORGANS = 5\n        input_shape = (IMG_SIZE[0], IMG_SIZE[1], 1)  # mask为1通道\n        loss_fn = weighted_cross_entropy()  # 使用加权交叉熵损失函数\n        optimizer = tf.keras.optimizers.Adam(learning_rate=3e-4)  # 使用适当的学习率\n\n        # 创建新模型\n        classifier_model = build_25d_classifier(input_shape=input_shape, num_organs=NUM_ORGANS)\n        classifier_model.compile(optimizer=optimizer, loss=loss_fn, metrics=['accuracy'])\n            \n        classifier_model.summary()\n\n        # 设置回调函数\n        logger.log(f\"Setting up classifier training callbacks for fold {fold+1}...\")\n        classifier_checkpoint = tf.keras.callbacks.ModelCheckpoint(\n            filepath=f'classifier_fold_{fold+1}_best.keras',\n            monitor='val_loss', save_best_only=True\n        )\n        \n        classifier_logger = tf.keras.callbacks.CSVLogger(f\"classifier_fold_{fold+1}_train.log\")\n        \n        early_stopping = tf.keras.callbacks.EarlyStopping(\n            monitor='val_loss', patience=5, restore_best_weights=True\n        )\n        \n        reduce_lr = tf.keras.callbacks.ReduceLROnPlateau(\n            monitor='val_loss', factor=0.2, patience=2, min_lr=1e-6\n        )\n        \n        # 使用我们定义的学习率调度器\n        lr_scheduler = get_lr_schedule()\n        \n        # 创建TensorBoard回调\n        tensorboard_callback = DetailedTensorBoard(\n            log_dir=f\"./logs/classifier_fold_{fold+1}\", histogram_freq=1\n        )\n        \n        # 创建类别平衡监控器\n        balance_monitor = ClassBalanceMonitor(train_dataset, min(5, classifier_steps_per_epoch))\n        \n        # 组合所有回调\n        callbacks = [\n            classifier_checkpoint, \n            classifier_logger,\n            early_stopping,\n            reduce_lr,\n            lr_scheduler,\n            tensorboard_callback,\n            balance_monitor\n        ]\n        \n        # 可视化类别分布\n        visualize_class_distribution(train_dataset, min(10, classifier_steps_per_epoch), f\"Training Set Class Distribution (After Balancing) - Fold {fold+1}\")\n        visualize_class_distribution(val_dataset, min(10, classifier_validation_steps), f\"Validation Set Class Distribution - Fold {fold+1}\")\n        \n        # 训练分类器模型\n        logger.log(f\"Starting classifier training for fold {fold+1}...\")\n        history_classifier = classifier_model.fit(\n            train_dataset,\n            validation_data=val_dataset,\n            epochs=15,  # 增加到15轮\n            steps_per_epoch=classifier_steps_per_epoch,\n            validation_steps=classifier_validation_steps,\n            callbacks=callbacks,\n            verbose=1\n        )\n        \n        logger.log(f\"Saving classifier model for fold {fold+1}...\")\n        classifier_model.save(f'classifier_fold_{fold+1}.keras')\n        \n        # 可视化训练历史记录\n        plt.figure(figsize=(12, 5))\n        \n        plt.subplot(1, 2, 1)\n        plt.plot(history_classifier.history['loss'], label='Training Loss')\n        plt.plot(history_classifier.history['val_loss'], label='Validation Loss')\n        plt.title('Model Loss')\n        plt.ylabel('Loss')\n        plt.xlabel('Epoch')\n        plt.legend()\n        \n        plt.subplot(1, 2, 2)\n        plt.plot(history_classifier.history['accuracy'], label='Training Accuracy')\n        plt.plot(history_classifier.history['val_accuracy'], label='Validation Accuracy')\n        plt.title('Model Accuracy')\n        plt.ylabel('Accuracy')\n        plt.xlabel('Epoch')\n        plt.legend()\n        \n        plt.tight_layout()\n        plt.savefig(os.path.join(OUTPUT_DIR, f'training_history_fold_{fold+1}.png'), dpi=300)\n        plt.show()\n        \n        # 6. 分类模型评估\n        logger.log(f\"Evaluating classifier model performance for fold {fold+1}...\")\n        evaluation_results = evaluate_model_by_organ(classifier_model, val_dataset, classifier_validation_steps)\n        \n        # 绘制ROC曲线\n        plot_roc_curves(evaluation_results)\n        \n        # 绘制混淆矩阵\n        plot_confusion_matrices(evaluation_results)\n        \n        # 保存评估结果\n        results_df = pd.DataFrame([\n            {\n                'fold': fold + 1,\n                'organ': result['organ'],\n                'accuracy': result['accuracy'],\n                'precision': result['precision'],\n                'recall': result['recall'],\n                'specificity': result['specificity'],\n                'ppv': result['ppv'],\n                'npv': result['npv'],\n                'auc': result['auc']\n            }\n            for result in evaluation_results\n        ])\n        results_df.to_csv(os.path.join(OUTPUT_DIR, f'evaluation_results_fold_{fold+1}.csv'), index=False)\n        logger.log(f\"Evaluation results saved to {os.path.join(OUTPUT_DIR, f'evaluation_results_fold_{fold+1}.csv')}\")\n        \n        # 添加到总结果\n        fold_results.append(evaluation_results)\n        \n        # 7. 可视化分类结果\n        logger.log(f\"Visualizing classification results for fold {fold+1}...\")\n        # 随机选择3个患者进行可视化\n        if isinstance(val_patient_ids, np.ndarray):\n            sample_patients = random.sample(val_patient_ids.tolist(), min(3, len(val_patient_ids)))\n        else:\n            sample_patients = random.sample(list(val_patient_ids), min(3, len(val_patient_ids)))\n        \n        for patient_id in sample_patients:\n            try:\n                visualize_classification_results(classifier_model, unet_model, patient_id, meta_df, \n                                            dicom_tags_df, segmentation_map, organ_status_map)\n            except Exception as e:\n                logger.log(f\"Error visualizing results for patient {patient_id}: {str(e)}\")\n        \n        # 清理内存\n        gc.collect()\n    \n    # 计算所有fold的平均性能\n    logger.log(\"Calculating average performance across all folds...\")\n    \n    # 提取所有fold的指标\n    all_metrics = []\n    for fold, results in enumerate(fold_results):\n        for result in results:\n            all_metrics.append({\n                'fold': fold + 1,\n                'organ': result['organ'],\n                'accuracy': result['accuracy'],\n                'precision': result['precision'],\n                'recall': result['recall'],\n                'specificity': result['specificity'],\n                'ppv': result['ppv'],\n                'npv': result['npv'],\n                'auc': result['auc']\n            })\n    \n    # 创建DataFrame并计算平均值\n    all_metrics_df = pd.DataFrame(all_metrics)\n    avg_metrics = all_metrics_df.groupby('organ').mean().reset_index()\n    avg_metrics['fold'] = 'Average'\n    \n    # 保存平均指标\n    avg_metrics.to_csv(os.path.join(OUTPUT_DIR, 'average_metrics.csv'), index=False)\n    \n    # 打印平均性能\n    print(\"\\nAverage Performance Across All Folds:\")\n    print(\"=\" * 60)\n    for _, row in avg_metrics.iterrows():\n        organ = row['organ']\n        print(f\"\\nOrgan: {organ}\")\n        print(f\"  Accuracy: {row['accuracy']:.4f}\")\n        print(f\"  AUC: {row['auc']:.4f}\")\n        print(f\"  Precision: {row['precision']:.4f}\")\n        print(f\"  Recall (Sensitivity): {row['recall']:.4f}\")\n        print(f\"  Specificity: {row['specificity']:.4f}\")\n        print(f\"  PPV: {row['ppv']:.4f}\")\n        print(f\"  NPV: {row['npv']:.4f}\")\n    print(\"=\" * 60)\n    \n    # 生成最终的结果表格\n    fig, ax = plt.subplots(figsize=(12, 6))\n    \n    # 隐藏轴线\n    ax.axis('tight')\n    ax.axis('off')\n    \n    # 创建表格数据\n    table_data = []\n    for _, row in avg_metrics.iterrows():\n        table_data.append([\n            row['organ'].capitalize(),\n            f\"{row['auc']:.3f}\",\n            f\"{row['accuracy']:.3f}\",\n            f\"{row['ppv']:.3f}\",\n            f\"{row['npv']:.3f}\",\n            f\"{row['recall']:.3f}\",\n            f\"{row['specificity']:.3f}\"\n        ])\n    \n    # 创建表格\n    table = ax.table(\n        cellText=table_data,\n        colLabels=['Organ', 'AUC', 'Accuracy', 'PPV', 'NPV', 'Sensitivity', 'Specificity'],\n        loc='center',\n        cellLoc='center',\n        colColours=['#f2f2f2'] * 7\n    )\n    \n    # 设置表格样式\n    table.auto_set_font_size(False)\n    table.set_fontsize(10)\n    table.scale(1.2, 1.5)\n    \n    # 设置标题\n    plt.title('Average Performance Metrics Across All Folds', fontsize=14, pad=20)\n    plt.tight_layout()\n    \n    # 保存表格\n    plt.savefig(os.path.join(OUTPUT_DIR, 'performance_summary_table.png'), dpi=300, bbox_inches='tight')\n    plt.show()\n    \n    logger.log(\"Training and evaluation complete!\")\n    gc.collect()\n    \n    return logger  # 返回日志记录器以便在主脚本中使用\n\n\nif __name__ == \"__main__\":\n    # 设置TensorFlow日志级别，只显示警告和错误\n    import os\n    os.environ['TF_CPP_MIN_LOG_LEVEL'] = '1'  # 0=全部显示, 1=不显示INFO, 2=不显示INFO和WARNING\n    \n    # 清理内存\n    gc.collect()\n    \n    # 创建日志记录器\n    logger = ProgressLogger()\n    \n    # 记录开始时间\n    start_time = time.time()\n    logger.log(\"Program execution started\")\n    \n    try:\n        # 执行主函数，获取返回的logger\n        main_logger = main()\n        \n        # 如果main函数返回了logger，使用它；否则使用当前logger\n        active_logger = main_logger if main_logger else logger\n    except Exception as e:\n        logger.log(f\"Error during execution: {str(e)}\")\n        import traceback\n        logger.log(traceback.format_exc())\n        raise\n    finally:\n        # 记录结束时间\n        end_time = time.time()\n        elapsed = end_time - start_time\n        hours, remainder = divmod(elapsed, 3600)\n        minutes, seconds = divmod(remainder, 60)\n        \n        logger.log(f\"Program execution complete, total time: {int(hours)} hours {int(minutes)} minutes {seconds:.2f} seconds\")\n        \n        # 最终清理内存\n        gc.collect()\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-22T15:42:16.141045Z","iopub.execute_input":"2025-04-22T15:42:16.141297Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 生成预测文件","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}