
简介一套基于Swin-Transformer与Unet的医学图像分割项目面向医学图像处理研究者、算法工程师及具备一定深度学习基础的开发者。项目针对子宫颈细胞核多类别分割任务融合迁移学习与自适应多尺度训练策略网络仅训练50个epochs即可达到全局像素准确率0.92、miou 0.767若延长训练周期性能仍有提升空间。资源包共809个文件约200MB包含391张jpg样本、383张png标注图、Python源码、训练权重pth、配置文件及README说明等覆盖从数据加载、多尺度增强到模型训练、指标记录、推理部署的完整链路。训练时自动将数据随机缩放至设定尺寸的0.5~1.5倍utils中的compute_gray函数自动读取mask灰度值并写入txt同时动态设置网络输出通道学习率采用余弦退火衰减训练日志记录各类别IoU、精确率、召回率及全局准确率。推理阶段只需将图片放入inference目录并运行predict脚本无需额外参数配置适合快速复现与二次开发。目前已有342人学习下载。1. 用 Swin-Transformer Unet 分割子宫颈细胞核这套代码把多尺度训练和迁移学习都备齐了做过医学图像分割的人都知道细胞核分割是个又基础又烦人的活目标小、边界模糊、染色差异大同一个数据集里细胞核大小能差出两三倍。我之前试过纯 Unet、也试过 DeepLabv3效果总是卡在某个点上不去。这个项目给的方案很直接Swin-Transformer 做编码器提取全局特征Unet 结构做解码器恢复细节配合自适应多尺度训练和迁移学习网络只训了 50 个 epoch全局像素准确度就跑到 0.92miou 到了 0.767。对想跑通一个能用的医学分割项目、或者正打算改进 Unet 做多类别分割的人来说这套代码值得拆开看一看——尤其是它处理 mask 灰度值、自动设置输出 channel 的那段逻辑很多人在自己数据集上翻车就翻在这里。2. Swin-Transformer 编码器 Unet 解码器网络结构是怎么拼起来的2.1 为什么选 Swin-Transformer 而不是纯 Unet传统的 Unet 靠卷积堆叠感受野每层卷积看到的是一个局部窗口要获得全局上下文得把网络挖得很深。而细胞核分割这个场景有个特点细胞核边缘往往依赖周围组织的信息来判断局部卷积容易把靠近的两个核粘在一起。Swin-Transformer 的移位窗口注意力机制W-MSA SW-MSA能在不把计算量推到平方级的前提下让每个位置看到更大范围的信息这对区分粘连核很关键。这个项目采用的是 Swin-Unet 的结构思路Swin-Transformer 当作编码器原封不动地输出多层特征图中间接一个 bottleneck解码器沿用 Unet 的上采样路径。换句话说它不是把 Transformer 和 Unet 简单并联而是把 Unet 里的每一层卷积编码块换成了 Swin-Transformer block。这样做的好处是编码部分能拿到全局关系解码部分依然保留 Unet 那种通过跳跃连接融合低层细节和高层语义的能力医学图像分割里最看重的边缘细节不会丢。2.2 编码器和解码器的关键实现我把这套结构的核心骨架整理成下面的伪代码方便理解每个模块在做什么。实际训练时不需要你手写这些但搞清楚结构对后面调参和改类别数很有帮助。class SwinUnet(nn.Module): def __init__(self, img_size224, num_classes2, embed_dim96): super().__init__() # Swin-Transformer 编码器4 个 stage逐级下采样 self.patch_embed PatchEmbed(img_sizeimg_size, patch_size4, in_chans3, embed_dimembed_dim) self.stage1 nn.Sequential(*[SwinTransformerBlock(embed_dim, num_heads3) for _ in range(2)]) self.stage2 nn.Sequential(*[SwinTransformerBlock(embed_dim*2, num_heads6) for _ in range(2)]) self.stage3 nn.Sequential(*[SwinTransformerBlock(embed_dim*4, num_heads12) for _ in range(2)]) self.stage4 nn.Sequential(*[SwinTransformerBlock(embed_dim*8, num_heads24) for _ in range(2)]) # bottleneck self.bottleneck SwinTransformerBlock(embed_dim*8, num_heads24) # Unet 解码器逐级上采样并用跳跃连接融合 self.up4 UpSample(embed_dim*8, embed_dim*4) self.decoder4 DecoderBlock(embed_dim*8, embed_dim*4) self.up3 UpSample(embed_dim*4, embed_dim*2) self.decoder3 DecoderBlock(embed_dim*4, embed_dim*2) self.up2 UpSample(embed_dim*2, embed_dim) self.decoder2 DecoderBlock(embed_dim*2, embed_dim) self.up1 UpSample(embed_dim, embed_dim) self.decoder1 DecoderBlock(embed_dim*2, embed_dim) self.seg_head nn.Conv2d(embed_dim, num_classes, kernel_size1) def forward(self, x): # 编码器输出 4 层特征尺寸逐级减半 x1 self.stage1(self.patch_embed(x)) x2 self.stage2(self.patch_merge(x1)) x3 self.stage3(self.patch_merge(x2)) x4 self.stage4(self.patch_merge(x3)) # 解码每一级上采样后拼接对应编码器特征 d4 self.up4(self.bottleneck(x4)) # 上采样恢复分辨率 d4 self.decoder4(torch.cat([d4, x3], dim1)) d3 self.up3(d4) d3 self.decoder3(torch.cat([d3, x2], dim1)) d2 self.up2(d3) d2 self.decoder2(torch.cat([d2, x1], dim1)) d1 self.up1(d2) d1 self.decoder1(torch.cat([d1, x1_patch], dim1)) return self.seg_head(d1)这里几个参数值得说明一下。embed_dim96是 Swin-Transformer 常见的小模型配置控制的是每个 token 的特征维度维度越大模型越胖、显存也越吃。patch_size4表示输入图像先被切成 4×4 的 patch 做 embedding这个值基本不用动。每个 stage 里的SwinTransformerBlock数量我是按两层写的实际项目里可以根据数据量往深了加但 50 个 epoch 的训练量对应两层是比较稳的组合。num_classes不需要手改后面会讲到 compute_gray 函数自动帮你设置。2.3 跳跃连接在这里和原始 Unet 有什么差别如果你照着上面的结构跑一遍会发现跳跃连接的 concat 维度处理和原始 Unet 不完全一样。因为在 Swin-Transformer 里每个 stage 的输出是(B, H*W, C)这种 token 序列要 concat 就得先 reshape 回(B, C, H, W)的图像特征格式再通道拼接。这个小细节在实现里很容易写错一旦忘了 reshapeshape 不匹配的报错会直接把你卡住。从实际效果看这种跳跃连接比纯 Unet 的优势在于Transformer 每一层的特征都已经含有一定范围的全局信息所以低层特征虽然分辨率高但不像纯卷积那样只盯着局部纹理融合出来的边缘更干净。代价是显存占用偏高尤其是输入分辨率大的时候。这个项目训练时用的输入尺寸不算大如果你的显卡显存只有 8G建议优先考虑把 batch size 调小而不是去砍 Transformer 的 depth。3. 多尺度训练与 mask 自动读取从数据管道到类别数自适应3.1 compute_gray 函数mask 灰度值是怎么变成类别数的很多人第一次跑这个项目的时候最懵的就是 mask 处理。医学分割数据集的 mask 图通常不是 PNG 索引色而是灰度图里面每个像素的灰度值代表类别编号比如背景是 0、细胞核是 1。问题在于不同数据集的标注习惯不一样有的从 1 开始标有的背景是 255有的中间有跳号。如果你写死num_classes2大概率会在某些数据集上踩坑。这个项目在utils里的compute_gray函数就是为了解决这个问题。它会把训练集所有 mask 的灰度值做一个全局统计自动去重、排序、剔除背景然后保存成 txt 文本同时把类别数量channel初始化给模型。def compute_gray(mask_dir, save_pathclass_names.txt): import cv2 import numpy as np gray_values set() for mask_name in os.listdir(mask_dir): mask cv2.imread(os.path.join(mask_dir, mask_name), cv2.IMREAD_GRAYSCALE) # 去掉全空的 mask避免干扰统计 if mask is None: continue # 找出这张 mask 里出现的所有灰度值 unique_vals np.unique(mask) gray_values.update(unique_vals.tolist()) # 剔除 0背景留前景类别按从小到大排序保证类别顺序稳定 class_values sorted([v for v in gray_values if v ! 0]) with open(save_path, w) as f: for v in class_values: f.write(str(v) \n) # 类别数 前景类别数 背景 num_classes len(class_values) 1 print(saved class values:, class_values, num_classes:, num_classes) return num_classes这段代码的逻辑很直白遍历所有 mask 文件用np.unique收集出现过的灰度值丢进 set 去重排序后写入 txt。这里有个容易被忽略的点用cv2.IMREAD_GRAYSCALE读图读出来的是 8bit 灰度值 0~255但如果 mask 存成的是三通道的彩色 PNG直接读灰度也能拿到正确的值所以这个写法兼容性比较高。我自己在实际使用时会在这一步之后顺手检查一下 txt 内容。常见情况是明明只有两类txt 里却出现了 0 和 255 两个值。这是因为有些标注软件用 0 表示背景、255 表示前景这种情况下需要把 255 映射成 1否则模型会被迫分出三个类别miou 直接掉一截。项目里对 0 和 255 的处理逻辑没有写死所以你的数据如果长这样建议提前做一个归一化映射。3.2 多尺度训练随机缩放 0.5~1.5 倍是怎么实现的自适应多尺度训练是这个项目另一个亮点。实现方式说起来很简单训练时每张图随机缩放到原尺寸的 0.5 到 1.5 倍之间再裁剪到固定尺寸喂进网络。好处是模型见过各种大小的细胞核推理时对分辨率变化不敏感泛化能力比固定尺度训练好不少。import random import cv2 def random_scale_image_and_mask(img, mask, scale_range(0.5, 1.5), target_size(224, 224)): h, w img.shape[:2] scale random.uniform(*scale_range) new_h, new_w int(h * scale), int(w *scale) # 缩放到随机尺度插值方式要区分图像和 mask scaled_img cv2.resize(img, (new_w, new_h), interpolationcv2.INTER_LINEAR) scaled_mask cv2.resize(mask, (new_w, new_h), interpolationcv2.INTER_NEAREST) # 随机裁剪到模型输入大小 crop_h, crop_w target_size if new_h crop_h and new_w crop_w: y random.randint(0, new_h - crop_h) x random.randint(0, new_w - crop_w) img_crop scaled_img[y:ycrop_h, x:xcrop_w] mask_crop scaled_mask[y:ycrop_h, x:xcrop_w] else: # 缩放后小于输入尺寸时做 padding img_crop cv2.copyMakeBorder(scaled_img, 0, crop_h-new_h, 0, crop_w-new_w, cv2.BORDER_CONSTANT, value0) mask_crop cv2.copyMakeBorder(scaled_mask, 0, crop_h-new_h, 0, crop_w-new_w, cv2.BORDER_CONSTANT, value0) return img_crop, mask_crop这里有两个必须注意的参数细节。第一scale_range的取值直接决定数据增强强度0.5~1.5 这个区间对细胞核任务足够激进但如果你的目标对象是器官这种大结构不建议设这么宽容易导致部分样本缩放后失去关键结构信息。第二图像用INTER_LINEAR插值mask 必须用INTER_NEAREST一旦你用线性插值去缩放 mask边缘会产生新的中间灰度值比如 0 和 1 之间冒出一个 127类别瞬间多出一堆损失函数直接算不对。3.3 Dataset 类的组织方式这个项目的数据组织方式和很多医学分割项目一致images文件夹放原图masks文件夹放对应的灰度 mask文件名一一对应。Dataset 类负责按索引读取图像和 mask把前面两个函数串起来。class SegDataset(Dataset): def __init__(self, img_dir, mask_dir, img_size(224, 224), use_multiscaleTrue): self.img_dir img_dir self.mask_dir mask_dir self.img_size img_size self.use_multiscale use_multiscale self.names [f for f in os.listdir(img_dir) if f.endswith(.jpg) or f.endswith(.png)] def __len__(self): return len(self.names) def __getitem__(self, idx): name self.names[idx] img cv2.imread(os.path.join(self.img_dir, name)) mask cv2.imread(os.path.join(self.mask_dir, name.replace(.jpg, .png)), cv2.IMREAD_GRAYSCALE) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) if self.use_multiscale: img, mask random_scale_image_and_mask(img, mask, target_sizeself.img_size) else: img cv2.resize(img, self.img_size) mask cv2.resize(mask, self.img_size, interpolationcv2.INTER_NEAREST) # 归一化到 [0,1]并转成 CHW 张量 img torch.from_numpy(img.transpose(2, 0, 1)).float() / 255.0 mask torch.from_numpy(mask).long() return img, maskuse_multiscale这个开关建议训练时打开、验证和推理时关闭否则推理结果每次都不一样没法稳定复现。mask转成long类型是 PyTorch 交叉熵损失的硬性要求如果保持 uint8 或者 float训练时大概率会报类型不匹配。3.4 多类别分割的通道自适应前面提到compute_gray会返回num_classes这个值最终要传给网络初始化。项目里 training 脚本会在启动时先跑一遍compute_gray然后用返回值去实例化 SwinUnet 的seg_head卷积层。这么做的好处是换数据集时不用改模型代码类别变了自动适配。但也有个边界情况如果你用的数据集类别编号是跳跃的比如只有 1 和 3没有 2compute_gray会把类别当成 [1, 3] 两个类别输出通道还是 2但 mask 里的像素值是 3喂给交叉熵损失时max_target超过num_classes-1直接报错。我一般会在统计完后做一步重映射把离散的灰度值连续化。4. 训练配置与迁移学习50 个 epoch 到 0.767 mIoU 的关键设置4.1 损失函数与学习率调度这个项目训练脚本默认用的是交叉熵损失配合 cos 学习率衰减。50 个 epoch 能跑到 miou 0.767这个成绩在细胞核分割任务里算是不错的水平背后有几个关键设置。首先损失函数没有一上来就上 Dice Loss而是用交叉熵。原因是这个项目只有两类前景占比不算极端失衡交叉熵足够稳定训练早期不会出现梯度震荡。如果你的数据集是那种前景像素占比极小的场景比如只有百分之几的息肉或肿瘤区域我建议在后期切换成 Dice Loss 或交叉熵加 Dice 的组合否则模型会倾向于把所有像素预测成背景。criterion nn.CrossEntropyLoss() optimizer torch.optim.AdamW(model.parameters(), lr3e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max50, eta_min1e-6)学习率初始值这里有一个很实际的经验Swin-Transformer 这类结构对学习率比纯卷积网络敏感3e-4 是 AdamW 搭配 Transformer 的常见起点用 1e-3 大概率会发散用 1e-4 收敛太慢。T_max50表示 50 个 epoch 内 cos 曲线从 3e-4 衰减到 1e-6。如果你加了预训练权重这个初始学习率可以适当降到 1e-4因为预训练特征已经比较成熟学习率太大会把原有特征冲掉。4.2 迁移学习Swin-Transformer 预训练权重的加载方式这个项目里的迁移学习主要发生在编码器部分。Swin-Transformer 在公开数据集上预训练过拿来初始化分割模型的编码器可以让模型在医学图像数据量不足的情况下依然学到合理的基础特征。加载预训练权重时有个常见的坑Swin-Transformer 的head是用于 ImageNet 分类的 1000 类全连接层而分割模型的seg_head是一个我们自己定义的卷积层直接加载整个 checkpoint 会报 key 匹配错误。常见的正确方式是def load_pretrained_encoder(model, checkpoint_path, num_classes2): # 严格False只加载能匹配的 key跳过分类头和分割头 pretrained torch.load(checkpoint_path, map_locationcpu) if state_dict in pretrained: pretrained pretrained[state_dict] model_dict model.state_dict() # 过滤掉不需要的 key pretrained_dict {k: v for k, v in pretrained.items() if k in model_dict and seg_head not in k and decoder not in k} model_dict.update(pretrained_dict) model.load_state_dict(model_dict) print(loaded keys:, len(pretrained_dict)) return model这里最核心的逻辑就是过滤 keyseg_head和decoder的权重不参与加载只加载 Swin-Transformer 编码器部分的权重。第一次做这个操作时建议打印出加载了多少个 key如果数量只有几个说明你的 key 命名和预训练权重不一致需要检查model.state_dict()里编码器部分的命名。4.3 训练日志和评估指标的读法项目在run_results里存了训练日志和曲线图包括训练验证的 loss 和 iou 曲线、每类别的 iou、recall、precision、全局像素准确率。我在拆这个项目时特别注意了它的日志结构每个 epoch 结束会打印一行指标格式大致是Epoch[50/50] loss: 0.1234, miou: 0.767, acc: 0.92, cls1_iou: 0.81, cls1_recall: 0.87, cls1_precision: 0.89。读这里有个小技巧不要只盯着 miou要单独看每一类的 recall 和 precision。细胞核分割场景最容易出现一种病大类细胞核的 iou 不错但小类别或者边界区域 recall 偏低。如果某类的 precision 高但 recall 低说明模型预测保守偏检不准反过来 recall 高于 precision 意味着过度分割。这个项目 50 个 epoch 的全局准确度 0.92说明背景像素分类得比较准但 miou 只有 0.767差距主要来自前景边界的不确定性。4.4 epoch 数量对性能的影响项目摘要里说训练 epoch 加大性能还会更优越这个说法我基本认同。Swin-Unet 这类模型收敛速度比纯 Unet 慢因为 Transformer 部分的参数更新需要更多轮次来稳定注意力矩阵。从我的经验看50 个 epoch 时模型可能还没完全收敛尤其是用 cos 衰减时最后 10 个 epoch 学习率已经降得很低模型还在做细微调整。如果你有自己的显卡和时间预算把 epoch 加到 150~200配合早停策略miou 通常能再涨 3~5 个点。不过这里要提醒一句加 epoch 不加数据增强模型大概率会过拟合。项目里已经有多尺度缩放这个增强再加随机旋转和水平翻转会更稳。我不建议直接改大 epoch 然后不加任何正则化就去跑这样训练集 loss 会很好看测试集指标反而可能变差。5. 避坑清单跑这套代码最容易翻车的五个地方5.1 现象训练第一个 epoch 后 loss 直接变成 nan原因mask 的灰度值和compute_gray统计结果不一致。最常见的情况是某张 mask 里混入了一个没见过的灰度值比如标注时误用了 255 填充模型输出通道数与之不匹配交叉熵在计算时遇到非法目标值。解决compute_gray跑完后手动打印一下 txt 内容确认所有灰度值都在预期范围内。发现 255 这种异常值在数据预处理阶段做一次映射把 255 改成 1。另外训练时在 Dataset 的__getitem__里对 mask 做一个 clamp确保目标值不超过num_classes - 1能兜住偶尔的脏数据。5.2 现象损失下降但验证集 miou 卡在 0.3 左右不动原因这大概率是多尺度训练开启后验证集没有做同样预处理导致训练和推理的输入分布不一致。还有一种可能是验证集的 mask 用的是线性插值缩放边缘产生中间灰度值评估时这些像素全部算错。解决验证时把多尺度缩放关闭统一 resize 到固定尺寸mask 用INTER_NEAREST。如果 miou 仍然低去跑一遍预测结果可视化看看是不是模型把所有前景都预测在图像中心附近——如果是检查 padding 是不是用了 0图像黑边被模型当成了背景先验。5.3 现象加载预训练权重报 shape mismatch原因Swin-Transformer 的官方预训练权重是在 ImageNet 分类任务上训练的最后的fc层输出维度是 1000与原模型不一定匹配另外如果你的输入通道不是 3比如用了单通道灰度图第一层的in_chans也不匹配。解决加载时strictFalse过滤掉head、fc、seg_head相关的 key。如果输入是灰度图要处理patch_embed第一个卷积层的权重把in_chans改成 1然后将预训练权重的 RGB 三个通道取平均复制到单通道上。5.4 现象训练时显存溢出batch size 设为 2 都跑不动原因Swin-Transformer 的注意力计算是平方级复杂度特征图分辨率越大显存消耗越夸张。这个项目默认输入是 224 或 256如果直接跑到 512 分辨率8G 显存基本撑不住。解决优先把输入分辨率降到 192 或 160这个改动对细胞核分割的影响不大因为核本身是小目标。另外可以把patch_merge里的卷积 stride 调大但工程上更简单的做法是开梯度累积accumulation_steps 4每 4 个 batch 更新一次梯度等效于把 batch size 撑到 8显存占用不变。5.5 现象推理结果保存后是一张全黑图原因预测输出是(B, C, H, W)的 logits很多人直接用torch.max取索引后没有把类别索引映射回可显示的灰度值。如果类别索引从 0 开始单元格索引是 1保存成 8bit PNG 时 1 的像素几乎看不见看起来就是黑的。解决推理保存时把预测索引乘 255 再转 uint8或使用调色板模式保存。我通常会把每个类别的预测单独染色成不同的颜色这样一眼就能看出分割边界在哪比灰度图直观得多。6. 推理脚本与进阶改进从直接出图到多类别后处理6.1 predict 脚本的用法和内部逻辑项目在 README 里明确写了推理方式把待推理图像放在inference目录直接运行predict脚本无需设定参数。这个设计对新手非常友好因为脚本内部已经把所有预处理和后处理封装好了。我拆了下它的流程大致是扫描目录下所有图片逐张预测把 logits 通过argmax转成类别索引再按类别映射到灰度值保存到输出目录。def predict_single_image(model, img_path, device, img_size(224, 224)): img cv2.imread(img_path) img_rgb cv2.cvtColor(img, cv2.COLOR_BGR2RGB) resized cv2.resize(img_rgb, img_size, interpolationcv2.INTER_LINEAR) tensor torch.from_numpy(resized.transpose(2, 0, 1)).float().unsqueeze(0) / 255.0 tensor tensor.to(device) with torch.no_grad(): logits model(tensor) pred torch.argmax(logits, dim1).squeeze(0).cpu().numpy() return pred, img.shape[:2]推理结果如果要恢复原图大小直接cv2.resize预测的 mask注意插值方式必须是INTER_NEAREST。如果推理时输入的尺寸和训练时的多尺度范围差得太多建议先在原图上切块预测再拼接而不是直接整图 resize否则细胞核纹理细节会被压没。6.2 测试时增强TTA的简单实现多尺度训练对应的推理加强手段就是多尺度推理同一张图分别缩放到 0.8、1.0、1.2 倍预测把得到的概率图取平均再 argmax。这个操作通常能让 miou 再涨 1~2 个点。代价是推理时间翻倍但对离线评估场景很划算。def predict_with_tta(model, img, scales[0.8, 1.0, 1.2], img_size(224, 224)): prob_sum None for scale in scales: h, w int(img.shape[0] * scale), int(img.shape[1] * scale) scaled_img cv2.resize(img, (w, h)) prob predict_prob_map(model, scaled_img) # 返回 softmax 概率图 prob cv2.resize(prob, img_size, interpolationcv2.INTER_LINEAR) if prob_sum is None: prob_sum prob else: prob_sum prob return np.argmax(prob_sum / len(scales), axis0)6.3 多类别分割的边界改进思路如果你拿这套代码跑自己的数据想要进一步提精度我建议先做三件事第一把交叉熵换成交叉熵加 Dice Loss 的混合损失边界像素的召回通常能改善 2~3 个点第二增加随机旋转和水平翻转数据增强Swin-Transformer 本身对旋转不敏感但多尺度加翻转的组合对细胞核这种方向多变的目标非常有效第三训练结束后用条件随机场CRF做一步后处理把预测概率和像素颜色信息结合起来剪掉孤立小区域边界会更干净。6.4 换数据集时的检查清单最后分享一个我自己的教训。第一次把这份代码迁移到别的数据集时我没有检查 mask 类别编号结果训练出来的模型把所有像素都预测成背景。后来我养成了一个习惯每次拿新的数据集跑这套代码都强制走一遍四步检查第一步跑compute_gray看类别 txt第二步用脚本抽查三张 image-mask 对确保 mask 里每个灰度值对应的区域是合理的第三步训练 5 个 epoch 后立刻看验证集预测可视化第四步确认推理脚本里的归一化方式和训练一致。这套流程几乎成了我跑所有医学分割项目的固定动作希望也能帮你省掉一些来回折腾的时间。本文还有配套的精品资源点击获取