Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model
0
1# Copyright (c) Facebook, Inc. and its affiliates.2import torch3import torch.distributed as dist4from fvcore.nn.distributed import differentiable_all_reduce5from torch import nn6from torch.nn import functional as F7 8from detectron2.utils import comm, env9 10from .wrappers import BatchNorm2d11 12 13class FrozenBatchNorm2d(nn.Module):14 """15 BatchNorm2d where the batch statistics and the affine parameters are fixed.16 17 It contains non-trainable buffers called18 "weight" and "bias", "running_mean", "running_var",19 initialized to perform identity transformation.20 21 The pre-trained backbone models from Caffe2 only contain "weight" and "bias",22 which are computed from the original four parameters of BN.23 The affine transform `x * weight + bias` will perform the equivalent24 computation of `(x - running_mean) / sqrt(running_var) * weight + bias`.25 When loading a backbone model from Caffe2, "running_mean" and "running_var"26 will be left unchanged as identity transformation.27 28 Other pre-trained backbone models may contain all 4 parameters.29 30 The forward is implemented by `F.batch_norm(..., training=False)`.31 """32 33 _version = 334 35 def __init__(self, num_features, eps=1e-5):36 super().__init__()37 self.num_features = num_features38 self.eps = eps39 self.register_buffer("weight", torch.ones(num_features))40 self.register_buffer("bias", torch.zeros(num_features))41 self.register_buffer("running_mean", torch.zeros(num_features))42 self.register_buffer("running_var", torch.ones(num_features) - eps)43 self.register_buffer("num_batches_tracked", None)44 45 def forward(self, x):46 if x.requires_grad:47 # When gradients are needed, F.batch_norm will use extra memory48 # because its backward op computes gradients for weight/bias as well.49 scale = self.weight * (self.running_var + self.eps).rsqrt()50 bias = self.bias - self.running_mean * scale51 scale = scale.reshape(1, -1, 1, 1)52 bias = bias.reshape(1, -1, 1, 1)53 out_dtype = x.dtype # may be half54 return x * scale.to(out_dtype) + bias.to(out_dtype)55 else:56 # When gradients are not needed, F.batch_norm is a single fused op57 # and provide more optimization opportunities.58 return F.batch_norm(59 x,60 self.running_mean,61 self.running_var,62 self.weight,63 self.bias,64 training=False,65 eps=self.eps,66 )67 68 def _load_from_state_dict(69 self,70 state_dict,71 prefix,72 local_metadata,73 strict,74 missing_keys,75 unexpected_keys,76 error_msgs,77 ):78 version = local_metadata.get("version", None)79 80 if version is None or version < 2:81 # No running_mean/var in early versions82 # This will silent the warnings83 if prefix + "running_mean" not in state_dict:84 state_dict[prefix + "running_mean"] = torch.zeros_like(self.running_mean)85 if prefix + "running_var" not in state_dict:86 state_dict[prefix + "running_var"] = torch.ones_like(self.running_var)87 88 super()._load_from_state_dict(89 state_dict,90 prefix,91 local_metadata,92 strict,93 missing_keys,94 unexpected_keys,95 error_msgs,96 )97 98 def __repr__(self):99 return "FrozenBatchNorm2d(num_features={}, eps={})".format(self.num_features, self.eps)100 101 @classmethod102 def convert_frozen_batchnorm(cls, module):103 """104 Convert all BatchNorm/SyncBatchNorm in module into FrozenBatchNorm.105 106 Args:107 module (torch.nn.Module):108 109 Returns:110 If module is BatchNorm/SyncBatchNorm, returns a new module.111 Otherwise, in-place convert module and return it.112 113 Similar to convert_sync_batchnorm in114 https://github.com/pytorch/pytorch/blob/master/torch/nn/modules/batchnorm.py115 """116 bn_module = nn.modules.batchnorm117 bn_module = (bn_module.BatchNorm2d, bn_module.SyncBatchNorm)118 res = module119 if isinstance(module, bn_module):120 res = cls(module.num_features)121 if module.affine:122 res.weight.data = module.weight.data.clone().detach()123 res.bias.data = module.bias.data.clone().detach()124 res.running_mean.data = module.running_mean.data125 res.running_var.data = module.running_var.data126 res.eps = module.eps127 res.num_batches_tracked = module.num_batches_tracked128 else:129 for name, child in module.named_children():130 new_child = cls.convert_frozen_batchnorm(child)131 if new_child is not child:132 res.add_module(name, new_child)133 return res134 135 136def get_norm(norm, out_channels):137 """138 Args:139 norm (str or callable): either one of BN, SyncBN, FrozenBN, GN;140 or a callable that takes a channel number and returns141 the normalization layer as a nn.Module.142 143 Returns:144 nn.Module or None: the normalization layer145 """146 if norm is None:147 return None148 if isinstance(norm, str):149 if len(norm) == 0:150 return None151 norm = {152 "BN": BatchNorm2d,153 # Fixed in https://github.com/pytorch/pytorch/pull/36382154 "SyncBN": NaiveSyncBatchNorm if env.TORCH_VERSION <= (1, 5) else nn.SyncBatchNorm,155 "FrozenBN": FrozenBatchNorm2d,156 "GN": lambda channels: nn.GroupNorm(32, channels),157 # for debugging:158 "nnSyncBN": nn.SyncBatchNorm,159 "naiveSyncBN": NaiveSyncBatchNorm,160 # expose stats_mode N as an option to caller, required for zero-len inputs161 "naiveSyncBN_N": lambda channels: NaiveSyncBatchNorm(channels, stats_mode="N"),162 "LN": lambda channels: LayerNorm(channels),163 }[norm]164 return norm(out_channels)165 166 167class NaiveSyncBatchNorm(BatchNorm2d):168 """169 In PyTorch<=1.5, ``nn.SyncBatchNorm`` has incorrect gradient170 when the batch size on each worker is different.171 (e.g., when scale augmentation is used, or when it is applied to mask head).172 173 This is a slower but correct alternative to `nn.SyncBatchNorm`.174 175 Note:176 There isn't a single definition of Sync BatchNorm.177 178 When ``stats_mode==""``, this module computes overall statistics by using179 statistics of each worker with equal weight. The result is true statistics180 of all samples (as if they are all on one worker) only when all workers181 have the same (N, H, W). This mode does not support inputs with zero batch size.182 183 When ``stats_mode=="N"``, this module computes overall statistics by weighting184 the statistics of each worker by their ``N``. The result is true statistics185 of all samples (as if they are all on one worker) only when all workers186 have the same (H, W). It is slower than ``stats_mode==""``.187 188 Even though the result of this module may not be the true statistics of all samples,189 it may still be reasonable because it might be preferrable to assign equal weights190 to all workers, regardless of their (H, W) dimension, instead of putting larger weight191 on larger images. From preliminary experiments, little difference is found between such192 a simplified implementation and an accurate computation of overall mean & variance.193 """194 195 def __init__(self, *args, stats_mode="", **kwargs):196 super().__init__(*args, **kwargs)197 assert stats_mode in ["", "N"]198 self._stats_mode = stats_mode199 200 def forward(self, input):201 if comm.get_world_size() == 1 or not self.training:202 return super().forward(input)203 204 B, C = input.shape[0], input.shape[1]205 206 half_input = input.dtype == torch.float16207 if half_input:208 # fp16 does not have good enough numerics for the reduction here209 input = input.float()210 mean = torch.mean(input, dim=[0, 2, 3])211 meansqr = torch.mean(input * input, dim=[0, 2, 3])212 213 if self._stats_mode == "":214 assert B > 0, 'SyncBatchNorm(stats_mode="") does not support zero batch size.'215 vec = torch.cat([mean, meansqr], dim=0)216 vec = differentiable_all_reduce(vec) * (1.0 / dist.get_world_size())217 mean, meansqr = torch.split(vec, C)218 momentum = self.momentum219 else:220 if B == 0:221 vec = torch.zeros([2 * C + 1], device=mean.device, dtype=mean.dtype)222 vec = vec + input.sum() # make sure there is gradient w.r.t input223 else:224 vec = torch.cat(225 [226 mean,227 meansqr,228 torch.ones([1], device=mean.device, dtype=mean.dtype),229 ],230 dim=0,231 )232 vec = differentiable_all_reduce(vec * B)233 234 total_batch = vec[-1].detach()235 momentum = total_batch.clamp(max=1) * self.momentum # no update if total_batch is 0236 mean, meansqr, _ = torch.split(vec / total_batch.clamp(min=1), C) # avoid div-by-zero237 238 var = meansqr - mean * mean239 invstd = torch.rsqrt(var + self.eps)240 scale = self.weight * invstd241 bias = self.bias - mean * scale242 scale = scale.reshape(1, -1, 1, 1)243 bias = bias.reshape(1, -1, 1, 1)244 245 self.running_mean += momentum * (mean.detach() - self.running_mean)246 self.running_var += momentum * (var.detach() - self.running_var)247 ret = input * scale + bias248 if half_input:249 ret = ret.half()250 return ret251 252 253class CycleBatchNormList(nn.ModuleList):254 """255 Implement domain-specific BatchNorm by cycling.256 257 When a BatchNorm layer is used for multiple input domains or input258 features, it might need to maintain a separate test-time statistics259 for each domain. See Sec 5.2 in :paper:`rethinking-batchnorm`.260 261 This module implements it by using N separate BN layers262 and it cycles through them every time a forward() is called.263 264 NOTE: The caller of this module MUST guarantee to always call265 this module by multiple of N times. Otherwise its test-time statistics266 will be incorrect.267 """268 269 def __init__(self, length: int, bn_class=nn.BatchNorm2d, **kwargs):270 """271 Args:272 length: number of BatchNorm layers to cycle.273 bn_class: the BatchNorm class to use274 kwargs: arguments of the BatchNorm class, such as num_features.275 """276 self._affine = kwargs.pop("affine", True)277 super().__init__([bn_class(**kwargs, affine=False) for k in range(length)])278 if self._affine:279 # shared affine, domain-specific BN280 channels = self[0].num_features281 self.weight = nn.Parameter(torch.ones(channels))282 self.bias = nn.Parameter(torch.zeros(channels))283 self._pos = 0284 285 def forward(self, x):286 ret = self[self._pos](x)287 self._pos = (self._pos + 1) % len(self)288 289 if self._affine:290 w = self.weight.reshape(1, -1, 1, 1)291 b = self.bias.reshape(1, -1, 1, 1)292 return ret * w + b293 else:294 return ret295 296 def extra_repr(self):297 return f"affine={self._affine}"298 299 300class LayerNorm(nn.Module):301 """302 A LayerNorm variant, popularized by Transformers, that performs point-wise mean and303 variance normalization over the channel dimension for inputs that have shape304 (batch_size, channels, height, width).305 https://github.com/facebookresearch/ConvNeXt/blob/d1fa8f6fef0a165b27399986cc2bdacc92777e40/models/convnext.py#L119 # noqa B950306 """307 308 def __init__(self, normalized_shape, eps=1e-6):309 super().__init__()310 self.weight = nn.Parameter(torch.ones(normalized_shape))311 self.bias = nn.Parameter(torch.zeros(normalized_shape))312 self.eps = eps313 self.normalized_shape = (normalized_shape,)314 315 def forward(self, x):316 u = x.mean(1, keepdim=True)317 s = (x - u).pow(2).mean(1, keepdim=True)318 x = (x - u) / torch.sqrt(s + self.eps)319 x = self.weight[:, None, None] * x + self.bias[:, None, None]320 return x321 