ARTICLE · INTELLIGENCE

战地情报 · 详情页

来自尧图项目组的一线实战观察与深度解析

Swin Transformer Unet实战:医学影像分割从训练到推理全流程

Swin Transformer Unet实战:医学影像分割从训练到推理全流程 简介Swing Transformer Unet 源代码资源面向深度学习与计算机视觉方向的研究者、算法工程师及学生聚焦图像分割任务将 Transformer 的长距离依赖建模能力与 U-Net 的编码器—解码器结构相结合兼顾局部细节与全局上下文理解。压缩包共 227 个文件约 3.35MB以 141 个 png 图像数据、23 个 py 脚本为主另含 nbi/nbc 笔记本文件、pyc 缓存、xml 配置、m 脚本、md 说明及 yaml、csv、mat 等覆盖网络结构定义、训练与评估脚本、数据预处理及配置模块。该版本相较 GitHub 其他实现已做优化可直接运行省去环境调试时间便于快速复现训练、评估 IoU 与 Dice 等指标。目前已有 1806 人学习下载适合作为图像分割项目的研究起点与二次开发基础。1. 从一张 512×512 的医学影像说起Swin Transformer Unet 到底解决了什么如果你手头有一批 512×512 的医学影像或者遥感切片用经典 U-Net 跑分割大概率会遇到两个问题一是病灶边界模糊、细小结构丢失二是显存吃紧、batch size 只能开到 2。换成 Swin Transformer Unet 之后同样的数据边界连续性明显改善显存占用反而更可控。这不是玄学是窗口注意力机制带来的实际收益。Swin Transformer Unet 本质上是把 U-Net 的编码器-解码器骨架保留把里面的卷积块替换成 Swin Transformer Block。编码器逐级下采样提取多尺度特征解码器逐级上采样恢复分辨率中间用跳跃连接把浅层细节和深层语义拼在一起。区别在于Swin 的注意力是在不重叠的窗口内计算的窗口之间通过移位操作实现跨窗口信息交换计算复杂度从全局注意力的平方级降到线性级。这套结构适合谁做医学图像分割、遥感地物提取、工业缺陷检测的从业者尤其是那些数据量中等几千到几万张、标注成本高、对边界精度有要求的场景。如果你只是做简单的二分类分割经典 U-Net 够用但当你发现模型对细小目标不敏感、边界总是差几个像素时Swin Transformer Unet 值得一试。下面从环境搭建到训练调参把能直接运行的路径拆开讲。2. 把 Swin Transformer Unet 跑起来环境、代码结构与最小训练闭环2.1 环境依赖与版本锁定这套代码对 PyTorch 和 timm 的版本比较敏感我一般会锁死版本避免因为 Swin 实现差异导致加载权重失败。以下是经过验证的组合组件版本说明Python3.93.10 也稳定PyTorch1.13.12.x 部分算子有变动torchvision0.14.1与 PyTorch 配套timm0.6.13提供 Swin 预训练权重einops0.6.1张量重排opencv-python4.8.0数据增强与后处理安装命令如下pip install torch1.13.1 torchvision0.14.1 --index-url https://download.pytorch.org/whl/cu117 pip install timm0.6.13 einops0.6.1 opencv-python4.8.0提示如果你用的是 30 系或 40 系显卡cu117 对应 CUDA 11.7驱动版本建议 515 以上。版本不对会出现undefined symbol或no kernel image报错。2.2 代码目录结构与核心模块一份能直接运行的 Swin Transformer Unet 源码通常包含以下文件swin_unet/ ├── models/ │ ├── swin_transformer.py # Swin Block、Patch Embedding、Window Attention │ └── swin_unet.py # 整体网络组装 ├── datasets/ │ └── segmentation.py # 数据加载与增强 ├── utils/ │ ├── losses.py # Dice BCE 混合损失 │ └── metrics.py # IoU、Dice 计算 ├── train.py # 训练入口 ├── predict.py # 推理脚本 └── config.yaml # 超参数配置核心在swin_unet.py里编码器用 Swin Transformer 的四个 stage解码器用双线性插值加卷积。跳跃连接不是简单 concat而是先经过一个 1×1 卷积对齐通道数再和上采样特征相加。这个细节在源码里容易看漏但直接影响收敛速度。2.3 最小训练闭环从数据到第一个 checkpoint假设你的数据是图像和掩码成对存放目录结构为images/和masks/文件名一一对应。数据加载部分import cv2 import torch from torch.utils.data import Dataset class SegDataset(Dataset): def __init__(self, img_dir, mask_dir, img_size512, augmentTrue): self.img_dir img_dir self.mask_dir mask_dir self.img_size img_size self.augment augment self.names [f for f in os.listdir(img_dir) if f.endswith(.png)] def __getitem__(self, idx): name self.names[idx] img cv2.imread(os.path.join(self.img_dir, name)) img cv2.cvtColor(img, cv2.COLOR_BGR2RGB) mask cv2.imread(os.path.join(self.mask_dir, name), 0) img cv2.resize(img, (self.img_size, self.img_size)) mask cv2.resize(mask, (self.img_size, self.img_size), interpolationcv2.INTER_NEAREST) # 归一化到 [0,1]掩码二值化 img img.astype(float32) / 255.0 mask (mask 127).astype(float32) if self.augment: # 随机水平翻转 if np.random.rand() 0.5: img np.fliplr(img).copy() mask np.fliplr(mask).copy() img torch.from_numpy(img).permute(2, 0, 1) mask torch.from_numpy(mask).unsqueeze(0) return img, mask逻辑说明读取图像和掩码统一 resize 到 512×512。掩码用最近邻插值避免引入灰度值。归一化后转成张量通道顺序从 HWC 转成 CHW。随机翻转是最基础的增强医学图像里还可以加弹性形变但先跑通再说。训练循环的关键参数model SwinUnet(img_size512, in_chans3, num_classes1) optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-5) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_max50) criterion DiceBCELoss() for epoch in range(50): model.train() for img, mask in train_loader: img, mask img.cuda(), mask.cuda() pred model(img) loss criterion(pred, mask) optimizer.zero_grad() loss.backward() optimizer.step() scheduler.step() # 每个 epoch 后在验证集上算 Dice dice evaluate(model, val_loader) print(fEpoch {epoch}, Dice: {dice:.4f}) if dice best_dice: torch.save(model.state_dict(), best_swin_unet.pth)参数说明学习率 1e-4 是 AdamW 的常用起点如果 loss 震荡明显降到 5e-5。weight_decay 1e-5 防止过拟合。CosineAnnealingLR 让学习率从 1e-4 平滑降到接近 0比 StepLR 更稳。DiceBCELoss 是 Dice Loss 和 BCE Loss 按 1:1 相加前者管区域重叠后者管像素级分类。注意Swin Transformer 的预训练权重加载时如果num_classes不是 1000分类头会随机初始化这是正常的。编码器部分会加载 ImageNet 预训练参数解码器从头训练。3. 参数怎么调Swin Transformer Unet 的 4 个关键配置与显存优化3.1 窗口大小与输入分辨率的匹配关系Swin Transformer 的窗口大小默认是 7×7但这是针对 224×224 输入设计的。当你把输入改成 512×512特征图经过四次下采样变成 32×32窗口 7 在 32 上不能整除会触发 padding。padding 本身不影响正确性但会浪费计算。我一般会把窗口大小改成 8这样 512 输入下每层特征图尺寸都是 8 的倍数512→256→128→64→32窗口 8 在每一层都能整除。修改位置在swin_transformer.py的window_size参数# 原始配置 model SwinUnet(img_size512, window_size7, ...) # 调整后 model SwinUnet(img_size512, window_size8, ...)实测下来窗口 8 比窗口 7 在 512 输入上快 8% 左右Dice 基本持平。如果你的输入是 384×384窗口 6 更合适256×256 用窗口 8 也行但特征图最后一层只有 8×8窗口覆盖整张图退化成全局注意力显存会涨。3.2 嵌入维度与深度的取舍Swin Transformer 有四个 stage嵌入维度通常是 [96, 192, 384, 768]深度 [2, 2, 6, 2]。这是 Swin-Tiny 的配置参数量约 28M。如果你显存只有 8G建议把第一个 stage 的维度降到 64深度保持 2这样参数量降到 22M 左右Dice 掉不到 1 个点。# 轻量配置 embed_dim 64 depths [2, 2, 4, 2] num_heads [2, 4, 8, 16]如果显存 12G 以上可以用 Swin-Small 配置嵌入维度 [96, 192, 384, 768]深度 [2, 2, 18, 2]参数量 50M。但要注意深度加到 18 之后训练时间翻倍小数据集上容易过拟合。3.3 混合精度训练与梯度累积Swin Transformer Unet 在 512×512 输入下batch size 设为 4 时显存占用约 10G。如果你只有 8G 显存有两个选择一是用混合精度二是用梯度累积。混合精度开启方式from torch.cuda.amp import autocast, GradScaler scaler GradScaler() for img, mask in train_loader: with autocast(): pred model(img) loss criterion(pred, mask) scaler.scale(loss).backward() scaler.step(optimizer) scaler.update() optimizer.zero_grad()混合精度能把显存降到 6G 左右速度提升 20%30%。但要注意Dice Loss 在 fp16 下可能溢出建议在计算 loss 前把 pred 转成 fp32with autocast(): pred model(img) loss criterion(pred.float(), mask)梯度累积适合 batch size 想开大但显存不够的情况。比如你想用 batch size 8但只能开 2就累积 4 次梯度再更新accum_steps 4 for i, (img, mask) in enumerate(train_loader): pred model(img) loss criterion(pred, mask) / accum_steps loss.backward() if (i 1) % accum_steps 0: optimizer.step() optimizer.zero_grad()提示梯度累积时BatchNorm 的统计量是按实际 batch 算的不是累积后的。如果模型里有 BN 层建议换成 GroupNorm 或 LayerNormSwin 本身用的是 LayerNorm问题不大。3.4 学习率与损失函数的配合Swin Transformer Unet 的编码器有预训练权重解码器是随机初始化。如果全局用同一个学习率编码器可能被“带偏”。常见做法是分层设置学习率编码器用 1e-5解码器用 1e-4。encoder_params [] decoder_params [] for name, param in model.named_parameters(): if encoder in name: encoder_params.append(param) else: decoder_params.append(param) optimizer torch.optim.AdamW([ {params: encoder_params, lr: 1e-5}, {params: decoder_params, lr: 1e-4} ], weight_decay1e-5)损失函数方面Dice Loss 对类别不平衡不敏感BCE Loss 对边界像素更敏感。两者按 1:1 相加是稳妥起点。如果边界仍然模糊把 Dice 权重降到 0.3BCE 升到 0.7让模型更关注像素级分类。如果小目标漏检严重加一个 Focal Loss 分支权重 0.2 左右。4. 避坑与排查Swin Transformer Unet 训练中最容易翻车的 5 个地方4.1 现象loss 从第一个 epoch 就不降Dice 一直在 0.1 附近原因掩码值没有二值化或者归一化时把掩码也除了 255。Swin Transformer Unet 的输出是 logits如果掩码是 0255 的灰度值BCE Loss 会计算出巨大的梯度直接让模型崩溃。解决在 Dataset 里确认掩码只包含 0 和 1。用np.unique(mask)检查如果出现 0 和 255执行mask (mask 127).astype(float32)。另外图像归一化时只对图像做掩码不要除 255。4.2 现象训练到第 10 个 epoch 左右loss 突然变成 NaN原因混合精度下 Dice Loss 溢出或者学习率太大导致梯度爆炸。Swin Transformer 的注意力计算中有 softmaxfp16 下容易产生 inf。解决把 loss 计算强制转成 fp32或者在autocast上下文外计算 loss。如果已经出现 NaN把学习率降到 5e-5加梯度裁剪torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm1.0)4.3 现象验证集 Dice 比训练集低 20 个点以上原因过拟合。Swin Transformer Unet 参数量大小数据集上容易记住训练样本。另外如果训练集和验证集的分布差异大比如不同设备采集的影像也会出现这种情况。解决加数据增强除了翻转还可以加随机旋转 90 度、随机缩放 0.91.1、颜色抖动。如果增强后仍然过拟合把编码器冻结前两个 stage只训练后两个 stage 和解码器。验证集分布差异大的话考虑用域适应方法但那是另一个话题了。4.4 现象推理时显存爆了但训练时正常原因推理时没有用torch.no_grad()或者输入尺寸比训练时大。Swin Transformer 的显存占用和输入分辨率成正比512 训练、1024 推理显存直接翻四倍。解决推理脚本里加with torch.no_grad():并且把输入 resize 到训练时的尺寸。如果必须用大图用滑动窗口推理每次取 512×512 的块重叠 64 像素最后拼接。model.eval() with torch.no_grad(): for img, mask in val_loader: pred model(img) pred torch.sigmoid(pred) pred (pred 0.5).float()4.5 现象加载预训练权重时报错提示 key 不匹配原因timm 的 Swin 实现和源码里的 Swin 实现命名规则不同。比如 timm 用layers.0.blocks.0.attn.relative_position_bias_table源码可能用encoder.layers.0.blocks.0.attn.relative_position_bias_table。解决用model.load_state_dict(state_dict, strictFalse)然后打印缺失的 key 和多余的 key手动映射。常见做法是只加载编码器部分解码器随机初始化。如果 key 前缀不一致写一个字典推导式去掉前缀new_state_dict {k.replace(encoder., ): v for k, v in state_dict.items()} model.load_state_dict(new_state_dict, strictFalse)注意strictFalse 会静默忽略不匹配的 key一定要打印出来确认否则可能加载了错误的权重还不自知。5. 进阶技巧用滑动窗口推理和 TTA 把 Dice 再提 2 个点训练完之后推理阶段还有提升空间。我一般会做两件事滑动窗口推理和测试时增强。滑动窗口推理针对大图。假设你的测试图像是 2048×2048直接 resize 到 512 会丢失细节。做法是每次取 512×512 的窗口步长 256重叠 256 像素每个像素的预测值取平均。这样显存占用不变但能保留原始分辨率的信息。def sliding_window_inference(model, image, window_size512, stride256): model.eval() h, w image.shape[:2] pred_sum np.zeros((h, w), dtypenp.float32) count np.zeros((h, w), dtypenp.float32) with torch.no_grad(): for y in range(0, h - window_size 1, stride): for x in range(0, w - window_size 1, stride): patch image[y:ywindow_size, x:xwindow_size] patch torch.from_numpy(patch).permute(2,0,1).unsqueeze(0).cuda() pred model(patch) pred torch.sigmoid(pred).squeeze().cpu().numpy() pred_sum[y:ywindow_size, x:xwindow_size] pred count[y:ywindow_size, x:xwindow_size] 1 return pred_sum / np.maximum(count, 1)参数说明window_size 和训练时一致stride 一般取 window_size 的一半。重叠区域越大拼接越平滑但计算量也越大。2048×2048 的图stride 256 需要跑 49 个窗口耗时约 3 秒V100可以接受。测试时增强TTA是另一个技巧。对同一张图做水平翻转、垂直翻转、旋转 90 度分别推理然后把结果反变换回来取平均。TTA 能把 Dice 提升 12 个点代价是推理时间翻 4 倍。def tta_inference(model, image): preds [] # 原始 preds.append(inference(model, image)) # 水平翻转 preds.append(np.fliplr(inference(model, np.fliplr(image)))) # 垂直翻转 preds.append(np.flipud(inference(model, np.flipud(image)))) # 旋转 90 度 preds.append(np.rot90(inference(model, np.rot90(image)), -1)) return np.mean(preds, axis0)这两个技巧叠加在医学影像分割任务上Dice 从 0.82 提到 0.85 左右。但要注意TTA 对边界模糊的样本提升明显对已经分割很好的样本可能引入噪声。建议在验证集上先试确认有提升再上。最后说一个我踩过的坑滑动窗口推理时如果图像边缘不足一个窗口直接跳过会导致边缘区域没有预测值。我的习惯是先把图像 padding 到窗口大小的整数倍推理完再裁掉。这个细节在源码里通常不写但实际部署时一定会遇到。希望帮到你。本文还有配套的精品资源点击获取
RELATED READING

延伸阅读

更多一线实战笔记与深度复盘,助您持续精进