十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

SwinUNETR细胞核分割实战:MoNuSeg数据管线与训练推理

SwinUNETR细胞核分割实战:MoNuSeg数据管线与训练推理 简介面向具备PyTorch基础的医学图像分析与深度学习研究者这份资源以Swin-Transformer为核心提供了一套在MoNuSeg数据集上端到端训练与推理的细胞核分割方案。内容覆盖数据自动下载与解压、512×512滑窗切片、albumentations增强、SwinUNETR模型搭建以及Focal-Tversky损失函数的实现细节并支持对整张全切片图像进行滑窗预测。文档同步整理了训练、验证、推理的完整流程与目录结构实测Dice分数达0.923可直接作为Kaggle/科研竞赛的分割技术baseline。资源为docx格式共1个文件压缩包整体仅35KB轻量易读便于快速查阅与本地调试。已有136人浏览学习适合需要复现SOTA模型或深入掌握医学图像分割关键技术点的研究人员与开发者。1. 从 MoNuSeg 到 Swin-Transformer为什么细胞核分割要换掉纯卷积细胞核分割是病理图像分析里的小目标密集任务。MoNuSeg 仅有 30 张训练图、14 张测试图每张都是高倍 HE 染色全景图核密度高、边界重叠、染色不均。传统 U-Net 收敛快但边界细节易被下采样抹平紧贴的核常被预测成一个连通域。Swin-Transformer 在窗口内做自注意力shifted window 让信息跨窗口流动上下文建模强于同参数卷积主干。这套项目用 MONAI 的 SwinUNETR 接 MoNuSeg配 Focal-Tversky 损失压类别不平衡端到端跑通下载、训练、推理实测 Dice 0.923、IoU 0.864、AJI 0.816。适合想在医学分割里验证 Transformer 收益的工程师也适合拿 MoNuSeg 当 baseline 的竞赛选手。下文按数据、模型、训练、推理四部分拆解。2. 数据管线MoNuSeg 下载、滑窗切片与语义掩码清洗2.1 下载脚本与目录约定MoNuSeg 官方数据托管在 grand-challenge 上训练集和测试集是两个独立 zip 包。仓库里的 download.sh 把下载和解压合并成一条命令避免每次手工点网页。需要注意 URL 里的空格被转义成%20在 bash 中必须用引号包住整个链接否则路径会被拆成两段导致 wget 失败。mkdir -p data/raw wget https://monuseg.grand-challenge.org/static/dataset/MoNuSeg%20Training%20Data.zip -O data/raw/train.zip wget https://monuseg.grand-challenge.org/static/dataset/MoNuSeg%20Test%20Data.zip -O data/raw/test.zip unzip data/raw/train.zip -d data/raw unzip data/raw/test.zip -d data/raw python src/slice_patches.py这里有两个细节值得注意。第一MoNuSeg 测试集官方只提供图像标签需要从 TCGA 的 XML 注释文件转换仓库里的脚本默认只拉图像部分如果你在离线环境复现建议直接找社区转换好的 mask 版本省去注释解析的额外工作。第二解压后目录名带空格Training Images和Test Images这种命名在 shell 循环里容易踩坑脚本里做了显式处理。提示图片命名是TCGA-xx-xxxx-01Z-00-DX1.png这种病理编号滑窗脚本依赖前缀做 train/val 划分不要重命名。2.2 512×512 滑窗切片MoNuSeg 原始图尺寸从 1000×1000 到 2000×2000 不等直接整图训练显存不够而且 SwinUNETR 的窗口划分要求输入尺寸能被 patch size 整除。常见做法是滑窗切成 512×512步长 256相邻 patch 保留 50% 重叠这样核落在 patch 边缘时不会丢失上下文。# src/slice_patches.py核心逻辑 import cv2, os patch_size, stride 512, 256 for split in [train, test]: img_dir fdata/raw/{split} out_img fdata/patches/{split}/images out_mask fdata/patches/{split}/masks os.makedirs(out_img, exist_okTrue) os.makedirs(out_mask, exist_okTrue) for name in os.listdir(img_dir): if not name.endswith(.png): continue img cv2.imread(os.path.join(img_dir, name)) mask cv2.imread(os.path.join(img_dir, name.replace(.png, _mask.png)), 0) h, w img.shape[:2] for y in range(0, h - patch_size 1, stride): for x in range(0, w - patch_size 1, stride): patch_img img[y:ypatch_size, x:xpatch_size] patch_mask mask[y:ypatch_size, x:xpatch_size] cv2.imwrite(f{out_img}/{name[:-4]}_{y}_{x}.png, patch_img) cv2.imwrite(f{out_mask}/{name[:-4]}_{y}_{x}.png, patch_mask)滑窗参数直接影响样本量和边界效应。步长 256 时一张 1000×1000 的图能切出大约 9 个 patch30 张训练图扩到两三百个样本对 Transformer 这种数据饥渴的模型勉强够用。显存允许的话可以把步长降到 128样本量接近翻倍但相邻 patch 高度相似验证集划分不当会让 Dice 虚高。我一般把同一个原始图切出的 patch 全部放进同一个 split避免训练和验证之间存在数据泄漏。2.3 训练集与验证集的增强差异MoNuSeg 标注包含人工核轮廓细胞核在 HE 染色下会有轻微形态形变albumentations 的 ElasticTransform 对这种扰动非常有效。训练管线是 Resize HorizontalFlip RandomRotate90 ElasticTransform Normalize验证集只做 Resize 和 Normalize保证指标不受随机性干扰。增强操作训练验证说明Resize(512,512)是是统一输入尺寸配合 SwinUNETR 窗口划分HorizontalFlip(p0.5)是否病理图无方向先验翻转安全RandomRotate90(p0.5)是否90 度倍数旋转不产生插值伪影ElasticTransform是否alpha120模拟组织切片形变Normalize是是ImageNet 均值方差配合预训练权重self.tf A.Compose([ A.Resize(img_size, img_size), A.HorizontalFlip(p0.5), A.RandomRotate90(p0.5), A.ElasticTransform(alpha120, sigma120 * 0.05, alpha_affine120 * 0.03, p0.5), A.Normalize(mean(0.485, 0.456, 0.406), std(0.229, 0.224, 0.225)), ToTensorV2() ])ElasticTransform 的三个参数 alpha、sigma、alpha_affine 必须一起调。alpha 控制位移幅度sigma 控制形变平滑程度alpha_affine 控制全局仿射分量。仓库里alpha120, sigma6, alpha_affine3.6这组值对细胞核这种小目标不会把结构扭曲到不可识别如果用 albumentations 默认的 alpha1形变效果几乎等于没有。注意 Normalize 必须在 ToTensorV2 之前顺序反了会直接抛异常因为 Normalize 期望 HWC 格式的 numpy 数组。还有一处隐蔽的坑MoNuSeg 原始 mask 是 RGB 三通道注释图直接cv2.imread(..., 0)读成灰度后核轮廓线和填充区域的灰度值可能不同。dataset.py 强制做了(mask0).astype(np.uint8)二值化把一切非零像素统一成前景。这个操作必须在增强之前完成否则 ElasticTransform 同时作用在图像和掩码上时掩码里的灰度不一致会被模型当成类别信息学进去。3. SwinUNETR 模型构建与 Focal-Tversky 损失3.1 SwinUNETR 主干结构与参数选择模型部分直接复用 MONAI 的 SwinUNETR不需要自己写 attention。SwinUNETR 是 U 型骨架encoder 是 Swin-Transformer 层级堆叠decoder 逐级反卷积上采样层间用 skip connection 把各阶段特征接回 decoder。相比纯 CNN 的 U-Net它在深层特征里保留了窗口间全局依赖核与周围组织的相对位置关系能被 attention 显式建模这一点对边界重叠严重的细胞核场景尤其关键。from monai.networks.nets import SwinUNETR def get_model(img_size512, out_channels2): return SwinUNETR( img_sizeimg_size, in_channels3, out_channelsout_channels, feature_size48, drop_rate0.1, attn_drop_rate0.1, dropout_path_rate0.1, use_checkpointTrue, )参数值影响img_size512必须与输入 patch 尺寸一致决定窗口数量feature_size48通道基数48 是显存与精度的折中drop_rate / attn_drop_rate0.1特征与 attention 权重的 dropoutdropout_path_rate0.1随机丢弃整个 transformer block 的概率use_checkpointTrue梯度检查点换显存训练变慢但显存省 30%-40%feature_size 是最值得调的参数。MONAI 官方预训练权重默认 feature_size48显存有富余可以提到 96分割精度会有可感知提升但显存占用接近翻倍。dropout_path_rate 设 0.1 是 Swin 系列常见默认值MoNuSeg 这种 30 张训练图的小规模数据如果想进一步压过拟合可以提到 0.2但要从 0.1 起步观察验证集曲线。use_checkpointTrue 对显存不充裕的机器很关键。SwinUNETR 的 attention 中间激活值很大开启后反向传播会重新计算前向结果以增加约 30% 耗时换回 30%-40% 显存。如果你的 batch size 只能开到 4优先开这个选项而不是降低输入分辨率因为 512 的输入对窗口划分是硬性要求缩到 256 虽然能训练但 attention 窗口的覆盖范围会明显变小。3.2 为什么用 Focal-Tversky 而不是 Dice 或 CE细胞核分割的不平衡体现在两个层面前景核占整张图像素可能只有 5% 到 15%背景主导核边界一两个像素的错误在大多数损失函数里被当作普通错误但在病理分析里边界差一个像素就影响核形态学测量。CrossEntropy 对像素独立计算天然不适合结构化目标Dice loss 对前景占比不敏感但梯度在预测接近 0 或 1 时容易饱和。Focal-Tversky 把 Tversky index 和 Focal 思想合并。Tversky index 用 alpha 和 beta 分别控制假阳性和假阴性的惩罚权重对细胞核这种小目标漏检核像素的代价更高Focal 部分用 gamma 指数对易分类样本降权让模型注意力集中在困难边界上。仓库里alpha0.7, beta0.3看起来 alpha 更大但注意 Tversky 定义里 alpha 乘的是假阳性这个设置结合 gamma 是对背景干扰和边界困难样本的平衡不是直觉上更重视哪一类的简单对应。class FocalTverskyLoss(nn.Module): def __init__(self, alpha0.7, beta0.3, gamma4/3): super().__init__() self.alpha, self.beta, self.gamma alpha, beta, gamma def forward(self, pred, target): pred pred.softmax(dim1)[:, 1] target target.squeeze(1) tp (pred * target).sum(dim[1, 2]) fp (pred * (1 - target)).sum(dim[1, 2]) fn ((1 - pred) * target).sum(dim[1, 2]) tversky (tp 1e-7) / (tp self.alpha * fp self.beta * fn 1e-7) loss (1 - tversky) ** self.gamma return loss.mean()forward 里的处理顺序逐行拆解pred 先 softmax 取第 1 通道也就是前景类概率对应 SwinUNETR 的两通道输出target 是 (B,1,H,W) 浮点掩码squeeze(1) 去掉通道维变成 (B,H,W)。tp/fp/fn 按 batch 内每个样本独立在 H、W 上求和最后用 1e-7 做数值稳定。gamma 取 4/3 而不是 Focal Loss 常用的 2因为 gamma 太大会让难样本梯度被过度放大细胞核边界像素在训练后期会振荡。想验证边界改善可以尝试把 gamma 提到 1.5但要同时盯住验证 Dice 的波动幅度。3.3 Dataset 输出格式与模型输入的对接SwinUNETR 期望输入是 (B, 3, H, W) 浮点张量输出 (B, 2, H, W) logits。dataset.py 的__getitem__返回的 mask 被转成 (1, H, W) float32损失函数里再 squeeze。这里有一个常见报错MONAI 的 SwinUNETR 在 forward 内部会校验 img_size改了 patch 尺寸但忘了同步 get_model 的 img_size会直接抛 shape mismatch。另外DataLoader 的collate_fn默认按 batch 堆叠如果 Dataset 返回的 img 和 mask 尺寸不一致堆叠时会报错所以 Resize 必须同时作用于 img 和 mask不能只处理图像。4. 训练循环、学习率调度与指标验证4.1 训练脚本结构与超参入口train.py 把数据加载、模型、损失、优化器、验证串成主流程命令行暴露 batch、lr、epochs、save 四个最常改的选项。默认 batch8、lr1e-4、epochs150这套组合在单张 24G 显存卡上能完整跑完。parser.add_argument(--data, defaultdata/patches) parser.add_argument(--batch, typeint, default8) parser.add_argument(--lr, typefloat, default1e-4) parser.add_argument(--epochs, typeint, default150) parser.add_argument(--save, defaultweights)训练主循环里有一个容易踩的坑pred model(img)输出 logits而 FocalTverskyLoss 内部自己做了 softmax所以训练时不要再对 pred 手动过 softmax否则概率被软化两次梯度数值明显变小表现为 loss 下降极慢。验证阶段则相反DiceMetric 需要 one-hot 硬标签验证循环里用AsDiscrete(argmaxTrue, to_onehot2)把 logits 转成 one-hot 再喂入。注意模型输出是 logits 时损失函数要接收 logits模型输出是概率时损失函数要接收概率。两种写法对应不同的前向处理混用是最常见的训练不收敛原因之一。4.2 优化器与学习率调度优化器选 AdamW 而不是 Adam核心原因是 Swin 这类 Transformer 对权重衰减更敏感AdamW 把 weight decay 从梯度更新里解耦出来可以独立控制。weight_decay1e-4 是 Swin 论文在 ImageNet 上的常用设置迁移到医学小数据集不需要大改。optimizer torch.optim.AdamW(model.parameters(), lrargs.lr, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_050, T_mult2)CosineAnnealingWarmRestarts 是训练 Transformer 分割模型时容易出效果的调度策略。T_050 表示第一个周期 50 个 epochT_mult2 表示后续周期长度翻倍即 50、100、200。每次重启时学习率回到初始值模型有机会跳出局部最优。150 个 epoch 会经历两次完整重启加半程前期较高学习率快速收敛后期余弦退火精细调整。如果验证 Dice 在 epoch 50 附近出现一次下降尖峰那是重启后的正常现象不是模型崩了看下一个周期的后半段即可。4.3 Dice、IoU、AJI 的验证逻辑验证用 MONAI 的 DiceMetricreduction 设 meaninclude_backgroundFalse 表示只算前景类。MoNuSeg 官方榜单更看重 AJIAggregated Jaccard Index它会把一个核分成两半和两个核合并成一个都算作错误比 Dice 严格很多也更贴近病理分析的实际需求。post_trans AsDiscrete(argmaxTrue, to_onehot2) dice_metric DiceMetric(include_backgroundFalse, reductionmean) # validation loop 核心 dice_metric(y_pred[post_trans(i) for i in decollate_batch(pred)], y[post_trans(i) for i in decollate_batch(mask)]) dice_vals.append(dice_metric.aggregate().item())decollate_batch是 MONAI 特有操作把 batch 维度的张量拆成单样本列表因为 DiceMetric 在默认配置下要求 y_pred 和 y 是样本列表而不是整 batch 张量。用惯了其他库的话这里直接传 (B,2,H,W) 会报错。实测 Dice 0.923、IoU 0.864、AJI 0.816 是单模型、无 TTA 的成绩说明 SwinUNETR 在这个任务上还有提升空间也说明指标本身没有虚高。4.4 权重保存与断点续训train.py 只在验证 Dice 高于历史最优时保存权重并把 best_dice 写入 log.json。这个策略比固定每 N 个 epoch 保存更实用因为训练后期 Dice 波动很小固定保存会留下大量同质化 checkpoint。当前脚本的局限是中断后从头重跑实际使用时建议每 10 个 epoch 额外存一个包含 optimizer 和 scheduler state_dict 的 checkpoint。恢复时要同时恢复 optimizer.state_dict() 和 scheduler 的 last_epoch只恢复模型权重会导致学习率从初始值重新起步后续几个 epoch 的 loss 会异常偏高。5. 推理、WSI 滑窗与 0.923 Dice 的复现要点5.1 单图推理的最小实现infer.py 的流程足够精简加载权重、读图、走预处理、softmax 取前景通道、阈值 0.5 二值化。torch.load时 map_location 参数决定权重能否跨设备加载CPU 机器加载 GPU 权重时不写会直接报 CUDA 错误model.eval()必须调用它关闭 dropout 和 attention 里的随机路径否则同一张图每次输出的掩码可能不同。def predict_patch(model, img): img transform(imageimg)[image].unsqueeze(0).cuda() with torch.no_grad(): pred model(img).softmax(dim1)[:, 1] return pred.squeeze().cpu().numpy() def load_model(weight_path, device): model get_model() model.load_state_dict(torch.load(weight_path, map_locationdevice)) model.eval() return model.to(device)5.2 WSI 整片滑窗推理的三个细节MoNuSeg 单张测试图已经算大图临床 WSI 是十万像素级别必须用 openslide 按金字塔层级读取在目标倍率下切 patch 推理再拼回大坐标。仓库的 infer_wsi.py 只给了框架真正落地时三个细节决定成败。import openslide, cv2, numpy as np, torch from src.infer import predict_patch wsi openslide.OpenSlide(xxx.svs) # 在 20x 或 40x 倍率下计算 patch 数量按 stride 滑窗 # 组织区域检测下采样缩略图 OTSU 过滤背景 # 每个 patch 推理后写入对应坐标重叠区域做平均池化第一WSI 背景远多于组织。先下采样读缩略图用 OTSU 阈值找出组织 bounding box只在组织区域内滑窗能省掉一半以上推理时间。第二patch 边缘预测质量差滑窗步长取 patch 一半重叠区域做平均或加权平均融合能显著减少拼接缝隙。第三OpenSlide 读出来是 BGRA必须转 RGB 再走预处理通道顺序颠倒会让模型输出完全错误的掩码这个 bug 视觉上很难发现因为整体形状仍然像细胞核。5.3 从 0.923 往更高分刷的三个方向想在这套代码上继续提分优先级排序是 TTA、后处理、模型集成。TTA 把输入做水平翻转和 90 度旋转同一 patch 预测四次取平均通常能带来 0.005 到 0.01 的 Dice 提升。后处理方面预测概率图用 scipy.ndimage 做一次形态学开运算过滤面积小于 20 像素的连通域真实细胞核不会小到那个程度。模型集成如果显存允许训练两个不同 seed 的 SwinUNETR一个 lr1e-4、一个 lr5e-5概率平均后 AJI 的提升往往比 Dice 更明显因为集成能减少紧贴核被合并这种结构化错误。复现时不要只看训练集和验证集 Dice单独跑一遍测试集把预测掩码和 Ground Truth 叠图保存视觉检查边界是否过分割或欠分割。数值指标在高分段对几个像素的差异不敏感但病理分析的下游任务对核形态特征要求高叠图检查是数值之外不可替代的一步。本文还有配套的精品资源点击获取
返回列表