十年匠心定制 · 商业建站与技术教学双线并行 咨询热线:400-886-1026 service@lmnt.cn
ARTICLE DETAIL

资讯详情

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

Y-Net双编码器图像分割网络:结构解析与PyTorch实现

Y-Net双编码器图像分割网络:结构解析与PyTorch实现 简介面向体内光声成像的 Y-Net 开源实现是一套基于 U-Net 改进的深度学习重建框架专门解决传统光声重建方法计算量大、易产生伪影的问题。资源共 9 个文件以 8 个 Python 源码文件为主涵盖网络模型、数据加载、训练入口等模块另有 1 份 README 说明文档压缩包仅 14KB轻量紧凑结合说明文档可快速上手便于模块化阅读和二次开发。目前已有 499 人学习下载适合有一定深度学习基础、希望将神经网络用于医学图像重建的研究人员和开发者。通过源码可学习 Y-Net 的模型设计与训练流程理解卷积、上采样与跳跃连接如何捕捉多尺度信息完成光声信号到图像的重建也能将框架迁移到 OCT、超声成像等相近领域或调整参数、扩展网络结构为课题研究提供参考。1. Y-Net项目概述与源代码学习价值1.1 Y-Net是什么Y-Net这个名字我第一次看到的时候第一反应是“又在U-Net后面加了个字母”。但仔细把网络结构图展开后会发现它确实长得像个Y字两条编码分支各自独立地从输入中抽取特征在底部汇合再共用一条解码路径生成最终结果。这种结构天然适合两类场景一类是输入本身就有多个模态或多个视角的任务比如医学影像里的T1加权像和T2加权像同时输入另一类是同一个输入需要拆成不同表征来学习比如原始图和对应的边缘响应图一起进网络。源代码学习的重点恰恰是理解这两条分支为什么存在、它们从哪里分岔、又在哪里汇合——搞清楚了这个问题其他模块基本都是U-Net那套熟练工。市面上叫Y-Net的实现并不少不同论文和开源仓库里的细节差异还挺大。我这次用的版本是以双编码器加共享解码器为骨架、在融合处加了一个简单的注意力门控的PyTorch实现。这套代码适合正在做图像分割、医学影像分析、以及想学习多分支模型如何组织训练逻辑的开发者。它能解决的核心问题很简单当单一输入无法完整描述目标的边界与上下文时怎么用两条并行的编码路径把互补信息揉进同一个分割结果里。1.2 这套源代码适合谁来读如果你是第一次接触这种更复杂的分割网络我建议不要把全部精力放在凑指标上先跟着代码把数据流走一遍。这份源代码比较适合三类人一是熟悉U-Net但想了解多分支扩展的初学者二是需要处理多模态输入的算法工程师三是想借鉴双编码器结构做特征融合的科研党。读代码的时候优先看数据加载器和模型forward里张量尺寸的变化。Y型结构最怕两个分支的feature map对齐不上很多实现写出来跑不通十有八九都死在这个地方。后面我会把每一层的输入输出尺寸都列出来方便你对照自己的数据做调整。提示如果你想直接跳到自己动手复现的部分可以先看第4节的依赖环境想弄明白设计原理就先看第2节。两者不冲突只是阅读顺序不同。2. 网络结构设计思路拆解2.1 双编码器的本质分支做的不是重复劳动很多人看到双编码器第一反应是“参数翻倍了网络是不是更重了”。其实这个理解并不准确。两条编码分支之所以存在往往是为了处理不同性质的输入而不是为了把同一个输入跑两遍。以我的使用场景为例一条分支输入原始灰度图像另一条分支输入Canny边缘图。原始图给网络提供区域纹理和灰度对比度信息边缘图则直接告诉网络哪里是边界、哪里有结构断裂。这种分工其实很像人类医生看CT片子时的操作先整体观察器官轮廓再放大看局部细节最后把两方面的印象合在一起做判断。如果只用一条编码器去同时承担这两种任务模型内部可能需要更多的层才能学到等价的特征解耦反而不容易收敛。从代码层面看两条编码器是完全独立的权重复制品吗并不必然。有的实现会做权值共享或半共享只在输入头部做分支有的实现则让两个分支完全独立。我建议初期学习时选完全独立的那一版至少训练曲线更直观出现问题也好排查。等跑稳定了再尝试共享部分权重来降参数。2.2 特征融合与跳跃连接怎么保住细节编码器一路向下采样特征图分辨率越来越低语义信息越来越强但空间细节也在同步丢失。U-Net的经典解法是跳跃连接把编码器每一层的输出拼到解码器对应的层去相当于把高分辨率浅层特征直接递到解码阶段。Y-Net在这个基础上多了一个问题——两个编码器每层都有输出是两条都拼还是只拼一条不同实现选择不一样。我用的这份代码把两个编码器的浅层输出都通过跳跃连接送入解码器用concat操作叠在一起然后接一个1x1卷积把通道数压回来。这样做的好处是两头的信息都不会丢缺点是解码器第一层的输入通道数会比较夸张。假设单分支编码器第3层输出是128通道两个分支拼接后就是256通道再加一个跳跃连接解码端通道压力会明显变大。为了解决通道膨胀一些实现会在融合后加1x1卷积降维或者在跳跃连接前加额外的attention模块做权重筛选。我实测下来直接用1x1卷积降维最省事而且效果差距不大。真正的瓶颈反而不在通道数而在训练时两个分支是否收敛得均匀——这个坑留在第4节细说。2.3 为什么选择Y型而不是双解码器这是我在学习过程中反复问过自己的问题。既然有两个输入为什么不干脆做成两个独立的编码-解码网络最后再把预测结果融合这样不是更省心吗答案是分割任务不是单纯地做“两张图的加权平均”。两个模态或两种特征之间可能存在强耦合关系比如原始图的某些纹理说明这里是血管边缘图中的某条闭合曲线也佐证这是一个完整结构。如果两套网络各跑各的融合阶段就只能看到最终概率图它们之间是否有过深层的交互网络并不关心。Y型结构通过共享解码器强制两个分支的特征在底层完成信息交换后再一起向上恢复分辨率相当于逼着模型在早期就开始整合不同来源的证据。这种设计在参数量上其实比两个独立网络要省的因为解码器只有一套。缺点是底层融合之后如果某个分支特征质量较差会直接污染整个解码阶段。所以在训练策略上有的实现会先单独预训练两个编码器再联合训练解码器。我建议至少在前几个epoch观察两个分支的loss下降趋势再做调整。3. 源代码核心模块逐项解析3.1 数据加载与增强两个输入一个标签数据加载器是这个项目里最不起眼却最容易出错的部分。Y-Net的每个训练样本包括input_a、input_b和一整张对应的mask。在医学图像场景里input_a常是原始灰度图input_b是Canny边缘图或梯度幅值图在其他场景里也可以换成RGB图和深度图、或者两张不同模态的配准图。读取数据时要留意两点。第一input_a和input_b必须做完全相同的几何变换比如旋转、翻转、缩放否则两张图的空间位置就对不上了。用albumentations时通常是定义一个transform pipeline然后对a和b分别调用同一个增强对象或者用支持多输入的接口统一处理。第二Canny边缘图是在原图上算出来的做增强之后再计算更准确因为旋转和缩放会改变边缘形态先算好再变换会引入重采样噪声。下面是我常用的一段核心代码结构import cv2 import albumentations as A from torch.utils.data import Dataset class DualInputDataset(Dataset): def __init__(self, image_paths, mask_paths, trainTrue): self.image_paths image_paths self.mask_paths mask_paths self.train train self.aug A.Compose([ A.HorizontalFlip(p0.5), A.VerticalFlip(p0.2), A.RandomRotate90(p0.3), A.RandomBrightnessContrast(p0.2), ]) def __getitem__(self, idx): img cv2.imread(self.image_paths[idx], cv2.IMREAD_GRAYSCALE) mask cv2.imread(self.mask_paths[idx], cv2.IMREAD_GRAYSCALE) mask (mask 127).astype(float32) # 在原图上先算边缘再做统一增强 edge cv2.Canny(img, 50, 150) if self.train: auged self.aug(imageimg, maskmask) img, mask auged[image], auged[mask] edge_auged self.aug(imageedge) edge edge_auged[image] img torch.from_numpy(img).float().unsqueeze(0) / 255.0 edge torch.from_numpy(edge).float().unsqueeze(0) / 255.0 mask torch.from_numpy(mask).float().unsqueeze(0) return img, edge, mask这套思路好在把数据预处理全部收敛在Dataset内部模型侧不用关心输入是怎么算出来的。注意Canny边缘计算完以后我单独又调用了一次增强器但只做了图像相关的变换没有把mask传进去这么做是为了让旋转翻转操作保持一致因为同一张原图翻转后其边缘图也应该跟着翻转。比较保险的做法是写一个辅助函数把img、edge、mask一起传入同一个A.Compose对象用additional_targets声明edge字段这样所有变换只执行一次保证三者绝对对齐。3.2 Y型模型的主体实现模型部分是整个项目的核心但拆开看逻辑并不复杂。我一开始看到两个编码器有点懵后来把forward画成一条数据流就清楚了左侧分支输入原始图右侧分支输入边缘图两个编码器各自的5层输出都保存下来在底部做concat加注意力融合解码器按U-Net的方式逐层上采样每一步接收上一层输出以及两个编码器对应的跳跃连接输出。这里给出一个简化但结构完整的PyTorch实现框架import torch import torch.nn as nn import torch.nn.functional as F class DoubleConv(nn.Module): def __init__(self, in_ch, out_ch): super().__init__() self.conv nn.Sequential( nn.Conv2d(in_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), nn.Conv2d(out_ch, out_ch, 3, padding1), nn.BatchNorm2d(out_ch), nn.ReLU(inplaceTrue), ) def forward(self, x): return self.conv(x) class Encoder(nn.Module): def __init__(self, in_ch, base_ch32): super().__init__() self.conv1 DoubleConv(in_ch, base_ch) self.pool1 nn.MaxPool2d(2) self.conv2 DoubleConv(base_ch, base_ch * 2) self.pool2 nn.MaxPool2d(2) self.conv3 DoubleConv(base_ch * 2, base_ch * 4) self.pool3 nn.MaxPool2d(2) self.conv4 DoubleConv(base_ch * 4, base_ch * 8) self.pool4 nn.MaxPool2d(2) def forward(self, x): f1 self.conv1(x) f2 self.conv2(self.pool1(f1)) f3 self.conv3(self.pool2(f2)) f4 self.conv4(self.pool3(f3)) f5 self.pool4(f4) return [f1, f2, f3, f4, f5] class YNet(nn.Module): def __init__(self, in_ch1, out_ch1, base_ch32): super().__init__() self.enc_a Encoder(in_ch, base_ch) self.enc_b Encoder(in_ch, base_ch) self.fusion nn.Sequential( nn.Conv2d(base_ch * 16, base_ch * 8, 1), nn.BatchNorm2d(base_ch * 8), nn.ReLU(inplaceTrue), ) self.up1 nn.ConvTranspose2d(base_ch * 8, base_ch * 8, 2, stride2) self.dec1 DoubleConv(base_ch * 8 base_ch * 4 * 2, base_ch * 4) self.up2 nn.ConvTranspose2d(base_ch * 4, base_ch * 4, 2, stride2) self.dec2 DoubleConv(base_ch * 4 base_ch * 2 * 2, base_ch * 2) self.up3 nn.ConvTranspose2d(base_ch * 2, base_ch * 2, 2, stride2) self.dec3 DoubleConv(base_ch * 2 base_ch * 2, base_ch) self.up4 nn.ConvTranspose2d(base_ch, base_ch, 2, stride2) self.dec4 DoubleConv(base_ch in_ch * 2, base_ch) self.out nn.Conv2d(base_ch, out_ch, 1) def forward(self, x_a, x_b): fa self.enc_a(x_a) # [f1, f2, f3, f4, f5] fb self.enc_b(x_b) f_fuse self.fusion(torch.cat([fa[-1], fb[-1]], dim1)) d self.dec1(torch.cat([self.up1(f_fuse), fa[3], fb[3]], dim1)) d self.dec2(torch.cat([self.up2(d), fa[2], fb[2]], dim1)) d self.dec3(torch.cat([self.up3(d), fa[1], fb[1]], dim1)) d self.dec4(torch.cat([self.up4(d), fa[0], fb[0]], dim1)) return self.out(d)注意dec2这一层输入通道计算是上采样来的base_ch * 4加两条跳跃连接各自的base_ch * 2所以实际是base_ch * 4 base_ch * 2 * 2。dec3是base_ch * 2 base_ch * 1 * 2。dec4是base_ch in_ch * 2因为两个编码器的第一个卷积输出f1是base_ch而原始输入通道被作为直连补充。这里面的通道数不是固定的你可以根据自己显卡显存调整base_ch。我建议第一次跑通时不要贪大base_ch设为32就够了一张12GB显存的显卡能轻松吃下256x256的输入。3.3 损失函数与训练策略分割任务里最常用的就是Dice Loss和Cross Entropy的加权组合。Dice Loss解决正负样本不平衡问题Cross Entropy帮助梯度更平稳地传播。我用的组合如下class DiceCEloss(nn.Module): def __init__(self, weight0.5): super().__init__() self.weight weight self.ce nn.BCEWithLogitsLoss() def forward(self, pred, mask): ce self.ce(pred, mask) p torch.sigmoid(pred) inter (p * mask).sum(dim(2, 3)) union p.sum(dim(2, 3)) mask.sum(dim(2, 3)) dice 1 - (2 * inter 1) / (union 1) return self.weight * ce (1 - self.weight) * dice.mean()训练时有两个细节值得关注。第一BatchNorm在双分支结构里要格外小心。两个编码器虽然结构相同但喂进去的数据分布差异可能很大比如原始图是0到255的灰度边缘图则是0到1的二值边缘响应。如果BatchNorm参数在初始化时没有对齐两个分支的归一化统计量会互相拉扯。稳妥的做法是对每个编码器使用独立的BatchNorm层而不是共享。第二损失函数的权重分配最好做一个简单实验。我在早期版本里把Dice和CE的权重设为0.5和0.5训练出来的mask边界偏模糊后来改成0.4和0.6边界清晰了不少但召回率下降。这个没有绝对标准建议你在验证集上多跑几次看曲线。4. 实操复现要点与踩坑记录4.1 环境配置与依赖选择我用的是PyTorch 1.13加CUDA 11.7Python版本3.9。albumentations和opencv-python是数据增强和图像读取的标配建议装最新稳定版。训练时没用复杂的分布式配置单卡完全够用。显存方面以256x256输入、base_ch32为例每个batch的显存占用大约4到6GB。如果你只有8GB显存建议把batch size设为8或者把图片缩放到224x224。值得注意的是Y型结构因为有两个编码器前向计算量几乎是同尺寸U-Net的1.8倍左右反向传播会再放大一些。如果想省显存可以考虑将其中一个编码器的特征提取精度降低或者使用混合精度训练。我的训练环境配置如下项目配置CUDA11.7PyTorch1.13.0Python3.9.16GPURTX 3080 10GB输入尺寸256 x 256batch size8优化器AdamW初始学习率1e-4学习率调整CosineAnnealing训练轮次604.2 关键参数与显存计算训练过程中我把大部分精力花在了学习率和batch size的搭配上。Y型网络因为参数量大尤其两个编码器的梯度下降幅度如果有差异很容易出现分支偏离。我的做法是前10个epoch冻结解码器只让两个编码器适应输入分布等到第10个epoch再整体训练。这个方法在医学图像上特别有效因为边界图和原始图的分布差异大早让解码器介入反而容易让模型陷入局部解。显存计算有个粗略公式可以估:模型总显存约等于输入特征图体积 中间激活值 梯度之和Y型网络激活值约是单U-Net的1.6到2倍。如果训练时报CUDA OOM先不要急着加显卡优先检查base_ch和batch size。另外把torch.cuda.amp.autocast加上往往能省下30%以上显存代价是少量精度损失在分割任务里通常可以接受。4.3 三种典型问题和排查思路第一个高频问题两个分支收敛不均衡。表现是一个分支对应的loss下降很快另一个却几乎不动。我排查后发现是边缘图分支的特征方差太小时BatchNorm把信息都压掉了。解决方法是把边缘图乘以一个可学习的缩放系数再喂进去或者调整BatchNorm的eps参数。第二个问题预测结果出现网格状伪影。这通常和转置卷积跳跃连接的通道数不匹配有关。检查一下解码器每一层的输入通道是否和我给的代码一致多数情况是某个concat的维度对不上。还有就是转置卷积的kernel和stride设置不当导致上采样时产生周期性的重叠或空洞。第三个问题验证集Dice很高但实际分割效果边缘很碎。这种一般发生在训练数据很少、边缘信息又被增强过度的情况。我给边缘分支加了高斯模糊做数据增强降低边缘响应的锐利程度模型泛化性会明显提升。注意Y-Net不是万能结构。如果你的输入本身就是单模态且特征很均匀双分支非但不能提升效果反而会把噪声引入。不要为了用这个结构而强行加一条分支先想清楚第二条输入到底提供了什么互补信息。我在实际复现中也试过把第二个输入换成原始图的高频滤波结果效果反而不如边缘图。后来想明白了边缘图是语义级的高层抽象高频滤波还停留在像素级的噪声放大两者对分割的贡献完全不可同日而语。真正有效的第二条分支必须是能提供主分支缺少的某种结构化信息而不是简单的数值变换。这一点在你自己设计分支输入时值得反复斟酌。本文还有配套的精品资源点击获取
返回列表