首页 / 资讯中心 / 文章详情

TensorFlow全切片图像癌细胞检测实战:WSI处理、ResNet改造与临床部署

TensorFlow全切片图像癌细胞检测实战:WSI处理、ResNet改造与临床部署 ★ FEATURED ARTICLE
简介本资源是一份面向AI医疗开发者、病理学研究者及医学图像分析工程师的实战技术文档聚焦基于TensorFlow构建全切片图像WSI癌细胞检测系统的核心流程与工程实践。文档覆盖数字病理学背景、TensorFlow环境配置、WSI数据预处理含标注、增强、归一化、ResNet模型选型与调优、分层系统架构设计以及乳腺癌、肺癌、结直肠癌三大临床场景的端到端案例验证兼具理论深度与落地指导性。资源为单个PDF文件共32页大小1.87MB内容结构完整目录清晰呈现从引言、基础框架、需求设计、模型训练到部署评估的十一章逻辑主线便于按模块精读与复现。目前已有98人学习下载读者可直接获取可复用的数据处理管道构建方法、模型训练监控策略、系统集成要点及临床效果评估指标体系显著降低医学AI项目从算法到应用的转化门槛。1. 这不是又一个“YOLO跑个COCO”的DemoTensorFlow数字病理学实录专治全切片图像癌细胞检测的“内存爆炸、漏检率高、部署即崩”三大玄学病你手头有一张 100,000×80,000 像素、2–5 GB 大小的 SVS 全切片图像WSI想用 TensorFlow 检出散落在几十亿像素里的几十个癌细胞核——结果刚cv2.imread()就 OOMtf.data.Dataset.from_generator()卡死在prefetch模型训完一推理GPU 显存瞬间飙到 98%但输出热图全是噪点。这不是理论跑通就能交差的课程设计这是真实病理实验室里每天发生的“算力窒息”。这份《TensorFlow 数字病理学全切片图像癌细胞检测系统开发实录》PDF不是概念稿是作者在三甲医院病理科驻场 4 个月、踩过 17 次显存溢出、重写 3 轮数据管道、把 QuPath 标注导出逻辑硬抠进 Python 的血泪复盘。它不讲“深度学习有多酷”只解决三个硬问题怎么把一张 4GB 的 SVS 文件切成可喂给 GPU 的 patch 而不炸内存怎么让 ResNet 在 224×224 输入下依然抓住 5μm 级别的核仁畸变怎么把训练好的模型塞进 Docker 容器让病理科医生点开浏览器就能上传、检测、看报告而不是找你调三天环境适合正在做医学图像落地的工程师、想用 AI 辅助诊断的病理医生以及被导师催着“把 WSI 拿 TensorFlow 跑通”的研二学生——如果你的痛点是“数据太大、模型太糙、上线太难”这篇就是你的后悔药。1.1 为什么传统 CV 流水线在这儿全失效常规图像处理流程OpenCV 读图 → resize → augment → feed to model在 WSI 面前直接崩溃。原因很实在一张中等分辨率 WSI如 40× 扫描原始尺寸常达 80,000×60,000 像素按 RGB 三通道 uint8 存储内存占用 80,000 × 60,000 × 3 ≈ 14.4 GB。而主流训练机显存多为 24GBRTX 4090或 40GBA100连整图加载都做不到更别说做旋转/CLAHE 增强。有人会说“那我用 OpenSlideread_region切 patch 啊”但问题来了read_region((x,y), level, (w,h))返回的是 RGBA PIL Imagenp.array(pil_img)会触发全图解码若 level0最高分辨率一次切 224×224 patch 仍需解码整个金字塔底层IO 和内存压力巨大。我们实测过用默认openslide参数读取 100 个 patch平均耗时 3.2 秒/patchCPU 占用 92%根本没法进tf.datapipeline。这不是代码写得不够 Pythonic是 WSI 的数据结构金字塔 压缩编码和深度学习框架的 tensor 流水线存在天然代沟。这份实录的核心价值就是用 6 个可抄作业的代码块把这道沟填平。1.2 它不是教你怎么搭 ResNet而是告诉你 ResNet 在 WSI 上为什么必须砍掉最后两层文档里反复出现的 “ResNet50 Global Average Pooling” 不是拍脑袋选的。我们在对比了 DenseNet121、EfficientNet-B3、ViT-Base 在 Camelyon16 数据集上的表现后发现当输入 patch 尺寸固定为 224×224 时ResNet50 的 top-1 准确率比 DenseNet121 高 2.3%但参数量少 37%而 ViT 在小 patch 上因缺乏归纳偏置F1-score 反而低 5.8%。更关键的是 ResNet 的 stage3 输出特征图尺寸为 28×28恰好能与 WSI 中癌细胞簇的典型空间尺度100–500 μm对应 28–140 pixel 40×匹配便于后续的 attention 或 mask head 定位。但直接套用 Keras 官方ResNet50(weightsimagenet)会翻车——它的GlobalAveragePooling2D层强制把 7×7 特征图压成 2048-dim 向量丢失全部空间信息导致无法生成像素级热图。实录第 6 章明确要求必须删掉model.layers[-2:]即 GAP 和 Dense接一个Conv2D(1, 1)sigmoid让模型输出与输入 patch 等尺寸的 224×224 概率图。这个改动看似微小却是从“分类器”蜕变为“检测器”的分水岭。没有这一步你永远只能知道“这张图有癌”而不知道“癌在哪”。1.3 它解决的不是“能不能跑”而是“敢不敢让医生用”很多开源项目止步于 Jupyter Notebook 里model.predict()出来一个 0.92 的概率值就宣告成功。但在病理科这毫无意义。医生需要的是① 上传一张 SVS 文件30 秒内返回带红色框标注的缩略图② 点击任意框弹出该区域的 HE 染色细节、预测置信度、与邻近正常组织的对比直方图③ 所有结果存入 SQLite支持按患者 ID、日期、癌种筛选。实录第 8 章的 Flask 集成方案不是简单app.route(/predict)而是用threading.Lock()控制 OpenSlide 实例复用避免多请求并发打开同一 SVS 导致文件句柄泄漏用multiprocessing.Queue异步处理 patch 推理防止长耗时阻塞 HTTP 请求并内置了 DICOM-SR 兼容的结构化报告生成器。这意味着你照着文档走完第 8 章得到的不是一个 demo而是一个能放进医院内网、经得起临床验证的最小可行产品MVP。它不承诺替代病理医生但能确保医生点开网页的那一刻看到的是可操作、可追溯、可审计的结果而不是一行tensor([0.921], dtypefloat32)。2. 把 4GB SVS 文件变成 GPU 友好型 Patch 流OpenSlide PyVips tf.data 的三段式流水线WSI 预处理的核心矛盾在于既要保留原始信息的完整性不能盲目 resize 丢细节又要满足 GPU 显存的物理限制不能全图加载。常见错误是试图用cv2.imread或PIL.Image.open直接读 SVS——它们根本不支持多级金字塔会直接报错或只读第一层模糊图。正确解法是分三步走先用 OpenSlide 定位高分辨率区域再用 PyVips 高效解码最后用tf.data构建零拷贝 pipeline。这三者不是简单串联而是有严格依赖关系的协同。2.1 为什么必须用 OpenSlide 定位而不能靠“随机采样”全切片图像中癌细胞只占极小区域常 0.1%其余是正常组织、坏死区、背景空白。若对整图做均匀网格采样如每 224×224 像素取一个 patch99% 的 patch 是负样本训练效率极低且模型易过拟合背景纹理。实录第 5.4.1 节强调必须基于组织分割tissue segmentation做前景采样。OpenSlide 提供了slide.associated_images和slide.read_region()的精确坐标控制能力可先用低分辨率 level如 level6尺寸约 1000×800快速生成组织掩膜tissue mask再反推高分辨率 levellevel0的坐标范围。以下是生成 tissue mask 的核心代码import openslide import numpy as np from scipy import ndimage def get_tissue_mask(slide_path, level6, threshold0.8): 生成组织掩膜在指定 level 读取缩略图通过 HSV 颜色空间分离组织区域 Args: slide_path: SVS 文件路径 level: OpenSlide pyramid level (0最高分辨率) threshold: 二值化阈值控制组织区域敏感度 Returns: mask: 二值掩膜 (H, W)True 表示组织区域 slide openslide.OpenSlide(slide_path) # 获取 level6 的尺寸和图像 dims slide.level_dimensions[level] img np.array(slide.read_region((0,0), level, dims)) slide.close() # 转 HSV 并提取饱和度通道组织区域饱和度更高 hsv cv2.cvtColor(img[:,:,:3], cv2.COLOR_RGB2HSV) s_channel hsv[:,:,1] # 自适应阈值 形态学闭运算去噪 blurred cv2.GaussianBlur(s_channel, (5,5), 0) _, mask cv2.threshold(blurred, 0, 255, cv2.THRESH_BINARY cv2.THRESH_OTSU) kernel np.ones((5,5), np.uint8) mask cv2.morphologyEx(mask, cv2.MORPH_CLOSE, kernel) # 填充小孔洞 mask ndimage.binary_fill_holes(mask).astype(np.uint8) return mask # 使用示例生成 Camelyon16 训练集的 tissue mask mask get_tissue_mask(train_001.svs, level6) print(fTissue mask shape: {mask.shape}, tissue ratio: {mask.sum() / mask.size:.3f})提示get_tissue_mask中level6是经验值。Camelyon16 数据集总共有 10 级 pyramidlevel6 对应约 1/64 缩放既能保证组织轮廓清晰又能在 1 秒内完成计算。若你的扫描仪分辨率更高如 60×可将level设为 7 或 8。threshold0.8不是固定值需根据染色深浅调整——HE 染色深的切片用 0.7浅的用 0.85。2.2 为什么 PyVips 比 OpenSlideread_region快 3.7 倍OpenSlide 的read_region是单线程、同步阻塞的每次调用都要重新解析 TIFF 标签、定位 tile 偏移、解码 JPEG2000。而 PyVips 基于 libvips 的 demand-driven processing它不立即解码而是构建一个计算图computation graph只有在.write_to_memory()或.numpy()时才触发实际解码且支持多线程 tile 解码。实录第 5.2.1 节给出的 benchmark在相同硬件上读取 100 个 224×224 patchOpenSlide 平均耗时 3.2s/patchPyVips 仅 0.86s/patch。关键在于 PyVips 的extract_area方法可直接从压缩流中裁剪无需全图解码。以下是高效 patch 读取器import pyvips import numpy as np class WsiPatchReader: def __init__(self, svs_path, level0): 初始化 WSI patch 读取器 Args: svs_path: SVS 文件路径 level: 目标金字塔层级0最高分辨率 self.image pyvips.Image.new_from_file(svs_path, accesssequential) self.level level # 获取 level0 的原始尺寸 self.width, self.height self.image.width, self.image.height def read_patch(self, x, y, width, height): 从指定坐标读取 patch返回 numpy array (H,W,3) 注意x,y 是 level0 坐标自动按比例缩放到当前 level # 计算当前 level 下的实际坐标OpenSlide 的 level_downsample 是近似值PyVips 更准 downsample 2 ** self.level x_level int(x / downsample) y_level int(y / downsample) w_level int(width / downsample) h_level int(height / downsample) try: # PyVips extract_area 是零拷贝裁剪极快 patch self.image.extract_area(x_level, y_level, w_level, h_level) # 转 numpy注意顺序pyvips 默认 BGR需转 RGB patch_np patch.numpy() if len(patch_np.shape) 3 and patch_np.shape[2] 3: patch_np patch_np[:, :, ::-1] # BGR - RGB return patch_np except Exception as e: # 若坐标越界返回黑图避免 pipeline 中断 print(fWarning: patch at ({x},{y}) out of bounds, returning black) return np.zeros((height, width, 3), dtypenp.uint8) # 使用示例从 tissue mask 中随机采样 10 个组织区域 patch reader WsiPatchReader(train_001.svs, level0) mask get_tissue_mask(train_001.svs, level6) # 将 mask 坐标映射回 level0 coords np.argwhere(mask) * (2**6) # level6 到 level0 的缩放因子 samples coords[np.random.choice(len(coords), 10, replaceFalse)] for i, (y, x) in enumerate(samples): patch reader.read_patch(x, y, 224, 224) print(fPatch {i} shape: {patch.shape}) # 应输出 (224, 224, 3)参数说明accesssequential告诉 PyVips 按顺序读取优化 IOextract_area的(x,y,w,h)是 level-specific 坐标必须按2**level缩放patch.numpy()是唯一触发解码的操作且返回的是 C-contiguous array可直接送入tf.data。若遇到libvips报错VipsJpeg: Invalid JPEG data, 说明 SVS 文件含非标准 JPEG2000需在new_from_file中加fail-on-errorFalse并预处理。2.3 构建零拷贝tf.dataPipeline绕过np.array的内存地狱最致命的坑是把WsiPatchReader.read_patch()返回的np.ndarray直接塞进tf.data.Dataset.from_tensor_slices()。这会导致每个 patch 被复制 3 次Python list → tf.Tensor → GPU memory1000 个 patch 就吃掉 15GB 内存。实录第 5.2.2 节的解法是用tf.data.Dataset.from_generator()tf.py_function让 patch 生成和 tensor 转换在同一个内存空间完成。关键技巧是tf.py_function的Tout参数指定输出类型tf.numpy_function已废弃必须用py_function。以下是完整 pipelineimport tensorflow as tf import numpy as np def _parse_patch_function(x, y, width, height, svs_path, level0): tf.py_function 包装的 patch 读取函数 注意此函数在 graph mode 下运行不能用 print错误用 tf.print try: # 复用上面的 WsiPatchReader但需在函数内初始化因多进程 reader WsiPatchReader(svs_path, levellevel) patch reader.read_patch(x, y, width, height) # 确保尺寸正确否则 tf.py_function 会报错 if patch.shape ! (height, width, 3): patch np.zeros((height, width, 3), dtypenp.uint8) # 归一化到 [0,1] float32为模型输入做准备 patch patch.astype(np.float32) / 255.0 return patch except Exception as e: tf.print(fError in _parse_patch_function: {e}) return np.zeros((height, width, 3), dtypenp.float32) def create_patch_dataset(svs_path, coords, width224, height224, level0, batch_size16): 创建 WSI patch dataset Args: svs_path: SVS 文件路径 coords: 组织区域坐标列表格式为 [(x1,y1), (x2,y2), ...] width/height: patch 尺寸 level: pyramid level batch_size: batch size Returns: tf.data.Dataset: 可直接用于 model.fit() 的 dataset # 将 coords 转为 tf.Tensor x_coords tf.constant([c[0] for c in coords], dtypetf.int32) y_coords tf.constant([c[1] for c in coords], dtypetf.int32) # 创建 dataset每个元素是 (x, y) 坐标对 dataset tf.data.Dataset.from_tensor_slices((x_coords, y_coords)) # 使用 py_function 读取 patch注意 Tout 必须匹配返回类型 dataset dataset.map( lambda x, y: tf.py_function( func_parse_patch_function, inp[x, y, width, height, svs_path, level], Touttf.float32 ), num_parallel_callstf.data.AUTOTUNE ) # 设置输出形状否则 model 无法 infer shape dataset dataset.map( lambda x: tf.reshape(x, (height, width, 3)), num_parallel_callstf.data.AUTOTUNE ) # 批处理、预取、缓存若内存充足 dataset dataset.batch(batch_size) dataset dataset.prefetch(tf.data.AUTOTUNE) # dataset dataset.cache() # 仅当所有 patch 能放进内存时启用 return dataset # 使用示例为一张 SVS 创建训练 dataset coords [(1000, 2000), (5000, 3000), (8000, 1000)] # 从 tissue mask 得到的真实坐标 train_ds create_patch_dataset( svs_pathtrain_001.svs, coordscoords, width224, height224, level0, batch_size8 ) # 验证 pipeline for batch in train_ds.take(1): print(fBatch shape: {batch.shape}) # 应输出 (8, 224, 224, 3) break逻辑说明tf.py_function是 bridge它允许你在 eager mode 下执行任意 Python 代码如调用WsiPatchReader并将结果无缝转为tf.Tensor。num_parallel_callstf.data.AUTOTUNE让 TensorFlow 自动分配线程数实测在 8 核 CPU 上prefetch后吞吐量提升 4.2 倍。tf.reshape是必须的因为py_function返回的 tensor shape 是NoneKeras 模型需要明确的(B,H,W,C)。若你发现dataset运行缓慢90% 的概率是WsiPatchReader.__init__()中pyvips.Image.new_from_file()被重复调用——解决方案是将reader提升为全局变量或使用functools.lru_cache缓存。3. ResNet50 改装指南从分类器到检测器的四步手术刀式改造直接tf.keras.applications.ResNet50(weightsimagenet)在 WSI 上效果差不是因为 ResNet 不行而是它被设计为“整图分类”而 WSI 检测需要“局部定位”。实录第 6.1.3 节提出的“四步改造法”是经过 Camelyon16 和 BreakHis 数据集验证的最小改动方案。它不重写 backbone只动连接层确保迁移学习的有效性。3.1 第一步砍掉 GlobalAveragePooling2D 和顶层 Dense接 Conv2D(1,1)这是最核心的改动。原 ResNet50 的GlobalAveragePooling2D将Conv2D输出的(7,7,2048)特征图压成(2048,)向量彻底丢失空间位置信息。而癌细胞检测需要知道“癌在哪”必须保留(H,W,C)的空间维度。实录要求取base_model.get_layer(conv5_block3_out).output即 stage4 的输出尺寸为(7,7,2048)在其后接Conv2D(256, 1)降维再Conv2D(1, 1)输出(7,7,1)的 logits 图最后UpSampling2D(size32)插值回(224,224,1)。以下是可直接运行的模型构建代码import tensorflow as tf from tensorflow.keras import layers, models from tensorflow.keras.applications import ResNet50 def build_wsi_detection_model(input_shape(224, 224, 3), num_classes1): 构建 WSI 癌细胞检测模型ResNet50 backbone detection head Args: input_shape: 输入 patch 尺寸 num_classes: 分类数此处为 1癌/非癌 Returns: model: Keras Model # 加载预训练 ResNet50不包括顶层 base_model ResNet50( weightsimagenet, include_topFalse, input_shapeinput_shape ) # 冻结 backbone 前 100 层保留低层纹理特征微调高层语义 base_model.trainable True for layer in base_model.layers[:100]: layer.trainable False # 获取 stage4 输出conv5_block3_out尺寸 (7,7,2048) x base_model.get_layer(conv5_block3_out).output # Detection head1x1 conv 降维 upsample 回原尺寸 x layers.Conv2D(256, 1, activationrelu, namedetection_conv1)(x) x layers.Dropout(0.3)(x) # 防止过拟合小 patch 数据 x layers.Conv2D(num_classes, 1, namedetection_conv2)(x) # (7,7,1) # 上采样到 (224,224,1)使用 bilinear 插值比 nearest 更平滑 x layers.UpSampling2D(size32, interpolationbilinear, nameupsample)(x) # 输出层sigmoid 激活输出每个像素的癌概率 outputs layers.Activation(sigmoid, nameprediction)(x) model models.Model(inputsbase_model.input, outputsoutputs) return model # 构建模型并查看结构 model build_wsi_detection_model() model.summary()参数说明include_topFalse是必须的否则会自动加 GAPlayer.trainable False冻结前 100 层实测在 Camelyon16 上比全量微调收敛快 2.1 倍且 val_loss 更稳定UpSampling2D(size32)的 32 来自224/732是精确上采样不是近似interpolationbilinear比nearest更适合医学图像减少棋盘效应。若你用的是 ResNet101size应为224/732因 ResNet101 的 stage4 输出也是 7×7。3.2 第二步用 Dice Loss 替代 Binary Crossentropy专治类别极度不平衡WSI patch 中癌细胞像素占比常低于 0.01%如 224×22450176 像素中只有 10 个癌细胞核Binary CrossentropyBCE会因大量负样本主导梯度导致模型学会“全预测为负”而 loss 仍很低。实录第 7.2.2 节推荐 Dice Loss它直接优化预测 mask 与真值 mask 的重叠度IoU对正样本更敏感。以下是 Dice Loss 的 TensorFlow 实现含 smooth 项防除零import tensorflow as tf def dice_loss(y_true, y_pred, smooth1e-6): Dice loss for binary segmentation Args: y_true: ground truth mask (B, H, W, 1) y_pred: predicted mask (B, H, W, 1) smooth: smoothing factor to avoid division by zero Returns: loss: scalar dice loss # Flatten tensors to 2D: (B*H*W, 1) y_true_f tf.reshape(y_true, [-1]) y_pred_f tf.reshape(y_pred, [-1]) # 计算交集和并集 intersection tf.reduce_sum(y_true_f * y_pred_f) union tf.reduce_sum(y_true_f) tf.reduce_sum(y_pred_f) # Dice coefficient dice (2. * intersection smooth) / (union smooth) # Dice loss 1 - Dice coefficient return 1. - dice # 编译模型使用 dice_loss model.compile( optimizertf.keras.optimizers.Adam(learning_rate1e-4), lossdice_loss, metrics[accuracy] ) # 验证 loss 计算 import numpy as np y_true np.zeros((1, 224, 224, 1)) y_true[0, 100, 100, 0] 1.0 # 一个癌细胞像素 y_pred np.zeros((1, 224, 224, 1)) y_pred[0, 100, 100, 0] 0.9 # 模型预测对了 loss_val dice_loss(y_true, y_pred).numpy() print(fDice loss for perfect prediction: {loss_val:.6f}) # 应接近 0.000001为什么不用 Focal Loss实录第 7.4.2 节做了对比实验在 Camelyon16 上Focal Lossγ2的最终 Dice Score 为 0.782而 Dice Loss 为 0.815。原因是 Focal Loss 通过降低易分样本权重来聚焦难样本但 WSI 中“难样本”往往是组织边缘或染色不均区域这些并非癌细胞反而干扰定位。Dice Loss 直接优化空间重叠更契合检测任务目标。3.3 第三步添加 Class Activation MappingCAM可视化让医生信你的模型医生不会相信一个黑匣子输出的 0.92 概率。他们需要看到模型“看”到了什么。实录第 6.4.2 节要求在训练好的模型上用 Grad-CAM 生成热图叠加在原始 patch 上直观显示模型关注的癌细胞区域。这不仅是调试工具更是临床信任的基石。以下是 Grad-CAM 实现兼容 TensorFlow 2.ximport numpy as np import cv2 import matplotlib.pyplot as plt from tensorflow.keras import backend as K def make_gradcam_heatmap(img_array, model, last_conv_layer_name, pred_indexNone): 生成 Grad-CAM 热图 Args: img_array: 输入 patch (1, H, W, 3)已归一化 model: 训练好的模型 last_conv_layer_name: 最后一个卷积层名如 conv5_block3_out pred_index: 预测类别索引此处为 0因是单类 Returns: heatmap: 热图 (H, W) # 创建模型输入为 img_array输出为 last_conv_layer 和 model.output grad_model tf.keras.models.Model( [model.inputs], [model.get_layer(last_conv_layer_name).output, model.output] ) with tf.GradientTape() as tape: conv_outputs, predictions grad_model(img_array) if pred_index is None: pred_index tf.argmax(predictions[0]) loss predictions[:, pred_index] # 计算梯度 grads tape.gradient(loss, conv_outputs) pooled_grads tf.reduce_mean(grads, axis(0, 1, 2)) # 权重乘以特征图 conv_outputs conv_outputs[0] heatmap conv_outputs pooled_grads[..., tf.newaxis] heatmap tf.squeeze(heatmap) # ReLU 激活归一化 heatmap tf.maximum(heatmap, 0) heatmap / tf.reduce_max(heatmap) K.epsilon() return heatmap.numpy() # 使用示例为一个 patch 生成热图 test_patch next(iter(train_ds.take(1)))[0] # 取一个 batch 的第一个 patch test_patch_expanded np.expand_dims(test_patch[0].numpy(), axis0) # (1,224,224,3) heatmap make_gradcam_heatmap( test_patch_expanded, model, last_conv_layer_nameconv5_block3_out ) # 可视化 plt.figure(figsize(10, 4)) plt.subplot(1, 2, 1) plt.imshow(test_patch[0]) plt.title(Original Patch) plt.axis(off) plt.subplot(1, 2, 2) plt.imshow(test_patch[0]) plt.imshow(heatmap, cmapjet, alpha0.4) plt.title(Grad-CAM Heatmap) plt.axis(off) plt.show()关键点last_conv_layer_name必须是conv5_block3_out因为它是 stage4 的输出感受野最大约 224px能覆盖整个癌细胞簇tf.reduce_mean(grads, axis(0,1,2))计算全局梯度权重是 Grad-CAM 的核心alpha0.4控制热图透明度确保医生能看清底层组织结构。若热图全黑检查model.trainableTrue是否生效或grad_model是否正确指向了输出层。3.4 第四步集成 Non-Maximum SuppressionNMS从热图到检测框模型输出的是(224,224,1)的概率热图但医生需要的是矩形框bounding box。实录第 8.2.2 节要求在推理时对热图做阈值分割如 0.5连通域分析cv2.connectedComponents再对每个连通域拟合最小外接矩形最后用 NMS 合并重叠框。以下是端到端推理函数import cv2 import numpy as np def predict_bboxes(model, patch, conf_threshold0.5, nms_threshold0.3): 从 patch 预测癌细胞 bounding boxes Args: model: 训练好的模型 patch: 输入 patch (H,W,3)uint8 conf_threshold: 热图阈值 nms_threshold: NMS IoU 阈值 Returns: bboxes: list of [x1, y1, x2, y2, score] # 预处理归一化、扩维 patch_norm patch.astype(np.float32) / 255.0 patch_input np.expand_dims(patch_norm, axis0) # (1,H,W,3) # 模型预测热图 heatmap model.predict(patch_input)[0, :, :, 0] # (H,W) # 阈值分割 binary_map (heatmap conf_threshold).astype(np.uint8) # 连通域分析 num_labels, labels cv2.connectedComponents(binary_map) bboxes [] for i in range(1, num_labels): # 跳过背景 label 0 # 获取连通域像素坐标 y_coords, x_coords np.where(labels i) if len(x_coords) 10: # 过滤太小的噪声区域 continue # 拟合最小外接矩形 x1, y1, w, h cv2.boundingRect(np.column_stack((x_coords, y_coords))) x2, y2 x1 w, y1 h # 计算该区域的平均置信度作为 score score np.mean(heatmap[y1:y2, x1:x2]) bboxes.append([x1, y1, x2, y2, score]) # NMS 合并重叠框 if len(bboxes) 0: return [] bboxes np.array(bboxes) scores bboxes[:, 4] x1 bboxes[:, 0] y1 bboxes[:, 1] x2 bboxes[:, p a hrefhttps://download.csdn.net/download/ashyyyy/90199711 stylecolor:#ec7500;font-size:14px; 本文还有配套的精品资源点击获取 /a img altmenu-r.4af5f7ec.gif srchttps://csdnimg.cn/release/wenkucmsfe/public/img/menu-r.4af5f7ec.gif stylewidth:16px;margin-left:4px;vertical-align:text-bottom;cursor:text; /p
阅读完成 · 觉得有帮助?
咨询建站