RabbitRUI/ruispace
0
1import torch.nn as nn2import torch3 4 5class ResidualConv(nn.Module):6 def __init__(self, input_dim, output_dim, stride, padding):7 super(ResidualConv, self).__init__()8 9 self.conv_block = nn.Sequential(10 nn.BatchNorm2d(input_dim),11 nn.ReLU(),12 nn.Conv2d(13 input_dim, output_dim, kernel_size=3, stride=stride, padding=padding14 ),15 nn.BatchNorm2d(output_dim),16 nn.ReLU(),17 nn.Conv2d(output_dim, output_dim, kernel_size=3, padding=1),18 )19 self.conv_skip = nn.Sequential(20 nn.Conv2d(input_dim, output_dim, kernel_size=3, stride=stride, padding=1),21 nn.BatchNorm2d(output_dim),22 )23 24 def forward(self, x):25 26 return self.conv_block(x) + self.conv_skip(x)27 28 29class Upsample(nn.Module):30 def __init__(self, input_dim, output_dim, kernel, stride):31 super(Upsample, self).__init__()32 33 self.upsample = nn.ConvTranspose2d(34 input_dim, output_dim, kernel_size=kernel, stride=stride35 )36 37 def forward(self, x):38 return self.upsample(x)39 40 41class Squeeze_Excite_Block(nn.Module):42 def __init__(self, channel, reduction=16):43 super(Squeeze_Excite_Block, self).__init__()44 self.avg_pool = nn.AdaptiveAvgPool2d(1)45 self.fc = nn.Sequential(46 nn.Linear(channel, channel // reduction, bias=False),47 nn.ReLU(inplace=True),48 nn.Linear(channel // reduction, channel, bias=False),49 nn.Sigmoid(),50 )51 52 def forward(self, x):53 b, c, _, _ = x.size()54 y = self.avg_pool(x).view(b, c)55 y = self.fc(y).view(b, c, 1, 1)56 return x * y.expand_as(x)57 58 59class ASPP(nn.Module):60 def __init__(self, in_dims, out_dims, rate=[6, 12, 18]):61 super(ASPP, self).__init__()62 63 self.aspp_block1 = nn.Sequential(64 nn.Conv2d(65 in_dims, out_dims, 3, stride=1, padding=rate[0], dilation=rate[0]66 ),67 nn.ReLU(inplace=True),68 nn.BatchNorm2d(out_dims),69 )70 self.aspp_block2 = nn.Sequential(71 nn.Conv2d(72 in_dims, out_dims, 3, stride=1, padding=rate[1], dilation=rate[1]73 ),74 nn.ReLU(inplace=True),75 nn.BatchNorm2d(out_dims),76 )77 self.aspp_block3 = nn.Sequential(78 nn.Conv2d(79 in_dims, out_dims, 3, stride=1, padding=rate[2], dilation=rate[2]80 ),81 nn.ReLU(inplace=True),82 nn.BatchNorm2d(out_dims),83 )84 85 self.output = nn.Conv2d(len(rate) * out_dims, out_dims, 1)86 self._init_weights()87 88 def forward(self, x):89 x1 = self.aspp_block1(x)90 x2 = self.aspp_block2(x)91 x3 = self.aspp_block3(x)92 out = torch.cat([x1, x2, x3], dim=1)93 return self.output(out)94 95 def _init_weights(self):96 for m in self.modules():97 if isinstance(m, nn.Conv2d):98 nn.init.kaiming_normal_(m.weight)99 elif isinstance(m, nn.BatchNorm2d):100 m.weight.data.fill_(1)101 m.bias.data.zero_()102 103 104class Upsample_(nn.Module):105 def __init__(self, scale=2):106 super(Upsample_, self).__init__()107 108 self.upsample = nn.Upsample(mode="bilinear", scale_factor=scale)109 110 def forward(self, x):111 return self.upsample(x)112 113 114class AttentionBlock(nn.Module):115 def __init__(self, input_encoder, input_decoder, output_dim):116 super(AttentionBlock, self).__init__()117 118 self.conv_encoder = nn.Sequential(119 nn.BatchNorm2d(input_encoder),120 nn.ReLU(),121 nn.Conv2d(input_encoder, output_dim, 3, padding=1),122 nn.MaxPool2d(2, 2),123 )124 125 self.conv_decoder = nn.Sequential(126 nn.BatchNorm2d(input_decoder),127 nn.ReLU(),128 nn.Conv2d(input_decoder, output_dim, 3, padding=1),129 )130 131 self.conv_attn = nn.Sequential(132 nn.BatchNorm2d(output_dim),133 nn.ReLU(),134 nn.Conv2d(output_dim, 1, 1),135 )136 137 def forward(self, x1, x2):138 out = self.conv_encoder(x1) + self.conv_decoder(x2)139 out = self.conv_attn(out)140 return out * x2