Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
vit.py525 linesDownload Raw Back to backbone
1import logging2import math3import fvcore.nn.weight_init as weight_init4import torch5import torch.nn as nn6 7from detectron2.layers import CNNBlockBase, Conv2d, get_norm8from detectron2.modeling.backbone.fpn import _assert_strides_are_log2_contiguous9 10from .backbone import Backbone11from .utils import (12    PatchEmbed,13    add_decomposed_rel_pos,14    get_abs_pos,15    window_partition,16    window_unpartition,17)18 19logger = logging.getLogger(__name__)20 21 22__all__ = ["ViT", "SimpleFeaturePyramid", "get_vit_lr_decay_rate"]23 24 25class Attention(nn.Module):26    """Multi-head Attention block with relative position embeddings."""27 28    def __init__(29        self,30        dim,31        num_heads=8,32        qkv_bias=True,33        use_rel_pos=False,34        rel_pos_zero_init=True,35        input_size=None,36    ):37        """38        Args:39            dim (int): Number of input channels.40            num_heads (int): Number of attention heads.41            qkv_bias (bool:  If True, add a learnable bias to query, key, value.42            rel_pos (bool): If True, add relative positional embeddings to the attention map.43            rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.44            input_size (int or None): Input resolution for calculating the relative positional45                parameter size.46        """47        super().__init__()48        self.num_heads = num_heads49        head_dim = dim // num_heads50        self.scale = head_dim**-0.551 52        self.qkv = nn.Linear(dim, dim * 3, bias=qkv_bias)53        self.proj = nn.Linear(dim, dim)54 55        self.use_rel_pos = use_rel_pos56        if self.use_rel_pos:57            # initialize relative positional embeddings58            self.rel_pos_h = nn.Parameter(torch.zeros(2 * input_size[0] - 1, head_dim))59            self.rel_pos_w = nn.Parameter(torch.zeros(2 * input_size[1] - 1, head_dim))60 61            if not rel_pos_zero_init:62                nn.init.trunc_normal_(self.rel_pos_h, std=0.02)63                nn.init.trunc_normal_(self.rel_pos_w, std=0.02)64 65    def forward(self, x):66        B, H, W, _ = x.shape67        # qkv with shape (3, B, nHead, H * W, C)68        qkv = self.qkv(x).reshape(B, H * W, 3, self.num_heads, -1).permute(2, 0, 3, 1, 4)69        # q, k, v with shape (B * nHead, H * W, C)70        q, k, v = qkv.reshape(3, B * self.num_heads, H * W, -1).unbind(0)71 72        attn = (q * self.scale) @ k.transpose(-2, -1)73 74        if self.use_rel_pos:75            attn = add_decomposed_rel_pos(attn, q, self.rel_pos_h, self.rel_pos_w, (H, W), (H, W))76 77        attn = attn.softmax(dim=-1)78        x = (attn @ v).view(B, self.num_heads, H, W, -1).permute(0, 2, 3, 1, 4).reshape(B, H, W, -1)79        x = self.proj(x)80 81        return x82 83 84class ResBottleneckBlock(CNNBlockBase):85    """86    The standard bottleneck residual block without the last activation layer.87    It contains 3 conv layers with kernels 1x1, 3x3, 1x1.88    """89 90    def __init__(91        self,92        in_channels,93        out_channels,94        bottleneck_channels,95        norm="LN",96        act_layer=nn.GELU,97    ):98        """99        Args:100            in_channels (int): Number of input channels.101            out_channels (int): Number of output channels.102            bottleneck_channels (int): number of output channels for the 3x3103                "bottleneck" conv layers.104            norm (str or callable): normalization for all conv layers.105                See :func:`layers.get_norm` for supported format.106            act_layer (callable): activation for all conv layers.107        """108        super().__init__(in_channels, out_channels, 1)109 110        self.conv1 = Conv2d(in_channels, bottleneck_channels, 1, bias=False)111        self.norm1 = get_norm(norm, bottleneck_channels)112        self.act1 = act_layer()113 114        self.conv2 = Conv2d(115            bottleneck_channels,116            bottleneck_channels,117            3,118            padding=1,119            bias=False,120        )121        self.norm2 = get_norm(norm, bottleneck_channels)122        self.act2 = act_layer()123 124        self.conv3 = Conv2d(bottleneck_channels, out_channels, 1, bias=False)125        self.norm3 = get_norm(norm, out_channels)126 127        for layer in [self.conv1, self.conv2, self.conv3]:128            weight_init.c2_msra_fill(layer)129        for layer in [self.norm1, self.norm2]:130            layer.weight.data.fill_(1.0)131            layer.bias.data.zero_()132        # zero init last norm layer.133        self.norm3.weight.data.zero_()134        self.norm3.bias.data.zero_()135 136    def forward(self, x):137        out = x138        for layer in self.children():139            out = layer(out)140 141        out = x + out142        return out143 144 145class Block(nn.Module):146    """Transformer blocks with support of window attention and residual propagation blocks"""147 148    def __init__(149        self,150        dim,151        num_heads,152        mlp_ratio=4.0,153        qkv_bias=True,154        drop_path=0.0,155        norm_layer=nn.LayerNorm,156        act_layer=nn.GELU,157        use_rel_pos=False,158        rel_pos_zero_init=True,159        window_size=0,160        use_residual_block=False,161        input_size=None,162    ):163        """164        Args:165            dim (int): Number of input channels.166            num_heads (int): Number of attention heads in each ViT block.167            mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.168            qkv_bias (bool): If True, add a learnable bias to query, key, value.169            drop_path (float): Stochastic depth rate.170            norm_layer (nn.Module): Normalization layer.171            act_layer (nn.Module): Activation layer.172            use_rel_pos (bool): If True, add relative positional embeddings to the attention map.173            rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.174            window_size (int): Window size for window attention blocks. If it equals 0, then not175                use window attention.176            use_residual_block (bool): If True, use a residual block after the MLP block.177            input_size (int or None): Input resolution for calculating the relative positional178                parameter size.179        """180        super().__init__()181        self.norm1 = norm_layer(dim)182        self.attn = Attention(183            dim,184            num_heads=num_heads,185            qkv_bias=qkv_bias,186            use_rel_pos=use_rel_pos,187            rel_pos_zero_init=rel_pos_zero_init,188            input_size=input_size if window_size == 0 else (window_size, window_size),189        )190 191        from timm.models.layers import DropPath, Mlp192 193        self.drop_path = DropPath(drop_path) if drop_path > 0.0 else nn.Identity()194        self.norm2 = norm_layer(dim)195        self.mlp = Mlp(in_features=dim, hidden_features=int(dim * mlp_ratio), act_layer=act_layer)196 197        self.window_size = window_size198 199        self.use_residual_block = use_residual_block200        if use_residual_block:201            # Use a residual block with bottleneck channel as dim // 2202            self.residual = ResBottleneckBlock(203                in_channels=dim,204                out_channels=dim,205                bottleneck_channels=dim // 2,206                norm="LN",207                act_layer=act_layer,208            )209 210    def forward(self, x):211        shortcut = x212        x = self.norm1(x)213        # Window partition214        if self.window_size > 0:215            H, W = x.shape[1], x.shape[2]216            x, pad_hw = window_partition(x, self.window_size)217 218        x = self.attn(x)219        # Reverse window partition220        if self.window_size > 0:221            x = window_unpartition(x, self.window_size, pad_hw, (H, W))222 223        x = shortcut + self.drop_path(x)224        x = x + self.drop_path(self.mlp(self.norm2(x)))225 226        if self.use_residual_block:227            x = self.residual(x.permute(0, 3, 1, 2)).permute(0, 2, 3, 1)228 229        return x230 231 232class ViT(Backbone):233    """234    This module implements Vision Transformer (ViT) backbone in :paper:`vitdet`.235    "Exploring Plain Vision Transformer Backbones for Object Detection",236    https://arxiv.org/abs/2203.16527237    """238 239    def __init__(240        self,241        img_size=1024,242        patch_size=16,243        in_chans=3,244        embed_dim=768,245        depth=12,246        num_heads=12,247        mlp_ratio=4.0,248        qkv_bias=True,249        drop_path_rate=0.0,250        norm_layer=nn.LayerNorm,251        act_layer=nn.GELU,252        use_abs_pos=True,253        use_rel_pos=False,254        rel_pos_zero_init=True,255        window_size=0,256        window_block_indexes=(),257        residual_block_indexes=(),258        use_act_checkpoint=False,259        pretrain_img_size=224,260        pretrain_use_cls_token=True,261        out_feature="last_feat",262    ):263        """264        Args:265            img_size (int): Input image size.266            patch_size (int): Patch size.267            in_chans (int): Number of input image channels.268            embed_dim (int): Patch embedding dimension.269            depth (int): Depth of ViT.270            num_heads (int): Number of attention heads in each ViT block.271            mlp_ratio (float): Ratio of mlp hidden dim to embedding dim.272            qkv_bias (bool): If True, add a learnable bias to query, key, value.273            drop_path_rate (float): Stochastic depth rate.274            norm_layer (nn.Module): Normalization layer.275            act_layer (nn.Module): Activation layer.276            use_abs_pos (bool): If True, use absolute positional embeddings.277            use_rel_pos (bool): If True, add relative positional embeddings to the attention map.278            rel_pos_zero_init (bool): If True, zero initialize relative positional parameters.279            window_size (int): Window size for window attention blocks.280            window_block_indexes (list): Indexes for blocks using window attention.281            residual_block_indexes (list): Indexes for blocks using conv propagation.282            use_act_checkpoint (bool): If True, use activation checkpointing.283            pretrain_img_size (int): input image size for pretraining models.284            pretrain_use_cls_token (bool): If True, pretrainig models use class token.285            out_feature (str): name of the feature from the last block.286        """287        super().__init__()288        self.pretrain_use_cls_token = pretrain_use_cls_token289 290        self.patch_embed = PatchEmbed(291            kernel_size=(patch_size, patch_size),292            stride=(patch_size, patch_size),293            in_chans=in_chans,294            embed_dim=embed_dim,295        )296 297        if use_abs_pos:298            # Initialize absolute positional embedding with pretrain image size.299            num_patches = (pretrain_img_size // patch_size) * (pretrain_img_size // patch_size)300            num_positions = (num_patches + 1) if pretrain_use_cls_token else num_patches301            self.pos_embed = nn.Parameter(torch.zeros(1, num_positions, embed_dim))302        else:303            self.pos_embed = None304 305        # stochastic depth decay rule306        dpr = [x.item() for x in torch.linspace(0, drop_path_rate, depth)]307 308        self.blocks = nn.ModuleList()309        for i in range(depth):310            block = Block(311                dim=embed_dim,312                num_heads=num_heads,313                mlp_ratio=mlp_ratio,314                qkv_bias=qkv_bias,315                drop_path=dpr[i],316                norm_layer=norm_layer,317                act_layer=act_layer,318                use_rel_pos=use_rel_pos,319                rel_pos_zero_init=rel_pos_zero_init,320                window_size=window_size if i in window_block_indexes else 0,321                use_residual_block=i in residual_block_indexes,322                input_size=(img_size // patch_size, img_size // patch_size),323            )324            if use_act_checkpoint:325                # TODO: use torch.utils.checkpoint326                from fairscale.nn.checkpoint import checkpoint_wrapper327 328                block = checkpoint_wrapper(block)329            self.blocks.append(block)330 331        self._out_feature_channels = {out_feature: embed_dim}332        self._out_feature_strides = {out_feature: patch_size}333        self._out_features = [out_feature]334 335        if self.pos_embed is not None:336            nn.init.trunc_normal_(self.pos_embed, std=0.02)337 338        self.apply(self._init_weights)339 340    def _init_weights(self, m):341        if isinstance(m, nn.Linear):342            nn.init.trunc_normal_(m.weight, std=0.02)343            if isinstance(m, nn.Linear) and m.bias is not None:344                nn.init.constant_(m.bias, 0)345        elif isinstance(m, nn.LayerNorm):346            nn.init.constant_(m.bias, 0)347            nn.init.constant_(m.weight, 1.0)348 349    def forward(self, x):350        x = self.patch_embed(x)351        if self.pos_embed is not None:352            x = x + get_abs_pos(353                self.pos_embed, self.pretrain_use_cls_token, (x.shape[1], x.shape[2])354            )355 356        for blk in self.blocks:357            x = blk(x)358 359        outputs = {self._out_features[0]: x.permute(0, 3, 1, 2)}360        return outputs361 362 363class SimpleFeaturePyramid(Backbone):364    """365    This module implements SimpleFeaturePyramid in :paper:`vitdet`.366    It creates pyramid features built on top of the input feature map.367    """368 369    def __init__(370        self,371        net,372        in_feature,373        out_channels,374        scale_factors,375        top_block=None,376        norm="LN",377        square_pad=0,378    ):379        """380        Args:381            net (Backbone): module representing the subnetwork backbone.382                Must be a subclass of :class:`Backbone`.383            in_feature (str): names of the input feature maps coming384                from the net.385            out_channels (int): number of channels in the output feature maps.386            scale_factors (list[float]): list of scaling factors to upsample or downsample387                the input features for creating pyramid features.388            top_block (nn.Module or None): if provided, an extra operation will389                be performed on the output of the last (smallest resolution)390                pyramid output, and the result will extend the result list. The top_block391                further downsamples the feature map. It must have an attribute392                "num_levels", meaning the number of extra pyramid levels added by393                this block, and "in_feature", which is a string representing394                its input feature (e.g., p5).395            norm (str): the normalization to use.396            square_pad (int): If > 0, require input images to be padded to specific square size.397        """398        super(SimpleFeaturePyramid, self).__init__()399        assert isinstance(net, Backbone)400 401        self.scale_factors = scale_factors402 403        input_shapes = net.output_shape()404        strides = [int(input_shapes[in_feature].stride / scale) for scale in scale_factors]405        _assert_strides_are_log2_contiguous(strides)406 407        dim = input_shapes[in_feature].channels408        self.stages = []409        use_bias = norm == ""410        for idx, scale in enumerate(scale_factors):411            out_dim = dim412            if scale == 4.0:413                layers = [414                    nn.ConvTranspose2d(dim, dim // 2, kernel_size=2, stride=2),415                    get_norm(norm, dim // 2),416                    nn.GELU(),417                    nn.ConvTranspose2d(dim // 2, dim // 4, kernel_size=2, stride=2),418                ]419                out_dim = dim // 4420            elif scale == 2.0:421                layers = [nn.ConvTranspose2d(dim, dim // 2, kernel_size=2, stride=2)]422                out_dim = dim // 2423            elif scale == 1.0:424                layers = []425            elif scale == 0.5:426                layers = [nn.MaxPool2d(kernel_size=2, stride=2)]427            else:428                raise NotImplementedError(f"scale_factor={scale} is not supported yet.")429 430            layers.extend(431                [432                    Conv2d(433                        out_dim,434                        out_channels,435                        kernel_size=1,436                        bias=use_bias,437                        norm=get_norm(norm, out_channels),438                    ),439                    Conv2d(440                        out_channels,441                        out_channels,442                        kernel_size=3,443                        padding=1,444                        bias=use_bias,445                        norm=get_norm(norm, out_channels),446                    ),447                ]448            )449            layers = nn.Sequential(*layers)450 451            stage = int(math.log2(strides[idx]))452            self.add_module(f"simfp_{stage}", layers)453            self.stages.append(layers)454 455        self.net = net456        self.in_feature = in_feature457        self.top_block = top_block458        # Return feature names are "p<stage>", like ["p2", "p3", ..., "p6"]459        self._out_feature_strides = {"p{}".format(int(math.log2(s))): s for s in strides}460        # top block output feature maps.461        if self.top_block is not None:462            for s in range(stage, stage + self.top_block.num_levels):463                self._out_feature_strides["p{}".format(s + 1)] = 2 ** (s + 1)464 465        self._out_features = list(self._out_feature_strides.keys())466        self._out_feature_channels = {k: out_channels for k in self._out_features}467        self._size_divisibility = strides[-1]468        self._square_pad = square_pad469 470    @property471    def padding_constraints(self):472        return {473            "size_divisiblity": self._size_divisibility,474            "square_size": self._square_pad,475        }476 477    def forward(self, x):478        """479        Args:480            x: Tensor of shape (N,C,H,W). H, W must be a multiple of ``self.size_divisibility``.481 482        Returns:483            dict[str->Tensor]:484                mapping from feature map name to pyramid feature map tensor485                in high to low resolution order. Returned feature names follow the FPN486                convention: "p<stage>", where stage has stride = 2 ** stage e.g.,487                ["p2", "p3", ..., "p6"].488        """489        bottom_up_features = self.net(x)490        features = bottom_up_features[self.in_feature]491        results = []492 493        for stage in self.stages:494            results.append(stage(features))495 496        if self.top_block is not None:497            if self.top_block.in_feature in bottom_up_features:498                top_block_in_feature = bottom_up_features[self.top_block.in_feature]499            else:500                top_block_in_feature = results[self._out_features.index(self.top_block.in_feature)]501            results.extend(self.top_block(top_block_in_feature))502        assert len(self._out_features) == len(results)503        return {f: res for f, res in zip(self._out_features, results)}504 505 506def get_vit_lr_decay_rate(name, lr_decay_rate=1.0, num_layers=12):507    """508    Calculate lr decay rate for different ViT blocks.509    Args:510        name (string): parameter name.511        lr_decay_rate (float): base lr decay rate.512        num_layers (int): number of ViT blocks.513 514    Returns:515        lr decay rate for the given parameter.516    """517    layer_id = num_layers + 1518    if name.startswith("backbone"):519        if ".pos_embed" in name or ".patch_embed" in name:520            layer_id = 0521        elif ".blocks." in name and ".residual." not in name:522            layer_id = int(name[name.find(".blocks.") :].split(".")[2]) + 1523 524    return lr_decay_rate ** (num_layers + 1 - layer_id)525