CoolFace
Modelpublic

RGBD-SOD/dptdepth

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes13downloads
vit.py577 linesDownload Raw Back to root
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