{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.10.12","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":52254,"databundleVersionId":9674523,"sourceType":"competition"}],"dockerImageVersionId":30918,"isInternetEnabled":true,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"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","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T04:04:57.875485Z","iopub.execute_input":"2025-04-11T04:04:57.876908Z","iopub.status.idle":"2025-04-11T04:04:57.890967Z","shell.execute_reply.started":"2025-04-11T04:04:57.876841Z","shell.execute_reply":"2025-04-11T04:04:57.888637Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 检查GPU是否可用\nprint(\"TensorFlow版本:\", tf.__version__)\nprint(\"GPU是否可用:\", tf.config.list_physical_devices('GPU'))","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T04:04:57.893203Z","iopub.execute_input":"2025-04-11T04:04:57.893709Z","iopub.status.idle":"2025-04-11T04:04:57.925459Z","shell.execute_reply.started":"2025-04-11T04:04:57.893664Z","shell.execute_reply":"2025-04-11T04:04:57.923572Z"}},"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-11T04:04:57.92725Z","iopub.execute_input":"2025-04-11T04:04:57.927657Z","iopub.status.idle":"2025-04-11T04:04:58.012086Z","shell.execute_reply.started":"2025-04-11T04:04:57.927618Z","shell.execute_reply":"2025-04-11T04:04:58.010495Z"}},"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'\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-11T04:04:58.017381Z","iopub.execute_input":"2025-04-11T04:04:58.017781Z","iopub.status.idle":"2025-04-11T04:04:58.043828Z","shell.execute_reply.started":"2025-04-11T04:04:58.017751Z","shell.execute_reply":"2025-04-11T04:04:58.04187Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 加载训练标签并处理数据\ndef load_and_process_data():\n    \"\"\"Load and process training data\"\"\"\n    print(\"Loading training data...\")\n    train_df = pd.read_csv(TRAIN_CSV)\n    print(f\"Training data shape: {train_df.shape}\")\n    print(f\"Data type of patient_id in train_df: {train_df['patient_id'].dtype}\") #Add this line\n    print(f\"Sample rows from train_df: \\n{train_df.head()}\") #Add this line\n    train_df['patient_id'] = train_df['patient_id'].astype(str)\n\n    # 显示前几行数据\n    print(\"\\nTraining data sample:\")\n    print(train_df.head())\n\n    # 检查列名\n    print(\"\\nColumn names:\")\n    print(train_df.columns.tolist())\n\n    # 处理不同类型的器官标签\n    organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n\n    # 创建图表\n    plt.figure(figsize=(20, 15))\n\n    # 为每个器官创建饼图\n    for i, organ in enumerate(organs):\n        plt.subplot(2, 3, i+1)\n\n        # 收集该器官的所有状态数据\n        status_data = {}\n\n        # 检查健康状态\n        if f'{organ}_healthy' in train_df.columns:\n            healthy_count = train_df[f'{organ}_healthy'].sum()\n            total = len(train_df)\n            status_data['Healthy'] = healthy_count\n            print(f\"\\n{organ.capitalize()} - Healthy: {healthy_count} ({healthy_count/total*100:.2f}%)\")\n\n        # 检查损伤状态\n        if f'{organ}_injury' in train_df.columns:\n            injury_count = train_df[f'{organ}_injury'].sum()\n            total = len(train_df)\n            status_data['Injury'] = injury_count\n            print(f\"{organ.capitalize()} - Injury: {injury_count} ({injury_count/total*100:.2f}%)\")\n\n        # 检查低度损伤\n        if f'{organ}_low' in train_df.columns:\n            low_count = train_df[f'{organ}_low'].sum()\n            total = len(train_df)\n            status_data['Low-grade Injury'] = low_count\n            print(f\"{organ.capitalize()} - Low-grade Injury: {low_count} ({low_count/total*100:.2f}%)\")\n\n        # 检查高度损伤\n        if f'{organ}_high' in train_df.columns:\n            high_count = train_df[f'{organ}_high'].sum()\n            total = len(train_df)\n            status_data['High-grade Injury'] = high_count\n            print(f\"{organ.capitalize()} - High-grade Injury: {high_count} ({high_count/total*100:.2f}%)\")\n\n        # 绘制饼图\n        if status_data:\n            labels = status_data.keys()\n            sizes = status_data.values()\n\n            # 计算百分比\n            total = sum(sizes)\n            sizes_percent = [size/total*100 for size in sizes]\n\n            # 添加百分比到标签\n            labels_with_percent = [f'{label}: {size:.1f}%' for label, size in zip(labels, sizes_percent)]\n\n            # 设置颜色\n            colors = plt.cm.Paired(np.arange(len(sizes)) / len(sizes))\n\n            # 突出显示损伤部分\n            explode = [0] * len(sizes)\n            for j, label in enumerate(labels):\n                if 'Injury' in label:\n                    explode[j] = 0.1\n\n            # 绘制饼图\n            plt.pie(sizes, explode=explode, labels=labels_with_percent,\n                    colors=colors, autopct='%1.1f%%', shadow=True, startangle=90)\n            plt.axis('equal')  # 确保饼图是圆的\n            plt.title(f'{organ.capitalize()} Status Distribution', fontsize=15)\n\n    plt.tight_layout()\n    plt.savefig(os.path.join(OUTPUT_DIR, 'organ_status_distribution.png'), dpi=300)\n    plt.show()\n\n    # 绘制器官损伤比例的条形图\n    plt.figure(figsize=(12, 8))\n\n    # 收集所有器官的损伤百分比\n    organ_injury_percentages = []\n\n    for organ in organs:\n        # 计算损伤的总数（包括低度和高度损伤）\n        injury_count = 0\n\n        if f'{organ}_injury' in train_df.columns:\n            injury_count += train_df[f'{organ}_injury'].sum()\n\n        if f'{organ}_low' in train_df.columns:\n            injury_count += train_df[f'{organ}_low'].sum()\n\n        if f'{organ}_high' in train_df.columns:\n            injury_count += train_df[f'{organ}_high'].sum()\n\n        total = len(train_df)\n        injury_percentage = injury_count / total * 100\n        organ_injury_percentages.append(injury_percentage)\n\n    # 绘制条形图\n    bars = plt.bar(organs, organ_injury_percentages, color=plt.cm.viridis(np.linspace(0, 1, len(organs))))\n\n    # 添加数值标签\n    for bar, percentage in zip(bars, organ_injury_percentages):\n        plt.text(bar.get_x() + bar.get_width()/2, bar.get_height() + 0.5,\n                f'{percentage:.1f}%', ha='center', va='bottom', fontsize=12)\n\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(organ_injury_percentages) * 1.2)  # 设置y轴上限\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    # 如果数据集中有损伤程度的区分，绘制损伤严重程度分布图\n    has_severity_data = any(f'{organ}_low' in train_df.columns or f'{organ}_high' in train_df.columns for organ in organs)\n\n    if has_severity_data:\n        plt.figure(figsize=(14, 10))\n\n        # 为每个有严重程度区分的器官创建子图\n        severity_organs = [organ for organ in organs if f'{organ}_low' in train_df.columns or f'{organ}_high' in train_df.columns]\n\n        for i, organ in enumerate(severity_organs):\n            plt.subplot(2, 3, i+1)\n\n            # 收集该器官的损伤严重程度数据\n            severity_data = {}\n\n            if f'{organ}_low' in train_df.columns:\n                low_count = train_df[f'{organ}_low'].sum()\n                severity_data['Low-grade Injury'] = low_count\n\n            if f'{organ}_high' in train_df.columns:\n                high_count = train_df[f'{organ}_high'].sum()\n                severity_data['High-grade Injury'] = high_count\n\n            if severity_data:\n                # 绘制饼图\n                labels = severity_data.keys()\n                sizes = severity_data.values()\n\n                # 计算百分比\n                total = sum(sizes)\n                if total > 0:  # 避免除以零\n                    sizes_percent = [size/total*100 for size in sizes]\n\n                    # 添加百分比到标签\n                    labels_with_percent = [f'{label}: {size:.1f}%' for label, size in zip(labels, sizes_percent)]\n\n                    # 设置颜色\n                    colors = ['#ff9999', '#ff3333']  # 浅红色和深红色分别表示低度和高度损伤\n\n                    # 绘制饼图\n                    plt.pie(sizes, labels=labels_with_percent, colors=colors,\n                            autopct='%1.1f%%', shadow=True, startangle=90)\n                    plt.axis('equal')\n                    plt.title(f'{organ.capitalize()} Injury Severity Distribution', fontsize=15)\n\n        plt.tight_layout()\n        plt.savefig(os.path.join(OUTPUT_DIR, 'injury_severity_distribution.png'), dpi=300)\n        plt.show()\n\n    # 检查是否有缺失值\n    print(\"\\nChecking for missing values:\")\n    missing_values = train_df.isnull().sum()\n    if missing_values.sum() > 0:\n        print(missing_values[missing_values > 0])\n\n        # 可视化缺失值\n        plt.figure(figsize=(10, 6))\n        missing_cols = missing_values[missing_values > 0].index\n        plt.bar(missing_cols, missing_values[missing_values > 0], 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\")\n\n    # 创建一个器官状态的映射表，用于后续处理\n    organ_status_map = {}\n    for index, row in train_df.iterrows():\n        patient_id = row['patient_id']\n        organ_statuses = {}\n        for organ in organs:\n            # 检查该器官的所有可能状态列\n            status_columns = [col for col in train_df.columns if col.startswith(f'{organ}_')]\n\n            # 记录该器官的状态\n            organ_status = {}\n            for col in status_columns:\n                status_type = col.split('_')[1]  # 获取状态类型（healthy, injury, low, high）\n                organ_status[status_type] = row[col]\n\n            organ_statuses[organ] = organ_status\n\n        organ_status_map[patient_id] = organ_statuses\n\n    # 构建 segmentation_map\n    # 利用 train_series_meta 文件建立 series_id 到 patient_id 的映射\n    train_series_meta = pd.read_csv(TRAIN_META)\n    # 假设 train_series_meta 包含 'series_id' 和 'patient_id' 两列\n    series_to_patient = dict(zip(train_series_meta['series_id'], train_series_meta['patient_id']))\n    segmentation_map = {}\n    for f in glob.glob(os.path.join(SEGMENTATION_DIR, \"*.nii\")):\n        try:\n            series_id = int(os.path.splitext(os.path.basename(f))[0])\n            segmentation_map[str(series_id)] = f # force it to be string\n        except ValueError:\n            print(f\"Could not extract series_id from {f}\")\n\n    # 保存图表\n    print(f\"\\nVisualization results saved to {OUTPUT_DIR}\")\n\n    return train_df, organ_status_map, segmentation_map","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T04:04:58.045764Z","iopub.execute_input":"2025-04-11T04:04:58.046163Z","iopub.status.idle":"2025-04-11T04:04:58.100258Z","shell.execute_reply.started":"2025-04-11T04:04:58.04613Z","shell.execute_reply":"2025-04-11T04:04:58.097972Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"# 加载并处理数据\ntrain_df, organ_status_map, segmentation_map = load_and_process_data()\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-11T04:04:58.103962Z","iopub.execute_input":"2025-04-11T04:04:58.10439Z","iopub.status.idle":"2025-04-11T04:05:04.439583Z","shell.execute_reply.started":"2025-04-11T04:04:58.104357Z","shell.execute_reply":"2025-04-11T04:05:04.43715Z"}},"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\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    # 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)[0]\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","metadata":{"trusted":true},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class UNetDataGenerator(tf.keras.utils.Sequence):\n    def __init__(self, patient_ids, meta_df, dicom_tags_df, segmentation_map, organ_status_map,\n                 batch_size=8, img_size=(224, 224), shuffle=True, debug=False):\n        self.patient_ids = patient_ids\n        self.meta_df = meta_df\n        self.dicom_tags_df = dicom_tags_df\n        self.segmentation_map = segmentation_map\n        self.organ_status_map = organ_status_map\n        self.batch_size = batch_size\n        self.img_size = img_size\n        self.shuffle = shuffle\n        self.debug = debug\n        self.indexes = np.arange(len(self.patient_ids))\n        if self.shuffle:\n            np.random.shuffle(self.indexes)\n        self.patient_cache = {}\n\n    def __len__(self):\n        return int(np.ceil(len(self.patient_ids) / self.batch_size))\n    \n    def on_epoch_end(self):\n        self.indexes = np.arange(len(self.patient_ids))\n        if self.shuffle:\n            np.random.shuffle(self.indexes)\n            \n    def __getitem__(self, index):\n        indexes = self.indexes[index * self.batch_size:(index + 1) * self.batch_size]\n        batch_patient_ids = [self.patient_ids[i] for i in indexes]\n        return self._generate_unet_batch(batch_patient_ids)\n    \n    def _load_patient_data(self, patient_id):\n        \"\"\"This function is largely the same as the old function.\"\"\"\n        if patient_id in self.patient_cache:\n            return self.patient_cache[patient_id]\n        \n        try:\n            subset = self.meta_df[self.meta_df['patient_id'] == patient_id]\n            series_id = str(subset['series_id'].iloc[0])\n        except Exception as e:\n            print(f\"No series information found for patient {patient_id}. Error: {e}\")\n            return None\n\n        segmentation_file = self.segmentation_map.get(series_id)\n        \n        has_segmentation = False\n        segmentation_data = None\n        if segmentation_file:\n            try:\n                segmentation_data = nib.load(segmentation_file).get_fdata().astype(np.float32)\n                has_segmentation = True\n                resized_segmentation_data = resize(\n                    segmentation_data,\n                    self.img_size,\n                    order=0,\n                    preserve_range=True\n                ).astype(np.float32)\n            except Exception as e:\n                print(f\"Error loading the Segmentation for {series_id}: {e}\")\n                has_segmentation = False\n        \n        if self.dicom_tags_df is not None:\n            image_ids = self.dicom_tags_df[\n                (self.dicom_tags_df['PatientID'] == patient_id) &\n                (self.dicom_tags_df['series_id'] == series_id)\n            ]['InstanceNumber'].values\n            image_ids = sorted(image_ids)\n        else:\n            print(f\"No image_id information found for patient {patient_id} and series {series_id}.\")\n            return None\n\n        if len(image_ids) == 0:\n            print(f\"No image IDs found for patient {patient_id} and series {series_id}!\")\n            return None\n\n        patient_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, self.img_size)\n                max_val = np.max(image)\n                image = image / max_val if max_val != 0 else image.astype(np.float32)\n                organ_statuses = self.organ_status_map.get(patient_id)\n                \n                mask = np.zeros(self.img_size, dtype=np.float32)\n                if has_segmentation:\n                    if image_id <= segmentation_data.shape[2]:\n                        slice_mask = resize(\n                            segmentation_data[:, :, image_id-1],\n                            self.img_size,\n                            order=0,\n                            preserve_range=True\n                        ).astype(np.float32)\n                        mask = slice_mask\n                slice_data = {\n                    'patient_id': patient_id,\n                    'series_id': series_id,\n                    'image_id': image_id,\n                    'image': image,\n                    'mask': mask,\n                    'organ_statuses': organ_statuses,\n                    'instance_number': dicom.InstanceNumber\n                }\n                patient_slices.append(slice_data)\n                del dicom\n                del image\n            except Exception as e:\n                print(f\"Error reading DICOM file {image_file}: {e}\")\n                continue\n        print(f\"Patient {patient_id} loaded with {len(patient_slices)} slices\")\n        result = {\n            'patient_id': patient_id,\n            'series_id': series_id,\n            'slices': patient_slices,\n            'has_segmentation': has_segmentation,\n            'segmentation_data': segmentation_data if has_segmentation else None\n        }\n        self.patient_cache[patient_id] = result\n        return result\n\n    def _generate_unet_batch(self, batch_patient_ids):\n        batch_images = []\n        batch_masks = []\n        \n        for patient_id in batch_patient_ids:\n            print(f\"Generating UNet data for patient {patient_id}\")\n            patient_data = self._load_patient_data(patient_id)\n            if patient_data is None:\n                print(f\"Skipping patient {patient_id} due to missing data\")\n                continue\n                    \n            slices = patient_data['slices']\n            if len(slices) == 0:\n                print(f\"Skipping patient {patient_id} due to missing slices\")\n                continue\n            \n            slice_data = random.choice(slices)\n            image = slice_data['image']\n            mask = slice_data['mask']\n            if len(image.shape) == 2:\n                image = np.stack([image, image, image], axis=-1)\n            if len(mask.shape) == 2:\n                mask = np.expand_dims(mask, axis=-1)\n                    \n            image, mask = augment_data(image, mask)\n\n            batch_images.append(image)\n            batch_masks.append(mask)\n            del image, mask, slice_data, slices, patient_data\n        \n        if not batch_images:\n            return np.empty((0, *self.img_size, 3)), np.empty((0, *self.img_size, 1))\n                \n        batch_images = np.array(batch_images)\n        batch_masks = np.array(batch_masks)\n        print(f\"[U-Net] Generated batch: images shape={batch_images.shape}, masks shape={batch_masks.shape}\")\n        return batch_images, batch_masks\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T04:05:04.442351Z","iopub.execute_input":"2025-04-11T04:05:04.44292Z","iopub.status.idle":"2025-04-11T04:05:04.466086Z","shell.execute_reply.started":"2025-04-11T04:05:04.442874Z","shell.execute_reply":"2025-04-11T04:05:04.463785Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"class DataGenerator(tf.keras.utils.Sequence):\n    \"\"\"数据生成器，用于分批次加载和处理CT图像数据，使用U-Net分割结果\"\"\"\n\n    def __init__(self, patient_ids, meta_df, dicom_tags_df, segmentation_map, organ_status_map,\n                 unet_model, batch_size=8, img_size=(224, 224), shuffle=True, debug=False):\n        \"\"\"\n        初始化数据生成器\n\n        参数:\n            patient_ids: 患者ID列表\n            meta_df: 包含series_id信息的元数据DataFrame\n            dicom_tags_df: 包含DICOM标签的DataFrame\n            segmentation_map: 分割文件映射字典\n            organ_status_map: 器官状态映射字典\n            unet_model: 训练好的U-Net模型 (用于分割)\n            batch_size: 批次大小\n            img_size: 图像大小\n            shuffle: 是否打乱数据\n            mode: 'classifier'用于2.5D分类模型\n            debug: 是否为调试模式（此处为 False 时使用全部数据）\n        \"\"\"\n        self.patient_ids = patient_ids\n        self.meta_df = meta_df\n        self.dicom_tags_df = dicom_tags_df\n        self.segmentation_map = segmentation_map\n        self.organ_status_map = organ_status_map\n        self.unet_model = unet_model\n        self.batch_size = batch_size\n        self.img_size = img_size\n        self.shuffle = shuffle\n        self.mode = mode\n        self.debug = debug\n        self.organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n        self.indexes = np.arange(len(self.patient_ids))\n        if self.shuffle:\n            np.random.shuffle(self.indexes)\n        self.patient_cache = {}\n\n    def __len__(self):\n        return int(np.ceil(len(self.patient_ids) / self.batch_size))\n\n    def __getitem__(self, index):\n        print(f\"正在生成第 {index+1} 个批次\")\n        indexes = self.indexes[index * self.batch_size:(index + 1) * self.batch_size]\n        batch_patient_ids = [self.patient_ids[i] for i in indexes]\n        print(f\"本批次患者ID: {batch_patient_ids}\")\n\n        return self._generate_classifier_batch(batch_patient_ids)\n\n    def on_epoch_end(self):\n        self.indexes = np.arange(len(self.patient_ids))\n        if self.shuffle:\n            np.random.shuffle(self.indexes)\n\n    def _load_patient_data(self, patient_id):\n        if patient_id in self.patient_cache:\n            return self.patient_cache[patient_id]\n\n        try:\n            subset = self.meta_df[self.meta_df['patient_id'] == patient_id]\n            series_id = str(subset['series_id'].iloc[0])\n            print(f\"患者 {patient_id} 对应的 series_id: {series_id}\")\n        except Exception as e:\n            print(f\"No series information found for patient {patient_id}. Error: {e}\")\n            return None\n\n        segmentation_file = self.segmentation_map.get(series_id)\n\n        has_segmentation = False\n        segmentation_data = None\n        if segmentation_file:\n            try:\n                segmentation_data = nib.load(segmentation_file).get_fdata().astype(np.float32)\n                has_segmentation = True\n                resized_segmentation_data = resize(\n                    segmentation_data,\n                    self.img_size,\n                    order=0,\n                    preserve_range=True\n                ).astype(np.float32)\n                print(f\"患者 {patient_id} 成功加载分割数据。\")\n            except Exception as e:\n                print(f\"Error loading the Segmentation for {series_id}: {e}\")\n                has_segmentation = False\n\n        if self.dicom_tags_df is not None:\n            image_ids = self.dicom_tags_df[(self.dicom_tags_df['PatientID'] == patient_id) & \n                                      (self.dicom_tags_df['series_id'] == series_id)]['InstanceNumber'].values\n            image_ids = sorted(image_ids)\n        else:\n            print(f\"No image_id information found for patient {patient_id} and series {series_id}.\")\n            return None\n\n        if len(image_ids) == 0:\n            print(f\"No image IDs found for patient {patient_id} and series {series_id}!\")\n            return None\n\n        print(f\"患者 {patient_id} 序列 {series_id} 共找到 {len(image_ids)} 张切片\")\n\n        patient_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, self.img_size)\n                max_val = np.max(image)\n                image = image / max_val if max_val != 0 else image.astype(np.float32)\n                organ_statuses = self.organ_status_map.get(patient_id)\n                \n                mask = np.zeros(self.img_size, dtype=np.float32)\n                if has_segmentation:\n                    if image_id <= segmentation_data.shape[2]:\n                        slice_mask = resize(\n                            segmentation_data[:, :, image_id-1],\n                            self.img_size,\n                            order=0,\n                            preserve_range=True\n                        ).astype(np.float32)\n                        mask = slice_mask\n                \n                slice_data = {\n                    'patient_id': patient_id,\n                    'series_id': series_id,\n                    'image_id': image_id,\n                    'image': image,\n                    'mask': mask,\n                    'organ_statuses': organ_statuses,\n                    'instance_number': dicom.InstanceNumber\n                }\n                patient_slices.append(slice_data)\n                del dicom\n                del image\n            except Exception as e:\n                print(f\"Error reading DICOM file {image_file}: {e}\")\n                continue\n        print(f\"患者 {patient_id} 加载完成，切片数量: {len(patient_slices)} \")\n        result = {\n            'patient_id': patient_id,\n            'series_id': series_id,\n            'slices': patient_slices,\n            'has_segmentation': has_segmentation,\n            'segmentation_data': segmentation_data if has_segmentation else None\n        }\n        self.patient_cache[patient_id] = result\n        del segmentation_data,resized_segmentation_data\n        gc.collect()\n        return result\n\n    def _generate_classifier_batch(self, batch_patient_ids):\n        batch_sequences = []\n        batch_labels = []\n\n        for patient_id in batch_patient_ids:\n            print(f\"Processing patient {patient_id} for classifier batch\")\n            patient_data = self._load_patient_data(patient_id)\n            if patient_data is None:\n                print(f\"Skipping patient {patient_id} due to missing data\")\n                continue\n\n            slices = patient_data['slices']\n            if len(slices) < 32:\n                print(f\"Skipping patient {patient_id} because they have less than 32 slices\")\n                continue\n                # 确保列表长度大于等于32\n            if len(slices) < 32:\n                 print(f\"患者 {patient_id} 的切片数量不足 32 张，跳过\")\n                 continue\n\n            # 选择中间的32张连续切片\n            middle_index = len(slices) // 2\n            start_index = max(0, middle_index - 16)\n            sequence = slices[start_index:start_index + 32]\n\n            if len(sequence) < 32:\n                print(f\"提取连续32张切片失败，跳过患者{patient_id}\")\n                continue\n\n            processed_sequence = []\n            print(\"正在处理batch\")\n            for slice_data in sequence:\n                image = slice_data['image']\n                #将分割后的图像放进去, 并且使用unet分割后的图像，\n                #TODO 添加分割的代码，将unet_model作为参数传递进来后，实现分割功能。\n                if len(image.shape) == 2:\n                    image = np.stack([image, image, image], axis=-1)\n                \n                segmented_slice = segment_slice(image, self.unet_model, self.img_size) #call the segmentation process\n                processed_sequence.append(segmented_slice) #append the new segmented slice.\n\n                del image, segmented_slice\n                gc.collect()\n\n            print(\"处理完毕\")\n            organ_status = self.organ_status_map.get(patient_id)\n            if not organ_status:\n                print(f\"不存在{patient_id}的数据集\")\n                continue\n\n            labels = []\n            for organ in self.organs:\n                status = organ_status.get(organ, {})\n                organ_label = [status.get('healthy', 0),\n                               status.get('injury', 0),\n                               status.get('low', 0),\n                               status.get('high', 0)]\n                labels.append(organ_label)\n\n            processed_sequence_arr = np.array(processed_sequence, dtype=np.float32) #make sure this is the proper format before adding to sequence. \n            batch_sequences.append(processed_sequence_arr)\n            batch_labels.append(np.array(labels, dtype=np.float32))\n\n            del processed_sequence, labels, slice_data, slices, patient_data, sequence,processed_sequence_arr\n            gc.collect()\n\n        if not batch_sequences:\n            print(\"当前批次为空，返回空数组\")\n            return np.empty((0, 32, *self.img_size, 3),dtype=np.float32), np.empty((0, len(self.organs), 4),dtype=np.float32)\n        \n        batch_sequences = np.array(batch_sequences,dtype=np.float32)\n        batch_labels = np.array(batch_labels,dtype=np.float32)\n        print(f\"2.5D分类 成功生成批次，序列shape：{batch_sequences.shape}，标签shape：{batch_labels.shape}\")\n        gc.collect()\n        return batch_sequences, batch_labels\n","metadata":{"trusted":true},"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-11T04:05:04.467759Z","iopub.execute_input":"2025-04-11T04:05:04.468645Z","iopub.status.idle":"2025-04-11T04:05:04.502467Z","shell.execute_reply.started":"2025-04-11T04:05:04.46858Z","shell.execute_reply":"2025-04-11T04:05:04.499733Z"}},"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-11T04:05:04.504501Z","iopub.execute_input":"2025-04-11T04:05:04.504924Z","iopub.status.idle":"2025-04-11T04:05:04.525171Z","shell.execute_reply.started":"2025-04-11T04:05:04.504891Z","shell.execute_reply":"2025-04-11T04:05:04.52348Z"}},"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    # 输入层 - 32张连续切片\n    input_ct = layers.Input(shape=(32,) + input_shape, name='ct_input')\n    \n    # 使用TimeDistributed包装EfficientNetB1，处理每张切片\n    # 加载预训练的EfficientNetB1，但移除顶层\n    base_model = EfficientNetB1(\n        include_top=False, \n        weights='imagenet', \n        input_shape=input_shape,\n        pooling='avg'\n    )\n    base_model.trainable = False  # 冻结基础模型权重\n    \n    # 使用TimeDistributed应用到每张切片\n    x = layers.TimeDistributed(base_model)(input_ct)\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-11T04:05:04.526675Z","iopub.execute_input":"2025-04-11T04:05:04.527546Z","iopub.status.idle":"2025-04-11T04:05:04.552848Z","shell.execute_reply.started":"2025-04-11T04:05:04.527395Z","shell.execute_reply":"2025-04-11T04:05:04.551213Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 损失函数&训练","metadata":{}},{"cell_type":"code","source":"def weighted_cross_entropy(organ_weights):\n    \"\"\"\n    Returns a weighted cross-entropy loss function.\n    \"\"\"\n    def loss(y_true, y_pred):\n        \"\"\"\n        Calculates the weighted cross-entropy loss.\n        \"\"\"\n        loss = 0.0\n        for i in range(NUM_ORGANS):\n            # 提取每个器官的真实标签和预测值\n            y_true_organ = y_true[:, i, :]\n            y_pred_organ = y_pred[:, i, :]\n\n            # 计算交叉熵损失\n            cross_entropy = tf.keras.losses.binary_crossentropy(y_true_organ, y_pred_organ)\n\n            # 应用权重\n            loss += organ_weights[i] * tf.reduce_mean(cross_entropy)\n\n        return loss\n    return loss\n\ndef calculate_metrics(y_true, y_pred, threshold=0.5):\n    \"\"\"Calculates evaluation metrics.\"\"\"\n    y_pred_binary = (y_pred > threshold).astype(int)  # 将预测概率转换为二进制标签\n    accuracy = accuracy_score(y_true, y_pred_binary)\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    if len(np.unique(y_true)) > 1:  # 避免二分类混淆矩阵只有一类时出错\n        tn, fp, fn, tp = confusion_matrix(y_true, y_pred_binary).ravel()\n        specificity = tn / (tn + fp) if (tn + fp) > 0 else 0\n    else:\n        specificity = 0  # 如果只有一类，则特异度无法计算\n    ppv = precision\n    npv = 0 # 默认 npv 为 0 ， 在计算出 cm 后更新数值\n    if len(np.unique(y_true)) > 1:\n      npv = tn / (tn + fn) if (tn + fn) > 0 else 0 # 根据 tn 计算 npv\n    return accuracy, precision, recall, specificity, ppv, npv, confusion_matrix(y_true, y_pred_binary)\n","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T04:05:04.553848Z","iopub.execute_input":"2025-04-11T04:05:04.554191Z","iopub.status.idle":"2025-04-11T04:05:04.575248Z","shell.execute_reply.started":"2025-04-11T04:05:04.554167Z","shell.execute_reply":"2025-04-11T04:05:04.573819Z"}},"outputs":[],"execution_count":null},{"cell_type":"code","source":"def main():\n    # 1. 加载和预处理数据、器官状态和分割映射\n    train_df, organ_status_map, segmentation_map = load_and_process_data()\n    \n    # 2. 划分患者ID为训练集和验证集\n    from sklearn.model_selection import train_test_split\n    patient_ids = train_df['patient_id'].unique()\n    train_patient_ids, val_patient_ids = train_test_split(patient_ids, test_size=0.2, random_state=42)\n    \n    # 3. 加载元数据和DICOM标签\n    meta_df = pd.read_csv(TRAIN_META)\n    meta_df['patient_id'] = meta_df['patient_id'].astype(str)\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    \n\n    # 4. 构建U-Net模型\n    input_shape = (IMG_SIZE[0], IMG_SIZE[1], 3)\n    unet_model = build_unet(input_shape)\n    unet_model.compile(\n        optimizer=optimizers.Adam(learning_rate=1e-4),\n        loss='binary_crossentropy',\n        metrics=['accuracy']\n    )\n    unet_model.summary()\n\n    #create the data Gen and now train the U-net on It.\n    unet_train_generator = UNetDataGenerator(\n        train_patient_ids, meta_df, dicom_tags_df, segmentation_map, organ_status_map,\n        batch_size=2, img_size=IMG_SIZE, shuffle=True, # The UNet dataloader doesn't require UNET,\n    )\n    unet_val_generator = UNetDataGenerator(\n        val_patient_ids, meta_df, dicom_tags_df, segmentation_map, organ_status_map,\n        batch_size=2, img_size=IMG_SIZE, shuffle=False,\n    )\n    gc.collect() # Explicitly run garbage collection\n\n    history_unet = unet_model.fit(\n        unet_train_generator,\n        validation_data=unet_val_generator,\n        epochs=5, # Reduce for testing\n    )\n    unet_model.save('unet_model.h5') #Save the UNet after training the UNet\n    #deallocate memory, and remove them so that it reloads and reduces future memory use.\n    del train_patient_ids, val_patient_ids,unet_train_generator, unet_val_generator\n    gc.collect() # Run garbage collection\n\n    # 5. 重新加载需要的数据 (需要重新加载，因为之前删除了)\n    train_df, organ_status_map, segmentation_map = load_and_process_data()\n    meta_df = pd.read_csv(TRAIN_META)\n    meta_df['patient_id'] = meta_df['patient_id'].astype(str)\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    \n    patient_ids = train_df['patient_id'].unique()\n    train_patient_ids, val_patient_ids = train_test_split(patient_ids, test_size=0.2, random_state=42)\n\n\n    # 6: Building 2.5D Classifier\n    #Load Generator\n    train_generator_classifier = DataGenerator(\n        train_patient_ids, meta_df, dicom_tags_df, segmentation_map, organ_status_map, unet_model, #Pass unet to data generator to load the unet model\n        batch_size=2, img_size=IMG_SIZE, mode='classifier', debug=False\n    )\n    val_generator_classifier = DataGenerator(\n        val_patient_ids, meta_df, dicom_tags_df, segmentation_map, organ_status_map, unet_model,#Pass the unet model to the dataloader\n        batch_size=2, img_size=IMG_SIZE, mode='classifier', debug=False, shuffle=False\n    )\n\n    gc.collect()\n    del train_patient_ids, val_patient_ids, meta_df, dicom_tags_df, temp, train_df, organ_status_map, segmentation_map, patient_ids\n    gc.collect() # Explicitly run garbage collection\n    # 7. 构建 2.5D 分类模型\n    input_shape = (IMG_SIZE[0], IMG_SIZE[1], 3)\n    classifier_model = build_25d_classifier(input_shape=input_shape, num_organs=NUM_ORGANS)\n    organ_weights = [1.0, 1.0, 1.0, 1.0, 1.0]\n    loss_function = weighted_cross_entropy(organ_weights)\n    optimizer_inst = optimizers.Adam(learning_rate=1e-4)\n    classifier_model.compile(optimizer=optimizer_inst, loss=loss_function, metrics=['accuracy'])\n    classifier_model.summary()\n\n    # 8. 使用分割结果训练2.5D分类模型\n    history_classifier = classifier_model.fit(\n        train_generator_classifier,\n        validation_data=val_generator_classifier,\n        epochs=5,\n    )\n    classifier_model.save('classifier_model.h5')\n\n    del train_generator_classifier, val_generator_classifier #deallocate training generators\n    gc.collect() #Run garbage collection\n\n    #deallocate memory for UNet Model\n    del unet_model\n    gc.collect()\n\n    # 9. 评估 2.5D 分类模型\n    print(\"现在开始评估分类器\")\n    #ReLoading data sets again\n    train_df, organ_status_map, segmentation_map = load_and_process_data()\n    meta_df = pd.read_csv(TRAIN_META)\n    meta_df['patient_id'] = meta_df['patient_id'].astype(str)\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    patient_ids = train_df['patient_id'].unique()\n    \n    #Load the validation data Generator\n    val_generator_classifier = DataGenerator(\n        patient_ids, meta_df, dicom_tags_df, segmentation_map, organ_status_map,\n        unet_model, batch_size=2, img_size=IMG_SIZE, mode='classifier', debug=False, shuffle=False #Do not shuffle the eval set\n    )\n    \n        #deallocate memory we've used.\n    del train_df, meta_df, dicom_tags_df, temp #deallocate training generators\n    gc.collect() # Explicitly run garbage collection\n\n    val_ground_truth = []\n    val_predictions = []\n    for i in range(len(val_generator_classifier)):\n        X_val, y_val = val_generator_classifier[i]\n\n        #验证是否有batch\n        if X_val.shape[0] > 0:\n            y_pred = classifier_model.predict(X_val)\n            val_ground_truth.extend(y_val)\n            val_predictions.extend(y_pred)\n    \n    del X_val, y_val, val_generator_classifier #Explicit deallocation\n    gc.collect() #Garbage Collection\n    \n    val_ground_truth = np.array(val_ground_truth)\n    val_predictions = np.array(val_predictions)\n    gc.collect()#deallocate inside eval loop\n    \n    if len(val_predictions) > 0:\n        print(\"完成验证集评估\")\n        organs = ['bowel', 'extravasation', 'kidney', 'liver', 'spleen']\n    \n        for i, organ in enumerate(organs):\n            print(f\"器官: {organ}\")\n            y_true_organ = val_ground_truth[:, i, 1]\n            y_pred_organ = val_predictions[:, i, 1]\n            accuracy, precision, recall, specificity, ppv, npv, cm = calculate_metrics(y_true_organ, y_pred_organ)\n            print(f\"  准确率: {accuracy:.4f}\")\n            print(f\"  精确率: {precision:.4f}\")\n            print(f\"  召回率: {recall:.4f}\")\n            print(f\"  特异性: {specificity:.4f}\")\n            print(f\"  PPV: {ppv:.4f}\")\n            print(f\"  NPV: {npv:.4f}\")\n            print(f\"  混淆矩阵: \\n{cm}\")\n    \n    del val_predictions, val_ground_truth ,classifier_model,  loss_function, optimizer_inst#Final deallocation\n    gc.collect() #Garbage collection\ngc.collect()#garbage collection at the very beginning.\nif __name__ == \"__main__\":\n    gc.collect()\n    main()\n    gc.collect()#garbage collection at the very end.","metadata":{"trusted":true,"execution":{"iopub.status.busy":"2025-04-11T04:05:04.577155Z","iopub.execute_input":"2025-04-11T04:05:04.577596Z","execution_failed":"2025-04-11T04:07:20.352Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 可视化模块","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true,"execution":{"execution_failed":"2025-04-11T04:07:20.352Z"}},"outputs":[],"execution_count":null},{"cell_type":"markdown","source":"# 生成预测文件","metadata":{}},{"cell_type":"code","source":"","metadata":{"trusted":true},"outputs":[],"execution_count":null}]}