Team Ai
Apppublic

Arulkumar03/Fox_Sheep_Detector_Computer_Vision_model

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
batch_norm.py321 linesDownload Raw Back to layers
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