Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model
0
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 