{"metadata":{"kernelspec":{"language":"python","display_name":"Python 3","name":"python3"},"language_info":{"name":"python","version":"3.7.10","mimetype":"text/x-python","codemirror_mode":{"name":"ipython","version":3},"pygments_lexer":"ipython3","nbconvert_exporter":"python","file_extension":".py"},"kaggle":{"accelerator":"gpu","dataSources":[{"sourceId":30201,"databundleVersionId":2750748,"sourceType":"competition"},{"sourceId":2709672,"sourceType":"datasetVersion","datasetId":1650863},{"sourceId":2763328,"sourceType":"datasetVersion","datasetId":1686152},{"sourceId":2763687,"sourceType":"datasetVersion","datasetId":1654455},{"sourceId":2763706,"sourceType":"datasetVersion","datasetId":1686418},{"sourceId":2763719,"sourceType":"datasetVersion","datasetId":1686427}],"dockerImageVersionId":30140,"isInternetEnabled":false,"language":"python","sourceType":"notebook","isGpuEnabled":true}},"nbformat_minor":4,"nbformat":4,"cells":[{"cell_type":"markdown","source":"<br>\n\n<center><img src=\"https://rs1.chemie.de/images//128537-76.jpg\" width=60%></center>\n\n<h2 style=\"text-align: center; font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: underline; text-transform: none; letter-spacing: 2px; color: blue; background-color: #ffffff;\">Cell Instance Segmentation Challenge</h2>\n<h5 style=\"text-align: center; font-family: Verdana; font-size: 12px; font-style: normal; font-weight: bold; text-decoration: None; text-transform: none; letter-spacing: 1px; color: black; background-color: #ffffff;\">CREATED BY: DARIEN SCHETTLER</h5>\n\n<br>\n\n---\n\n<br>\n\n<center><div class=\"alert alert-block alert-danger\" style=\"margin: 2em; line-height: 1.7em; font-family: Verdana;\">\n    <b style=\"font-size: 18px;\">🛑 &nbsp; WARNING:</b><br><br><b>THIS IS A WORK IN PROGRESS</b><br>\n</div></center>\n\n\n<center><div class=\"alert alert-block alert-warning\" style=\"margin: 2em; line-height: 1.7em; font-family: Verdana;\">\n    <b style=\"font-size: 18px;\">👏 &nbsp; IF YOU FORK THIS OR FIND THIS HELPFUL &nbsp; 👏</b><br><br><b style=\"font-size: 22px; color: darkorange\">PLEASE UPVOTE!</b><br><br>This was a lot of work for me and while it may seem silly, it makes me feel appreciated when others like my work. 😅\n</div></center>\n\n\n","metadata":{}},{"cell_type":"markdown","source":"<p id=\"toc\"></p>\n\n<br><br>\n\n<h1 style=\"font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; color: blue; background-color: #ffffff;\">TABLE OF CONTENTS</h1>\n\n---\n\n<h3 style=\"text-indent: 10vw; font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\"><a href=\"#imports\">0&nbsp;&nbsp;&nbsp;&nbsp;IMPORTS</a></h3>\n\n---\n\n<h3 style=\"text-indent: 10vw; font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\"><a href=\"#background_information\">1&nbsp;&nbsp;&nbsp;&nbsp;BACKGROUND INFORMATION</a></h3>\n\n---\n\n<h3 style=\"text-indent: 10vw; font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\"><a href=\"#setup\">2&nbsp;&nbsp;&nbsp;&nbsp;SETUP</a></h3>\n\n---\n\n<h3 style=\"text-indent: 10vw; font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\"><a href=\"#helper_functions\">3&nbsp;&nbsp;&nbsp;&nbsp;HELPER FUNCTIONS</a></h3>\n\n---\n\n<h3 style=\"text-indent: 10vw; font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\"><a href=\"#create_dataset\">4&nbsp;&nbsp;&nbsp;&nbsp;DATASET CREATION AND EXPLORATION</a></h3>\n\n---\n\n<h3 style=\"text-indent: 10vw; font-family: Verdana; font-size: 20px; font-style: normal; font-weight: normal; text-decoration: none; text-transform: none; letter-spacing: 2px; color: navy; background-color: #ffffff;\"><a href=\"#modelling\">5&nbsp;&nbsp;&nbsp;&nbsp;MODELLING</a></h3>\n\n---","metadata":{}},{"cell_type":"markdown","source":"<br>\n\n<a id=\"imports\"></a>\n\n<h1 style=\"font-family: Verdana; font-size: 24px; font-style: normal; font-weight: bold; text-decoration: none; text-transform: none; letter-spacing: 3px; background-color: #ffffff; color: blue;\" id=\"imports\">0&nbsp;&nbsp;IMPORTS&nbsp;&nbsp;&nbsp;&nbsp;<a href=\"#toc\">&#10514;</a></h1>","metadata":{}},{"cell_type":"code","source":"# 打印开始导入模块的提示信息\nprint(\"\\n... IMPORTS STARTING ...\\n\")\n\n# 打印开始安装和下载依赖的提示信息\nprint(\"\\n... PIP/APT INSTALLS AND DOWNLOADS/ZIP STARTING ...\")\n# 尝试跳过并禁用网络以便无网络提交\n!pip install -q ../input/tensorflow-model-optimization/numpy-1.21.3-cp37-cp37m-manylinux_2_12_x86_64.manylinux2010_x86_64.whl\n!pip install -q ../input/tensorflow-model-optimization/dm_tree-0.1.6-cp37-cp37m-manylinux_2_24_x86_64.whl\n!pip install -q ../input/tensorflow-model-optimization/six-1.16.0-py2.py3-none-any.whl\n!pip install -q ../input/tensorflow-model-optimization/tensorflow_model_optimization-0.7.0-py2.py3-none-any.whl\n!pip install ../input/neural-structued-learning/neural_structured_learning-1.3.1-py2.py3-none-any.whl\n# !pip install -q --upgrade tensorflow_datasets\n# !pip install -q neural-structured-learning\nprint(\"... PIP/APT INSTALLS COMPLETE ...\\n\")\n\n# 打印版本信息\nprint(\"\\n\\tVERSION INFORMATION\")\n# 机器学习和数据科学模块导入\nimport tensorflow as tf; print(f\"\\t\\t– TENSORFLOW VERSION: {tf.__version__}\");\nimport tensorflow_addons as tfa; print(f\"\\t\\t– TENSORFLOW ADDONS VERSION: {tfa.__version__}\");\nimport pandas as pd; pd.options.mode.chained_assignment = None;\nimport numpy as np; print(f\"\\t\\t– NUMPY VERSION: {np.__version__}\");\nimport sklearn; print(f\"\\t\\t– SKLEARN VERSION: {sklearn.__version__}\");\nfrom sklearn.preprocessing import RobustScaler, PolynomialFeatures\nfrom pandarallel import pandarallel; pandarallel.initialize();\nfrom sklearn.model_selection import GroupKFold;\n\n# 内置模块导入\nfrom kaggle_datasets import KaggleDatasets\nfrom collections import Counter\nfrom datetime import datetime\nfrom glob import glob\nimport warnings\nimport requests\nimport hashlib\nimport imageio\nimport IPython\nimport sklearn\nimport urllib\nimport zipfile\nimport pickle\nimport random\nimport shutil\nimport string\nimport json\nimport math\nimport time\nimport gzip\nimport ast\nimport sys\nimport io\nimport os\nimport gc\nimport re\n\n# 可视化模块导入\nfrom matplotlib.colors import ListedColormap\nimport matplotlib.patches as patches\nimport plotly.graph_objects as go\nimport matplotlib.pyplot as plt\nfrom tqdm.notebook import tqdm; tqdm.pandas();\nimport plotly.express as px\nimport seaborn as sns\nfrom PIL import Image, ImageEnhance\nimport matplotlib; print(f\"\\t\\t– MATPLOTLIB VERSION: {matplotlib.__version__}\");\nimport plotly\nimport PIL\nimport cv2\n\n# 设置随机种子以实现可重复性\ndef seed_it_all(seed=7):\n    \"\"\"尝试使结果可重现的种子函数\"\"\"\n    os.environ['PYTHONHASHSEED'] = str(seed)\n    random.seed(seed)\n    np.random.seed(seed)\n    tf.random.set_seed(seed)\n\n# 打印导入模块完成的提示信息\nprint(\"\\n\\n... IMPORTS COMPLETE ...\\n\")\n\n# 打印开始EfficientDET设置的提示信息\nprint(\"\\n... EFFICIENTDET SETUP STARTING ...\")\n\n# 设置库目录\nLIB_DIR = \"/kaggle/input/google-automl-efficientdetefficientnet-oct-2021\"\n\n# 为了让automl文件可访问，将路径添加到sys.path中\nsys.path.insert(0, LIB_DIR)\nsys.path.insert(0, os.path.join(LIB_DIR, \"automl-master\"))\nsys.path.insert(0, os.path.join(LIB_DIR, \"automl-master\", \"efficientdet\"))\nsys.path.insert(0, os.path.join(LIB_DIR, \"automl-master\", \"efficientdet\", \"tf2\"))\n\n# EfficientDET模块导入,EfficientDET是用于目标检测的深度学习模型\nimport hparams_config\nfrom tf2 import efficientdet_keras\nfrom tf2 import train_lib\nfrom tf2 import anchors\nfrom tf2 import efficientdet_keras\nfrom tf2 import label_util\nfrom tf2 import postprocess\nfrom tf2 import util_keras\nfrom tf2.train import setup_model\nfrom efficientdet import dataloader\nfrom visualize import vis_utils\nfrom inference import visualize_image\nprint(\"... EFFICIENTDET SETUP COMPLETE ...\\n\")\n\n# 打印设置随机种子的提示信息\nprint(\"\\n... SEEDING FOR DETERMINISTIC BEHAVIOUR ...\\n\")\nseed_it_all()","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:14:58.743376Z","iopub.execute_input":"2023-12-21T03:14:58.743734Z","iopub.status.idle":"2023-12-21T03:17:45.754061Z","shell.execute_reply.started":"2023-12-21T03:14:58.743642Z","shell.execute_reply":"2023-12-21T03:17:45.753074Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 1 基本竞赛信息","metadata":{}},{"cell_type":"markdown","source":"## 1.1 主要任务描述","metadata":{}},{"cell_type":"markdown","source":"神经系统疾病，包括阿尔茨海默病和脑肿瘤等神经退行性疾病，是全球主要的死亡和致残原因。然而，`很难量化这些致命疾病对治疗的反应`。一种公认的方法是通过`光学显微镜`检查神经元细胞，这既具有可访问性又是无创的。然而，`在显微图像中分割单个神经元细胞`可能是具有挑战性且耗时的。`借助计算机视觉进行准确的细胞实例分割，有望带来新的、有效的药物发现，以治疗数百万患有这些疾病的人`。\n\n`当前的解决方案对神经元细胞的准确性有限`。在开发细胞实例分割模型的内部研究中，`神经母细胞瘤细胞系SH-SY5Y在八种不同的癌细胞类型中一直表现出最低的精度分数`。这可能是因为`神经元细胞具有非常独特、不规则且凹凸不平的形态`，使它们难以使用常用的掩膜头进行分割。\n\nSartorius是生命科学研究和生物制药行业的合作伙伴。他们赋予科学家和工程师简化和加速生命科学和生物加工进展的能力，促使新的、更好的疗法和更经济的药物的开发。他们是领域内的先锋和领先专家的磁铁和动态平台。他们将创造性的思维汇聚在一起，为共同的目标而努力：技术突破，为更多人带来更好的健康。\n\n`在这个竞赛中，您将检测和描绘在描绘神经系统疾病研究中常用的神经元细胞类型的生物图像中的感兴趣的不同对象。更具体地说，您将使用相位对比显微镜图像来训练和测试模型，实现神经元细胞的实例分割。成功的模型将以高准确度完成此任务`。\n\n如果成功，您将通过收集稳健的定量数据来帮助推动神经生物学研究。研究人员可能能够更容易地测量疾病和治疗条件对神经元细胞的影响。因此，可能会发现用于治疗数百万患有这些主要死亡和致残原因的新药物。 ","metadata":{}},{"cell_type":"markdown","source":"## 1.2 竞赛评估","metadata":{}},{"cell_type":"markdown","source":"总体评估信息\n\n本竞赛的评估依据不同交并比（Intersection over Union，IoU）阈值下的平均精度。提出的对象像素集与真实对象像素集之间的IoU计算如下：\n<br><center><b style=\"font-size: 20px;\">$IoU(A,B) = \\frac{A \\cap B}{ A \\cup B}$</b></center><br>\n \n该指标在一系列IoU阈值范围内进行扫描，每一点计算一个平均精度值。阈值的范围从0.5到0.95，步长为0.05:<br />\n*即0.5, 0.55, 0.6, 0.65, 0.7, 0.75, 0.8, 0.85, 0.9, 0.95。<br />\n*<mark><b>换句话说，在阈值为0.5时，如果预测对象与地面真实对象的IoU大于0.5，则认为是一个“命中”`。</b></mark>\n\n\n在每个阈值值 t 下，基于真正例 TP、假负例 FN 和假正例 FP 的数量计算精度值，这些数量是通过将预测对象与所有地面真实对象进行比较得到的：\n<br><center><b style=\"font-size: 24px;\">$\\frac{TP(t)}{TP(t) + FP(t) + FN(t)}$</b></center><br>\n\n当单个预测对象与地面真实对象的IoU超过阈值时，计为真正例。\n\n假正例表示预测对象没有关联的地面真实对象。假负例表示地面真实对象没有关联的预测对象。\n\n然后，单个图像的平均精度计算为在每个IoU阈值下上述精度值的平均值。\n\n<br><b>IoU threshold:</b>\n\n<br><center><b style=\"font-size: 24px;\">$\\frac{1}{|thresholds|} \\sum_t \\frac{TP(t)}{TP(t) + FP(t) + FN(t)}$</b></center><br>\n\n最后，竞赛指标返回的分数是测试数据集中每个图像的个体平均精度的均值。\n\n**提交文件信息**\n\n为了减小提交文件大小，我们的指标使用了基于运行长度编码（Run-Length Encoding，RLE）的像素值。您将提交一对值的空格分隔列表，其中包含起始位置和运行长度。\n\n例如，'1 3' 表示从像素1开始，运行总共3个像素1,2,3。竞赛格式要求一对值的空格分隔列表。例如，'1 3 10 5' 表示要包含掩码中的像素1,2,3,10,11,12,13,14。像素从1开始索引，按从上到下、从左到右的顺序编号：\n\n1 表示像素1,1<br />\n2 表示像素2,1\n\n指标检查这些对值是否已排序、为正值，并且解码后的像素值没有重复。还检查同一图像的两个预测掩码是否重叠。\n\n文件应包含标题，并采用以下格式。您提交的每一行代表给定ImageId的单个预测细胞核分割。\n\nImageId,EncodedPixels<br />\n0114f484a16c152baa2d82fdd43740880a762c93f436c8988ac461c5c9dbe7d5,1 1\n0999dab07b11bc85fb8464fc36c947fbd8b5d6ec49817361cb780659ca805eac,1 1\n0999dab07b11bc85fb8464fc36c947fbd8b5d6ec49817361cb780659ca805eac,2 3 8 9\n\n等等...","metadata":{}},{"cell_type":"markdown","source":"## 1.3 数据集概述","metadata":{}},{"cell_type":"markdown","source":"**一般信息**\n\n在这个竞赛中，我们要在图像中分割神经元细胞。训练注释以运行长度编码的掩码形式提供，图像以PNG格式呈现。图像数量较少，但注释的对象数量相当大。隐藏的测试集大约包含240张图像。\n\n**文件**\n\n`train.csv`\n\n**所有训练对象的ID和掩码。测试集不提供任何这些元数据。**\n\n`id`: 对象的唯一标识符<br />\n`annotation`: 识别的神经元细胞的运行长度编码像素<br />\n`width`: 源图像宽度<br />\n`height`: 源图像高度<br />\n`cell_type`: 细胞系<br />\n`plate_time`: 创建板的时间<br />\n`sample_date`: 创建样本的日期<br />\n`sample_id`: 样本标识符<br />\n`elapsed_timedelta`: 自第一张样本图像拍摄以来的时间<br />\n\n`sample_submission.csv`\n\n**正确格式的示例提交文件。**\n\n`train`\n\n**以PNG格式提供的训练图像。**\n\n`test`\n\n**以PNG格式提供的测试图像。只有少量测试集图像可供下载；<br />\n其余图像只能在提交时由您的笔记本访问。**\n\n`train_semi_supervised`\n\n**提供未标记的图像，以便您可以使用额外的数据进行半监督方法。**\n\n`LIVECell_dataset_2021`\n\n**LIVECell数据集的镜像。LIVECell是该竞赛的前身数据集。您将找到SH-SHY5Y细胞系的额外数据，以及竞赛数据集中未涵盖的其他细胞系的数据，这可能对迁移学习很有帮助。**","metadata":{}},{"cell_type":"markdown","source":"# 2 设置","metadata":{}},{"cell_type":"markdown","source":"## 2.1 加速器检测","metadata":{}},{"cell_type":"markdown","source":"为了使用TPU，我们使用TPUClusterResolver进行初始化，这对于连接到远程集群并初始化云TPU是必要的。让我们了解两个重要的点：\n\n1.\t在Kaggle上使用TPU时，您不需要为TPUClusterResolver指定参数。\n\n2.\t但是，在Google Compute Engine（GCE）上，您需要执行以下操作：\n\n```python\n# 您给予要使用的TPU的名称\nTPU_WORKER = 'my-tpu-name'\n\n# 或者您还可以直接指定grpc路径\n# TPU_WORKER = 'grpc://xxx.xxx.xxx.xxx:8470'\n\n# 您在创建用于GCP上的TPU时选择的区域。\nZONE = 'us-east1-b'\n\n# 在GCP上创建用于使用TPU的项目的名称。\nPROJECT = '我的TPU项目'\n\ntpu = tf.distribute.cluster_resolver.TPUClusterResolver(tpu=TPU_WORKER, zone=ZONE, project=PROJECT)\n\n```\n\n<div class=\"alert alert-block alert-danger\" style=\"margin: 2em; line-height: 1.7em; font-family: Verdana;\">\n    <b style=\"font-size: 16px;\">🛑 &nbsp; WARNING:</b><br><br>- 尽管Tensorflow文档表示project参数应提供项目名称，但实际上应提供的是项目ID。您可以在GCP项目仪表板页面上找到这个ID。<br>","metadata":{}},{"cell_type":"code","source":"print(f\"\\n... ACCELERATOR SETUP STARTING ...\\n\") #加速器设置开始\n\n# 检测硬件，返回适当的分布策略\ntry:\n    # TPU检测。如果设置了TPU_NAME环境变量，则不需要任何参数。在Kaggle上，这总是成立的。\n    TPU = tf.distribute.cluster_resolver.TPUClusterResolver()  \nexcept ValueError:\n    TPU = None\n\nif TPU:\n    print(f\"\\n... RUNNING ON TPU - {TPU.master()}...\")\n    tf.config.experimental_connect_to_cluster(TPU)\n    tf.tpu.experimental.initialize_tpu_system(TPU)\n    strategy = tf.distribute.experimental.TPUStrategy(TPU)\nelse:\n    print(f\"\\n... RUNNING ON CPU/GPU ...\")\n    # 在Tensorflow中使用默认的分布策略\n    #   --> 适用于CPU和单个GPU。\n    strategy = tf.distribute.get_strategy() \n\n# 什么是副本？\n#    --> 单个Cloud TPU设备由四个芯片组成，每个芯片都有两个TPU核心。\n#    --> 因此，为了有效利用Cloud TPU，程序应该使用每个芯片的所有八个（4x2）核心。\n#    --> 每个副本本质上是在每个核心上运行的训练图的副本，并训练包含总批次大小的1/8的小批次。\nN_REPLICAS = strategy.num_replicas_in_sync\n    \nprint(f\"... # OF REPLICAS: {N_REPLICAS} ...\\n\") #副本数量\n\nprint(f\"\\n... ACCELERATOR SETUP COMPLTED ...\\n\") #加速器设置完成\n","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:18:26.888706Z","iopub.execute_input":"2023-12-21T03:18:26.889585Z","iopub.status.idle":"2023-12-21T03:18:26.906064Z","shell.execute_reply.started":"2023-12-21T03:18:26.889537Z","shell.execute_reply":"2023-12-21T03:18:26.904973Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2.2 竞赛数据访问","metadata":{}},{"cell_type":"markdown","source":"TPU必须直接从Google Cloud Storage (GCS) 中读取数据。Kaggle提供了一个实用库——KaggleDatasets，其中包含一个实用函数`.get_gcs_path`，它允许我们访问位于GCS中的输入数据集的位置。\n\n<div class=\"alert alert-block alert-info\" style=\"margin: 2em; line-height: 1.7em; font-family: Verdana;\">\n    <b style=\"font-size: 16px;\">📌 &nbsp; 技巧:</b><br><br>- 如果您的笔记本上附加了多个数据集，您应该将特定数据集的名称传递给<b><code>`get_gcs_path()`</code></b>函数。在我们的情况下，数据集的名称就是数据集在其中挂载的目录的名称。</i><br><br>","metadata":{}},{"cell_type":"code","source":"print(\"\\n... DATA ACCESS SETUP STARTED ...\\n\") #数据访问设置开始\n\nif TPU:\n    # Google Cloud Dataset路径到训练和验证图像\n    DATA_DIR = KaggleDatasets().get_gcs_path('sartorius-cell-instance-segmentation')\n    save_locally = tf.saved_model.SaveOptions(experimental_io_device='/job:localhost')\nelse:\n    # 本地路径到训练和验证图像\n    DATA_DIR = \"/kaggle/input/sartorius-cell-instance-segmentation\"\n    save_locally = None\n    \nprint(f\"\\n... DATA DIRECTORY PATH IS:\\n\\t--> {DATA_DIR}\") #数据目录路径\n\nprint(f\"\\n... IMMEDIATE CONTENTS OF DATA DIRECTORY IS:\") #数据目录的直接内容\nfor file in tf.io.gfile.glob(os.path.join(DATA_DIR, \"*\")): print(f\"\\t--> {file}\")\n\n    \nprint(\"\\n\\n... DATA ACCESS SETUP COMPLETED ...\\n\") #数据访问设置完成","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:18:33.134981Z","iopub.execute_input":"2023-12-21T03:18:33.135557Z","iopub.status.idle":"2023-12-21T03:18:33.147707Z","shell.execute_reply.started":"2023-12-21T03:18:33.135511Z","shell.execute_reply":"2023-12-21T03:18:33.146667Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2.3 利用XLA优化","metadata":{}},{"cell_type":"markdown","source":"**XLA**（加速线性代数）是用于线性代数的特定领域编译器，可以加速TensorFlow模型，而且可能无需更改源代码。**其结果是在速度和内存使用方面的改进**。\n\n  \n\n当运行TensorFlow程序时，所有操作都由TensorFlow执行器单独执行。每个TensorFlow操作都有一个预编译的GPU/TPU内核实现，由执行器进行分派。\n\nXLA为我们提供了运行模型的另一种模式：它将TensorFlow图编译成为针对给定模型生成的一系列计算内核。因为这些内核是针对模型的唯一的，它们可以利用模型特定信息进行优化。\n\nWarning:\n\nXLA目前无法编译那些维度无法推断的函数：也就是说，如果无法在运行整个计算之前推断所有张量的维度。\n\nNote:\n\n·XLA编译仅应用于编译为图形的代码（在TF2中，这仅适用于`tf.function`内部的代码）。<br />\n·`jit_compile` API 具有“必须编译”的语义，即要么整个函数与XLA一起编译，要么抛出`errors.InvalidArgumentError`异常。","metadata":{}},{"cell_type":"code","source":"print(f\"\\n... XLA OPTIMIZATIONS STARTING ...\\n\") #XLA优化开始\n\nprint(f\"\\n... CONFIGURE JIT (JUST IN TIME) COMPILATION ...\\n\") #配置JIT（即时编译）\n# 启用XLA优化（在使用@tf.function调用时可提高10%的速度）\ntf.config.optimizer.set_jit(True)\n\nprint(f\"\\n... XLA OPTIMIZATIONS COMPLETED ...\\n\") #XLA优化完成","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:18:39.973859Z","iopub.execute_input":"2023-12-21T03:18:39.974177Z","iopub.status.idle":"2023-12-21T03:18:39.980889Z","shell.execute_reply.started":"2023-12-21T03:18:39.974140Z","shell.execute_reply":"2023-12-21T03:18:39.979870Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 2.4 基本数据定义与初始化","metadata":{}},{"cell_type":"code","source":"print(\"\\n... BASIC DATA SETUP STARTING ...\\n\\n\") #基本数据设置开始\n\nprint(\"\\n... SET PATH INFORMATION ..\\n\") #设置路径信息\nSEG_DIR = \"/kaggle/input/sartorius-segmentation-train-mask-dataset-npz\"\nLC_DIR = os.path.join(DATA_DIR, \"LIVECell_dataset_2021\")\nLC_ANN_DIR = os.path.join(LC_DIR, \"annotations\")\nLC_IMG_DIR = os.path.join(LC_DIR, \"images\")\nTRAIN_DIR = os.path.join(DATA_DIR, \"train\")\nTEST_DIR = os.path.join(DATA_DIR, \"test\")\nSEMI_DIR = os.path.join(DATA_DIR, \"train_semi_supervised\")\n\nprint(\"\\n... TRAIN DATAFRAME ...\\n\") #训练数据集\n\n# 修正训练数据集（将RLE聚合在一起）\nTRAIN_CSV = os.path.join(DATA_DIR, \"train.csv\")\ntrain_df = pd.read_csv(TRAIN_CSV)\ndisplay(train_df)\n\nprint(\"\\n... SS DATAFRAME ..\\n\")\nSS_CSV = os.path.join(DATA_DIR, \"sample_submission.csv\")\nss_df = pd.read_csv(SS_CSV)\nss_df[\"img_path\"] = ss_df[\"id\"].apply(lambda x: os.path.join(TEST_DIR, x+\".png\")) # 同时捕获图像路径\ndisplay(ss_df)\n\nCELL_TYPES = list(train_df.cell_type.unique())\nFIRST_SHSY5Y_IDX = 0\nFIRST_ASTRO_IDX  = 1\nFIRST_CORT_IDX   = 2\n\n# 这对于绘图是必需的，以便较小的分布位于顶部\nARB_SORT_MAP = {\"astro\":0, \"shsy5y\":1, \"cort\":2}\n\nprint(\"\\n... CELL TYPES ..\") #细胞类型\nfor x in CELL_TYPES: print(f\"\\t--> {x}\")\n    \nprint(\"\\n\\n... BASIC DATA SETUP FINISHING ...\\n\") # 基本数据设置完成","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:18:48.423939Z","iopub.execute_input":"2023-12-21T03:18:48.424256Z","iopub.status.idle":"2023-12-21T03:18:49.155599Z","shell.execute_reply.started":"2023-12-21T03:18:48.424222Z","shell.execute_reply":"2023-12-21T03:18:49.154760Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 3 辅助函数和类","metadata":{}},{"cell_type":"code","source":"# ref: https://www.kaggle.com/paulorzp/run-length-encode-and-decode\n# modified from: https://www.kaggle.com/inversion/run-length-decoding-quick-start\n\n# 函数：解码 Run-Length 编码的掩码\ndef rle_decode(mask_rle, shape, color=1):\n    \"\"\" TBD\n    \n    Args:\n        mask_rle (str): 以字符串格式表示的 Run-Length 编码（start length）\n        shape (tuple of ints): 返回数组的形状（height, width）\n        color (int): 表示掩码的值\n    Returns:\n        Mask (np.array)\n            - 1 表示掩码\n            - 0 表示背景\n\n    \"\"\"\n    # 将字符串通过空格分隔，然后转换为整数数组\n    s = np.array(mask_rle.split(), dtype=int)\n\n    # 每个偶数值表示起始位置，每个奇数值表示“run”长度\n    starts = s[0::2] - 1\n    lengths = s[1::2]\n    ends = starts + lengths\n\n    # 图像实际上是扁平化的，因为 RLE 是一维的“run”\n    if len(shape)==3:\n        h, w, d = shape\n        img = np.zeros((h * w, d), dtype=np.float32)\n    else:\n        h, w = shape\n        img = np.zeros((h * w,), dtype=np.float32)\n\n    # 这里的颜色实际上可以是任何你想要的整数!\n    for lo, hi in zip(starts, ends):\n        img[lo : hi] = color\n        \n    # 别忘了将图像还原为原始形状\n    return img.reshape(shape)\n\n\n\n\n# https://www.kaggle.com/namgalielei/which-reshape-is-used-in-rle\n# 函数：从顶到底先解码 Run-Length 编码的掩码\n\ndef rle_decode_top_to_bot_first(mask_rle, shape):\n    \"\"\" TBD\n    \n    Args:\n        mask_rle (str): 以字符串格式表示的 Run-Length 编码（start length）\n        shape (tuple of ints): 返回数组的形状（height, width）\n    Returns:\n        Mask (np.array)\n            - 1 表示掩码\n            - 0 表示背景\n\n    \"\"\"\n    s = mask_rle.split()\n    starts, lengths = [np.asarray(x, dtype=int) for x in (s[0:][::2], s[1:][::2])]\n    starts -= 1\n    ends = starts + lengths\n    img = np.zeros(shape[0]*shape[1], dtype=np.uint8)\n    for lo, hi in zip(starts, ends):\n        img[lo:hi] = 1\n    return img.reshape((shape[1], shape[0]), order='F').T  # 从上到下首先进行重新整形\n\n# ref.: https://www.kaggle.com/stainsby/fast-tested-rle\n# 函数：编码图像掩码为 Run-Length 格式\ndef rle_encode(img):\n    \"\"\" TBD\n    \n    Args:\n        img (np.array): \n            - 1 表示掩码\n            - 0 表示背景\n    Returns: \n        以字符串格式表示的 Run-Length 编码\n    \"\"\"\n    \n    pixels = img.flatten()\n    pixels = np.concatenate([[0], pixels, [0]])\n    runs = np.where(pixels[1:] != pixels[:-1])[0] + 1\n    runs[1::2] -= runs[::2]\n    return ' '.join(str(x) for x in runs)\n\n\n# 函数：将嵌套列表扁平化\ndef flatten_l_o_l(nested_list):\n    return [item for sublist in nested_list for item in sublist]\n\n# 函数：加载 JSON 文件到字典\ndef load_json_to_dict(json_path):\n    \"\"\" 从 JSON 文件加载数据到字典 \"\"\"\n    with open(json_path) as json_file:\n        data = json.load(json_file)\n    return data\n\n# 函数：获取轮廓的包围框\ndef grab_contours(cnts):\n    \"\"\" 获取轮廓的包围框 \"\"\"\n    \n    # 如果 cv2.findContours 返回的轮廓元组长度为 '2'，则使用 OpenCV v2.4、v4-beta 或 v4-official\n    if len(cnts) == 2:\n        cnts = cnts[0]\n\n    # 如果轮廓元组的长度为 '3'，则使用 OpenCV v3、v4-pre 或 v4-alpha\n    elif len(cnts) == 3:\n        cnts = cnts[1]\n\n    # 否则 OpenCV 又一次更改了 cv2.findContours 返回的签名，我完全不知道发生了什么\n    else:\n        raise Exception(\"Contours 元组的长度必须为 2 或 3，否则 OpenCV 又一次更改了 cv2.findContours 返回的签名。请参阅 OpenCV 的文档。\")\n\n    # 返回实际的轮廓数组\n    return cnts\n\n# 函数：获取给定掩码的边界框（tl, br）\ndef get_contour_bbox(msk):\n    \"\"\" 返回给定掩码的边界框 (tl, br) \"\"\"\n    \n    # 获取轮廓（应该只有一个）\n    cnts = cv2.findContours(msk.copy(), cv2.RETR_EXTERNAL, cv2.CHAIN_APPROX_SIMPLE)\n\n    contour = grab_contours(cnts)\n    \n    if len(contour)==0:\n        return None\n    else:\n        contour = contour[0]\n    \n    # 获取极端坐标\n    tl = (tuple(contour[contour[:, :, 0].argmin()][0])[0], \n          tuple(contour[contour[:, :, 1].argmin()][0])[1])\n    br = (tuple(contour[contour[:, :, 0].argmax()][0])[0], \n          tuple(contour[contour[:, :, 1].argmax()][0])[1])\n    return tl, br\n\n# 函数：使用 TensorFlow 加载 PNG 图像\ndef tf_load_png(img_path):\n    return tf.image.decode_png(tf.io.read_file(img_path), channels=3)\n\n# 函数：获取图像和掩码\ndef get_img_and_mask(img_path, annotation, width, height, mask_only=False, rle_fn=rle_decode):\n    \"\"\" 获取相关的图像数组和图像掩码 \"\"\"\n    img_mask = np.zeros((height, width), dtype=np.uint8)\n    for i, annot in enumerate(annotation): \n        img_mask = np.where(rle_fn(annot, (height, width))!=0, i, img_mask)\n    \n    # 提前退出\n    if mask_only:\n        return img_mask\n    \n    # 否则返回图像\n    img = tf_load_png(img_path)[..., 0]\n    return img, img_mask\n\n# 函数：绘制图像和掩码\ndef plot_img_and_mask(img, mask, bboxes=None, invert_img=True, boost_contrast=True):\n    \"\"\" 绘制图像和相应掩码的函数\n    \n    Args:\n        img (np.arr): 表示细胞结构图像的 1 通道 np 数组\n        mask (np.arr): 表示实例掩码的 1 通道 np 数组（递增 1）\n        bboxes (list of tuples, optional): 包围边界框的坐标 (tl, br)\n        invert_img (bool, optional): 是否反转基础图像\n        boost_contrast (bool, optional): 是否增强基础图像的对比度\n        \n    Returns:\n        None；绘制两个数组，并叠加它们以创建合并图像\n    \"\"\"\n    plt.figure(figsize=(20,10))\n    \n    plt.subplot(1,3,1)\n    _img = np.tile(np.expand_dims(img, axis=-1), 3)\n    \n    # 反转黑白，使黑色变白，白色变黑\n    if invert_img:\n        _img = _img.max()-_img\n    \n    if boost_contrast:\n        _img = np.asarray(ImageEnhance.Contrast(Image.fromarray(_img)).enhance(16))\n    \n    if bboxes:\n        for i, bbox in enumerate(bboxes):\n            mask = cv2.rectangle(mask, bbox[0], bbox[1], (i+1, 0, 0), thickness=2)\n    \n    plt.imshow(_img)\n    plt.axis(False)\n    plt.title(\"Cell Image\", fontweight=\"bold\")\n    \n    plt.subplot(1,3,2)\n    _mask = np.zeros_like(_img)\n    _mask[..., 0] = mask\n    plt.imshow(mask, cmap=\"inferno\")\n    plt.axis(False)\n    plt.title(\"Instance Segmentation Mask\", fontweight=\"bold\")\n    \n    merged = cv2.addWeighted(_img, 0.75, np.clip(_mask, 0, 1)*255, 0.25, 0.0,)\n    plt.subplot(1,3,3)\n    plt.imshow(merged)\n    plt.axis(False)\n    plt.title(\"Cell Image w/ Instance Segmentation Mask Overlay\", fontweight=\"bold\")\n    \n    plt.tight_layout()\n    plt.show()\n\n\ndef pd_get_bboxes(row):\n    \"\"\" 获取给定行/细胞图像的所有边界框 \"\"\"\n    mask = get_img_and_mask(row.img_path, row.annotation, row.width, row.height, mask_only=True)\n    return [get_contour_bbox(np.where(mask==i, 1, 0).astype(np.uint8)) for i in range(1, mask.max()+1)]\n\ndef get_bbox_stats(bbox_list, style=\"area\"): \n    \"\"\" 获取边界框列表的统计信息 \n    \n    Args:\n        bbox_list(): 边界框列表\n        style (str, optional): 统计风格，可以是 \"area\"、\"width\" 或其他\n    Returns:\n        统计信息列表\n    \"\"\"\n    bbox_stats = []\n    for box in bbox_list:\n        try:\n            if style==\"area\":\n                bbox_stats.append(float((box[1][0]-box[0][0])*(box[1][1]-box[0][1])))\n            elif style==\"width\":\n                bbox_stats.append(float(box[1][0]-box[0][0]))\n            else:\n                bbox_stats.append(float(box[1][1]-box[0][1]))\n        except:\n            bbox_stats.append(0.0)\n    return bbox_stats\n","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:18:55.670815Z","iopub.execute_input":"2023-12-21T03:18:55.671403Z","iopub.status.idle":"2023-12-21T03:18:55.718283Z","shell.execute_reply.started":"2023-12-21T03:18:55.671364Z","shell.execute_reply":"2023-12-21T03:18:55.717151Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 4 数据集创建与探索","metadata":{}},{"cell_type":"markdown","source":"## 4.0 更新训练数据框架","metadata":{}},{"cell_type":"markdown","source":"我们需要对训练数据框架进行一些更改\n\n 根据 id 进行聚合\n \n 添加特定列","metadata":{}},{"cell_type":"code","source":"# 在训练数据下进行聚合\ntrain_df[\"img_path\"] = train_df[\"id\"].apply(lambda x: os.path.join(TRAIN_DIR, x+\".png\"))  # 同时捕获图像路径\ntmp_df = train_df.drop_duplicates(subset=[\"id\", \"img_path\"]).reset_index(drop=True)\ntmp_df[\"annotation\"] = train_df.groupby(\"id\")[\"annotation\"].agg(list).reset_index(drop=True)\ntrain_df = tmp_df.copy()\ntrain_df[\"seg_path\"] = train_df.id.apply(lambda x: os.path.join(SEG_DIR, f\"{x}.npz\"))\ndisplay(train_df)\n\n\"\"\"\n这段代码对训练数据框架进行了一些处理。\n首先，它为每个图像的ID构造了图像路径。\n然后，通过将ID和图像路径去重，创建了一个新的数据框架（tmp_df）。\n接着，使用groupby和agg方法，将相同ID的注释聚合为列表，\n并将其添加到tmp_df中的新列“annotation”中。\n最后，更新了训练数据框架，并添加了一个新列“seg_path”，\n该列包含分割图的路径。\n\"\"\"","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:19:06.865946Z","iopub.execute_input":"2023-12-21T03:19:06.866863Z","iopub.status.idle":"2023-12-21T03:19:07.215267Z","shell.execute_reply.started":"2023-12-21T03:19:06.866821Z","shell.execute_reply":"2023-12-21T03:19:07.214285Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4.1 可视化训练数据","metadata":{}},{"cell_type":"markdown","source":"\"\"\"创建一个函数，以绘制单个示例（数据框架的一行）的相关信息\"\"\"","metadata":{}},{"cell_type":"code","source":"for i in range(2, 70, 8):\n    print(f\"\\n\\n\\n\\n... RELEVANT DATAFRAME ROW - INDEX={i} ...\\n\")\n    display(train_df.iloc[i:i+1])\n    img, msk = get_img_and_mask(**train_df[[\"img_path\", \"annotation\", \"width\", \"height\"]].iloc[i].to_dict())\n    plot_img_and_mask(img, msk)","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:19:15.312714Z","iopub.execute_input":"2023-12-21T03:19:15.313479Z","iopub.status.idle":"2023-12-21T03:19:25.627412Z","shell.execute_reply.started":"2023-12-21T03:19:15.313435Z","shell.execute_reply":"2023-12-21T03:19:25.626448Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"<b><font color=\"red\">快速双重检查 RLE 解码函数</font></b>\n* https://www.kaggle.com/namgalielei/which-reshape-is-used-in-rle","metadata":{}},{"cell_type":"code","source":"# 使用两个不同的 RLE 解码函数对单个细胞进行比较\nx1 = rle_decode_top_to_bot_first(train_df.iloc[0].annotation[0], (train_df.iloc[0].height, train_df.iloc[0].width))\nx2 = rle_decode(train_df.iloc[0].annotation[0], (train_df.iloc[0].height, train_df.iloc[0].width))\n\n# 绘制对比图\nplt.figure(figsize=(15,6))\nplt.subplot(1,2,1)\nplt.imshow(x1, cmap=\"inferno\")\nplt.axis(False)\nplt.title(\"NamGalielei RLE Decode Function\", fontweight=\"bold\") #NamGalielei RLE 解码函数\nplt.subplot(1,2,2)\nplt.imshow(x2, cmap=\"inferno\")\nplt.axis(False)\nplt.title(\"Original RLE Decode Function\", fontweight=\"bold\") #原始 RLE 解码函数\nplt.tight_layout()\nplt.show()\nprint(f\"\\n... 在单个细胞上使用两个函数时，存在 {(x1!=x2).sum()} 个不一致的像素...\\n\")\n\nimg1, msk1 = get_img_and_mask(**train_df[[\"img_path\", \"annotation\", \"width\", \"height\"]].iloc[0].to_dict())\nimg2, msk2 = get_img_and_mask(**train_df[[\"img_path\", \"annotation\", \"width\", \"height\"]].iloc[0].to_dict(), rle_fn=rle_decode_top_to_bot_first)\n\nplot_img_and_mask(img1, msk1)\nplot_img_and_mask(img2, msk2)\n\nprint(f\"\\n... 在所有细胞掩码上使用两个函数时，存在 {(msk2!=msk1).sum()} 个不一致的像素 ...\\n\")\n","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:19:37.600346Z","iopub.execute_input":"2023-12-21T03:19:37.601216Z","iopub.status.idle":"2023-12-21T03:19:39.831148Z","shell.execute_reply.started":"2023-12-21T03:19:37.601172Z","shell.execute_reply":"2023-12-21T03:19:39.830336Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4.2 调查训练数据框架","metadata":{}},{"cell_type":"code","source":"# 调查训练数据框架\nprint(\"\\n\\n... 宽度值计数 ...\")\nfor k, v in train_df.width.value_counts().items():\n    print(f\"\\t--> 有 {v} 张图片的 WIDTH={k}\")\n\nprint(\"\\n\\n... 高度值计数 ...\")\nfor k, v in train_df.height.value_counts().items():\n    print(f\"\\t--> 有 {v} 张图片的 HEIGHT={k}\")\n\nprint(\"\\n\\n... 区域计数 ...\")\nfor k, v in (train_df.width * train_df.height).value_counts().items():\n    print(f\"\\t--> 有 {v} 张图片的 AREA={k}\")\n\nprint(\"\\n\\n... 注意: 所有图片大小相同 ...\\n\")\n\nprint(\"\\n\\n... PLATE TIME 值计数 ...\")\nfor k, v in train_df.plate_time.value_counts().items():\n    print(f\"\\t--> 有 {v} 张图片的 PLATE_TIME={k}\")\nfig = px.histogram(train_df, x=\"plate_time\", color=\"cell_type\", title=\"<b>Plate Time 直方图</b>\")\nfig.show()\n\nprint(\"\\n\\n... 样本日期值计数 ...\")\nfor k, v in train_df.sample_date.value_counts().items():\n    print(f\"\\t--> 有 {v} 张图片的 SAMPLE_DATE={k}\")\nfig = px.histogram(train_df, train_df.sample_date.apply(lambda x: x.replace(\"-\", \"_\")), color=\"cell_type\", title=\"<b>样本日期值直方图</b>\")\nfig.show()\n\nprint(\"\\n\\n... 经过的时间差值计数 ...\")\nfor k, v in train_df.elapsed_timedelta.value_counts().items():\n    print(f\"\\t--> 有 {v} 张图片的 SAMPLE_DATE={k}\")\nfig = px.histogram(train_df, \"elapsed_timedelta\", color=\"cell_type\", title=\"<b>经过的时间差值直方图</b>\")\nfig.show()\n    \nprint(\"\\n\\n... 样本 ID 值计数 (>1) ...\")\nprint(f\"\\t--> 有 {len(train_df[train_df.sample_id.isin([x for x,v in train_df.sample_id.value_counts().items() if v>1])])} 个 SAMPLE_ID 具有多于一张图片\\n\")\nfor k, v in train_df[train_df.sample_id.isin([x for x, v in train_df.sample_id.value_counts().items() if v>1])].reset_index()[\"sample_id\"].value_counts().items():\n    print(f\"\\t--> 有 {v} 张图片的 SAMPLE_ID={k}\")\nfig = px.histogram(train_df[train_df.sample_id.isin([x for x,v in train_df.sample_id.value_counts().items() if v>1])].reset_index(), \"sample_id\", color=\"cell_type\", title=\"<b>SAMPLE_ID 值直方图</b>\")\nfig.show()\n\nprint(\"\\n\\n... CELL_TYPE 值计数 ...\")\nfor k, v in train_df.cell_type.value_counts().items():\n    print(f\"\\t--> 有 {v} 张图片的 CELL_TYPE={k}\")\n    \nfig = px.histogram(train_df, x=\"cell_type\", title=\"<b>Cell Type 直方图</b>\")\nfig.show()\n\nfor ct in CELL_TYPES:\n    print(f\"\\n\\n... 显示 CELL_TYPE {ct.upper()} 的三个示例 ...\\n\")\n    for i in range(3):\n        img, msk = get_img_and_mask(**train_df[train_df.cell_type==ct][[\"img_path\", \"annotation\", \"width\", \"height\"]].sample(3).reset_index(drop=True).iloc[i].to_dict())\n        plot_img_and_mask(img, msk)","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:20:30.603371Z","iopub.execute_input":"2023-12-21T03:20:30.603734Z","iopub.status.idle":"2023-12-21T03:20:38.231346Z","shell.execute_reply.started":"2023-12-21T03:20:30.603699Z","shell.execute_reply":"2023-12-21T03:20:38.230166Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4.3 可视化和调查 LIVECell 数据","metadata":{}},{"cell_type":"code","source":"DEFER = True\n\nif not DEFER:\n    LC_CELL_TYPES = os.listdir(os.path.join(LC_ANN_DIR, \"LIVECell_single_cells\"))\n\n    print(\"\\n... 加载训练 COCO JSON ...\\n\")\n    LC_COCO_TRAIN = os.path.join(LC_ANN_DIR, \"LIVECell\", \"livecell_coco_train.json\")\n\n    print(\"\\n... 加载验证 COCO JSON ...\\n\")\n    LC_COCO_VAL = os.path.join(LC_ANN_DIR, \"LIVECell\", \"livecell_coco_val.json\")\n\n    print(\"\\n... 加载测试 COCO JSON ...\\n\")\n    LC_COCO_TEST = os.path.join(LC_ANN_DIR, \"LIVECell\", \"livecell_coco_test.json\")\n\n    LC_SC_TRAIN = {\n        lc_ct: os.path.join(LC_ANN_DIR, \"LIVECell_single_cells\", lc_ct, f\"livecell_{lc_ct}_train.json\") \\\n        for lc_ct in LC_CELL_TYPES\n    }\n    LC_SC_VAL = {\n        lc_ct: os.path.join(LC_ANN_DIR, \"LIVECell_single_cells\", lc_ct, f\"livecell_{lc_ct}_val.json\") \\\n        for lc_ct in LC_CELL_TYPES\n    }\n    LC_SC_TEST = {\n        lc_ct: os.path.join(LC_ANN_DIR, \"LIVECell_single_cells\", lc_ct, f\"livecell_{lc_ct}_test.json\") \\\n        for lc_ct in LC_CELL_TYPES\n    }\n\n    print(LC_SC_TRAIN)\n    print(LC_SC_VAL)\n    print(LC_SC_TEST)\n","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:20:56.696947Z","iopub.execute_input":"2023-12-21T03:20:56.697809Z","iopub.status.idle":"2023-12-21T03:20:56.707717Z","shell.execute_reply.started":"2023-12-21T03:20:56.697757Z","shell.execute_reply":"2023-12-21T03:20:56.706710Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4.4 可视化和调查半监督训练数据","metadata":{}},{"cell_type":"code","source":"semi_df = pd.DataFrame()\n\nsemi_df[\"cell_type\"] = [x.split(\"[\", 1)[0] for x in tf.io.gfile.listdir(SEMI_DIR)]\nsemi_df[\"compound\"] = [x.split(\"]\", 1)[0].split(\"[\", 1)[-1] for x in tf.io.gfile.listdir(SEMI_DIR)]\nsemi_df[\"img_path\"] = tf.io.gfile.glob(os.path.join(SEMI_DIR, \"**\"))\n\nfig = px.histogram(semi_df, \"cell_type\", color=\"compound\")\nfig.show()\n\nfig = px.histogram(semi_df, \"compound\", color=\"cell_type\")\nfig.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:21:39.410485Z","iopub.execute_input":"2023-12-21T03:21:39.410844Z","iopub.status.idle":"2023-12-21T03:21:41.952713Z","shell.execute_reply.started":"2023-12-21T03:21:39.410772Z","shell.execute_reply":"2023-12-21T03:21:41.951906Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"plt.figure(figsize=(20,26))\nfor i, img_path in zip(range(15), semi_df.img_path.to_list()):\n    plt.subplot(5,3,i+1)\n    plt.imshow((255-np.asarray(ImageEnhance.Contrast(Image.fromarray(tf_load_png(img_path).numpy())).enhance(16))), cmap=\"inferno\")\n    plt.axis(False)\n    plt.title(img_path.rsplit(\"/\", 1)[-1].rsplit(\".\", 1)[0], fontweight=\"bold\")\n    \nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:21:52.521911Z","iopub.execute_input":"2023-12-21T03:21:52.522227Z","iopub.status.idle":"2023-12-21T03:21:56.778885Z","shell.execute_reply.started":"2023-12-21T03:21:52.522192Z","shell.execute_reply":"2023-12-21T03:21:56.777401Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4.5 可视化和调查单个细胞","metadata":{}},{"cell_type":"code","source":"DEMO_IDX = 11\nimg, msk = get_img_and_mask(**train_df[[\"img_path\", \"annotation\", \"width\", \"height\"]].iloc[DEMO_IDX].to_dict())\nplot_img_and_mask(img, msk)\n\nplt.figure(figsize=(20, min(80, msk.max()//2)))\nfor i in range(1, msk.max()+1):\n    plt.subplot(10,10,i)\n    tl, br = get_contour_bbox(np.where(msk==i, 1, 0).astype(np.uint8))\n    plt.imshow(np.asarray(ImageEnhance.Contrast(Image.fromarray(255-img.numpy())).enhance(16))[tl[1]:br[1], tl[0]:br[0]], cmap=\"magma\")\n    plt.axis(False)\n    plt.title(f\"{i}\", fontweight=\"bold\")\n    if i==100:\n        break\n    \nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:22:38.767212Z","iopub.execute_input":"2023-12-21T03:22:38.767544Z","iopub.status.idle":"2023-12-21T03:22:47.111397Z","shell.execute_reply.started":"2023-12-21T03:22:38.767500Z","shell.execute_reply":"2023-12-21T03:22:47.110500Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4.6 将边界框添加到数据框中\n\n---\n\n让我们遍历图像，按顺序将边界框坐标添加到数据框中（与 RLE 匹配）\n\n我们还将边界框的0-1归一化版本添加到数据框中，以供以后进行训练\n","metadata":{}},{"cell_type":"code","source":"# 大约需要1分钟\nprint(\"\\n... 创建完整尺寸的边界框 ...\\n\")\ntrain_df[\"bboxes\"] = train_df.parallel_apply(pd_get_bboxes, axis=1)\ndisplay(train_df.head())\n\nprint(\"\\n... 创建缩小尺寸（0-1）的边界框 ...\\n\")\nIMG_O_W, IMG_O_H = train_df.iloc[0].width, train_df.iloc[0].height\ntrain_df[\"scaled_bboxes\"] = train_df.bboxes.progress_apply(lambda box_list: [((box[0][0]/IMG_O_W, box[0][1]/IMG_O_H), (box[1][0]/IMG_O_W,box[1][1]/IMG_O_H)) if box else None for box in box_list])\n\n# CORT\nimg, msk = get_img_and_mask(**train_df[[\"img_path\", \"annotation\", \"width\", \"height\"]].iloc[FIRST_SHSY5Y_IDX].to_dict())\nplot_img_and_mask(img, msk, bboxes=train_df.iloc[FIRST_SHSY5Y_IDX].bboxes)\n\n# ASTRO\nimg, msk = get_img_and_mask(**train_df[[\"img_path\", \"annotation\", \"width\", \"height\"]].iloc[FIRST_ASTRO_IDX].to_dict())\nplot_img_and_mask(img, msk, bboxes=train_df.iloc[FIRST_ASTRO_IDX].bboxes)\n\n# SHSY5Y\nimg, msk = get_img_and_mask(**train_df[[\"img_path\", \"annotation\", \"width\", \"height\"]].iloc[FIRST_CORT_IDX].to_dict())\nplot_img_and_mask(img, msk, bboxes=train_df.iloc[FIRST_CORT_IDX].bboxes)\n","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:23:15.969429Z","iopub.execute_input":"2023-12-21T03:23:15.970145Z","iopub.status.idle":"2023-12-21T03:24:08.236060Z","shell.execute_reply.started":"2023-12-21T03:23:15.970107Z","shell.execute_reply":"2023-12-21T03:24:08.235121Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 4.7 获取关于边界框的统计信息","metadata":{}},{"cell_type":"code","source":"train_df[\"bbox_widths\"] = train_df.bboxes.apply(lambda x: get_bbox_stats(x, style=\"width\"))\ntrain_df[\"bbox_heights\"] = train_df.bboxes.apply(lambda x: get_bbox_stats(x, style=\"height\"))\ntrain_df[\"bbox_areas\"] = train_df.bboxes.apply(lambda x: get_bbox_stats(x, style=\"area\"))\n\ntrain_df[\"scaled_bbox_widths\"] = train_df.scaled_bboxes.apply(lambda x: get_bbox_stats(x, style=\"width\"))\ntrain_df[\"scaled_bbox_heights\"] = train_df.scaled_bboxes.apply(lambda x: get_bbox_stats(x, style=\"height\"))\ntrain_df[\"scaled_bbox_areas\"] = train_df.scaled_bboxes.apply(lambda x: get_bbox_stats(x, style=\"area\"))\n\ndisplay(train_df.head())\n\n# Plot\npx.scatter(train_df.sort_values(by=\"cell_type\", key=lambda x: x.map(ARB_SORT_MAP))[[\"cell_type\", \"bbox_widths\", \"bbox_heights\", \"bbox_areas\"]].explode(column=[\"bbox_widths\",\"bbox_heights\", \"bbox_areas\"]), x=\"bbox_widths\", y=\"bbox_heights\", color=\"cell_type\", title=\"<b>Cell Bounding Box Sizes (WxH)</b>\")","metadata":{"execution":{"iopub.status.busy":"2023-12-21T03:24:53.335428Z","iopub.execute_input":"2023-12-21T03:24:53.335757Z","iopub.status.idle":"2023-12-21T03:24:55.692009Z","shell.execute_reply.started":"2023-12-21T03:24:53.335723Z","shell.execute_reply":"2023-12-21T03:24:55.691173Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 5 modelling\n\n---\nEfficientDET 用于目标检测、分类和分割","metadata":{}},{"cell_type":"markdown","source":"## 5.0 EfficientDet 实用函数和常量","metadata":{}},{"cell_type":"code","source":"IMAGE_SHAPE = (train_df.iloc[0].height, train_df.iloc[0].width, 3)\nINPUT_SHAPE = (640, 640, 3)\nSEG_SHAPE = (INPUT_SHAPE[0] // 4, INPUT_SHAPE[1] // 4, 1)\nMODEL_LEVEL = \"d1\"\nMODEL_NAME = f\"efficientdet-{MODEL_LEVEL}\"\nBATCH_SIZE = 8\nN_EVAL = 50\nN_TRAIN = len(train_df) - N_EVAL\nN_TEST = len(ss_df)\nDEBUG = N_TEST == 3\nN_EPOCH = 40\nN_EX_PER_REC = 280\nCLASS_LABELS = list(train_df.cell_type.unique())\nN_CLASSES_OD = len(CLASS_LABELS) + 1  # 背景 + 3 种细胞类型\nN_CLASSES_SEG = 2  # 背景 + 前景（细胞）\nMAX_N_INSTANCES = int(100 * np.ceil(train_df.bboxes.apply(len).max() / 100))\n\n# 是否从头开始训练还是加载模型\nDO_TRAIN = False\nPRETRAINED_MODEL_DIR = \"/kaggle/input/model-weights-40-epoch-efficientdet-d1-640\"\n\nprint(\"\\n ... 超参数常数 ...\")\nprint(f\"\\t--> 模型名称          : {MODEL_NAME}\")\nprint(f\"\\t--> 批量大小          : {BATCH_SIZE}\")\nprint(f\"\\t--> 图像形状          : {IMAGE_SHAPE}\")\nprint(f\"\\t--> 输入形状          : {INPUT_SHAPE}\")\nprint(f\"\\t--> 分割形状          : {SEG_SHAPE}\")\n","metadata":{"execution":{"iopub.status.busy":"2023-12-14T10:30:08.231152Z","iopub.execute_input":"2023-12-14T10:30:08.231411Z","iopub.status.idle":"2023-12-14T10:30:08.247827Z","shell.execute_reply.started":"2023-12-14T10:30:08.231376Z","shell.execute_reply":"2023-12-14T10:30:08.246998Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.1 加载EfficientDet模型并初始化","metadata":{}},{"cell_type":"code","source":"# 获取 EfficientDET 模型的配置信息\nconfig = hparams_config.get_efficientdet_config(MODEL_NAME)\n# 关键配置项\nKEY_CONFIGS = [\n    \"name\", \"image_size\", \"num_classes\", \"seg_num_classes\", \"heads\", \"train_file_pattern\",\n    \"val_file_pattern\", \"model_name\", \"model_dir\", \"pretrained_ckpt\", \"batch_size\", \"eval_samples\",\n    \"num_examples_per_epoch\", \"num_epochs\", \"steps_per_execution\", \"steps_per_epoch\", \n    \"profile\", \"val_json_file\", \"max_instances_per_image\", \"mixed_precision\", \n    \"learning_rate\", \"lr_warmup_init\", \"mean_rgb\", \"stddev_rgb\",\"scale_range\",\n              ]\n\n# 打印配置信息\nfor k in config.keys():\n    if k==\"model_optimizations\":\n        continue\n    elif k==\"nms_configs\":\n        for _k, _v in dict(config[k]).items():\n            print(f\"PARAMETER: {'     ' if _k not in KEY_CONFIGS else ' *** '}nms_config_{_k: <16}  ---->    VALUE:  {_v}\")\n        \n    else:\n        print(f\"PARAMETER: {'     ' if k not in KEY_CONFIGS else ' *** '}{k: <27}  ---->    VALUE:  {config[k]}\")","metadata":{"execution":{"iopub.status.busy":"2023-12-14T10:30:08.249305Z","iopub.execute_input":"2023-12-14T10:30:08.249587Z","iopub.status.idle":"2023-12-14T10:30:08.269326Z","shell.execute_reply.started":"2023-12-14T10:30:08.249555Z","shell.execute_reply":"2023-12-14T10:30:08.268397Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"DO_ADV_PROP = True\nMODEL_DIR = f\"/kaggle/working/{MODEL_NAME}-finetune\"\n\nif TPU:\n    TFRECORD_DIR = os.path.join(KaggleDatasets().get_gcs_path('effdet-d5-dataset-sartorius'), \"tfrecords\")\nelse:\n    TFRECORD_DIR = \"/kaggle/working/tfrecords\"\n\nos.makedirs(MODEL_DIR, exist_ok=True)\nconfig = hparams_config.get_efficientdet_config(MODEL_NAME)\noverrides = dict(\n    train_file_pattern=os.path.join(TFRECORD_DIR, \"train\", \"*.tfrec\"),\n    val_file_pattern=os.path.join(TFRECORD_DIR, \"val\", \"*.tfrec\"),\n    test_file_pattern=os.path.join(TFRECORD_DIR, \"test\", \"*.tfrec\"),\n    model_name=MODEL_NAME,\n    model_dir=MODEL_DIR,\n    pretrained_ckpt=MODEL_NAME,\n    batch_size=BATCH_SIZE,\n    eval_samples=N_EVAL,\n    num_examples_per_epoch=N_TRAIN,\n    num_epochs=N_EPOCH,\n    steps_per_execution=1,\n    steps_per_epoch=N_TRAIN//BATCH_SIZE,\n    profile=None, val_json_file=None,\n    heads=['object_detection', 'segmentation'],\n    image_size=INPUT_SHAPE[:-1],\n    num_classes=N_CLASSES_OD,\n    seg_num_classes=N_CLASSES_SEG,\n    max_instances_per_image=MAX_N_INSTANCES,\n    input_rand_hflip=False, jitter_min=0.99, jitter_max=1.01,\n    skip_crowd_during_training=False,\n)\nconfig.override(overrides, True)\nconfig.nms_configs.max_output_size = MAX_N_INSTANCES\n\n# 修改输入预处理方式\nif DO_ADV_PROP:\n    config.override(dict(mean_rgb=0.0, stddev_rgb=1.0, scale_range=True), True)\n\ntf.keras.backend.clear_session()\n\nmodel = efficientdet_keras.EfficientDetModel(config=config)\nmodel.build((1, *INPUT_SHAPE))\n\nprint(\"\\n... 模型预测 ...\\n\")\npreds = model.predict(np.zeros((1, *INPUT_SHAPE)))\nfor i, name in enumerate([\"bboxes\", \"confidences\", \"classes\", \"valid_len\", \"segmentation map\"]):\n    print(name)\n    print(preds[i].shape)\n    try:\n        if preds[i].shape[-2] == 64:\n            print(preds[i][0, 0, 0, :5])\n        else:\n            print(preds[i][0, :5])\n    except:\n        print(preds[i][0])\n    print()\n","metadata":{"execution":{"iopub.status.busy":"2023-12-14T10:30:08.270743Z","iopub.execute_input":"2023-12-14T10:30:08.270957Z","iopub.status.idle":"2023-12-14T10:30:34.057453Z","shell.execute_reply.started":"2023-12-14T10:30:08.270930Z","shell.execute_reply":"2023-12-14T10:30:34.056530Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.2 创建具有正确结构的数据集\n---\n\n* **INPUT**\n    * Raw Image (256x256x3)\n\n\n* **OUTPUT/TARGET**\n    * Bounding Boxes\n    * Instance Classes\n    * Segmented Image (64x64x3)\n    \n---\n\n首先创建TFRecord数据集，然后实例化数据加载器","metadata":{}},{"cell_type":"code","source":"def create_id_to_iloc_map(df):\n    \"\"\"\n    创建映射以允许数字文件名\n        --> 原始 df 中的索引 --> id\n    \"\"\"\n    return {v: k for k, v in df.id.to_dict().items()}\n\nTRAIN_ID_2_ILOC = create_id_to_iloc_map(train_df)\nTEST_ID_2_ILOC = create_id_to_iloc_map(ss_df)\n\n\ndef tf_load_image(path, resize_to=INPUT_SHAPE):\n    \"\"\"使用 TF 仅加载图像并获得正确形状\n\n    Args:\n        path (tf.string): 要加载的图像的路径\n        resize_to (tuple, optional): 调整图像大小\n    \n    Returns:\n        3 通道的 tf.Constant 图像，准备用于训练/推理\n    \n    \"\"\"\n    img_bytes = tf.io.read_file(path)\n    img = tf.image.decode_png(img_bytes, channels=resize_to[-1])\n    img = tf.image.resize(img, resize_to[:-1])\n    img = tf.cast(img, tf.uint8)\n    \n    return img\n\ndef load_npz(path, resize_to=SEG_SHAPE, to_binary=True):\n    np_arr = np.load(path)[\"arr_0\"]\n    if to_binary:\n        return np.where(cv2.resize(np_arr, resize_to[:-1]) > 0, 1, 0).reshape(resize_to).astype(np.uint8)\n    else:\n        return cv2.resize(np_arr, resize_to[:-1]).reshape(resize_to).astype(np.int32)\n\ndef image_preprocess(image, image_size, mean_rgb=config.mean_rgb, stddev_rgb=config.stddev_rgb):\n    \"\"\"对推理进行图像预处理。\n\n    Args:\n        image: 输入图像，可以是张量或 numpy 数组。\n        image_size: 图像大小的单个整数或两个整数的元组，格式为 (image_height, image_width)。\n        mean_rgb: RGB 的均值，可以是浮点列表或浮点值。\n        stddev_rgb: RGB 的标准差，可以是浮点列表或浮点值。\n    \n    Returns:\n        (image, scale): 处理后的图像和其缩放比例的元组。\n  \"\"\"\n    input_processor = dataloader.DetectionInputProcessor(image, image_size)\n    input_processor.normalize_image(mean_rgb, stddev_rgb)\n    input_processor.set_scale_factors_to_output_size()\n    image = input_processor.resize_and_crop_image()\n    image_scale = input_processor.image_scale_to_original\n    return image, image_scale\n\n\ndef _bytes_feature(value, is_list=False):\n    \"\"\"从字符串/字节返回 bytes_list。\"\"\"\n    if isinstance(value, type(tf.constant(0))):\n        value = value.numpy()  # BytesList 不会从 EagerTensor 中解包字符串。\n    \n    if not is_list:\n        value = [value]\n    \n    return tf.train.Feature(bytes_list=tf.train.BytesList(value=value))\n\ndef _float_feature(value, is_list=False):\n    \"\"\"从浮点/双精度浮点返回 float_list。\"\"\"\n        \n    if not is_list:\n        value = [value]\n        \n    return tf.train.Feature(float_list=tf.train.FloatList(value=value))\n\ndef _int64_feature(value, is_list=False):\n    \"\"\"从布尔值/枚举/整数/无符号整数返回 int64_list。\"\"\"\n        \n    if not is_list:\n        value = [value]\n        \n    return tf.train.Feature(int64_list=tf.train.Int64List(value=value))\n\ndef serialize_raw(example_data):\n    \"\"\"\n    从 4 个特征创建 tf.Example 消息，准备写入文件。\n\n    Args:\n        example_data: 来自 pandas 行的所有内容\n        style (str, optional): 要执行的子集是什么... [train|val]\n            [test] 将通过不同的函数处理\n    \n    Returns:\n        准备写入文件的 tf.Example 消息\n    \"\"\"\n    \n    image_object_mask = tf.io.encode_png(load_npz(example_data[\"seg_path\"]))\n    \n    image_height = INPUT_SHAPE[0]\n    image_width = INPUT_SHAPE[1]\n    image_source_id = image_filename = f\"{TRAIN_ID_2_ILOC[example_data['id']]:>05}\".encode('utf8')\n    \n    image_encoded = tf.io.encode_png(tf_load_image(example_data[\"img_path\"]))\n    image_key_sha256 = hashlib.sha256(image_encoded).hexdigest().encode('utf8')\n    image_format = example_data[\"img_path\"][-4:].encode('utf8')  # png\n    \n    image_object_bbox_xmins, image_object_bbox_xmaxs  = [], []\n    image_object_bbox_ymins, image_object_bbox_ymaxs  = [], []\n    image_object_class_text, image_object_class_label = [], []\n    image_object_is_crowd, image_object_area = [], []\n    for i, box in enumerate(example_data[\"scaled_bboxes\"]):\n        if box and example_data[\"bbox_areas\"][i] > 0.0:\n            image_object_bbox_xmins.append(box[0][0])\n            image_object_bbox_xmaxs.append(box[1][0])\n            image_object_bbox_ymins.append(box[0][1])\n            image_object_bbox_ymaxs.append(box[1][1])\n            image_object_class_text.append(example_data[\"cell_type\"].encode('utf8'))\n            image_object_class_label.append(ARB_SORT_MAP[example_data[\"cell_type\"]])\n            image_object_is_crowd.append(0)\n            image_object_area.append(example_data[\"scaled_bbox_areas\"][i])\n    \n    # 创建一个字典，将特征名称映射到 tf.Example 兼容的数据类型。\n    feature_dict = {\n        'image/height': _int64_feature(image_height),\n        'image/width': _int64_feature(image_width),\n        'image/filename': _bytes_feature(image_filename),\n        'image/source_id': _bytes_feature(image_source_id),\n        'image/key/sha256': _bytes_feature(image_key_sha256),\n        'image/encoded': _bytes_feature(image_encoded),\n        'image/format': _bytes_feature(image_format),\n        'image/object/bbox/xmin': _float_feature(image_object_bbox_xmins, is_list=True),\n        'image/object/bbox/xmax': _float_feature(image_object_bbox_xmaxs, is_list=True),\n        'image/object/bbox/ymin': _float_feature(image_object_bbox_ymins, is_list=True),\n        'image/object/bbox/ymax': _float_feature(image_object_bbox_ymaxs, is_list=True),\n        'image/object/class/text': _bytes_feature(image_object_class_text, is_list=True),\n        'image/object/class/label': _int64_feature(image_object_class_label, is_list=True),\n        'image/object/is_crowd': _int64_feature(image_object_is_crowd, is_list=True),\n        'image/object/area': _float_feature(image_object_area, is_list=True),\n        'image/object/mask': _bytes_feature(image_object_mask),\n    }\n       \n    # 使用 tf.train.Example 创建一个 Features 消息。\n    example_proto = tf.train.Example(features=tf.train.Features(feature=feature_dict))\n    return example_proto.SerializeToString()\n\n\ndef serialize_test_raw(example_data):\n    \"\"\"\n    创建 tf.Example 消息，准备写入文件\n\n    Args:\n        example_data: 来自 pandas 行的所有内容\n    \n    Returns:\n        准备写入文件的 tf.Example 消息\n    \"\"\"\n    \n    image_height = INPUT_SHAPE[0]\n    image_width = INPUT_SHAPE[1]\n    image_source_id = image_filename = f\"{TEST_ID_2_ILOC[example_data['id']]:>05}\".encode('utf8')\n    \n    image_encoded = tf.io.encode_png(tf_load_image(example_data[\"img_path\"]))\n    image_key_sha256 = hashlib.sha256(image_encoded).hexdigest().encode('utf8')\n    image_format = example_data[\"img_path\"][-4:].encode('utf8')  # png\n    \n    # 创建一个字典，将特征名称映射到 tf.Example 兼容的数据类型。\n    feature_dict = {\n        'image/height': _int64_feature(image_height),\n        'image/width': _int64_feature(image_width),\n        'image/filename': _bytes_feature(image_filename),\n        'image/source_id': _bytes_feature(image_source_id),\n        'image/key/sha256': _bytes_feature(image_key_sha256),\n        'image/encoded': _bytes_feature(image_encoded),\n        'image/format': _bytes_feature(image_format),\n    }\n       \n    # 使用 tf.train.Example 创建一个 Features 消息。\n    example_proto = tf.train.Example(features=tf.train.Features(feature=feature_dict))\n    return example_proto.SerializeToString()\n\n\ndef write_tfrecords(df, n_ex, n_ex_per_rec=50, serialize_fn=serialize_raw, out_dir=\"/kaggle/working/tfrecords\", ds_type=\"train\"):\n    \"\"\"\"\"\"\n    n_recs = int(np.ceil(n_ex / n_ex_per_rec))\n    \n    # 使 DataFrame 可迭代\n    iter_df = df.iterrows()\n        \n    out_dir = os.path.join(out_dir, ds_type)\n    # 创建文件夹\n    if not os.path.isdir(out_dir):\n        os.makedirs(out_dir, exist_ok=True)\n        \n    # 创建 tfrecords\n    for i in tqdm(range(n_recs), total=n_recs):\n        print(f\"\\n... 正在写入 {ds_type.title()} TFRecord {i+1} of {n_recs} ...\\n\")\n        tfrec_path = os.path.join(out_dir, f\"{ds_type}__{(i+1):02}_{n_recs:02}.tfrec\")\n        \n        # 这将创建 tfrecord\n        with tf.io.TFRecordWriter(tfrec_path) as writer:\n            for ex in tqdm(range(n_ex_per_rec), total=n_ex_per_rec):\n                try:\n                    example = serialize_fn(next(iter_df)[1])\n                    writer.write(example)\n                except:\n                    break\n\n# 训练\nwrite_tfrecords(train_df.iloc[:-N_EVAL], N_TRAIN, n_ex_per_rec=N_EX_PER_REC, serialize_fn=serialize_raw, out_dir=TFRECORD_DIR, ds_type=\"train\")\n    \n# 验证\nwrite_tfrecords(train_df[-N_EVAL:], N_EVAL, n_ex_per_rec=N_EX_PER_REC, serialize_fn=serialize_raw, out_dir=TFRECORD_DIR, ds_type=\"val\")\n\n# 测试\nwrite_tfrecords(ss_df, N_TEST, n_ex_per_rec=N_EX_PER_REC, serialize_fn=serialize_test_raw, out_dir=TFRECORD_DIR, ds_type=\"test\")\n","metadata":{"execution":{"iopub.status.busy":"2023-12-14T10:30:34.059151Z","iopub.execute_input":"2023-12-14T10:30:34.059461Z","iopub.status.idle":"2023-12-14T10:32:06.391353Z","shell.execute_reply.started":"2023-12-14T10:30:34.059421Z","shell.execute_reply":"2023-12-14T10:32:06.390594Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.3 实例化我们的数据加载器\n---\n\n由于增强操作破坏了蒙版，暂时禁用它们","metadata":{}},{"cell_type":"code","source":"train_dl = dataloader.InputReader(\n    file_pattern=config.train_file_pattern,\n    is_training=\"train\" in config.train_file_pattern,\n    max_instances_per_image=config.max_instances_per_image\n)(config.as_dict())\n\nval_dl = dataloader.InputReader(\n    file_pattern=config.val_file_pattern,\n    is_training=\"train\" in config.val_file_pattern,\n    max_instances_per_image=config.max_instances_per_image\n)(config.as_dict())\n\ntest_dl = dataloader.InputReader(\n    file_pattern=config.test_file_pattern,\n    is_training=\"train\" in config.test_file_pattern,\n    max_instances_per_image=config.max_instances_per_image\n)(config.as_dict(), batch_size=1)\n\n\nprint(\"\\n... 训练数据加载器 ...\\n\")\nprint(train_dl)\n\nprint(\"\\n\\n... 验证数据加载器 ...\\n\")\nprint(val_dl)\n\nprint(\"\\n\\n... 测试数据加载器 ...\\n\")\nprint(test_dl)\n\nprint(\"\\n\\n\\n\\n 让我们从训练数据加载器中看一个例子 ...\\n\\n\")\n\nx = next(iter(train_dl))\n\nprint(int(x[1][\"source_ids\"][0]))\nimg, msk = get_img_and_mask(**train_df[[\"img_path\", \"annotation\", \"width\", \"height\"]].iloc[int(x[1][\"source_ids\"][0])].to_dict(), )\nplot_img_and_mask(img, msk)\n\nplt.figure(figsize=(20,10))\n\nplt.subplot(1,3,1)\nplt.imshow(x[0][0])\nplt.axis(False)\nplt.title(\"Cell Image\", fontweight=\"bold\") #细胞图像\n\nplt.subplot(1,3,2)\nplt.imshow(x[1][\"image_masks\"][0][0])\nplt.axis(False)\nplt.title(\"Segmentation Mask Overlay\", fontweight=\"bold\") #分割蒙版叠加\n\nmerged = cv2.addWeighted(np.array(x[0][0]), 0.75, np.clip(cv2.resize(np.tile(np.expand_dims(x[1][\"image_masks\"][0][0], axis=-1), 3), INPUT_SHAPE[:-1]), 0, 1)*255, 0.25, 0.0,)\nplt.subplot(1,3,3)\nplt.imshow(merged)\nplt.axis(False)\nplt.title(\"Cell Image w/ Instance Segmentation Mask Overlay\", fontweight=\"bold\") #带实例分割蒙版叠加的细胞图像\n\nplt.tight_layout()\nplt.show()","metadata":{"execution":{"iopub.status.busy":"2023-12-14T10:32:06.392813Z","iopub.execute_input":"2023-12-14T10:32:06.393367Z","iopub.status.idle":"2023-12-14T10:32:13.983608Z","shell.execute_reply.started":"2023-12-14T10:32:06.393323Z","shell.execute_reply":"2023-12-14T10:32:13.982774Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.4 创建模型并加载预训练权重\n---\n\nCOCO weights","metadata":{}},{"cell_type":"code","source":"if DO_TRAIN:\n    if not os.path.isdir(MODEL_NAME):\n        if DO_ADV_PROP:\n            !wget https://storage.googleapis.com/cloud-tpu-checkpoints/efficientdet/advprop/{MODEL_NAME}.tar.gz\n        else:\n            !wget https://storage.googleapis.com/cloud-tpu-checkpoints/efficientdet/coco2/{MODEL_NAME}.tar.gz\n        !tar -zxf {MODEL_NAME}.tar.gz\n        !rm -rf {MODEL_NAME}.tar.gz\n    \nwith strategy.scope():\n    model = train_lib.EfficientDetNetTrain(config=config)\n    model = setup_model(model, config)\n    \n    if DO_TRAIN:\n        util_keras.restore_ckpt(\n          model=model,\n          ckpt_path_or_file=tf.train.latest_checkpoint(MODEL_NAME),\n          ema_decay=config.moving_average_decay,\n          exclude_layers=['class_net']\n        )\n        ckpt_cb = tf.keras.callbacks.ModelCheckpoint(\n            os.path.join(MODEL_DIR, 'ckpt-{epoch:d}'),\n            verbose=1, save_freq=\"epoch\", save_weights_only=True)\n    else:\n        model.load_weights(os.path.join(PRETRAINED_MODEL_DIR, \"ckpt\"))\nmodel.summary()","metadata":{"execution":{"iopub.status.busy":"2023-12-14T10:32:13.985001Z","iopub.execute_input":"2023-12-14T10:32:13.985338Z","iopub.status.idle":"2023-12-14T10:32:23.285093Z","shell.execute_reply.started":"2023-12-14T10:32:13.985303Z","shell.execute_reply":"2023-12-14T10:32:23.284263Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.5 训练模型\n\n---\n\n由于我在先前的版本中训练了模型，我将加载从那个预先训练的检查点中得到的模型。\n请参阅笔记本的先前版本（版本33）以查看训练运行。","metadata":{}},{"cell_type":"code","source":"if DO_TRAIN:\n    history = model.fit(\n        train_dl,\n        epochs=config.num_epochs,\n        steps_per_epoch=config.steps_per_epoch,\n        callbacks=[ckpt_cb,],\n        validation_data=val_dl,\n        validation_steps=N_EVAL//BATCH_SIZE\n    )\nelse:\n    print(model.evaluate(train_dl, steps=config.steps_per_epoch))\n    print(model.evaluate(val_dl, steps=N_EVAL//BATCH_SIZE))","metadata":{"execution":{"iopub.status.busy":"2023-12-14T10:32:23.286512Z","iopub.execute_input":"2023-12-14T10:32:23.286830Z","iopub.status.idle":"2023-12-14T10:35:01.200210Z","shell.execute_reply.started":"2023-12-14T10:32:23.286791Z","shell.execute_reply":"2023-12-14T10:35:01.199394Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.6 验证模型是否在学习","metadata":{}},{"cell_type":"code","source":"\n# 设置中文字符的字体\nplt.rcParams['font.sans-serif'] = ['SimHei']  # 使用宋体作为中文字体\nplt.rcParams['axes.unicode_minus'] = False  # 处理减号显示问题\n\ndef plot_gt(_image, _gt_classes, _gt_boxes, _gt_mask):\n    img_class = int(_gt_classes.numpy()[0])\n    img_boxes = _gt_boxes.numpy().astype(np.int32)[np.where(_gt_classes!=-1)[0]]    \n    _image = _image.numpy()\n    _gt_dummy_mask = np.zeros_like(_image)\n    _gt_dummy_mask[..., img_class] = cv2.resize(np.expand_dims(_gt_mask, axis=-1), INPUT_SHAPE[:-1])\n    _gt_mask = _gt_dummy_mask\n    \n    \n    plt.figure(figsize=(20,7))\n    \n    plt.subplot(1,3,1)\n    plt.imshow(_image, cmap=\"inferno\")\n    plt.axis(False)\n    plt.title(\"预处理后的原始图像\", fontweight=\"bold\")\n    \n    mask_merged = cv2.addWeighted(_image, 0.55, _gt_mask, 1.25, 0.0)\n    plt.subplot(1,3,2)\n    plt.imshow(mask_merged)\n    plt.axis(False)\n    plt.title(f\"原始图像蒙版  (类别={img_class})\", fontweight=\"bold\")\n    \n    plt.subplot(1,3,3)\n    box_image = np.zeros_like(_image)\n    for box in img_boxes:\n        ymin, xmin, ymax, xmax = box\n        box_image = cv2.rectangle(img=box_image, thickness=1,  pt1=(xmin, ymin), pt2=(xmax, ymax), \n                                  color=[0 if i!=img_class else 255 for i in range(3)])\n     \n    box_merged = cv2.addWeighted(_image, 0.55, box_image, 1.25 if img_class==2 else 0.45, 0.0,)\n    plt.imshow(box_merged)\n    plt.axis(False)\n    plt.title(f\"原始图像边界框  (类别={img_class})\", fontweight=\"bold\")\n\n    plt.tight_layout()\n    plt.show()\n    \ndef plot_pred(_image, _pred_boxes, _pred_scores, _pred_classes, _pred_mask, conf_thresh=0.25, iou_thresh=0.0001):\n    \"\"\"\"\"\"\n    \n    if iou_thresh is not None:\n        _indices, _pred_scores = tf.image.non_max_suppression_with_scores(\n            _pred_boxes, _pred_scores, 800, iou_threshold=iou_thresh,\n            score_threshold=conf_thresh/5, soft_nms_sigma=0.0\n        )\n        _pred_boxes = tf.gather(_pred_boxes, _indices)\n\n    \n    above_thresh_idx = np.where(_pred_scores.numpy()>conf_thresh)[0]\n    if len(above_thresh_idx)==0:\n        print(\"\\n... 超过置信度阈值的预测为空... 采样最多五十个样本 ...\\n\")\n        above_thresh_idx = np.arange(min(50, len(_pred_scores)))\n\n    _image = _image.numpy()\n    _pred_class = int(np.round(_pred_classes.numpy()[above_thresh_idx].mean()))\n\n    _pred_scores = _pred_scores.numpy()[above_thresh_idx]\n    _pred_boxes = _pred_boxes.numpy().astype(np.int32)[above_thresh_idx]\n    _pred_mask = np.where(_pred_mask[..., 1]>_pred_mask[..., 0], 1.0, 0.0)\n    _dummy_mask = np.zeros_like(_image)\n    _dummy_mask[..., _pred_class] = cv2.resize(np.expand_dims(_pred_mask, axis=-1), INPUT_SHAPE[:-1])\n    _pred_mask = _dummy_mask\n    \n    \n    plt.figure(figsize=(20,7))\n    \n    plt.subplot(1,3,1)\n    plt.imshow(_image, cmap=\"inferno\")\n    plt.axis(False)\n    plt.title(\"预处理后的原始图像\", fontweight=\"bold\")\n    \n    mask_merged = cv2.addWeighted(_image, 0.55, _pred_mask, 1.25, 0.0,)\n    plt.subplot(1,3,2)\n    plt.imshow(mask_merged)\n    plt.axis(False)\n    plt.title(f\"预测图像蒙版  (类别={_pred_class})\", fontweight=\"bold\")\n    \n    plt.subplot(1,3,3)\n    box_image = np.zeros_like(_image)\n    for box in _pred_boxes:\n        ymin, xmin, ymax, xmax = box\n        box_image = cv2.rectangle(img=box_image, thickness=1, pt1=(xmin, ymin), pt2=(xmax, ymax), \n                                  color=[0 if i!=_pred_class else 255 for i in range(3)])\n     \n    box_merged = cv2.addWeighted(_image, 0.55, box_image, 1.25 if _pred_class==2 else 0.45, 0.0,)\n    plt.imshow(box_merged)\n    plt.axis(False)\n    plt.title(f\"预测图像边界框  (类别={_pred_class})\", fontweight=\"bold\")\n\n    plt.tight_layout()\n    plt.show()\n\ndef plot_diff(_image, _gt_classes, _gt_boxes, _gt_mask, _pred_boxes, _pred_scores, _pred_classes, _pred_mask, conf_thresh=0.25, iou_thresh=0.0001):\n    \"\"\"\"\"\"\n    \n    if iou_thresh is not None:\n        _indices, _pred_scores = tf.image.non_max_suppression_with_scores(\n            _pred_boxes, _pred_scores, 800, iou_threshold=iou_thresh,\n            score_threshold=conf_thresh/5, soft_nms_sigma=0.0\n        )\n        _pred_boxes = tf.gather(_pred_boxes, _indices)\n    \n    _image = _image.numpy()\n    \n    above_thresh_idx = np.where(_pred_scores.numpy()>conf_thresh)[0]\n    gt_idxs = np.where(_gt_classes!=-1)[0]\n    \n    if len(above_thresh_idx)==0:\n        print(\"\\n... 超过置信度阈值的预测为空... 采样最多五十个样本 ...\\n\")\n        above_thresh_idx = np.arange(min(50, len(_pred_scores)))\n    \n    _img_class = int(_gt_classes.numpy()[0])\n    _pred_class = int(np.round(_pred_classes.numpy()[above_thresh_idx].mean()))\n    \n    img_boxes = _gt_boxes.numpy().astype(np.int32)[gt_idxs]\n    _pred_boxes = _pred_boxes.numpy().astype(np.int32)[above_thresh_idx]\n    \n    _pred_scores = _pred_scores.numpy()[above_thresh_idx]\n    \n    _combo_mask = np.zeros_like(_image)\n    _combo_mask[..., 0] = cv2.resize(np.expand_dims(_gt_mask, axis=-1), INPUT_SHAPE[:-1])        \n    _pred_mask = np.where(_pred_mask[..., -1]>_pred_mask[..., 0], 1.0, 0.0)\n    _combo_mask[..., 1] = cv2.resize(np.expand_dims(_pred_mask, axis=-1), INPUT_SHAPE[:-1])\n    \n    plt.figure(figsize=(20,7))\n    \n    plt.subplot(1,3,1)\n    plt.imshow(_image, cmap=\"inferno\")\n    plt.axis(False)\n    plt.title(\"预处理后的原始图像\", fontweight=\"bold\")\n    \n    mask_merged = cv2.addWeighted(_image, 0.55, _combo_mask, 1.25, 0.0,)\n    plt.subplot(1,3,2)\n    plt.imshow(mask_merged)\n    plt.axis(False)\n    plt.title(f\"混合图像蒙版\\n(红色=GT, 绿色=PRED, 黄色=一致)\", fontweight=\"bold\")\n    \n    plt.subplot(1,3,3)\n    box_image = np.zeros_like(_image)\n    for box in img_boxes:\n        ymin, xmin, ymax, xmax = box\n        box_image = cv2.rectangle(img=box_image, thickness=1, pt1=(xmin, ymin), pt2=(xmax, ymax), \n                                  color=(255,0,0))\n    for box in _pred_boxes:\n        ymin, xmin, ymax, xmax = box\n        box_image = cv2.rectangle(img=box_image, thickness=1, pt1=(xmin, ymin), pt2=(xmax, ymax), \n                                  color=(0,255,0))\n     \n    box_merged = cv2.addWeighted(_image, 0.55, box_image, 1.25, 0.0)\n    plt.imshow(box_merged)\n    plt.axis(False)\n    plt.title(f\"预测图像边界框\\n(红色=GT, 绿色=PRED)\", fontweight=\"bold\")\n\n    plt.tight_layout()\n    plt.show()\n    \n    \ndef compute_iou(labels, y_pred):\n    \"\"\"\n    计算实例标签和预测之间的IoU。\n\n    Args:\n        labels (np array): 标签。\n        y_pred (np array): 预测。\n\n    Returns:\n        np array: IoU 矩阵，大小为 true_objects x pred_objects。\n    \"\"\"\n\n    true_objects = len(np.unique(labels))\n    pred_objects = len(np.unique(y_pred))\n\n    # 计算所有对象之间的交集\n    intersection = np.histogram2d(\n        labels.flatten(), y_pred.flatten(), bins=(true_objects, pred_objects)\n    )[0]\n\n    # 计算面积（用于找到所有对象之间的并集）\n    area_true = np.histogram(labels, bins=true_objects)[0]\n    area_pred = np.histogram(y_pred, bins=pred_objects)[0]\n    area_true = np.expand_dims(area_true, -1)\n    area_pred = np.expand_dims(area_pred, 0)\n\n    # 计算并集\n    union = area_true + area_pred - intersection\n    iou = intersection / union\n    \n    return iou[1:, 1:]  # 排除背景\n\n\ndef precision_at(threshold, iou):\n    \"\"\"\n    计算给定阈值的精度。\n\n    Args:\n        threshold (float): 阈值。\n        iou (np array): IoU 矩阵。\n\n    Returns:\n        int: 真正例的数量，\n        int: 假正例的数量，\n        int: 假负例的数量。\n    \"\"\"\n    matches = iou > threshold\n    true_positives = np.sum(matches, axis=1) == 1  # 正确的对象\n    false_positives = np.sum(matches, axis=0) == 0  # 丢失的对象\n    false_negatives = np.sum(matches, axis=1) == 0  # 额外的对象\n    tp, fp, fn = (\n        np.sum(true_positives),\n        np.sum(false_positives),\n        np.sum(false_negatives),\n    )\n    return tp, fp, fn\n\n\ndef iou_map(truths, preds, verbose=1):\n    \"\"\"\n    计算比赛的度量标准。\n    Masks 包含了每个对象具有一个关联值的分段像素，\n    并且0是背景。\n\n    Args:\n        truths (list of masks): 真实值。\n        preds (list of masks): 预测值。\n        verbose (int, optional): 是否打印信息。默认为 0。\n\n    Returns:\n        float: mAP。\n    \"\"\"\n    ious = [compute_iou(truth, pred) for truth, pred in zip(truths, preds)]\n\n    if verbose:\n        print(\"阈值\\tTP\\tFP\\tFN\\tPrec.\")\n\n    prec = []\n    for t in np.arange(0.5, 1.0, 0.05):\n        tps, fps, fns = 0, 0, 0\n        for iou in ious:\n            tp, fp, fn = precision_at(t, iou)\n            tps += tp\n            fps += fp\n            fns += fn\n\n        p = tps / (tps + fps + fns)\n        prec.append(p)\n\n        if verbose:\n            print(\"{:1.3f}\\t{}\\t{}\\t{}\\t{:1.3f}\".format(t, tps, fps, fns, p))\n\n\n    if verbose:\n        print(\"AP\\t-\\t-\\t-\\t{:1.3f}\".format(np.mean(prec)))\n\n    return np.mean(prec)\n\n\ndef get_pred_instance_mask(_pred_boxes, _pred_scores, _pred_mask, iou_thresh=0.0, conf_thresh=0.25):\n    _indices, _pred_scores = tf.image.non_max_suppression_with_scores(\n        _pred_boxes, _pred_scores, 800, iou_threshold=iou_thresh,\n        score_threshold=conf_thresh/5, soft_nms_sigma=0.0\n    )\n    _pred_boxes = tf.gather(_pred_boxes, _indices)\n    \n    above_thresh_idx = np.where(_pred_scores.numpy()>conf_thresh)[0]\n    if len(above_thresh_idx)==0:\n        above_thresh_idx = np.arange(min(50, len(_pred_scores)))\n\n    _pred_scores = _pred_scores.numpy()[above_thresh_idx]\n    _pred_boxes = _pred_boxes.numpy().astype(np.int32)[above_thresh_idx]\n    _pred_mask = cv2.resize(_pred_mask, INPUT_SHAPE[:-1], interpolation=cv2.INTER_NEAREST)\n    _pred_mask = np.where(_pred_mask[..., 1]>_pred_mask[..., 0], 1.0, 0.0)\n    _instance_mask = np.zeros_like(_pred_mask)\n    for i, _box in enumerate(_pred_boxes):\n        _instance_mask[_box[0]:_box[2], _box[1]:_box[3]] = (i+1)*_pred_mask[_box[0]:_box[2], _box[1]:_box[3]]\n    _instance_mask = cv2.resize(_instance_mask, IMAGE_SHAPE[-2::-1], interpolation=cv2.INTER_NEAREST)\n    return _instance_mask","metadata":{"execution":{"iopub.status.busy":"2023-12-14T10:35:01.201932Z","iopub.execute_input":"2023-12-14T10:35:01.202466Z","iopub.status.idle":"2023-12-14T10:35:01.264528Z","shell.execute_reply.started":"2023-12-14T10:35:01.202419Z","shell.execute_reply":"2023-12-14T10:35:01.263720Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.7 在验证数据集上的交并比 (IOU)","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"## 5.8 在验证数据集上的混淆矩阵","metadata":{}},{"cell_type":"code","source":"","metadata":{},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**训练数据的第一批次**","metadata":{}},{"cell_type":"code","source":"# 我们希望... 待定 --> 封装成函数\n#    GT（Ground Truth）：\n#        - 边界框\n#        - 置信度分数\n#        - 分割蒙版\n#    PRED（Prediction）：\n#        - 边界框\n#        - 置信度分数\n#        - 实例类别\n#        - 分割蒙版\nfor _image_batch, _label_batch in train_dl.take(1):\n    gt_mask = _label_batch[\"image_masks\"][:, 0]\n    gt_boxes = _label_batch[\"groundtruth_data\"][..., :4]\n    gt_is_crowds = _label_batch[\"groundtruth_data\"][..., 4]\n    gt_areas = _label_batch[\"groundtruth_data\"][..., 5]\n    gt_classes = _label_batch[\"groundtruth_data\"][..., 6]\n    \n    pred_classes, pred_boxes, pred_mask = model(_image_batch, training=False)\n    pred_boxes, pred_scores, pred_classes, valid_len = postprocess.postprocess_global(config, pred_classes, pred_boxes)\n    gt_instance_masks, pred_instance_masks = [], []\n    for i in range(BATCH_SIZE):\n        print(\"\\n\\n... 原始显示图像 ...\\n\")\n        _img, _mask = get_img_and_mask(**train_df.iloc[int(_label_batch[\"source_ids\"][i])][[\"img_path\", \"annotation\", \"width\", \"height\"]])\n        plot_img_and_mask(_img, _mask)\n        gt_instance_masks.append(_mask)\n\n        print(\"\\n... 真实值图 ...\\n\")\n        plot_gt(_image_batch[i], gt_classes[i], gt_boxes[i], gt_mask[i])\n\n        print(f\"\\n... 预测图 (NMS={'是' if i<4 else '否'}) ...\\n\")\n        plot_pred(_image_batch[i], pred_boxes[i], pred_scores[i], pred_classes[i], pred_mask[i], iou_thresh=0.0 if i<4 else None)\n\n        print(f\"\\n... 真实值与预测值图 (NMS={'是' if i<4 else '否'}) ...\\n\")\n        plot_diff(_image_batch[i], gt_classes[i], gt_boxes[i], gt_mask[i], pred_boxes[i], pred_scores[i], pred_classes[i], pred_mask[i], iou_thresh=0.0 if i<4 else None)\n        \n        pred_instance_masks.append(get_pred_instance_mask(pred_boxes[i], pred_scores[i], pred_mask[i].numpy(), iou_thresh=0.0, conf_thresh=0.1))\n        \n        print(\"\\n\\n\\n\\n\")\n        print(\"-\"*50)\n        print(\"\\n\\n\")\n        \n    print(\"\\n批次评估:\\n\")\n    iou_map(gt_instance_masks, pred_instance_masks)\n","metadata":{"execution":{"iopub.status.busy":"2023-12-14T10:35:01.265616Z","iopub.execute_input":"2023-12-14T10:35:01.265826Z"},"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"**验证数据的第一批次**","metadata":{}},{"cell_type":"code","source":"for _image_batch, _label_batch in val_dl.take(1):\n    gt_mask = _label_batch[\"image_masks\"][:, 0]\n    gt_boxes = _label_batch[\"groundtruth_data\"][..., :4]\n    gt_is_crowds = _label_batch[\"groundtruth_data\"][..., 4]\n    gt_areas = _label_batch[\"groundtruth_data\"][..., 5]\n    gt_classes = _label_batch[\"groundtruth_data\"][..., 6]\n    \n    pred_classes, pred_boxes, pred_mask = model(_image_batch, training=False)\n    pred_boxes, pred_scores, pred_classes, valid_len = postprocess.postprocess_global(config, pred_classes, pred_boxes)\n    gt_instance_masks, pred_instance_masks = [], []\n        \n    for i in range(BATCH_SIZE):\n        print(\"\\n\\n... 原始图像显示图 ...\\n\")\n        _img, _mask = get_img_and_mask(**train_df.iloc[int(_label_batch[\"source_ids\"][i])][[\"img_path\", \"annotation\", \"width\", \"height\"]])\n        plot_img_and_mask(_img, _mask)\n        gt_instance_masks.append(_mask)\n\n        print(\"\\n... 真实值图 ...\\n\")\n        plot_gt(_image_batch[i], gt_classes[i], gt_boxes[i], gt_mask[i])\n\n        print(f\"\\n... 预测图 (NMS={'是' if i<4 else '否'}) ...\\n\")\n        plot_pred(_image_batch[i], pred_boxes[i], pred_scores[i], pred_classes[i], pred_mask[i], iou_thresh=0.0 if i<4 else None)\n\n        print(f\"\\n... 真实值与预测值图 (NMS={'是' if i<4 else '否'}) ...\\n\")\n        plot_diff(_image_batch[i], gt_classes[i], gt_boxes[i], gt_mask[i], pred_boxes[i], pred_scores[i], pred_classes[i], pred_mask[i], iou_thresh=0.0 if i<4 else None)\n        \n        pred_instance_masks.append(get_pred_instance_mask(pred_boxes[i], pred_scores[i], pred_mask[i].numpy(), iou_thresh=0.0, conf_thresh=0.1))\n        \n        print(\"\\n\\n\\n\\n\")\n        print(\"-\"*50)\n        print(\"\\n\\n\")\n        \n    print(\"\\n批次评估:\\n\")\n    iou_map(gt_instance_masks, pred_instance_masks)\n","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"pred_instance_masks = []\nfor _image_batch, _label_batch in test_dl:\n    pred_classes, pred_boxes, pred_mask = model(_image_batch, training=False)\n    pred_boxes, pred_scores, pred_classes, valid_len = postprocess.postprocess_global(config, pred_classes, pred_boxes)\n    pred_instance_masks.append(get_pred_instance_mask(pred_boxes[0], pred_scores[0], pred_mask[0].numpy(), iou_thresh=0.0, conf_thresh=0.075))","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"markdown","source":"# 6 Submission\n---\n\nWIP","metadata":{}},{"cell_type":"code","source":"submission_dfs = []\nfor i, _id in enumerate(ss_df.id.to_list()):\n    tmp_df = pd.DataFrame([_id,]*int(pred_instance_masks[i].max()), columns=[\"id\"])\n    tmp_df[\"predicted\"] = [rle_encode(np.where(pred_instance_masks[i]==j, 1.0, 0.0)) for j in range(1, int(pred_instance_masks[i].max()+1))]\n    submission_dfs.append(tmp_df)\nsubmission_df = pd.concat(submission_dfs).reset_index(drop=True)\nsubmission_df.to_csv(\"submission.csv\", index=False)","metadata":{"trusted":true},"execution_count":null,"outputs":[]},{"cell_type":"code","source":"import numpy as np\nimport pandas as pd\n\npredict = pd.read_csv(\"submission.csv\")\npredict.head(10)","metadata":{"trusted":true},"execution_count":null,"outputs":[]}]}