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