简介:面向医学图像分割与深度学习研究者,这套方案提供基于Swin-Transformer与U-Net融合的宫颈细胞核分割项目,支持自适应多尺度训练、多类别分割与迁移学习,可直接用于细胞核区域提取与分割实验。资源包约200.84MB,共809个文件;其中391张jpg原始图像、383张png标注或结果图,另有Python训练/推理脚本、pyc缓存、pth权重、txt训练记录与xml配置,结构清晰,便于对照学习。目前已有342人学习,项目实测仅训练50个epochs,全局像素准确率约0.92、miou约0.767;训练脚本会将输入随机缩放至设定尺寸的0.5~1.5倍,实现多尺度增强,utils中的compute_gray会自动将mask灰度值写入txt并匹配输出通道。训练采用cos学习率衰减,run_results中保存了训练集与测试集的损失、IoU曲线及每类IoU、recall、precision和全局准确率;推理时只需将待测图像放入inference目录运行predict脚本即可,具体流程见README,方便快速复现与迁移学习实践。
1. 为什么子宫颈细胞核分割偏偏要用 Swin-Transformer + Unet
做宫颈液基细胞学(TCT)自动分析的人,最先碰到的就是细胞核分割。细胞核是后续分级、判读的基础,但也是一个“看着简单、跑起来翻车”的任务:一张 2048×2048 的涂片里,细胞核可能只有 80 到 300 个,每个核直径从 8 像素到 80 像素不等;类别上至少分正常、低度病变、高度病变、炎症等,类别之间边界模糊,染色差异又大。纯 Unet 在这种密集小目标+大尺度差异的任务上,分割结果总是“差不多但不够用”。而 Swin-Transformer 做编码器、Unet 做解码器,正好把前者的窗口注意力特征表示能力和后者的多分辨率细节恢复能力拼在一起,是目前这类任务里最稳的组合方案。本文按这个思路,把网络结构、自适应多尺度训练、多类别分割损失和迁移学习策略一条线拆开讲,给出能直接改、能跑通的实现细节。
2. 网络结构怎么搭:Swin-Transformer 做编码器,Unet 解码器做恢复
2.1 为什么不用纯 Transformer 或纯 Unet
先做一个对比,帮助选型。纯 Unet(包括 Unet++、Attention Unet)在细胞核分割上的问题不在表达能力,而在感受野受限。细胞核本身是小目标,但“判断一个核是不是病变”往往需要看周围细胞的排列、核浆比这些上下文信息,Unet 的卷积要堆很深才能把上下文范围拉大,堆深了又丢失边缘细节。
纯 Transformer(例如 TransUNet 的直接替换)也不省心。Vision Transformer 的全局自注意力在 512×512 输入上计算量大,而且全局注意力会被大面积细胞质背景稀释——细胞核只占整张图的 5% 左右,注意力权重很容易被背景主导,导致小核漏检。
Swin-Transformer 是折中方案里效果最好的一个:它用窗口(window)限制自注意力范围,窗口内做自注意力,窗口之间通过 shift 操作跨窗交互。这样既控制了计算量,又让注意力集中在小区域内的细胞形态上,天然适配细胞核这种密集目标。Swin-T 的层级结构(4 个 stage,输出 1/4、1/8、1/16、1/32 分辨率特征)刚好可以和 Unet 解码器逐级跳跃连接对齐。所以常见做法是:Swin-T/Swin-B 做编码器,Unet 风格解码器做像素级恢复,中间用卷积层把通道对齐。
2.2 完整结构代码:基于 timm 的 Swin-Transformer + Unet 解码器
我用 timm 库加载 Swin 预训练权重,比自己手写 Swin 省事且不容易写错。下面是一个可直接跑通的完整网络定义,输入 512×512 单模态灰度图(或 3 通道 RGB,由你的数据决定),输出 N 类分割概率图。
import torch import torch.nn as nn import timm class SwinUnet(nn.Module): def __init__(self, num_classes=4, in_chans=3, img_size=512, pretrained=True, decoder_channels=[256, 128, 64, 32]): super().__init__() # 编码器:Swin-T,去掉分类头 self.encoder = timm.create_model( 'swin_tiny_patch4_window7_224', pretrained=pretrained, in_chans=in_chans, num_classes=0, # 只保留特征提取部分 img_size=img_size ) # 取 Swin 每个 stage 的中间特征 self.stage_channels = [96, 192, 384, 768] # 用 1x1 卷积把 Swin 各 stage 通道对齐到解码器通道 self.proj1 = nn.Conv2d(self.stage_channels[0], decoder_channels[0], 1) self.proj2 = nn.Conv2d(self.stage_channels[1], decoder_channels[1], 1) self.proj3 = nn.Conv2d(self.stage_channels[2], decoder_channels[2], 1) self.proj4 = nn.Conv2d(self.stage_channels[3], decoder_channels[3], 1) # Unet 风格解码器 self.up4 = nn.ConvTranspose2d(decoder_channels[3], decoder_channels[2], 2, stride=2) self.up3 = nn.ConvTranspose2d(decoder_channels[2], decoder_channels[1], 2, stride=2) self.up2 = nn.ConvTranspose2d(decoder_channels[1], decoder_channels[0], 2, stride=2) self.up1 = nn.ConvTranspose2d(decoder_channels[0], decoder_channels[0], 2, stride=2) self.conv4 = nn.Conv2d(decoder_channels[2]*2, decoder_channels[2], 3, padding=1) self.conv3 = nn.Conv2d(decoder_channels[1]*2, decoder_channels[1], 3, padding=1) self.conv2 = nn.Conv2d(decoder_channels[0]*2, decoder_channels[0], 3, padding=1) self.conv1 = nn.Conv2d(decoder_channels[0]*2, decoder_channels[0], 3, padding=1) self.head = nn.Conv2d(decoder_channels[0], num_classes, 1) self.relu = nn.ReLU(inplace=True) def forward(self, x): # 编码器。features 依次为 1/4, 1/8, 1/16, 1/32 分辨率 features = self.encoder(x) f1, f2, f3, f4 = features # 都是 (B, C, H, W) 形式 # 通道对齐 e1 = self.relu(self.proj1(f1)) # 1/4 e2 = self.relu(self.proj2(f2)) # 1/8 e3 = self.relu(self.proj3(f3)) # 1/16 e4 = self.relu(self.proj4(f4)) # 1/32 # 解码 d = self.relu(self.up4(e4)) # 1/16 d = self.conv4(torch.cat([d, e3], dim=1)) d = self.relu(self.up3(d)) # 1/8 d = self.conv3(torch.cat([d, e2], dim=1)) d = self.relu(self.up2(d)) # 1/4 d = self.conv2(torch.cat([d, e1], dim=1)) d = self.relu(self.up1(d)) # 1/2 d = self.conv1(d) # 恢复到原尺寸 d = nn.functional.interpolate(d, size=x.shape[2:], mode='bilinear', align_corners=False) return self.head(d)代码里的关键点都在注释里了。做三点额外说明:
encoder(x)返回的是一个 list,timm 的 Swin 会返回每个 stage 输出的特征张量,顺序是 1/4、1/8、1/16、1/32,不需要自己从中间层抠。这样写代码短、不易错。- 跳跃连接用
torch.cat拼接,通道数翻倍,所以每个conv的输入通道是decoder_channels[i]*2。这是 Unet 的经典结构,保留浅层空间细节。 - 最后用
interpolate上采样到输入尺寸。如果你想省掉这步,可以在forward里直接输出 1/2 分辨率,训练时把标签也降采样到 1/2,推理时再上采样。但我建议输出全分辨率,多类别小目标分割对分辨率很敏感,全分辨率输出在后处理阶段省去一次插值误差。
2.3 Swin 预训练权重和输入尺寸的适配
Swin-T 默认训练尺寸是 224×224,window_size=7。如果你的输入是 512×512,需要在timm.create_model里把img_size改为 512。Swin 位置编码是相对位置编码,理论上可以支持任意尺寸,但 window 的数量必须是偶数(因为 shift 操作),512/4=128,128/7 不是整数,Swin 会做 padding 自动补齐,不会有问题,但注意早期版本的 timm 在非 224 尺寸下可能报错,升级 timm 到 0.9+ 基本能规避。
另外,宫颈细胞学数据通常只有灰度或特定染色通道。我一般会把 3 通道预训练权重换成单通道输入:in_chans=1时,timm 会随机初始化第一层卷积的权重,这会导致预训练优势在浅层丢失。常见做法是保留 3 通道输入,把同一张灰度图复制成 3 个通道输入网络,这样第一层卷积权重基本派得上用场,训练也更稳。
3. 自适应多尺度训练:让模型在核大小差异极大的数据上不翻车
3.1 固定多尺度训练和自适应多尺度的区别
多尺度训练在医学图像分割里很常见,做法是把每个训练样本随机缩放到 0.75×、1.0×、1.5× 再输入网络,相当于数据增强+尺度泛化。但在宫颈细胞核这种场景,固定尺度列表有一个明显缺陷:不同涂片的细胞核平均直径差异巨大,同一批数据里有的核只有 10 像素,有的 60 像素。固定 0.75× 尺度会把本来就小的核缩得更小,标签几乎消失了;而固定 1.5× 尺度又会让大核超出感受野范围。
自适应多尺度的思路是:让尺度范围跟着当前样本的核大小分布走。具体做法是在 DataLoader 里预先统计每个训练样本掩膜中细胞核的面积分布,取每个核等效直径的中位数,再把它映射到预设的“目标核直径范围”上,据此算出这个样本的缩放系数。
3.2 自适应尺度采样器实现
我一般实现一个AdaptiveScaleSampler,在 Dataset 返回样本和标签后,根据标签实时算缩放因子。这样不用改网络结构,只改数据流。
import torch import numpy as np import torchvision.transforms.functional as F class AdaptiveScaleSampler: def __init__(self, min_scale=0.5, max_scale=2.0, target_diameter=(24, 48)): self.min_scale = min_scale self.max_scale = max_scale self.target_diameter = target_diameter def compute_scale(self, mask): # mask: (C, H, W) 或 (H, W),像素值 0 为背景,其余为类别 if len(mask.shape) == 3: mask_bin = (mask.sum(0) > 0).astype(np.uint8) else: mask_bin = (mask > 0).astype(np.uint8) # 连通域提取,得到每个核的像素面积 from skimage.measure import label, regionprops lab = label(mask_bin, connectivity=2) if lab.max() == 0: return 1.0 props = regionprops(lab) diameters = [2 * np.sqrt(p.area / np.pi) for p in props if p.area >= 5] if len(diameters) == 0: return 1.0 med_d = np.median(diameters) target_center = (self.target_diameter[0] + self.target_diameter[1]) / 2 scale = target_center / med_d # 限制缩放范围,防止过采样或欠采样 scale = np.clip(scale, self.min_scale, self.max_scale) return float(scale) def __call__(self, image, mask): scale = self.compute_scale(mask) new_h = int(round(image.shape[1] * scale)) new_w = int(round(image.shape[2] * scale)) new_h = max(new_h, 256) # 下限保护 new_w = max(new_w, 256) img = F.resize(torch.from_numpy(image), (new_h, new_w), interpolation=F.InterpolationMode.BILINEAR) msk = F.resize(torch.from_numpy(mask.astype(np.int64)), (new_h, new_w), interpolation=F.InterpolationMode.NEAREST) return img, msk, scale这段代码的逻辑是:每次迭代从当前样本的掩膜统计核直径中位数,然后算出一个缩放系数,使该样本的核中位直径落在 24–48 像素区间。这样大核样本会被缩小、小核样本会被放大,模型在每个 batch 里看到的核大小分布相对一致。
参数说明有三点值得注意:
target_diameter=(24, 48)不是拍脑袋定的。Swin 的 window 是 7×7 patch,每 patch 4 像素,窗口实际覆盖 28×28 像素区域。把核的中位直径控制在 24–48,意味着大多数核在窗口内完整可见,注意力能捕捉整个核的形态,而不是只看局部。min_scale和max_scale设置成 0.5–2.0,是考虑到过度缩放会把细胞形态扭曲。如果某个样本的核特别大(比如 100 像素),最大 2.0 倍的缩放依然到不了 48 像素,这时不要硬顶上去,保持 2.0 倍即可,让模型允许这种极端样本存在。- 掩膜必须用
NEAREST插值,不能用 BILINEAR。多类别掩膜用线性插值会产生类间混叠的中间值,比如类别 1 和类别 2 之间插出 0.5,损失函数直接报错或乱算。这是刚入门最容易踩的坑。
3.3 多尺度训练时的损失函数怎么设计
多尺度训练不只是数据层改动,损失函数也要配合。我常用的组合是CrossEntropy + Dice Loss + 尺度自适应权重。Dice Loss 天然缓解类别不平衡,交叉熵保证梯度稳定,而尺度自适应权重解决的是“小核类别被大核类别淹没”的问题。
class ScaleAwareDiceLoss(nn.Module): def __init__(self, n_classes, epsilon=1e-6): super().__init__() self.n_classes = n_classes self.epsilon = epsilon def forward(self, logits, targets, scale): # logits: (B, C, H, W), targets: (B, H, W), scale: (B,) probs = torch.softmax(logits, dim=1) # (B, C, H, W) targets_onehot = torch.nn.functional.one_hot(targets, self.n_classes) # (B, H, W, C) targets_onehot = targets_onehot.permute(0, 3, 1, 2).float() dims = (2, 3) intersect = (probs * targets_onehot).sum(dim=dims) union = probs.sum(dim=dims) + targets_onehot.sum(dim=dims) dice = (2 * intersect + self.epsilon) / (union + self.epsilon) # 每个样本按其缩放系数加权:缩得越狠的样本,小核越多,权重越大 weights = 1.0 / (scale + 1e-3) weights = weights / weights.mean() # 归一化,保持整体损失量级不变 loss = 1.0 - dice loss = (loss * weights.view(-1, 1)).mean(dim=0) return loss.mean()这里和固定多尺度训练最大的差别在weights的用法。固定多尺度训练时所有尺度样本一视同仁,而自适应尺度中,scale < 1 意味着这个样本被缩小了,说明它原本是大核样本;scale > 1 说明原本是小核样本。小核样本的分割难度更高,给它们更高的损失权重,相当于把模型的优化重心往难点样本上压。
Dice Loss 和交叉熵的配比,我建议0.4 * dice + 0.6 * ce,先跑 20 个 epoch,再看验证集的类别别 Dice 调整比例。如果不同类别 Dice 方差大,把 dice 权重提到 0.6,让模型优先平衡类别间精度。
4. 多类别分割落地:标签体系、迁移学习与 Unet 训练实战
4.1 多类别分割的标签体系怎么定
宫颈细胞核多类别分割的标签体系直接决定模型能回答什么问题。我在实际项目中用的是四类结构:
| 类别编号 | 名称 | 说明 |
|---|---|---|
| 0 | 背景 | 细胞质、黏液、杂质 |
| 1 | 正常细胞核 | 小且染色均匀,形态规则 |
| 2 | 低度病变核(LSIL) | 核略大,染色稍深,形态轻微不规则 |
| 3 | 高度病变核(HSIL) | 核明显增大,深染,形态显著不规则 |
注意这里有一个经常被忽略的点:核浆比这个关键诊断特征,不是靠把“细胞质”也分割出来算的,而是靠分割核本身、然后用核面积在同区域内占比间接估计。所以你不需要额外分割细胞质,否则标签工作量翻倍,分割难度也大幅增加,但下游特征并不直接依赖细胞质边界。
标签类别不要超过 6 类。病理数据天然长尾,类别越多,低度病变这类中间态类别的样本越少,迁移学习难度越大。如果你手上的涂片有 ASC-US(非典型鳞状细胞)这种“not sure”类别,建议先合并到相邻类别,等模型基础指标达标后再细分。
4.2 迁移学习策略:两阶段解冻 + 直推式伪标签
迁移学习在这个项目里有两个来源:一是 ImageNet 上预训练的 Swin 权重,二是同数据集其他涂片的半监督利用。前者解决训练收敛速度,后者解决标注不足。完整策略分三段走。
第一段:冻结编码器,只训练解码器。用 ImageNet 预训练权重初始化 Swin,把编码器所有参数挂到requires_grad=False,只更新 Unet 解码器部分。这时损失函数用上文的自适应尺度损失,学习率设为 1e-3。跑 30 个 epoch。这段的目的是让解码器先学会如何把 Swin 的特征图上采样成分割图,不急着调整特征提取器。
第二段:解冻编码器浅层,继续冻结深层。Swin 的前两个 stage 编码的是边缘、纹理、局部形态这类通用特征,对自然图像和病理图像都有用;后两个 stage 编码的是语义特征,自然图像和病理图像差异大,需要更多数据来适应。所以第二段只解冻 stage1 和 stage2,学习率降到 1e-4(编码器)和 1e-3(解码器),跑 20 个 epoch。
第三段:全量解冻,低学习率微调。全部参数requires_grad=True,整个网络学习率统一 5e-5,用 cosine 学习率调度器收尾。
直推式迁移学习(transductive 伪标签法)用在同数据集、无标注的涂片上。做法是拿当前训练好的模型,对无标注涂片做预测,只保留预测概率高于 0.9 的像素作为伪标签,把对应图像 patch 加入训练集再跑一轮第三段微调。注意伪标签要按类别分别设阈值,正常核的 0.9 阈值会导致大量低度病变核被过滤掉——低度病变在模型不确定时概率偏低,所以类别 2 的伪标签阈值要降到 0.8。
4.3 用 Unet 训练自己的数据集:数据格式与流程
训练数据格式建议直接用 PNG 图像 + PNG 标签掩膜,不要用 VOC JSON 或者 RLE 编码的中间格式。每个样本一对文件:img_001.png是原图,mask_001.png是单通道 PNG,像素值就是类别编号(0、1、2、3)。训练流程我通常写成四个独立脚本:prepare_data.py做格式转换和统计、train.py做训练循环、evaluate.py做验证集评估、inference.py做推理。下面是train.py里训练循环的核心片段:
for epoch in range(start_epoch, total_epochs): model.train() epoch_loss = 0.0 for images, masks in train_loader: images, masks = images.to(device), masks.to(device) sampler = AdaptiveScaleSampler() # 每 batch 动态计算尺度 images_aug, masks_aug, scales = [], [], [] for i in range(images.size(0)): img_np = images[i].cpu().numpy() msk_np = masks[i].cpu().numpy() img_aug, msk_aug, scale = sampler(img_np, msk_np) images_aug.append(img_aug) masks_aug.append(msk_aug) scales.append(scale) images = torch.stack(images_aug).to(device) masks = torch.stack(masks_aug).long().to(device) scales = torch.tensor(scales, device=device) outputs = model(images) ce = nn.CrossEntropyLoss()(outputs, masks) dice = ScaleAwareDiceLoss(n_classes=4)(outputs, masks, scales) loss = 0.6 * ce + 0.4 * dice optimizer.zero_grad() loss.backward() optimizer.step()注意一个容易翻车的细节:梯度裁剪。Swin 编码器在预训练状态下输出特征范数比较大,解码器刚初始化时梯度传导到编码器会出现梯度爆炸。我在训练循环里固定加一行torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=5.0),尤其第二段解冻编码器时必不可少。不加这一行,loss 大概率在前 10 个 iteration 变成 NaN。
5. 避坑与常见问题:从 loss 异常到验证指标虚高的 5 个排查点
5.1 loss 变成 NaN,且发生在解冻编码器之后
现象:第一段冻结编码器训练正常,第二段解冻 stage1 和 stage2 后,loss 在几十个 iteration 内突然变成 NaN,之后一直恢复不了。
原因:Swin 里的 LayerNorm 在混合精度(AMP)训练下,如果使用torch.cuda.amp.GradScaler,某些版本下 LayerNorm 的 fp16 计算精度不足,梯度数值溢出。加上解码器刚开始反向传播时梯度较大,超出 fp16 表示范围。
解决:两个方案任选。方案一是只对解码器使用 AMP,编码器强制 fp32,把forward中传给编码器的输入做.float(),并在autocast上下文外用with torch.cuda.amp.autocast(enabled=False):包住编码器调用。方案二是在 backward 之前加torch.nn.utils.clip_grad_norm_,把 max_norm 设为 1.0–5.0,同时把 GradScaler 的init_scale调小到 2**10。两个方案同时上最稳。
5.2 验证集 Dice 很高,但分割结果看起来“很怪”
现象:训练时验证集的 Dice 从 0.85 一路涨到 0.93,看起来很好;但打开推理图看,大核分割得很完整,小核几乎全丢了,或者病变核被分割成碎片。
原因:全图 Dice 被背景主导。宫颈细胞学图像里背景占 95% 以上,Dice 的分子在(2 * 交集)上,背景类别只要预测基本正确,整体 Dice 就被拉得很高。小核漏检对全图 Dice 影响极小,所以模型“学会了偷懒”。
解决:换成实例级指标。用 Aggregated Jaccard Index(AJI)和 Panoptic Quality(PQ)监控训练。AJI 的核心是把每个真实核作为一个实例,要求预测的核实例和真实核匹配后计算 IoU,单个核漏检会显著拉低 AJI。训练时虽然不能直接用 AJI 做损失,但每个 epoch 结束跑一次验证集算 AJI,用它来判断模型是否真的在进步。0.2 的 AJI 提升比 0.05 的 Dice 提升更有价值。
5.3 多尺度训练后推理时小核反而变差
现象:加了自适应尺度训练后,训练集上小核分割变好,但验证集(不经过多尺度)推理时小核召回率不升反降。
原因:自适应尺度训练把大核样本缩小、小核样本放大,模型在训练期间看到的“尺度分布”被过度拉平了。但验证集是原始尺度,很多小核的真实尺寸比训练时见过的最小尺寸还要小,模型对过小小目标的外推能力不足。
解决:推理时使用多尺度测试增强(TTA),把输入缩放到 1.0×、1.25×、1.5× 三个尺度分别推理,把三个尺度的输出概率图 resize 回原尺寸取平均。小核在 1.5× 尺度下相当于被放大了,模型更容易识别。这个技巧在训练和推理不对称时尤其有效。
5.4 加载 Swin 预训练权重报 size mismatch
现象:timm.create_model(..., pretrained=True)后用model.load_state_dict(state_dict, strict=False)加载时报一堆size mismatch for encoder.norm.weight之类的错误,或者加载成功但 loss 不下降。
原因:Swin 预训练权重的num_classes=1000,分类头的最后全连接层结构和你的分割任务完全不匹配。如果用strict=True直接加载必然报错;用strict=False虽然跳过不匹配的层,但如果 timm 的 Swin 输出特征是 (B, C, H, W) 形状,而预训练权重的最后是head层,特征提取部分基本没问题,可偶尔会遇到encoder.norm的 shape 不匹配——这是因为不同版本 timm 对norm层的定义不同。
解决:加载时过滤掉所有包含head和norm的关键词,只加载卷积和窗口注意力部分。通用做法是:
def load_swin_weights(model, state_dict): filtered = {k: v for k, v in state_dict.items() if 'head' not in k and 'norm' not in k} missing, unexpected = model.load_state_dict(filtered, strict=False) print('missing:', missing) print('unexpected:', unexpected)missing是解码器部分,本来就没有预训练权重,属正常;unexpected是 Swin 的 head 层,被我们主动过滤掉了,也正常。只要这两栏里不出现encoder.stage1.blocks.*之类的键,说明编码器加载是成功的。
5.5 类别严重不平衡时,低度病变类别 Dice 始终为零
现象:训练 50 个 epoch,类别 2(低度病变核)的 Dice 一直是 0,模型把所有像素都预测成背景或正常核。
原因:低度病变核在数据集中占比可能只有 2%–5%,如果 batch size 为 8,可能好几个 batch 里完全没有类别 2 的样本,CrossEntropy 的梯度被背景类完全淹没。这不是模型能力问题,是采样问题。
解决:在 DataLoader 里做类别平衡采样。具体做法是先统计每个样本的掩膜中各类别像素占比,把包含稀有类别的样本权重提高。我用WeightedRandomSampler,权重公式为weight = 1.0 / np.log(1.0 + rare_class_ratio),稀有类别占比越小,权重越高。注意不要用1.0 / rare_class_ratio,这个权重增长太陡,会让模型过拟合到几百张稀有样本上,用 log 形式平滑一些。
6. 多尺度推理与粘连核分离:验证和推理的 3 个实战技巧
6.1 多尺度 TTA 融合:概率平均比投票可靠
推理阶段的多尺度测试增强,常见做法是对同一张图做 0.75×、1.0×、1.25× 缩放,分别推理后取概率平均。注意不要对最终标签做投票,因为多类别标签的投票会造成小类别被多数类别淹没,例如在 0.75× 尺度下类别 3 因为核太小没被识别,投票时类别 1 会赢。正确做法是把每个尺度的 softmax 输出概率图都 resize 回原始尺寸后再做逐元素平均,最后argmax。
实现的细节是:resize 概率图时用BILINEAR插值,不能像掩膜一样用NEAREST,否则概率分布会被扭曲,导致本来 0.8 的置信度被插值成 0.7 和 0.9 的阶梯状分布。另外三个尺度的权重不需要都设为 1,我一般给中间尺度更高权重,因为它是训练时见过最多的尺度。例如 0.75× 权重 0.8、1.0× 权重 1.2、1.25× 权重 1.0。
6.2 粘连核分离:距离变换 + Watershed
细胞核分割的最终输出往往需要按实例区分,深度学习模型直接输出的语义分割图会把两个相邻的核连成一个连通域。常见的处理是用距离变换找核中心,再做 watershed 分离。这个方法不需要额外训练一个实例分割模型,适合病理图像里核密集的场景。
import cv2 import numpy as np from scipy import ndimage as ndi from skimage.feature import peak_local_max from skimage.segmentation import watershed def split_instances(binary_mask, min_distance=6): # binary_mask: 单个类别的二值掩膜,0/1 dist = ndi.distance_transform_edt(binary_mask) coords = peak_local_max(dist, min_distance=min_distance, labels=binary_mask) mask_peaks = np.zeros(dist.shape, dtype=bool) mask_peaks[tuple(coords.T)] = True markers, _ = ndi.label(mask_peaks) labels = watershed(-dist, markers, mask=binary_mask) return labelsmin_distance是最关键的参数。宫颈正常细胞核的直径通常 15–25 像素(在 40× 物镜下),相邻核中心距一般不小于 12 像素,所以min_distance=6是一个比较合理的下限。取值过小会把一个核内部的高原区域误检为多个种子点,把核切碎;取值过大会漏掉粘连核之间的边界,导致两个核分不开。建议在验证集上从 5 试到 10,用 AJI 指标选最佳值。
6.3 验证指标的正确计算方式
最后给一个新的习惯:训练结束后不要只看验证集 Dice,用下面这张表同时计算三类指标并对比:
| 指标 | 计算方式 | 用途 |
|---|---|---|
| mDice(类别平均) | 每个类别的 Dice 算术平均 | 粗略看整体 |
| AJI(Aggregated Jaccard Index) | 实例级匹配后加和 IoU | 判断核是否漏检/过分割 |
| PQ(Panoptic Quality) | 检测质量 × 分割质量 | 兼顾检测和分割的综合指标 |
其中 AJI 的计算并不复杂:把预测连通域和真实连通域做匹配,每个真实核最多匹配一个预测核,匹配条件是两个区域的 IoU 大于 0.5,然后把匹配成功的 IoU 总和除以所有真实核和所有未匹配预测核覆盖面积的总和。这个指标对小核漏检非常敏感,能真正反映“这个模型能不能用于临床辅助判读”。
一个经验:mDice 达到 0.85 而 AJI 只有 0.3–0.4,这是常见现象。说明模型“大致能分割”但实例级精度不够,问题往往出在粘连核和极小核上。先跑 watershed 后处理,AJI 通常能涨 0.1 左右,如果还提不上去,就要考虑是标注质量问题——让病理医生复核那些面积小于 15 像素的核标注,很多是标注漏标,不是模型问题。
我做这类项目收尾时有个习惯:最后一周不再改模型结构,只做阈值搜索和连通域参数搜索。把每个类别的置信度阈值从 0.3 到 0.7 按 0.05 步长扫一遍,同时把 watershed 的 min_distance 从 4 到 12 按 1 步长扫一遍,在验证集上选 AJI 最高的组合。这比多调 10 个 epoch 的效果更直接,也算给项目一个可复现的“后悔药”:模型权重归档后,这些后处理参数还能单独调优,不必重新训练。希望帮到你,祝你的细胞核分割模型早日达到能上病理辅助诊断的精度。
本文还有配套的精品资源,点击获取