Swin-Transformer在MoNuSeg细胞核分割中的工程实践
发布时间:2026/9/18 1:30:01来源:尧图网络
简介面向具备PyTorch基础、从事医学图像分析或图像分割的研究人员与开发工程师这份基于Swin-Transformer的细胞核分割端到端复现方案覆盖了MoNuSeg公开数据集从自动下载、滑窗切块、数据预处理、增强、模型训练到整张全切片图像推理的完整流程。实现以SwinUNETR为主干采用Focal-Tversky损失函数有效缓解前景背景类别不平衡配合albumentations增强库中的水平翻转、随机旋转与弹性变换提升分割泛化能力实测Dice分数达到零点九二三性能表现亮眼。资源虽仅包含一个docx说明文档、压缩包大小约三十五KB但浓缩了仓库目录结构说明、环境依赖清单、数据下载与切片脚本、模型构建关键代码、训练验证与推理命令等可直接落地的核心要点并附有Kaggle或Colab运行方式提示便于快速上手。目前已有136人学习适合作为Kaggle或科研竞赛的SOTA复现基线也可快速迁移到真实病理图像分割项目中作为二次开发的出发模板。1. Swin-Transformer 分割细胞核MoNuSeg 为什么不沿用纯卷积 U-Net在医学图像分割里Swin-Transformer 适不适用要看任务是否密集、边界是否敏感。细胞核恰好最吃边界HE 染色切片核密度极高一个 256×256 的 patch 常挤上百个核彼此粘连、染色不均。纯卷积 U-Net 的池化容易抹平小目标边界而 Swin 的窗口自注意力在窗口内直接建立任意两点间的依赖shifted window 让这种依赖逐层向外扩散分层结构又天然对齐 UNet 的跳跃连接。落地路径基本固定Swin 编码器加卷积解码器拼成 Swin-UNet在 MoNuSeg 的 30 张 1000×1000 训练切片上做端到端训练推理时滑窗切大图、阈值出语义掩码、watershed 拆粘连核。数据量小真正难点不在模型结构而在数据管线和训练超参。后面按数据、模型、训练、推理四条线展开代码可直接抄改参数旁边会标坑。2. Swin-Transformer 分割骨干的架构拆解与参数选型2.1 窗口注意力与分层特征分割任务要什么Swin 给什么Swin-Transformer 和 ViT 最大的差别是它不做全局自注意力而是把特征图切成互不重叠的 window在每个 window 内部算多头自注意力。以 window_size7 为例224×224 输入经过 patch_embed 变成 56×56 的 token 网格再切成 8×8 个窗口每个窗口只在自己的 7×7 范围内做注意力。这样做复杂度从 O(N²) 降到 O(N)而且窗口越小局部归纳偏置越强——这一点对细胞核分割反而是优点因为核的边界语义本来就是局部的。单纯切窗会让信息被锁死在窗口内所以 Swin 每隔一层把窗口偏移 (⌊W/2⌋, ⌊W/2⌋)也就是 shifted window。第奇数层的窗口边界落在第偶数层的窗口内部像素间的依赖就能逐层跨窗口传播。这个设计在代码里表现为shift_size参数从 0 切到 W/2 再切回 0配套的(2W−1)²大小的相对位置偏置表在初始化后就作为可学习参数参与训练。相对位置偏置是个常被忽略但很关键的细节它给注意力加了平移等变性对肿瘤切片这种没有固定朝向、只有相对位置关系的输入比 ViT 的绝对位置编码合理得多。分割还需要一个 ViT 给不了的东西多尺度。Swin 每个 stage 末尾用 Patch Merging 把 2×2 的 token 邻域合并成一个空间减半通道翻倍。Swin-T 四个 stage 走完通道从 96 到 192、384、768分辨率从输入的 1/4 一路降到 1/32。这正好对应 UNet 里 skip connection 要的那几级特征图所以 Swin 做分割不需要像 ViT 那样额外插一个特征金字塔结构上自带。2.2 解码器装配方式Swin-UNet 跳跃连接与 UPerNet 的差异用 Swin 做分割主要有两条路线。一条是完全 Transformer 化的 Swin-UNet解码器用 Patch Expanding 逐级上采样每一级和编码器对应层做 skip concatenation全程没有卷积。另一条是 Swin 只当 backbone解码器用卷积也就是 UPerNet 或 UNet-like 的混合结构。两条路线在 MoNuSeg 上都有人跑通但数据集只有 30 张训练图从零训练一个纯 Transformer 解码器很容易过拟合。我一般走混合路线backbone 用 timm 加载 ImageNet 预训练的 Swin-T解码器用轻量卷积。预训练权重虽然来自自然图像但它学到的边缘、纹理滤波器对 HE 染色图仍有迁移价值收敛速度和稳定性都比从零训好。下面这段代码验证 backbone 输出的特征图尺寸和通道数确认能和 UNet 解码器对接。# 用 timm 抽取 Swin-T 的 4 级特征验证 shape 是否符合 UNet 拼接需求 import torch import timm encoder timm.create_model( swin_tiny_patch4_window7_224, pretrainedTrue, features_onlyTrue, out_indices(0, 1, 2, 3), # 取 4 个 stage 的输出 ) x torch.randn(1, 3, 224, 224) outs encoder(x) for i, f in enumerate(outs): print(fstage{i}: {f.shape})features_onlyTrue会去掉分类头直接返回各 stage 的特征图out_indices(0,1,2,3)对应分辨率为输入的 1/4、1/8、1/16、1/32 四层。输出布局有个版本差异要留意timm 0.9 之后Swin 这类模型在 features_only 模式下默认返回 NHWC 而不是 NCHW接 Conv2d 之前要手动permute(0, 3, 1, 2)老版本不用。上面代码不做 permute就是为了先单纯看 shapeSwin-T 的四个输出分别是 56×56×96、28×28×192、14×14×384、7×7×768。如果输入换成 512×512第一层就是 128×128×96显存占用接近原来的 5 倍先把 224 的管线跑通再上大分辨率是更稳的顺序。2.3 Swin-T 关键参数表与输入分辨率对齐Swin-T 移植到 MoNuSeg 时需要动的主要参数就几个。下面的表给出默认值和针对细胞核任务的调整方向。参数Swin-T 默认MoNuSeg 推荐调整影响patch_size4保持 4改成 2 会让 token 数变成 4 倍显存和训练时间都扛不住embed_dim9696显存宽裕可 128决定第一层通道数模型规模的主要来源depths[2,2,6,2][2,2,6,2] 或 [2,2,2,2]30 张训练图时删深 stage 能明显压过拟合num_heads[3,6,12,24]随 embed_dim 等比缩放和 embed_dim 强绑定不要单独乱调window_size77 或 8必须和输入分辨率对齐见下方说明drop_path_rate0.10.2 ~ 0.3数据量小时加大随机深度正则效果显著window_size 和输入分辨率不是随便配的。patch_embed 之后 token 数是 H/4 × W/4这个数必须能被 window 整除。224 输入得到 56×56 token56/78正好对齐256 输入得到 64×64 token64 不能被 7 整除就要换 window8 或者把输入 pad 到 280。更隐蔽的坑是换 window_size 会让相对位置偏置表从 (2×7−1)² 变成 (2×8−1)²这一层参数全部重新初始化预训练权重在这部分直接失效。所以想完整吃到 ImageNet 预训练红利优先保持 window7只在输入分辨率上做文章。提示换 window_size 等于丢弃预训练的相对位置偏置表。能用输入分辨率对齐解决的就不要动 window。3. MoNuSeg 数据管线染色归一化、Patch 切分与边界权重3.1 标注形态与训练/验证划分MoNuSeg 是 MICCAI 2018 的细胞核分割挑战赛训练集 30 张 1000×1000 HE 全视野切片截图覆盖乳腺、肾、结肠、肺、膀胱、前列腺、胃 7 个器官测试集 14 张比训练集多出肝脏。官方的标注以细胞核边界坐标点形式下发网上流通的版本大多已经预处理成二值掩码。拿到数据第一步是确认掩码语义画的是边界线还是实心区域——训练要的是实心核掩码如果只有边界线需要先用多边形填充算法把轮廓内部填满。30 张图直接做 random split 风险很大。我一般按器官做分层抽样每个器官抽 1 张左右组成 5~6 张的验证集保证验证集里不含某个完全没见过的器官。否则某一器官的所有图都进了训练集验证指标会虚高测试集一换器官就现原形。MoNuSeg 的测试集包含训练集没有的器官肝脏这也意味着模型必须对染色变异和形态差异有一定的泛化能力而不是记住器官特征。3.2 HE 染色差异先做染色归一化再进模型不同医院、不同批次的 HE 切片染色深浅和色相差异非常大。同一个细胞核在一张图里是深紫色在另一张图里可能是浅粉色。如果不处理模型会把染色风格当成特征去学习泛化到新切片时掉点明显。常见的做法有两种完整版用 Macenko 染色归一化把每个样本映射到一张参考图的染色空间快速版用 CLAHE 做 per-image 对比度均衡。MoNuSeg 的 30 张训练图来自同一数据源染色差异没有多中心数据那么极端CLAHE 通常够用而且不需要参考图管线更简单。import cv2 import numpy as np def clahe_image(rgb: np.ndarray, clip: float 2.0, tile: int 8) - np.ndarray: 对 LAB 色彩空间的 L 通道做 CLAHE保留色相适合 HE 批量预处理。 lab cv2.cvtColor(rgb, cv2.COLOR_RGB2LAB) l, a, b cv2.split(lab) clahe cv2.createCLAHE(clipLimitclip, tileGridSize(tile, tile)) l clahe.apply(l) return cv2.cvtColor(cv2.merge([l, a, b]), cv2.LAB2RGB)在 LAB 空间只处理 L 通道是为了不破坏 HE 的色相信息——细胞核和细胞质的颜色区分是病理诊断的重要线索。clip控制对比度放大的上限设太大会把染色噪声一起放大2.0 比较稳tile是局部直方图网格大小8×8 是常规值。跑完整版 Macenko 的话参考图选 30 张里染色最接近中位水平的那一张而不是随便挑第一张。3.3 Patch 切分与数据增强参数1000×1000 的整图直接进 Swin-T 不现实训练时随机采 patch验证和推理时用滑窗切全图。patch 尺寸通常选 256×256 起步显存允许再上 512 微调。滑窗的 overlap 决定推理时接缝的严重程度64~128 是常用范围。下面这个函数是推理用的全图覆盖版本训练用的随机采样版不需要记录坐标。def extract_patches(img: np.ndarray, mask: np.ndarray, patch: int 256, overlap: int 64): 滑窗切 patch返回 patch 对和每个 patch 的左上角坐标。 stride patch - overlap h, w img.shape[:2] ys list(range(0, h - patch 1, stride)) xs list(range(0, w - patch 1, stride)) if ys[-1] patch h: # 边缘补一刀保证覆盖整图 ys.append(h - patch) if xs[-1] patch w: xs.append(w - patch) imgs, masks, coords [], [], [] for y in ys: for x in xs: imgs.append(img[y:y patch, x:x patch]) masks.append(mask[y:y patch, x:x patch]) coords.append((y, x)) return np.stack(imgs), np.stack(masks), coords1000×1000 的图256 patch 64 overlap 时 stride 是 192按 0、192、384、576、768 滑动最后一个位置 7682561024 超出边界所以代码里补一个h - patch的对齐刀每张图产出 36 个 patch。overlap 越大推理时重叠区域越多融合后接缝越不明显但推理量成倍增加一般 64 就能压住接缝。训练时的增强参数按下表配置注意病理图没有方向性旋转角度可以放开。增强操作参数说明随机旋转0/90/180/270 等概率再叠加 ±15°病理切片无固定朝向水平/垂直翻转各 p0.5常规操作弹性形变σ4α6模拟组织形变增强边界鲁棒性亮度/对比度抖动0.8 ~ 1.2缓解染色差异弹性形变同步掩码和权重图用同一插值网格必须和图像共用变换3.4 细胞核边界权重图细胞核分割的标签有一个特性核与核之间的边界极窄且相邻核几乎贴在一起边界像素在整张图里的占比很低。普通 BCE 会把梯度平均摊到所有像素上边界区域的损失被大面积核内部稀释。U-Net 原论文给出的解法是边界权重图让模型把学习能力集中在难分的边界上。MoNuSeg 场景里不需要严格照搬原公式一个简化版就够用from scipy import ndimage as ndi def border_weight_map(mask: np.ndarray, sigma: float 5.0, weight_bg: float 3.0) - np.ndarray: 给紧贴核边界的外围背景加权把训练梯度引到粘连区域。 bg (mask 0).astype(np.uint8) d ndi.distance_transform_edt(bg) # 背景像素到最近核的距离 w np.ones_like(mask, dtypenp.float32) w weight_bg * np.exp(-((d - sigma) ** 2) / (2 * (sigma / 2) ** 2)) w mask.astype(np.float32) * 0.5 return w核心逻辑是对背景做距离变换距离场在 sigma 附近的背景像素获得最高权重——这些像素正好是紧贴核边界的区域。sigma取 5 表示给核边界外约 5 像素宽的区域加权weight_bg是背景加权的倍率3.0 足够把边界区域的损失抬起来。原版的 w0·exp(−(d1d2)²/(2σ²)) 公式需要计算到最近和第二近细胞边界的距离实现复杂且收益有限。需要注意这个权重图必须和图像、掩码做完全相同的增强变换旋转翻转要对齐弹性形变要共用同一套插值网格否则边界权重和实际边界会错位。4. 端到端训练Swin-UNet 装配、损失函数与超参4.1 最小可跑的 Swin-UNettimm 骨干加卷积解码器编码器用 timm 的 Swin-T解码器用四层卷积上采样每一级做 up concat 卷积降维。下面的实现去掉了所有花活是能直接跑通端到端训练的最小骨架。import torch import torch.nn as nn import timm class DecoderBlock(nn.Module): def __init__(self, in_ch: int, skip_ch: int, out_ch: int): super().__init__() self.up nn.ConvTranspose2d(in_ch, skip_ch, kernel_size2, stride2) self.conv nn.Sequential( nn.Conv2d(skip_ch * 2, out_ch, 3, padding1, biasFalse), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x, skip): return self.conv(torch.cat([self.up(x), skip], dim1)) class SwinUNet(nn.Module): def __init__(self, num_classes: int 2): super().__init__() self.encoder timm.create_model( swin_tiny_patch4_window7_224, pretrainedTrue, features_onlyTrue, out_indices(0, 1, 2, 3), drop_path_rate0.2, ) # Swin-T 四层通道数: 96, 192, 384, 768 self.dec3 DecoderBlock(768, 384, 192) # 7 - 14 self.dec2 DecoderBlock(192, 192, 96) # 14 - 28 self.dec1 DecoderBlock(96, 96, 64) # 28 - 56 self.up0 nn.ConvTranspose2d(64, 64, kernel_size2, stride2) self.head nn.Sequential( nn.Conv2d(64, 64, 3, padding1), nn.ReLU(inplaceTrue), nn.Conv2d(64, num_classes, 1), nn.Upsample(scale_factor2, modebilinear, align_cornersFalse), ) def forward(self, x): f0, f1, f2, f3 self.encoder(x) # timm 较新版本返回 NHWC f0, f1, f2, f3 (f.permute(0, 3, 1, 2) for f in (f0, f1, f2, f3)) x self.dec3(f3, f2) x self.dec2(x, f1) x self.dec1(x, f0) x self.up0(x) # 56 - 112 return self.head(x) # 112 - 224forward 里的permute是给 timm 新版本输出 NHWC 布局用的老版本输出 NCHW 时 permute 也不会出错所以保留这一行最保险。drop_path_rate0.2在 30 张训练图的规模下属于必要正则不要省。解码器的通道数从 768 一路压到 64最后一个上采样把分辨率恢复到和输入一致。如果想把输入从 224 换成 512需要同步改 head 里的Upsample倍数和 window 对齐策略否则输出分辨率对不上。4.2 损失函数Dice 加加权 BCE细胞核分割的标签分布极不平衡核区域通常只占整张图的 10%~30%纯 BCE 会被大面积背景带偏纯 Dice 对小目标的梯度又不够稳定训练早期容易震荡。常见做法是 Dice BCE 相加再把 3.4 节的边界权重图乘到 BCE 的逐像素项上。import torch.nn.functional as F def weighted_dice_bce(logits: torch.Tensor, target: torch.Tensor, weight_map: torch.Tensor, smooth: float 1.0): # logits: (B, 2, H, W)取类别 1 的概率 probs torch.softmax(logits, dim1)[:, 1].float() target target.float() # 加权 BCE逐像素算完再按权重图归一化 bce F.binary_cross_entropy(probs, target, reductionnone) bce (bce * weight_map).sum() / (weight_map.sum() 1e-6) # Dice 不加权保持全局敏感性 inter (probs * target).sum() dice 1.0 - (2.0 * inter smooth) / (probs.sum() target.sum() smooth) return dice bcetarget 必须是单通道 0/1 掩码不要转 one-hotsoftmax 之后取第 1 个通道就是前景概率。BCE 除以weight_map.sum()而不是像素总数是为了避免加权后损失尺度被放大训练曲线看起来像在波动但实际没变。标签里如果存在未标注区域——MoNuSeg 的训练标注只覆盖肿瘤区域非肿瘤区域不标——不要把它们当普通背景硬学把对应位置的权重图置 0 是常见做法。4.3 优化器、学习率与参数分组预训练骨干加随机初始化解码器两者学习率必须分开。骨干权重已经收敛学习率过高会把预训练学到的特征冲掉解码器从零开始需要更快的收敛速度。常见做法是骨干 lr 取解码器的十分之一。优化器选 AdamW权重衰减和解耦逻辑比 Adam 干净Swin 官方也用它。超参推荐值说明骨干学习率2e-5预训练权重下压低防止灾难性遗忘解码器学习率2e-4从零训练十倍差weight_decay0.05Swin 官方常用值batch size8256×25624G 显存可上 16训练轮数100 ~ 200配合早停看 val Dice学习率调度warmup 10% cosine预热防止预训练权重被冲AMPfp16 开启显存直接省一半optimizer torch.optim.AdamW([ {params: model.encoder.parameters(), lr: 2e-5}, {params: [p for n, p in model.named_parameters() if not n.startswith(encoder.)], lr: 2e-4}, ], weight_decay0.05) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max150, eta_min1e-6) scaler torch.cuda.amp.GradScaler() for epoch in range(150): model.train() for imgs, masks, weights in loader: imgs, masks, weights imgs.cuda(), masks.cuda(), weights.cuda() with torch.cuda.amp.autocast(): logits model(imgs) loss weighted_dice_bce(logits, masks, weights) scaler.scale(loss).backward() scaler.unscale_(optimizer) torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) scaler.step(optimizer) scaler.update() optimizer.zero_grad() scheduler.step()参数分组通过named_parameters()的模块名前缀判断encoder.开头的进骨干组其余进解码器组。grad clip 的 1.0 不是随便拍的Swin 的窗口注意力在 shifted window 切换处梯度会有尖峰warmup 结束瞬间不 clip 很容易产出 NaN。AMP 三件套GradScaler、autocast、unscale是标准写法scaler.unscale_必须在 clip 之前执行否则 clip 作用在 fp32 缩放后的梯度上数值不对。4.4 训练监控与失效排查训练 Swin-UNet 最容易踩的坑按出现频率排一下。loss 卡在 0.69 左右不动这个值接近 BCE 在类别平衡时的熵上限基本可以断定 target 全是 0 或者形状没对齐检查数据加载器的 squeeze 和 dtype。val Dice 高而实例数虚高多数是标签是空心边界线模型学到了「描边」而不是「填核」。换输入分辨率之后训练直接炸八成是 2.3 节的 window 对齐问题回到输入尺寸表查一遍。过拟合显著train Dice 0.95 而 val 只有 0.75先加 drop_path_rate 到 0.3再加强染色抖动最后才考虑砍 depths。显存不足时先试timm.create_model(..., use_checkpointTrue)这个参数开启梯度检查点用一点训练时间换一半显存。5. 推理系统的滑窗融合与粘连核分离5.1 重叠区域概率融合推理时滑窗切出来的 patch 之间有 overlap同一像素会被预测多次。直接拼接会在 patch 边界留下接缝因为模型对窗口边缘的预测置信度天然偏低。常见做法是给每个 patch 生成一个线性斜坡权重图边缘为 0 中心为 1预测概率乘权重累加最后除以权重累计值。def linear_ramp(size: int, overlap: int) - np.ndarray: w np.ones(size, dtypenp.float32) if overlap 0: ramp np.linspace(0.0, 1.0, overlap) w[:overlap] ramp w[-overlap:] ramp[::-1] return w def patch_weight(size: int, overlap: int) - np.ndarray: wy linear_ramp(size, overlap) wx linear_ramp(size, overlap) return wy[:, None] * wx[None, :] # 二维斜坡叠加概率图时把每个 patch 的预测乘patch_weight(256, 64)再加到累积数组同时累积权重数组本身。最终概率图是两者相除。这个融合方式对病理图这类大图拼接尤其重要overlap 只有 32 时接缝还能隐约看到64 以上基本消失。5.2 Watershed 分离粘连核语义分割输出的是「哪些像素是核」但病理分析要的是「哪个核是哪个核」。MoNuSeg 官方评测按实例算粘连的核如果被语义掩码连成一个连通域实例数直接少一半。分离粘连核的标准做法是距离变换加 watershed先对语义掩码做距离变换核中心的距离值大、边界距离值小以局部极大值为种子做 watershed就能把粘连的核从边界处切开。import numpy as np from scipy import ndimage as ndi from skimage.feature import peak_local_max from skimage.segmentation import watershed def nuclei_instances(prob: np.ndarray, thresh: float 0.5, min_dist: int 5, min_area: int 30) - np.ndarray: sem prob thresh sem ndi.binary_opening(sem, iterations1) # 去掉孤立噪点 dist ndi.distance_transform_edt(sem) peaks peak_local_max(dist, min_distancemin_dist, exclude_borderFalse) markers, _ ndi.label(peaks) labels watershed(-dist, markers, masksem) ids, counts np.unique(labels, return_countsTrue) out np.zeros_like(labels) for lbl, cnt in zip(ids[1:], counts[1:]): if cnt min_area: out[labels lbl] lbl return outmin_distance是最关键的一个参数设太小一个大核内部会出现多个种子核被切碎设太大两个粘连核共用一个种子分离失败。MoNuSeg 核的半径一般在 8~15 像素min_dist 从 5 起步按实际数据集的核径分布调。binary_opening用来清掉概率图上低于阈值的散点避免 watershed 在噪声上生成大量碎片。5.3 指标验证与贴边策略跑完推理先别急着看 Dice。MoNuSeg 官方排行榜用的是 AJIAggregated Jaccard Index是实例级指标预测实例和真实实例先做匹配再在匹配对的基础上算 Jaccard。像素级 Dice 对粘连不敏感两个案例一个是完美分割、一个是把整片核连成一坨Dice 可能只差零点几个点AJI 能差出十几个点。复现时评估脚本至少同时报 Dice、IoU 和 AJI 三个数只用 Dice 调参会把实例分离能力调没。另一个常被忽略的是贴边策略。滑窗切出的 patch 边缘会有不完整的核验证时预测实例和 GT 实例在贴边的处理上必须一致都保留、都丢弃、还是按面积截断。MoNuSeg 的标注本身不包含被裁掉一半的核所以推理结果里接触 patch 边界的实例建议先单独统计和 GT 口径对齐后再计入指标不然 AJI 会莫名掉几个点。调整 watershed 参数前先花五分钟把训练集 30 张图的核面积直方图打出来min_dist和min_area按这个分布的第 10 百分位定比凭经验拍脑袋准确得多。本文还有配套的精品资源点击获取
网站建设高端定制企业官网