CoolFace
Apppublic

MLBench/Contours_Extraction

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
u2netp.py525 linesDownload Raw Back to root
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)