RGBD-SOD/dptdepth
113
1import torch2import torch.nn as nn3import timm4import types5import math6import torch.nn.functional as F7 8 9activations = {}10 11 12def get_activation(name):13 def hook(model, input, output):14 activations[name] = output15 16 return hook17 18 19attention = {}20 21 22def get_attention(name):23 def hook(module, input, output):24 x = input[0]25 B, N, C = x.shape26 qkv = (27 module.qkv(x)28 .reshape(B, N, 3, module.num_heads, C // module.num_heads)29 .permute(2, 0, 3, 1, 4)30 )31 q, k, v = (32 qkv[0],33 qkv[1],34 qkv[2],35 ) # make torchscript happy (cannot use tensor as tuple)36 37 attn = (q @ k.transpose(-2, -1)) * module.scale38 39 attn = attn.softmax(dim=-1) # [:,:,1,1:]40 attention[name] = attn41 42 return hook43 44 45def get_mean_attention_map(attn, token, shape):46 attn = attn[:, :, token, 1:]47 attn = attn.unflatten(2, torch.Size([shape[2] // 16, shape[3] // 16])).float()48 attn = torch.nn.functional.interpolate(49 attn, size=shape[2:], mode="bicubic", align_corners=False50 ).squeeze(0)51 52 all_attn = torch.mean(attn, 0)53 54 return all_attn55 56 57class Slice(nn.Module):58 def __init__(self, start_index=1):59 super(Slice, self).__init__()60 self.start_index = start_index61 62 def forward(self, x):63 return x[:, self.start_index :]64 65 66class AddReadout(nn.Module):67 def __init__(self, start_index=1):68 super(AddReadout, self).__init__()69 self.start_index = start_index70 71 def forward(self, x):72 if self.start_index == 2:73 readout = (x[:, 0] + x[:, 1]) / 274 else:75 readout = x[:, 0]76 return x[:, self.start_index :] + readout.unsqueeze(1)77 78 79class ProjectReadout(nn.Module):80 def __init__(self, in_features, start_index=1):81 super(ProjectReadout, self).__init__()82 self.start_index = start_index83 84 self.project = nn.Sequential(nn.Linear(2 * in_features, in_features), nn.GELU())85 86 def forward(self, x):87 readout = x[:, 0].unsqueeze(1).expand_as(x[:, self.start_index :])88 features = torch.cat((x[:, self.start_index :], readout), -1)89 90 return self.project(features)91 92 93class Transpose(nn.Module):94 def __init__(self, dim0, dim1):95 super(Transpose, self).__init__()96 self.dim0 = dim097 self.dim1 = dim198 99 def forward(self, x):100 x = x.transpose(self.dim0, self.dim1)101 return x102 103 104def forward_vit(pretrained, x):105 b, c, h, w = x.shape106 107 glob = pretrained.model.forward_flex(x)108 109 layer_1 = pretrained.activations["1"]110 layer_2 = pretrained.activations["2"]111 layer_3 = pretrained.activations["3"]112 layer_4 = pretrained.activations["4"]113 114 layer_1 = pretrained.act_postprocess1[0:2](layer_1)115 layer_2 = pretrained.act_postprocess2[0:2](layer_2)116 layer_3 = pretrained.act_postprocess3[0:2](layer_3)117 layer_4 = pretrained.act_postprocess4[0:2](layer_4)118 119 unflatten = nn.Sequential(120 nn.Unflatten(121 2,122 torch.Size(123 [124 h // pretrained.model.patch_size[1],125 w // pretrained.model.patch_size[0],126 ]127 ),128 )129 )130 131 if layer_1.ndim == 3:132 layer_1 = unflatten(layer_1)133 if layer_2.ndim == 3:134 layer_2 = unflatten(layer_2)135 if layer_3.ndim == 3:136 layer_3 = unflatten(layer_3)137 if layer_4.ndim == 3:138 layer_4 = unflatten(layer_4)139 140 layer_1 = pretrained.act_postprocess1[3 : len(pretrained.act_postprocess1)](layer_1)141 layer_2 = pretrained.act_postprocess2[3 : len(pretrained.act_postprocess2)](layer_2)142 layer_3 = pretrained.act_postprocess3[3 : len(pretrained.act_postprocess3)](layer_3)143 layer_4 = pretrained.act_postprocess4[3 : len(pretrained.act_postprocess4)](layer_4)144 145 return layer_1, layer_2, layer_3, layer_4146 147 148def _resize_pos_embed(self, posemb, gs_h, gs_w):149 posemb_tok, posemb_grid = (150 posemb[:, : self.start_index],151 posemb[0, self.start_index :],152 )153 154 gs_old = int(math.sqrt(len(posemb_grid)))155 156 posemb_grid = posemb_grid.reshape(1, gs_old, gs_old, -1).permute(0, 3, 1, 2)157 posemb_grid = F.interpolate(posemb_grid, size=(gs_h, gs_w), mode="bilinear")158 posemb_grid = posemb_grid.permute(0, 2, 3, 1).reshape(1, gs_h * gs_w, -1)159 160 posemb = torch.cat([posemb_tok, posemb_grid], dim=1)161 162 return posemb163 164 165def forward_flex(self, x):166 b, c, h, w = x.shape167 168 pos_embed = self._resize_pos_embed(169 self.pos_embed, h // self.patch_size[1], w // self.patch_size[0]170 )171 172 B = x.shape[0]173 174 if hasattr(self.patch_embed, "backbone"):175 x = self.patch_embed.backbone(x)176 if isinstance(x, (list, tuple)):177 x = x[-1] # last feature if backbone outputs list/tuple of features178 179 x = self.patch_embed.proj(x).flatten(2).transpose(1, 2)180 181 if getattr(self, "dist_token", None) is not None:182 cls_tokens = self.cls_token.expand(183 B, -1, -1184 ) # stole cls_tokens impl from Phil Wang, thanks185 dist_token = self.dist_token.expand(B, -1, -1)186 x = torch.cat((cls_tokens, dist_token, x), dim=1)187 else:188 cls_tokens = self.cls_token.expand(189 B, -1, -1190 ) # stole cls_tokens impl from Phil Wang, thanks191 x = torch.cat((cls_tokens, x), dim=1)192 193 x = x + pos_embed194 x = self.pos_drop(x)195 196 for blk in self.blocks:197 x = blk(x)198 199 x = self.norm(x)200 201 return x202 203 204def get_readout_oper(vit_features, features, use_readout, start_index=1):205 if use_readout == "ignore":206 readout_oper = [Slice(start_index)] * len(features)207 elif use_readout == "add":208 readout_oper = [AddReadout(start_index)] * len(features)209 elif use_readout == "project":210 readout_oper = [211 ProjectReadout(vit_features, start_index) for out_feat in features212 ]213 else:214 assert (215 False216 ), "wrong operation for readout token, use_readout can be 'ignore', 'add', or 'project'"217 218 return readout_oper219 220 221def _make_vit_b16_backbone(222 model,223 features=[96, 192, 384, 768],224 size=[384, 384],225 hooks=[2, 5, 8, 11],226 vit_features=768,227 use_readout="ignore",228 start_index=1,229 enable_attention_hooks=False,230):231 pretrained = nn.Module()232 233 pretrained.model = model234 pretrained.model.blocks[hooks[0]].register_forward_hook(get_activation("1"))235 pretrained.model.blocks[hooks[1]].register_forward_hook(get_activation("2"))236 pretrained.model.blocks[hooks[2]].register_forward_hook(get_activation("3"))237 pretrained.model.blocks[hooks[3]].register_forward_hook(get_activation("4"))238 239 pretrained.activations = activations240 241 if enable_attention_hooks:242 pretrained.model.blocks[hooks[0]].attn.register_forward_hook(243 get_attention("attn_1")244 )245 pretrained.model.blocks[hooks[1]].attn.register_forward_hook(246 get_attention("attn_2")247 )248 pretrained.model.blocks[hooks[2]].attn.register_forward_hook(249 get_attention("attn_3")250 )251 pretrained.model.blocks[hooks[3]].attn.register_forward_hook(252 get_attention("attn_4")253 )254 pretrained.attention = attention255 256 readout_oper = get_readout_oper(vit_features, features, use_readout, start_index)257 258 # 32, 48, 136, 384259 pretrained.act_postprocess1 = nn.Sequential(260 readout_oper[0],261 Transpose(1, 2),262 nn.Unflatten(2, torch.Size([size[0] // 16, size[1] // 16])),263 nn.Conv2d(264 in_channels=vit_features,265 out_channels=features[0],266 kernel_size=1,267 stride=1,268 padding=0,269 ),270 nn.ConvTranspose2d(271 in_channels=features[0],272 out_channels=features[0],273 kernel_size=4,274 stride=4,275 padding=0,276 bias=True,277 dilation=1,278 groups=1,279 ),280 )281 282 pretrained.act_postprocess2 = nn.Sequential(283 readout_oper[1],284 Transpose(1, 2),285 nn.Unflatten(2, torch.Size([size[0] // 16, size[1] // 16])),286 nn.Conv2d(287 in_channels=vit_features,288 out_channels=features[1],289 kernel_size=1,290 stride=1,291 padding=0,292 ),293 nn.ConvTranspose2d(294 in_channels=features[1],295 out_channels=features[1],296 kernel_size=2,297 stride=2,298 padding=0,299 bias=True,300 dilation=1,301 groups=1,302 ),303 )304 305 pretrained.act_postprocess3 = nn.Sequential(306 readout_oper[2],307 Transpose(1, 2),308 nn.Unflatten(2, torch.Size([size[0] // 16, size[1] // 16])),309 nn.Conv2d(310 in_channels=vit_features,311 out_channels=features[2],312 kernel_size=1,313 stride=1,314 padding=0,315 ),316 )317 318 pretrained.act_postprocess4 = nn.Sequential(319 readout_oper[3],320 Transpose(1, 2),321 nn.Unflatten(2, torch.Size([size[0] // 16, size[1] // 16])),322 nn.Conv2d(323 in_channels=vit_features,324 out_channels=features[3],325 kernel_size=1,326 stride=1,327 padding=0,328 ),329 nn.Conv2d(330 in_channels=features[3],331 out_channels=features[3],332 kernel_size=3,333 stride=2,334 padding=1,335 ),336 )337 338 pretrained.model.start_index = start_index339 pretrained.model.patch_size = [16, 16]340 341 # We inject this function into the VisionTransformer instances so that342 # we can use it with interpolated position embeddings without modifying the library source.343 pretrained.model.forward_flex = types.MethodType(forward_flex, pretrained.model)344 pretrained.model._resize_pos_embed = types.MethodType(345 _resize_pos_embed, pretrained.model346 )347 348 return pretrained349 350 351def _make_vit_b_rn50_backbone(352 model,353 features=[256, 512, 768, 768],354 size=[384, 384],355 hooks=[0, 1, 8, 11],356 vit_features=768,357 use_vit_only=False,358 use_readout="ignore",359 start_index=1,360 enable_attention_hooks=False,361):362 pretrained = nn.Module()363 364 pretrained.model = model365 366 if use_vit_only == True:367 pretrained.model.blocks[hooks[0]].register_forward_hook(get_activation("1"))368 pretrained.model.blocks[hooks[1]].register_forward_hook(get_activation("2"))369 else:370 pretrained.model.patch_embed.backbone.stages[0].register_forward_hook(371 get_activation("1")372 )373 pretrained.model.patch_embed.backbone.stages[1].register_forward_hook(374 get_activation("2")375 )376 377 pretrained.model.blocks[hooks[2]].register_forward_hook(get_activation("3"))378 pretrained.model.blocks[hooks[3]].register_forward_hook(get_activation("4"))379 380 if enable_attention_hooks:381 pretrained.model.blocks[2].attn.register_forward_hook(get_attention("attn_1"))382 pretrained.model.blocks[5].attn.register_forward_hook(get_attention("attn_2"))383 pretrained.model.blocks[8].attn.register_forward_hook(get_attention("attn_3"))384 pretrained.model.blocks[11].attn.register_forward_hook(get_attention("attn_4"))385 pretrained.attention = attention386 387 pretrained.activations = activations388 389 readout_oper = get_readout_oper(vit_features, features, use_readout, start_index)390 391 if use_vit_only == True:392 pretrained.act_postprocess1 = nn.Sequential(393 readout_oper[0],394 Transpose(1, 2),395 nn.Unflatten(2, torch.Size([size[0] // 16, size[1] // 16])),396 nn.Conv2d(397 in_channels=vit_features,398 out_channels=features[0],399 kernel_size=1,400 stride=1,401 padding=0,402 ),403 nn.ConvTranspose2d(404 in_channels=features[0],405 out_channels=features[0],406 kernel_size=4,407 stride=4,408 padding=0,409 bias=True,410 dilation=1,411 groups=1,412 ),413 )414 415 pretrained.act_postprocess2 = nn.Sequential(416 readout_oper[1],417 Transpose(1, 2),418 nn.Unflatten(2, torch.Size([size[0] // 16, size[1] // 16])),419 nn.Conv2d(420 in_channels=vit_features,421 out_channels=features[1],422 kernel_size=1,423 stride=1,424 padding=0,425 ),426 nn.ConvTranspose2d(427 in_channels=features[1],428 out_channels=features[1],429 kernel_size=2,430 stride=2,431 padding=0,432 bias=True,433 dilation=1,434 groups=1,435 ),436 )437 else:438 pretrained.act_postprocess1 = nn.Sequential(439 nn.Identity(), nn.Identity(), nn.Identity()440 )441 pretrained.act_postprocess2 = nn.Sequential(442 nn.Identity(), nn.Identity(), nn.Identity()443 )444 445 pretrained.act_postprocess3 = nn.Sequential(446 readout_oper[2],447 Transpose(1, 2),448 nn.Unflatten(2, torch.Size([size[0] // 16, size[1] // 16])),449 nn.Conv2d(450 in_channels=vit_features,451 out_channels=features[2],452 kernel_size=1,453 stride=1,454 padding=0,455 ),456 )457 458 pretrained.act_postprocess4 = nn.Sequential(459 readout_oper[3],460 Transpose(1, 2),461 nn.Unflatten(2, torch.Size([size[0] // 16, size[1] // 16])),462 nn.Conv2d(463 in_channels=vit_features,464 out_channels=features[3],465 kernel_size=1,466 stride=1,467 padding=0,468 ),469 nn.Conv2d(470 in_channels=features[3],471 out_channels=features[3],472 kernel_size=3,473 stride=2,474 padding=1,475 ),476 )477 478 pretrained.model.start_index = start_index479 pretrained.model.patch_size = [16, 16]480 481 # We inject this function into the VisionTransformer instances so that482 # we can use it with interpolated position embeddings without modifying the library source.483 pretrained.model.forward_flex = types.MethodType(forward_flex, pretrained.model)484 485 # We inject this function into the VisionTransformer instances so that486 # we can use it with interpolated position embeddings without modifying the library source.487 pretrained.model._resize_pos_embed = types.MethodType(488 _resize_pos_embed, pretrained.model489 )490 491 return pretrained492 493 494def _make_pretrained_vitb_rn50_384(495 pretrained,496 use_readout="ignore",497 hooks=None,498 use_vit_only=False,499 enable_attention_hooks=False,500):501 model = timm.create_model("vit_base_resnet50_384", pretrained=pretrained)502 503 hooks = [0, 1, 8, 11] if hooks == None else hooks504 return _make_vit_b_rn50_backbone(505 model,506 features=[256, 512, 768, 768],507 size=[384, 384],508 hooks=hooks,509 use_vit_only=use_vit_only,510 use_readout=use_readout,511 enable_attention_hooks=enable_attention_hooks,512 )513 514 515def _make_pretrained_vitl16_384(516 pretrained, use_readout="ignore", hooks=None, enable_attention_hooks=False517):518 model = timm.create_model("vit_large_patch16_384", pretrained=pretrained)519 520 hooks = [5, 11, 17, 23] if hooks == None else hooks521 return _make_vit_b16_backbone(522 model,523 features=[256, 512, 1024, 1024],524 hooks=hooks,525 vit_features=1024,526 use_readout=use_readout,527 enable_attention_hooks=enable_attention_hooks,528 )529 530 531def _make_pretrained_vitb16_384(532 pretrained, use_readout="ignore", hooks=None, enable_attention_hooks=False533):534 model = timm.create_model("vit_base_patch16_384", pretrained=pretrained)535 536 hooks = [2, 5, 8, 11] if hooks == None else hooks537 return _make_vit_b16_backbone(538 model,539 features=[96, 192, 384, 768],540 hooks=hooks,541 use_readout=use_readout,542 enable_attention_hooks=enable_attention_hooks,543 )544 545 546def _make_pretrained_deitb16_384(547 pretrained, use_readout="ignore", hooks=None, enable_attention_hooks=False548):549 model = timm.create_model("vit_deit_base_patch16_384", pretrained=pretrained)550 551 hooks = [2, 5, 8, 11] if hooks == None else hooks552 return _make_vit_b16_backbone(553 model,554 features=[96, 192, 384, 768],555 hooks=hooks,556 use_readout=use_readout,557 enable_attention_hooks=enable_attention_hooks,558 )559 560 561def _make_pretrained_deitb16_distil_384(562 pretrained, use_readout="ignore", hooks=None, enable_attention_hooks=False563):564 model = timm.create_model(565 "vit_deit_base_distilled_patch16_384", pretrained=pretrained566 )567 568 hooks = [2, 5, 8, 11] if hooks == None else hooks569 return _make_vit_b16_backbone(570 model,571 features=[96, 192, 384, 768],572 hooks=hooks,573 use_readout=use_readout,574 start_index=2,575 enable_attention_hooks=enable_attention_hooks,576 )577 