Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
mvit.py449 linesDownload Raw Back to backbone
1import logging2import numpy as np3import torch4import torch.nn as nn5 6from .backbone import Backbone7from .utils import (8    PatchEmbed,9    add_decomposed_rel_pos,10    get_abs_pos,11    window_partition,12    window_unpartition,13)14 15logger = logging.getLogger(__name__)16 17 18__all__ = ["MViT"]19 20 21def attention_pool(x, pool, norm=None):22    # (B, H, W, C) -> (B, C, H, W)23    x = x.permute(0, 3, 1, 2)24    x = pool(x)25    # (B, C, H1, W1) -> (B, H1, W1, C)26    x = x.permute(0, 2, 3, 1)27    if norm:28        x = norm(x)29 30    return x31 32 33class MultiScaleAttention(nn.Module):34    """Multiscale Multi-head Attention block."""35 36    def __init__(37        self,38        dim,39        dim_out,40        num_heads,41        qkv_bias=True,42        norm_layer=nn.LayerNorm,43        pool_kernel=(3, 3),44        stride_q=1,45        stride_kv=1,46        residual_pooling=True,47        window_size=0,48        use_rel_pos=False,49        rel_pos_zero_init=True,50        input_size=None,51    ):52        """53        Args:54            dim (int): Number of input channels.55            dim_out (int): Number of output channels.56            num_heads (int): Number of attention heads.57            qkv_bias (bool:  If True, add a learnable bias to query, key, value.58            norm_layer (nn.Module): Normalization layer.59            pool_kernel (tuple): kernel size for qkv pooling layers.60            stride_q (int): stride size for q pooling layer.61            stride_kv (int): stride size for kv pooling layer.62            residual_pooling (bool): If true, enable residual pooling.63            use_rel_pos (bool): If True, add relative postional embeddings to the attention map.64            rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.65            input_size (int or None): Input resolution.66        """67        super().__init__()68        self.num_heads = num_heads69        head_dim = dim_out // num_heads70        self.scale = head_dim**-0.571 72        self.qkv = nn.Linear(dim, dim_out * 3, bias=qkv_bias)73        self.proj = nn.Linear(dim_out, dim_out)74 75        # qkv pooling76        pool_padding = [k // 2 for k in pool_kernel]77        dim_conv = dim_out // num_heads78        self.pool_q = nn.Conv2d(79            dim_conv,80            dim_conv,81            pool_kernel,82            stride=stride_q,83            padding=pool_padding,84            groups=dim_conv,85            bias=False,86        )87        self.norm_q = norm_layer(dim_conv)88        self.pool_k = nn.Conv2d(89            dim_conv,90            dim_conv,91            pool_kernel,92            stride=stride_kv,93            padding=pool_padding,94            groups=dim_conv,95            bias=False,96        )97        self.norm_k = norm_layer(dim_conv)98        self.pool_v = nn.Conv2d(99            dim_conv,100            dim_conv,101            pool_kernel,102            stride=stride_kv,103            padding=pool_padding,104            groups=dim_conv,105            bias=False,106        )107        self.norm_v = norm_layer(dim_conv)108 109        self.window_size = window_size110        if window_size:111            self.q_win_size = window_size // stride_q112            self.kv_win_size = window_size // stride_kv113        self.residual_pooling = residual_pooling114 115        self.use_rel_pos = use_rel_pos116        if self.use_rel_pos:117            # initialize relative positional embeddings118            assert input_size[0] == input_size[1]119            size = input_size[0]120            rel_dim = 2 * max(size // stride_q, size // stride_kv) - 1121            self.rel_pos_h = nn.Parameter(torch.zeros(rel_dim, head_dim))122            self.rel_pos_w = nn.Parameter(torch.zeros(rel_dim, head_dim))123 124            if not rel_pos_zero_init:125                nn.init.trunc_normal_(self.rel_pos_h, std=0.02)126                nn.init.trunc_normal_(self.rel_pos_w, std=0.02)127 128    def forward(self, x):129        B, H, W, _ = x.shape130        # qkv with shape (3, B, nHead, H, W, C)131        qkv = self.qkv(x).reshape(B, H, W, 3, self.num_heads, -1).permute(3, 0, 4, 1, 2, 5)132        # q, k, v with shape (B * nHead, H, W, C)133        q, k, v = qkv.reshape(3, B * self.num_heads, H, W, -1).unbind(0)134 135        q = attention_pool(q, self.pool_q, self.norm_q)136        k = attention_pool(k, self.pool_k, self.norm_k)137        v = attention_pool(v, self.pool_v, self.norm_v)138 139        ori_q = q140        if self.window_size:141            q, q_hw_pad = window_partition(q, self.q_win_size)142            k, kv_hw_pad = window_partition(k, self.kv_win_size)143            v, _ = window_partition(v, self.kv_win_size)144            q_hw = (self.q_win_size, self.q_win_size)145            kv_hw = (self.kv_win_size, self.kv_win_size)146        else:147            q_hw = q.shape[1:3]148            kv_hw = k.shape[1:3]149 150        q = q.view(q.shape[0], np.prod(q_hw), -1)151        k = k.view(k.shape[0], np.prod(kv_hw), -1)152        v = v.view(v.shape[0], np.prod(kv_hw), -1)153 154        attn = (q * self.scale) @ k.transpose(-2, -1)155 156        if self.use_rel_pos:157            attn = add_decomposed_rel_pos(attn, q, self.rel_pos_h, self.rel_pos_w, q_hw, kv_hw)158 159        attn = attn.softmax(dim=-1)160        x = attn @ v161 162        x = x.view(x.shape[0], q_hw[0], q_hw[1], -1)163 164        if self.window_size:165            x = window_unpartition(x, self.q_win_size, q_hw_pad, ori_q.shape[1:3])166 167        if self.residual_pooling:168            x += ori_q169 170        H, W = x.shape[1], x.shape[2]171        x = x.view(B, self.num_heads, H, W, -1).permute(0, 2, 3, 1, 4).reshape(B, H, W, -1)172        x = self.proj(x)173 174        return x175 176 177class MultiScaleBlock(nn.Module):178    """Multiscale Transformer blocks"""179 180    def __init__(181        self,182        dim,183        dim_out,184        num_heads,185        mlp_ratio=4.0,186        qkv_bias=True,187        drop_path=0.0,188        norm_layer=nn.LayerNorm,189        act_layer=nn.GELU,190        qkv_pool_kernel=(3, 3),191        stride_q=1,192        stride_kv=1,193        residual_pooling=True,194        window_size=0,195        use_rel_pos=False,196        rel_pos_zero_init=True,197        input_size=None,198    ):199        """200        Args:201            dim (int): Number of input channels.202            dim_out (int): Number of output channels.203            num_heads (int): Number of attention heads in the MViT block.204            mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.205            qkv_bias (bool): If True, add a learnable bias to query, key, value.206            drop_path (float): Stochastic depth rate.207            norm_layer (nn.Module): Normalization layer.208            act_layer (nn.Module): Activation layer.209            qkv_pool_kernel (tuple): kernel size for qkv pooling layers.210            stride_q (int): stride size for q pooling layer.211            stride_kv (int): stride size for kv pooling layer.212            residual_pooling (bool): If true, enable residual pooling.213            window_size (int): Window size for window attention blocks. If it equals 0, then not214                use window attention.215            use_rel_pos (bool): If True, add relative postional embeddings to the attention map.216            rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.217            input_size (int or None): Input resolution.218        """219        super().__init__()220        self.norm1 = norm_layer(dim)221        self.attn = MultiScaleAttention(222            dim,223            dim_out,224            num_heads=num_heads,225            qkv_bias=qkv_bias,226            norm_layer=norm_layer,227            pool_kernel=qkv_pool_kernel,228            stride_q=stride_q,229            stride_kv=stride_kv,230            residual_pooling=residual_pooling,231            window_size=window_size,232            use_rel_pos=use_rel_pos,233            rel_pos_zero_init=rel_pos_zero_init,234            input_size=input_size,235        )236 237        from timm.models.layers import DropPath, Mlp238 239        self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()240        self.norm2 = norm_layer(dim_out)241        self.mlp = Mlp(242            in_features=dim_out,243            hidden_features=int(dim_out * mlp_ratio),244            out_features=dim_out,245            act_layer=act_layer,246        )247 248        if dim != dim_out:249            self.proj = nn.Linear(dim, dim_out)250 251        if stride_q > 1:252            kernel_skip = stride_q + 1253            padding_skip = int(kernel_skip // 2)254            self.pool_skip = nn.MaxPool2d(kernel_skip, stride_q, padding_skip, ceil_mode=False)255 256    def forward(self, x):257        x_norm = self.norm1(x)258        x_block = self.attn(x_norm)259 260        if hasattr(self, "proj"):261            x = self.proj(x_norm)262        if hasattr(self, "pool_skip"):263            x = attention_pool(x, self.pool_skip)264 265        x = x + self.drop_path(x_block)266        x = x + self.drop_path(self.mlp(self.norm2(x)))267 268        return x269 270 271class MViT(Backbone):272    """273    This module implements Multiscale Vision Transformer (MViT) backbone in :paper:'mvitv2'.274    """275 276    def __init__(277        self,278        img_size=224,279        patch_kernel=(7, 7),280        patch_stride=(4, 4),281        patch_padding=(3, 3),282        in_chans=3,283        embed_dim=96,284        depth=16,285        num_heads=1,286        last_block_indexes=(0, 2, 11, 15),287        qkv_pool_kernel=(3, 3),288        adaptive_kv_stride=4,289        adaptive_window_size=56,290        residual_pooling=True,291        mlp_ratio=4.0,292        qkv_bias=True,293        drop_path_rate=0.0,294        norm_layer=nn.LayerNorm,295        act_layer=nn.GELU,296        use_abs_pos=False,297        use_rel_pos=True,298        rel_pos_zero_init=True,299        use_act_checkpoint=False,300        pretrain_img_size=224,301        pretrain_use_cls_token=True,302        out_features=("scale2", "scale3", "scale4", "scale5"),303    ):304        """305        Args:306            img_size (int): Input image size.307            patch_kernel (tuple): kernel size for patch embedding.308            patch_stride (tuple): stride size for patch embedding.309            patch_padding (tuple): padding size for patch embedding.310            in_chans (int): Number of input image channels.311            embed_dim (int): Patch embedding dimension.312            depth (int): Depth of MViT.313            num_heads (int): Number of base attention heads in each MViT block.314            last_block_indexes (tuple): Block indexes for last blocks in each stage.315            qkv_pool_kernel (tuple): kernel size for qkv pooling layers.316            adaptive_kv_stride (int): adaptive stride size for kv pooling.317            adaptive_window_size (int): adaptive window size for window attention blocks.318            residual_pooling (bool): If true, enable residual pooling.319            mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.320            qkv_bias (bool): If True, add a learnable bias to query, key, value.321            drop_path_rate (float): Stochastic depth rate.322            norm_layer (nn.Module): Normalization layer.323            act_layer (nn.Module): Activation layer.324            use_abs_pos (bool): If True, use absolute positional embeddings.325            use_rel_pos (bool): If True, add relative postional embeddings to the attention map.326            rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.327            window_size (int): Window size for window attention blocks.328            use_act_checkpoint (bool): If True, use activation checkpointing.329            pretrain_img_size (int): input image size for pretraining models.330            pretrain_use_cls_token (bool): If True, pretrainig models use class token.331            out_features (tuple): name of the feature maps from each stage.332        """333        super().__init__()334        self.pretrain_use_cls_token = pretrain_use_cls_token335 336        self.patch_embed = PatchEmbed(337            kernel_size=patch_kernel,338            stride=patch_stride,339            padding=patch_padding,340            in_chans=in_chans,341            embed_dim=embed_dim,342        )343 344        if use_abs_pos:345            # Initialize absoluate positional embedding with pretrain image size.346            num_patches = (pretrain_img_size // patch_stride[0]) * (347                pretrain_img_size // patch_stride[1]348            )349            num_positions = (num_patches + 1) if pretrain_use_cls_token else num_patches350            self.pos_embed = nn.Parameter(torch.zeros(1, num_positions, embed_dim))351        else:352            self.pos_embed = None353 354        # stochastic depth decay rule355        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)]356        dim_out = embed_dim357        stride_kv = adaptive_kv_stride358        window_size = adaptive_window_size359        input_size = (img_size // patch_stride[0], img_size // patch_stride[1])360        stage = 2361        stride = patch_stride[0]362        self._out_feature_strides = {}363        self._out_feature_channels = {}364        self.blocks = nn.ModuleList()365        for i in range(depth):366            # Multiply stride_kv by 2 if it's the last block of stage2 and stage3.367            if i == last_block_indexes[1] or i == last_block_indexes[2]:368                stride_kv_ = stride_kv * 2369            else:370                stride_kv_ = stride_kv371            # hybrid window attention: global attention in last three stages.372            window_size_ = 0 if i in last_block_indexes[1:] else window_size373            block = MultiScaleBlock(374                dim=embed_dim,375                dim_out=dim_out,376                num_heads=num_heads,377                mlp_ratio=mlp_ratio,378                qkv_bias=qkv_bias,379                drop_path=dpr[i],380                norm_layer=norm_layer,381                qkv_pool_kernel=qkv_pool_kernel,382                stride_q=2 if i - 1 in last_block_indexes else 1,383                stride_kv=stride_kv_,384                residual_pooling=residual_pooling,385                window_size=window_size_,386                use_rel_pos=use_rel_pos,387                rel_pos_zero_init=rel_pos_zero_init,388                input_size=input_size,389            )390            if use_act_checkpoint:391                # TODO: use torch.utils.checkpoint392                from fairscale.nn.checkpoint import checkpoint_wrapper393 394                block = checkpoint_wrapper(block)395            self.blocks.append(block)396 397            embed_dim = dim_out398            if i in last_block_indexes:399                name = f"scale{stage}"400                if name in out_features:401                    self._out_feature_channels[name] = dim_out402                    self._out_feature_strides[name] = stride403                    self.add_module(f"{name}_norm", norm_layer(dim_out))404 405                dim_out *= 2406                num_heads *= 2407                stride_kv = max(stride_kv // 2, 1)408                stride *= 2409                stage += 1410            if i - 1 in last_block_indexes:411                window_size = window_size // 2412                input_size = [s // 2 for s in input_size]413 414        self._out_features = out_features415        self._last_block_indexes = last_block_indexes416 417        if self.pos_embed is not None:418            nn.init.trunc_normal_(self.pos_embed, std=0.02)419 420        self.apply(self._init_weights)421 422    def _init_weights(self, m):423        if isinstance(m, nn.Linear):424            nn.init.trunc_normal_(m.weight, std=0.02)425            if isinstance(m, nn.Linear) and m.bias is not None:426                nn.init.constant_(m.bias, 0)427        elif isinstance(m, nn.LayerNorm):428            nn.init.constant_(m.bias, 0)429            nn.init.constant_(m.weight, 1.0)430 431    def forward(self, x):432        x = self.patch_embed(x)433 434        if self.pos_embed is not None:435            x = x + get_abs_pos(self.pos_embed, self.pretrain_use_cls_token, x.shape[1:3])436 437        outputs = {}438        stage = 2439        for i, blk in enumerate(self.blocks):440            x = blk(x)441            if i in self._last_block_indexes:442                name = f"scale{stage}"443                if name in self._out_features:444                    x_out = getattr(self, f"{name}_norm")(x)445                    outputs[name] = x_out.permute(0, 3, 1, 2)446                stage += 1447 448        return outputs449