尧图网络科技YAOTU DIGITAL 获取报价
获取报价
首页 / 资讯中心 / 文章详情

基于U-Net和Attention U-Net的医学图像分割系统实战解析

发布时间:2026/9/28 14:32:37

资讯中心
01
ARTICLE

基于U-Net和Attention U-Net的医学图像分割系统实战解析

基于U-Net和Attention U-Net的医学图像分割系统实战解析
简介这套基于U-Net与Attention U-Net的医学图像分割系统面向医学影像分析研发人员与学生可用于CT等图像的语义分割解决多类别组织或病灶的自动标注问题适合作为基准方法或科研实验起点。包体共14个文件以5个Python源码dataset、model、train、predict、utils为核心搭配7个pyc编译缓存及readme、txt说明文档压缩包约16KB目前已有117人学习下载。系统包含完整的数据预处理、模型训练、评估与预测流程数据处理模块支持自定义图像掩码路径、格式与尺寸提供随机翻转、CT窗宽窗位调整及灰度值映射以适配多分类标签模型架构实现了标准U-Net和带注意力门控的Attention U-Net内含卷积块、上采样与循环卷积块利用跳跃连接增强特征融合。训练过程采用余弦衰减与AdamW记录损失并计算Dice、IoU、精确率、召回率、F1等指标自动保存最佳模型与JSON日志并可视化训练曲线。预测模块支持单张图像分割并叠加显示整体代码模块划分清晰可直接运行或扩展。1. 基于U-Net和Attention U-Net的医学图像分割系统先把数据链路理顺拿到这份基于U-Net和Attention U-Net的医学图像分割系统源码包时我第一反应不是看模型结构而是先翻dataset.py。做医学图像分割这几年最深的体感是真正卡人的往往不是网络不够深而是CT图像进来之后窗宽窗位怎么调、灰度值怎么映射成多分类标签、图像和掩码尺寸怎么对齐。这些细节没人跟你说清楚项目就卡在“数据入口”。这套代码把数据处理、模型训练、评估、预测串成了完整闭环尤其适合在CT等医学图像上做语义分割课题或者落地验证的工程师。下面我按数据、模型、训练、预测的顺序逐段拆开讲中间会把参数、边界和踩过的坑都交代明白。2. 数据入口先立规矩dataset.py的CT窗宽窗位与多类别标签处理2.1 从文件路径到张量dataset.py的读取与图像尺寸对齐dataset.py的定位是自定义医学图像数据集常见做法是继承PyTorch的Dataset类在__init__里接收图像路径、掩码路径、格式和尺寸参数在__getitem__里完成读取、预处理和张量返回。这个包支持的路径格式灵活PNG、JPG这类常规格式都能读关键是图像与掩码的路径列表必须按相同顺序排序否则训练时输入和标签就错位了。# 核心伪代码dataset.py的初始化参数和数据读取 class MedImageDataset(Dataset): def __init__(self, image_dir, mask_dir, image_size(256, 256), image_formatpng, mask_formatpng, use_windowTrue): self.image_paths sorted(glob.glob(f{image_dir}/*.{image_format})) self.mask_paths sorted(glob.glob(f{mask_dir}/*.{mask_format})) self.image_size image_size self.use_window use_window assert len(self.image_paths) len(self.mask_paths) def __getitem__(self, idx): image cv2.imread(self.image_paths[idx], cv2.IMREAD_GRAYSCALE) mask cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) image cv2.resize(image, self.image_size, interpolationcv2.INTER_LINEAR) mask cv2.resize(mask, self.image_size, interpolationcv2.INTER_NEAREST) if self.use_window: image apply_window_level(image, window_width400, window_level40) image torch.from_numpy(image).unsqueeze(0).float() / 255.0 mask torch.from_numpy(mask).long() return image, mask重点看两个设计和参数掩码缩放强制用INTER_NEAREST不能用线性插值。线性插值会在类别边界产生灰阶渐变训练时模型会看到不存在的中间标签导致恢复出来的掩码边缘发花。图像除以255归一化到0到1掩码保持原始灰度值读进来后续再通过灰度值映射变成类别索引。注意掩码这里存的是long类型PyTorch的CrossEntropyLoss要求标签是整数索引不能是浮点类型。2.2 窗宽窗位调整同一张CT不同的窗口给出不同的对比度窗宽窗位WW/WL是CT图像特有的处理步骤。CT原始值范围很大常见的有-1024到3071直接归一化会把软组织、病灶的对比度压缩到几乎看不见。窗宽窗位的本质是把感兴趣的灰度区间映射到整个0到1范围超过区间上限的置1低于下限的置0。# 常见做法通过窗宽窗位做对比度增强 def apply_window_level(image, window_width400, window_level40): min_val window_level - window_width * 0.5 max_val window_level window_width * 0.5 image np.clip((image.astype(np.float32) - min_val) / (max_val - min_val), 0, 1) return image这里参数选择直接决定网络能学到什么腹部软组织常用窗宽400、窗位40肺部观察常用窗宽1500、窗位-600骨窗常用窗宽1500、窗位300。同一次扫描用不同窗宽窗位看完全是两种对比度。这个代码包把窗宽窗位放在数据预处理里说明设计者是拿真实CT数据在磨合而不是拿普通自然图像凑数。我一般会建议在dataset.py里把use_window做成开关并保留不同窗宽窗位配置。调试阶段可以多试几个窗口看哪个窗口下目标器官和病灶的边缘最清楚选定后再固定参数训练。这个切换成本很低但能直接影响分割精度尤其是在软组织对比度不高的场景。2.3 灰度值映射与随机翻转多分类标签进入网络前的最后一道转换医学分割数据集里掩码图通常是单通道灰度PNG不同组织用不同灰度值标记比如0是背景、85是某个器官、170是另一个器官。模型训练时不能直接拿这些原始灰度值当标签必须映射成0到类别数减1的整数索引。# 将原始灰度值映射为类别索引 def gray_to_label(mask, class_values[0, 85, 170]): label np.zeros(mask.shape, dtypenp.int64) for idx, value in enumerate(class_values): label[mask value] idx return label这段代码的逻辑是遍历class_values里每个灰度值在原mask上找到对应像素点赋值为类别索引。参数class_values是数组顺序决定了索引编号索引0对应class_values[0]索引1对应class_values[1]以此类推。两点注意class_values里不要漏掉背景灰度值0否则所有背景像素会被默认归到索引0但如果掩码里还有别的未列出的灰度值那些点会始终是0导致类别错乱。做随机翻转时图像和掩码必须用同一套变换参数。这个包支持随机翻转常用的是水平翻转。但我实际用下来水平翻转在某些器官分割任务上要慎重——比如左右对称性强的器官翻转后解剖方位就变了如果训练集和测试集方位分布不一致模型会学到错误的先验。3. 模型结构拆开看标准U-Net与Attention U-Net的注意力门控差异3.1 卷积块与上采样U-Net主干的编码器解码器骨架model.py里实现的主干是标准U-Net结构编码器逐层下采样提取越来越抽象的特征解码器逐层上采样把空间分辨率恢复回来跳跃连接把同尺度的编码器特征拼接到解码器弥补下采样丢失的空间细节。这个包把基础模块拆得很细核心是卷积块、上采样模块和循环卷积块。# 标准双卷积块U-Net中的基础特征提取单元 class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.block nn.Sequential( nn.Conv2d(in_ch, out_ch, kernel_size3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, kernel_size3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue) ) def forward(self, x): return self.block(x)这个卷积块的设计逻辑是两个3×3卷积堆叠等价于一个5×5卷积的感受野但参数量更少、非线性更强。padding1保证输入输出分辨率不变方便跳跃连接直接拼接。BatchNorm层对医学图像这种小batch训练尤其重要能稳定分布避免梯度震荡。上采样模块常见的做法是先转置卷积放大两倍再与对应编码器层的特征拼接。转置卷积的kernel size通常是2、stride是2输出尺寸正好翻倍。拼接之后再做一次双卷积融合让编码器和解码器特征真正交互。3.2 Attention Gate把跳跃连接从拼接升级为筛选Attention U-Net最大的改动是在每一条跳跃连接上插入一个Attention Gate注意力门。标准U-Net的跳跃连接是“无差别拼接”不管当前区域是背景还是目标器官编码器特征全量传给解码器。Attention机制做的事情是给特征图加一个空间权重让模型自己判断哪些位置的编码器特征应该被放大、哪些应该被抑制最终输出是加权后的特征而不是原始拼接结果。# Attention Gate给跳跃连接特征做空间加权 class AttentionGate(nn.Module): def __init__(self, F_g, F_l, F_int): super().__init__() self.W_g nn.Sequential(nn.Conv2d(F_g, F_int, kernel_size1), nn.BatchNorm2d(F_int)) self.W_x nn.Sequential(nn.Conv2d(F_l, F_int, kernel_size1), nn.BatchNorm2d(F_int)) self.psi nn.Sequential(nn.Conv2d(F_int, 1, kernel_size1), nn.BatchNorm2d(1), nn.Sigmoid()) def forward(self, gating, skip): g1 self.W_g(gating) x1 self.W_x(skip) psi self.psi(torch.relu(g1 x1)) return skip * psi核心逻辑在forward里gating来自解码器的上采样特征skip来自编码器同层特征。两个1×1卷积先把通道数统一到F_int然后相加、ReLU、再经1×1卷积压缩到单通道最后Sigmoid输出0到1之间的空间权重图psi。psi与skip逐元素相乘就是过滤后的跳跃连接特征。用大白话说Attention Gate学会了“哪里有用看哪里”腹部CT里肝脏区域权重高肠道气体区域权重低传给解码器的时候目标区域的特征被保留无关区域被压暗。这比标准U-Net更省容量也让解码器少学一些无用特征。实际使用时F_int一般取F_g和F_l的四分之一到一半参数过多反而容易过拟合。3.3 循环卷积块在有限分辨率下多挣一点感受野这个包里还有一个容易被忽略的模块循环卷积块。它的想法是在同一层内重复做多次卷积操作相当于把卷积层在时间维度上展开每次输出又作为下一次输入。循环次数t通常取2到3参数量不会翻倍但有效感受野能扩大不少。# 循环卷积块同一层内多次卷积扩大感受野 class RecurrentConvBlock(nn.Module): def __init__(self, in_ch, out_ch, t2): super().__init__() self.conv nn.Conv2d(in_ch, out_ch, kernel_size3, padding1) self.bn nn.BatchNorm2d(out_ch) self.t t def forward(self, x): out torch.relu(self.bn(self.conv(x))) for _ in range(self.t - 1): out torch.relu(out self.bn(self.conv(out))) return out注意这里用了一个残差式写法out bn(conv(out))目的是让循环加深时不至于梯度消失。第一次卷积把通道数从in_ch转到out_ch后续循环保持通道数不变。这个模块在CT图像上的实际收益是不需要堆到更深的网络就能让每个输出像素看到更大的上下文范围对器官边缘模糊的情况有一定帮助。我拆这个包的时候体会是model.py把两种网络放在同一套代码框架里切换成本就是构造函数里一个参数这一点对做对比实验非常友好。4. 训练流程的参数门道AdamW、余弦退火与Dice/IoU指标4.1 优化器与学习率调度医疗分割为什么倾向AdamW配合余弦衰减train.py里的训练配置优化器用的是AdamW学习率走的是余弦衰减。AdamW和经典Adam的区别在于权重衰减的实现方式AdamW把权重衰减从梯度更新中解耦单独作用于参数本身这样不会因为Adam的梯度二阶矩估计把正则项缩放得过小。对小数据集上的医学分割模型这一点区别在精调阶段能看出来尤其当数据集只有几十上百例时权重衰减能有效抑制过拟合。# 常见做法AdamW 余弦退火 optimizer torch.optim.AdamW(model.parameters(), lr1e-4, weight_decay1e-4) scheduler torch.optim.lr_scheduler.CosineAnnealingLR( optimizer, T_max100, eta_min1e-6)参数设定方面初始学习率1e-4适合大多数基于U-Net的分割任务太大容易震荡太小收敛慢weight_decay取1e-4到5e-4之间视数据量调整T_max是退火周期我一般让T_max等于计划训练的总epoch数这样学习率从1e-4平滑降到1e-6。余弦退火的思路是前期大步探索、后期小步精调比固定学习率或阶梯下降更平顺。如果发现验证集Dice在某个epoch后掉头往下走常见做法是调大weight_decay或提前让学习率降得更低而不是直接砍训练轮数。这个包把训练曲线和学习率衰减曲线都存了下来就是为了能判断是模型容量不够还是学习率没降到位。4.2 Dice、IoU与混淆矩阵多类别评估到底看哪个指标医学图像分割的评估不能只看准确率因为背景像素往往占绝大多数模型全预测成背景也能有很高的准确率。train.py里重点用Dice和IoU这俩指标对类别不平衡更敏感。# 计算Dice系数和IoU def dice_coef(pred, target, eps1e-7): pred pred.view(pred.size(0), -1) target target.view(target.size(0), -1) intersection (pred * target).sum(dim1) dice (2.0 * intersection eps) / (pred.sum(dim1) target.sum(dim1) eps) return dice.mean() def iou_coef(pred, target, eps1e-7): pred pred.view(pred.size(0), -1) target target.view(target.size(0), -1) intersection (pred * target).sum(dim1) union pred.sum(dim1) target.sum(dim1) - intersection return (intersection eps) / (union eps)Dice的计算逻辑是两倍交集除以两个集合的元素总数之和它的范围和IoU的关系是Dice 2×IoU / (1IoU)所以Dice数值上总是略高于IoU比较时要用同一种指标横比别拿Dice和IoU直接比大小。多类别分割时confuse_matrix.py维护一个类别数×类别数的混淆矩阵cm[真实类][预测类]记录计数。从混淆矩阵能算每个类别的精确率、召回率、F1分数精确率Precision 该类预测正确的数量 / 所有被预测为该类的数量。数值低说明误检多。召回率Recall 该类预测正确的数量 / 该类的真实总数。数值低说明漏检多。医疗场景里漏检往往比误检更危险所以看指标时我习惯优先关注每个类别的召回率尤其是病灶类别。全局准确率高不代表小目标分割得好这个包保存每个类别的指标曲线就是为了方便排查“是哪一类拖了后腿”。4.3 日志记录与模型保存断点恢复和训练曲线从哪来train.py会把每个epoch的损失、Dice、IoU、学习率写入JSON格式日志并自动保存最佳模型。最佳模型的判定逻辑一般是验证集Dice最高的一次而不是最后一个epoch因为后期学习率太低时指标容易反复波动。# 训练日志和checkpoint保存 log_entry { epoch: epoch, train_loss: round(train_loss, 4), val_dice: round(val_dice, 4), val_iou: round(val_iou, 4), lr: float(scheduler.get_last_lr()[0]) } with open(train_log.json, a) as f: f.write(json.dumps(log_entry) \n) torch.save({ epoch: epoch, state_dict: model.state_dict(), best_dice: val_dice, }, best_model.pth)checkpoint里除了state_dict还保存epoch和best_dice这样可以从断点恢复训练也能直接拿到当时的指标值。我一般会在这段逻辑后面补一个“日志重放”的代码块把JSON读取出来画训练曲线观察学习率下降和张量指标的关系比只看最后一个数值更能定位问题。5. 避坑指南从数据读取到注意力模型的五个翻车现场5.1 现象验证集Dice很高预测出来的掩码却对不上原始标签有次我拿这套流程跑一个三分类任务验证Dice在0.9以上但把预测结果还原成灰度图后肝脏区域整体往右偏了几个像素怎么调都调不回来。后来发现是预测时把类别索引直接乘以255当灰度值输出而原始标签里标的灰度是0、85、170这种离散值索引1乘255出来是255等于把标签体系换了一套。原因本质上是训练标签映射和预测恢复没有走同一套映射表。train.py里把掩码灰度值映射成索引predict.py里必须用完全相反的映射把索引还原成灰度。只是一头有映射另一头没有或者映射表顺序不一致都会出现“模型没错、数据端口错位”的假象。解决方式在utils.py里统一维护一张映射表训练和预测都从这张表取灰度值和索引的对应关系不要各写各的。我从那以后每次写分割预测代码都会先在原图上画出预测掩码和真实掩码叠放肉眼看一眼对齐情况。5.2 现象训练损失能降但恢复出的分割图全黑或整片同一类损失降得很顺但预测结果整片都是一个值通常不是模型没学好而是数据读取阶段就把标签搞坏了。常见原因是掩码文件用cv2.imread按三通道彩色图读入然后直接把整张图当单通道mask用或者把三通道压缩成灰度时把不同类别的灰度值混到了一起。检查方法在训练前先打印数据集中掩码的唯一灰度值集合用np.unique(mask)看一眼有哪些值和class_values里的值逐一比对。如果是三通道问题改成cv2.IMREAD_GRAYSCALE读取。还有一个隐藏点resize掩码时若用了默认的线性插值类别边界会出现中间灰度这类“伪标签”会让模型学到模糊边界最终输出整片混乱。解决就是前面强调的INTER_NEAREST。5.3 现象加载保存的模型时报错state_dict键不匹配用多卡训练后单独保存模型加载时经常报Missing key(s) in state_dict或size mismatch。根本原因是nn.DataParallel会把模型包一层state_dict里的键名前面多出module.前缀单卡加载时自然对不上。现象出现时错误信息里会列出缺失键和多余键一眼就能看出来是多卡前缀问题。解决方式是加载前处理键名# 兼容DataParallel保存的checkpoint checkpoint torch.load(best_model.pth, map_locationcpu) state_dict checkpoint[state_dict] if module. in next(iter(state_dict.keys())): state_dict {k.replace(module., ): v for k, v in state_dict.items()} model.load_state_dict(state_dict)这段代码的逻辑是检查state_dict第一个键是否带module.前缀有就去掉再加载这样单卡多卡切换都不会翻车。我一般会在utils.py里封装一个load_checkpoint函数统一处理这个问题。5.4 现象Attention U-Net在小数据集上反而比普通U-Net差我拿几十例CT数据跑对比实验Attention U-Net的Dice比标准U-Net低了将近0.03一开始以为代码写错了。后来把训练曲线拉出来看发现注意力模块的权重收敛很慢数据量太小Attention Gate学出来的权重近似于均匀分布等于白加了一层参数。原因在于注意力机制也是需要数据喂出来的数据量不够时它学不到“哪里是重点”反而增加了过拟合风险。解决方向通常有三个一是把学习率调低一点让注意力参数更稳定更新二是加强数据增强特别是随机翻转和窗宽窗位扰动相当于变相扩充数据三是先用普通U-Net训练一个合适的初始化再加载权重微调Attention版本。如果数据量实在小我建议回归标准U-Net不要为了用注意力而用注意力。这个包支持切换两种网络做对比实验时把“Attention不如普通U-Net”这种结论写进报告一样是有价值的。5.5 现象predict.py推理速度慢显存占用也不正常预测单张图像时显存占用接近训练水平速度也慢得离谱常见原因是模型处于训练模式。代码里漏掉model.eval()的话BatchNorm会继续用当前batch的统计量Dropout随机屏蔽还会生效结果就是每次预测都有随机性指标自然不稳定。解决方式很简单预测入口处统一加上这两行model.eval() with torch.no_grad(): logits model(image.to(device)) pred torch.argmax(logits, dim1)eval()让BatchNorm和Dropout切到推理模式torch.no_grad()关闭梯度计算显存占用能少三分之一以上。另外如果单卡显存吃紧可以减小输入尺寸或者用半精度推理model.half()配合image.half()在Pytorch 1.10以上版本常用速度能再快一截。6. 预测验证的最后一公里灰度还原与批量推理技巧6.1 批量预测的关键路径predict.py里预测单张图像的流程核心就三步加载checkpoint、网络前向、argmax取类别索引。批量预测时除了加一个循环还要保证每个batch都走eval和no_grad模式。argmax是对logits在通道维度上取最大值所在位置通道数等于类别数输出是单通道索引图。# 批量预测关键路径 model.eval() with torch.no_grad(): logits model(batch_images.to(device)) pred_indices torch.argmax(logits, dim1).cpu().numpy()torch.argmax(logits, dim1)的意思是第二个维度是类别维度每个像素取概率最高的那个类作为预测结果。得到的pred_indices是整数索引数组不能直接当成掩码输出还需要做灰度还原。6.2 灰度还原表把类别索引变回原始标签的约定predict.py在输出前有一段灰度值映射还原我的习惯是把它做成一张显式的对应表放在代码注释或配置文件里一目了然网络输出索引原始灰度值实际含义00背景185目标器官A2170目标器官B还原逻辑就是反向遍历这张表# 将类别索引恢复为原始灰度值 result_mask np.zeros(pred_indices.shape, dtypenp.uint8) for cls_idx, gray_value in enumerate(class_values): result_mask[pred_indices cls_idx] gray_value这段代码遍历每个类别索引找到对应像素位置赋值成原始灰度值输出的就是和训练标签同一套符号体系的掩码。可视化叠加时把掩码转成彩色或半透明叠加在原图上能直观判断分割边界是否贴合器官轮廓。这个包我拆完最大的感受是模型结构谁都能写真正决定交付质量的是数据处理和预测恢复这两端是否严格对得上。从那以后我每次拿到医学图像分割源码包都强制走一遍冷启动验证——从一张没见过的CT切片开始跑通训练和预测再拿预测结果和原图叠加、人工核对类别边界全部确认无误才敢说这个包能用。希望帮到你。本文还有配套的精品资源点击获取
02
RELATED NEWS

相关资讯

更多网站建设与数字化升级内容

03
WHY YAOTU

想打造同款高转化官网?

懂行业、懂生意,从建站到增长一站式陪跑

◈

场景化定制

不做模板站,围绕你的业务场景量身设计,小众不撞款。

◐

营销型架构

以转化目标组织内容与路径,让官网真正带来询盘。

▲

全周期服务

设计、开发、运营、运维一体,上线只是开始。

免费获取你的建站方案

留下需求,专属顾问 24 小时内为你输出方案建议。