import torch
from torch import nn

# 一次卷积操作,包括 卷积 + BN(可选) + ReLU
class ConvBNReLU(nn.Module):
    """
    Conv + BN[optional] + ReLU
    """
    def __init__(self, in_ch, out_ch, isBN=True):
        super(ConvBNReLU, self).__init__()
        self.isBN = isBN
        self.conv = nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1)
        self.relu = nn.ReLU(inplace=True)
        if isBN:
            self.bn = nn.BatchNorm2d(out_ch)

    def forward(self, x):
        x = self.conv(x)
        if self.isBN:
            x = self.bn(x)
        x = self.relu(x)
        return x

# 两次卷积,可以选择是否预先进行MaxPool
class DoubleConv(nn.Module):
    """
    MaxPool[optional] + ConvBNReLU + ConvBNReLU
    """
    def __init__(self, in_ch, out_ch, isBN=True, is_pool=False):
        super(DoubleConv, self).__init__()
        self.is_pool = is_pool
        if is_pool:
            self.maxpool = nn.MaxPool2d(kernel_size=2, stride=2)
        self.conv1 = ConvBNReLU(in_ch, out_ch, isBN)
        self.conv2 = ConvBNReLU(out_ch, out_ch, isBN)

    def forward(self, x):
        if self.is_pool:
            x = self.maxpool(x)
        x = self.conv1(x)
        x = self.conv2(x)
        return x


class UNetPlusPlus(nn.Module):
    def __init__(self, img_ch, base_ch, num_class):
        super().__init__()
        c1, c2, c3 = base_ch, base_ch * 2, base_ch * 4
        c4, c5 = base_ch * 8, base_ch * 16
        # 每个level的第1个block
        # 注意:只有第一层需要下采样(is_pool=True)
        self.conv0_0 = DoubleConv(img_ch, c1, is_pool=False)
        self.conv1_0 = DoubleConv(c1, c2, is_pool=True)
        self.conv2_0 = DoubleConv(c2, c3, is_pool=True)
        self.conv3_0 = DoubleConv(c3, c4, is_pool=True)
        self.conv4_0 = DoubleConv(c4, c5, is_pool=True)
        # 每层的第2个block,level越深中间节点越少
        self.conv0_1 = DoubleConv(c1+c2, c1, is_pool=False)
        self.conv1_1 = DoubleConv(c2+c3, c2, is_pool=False)
        self.conv2_1 = DoubleConv(c3+c4, c3, is_pool=False)
        self.conv3_1 = DoubleConv(c4+c5, c4, is_pool=False)
        # 每层的第3个block
        self.conv0_2 = DoubleConv(c1*2+c2, c1, is_pool=False)
        self.conv1_2 = DoubleConv(c2*2+c3, c2, is_pool=False)
        self.conv2_2 = DoubleConv(c3*2+c4, c3, is_pool=False)
        # 每层的第3个block
        self.conv0_3 = DoubleConv(c1*3+c2, c1, is_pool=False)
        self.conv1_3 = DoubleConv(c2*3+c3, c2, is_pool=False)
        # 每层的第4个block
        self.conv0_4 = DoubleConv(c1*4+c2, c1, is_pool=False)
        # 输出 conv 层
        self.conv_out = nn.Conv2d(c1, num_class, kernel_size=1)
        self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)

    def forward(self, x):
        # 由于每个节点的计算都需要依赖于左边和左下角的节点,因此
        # 需要从每个level的第一个节点开始,逐层向右上方计算
        x0_0 = self.conv0_0(x)
        x1_0 = self.conv1_0(x0_0)
        x0_1 = self.conv0_1(torch.cat([x0_0, self.up(x1_0)], dim=1))
        # 以上第1个子网络
        x2_0 = self.conv2_0(x1_0)
        x1_1 = self.conv1_1(torch.cat([x1_0, self.up(x2_0)], dim=1))
        x0_2 = self.conv0_2(torch.cat([x0_0, x0_1, self.up(x1_1)], dim=1))
        # 以上第2个子网络
        x3_0 = self.conv3_0(x2_0)
        x2_1 = self.conv2_1(torch.cat([x2_0, self.up(x3_0)], dim=1))
        x1_2 = self.conv1_2(torch.cat([x1_0, x1_1, self.up(x2_1)], dim=1))
        x0_3 = self.conv0_3(torch.cat([x0_0, x0_1, x0_2, self.up(x1_2)], dim=1))
        # 以上第3个子网络
        x4_0 = self.conv4_0(x3_0)
        x3_1 = self.conv3_1(torch.cat([x3_0, self.up(x4_0)], dim=1))
        x2_2 = self.conv2_2(torch.cat([x2_0, x2_1, self.up(x3_1)], dim=1))
        x1_3 = self.conv1_3(torch.cat([x1_0, x1_1, x1_2, self.up(x2_2)], dim=1))
        x0_4 = self.conv0_4(torch.cat([x0_0, x0_1, x0_2, x0_3, self.up(x1_3)], dim=1))
        # 最终的UNet++网络
        output = self.conv_out(x0_4)
        return output

x_in = torch.randn((1, 3, 224, 224))
unetpp = UNetPlusPlus(img_ch=3, base_ch=32, num_class=3)
print("UNet++ \n", unetpp)
x_out = unetpp(x_in)
print("UNet++ input size : ", x_in.size())
print("UNet++ output size : ", x_out.size())

代码采用PyTorch实现了U-Net++模型,这是一种U-Net的改进变体,通过嵌套的跳跃连接(nested skip connections)和多子网络结构来提升语义分割的精度和鲁棒性。U-Net++的核心思想是构建多个嵌套的U-Net子网络,每个子网络共享编码器,但通过密集的特征融合路径来减少语义差距,提高边界分割效果。该实现使用模块化设计(如ConvBNReLU和DoubleConv),支持自定义输入通道(img_ch)、基础通道(base_ch)和输出类别(num_class)。假设输入图像为[1, 3, 224, 224](batch=1,RGB通道=3,高度/宽度=224),base_ch=32,num_class=3。

下面将按照代码的逻辑结构,从导入库开始,一步一步逐段分析。每段解释包括代码的目的、关键概念、数学公式(如果适用)、实现细节,以及为什么这样设计。代码整体分为辅助模块定义(ConvBNReLU、DoubleConv)和主模型类(UNetPlusPlus),以及测试部分。基于计算,特征图尺寸变化如下(所有尺寸均为[batch, channels, height, width],假设无BN影响尺寸):

  • x0_0: [1, 32, 224, 224]
  • x1_0: [1, 64, 112, 112]
  • x0_1: [1, 32, 224, 224]
  • x2_0: [1, 128, 56, 56]
  • x1_1: [1, 64, 112, 112]
  • x0_2: [1, 32, 224, 224]
  • x3_0: [1, 256, 28, 28]
  • x2_1: [1, 128, 56, 56]
  • x1_2: [1, 64, 112, 112]
  • x0_3: [1, 32, 224, 224]
  • x4_0: [1, 512, 14, 14]
  • x3_1: [1, 256, 28, 28]
  • x2_2: [1, 128, 56, 56]
  • x1_3: [1, 64, 112, 112]
  • x0_4: [1, 32, 224, 224]
  • output: [1, 3, 224, 224]

这些尺寸通过卷积公式计算得出:卷积输出尺寸 Hout=Hin+2×padding−kernel_sizestride+1H_{out} = \frac{H_{in} + 2 \times padding - kernel\_size}{stride} + 1Hout=strideHin+2×paddingkernel_size+1(stride默认为1);池化/上采样为2倍变化。代码运行时,print(“UNet++ \n”, unetpp)会输出模型的层结构(nn.Module的__str__表示),print尺寸验证输入输出一致。

1. 导入库
import torch
from torch import nn
  • 目的:导入PyTorch核心库,用于定义神经网络模型。
  • 详细解释
    • import torch:PyTorch主库,提供张量操作(如torch.cat、torch.randn)和自动微分。
    • from torch import nn:神经网络模块,提供Conv2d、ReLU、BatchNorm2d、MaxPool2d、Upsample等层。
  • 为什么需要这些库:U-Net++是一个深度学习模型,需要PyTorch的模块化构建。代码中未导入torch.nn.functional as F,但torch.cat是torch的内置函数,无需F。
2. 单次卷积模块:ConvBNReLU
# 一次卷积操作,包括 卷积 + BN(可选) + ReLU
class ConvBNReLU(nn.Module):
    """
    Conv + BN[optional] + ReLU
    """
    def __init__(self, in_ch, out_ch, isBN=True):
        super(ConvBNReLU, self).__init__()
        self.isBN = isBN
        self.conv = nn.Conv2d(in_ch, out_ch, kernel_size=3, padding=1)
        self.relu = nn.ReLU(inplace=True)
        if isBN:
            self.bn = nn.BatchNorm2d(out_ch)

    def forward(self, x):
        x = self.conv(x)
        if self.isBN:
            x = self.bn(x)
        x = self.relu(x)
        return x
  • 目的:定义基本卷积单元,包括卷积、可选BN和ReLU激活,用于构建更复杂的块。
  • 详细解释
    • __init__:初始化层。nn.Conv2d(in_channels, out_channels, kernel_size=3, padding=1):3x3卷积,padding=1保持尺寸。isBN控制添加nn.BatchNorm2dnn.ReLU(inplace=True):就地激活。
    • forward:Conv → [BN] → ReLU。
  • 数学概念:卷积同上文。BN标准化输入。
  • 为什么:模块化复用。isBN可选(U-Net++原版常添加BN提升稳定性)。
3. 双卷积模块:DoubleConv
# 两次卷积,可以选择是否预先进行MaxPool
class DoubleConv(nn.Module):
    """
    MaxPool[optional] + ConvBNReLU + ConvBNReLU
    """
    def __init__(self, in_ch, out_ch, isBN=True, is_pool=False):
        super(DoubleConv, self).__init__()
        self.is_pool = is_pool
        if is_pool:
            self.maxpool = nn.MaxPool2d(kernel_size=2, stride=2)
        self.conv1 = ConvBNReLU(in_ch, out_ch, isBN)
        self.conv2 = ConvBNReLU(out_ch, out_ch, isBN)

    def forward(self, x):
        if self.is_pool:
            x = self.maxpool(x)
        x = self.conv1(x)
        x = self.conv2(x)
        return x
  • 目的:定义U-Net++的基本构建块,双Conv单元,可选预池化,用于所有节点。
  • 详细解释
    • __init__is_pool添加池化(下采样)。两个ConvBNReLU:通道变化后保持。
    • forward:[MaxPool] → Conv1 → Conv2。
  • 尺寸变化:is_pool=True减半;Conv保持。
  • 为什么:U-Net++每个节点(x{i,j})都是DoubleConv。is_pool仅用于编码器第一列(下采样)。
4. 主模型:UNetPlusPlus
class UNetPlusPlus(nn.Module):
    def __init__(self, img_ch, base_ch, num_class):
        super().__init__()
        c1, c2, c3 = base_ch, base_ch * 2, base_ch * 4
        c4, c5 = base_ch * 8, base_ch * 16
        # 每个level的第1个block
        # 注意:只有第一层需要下采样(is_pool=True)
        self.conv0_0 = DoubleConv(img_ch, c1, is_pool=False)
        self.conv1_0 = DoubleConv(c1, c2, is_pool=True)
        self.conv2_0 = DoubleConv(c2, c3, is_pool=True)
        self.conv3_0 = DoubleConv(c3, c4, is_pool=True)
        self.conv4_0 = DoubleConv(c4, c5, is_pool=True)
        # 每层的第2个block,level越深中间节点越少
        self.conv0_1 = DoubleConv(c1+c2, c1, is_pool=False)
        self.conv1_1 = DoubleConv(c2+c3, c2, is_pool=False)
        self.conv2_1 = DoubleConv(c3+c4, c3, is_pool=False)
        self.conv3_1 = DoubleConv(c4+c5, c4, is_pool=False)
        # 每层的第3个block
        self.conv0_2 = DoubleConv(c1*2+c2, c1, is_pool=False)
        self.conv1_2 = DoubleConv(c2*2+c3, c2, is_pool=False)
        self.conv2_2 = DoubleConv(c3*2+c4, c3, is_pool=False)
        # 每层的第3个block
        self.conv0_3 = DoubleConv(c1*3+c2, c1, is_pool=False)
        self.conv1_3 = DoubleConv(c2*3+c3, c2, is_pool=False)
        # 每层的第4个block
        self.conv0_4 = DoubleConv(c1*4+c2, c1, is_pool=False)
        # 输出 conv 层
        self.conv_out = nn.Conv2d(c1, num_class, kernel_size=1)
        self.up = nn.Upsample(scale_factor=2, mode='bilinear', align_corners=True)
  • 目的:定义U-Net++模型,构建嵌套结构。
  • 详细解释
    • __init__:定义通道c1-c5(倍增)。节点命名x{i,j},i=深度(0-4),j=路径(0为编码,>0为解码)。第一列(j=0):下采样,除了conv0_0无池化。第二列(j=1):融合1个上采样+跳连。依此类推,输入通道为累加(e.g., conv0_2: c1*2 + c2 = x0_0 + x0_1 + up(x1_1))。nn.Upsample:双线性上采样2倍。conv_out:1x1卷积输出num_class通道。
  • 为什么:嵌套设计形成4个子U-Net(j=1到4),共享编码器。通道倍增捕捉多尺度特征。
    def forward(self, x):
        # 由于每个节点的计算都需要依赖于左边和左下角的节点,因此
        # 需要从每个level的第一个节点开始,逐层向右上方计算
        x0_0 = self.conv0_0(x)
        x1_0 = self.conv1_0(x0_0)
        x0_1 = self.conv0_1(torch.cat([x0_0, self.up(x1_0)], dim=1))
        # 以上第1个子网络
        x2_0 = self.conv2_0(x1_0)
        x1_1 = self.conv1_1(torch.cat([x1_0, self.up(x2_0)], dim=1))
        x0_2 = self.conv0_2(torch.cat([x0_0, x0_1, self.up(x1_1)], dim=1))
        # 以上第2个子网络
        x3_0 = self.conv3_0(x2_0)
        x2_1 = self.conv2_1(torch.cat([x2_0, self.up(x3_0)], dim=1))
        x1_2 = self.conv1_2(torch.cat([x1_0, x1_1, self.up(x2_1)], dim=1))
        x0_3 = self.conv0_3(torch.cat([x0_0, x0_1, x0_2, self.up(x1_2)], dim=1))
        # 以上第3个子网络
        x4_0 = self.conv4_0(x3_0)
        x3_1 = self.conv3_1(torch.cat([x3_0, self.up(x4_0)], dim=1))
        x2_2 = self.conv2_2(torch.cat([x2_0, x2_1, self.up(x3_1)], dim=1))
        x1_3 = self.conv1_3(torch.cat([x1_0, x1_1, x1_2, self.up(x2_2)], dim=1))
        x0_4 = self.conv0_4(torch.cat([x0_0, x0_1, x0_2, x0_3, self.up(x1_3)], dim=1))
        # 最终的UNet++网络
        output = self.conv_out(x0_4)
        return output
  • 目的:定义前向传播,逐节点计算嵌套特征。
  • 详细解释:从底向上、左到右计算(依赖关系)。每个节点:cat左边所有+up(下角) → DoubleConv。最终x0_4融合所有浅层路径,conv_out输出logits。
  • 数学概念:cat沿dim=1(通道)融合,up 2倍放大。
  • 为什么:顺序计算确保依赖满足。x0_4作为最终输出,聚合多子网络特征。
5. 测试部分
x_in = torch.randn((1, 3, 224, 224))
unetpp = UNetPlusPlus(img_ch=3, base_ch=32, num_class=3)
print("UNet++ \n", unetpp)
x_out = unetpp(x_in)
print("UNet++ input size : ", x_in.size())
print("UNet++ output size : ", x_out.size())
  • 目的:测试模型,打印结构和尺寸。
  • 详细解释:随机输入。实例化(base_ch=32)。print(unetpp)显示层结构(如UNetPlusPlus( (conv0_0): DoubleConv(…) … ))。x_out验证尺寸。
  • 为什么:调试,确保端到端工作。输出通道=num_class。
总体总结
  • 代码流程:辅助模块 → UNetPlusPlus组装 → 测试。
  • 适用场景:语义分割,提升U-Net精度(如医疗图像)。
  • 潜在改进:添加BN(当前isBN=True但代码中未传);深监督(当前无);剪枝优化推理。
  • 模型结构示例(基于print(unetpp)的典型输出):
    UNetPlusPlus(
    (conv0_0): DoubleConv(
    (conv1): ConvBNReLU(
      (conv): Conv2d(3, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(32, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv1_0): DoubleConv(
    (maxpool): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False)
    (conv1): ConvBNReLU(
      (conv): Conv2d(32, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv2_0): DoubleConv(
    (maxpool): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False)
    (conv1): ConvBNReLU(
      (conv): Conv2d(64, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv3_0): DoubleConv(
    (maxpool): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False)
    (conv1): ConvBNReLU(
      (conv): Conv2d(128, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv4_0): DoubleConv(
    (maxpool): MaxPool2d(kernel_size=2, stride=2, padding=0, dilation=1, ceil_mode=False)
    (conv1): ConvBNReLU(
      (conv): Conv2d(256, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(512, 512, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(512, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv0_1): DoubleConv(
    (conv1): ConvBNReLU(
      (conv): Conv2d(96, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(32, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv1_1): DoubleConv(
    (conv1): ConvBNReLU(
      (conv): Conv2d(192, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv2_1): DoubleConv(
    (conv1): ConvBNReLU(
      (conv): Conv2d(384, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv3_1): DoubleConv(
    (conv1): ConvBNReLU(
      (conv): Conv2d(768, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(256, 256, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(256, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv0_2): DoubleConv(
    (conv1): ConvBNReLU(
      (conv): Conv2d(128, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(32, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv1_2): DoubleConv(
    (conv1): ConvBNReLU(
      (conv): Conv2d(256, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv2_2): DoubleConv(
    (conv1): ConvBNReLU(
      (conv): Conv2d(512, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(128, 128, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(128, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv0_3): DoubleConv(
    (conv1): ConvBNReLU(
      (conv): Conv2d(160, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(32, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv1_3): DoubleConv(
    (conv1): ConvBNReLU(
      (conv): Conv2d(320, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(64, 64, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(64, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv0_4): DoubleConv(
    (conv1): ConvBNReLU(
      (conv): Conv2d(192, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    (conv2): ConvBNReLU(
      (conv): Conv2d(32, 32, kernel_size=(3, 3), stride=(1, 1), padding=(1, 1))
      (relu): ReLU(inplace=True)
      (bn): BatchNorm2d(32, eps=1e-05, momentum=0.1, affine=True, track_running_stats=True)
    )
    )
    (conv_out): Conv2d(32, 3, kernel_size=(1, 1), stride=(1, 1))
    (up): Upsample(scale_factor=2.0, mode='bilinear')
    )
    
Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐