ARTICLE DETAIL

资讯详情

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

基于Transformer的皮肤病变分割:编解码结构与PyTorch实现解析

基于Transformer的皮肤病变分割:编解码结构与PyTorch实现解析 简介面向深度学习与医学图像处理的毕业设计场景提供一套基于Transformer实现语义分割的完整工程核心任务是皮肤病变区域分割。压缩包内含2000个文件共59.39MB其中以5447张JPG皮肤图像作为训练/验证数据集并配有像素级标注20个Python源码覆盖Transformer模型搭建、改进编码器-解码器结构、训练验证与评估流程2个.pth预训练权重便于直接部署或微调另有YAML配置、Markdown说明、PDF文档及分割结果可视化图。已有1054人学习使用。项目将Transformer自注意力机制与CNN融合针对皮肤病变图像中复杂纹理和细微差异进行优化可计算IoU、Precision、Recall等指标并支持更换数据集迁移到其他语义分割任务适合用于课程设计、论文复现或进一步研究。1. 皮肤病变分割为什么盯上Transformer在皮肤镜图像上普通UNet的Dice往往卡在0.78左右而把编码器换成Transformer后很多脱敏数据集上能明显推高到0.85以上。这个差距对临床来说就是漏检边界的位置变化。语义分割在医疗影像里的位置很特殊它既不是目标检测那样的框级粗粒度也不是分类的图级语义而是每个像素都要判断。皮肤病变分割就是典型病灶边缘颜色接近、形状不规则局部卷积很容易被纹理欺骗Transformer的全局自注意力可以缓解这类问题。这篇文章拆解一份名为Medical-Transformer-Improvement的工程包包含完整PyTorch代码、皮肤病变图像与像素级标注、训练配置文件。围绕编解码结构、数据集组织、训练评估闭环和最后的微创新把一整套可落地的流程讲清楚适合正在做毕业设计或想快速从CNN切到Transformer的研究生。2. 从CNN到Transformer编解码结构怎么搭2.1 为什么医疗图像分割需要全局注意力卷积神经网络在图像分割里统治了很多年。UNet、DeepLabV3实例都依赖局部感受野一层层堆叠虽然3x3卷积堆到十几层以后理论上能看到全局但实际操作中远距离信息要经过多层才能传播而且浅层特征里包含的全局上下文很少。对皮肤病变分割来说这不够用一张皮肤镜图像里病变区域可能只占几个百分点但判断它是良性还是恶性、边界落在哪里需要参考周围皮肤纹理、血管走向和肤色渐变而不是只看局部颜色。Transformer的核心是自注意力。理解Transformer模型详解会发现自注意力层里每个token都可以和所有其他token计算关联权重因此在一层内就能把视野拉满。Vision TransformerViT把图像切成一串patch token直接送入标准Transformer encoder。这套思路被移植到语义分割后最直接的表现就是长距离依赖处理得更干净。在项目中常见的做法是保留UNet的编码器-解码器整体框架但把编码器换成Transformer模块这样既能继承Transformer的全局建模又保留了解码器的上采样和跳跃连接能力。我拿到的这个工程包目录结构中可以看到model/encoder.py、model/decoder.py这类文件。从命名上看它做的正是类似TransUNet的改进CNN下采样加Transformer编码器再配一个逐步上采样的解码器。2.2 编码器侧Patch Embedding与位置编码Transformer不直接吃原始像素它吃的是patch序列。以16x16的patch大小为例一张512x512的图像会被切成1024个patch每个patch经过线性投影变成256维的token。代码实现并不复杂常见做法是直接用一个stride等于patch_size的卷积完成embedding。import torch import torch.nn as nn class PatchEmbedding(nn.Module): def __init__(self, in_channels3, embed_dim256, patch_size16): super().__init__() self.patch_size patch_size # 用卷积把每个patch投影成embedding向量 self.proj nn.Conv2d(in_channels, embed_dim, kernel_sizepatch_size, stridepatch_size) def forward(self, x): # x: [B, 3, H, W] - [B, embed_dim, H/patch, W/patch] x self.proj(x) # 展平把H/patch * W/patch压缩成长度为N的序列 x x.flatten(2) # [B, D, N] x x.transpose(1, 2) # [B, N, D] return x这里的关键参数是patch_size。patch越小序列越长计算量越大但细节保留得越好。在皮肤病变分割上我一般会用16如果显存不够再换成32。输入图像尺寸也直接影响序列长度例如512x512、patch16序列长度是32x321024这个长度对于Transformer encoder是完全可以接受的。embedding之后还要加位置编码否则模型分辨不出patch在图像中的位置。项目里常见的是用sine-cosine绝对位置编码或者干脆用可学习的nn.Parameter。在PyTorch中import math import torch import torch.nn as nn class PositionalEncoding(nn.Module): def __init__(self, max_len2048, embed_dim256, learnableTrue): super().__init__() if learnable: self.pos nn.Parameter(torch.randn(1, max_len, embed_dim)) else: self.register_buffer(pos, self._build_sincos(max_len, embed_dim)) def _build_sincos(self, max_len, d_model): pe torch.zeros(max_len, d_model) position torch.arange(0, max_len).unsqueeze(1).float() div_term torch.exp(torch.arange(0, d_model, 2).float() * (-math.log(10000.0) / d_model)) pe[:, 0::2] torch.sin(position * div_term) pe[:, 1::2] torch.cos(position * div_term) return pe.unsqueeze(0) def forward(self, x): # x: [B, N, D] return x self.pos[:, :x.size(1)]这里我更推荐可学习位置编码。它在训练集上可以自由调整对图像平移和尺度变化比固定编码更宽容对于皮肤病变图像里的随机裁剪场景尤其友好。注意max_len要大于实际序列长度一般给到2048足够。2.3 解码器从token序列回到像素掩码编码器输出是[B, N, D]的token序列要得到像素级分割图必须把它reshape回二维特征图再做多次上采样。比较省事的方案是先把token reshape成[B, H/patch, W/patch, D]转成通道在后然后用转置卷积或双线性上采样卷积把分辨率逐级抬高到原图大小。import torch import torch.nn as nn import torch.nn.functional as F class SimpleDecoder(nn.Module): def __init__(self, embed_dim256, num_classes1, img_size512, patch_size16): super().__init__() self.img_size img_size self.patch_size patch_size h, w img_size // patch_size, img_size // patch_size self.hw (h, w) self.head nn.Sequential( nn.Conv2d(embed_dim, 128, kernel_size3, padding1), nn.BatchNorm2d(128), nn.ReLU(inplaceTrue), nn.ConvTranspose2d(128, 64, kernel_size2, stride2), nn.ReLU(inplaceTrue), nn.ConvTranspose2d(64, 32, kernel_size2, stride2), nn.ReLU(inplaceTrue), nn.Conv2d(32, num_classes, kernel_size1) ) def forward(self, x): # x: [B, N, D] B, N, D x.shape x x.transpose(1, 2).reshape(B, D, *self.hw) # [B, D, H, W] x self.head(x) # 上采样到原图尺寸 x F.interpolate(x, size(self.img_size, self.img_size), modebilinear, align_cornersFalse) return x这个解码器比较朴素适合快速跑通。真正效果好的工程包会引入UNet风格的跳跃连接也就是把浅层CNN特征和Transformer输出在对应尺度上拼接。比如先用一个CNN stem生成1/4分辨率的特征图再把它和Transformer恢复出来的特征拼接这样能弥补Transformer丢失的局部细节。这也是工程包中“Improvement”这层含义的核心不是简单地用Transformer替代所有卷积而是让它和CNN特征互补。2.4 改进方向CNN Stem与分层注意力纯粹堆叠Transformer encoder并不总是最优。在医学分割上常见的几个改进思路是CNN stem先用几个卷积下采样到1/4再切patch分层注意力像Swin Transformer那样在不同stage间改变分辨率以及混合编码器浅层用CNN深层用Transformer。下面这张表列出了几个常见backbone在皮肤病变分割场景下的特点Backbone感受野计算量典型用途ViT-Base全局从第一层开始高长依赖优先适合大图Swin-Tiny窗口局部跨窗口低显存受限兼顾局部CNN stem ViT浅层局部深层全局中医疗分割常见改进TransUNet式全局多尺度跳跃中高复杂病灶边界恢复从经验上看皮肤病变中病毒疣、黑色素瘤等样本边界不规则窗口注意力如果窗口太小效果会打折。所以在工程包里我建议优先尝试CNN stem ViT的组合它既保留了卷积对边缘细节的敏感又引入Transformer对整体结构的目标建模。如果你机器显存不够大再考虑Swin-Tiny。3. 皮肤病变数据集的组织与预处理3.1 目录结构与读取方式工程包里的数据集一般按images和masks两个目录组织对应原图和像素级标注。标注通常是PNG灰度图背景为0病灶区域为1或255。还会有train.txt、val.txt、test.txt三个文本文件记录每个样本的名称或路径。这种组织方式比把所有数据塞进一个npz更直观也方便增量式训练。dataset/ ├── images/ │ ├── 0001.jpg │ └── ... ├── masks/ │ ├── 0001.png │ └── ... ├── train.txt ├── val.txt └── test.txt其中train.txt每一行是文件名比如0001不含扩展名。读取时用os.listdir或者按txt逐行收集。一般我会写一个函数读取这些列表再和images、masks拼接出完整路径。在皮肤病变数据上还需注意图像格式一些公开数据集会有_segmentation.png这类后缀要统一处理。3.2 自定义Dataset的实现在PyTorch中自定义Dataset是标准做法。代码里需要同时返回图像和对应的mask并且mask要与图像做相同的几何变换否则训练时错位评价指标直接崩。import os import cv2 import torch from torch.utils.data import Dataset class SkinLesionDataset(Dataset): def __init__(self, img_dir, mask_dir, file_list, transformNone): self.img_dir img_dir self.mask_dir mask_dir self.file_list [line.strip() for line in open(file_list)] self.transform transform def __len__(self): return len(self.file_list) def __getitem__(self, idx): name self.file_list[idx] img_path os.path.join(self.img_dir, name .jpg) mask_path os.path.join(self.mask_dir, name .png) image cv2.imread(img_path) image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) # 统一尺寸避免每张图像大小不一致带来的batch问题 image cv2.resize(image, (512, 512), interpolationcv2.INTER_LINEAR) mask cv2.resize(mask, (512, 512), interpolationcv2.INTER_NEAREST) # mask归一化到0/1 mask (mask 0).astype(float32) if self.transform: # 这里只做了非几何增强几何增强在外部定义 image self.transform(image) # 转成Tensor像素值归一到0-1 image torch.from_numpy(image.transpose(2,0,1)).float() / 255.0 mask torch.from_numpy(mask).unsqueeze(0).float() return image, mask这个类里有个容易踩坑的点cv2.resize对mask必须用最近邻插值INTER_NEAREST如果用了线性插值mask会变成0到1之间的浮点训练时的损失函数就会在边界附近产生错误梯度。另外mask 0这一步是为了兼容标注是255灰度值的情况把255统一成1。图像归一化到0-1但没有做标准化如果你用预训练backbone建议加载一个ImageNet均值方差做标准化否则预训练权重相当于没利用上。3.3 数据增强医学图像要克制语义分割的数据增强要同时作用于图像和mask。常见增强有水平翻转、垂直翻转、旋转、随机裁剪、亮度对比度微调。在医学数据集上我一般不用过于夸张的色彩扰动因为皮肤颜色本身就是诊断依据。增强的主要作用是提高模型对旋转和平移的鲁棒性而不是改变颜色语义。下面给出一个针对皮肤病变的增强策略表增强操作概率参数范围作用随机水平翻转0.5-增加对称不变性随机旋转0.60~30度适应临床拍摄角度差异随机缩放0.30.9~1.1倍模拟镜头远近变化亮度/对比度0.30.95~1.05轻度色彩扰动随机裁剪0.20.8倍区域增大有效样本数这些增强可以用albumentations库或者PyTorch的torchvision.transforms实现。注意在增强时mask和image的种子要一致也就是说同一张图像的随机参数绑定。用albumentations的Compose很方便因为它自动处理重采样和mask同步。如果不想引入额外依赖也可以用torchvision.transforms.functional自己写同步旋转但容易在mask上留下黑边需要配合随机中心裁剪一起用。3.4 训练/验证/测试划分与类别平衡数据集规模不大时按7:2:1或8:1:1划分即可。用一个简单脚本就可以完成划分。python -c import os, random ids [f.split(.)[0] for f in os.listdir(dataset/images)] random.seed(42) random.shuffle(ids) n len(ids) with open(dataset/train.txt,w) as f: f.write(\n.join(ids[:int(n*0.7)])) with open(dataset/val.txt,w) as f: f.write(\n.join(ids[int(n*0.7):int(n*0.8)])) with open(dataset/test.txt,w) as f: f.write(\n.join(ids[int(n*0.8):])) 这个命令会生成三个文本文件但严格看这里30%划分可能太小。更推荐的做法是在进入训练循环前统计mask中前景像素占比如果小于5%就属于严重类别不平衡。应对方法有两种一是用加权损失函数二是用torch.utils.data.WeightedRandomSampler给前景样本更高的采样权重。皮肤病变分割通常前景占比在10%上下Dice Loss本身就能缓解不平衡不必每次都用WeightedRandomSampler否则训练波动会变大。4. 训练评估闭环损失、指标和调参4.1 损失函数选择Dice Loss还是Focal Loss语义分割损失函数里最常用的三类交叉熵、Dice Loss、Focal Loss。交叉熵对每个像素独立计算在类别不平衡时会被多数类主导Dice Loss直接优化Dice系数天然对前景占比不敏感Focal Loss降低易分类像素的损失权重让模型更关注难例。在皮肤病变分割上我一般不会单独用Dice Loss因为它和Softmax的数值范围有冲突容易训练震荡。常见组合是Dice Loss加一个带gamma的Focal Loss或者Dice和交叉熵各占一半权重。下面是一段简单但实用的组合损失代码import torch import torch.nn as nn import torch.nn.functional as F class DiceFocalLoss(nn.Module): def __init__(self, focal_gamma2.0, dice_weight0.5, focal_weight0.5): super().__init__() self.gamma focal_gamma self.dice_weight dice_weight self.focal_weight focal_weight def forward(self, pred_logits, target): # pred_logits: [B, 1, H, W] 未经过sigmoid # target: [B, 1, H, W] 0/1 pred torch.sigmoid(pred_logits) b, c, h, w pred.shape target target.view(b, 1, h, w) # Dice Loss smooth 1.0 intersection (pred * target).sum(dim(2,3)) dice (2.0 * intersection smooth) / (pred.sum(dim(2,3)) target.sum(dim(2,3)) smooth) dice_loss 1.0 - dice.mean() # Focal Loss (二分类) pt pred * target (1 - pred) * (1 - target) focal_weight_tensor ((1 - pt) ** self.gamma).detach() bce F.binary_cross_entropy_with_logits(pred_logits, target, reductionnone) focal_loss (focal_weight_tensor * bce).mean() return self.dice_weight * dice_loss self.focal_weight * focal_loss这段代码里dice_weight和focal_weight控制两个损失的占比。我的经验是让两者各0.5如果发现训练中期Dice值卡住可以把focal_weight降一点让Dice主导。注意pt和预测概率相关但focal_weight_tensor这里用detach()这样梯度只作用于BCE部分实现更稳定。4.2 训练脚本核心逻辑项目工程包通常提供一个train.py内部核心循环并不复杂。需要关注的是优化器选择和调度器。Transformer模型对学习率极其敏感常见做法是用AdamW配合Warmup Cosine。常规的初始学习率可以设在1e-4到5e-5之间warmup步数占训练总步数的5%左右。下面给出一段训练循环的核心代码optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay0.01) scheduler torch.optim.lr_scheduler.CosineAnnealingLR(optimizer, T_maxepochs) for epoch in range(epochs): model.train() train_loss 0.0 for images, masks in train_loader: images, masks images.cuda(), masks.cuda() pred model(images) loss criterion(pred, masks) optimizer.zero_grad() loss.backward() optimizer.step() train_loss loss.item() scheduler.step() # 每个epoch后在验证集上计算IoU和Dice iou, dice validate(model, val_loader) print(fEpoch {epoch}: loss{train_loss/len(train_loader):.4f} f IoU{iou:.3f} Dice{dice:.3f})代码中用了CosineAnnealingLR能有效避免后期学习率过大导致的参数震荡。训练时一般先冻结预训练backbone只训练解码器几个epoch然后全量微调效果比一上来就全量训练好。也可以使用混合精度训练在PyTorch里就是torch.cuda.amp.autocast()和GradScaler()能省接近一半显存尤其适合Transformer这种大模型。4.3 评估指标IoU、Dice、Precision/Recall毕业设计最常被问到的就是IoU和Dice。IoU也叫Jaccard系数是交集比上并集Dice是两倍交集比上两个面积之和。代码实现时要小心平滑因子。官方比赛里常用不加平滑的版本但在训练损失中通常加平滑。下面是二分类指标计算的一个可复用函数def segmentation_metrics(pred_mask, true_mask, thresh0.5, smooth1e-6): # pred_mask: [B, H, W] 概率值true_mask: [B, H, W] 0/1 pred (pred_mask thresh).float() true true_mask.float() intersection (pred * true).sum(dim(1,2)) union pred.sum(dim(1,2)) true.sum(dim(1,2)) - intersection dice (2.0 * intersection smooth) / (pred.sum(dim(1,2)) true.sum(dim(1,2)) smooth) iou intersection / union return iou.mean().item(), dice.mean().item()在验证集上通常把验证集所有样本的预测mask按阈值0.5转成二值然后累计计算全局IoU。注意不要在计算时把批次平均和全局平均混淆。有些同学直接在循环里打印每个batch的IoU再取平均会受样本数量影响导致和最终测试集指标不一致。最好维护一个全局的intersection和union累加器在循环结束后一次性计算。4.4 踩坑记录与常见错误第一个坑是输入尺寸不一致。很多原始皮肤镜图像是1920x1080的高分辨率直接resize到512x512虽然能跑但细节丢失严重。我的经验是先用大尺寸如768x768训练如果显存不够再用patch_size32或者用随机裁剪配合多尺度训练。第二个坑是预训练权重加载不匹配Transformer的encoder通常加载ImageNet预训练权重但patch_embedding的卷积核大小不同会导致shape不匹配需要手动截断或随机初始化。第三个坑是类别不平衡导致验证Loss很低但视觉上漏检明显这时应该观察Precision和Recall而不只是看Dice。下面给出一个超参数参考表超参数推荐值说明输入尺寸512x512显存与精度的折中patch_size16细节要求高用16编码器层数8~12受训练数据规模限制学习率1e-4 (预热后)过高导致不收敛批大小83090Transformer显存开销大训练epoch80~120配合验证集早停如果训练时损失曲线没有下降先检查数据加载有没有问题mask是否与图像对应、是否归一化到0/1。再检查损失函数里的pred是否经过sigmoid。常见错误是把DiceFocalLoss里对pred做了一次sigmoid后面BCE又用binary_cross_entropy_with_logits导致梯度翻转模型学不动。调用validate时也要记得torch.no_grad()否则显存会被验证过程占满。5. 把课题做成可答辩的微创新5.1 替换编码器为Swin-Tiny我遇到的很多工程包默认用ViT-Base但对毕业设计来说ViT-Base训练太慢。更实用的做法是换成Swin-Tiny它会大幅降低显存压力并提升小病灶的边界效果。替换编码器的主要工作量在patch embeddingSwin用窗口注意力不需要绝对位置编码所以要把模型里self.pos_embed相关的部分去掉。如果你用的原版ViT是预训练权重可以直接保留其encoder层只修改patch_embedding输入通道和预测头的分类维度。替换后通常不需要从零重新训练用同一套数据集和损失函数就能收敛。5.2 推理时用TTA和后处理拉高指标测试时增强TTA是最划算的涨点手段。常见做法是推理时对原图、水平翻转、垂直翻转分别预测把三个概率图取平均再阈值0.5。下面是一段推理脚本def predict_with_tta(model, image, patch512): model.eval() batch torch.from_numpy(image.transpose(2,0,1)).float().cuda().unsqueeze(0) flip_lr torch.flip(batch, dims[3]) flip_ud torch.flip(batch, dims[2]) with torch.no_grad(): p0 torch.sigmoid(model(batch)) p1 torch.sigmoid(model(flip_lr)) p2 torch.sigmoid(model(flip_ud)) p (p0 p1 p2) / 3.0 return p[0, 0].cpu().numpy()推理之后再加一步简单的形态学后处理用3x3或5x5的核做一次开运算去掉孤立的假阳性小点然后取最大连通域作为最终mask。对于皮肤病变最大连通域假设是成立的因为病灶通常是连续的一片区域。这一步在验证集上经常能拿到0.5-1个点的IoU提升。5.3 模型预测结果的可视化与答辩展示毕业设计答辩时一张图胜过十句话。建议把原图、真实mask、预测mask、以及差值图横向拼接成一张对比图便于评委看到模型改进的地方。输出可以用下面这段简单代码import matplotlib.pyplot as plt def show_prediction(image, true_mask, pred_mask, save_pathoutput.png): plt.figure(figsize(12, 4)) plt.subplot(1, 3, 1) plt.imshow(image) plt.title(Input) plt.subplot(1, 3, 2) plt.imshow(true_mask, cmapgray) plt.title(GT) plt.subplot(1, 3, 3) plt.imshow(pred_mask, cmapgray) plt.title(Prediction) plt.savefig(save_path, bbox_inchestight) plt.close()可以用cv2.imwrite或matplotlib保存注意在保存前把pred_mask的0-1浮点值乘255转成0-255灰度图否则输出是全黑的。同时建议把几个经典困难样本的预测结果单独整理成文件夹作为答辩PPT的素材证明模型对边界模糊区域的处理能力。这比单纯列出IoU数字更有说服力。最后别忘了把训练日志和指标曲线导出成CSV方便做消融实验的对比图。本文还有配套的精品资源点击获取
返回列表