MLBench/Contours_Extraction
0
1import torch2import torch.nn as nn3import torch.nn.functional as F4 5class REBNCONV(nn.Module):6 def __init__(self,in_ch=3,out_ch=3,dirate=1):7 super(REBNCONV,self).__init__()8 9 self.conv_s1 = nn.Conv2d(in_ch,out_ch,3,padding=1*dirate,dilation=1*dirate)10 self.bn_s1 = nn.BatchNorm2d(out_ch)11 self.relu_s1 = nn.ReLU(inplace=True)12 13 def forward(self,x):14 15 hx = x16 xout = self.relu_s1(self.bn_s1(self.conv_s1(hx)))17 18 return xout19 20## upsample tensor 'src' to have the same spatial size with tensor 'tar'21def _upsample_like(src,tar):22 23 src = F.upsample(src,size=tar.shape[2:],mode='bilinear')24 25 return src26 27 28### RSU-7 ###29class RSU7(nn.Module):#UNet07DRES(nn.Module):30 31 def __init__(self, in_ch=3, mid_ch=12, out_ch=3):32 super(RSU7,self).__init__()33 34 self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)35 36 self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)37 self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)38 39 self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)40 self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)41 42 self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)43 self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)44 45 self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)46 self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)47 48 self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)49 self.pool5 = nn.MaxPool2d(2,stride=2,ceil_mode=True)50 51 self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=1)52 53 self.rebnconv7 = REBNCONV(mid_ch,mid_ch,dirate=2)54 55 self.rebnconv6d = REBNCONV(mid_ch*2,mid_ch,dirate=1)56 self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)57 self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)58 self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)59 self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)60 self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)61 62 def forward(self,x):63 64 hx = x65 hxin = self.rebnconvin(hx)66 67 hx1 = self.rebnconv1(hxin)68 hx = self.pool1(hx1)69 70 hx2 = self.rebnconv2(hx)71 hx = self.pool2(hx2)72 73 hx3 = self.rebnconv3(hx)74 hx = self.pool3(hx3)75 76 hx4 = self.rebnconv4(hx)77 hx = self.pool4(hx4)78 79 hx5 = self.rebnconv5(hx)80 hx = self.pool5(hx5)81 82 hx6 = self.rebnconv6(hx)83 84 hx7 = self.rebnconv7(hx6)85 86 hx6d = self.rebnconv6d(torch.cat((hx7,hx6),1))87 hx6dup = _upsample_like(hx6d,hx5)88 89 hx5d = self.rebnconv5d(torch.cat((hx6dup,hx5),1))90 hx5dup = _upsample_like(hx5d,hx4)91 92 hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))93 hx4dup = _upsample_like(hx4d,hx3)94 95 hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))96 hx3dup = _upsample_like(hx3d,hx2)97 98 hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))99 hx2dup = _upsample_like(hx2d,hx1)100 101 hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))102 103 return hx1d + hxin104 105### RSU-6 ###106class RSU6(nn.Module):#UNet06DRES(nn.Module):107 108 def __init__(self, in_ch=3, mid_ch=12, out_ch=3):109 super(RSU6,self).__init__()110 111 self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)112 113 self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)114 self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)115 116 self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)117 self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)118 119 self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)120 self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)121 122 self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)123 self.pool4 = nn.MaxPool2d(2,stride=2,ceil_mode=True)124 125 self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=1)126 127 self.rebnconv6 = REBNCONV(mid_ch,mid_ch,dirate=2)128 129 self.rebnconv5d = REBNCONV(mid_ch*2,mid_ch,dirate=1)130 self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)131 self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)132 self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)133 self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)134 135 def forward(self,x):136 137 hx = x138 139 hxin = self.rebnconvin(hx)140 141 hx1 = self.rebnconv1(hxin)142 hx = self.pool1(hx1)143 144 hx2 = self.rebnconv2(hx)145 hx = self.pool2(hx2)146 147 hx3 = self.rebnconv3(hx)148 hx = self.pool3(hx3)149 150 hx4 = self.rebnconv4(hx)151 hx = self.pool4(hx4)152 153 hx5 = self.rebnconv5(hx)154 155 hx6 = self.rebnconv6(hx5)156 157 158 hx5d = self.rebnconv5d(torch.cat((hx6,hx5),1))159 hx5dup = _upsample_like(hx5d,hx4)160 161 hx4d = self.rebnconv4d(torch.cat((hx5dup,hx4),1))162 hx4dup = _upsample_like(hx4d,hx3)163 164 hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))165 hx3dup = _upsample_like(hx3d,hx2)166 167 hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))168 hx2dup = _upsample_like(hx2d,hx1)169 170 hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))171 172 return hx1d + hxin173 174### RSU-5 ###175class RSU5(nn.Module):#UNet05DRES(nn.Module):176 177 def __init__(self, in_ch=3, mid_ch=12, out_ch=3):178 super(RSU5,self).__init__()179 180 self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)181 182 self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)183 self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)184 185 self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)186 self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)187 188 self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)189 self.pool3 = nn.MaxPool2d(2,stride=2,ceil_mode=True)190 191 self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=1)192 193 self.rebnconv5 = REBNCONV(mid_ch,mid_ch,dirate=2)194 195 self.rebnconv4d = REBNCONV(mid_ch*2,mid_ch,dirate=1)196 self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)197 self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)198 self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)199 200 def forward(self,x):201 202 hx = x203 204 hxin = self.rebnconvin(hx)205 206 hx1 = self.rebnconv1(hxin)207 hx = self.pool1(hx1)208 209 hx2 = self.rebnconv2(hx)210 hx = self.pool2(hx2)211 212 hx3 = self.rebnconv3(hx)213 hx = self.pool3(hx3)214 215 hx4 = self.rebnconv4(hx)216 217 hx5 = self.rebnconv5(hx4)218 219 hx4d = self.rebnconv4d(torch.cat((hx5,hx4),1))220 hx4dup = _upsample_like(hx4d,hx3)221 222 hx3d = self.rebnconv3d(torch.cat((hx4dup,hx3),1))223 hx3dup = _upsample_like(hx3d,hx2)224 225 hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))226 hx2dup = _upsample_like(hx2d,hx1)227 228 hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))229 230 return hx1d + hxin231 232### RSU-4 ###233class RSU4(nn.Module):#UNet04DRES(nn.Module):234 235 def __init__(self, in_ch=3, mid_ch=12, out_ch=3):236 super(RSU4,self).__init__()237 238 self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)239 240 self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)241 self.pool1 = nn.MaxPool2d(2,stride=2,ceil_mode=True)242 243 self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=1)244 self.pool2 = nn.MaxPool2d(2,stride=2,ceil_mode=True)245 246 self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=1)247 248 self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=2)249 250 self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=1)251 self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=1)252 self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)253 254 def forward(self,x):255 256 hx = x257 258 hxin = self.rebnconvin(hx)259 260 hx1 = self.rebnconv1(hxin)261 hx = self.pool1(hx1)262 263 hx2 = self.rebnconv2(hx)264 hx = self.pool2(hx2)265 266 hx3 = self.rebnconv3(hx)267 268 hx4 = self.rebnconv4(hx3)269 270 hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))271 hx3dup = _upsample_like(hx3d,hx2)272 273 hx2d = self.rebnconv2d(torch.cat((hx3dup,hx2),1))274 hx2dup = _upsample_like(hx2d,hx1)275 276 hx1d = self.rebnconv1d(torch.cat((hx2dup,hx1),1))277 278 return hx1d + hxin279 280### RSU-4F ###281class RSU4F(nn.Module):#UNet04FRES(nn.Module):282 283 def __init__(self, in_ch=3, mid_ch=12, out_ch=3):284 super(RSU4F,self).__init__()285 286 self.rebnconvin = REBNCONV(in_ch,out_ch,dirate=1)287 288 self.rebnconv1 = REBNCONV(out_ch,mid_ch,dirate=1)289 self.rebnconv2 = REBNCONV(mid_ch,mid_ch,dirate=2)290 self.rebnconv3 = REBNCONV(mid_ch,mid_ch,dirate=4)291 292 self.rebnconv4 = REBNCONV(mid_ch,mid_ch,dirate=8)293 294 self.rebnconv3d = REBNCONV(mid_ch*2,mid_ch,dirate=4)295 self.rebnconv2d = REBNCONV(mid_ch*2,mid_ch,dirate=2)296 self.rebnconv1d = REBNCONV(mid_ch*2,out_ch,dirate=1)297 298 def forward(self,x):299 300 hx = x301 302 hxin = self.rebnconvin(hx)303 304 hx1 = self.rebnconv1(hxin)305 hx2 = self.rebnconv2(hx1)306 hx3 = self.rebnconv3(hx2)307 308 hx4 = self.rebnconv4(hx3)309 310 hx3d = self.rebnconv3d(torch.cat((hx4,hx3),1))311 hx2d = self.rebnconv2d(torch.cat((hx3d,hx2),1))312 hx1d = self.rebnconv1d(torch.cat((hx2d,hx1),1))313 314 return hx1d + hxin315 316 317##### U^2-Net ####318class U2NET(nn.Module):319 320 def __init__(self,in_ch=3,out_ch=1):321 super(U2NET,self).__init__()322 323 self.stage1 = RSU7(in_ch,32,64)324 self.pool12 = nn.MaxPool2d(2,stride=2,ceil_mode=True)325 326 self.stage2 = RSU6(64,32,128)327 self.pool23 = nn.MaxPool2d(2,stride=2,ceil_mode=True)328 329 self.stage3 = RSU5(128,64,256)330 self.pool34 = nn.MaxPool2d(2,stride=2,ceil_mode=True)331 332 self.stage4 = RSU4(256,128,512)333 self.pool45 = nn.MaxPool2d(2,stride=2,ceil_mode=True)334 335 self.stage5 = RSU4F(512,256,512)336 self.pool56 = nn.MaxPool2d(2,stride=2,ceil_mode=True)337 338 self.stage6 = RSU4F(512,256,512)339 340 # decoder341 self.stage5d = RSU4F(1024,256,512)342 self.stage4d = RSU4(1024,128,256)343 self.stage3d = RSU5(512,64,128)344 self.stage2d = RSU6(256,32,64)345 self.stage1d = RSU7(128,16,64)346 347 self.side1 = nn.Conv2d(64,out_ch,3,padding=1)348 self.side2 = nn.Conv2d(64,out_ch,3,padding=1)349 self.side3 = nn.Conv2d(128,out_ch,3,padding=1)350 self.side4 = nn.Conv2d(256,out_ch,3,padding=1)351 self.side5 = nn.Conv2d(512,out_ch,3,padding=1)352 self.side6 = nn.Conv2d(512,out_ch,3,padding=1)353 354 self.outconv = nn.Conv2d(6*out_ch,out_ch,1)355 356 def forward(self,x):357 358 hx = x359 360 #stage 1361 hx1 = self.stage1(hx)362 hx = self.pool12(hx1)363 364 #stage 2365 hx2 = self.stage2(hx)366 hx = self.pool23(hx2)367 368 #stage 3369 hx3 = self.stage3(hx)370 hx = self.pool34(hx3)371 372 #stage 4373 hx4 = self.stage4(hx)374 hx = self.pool45(hx4)375 376 #stage 5377 hx5 = self.stage5(hx)378 hx = self.pool56(hx5)379 380 #stage 6381 hx6 = self.stage6(hx)382 hx6up = _upsample_like(hx6,hx5)383 384 #-------------------- decoder --------------------385 hx5d = self.stage5d(torch.cat((hx6up,hx5),1))386 hx5dup = _upsample_like(hx5d,hx4)387 388 hx4d = self.stage4d(torch.cat((hx5dup,hx4),1))389 hx4dup = _upsample_like(hx4d,hx3)390 391 hx3d = self.stage3d(torch.cat((hx4dup,hx3),1))392 hx3dup = _upsample_like(hx3d,hx2)393 394 hx2d = self.stage2d(torch.cat((hx3dup,hx2),1))395 hx2dup = _upsample_like(hx2d,hx1)396 397 hx1d = self.stage1d(torch.cat((hx2dup,hx1),1))398 399 400 #side output401 d1 = self.side1(hx1d)402 403 d2 = self.side2(hx2d)404 d2 = _upsample_like(d2,d1)405 406 d3 = self.side3(hx3d)407 d3 = _upsample_like(d3,d1)408 409 d4 = self.side4(hx4d)410 d4 = _upsample_like(d4,d1)411 412 d5 = self.side5(hx5d)413 d5 = _upsample_like(d5,d1)414 415 d6 = self.side6(hx6)416 d6 = _upsample_like(d6,d1)417 418 d0 = self.outconv(torch.cat((d1,d2,d3,d4,d5,d6),1))419 420 return F.sigmoid(d0), F.sigmoid(d1), F.sigmoid(d2), F.sigmoid(d3), F.sigmoid(d4), F.sigmoid(d5), F.sigmoid(d6)421 422### U^2-Net small ###423class U2NETP(nn.Module):424 425 def __init__(self,in_ch=3,out_ch=1):426 super(U2NETP,self).__init__()427 428 self.stage1 = RSU7(in_ch,16,64)429 self.pool12 = nn.MaxPool2d(2,stride=2,ceil_mode=True)430 431 self.stage2 = RSU6(64,16,64)432 self.pool23 = nn.MaxPool2d(2,stride=2,ceil_mode=True)433 434 self.stage3 = RSU5(64,16,64)435 self.pool34 = nn.MaxPool2d(2,stride=2,ceil_mode=True)436 437 self.stage4 = RSU4(64,16,64)438 self.pool45 = nn.MaxPool2d(2,stride=2,ceil_mode=True)439 440 self.stage5 = RSU4F(64,16,64)441 self.pool56 = nn.MaxPool2d(2,stride=2,ceil_mode=True)442 443 self.stage6 = RSU4F(64,16,64)444 445 # decoder446 self.stage5d = RSU4F(128,16,64)447 self.stage4d = RSU4(128,16,64)448 self.stage3d = RSU5(128,16,64)449 self.stage2d = RSU6(128,16,64)450 self.stage1d = RSU7(128,16,64)451 452 self.side1 = nn.Conv2d(64,out_ch,3,padding=1)453 self.side2 = nn.Conv2d(64,out_ch,3,padding=1)454 self.side3 = nn.Conv2d(64,out_ch,3,padding=1)455 self.side4 = nn.Conv2d(64,out_ch,3,padding=1)456 self.side5 = nn.Conv2d(64,out_ch,3,padding=1)457 self.side6 = nn.Conv2d(64,out_ch,3,padding=1)458 459 self.outconv = nn.Conv2d(6*out_ch,out_ch,1)460 461 def forward(self,x):462 463 hx = x464 465 #stage 1466 hx1 = self.stage1(hx)467 hx = self.pool12(hx1)468 469 #stage 2470 hx2 = self.stage2(hx)471 hx = self.pool23(hx2)472 473 #stage 3474 hx3 = self.stage3(hx)475 hx = self.pool34(hx3)476 477 #stage 4478 hx4 = self.stage4(hx)479 hx = self.pool45(hx4)480 481 #stage 5482 hx5 = self.stage5(hx)483 hx = self.pool56(hx5)484 485 #stage 6486 hx6 = self.stage6(hx)487 hx6up = _upsample_like(hx6,hx5)488 489 #decoder490 hx5d = self.stage5d(torch.cat((hx6up,hx5),1))491 hx5dup = _upsample_like(hx5d,hx4)492 493 hx4d = self.stage4d(torch.cat((hx5dup,hx4),1))494 hx4dup = _upsample_like(hx4d,hx3)495 496 hx3d = self.stage3d(torch.cat((hx4dup,hx3),1))497 hx3dup = _upsample_like(hx3d,hx2)498 499 hx2d = self.stage2d(torch.cat((hx3dup,hx2),1))500 hx2dup = _upsample_like(hx2d,hx1)501 502 hx1d = self.stage1d(torch.cat((hx2dup,hx1),1))503 504 505 #side output506 d1 = self.side1(hx1d)507 508 d2 = self.side2(hx2d)509 d2 = _upsample_like(d2,d1)510 511 d3 = self.side3(hx3d)512 d3 = _upsample_like(d3,d1)513 514 d4 = self.side4(hx4d)515 d4 = _upsample_like(d4,d1)516 517 d5 = self.side5(hx5d)518 d5 = _upsample_like(d5,d1)519 520 d6 = self.side6(hx6)521 d6 = _upsample_like(d6,d1)522 523 d0 = self.outconv(torch.cat((d1,d2,d3,d4,d5,d6),1))524 525 return F.sigmoid(d0), F.sigmoid(d1), F.sigmoid(d2), F.sigmoid(d3), F.sigmoid(d4), F.sigmoid(d5), F.sigmoid(d6)