Team Ai
Apppublic

gulabpatel/First-Order-Motion

sourceHugging Faceupdated 5y agoView on Hugging Face
0likes
batchnorm.py316 linesDownload Raw Back to sync_batchnorm
1# -*- coding: utf-8 -*-2# File   : batchnorm.py3# Author : Jiayuan Mao4# Email  : maojiayuan@gmail.com5# Date   : 27/01/20186# 7# This file is part of Synchronized-BatchNorm-PyTorch.8# https://github.com/vacancy/Synchronized-BatchNorm-PyTorch9# Distributed under MIT License.10 11import collections12 13import torch14import torch.nn.functional as F15 16from torch.nn.modules.batchnorm import _BatchNorm17from torch.nn.parallel._functions import ReduceAddCoalesced, Broadcast18 19from .comm import SyncMaster20 21__all__ = ['SynchronizedBatchNorm1d', 'SynchronizedBatchNorm2d', 'SynchronizedBatchNorm3d']22 23 24def _sum_ft(tensor):25    """sum over the first and last dimention"""26    return tensor.sum(dim=0).sum(dim=-1)27 28 29def _unsqueeze_ft(tensor):30    """add new dementions at the front and the tail"""31    return tensor.unsqueeze(0).unsqueeze(-1)32 33 34_ChildMessage = collections.namedtuple('_ChildMessage', ['sum', 'ssum', 'sum_size'])35_MasterMessage = collections.namedtuple('_MasterMessage', ['sum', 'inv_std'])36 37 38class _SynchronizedBatchNorm(_BatchNorm):39    def __init__(self, num_features, eps=1e-5, momentum=0.1, affine=True):40        super(_SynchronizedBatchNorm, self).__init__(num_features, eps=eps, momentum=momentum, affine=affine)41 42        self._sync_master = SyncMaster(self._data_parallel_master)43 44        self._is_parallel = False45        self._parallel_id = None46        self._slave_pipe = None47 48    def forward(self, input):49        # If it is not parallel computation or is in evaluation mode, use PyTorch's implementation.50        if not (self._is_parallel and self.training):51            return F.batch_norm(52                input, self.running_mean, self.running_var, self.weight, self.bias,53                self.training, self.momentum, self.eps)54 55        # Resize the input to (B, C, -1).56        input_shape = input.size()57        input = input.view(input.size(0), self.num_features, -1)58 59        # Compute the sum and square-sum.60        sum_size = input.size(0) * input.size(2)61        input_sum = _sum_ft(input)62        input_ssum = _sum_ft(input ** 2)63 64        # Reduce-and-broadcast the statistics.65        if self._parallel_id == 0:66            mean, inv_std = self._sync_master.run_master(_ChildMessage(input_sum, input_ssum, sum_size))67        else:68            mean, inv_std = self._slave_pipe.run_slave(_ChildMessage(input_sum, input_ssum, sum_size))69 70        # Compute the output.71        if self.affine:72            # MJY:: Fuse the multiplication for speed.73            output = (input - _unsqueeze_ft(mean)) * _unsqueeze_ft(inv_std * self.weight) + _unsqueeze_ft(self.bias)74        else:75            output = (input - _unsqueeze_ft(mean)) * _unsqueeze_ft(inv_std)76 77        # Reshape it.78        return output.view(input_shape)79 80    def __data_parallel_replicate__(self, ctx, copy_id):81        self._is_parallel = True82        self._parallel_id = copy_id83 84        # parallel_id == 0 means master device.85        if self._parallel_id == 0:86            ctx.sync_master = self._sync_master87        else:88            self._slave_pipe = ctx.sync_master.register_slave(copy_id)89 90    def _data_parallel_master(self, intermediates):91        """Reduce the sum and square-sum, compute the statistics, and broadcast it."""92 93        # Always using same "device order" makes the ReduceAdd operation faster.94        # Thanks to:: Tete Xiao (http://tetexiao.com/)95        intermediates = sorted(intermediates, key=lambda i: i[1].sum.get_device())96 97        to_reduce = [i[1][:2] for i in intermediates]98        to_reduce = [j for i in to_reduce for j in i]  # flatten99        target_gpus = [i[1].sum.get_device() for i in intermediates]100 101        sum_size = sum([i[1].sum_size for i in intermediates])102        sum_, ssum = ReduceAddCoalesced.apply(target_gpus[0], 2, *to_reduce)103        mean, inv_std = self._compute_mean_std(sum_, ssum, sum_size)104 105        broadcasted = Broadcast.apply(target_gpus, mean, inv_std)106 107        outputs = []108        for i, rec in enumerate(intermediates):109            outputs.append((rec[0], _MasterMessage(*broadcasted[i*2:i*2+2])))110 111        return outputs112 113    def _compute_mean_std(self, sum_, ssum, size):114        """Compute the mean and standard-deviation with sum and square-sum. This method115        also maintains the moving average on the master device."""116        assert size > 1, 'BatchNorm computes unbiased standard-deviation, which requires size > 1.'117        mean = sum_ / size118        sumvar = ssum - sum_ * mean119        unbias_var = sumvar / (size - 1)120        bias_var = sumvar / size121 122        self.running_mean = (1 - self.momentum) * self.running_mean + self.momentum * mean.data123        self.running_var = (1 - self.momentum) * self.running_var + self.momentum * unbias_var.data124 125        return mean, bias_var.clamp(self.eps) ** -0.5126 127 128class SynchronizedBatchNorm1d(_SynchronizedBatchNorm):129    r"""Applies Synchronized Batch Normalization over a 2d or 3d input that is seen as a130    mini-batch.131 132    .. math::133 134        y = \frac{x - mean[x]}{ \sqrt{Var[x] + \epsilon}} * gamma + beta135 136    This module differs from the built-in PyTorch BatchNorm1d as the mean and137    standard-deviation are reduced across all devices during training.138 139    For example, when one uses `nn.DataParallel` to wrap the network during140    training, PyTorch's implementation normalize the tensor on each device using141    the statistics only on that device, which accelerated the computation and142    is also easy to implement, but the statistics might be inaccurate.143    Instead, in this synchronized version, the statistics will be computed144    over all training samples distributed on multiple devices.145    146    Note that, for one-GPU or CPU-only case, this module behaves exactly same147    as the built-in PyTorch implementation.148 149    The mean and standard-deviation are calculated per-dimension over150    the mini-batches and gamma and beta are learnable parameter vectors151    of size C (where C is the input size).152 153    During training, this layer keeps a running estimate of its computed mean154    and variance. The running sum is kept with a default momentum of 0.1.155 156    During evaluation, this running mean/variance is used for normalization.157 158    Because the BatchNorm is done over the `C` dimension, computing statistics159    on `(N, L)` slices, it's common terminology to call this Temporal BatchNorm160 161    Args:162        num_features: num_features from an expected input of size163            `batch_size x num_features [x width]`164        eps: a value added to the denominator for numerical stability.165            Default: 1e-5166        momentum: the value used for the running_mean and running_var167            computation. Default: 0.1168        affine: a boolean value that when set to ``True``, gives the layer learnable169            affine parameters. Default: ``True``170 171    Shape:172        - Input: :math:`(N, C)` or :math:`(N, C, L)`173        - Output: :math:`(N, C)` or :math:`(N, C, L)` (same shape as input)174 175    Examples:176        >>> # With Learnable Parameters177        >>> m = SynchronizedBatchNorm1d(100)178        >>> # Without Learnable Parameters179        >>> m = SynchronizedBatchNorm1d(100, affine=False)180        >>> input = torch.autograd.Variable(torch.randn(20, 100))181        >>> output = m(input)182    """183 184    def _check_input_dim(self, input):185        if input.dim() != 2 and input.dim() != 3:186            raise ValueError('expected 2D or 3D input (got {}D input)'187                             .format(input.dim()))188        super(SynchronizedBatchNorm1d, self)._check_input_dim(input)189 190 191class SynchronizedBatchNorm2d(_SynchronizedBatchNorm):192    r"""Applies Batch Normalization over a 4d input that is seen as a mini-batch193    of 3d inputs194 195    .. math::196 197        y = \frac{x - mean[x]}{ \sqrt{Var[x] + \epsilon}} * gamma + beta198 199    This module differs from the built-in PyTorch BatchNorm2d as the mean and200    standard-deviation are reduced across all devices during training.201 202    For example, when one uses `nn.DataParallel` to wrap the network during203    training, PyTorch's implementation normalize the tensor on each device using204    the statistics only on that device, which accelerated the computation and205    is also easy to implement, but the statistics might be inaccurate.206    Instead, in this synchronized version, the statistics will be computed207    over all training samples distributed on multiple devices.208    209    Note that, for one-GPU or CPU-only case, this module behaves exactly same210    as the built-in PyTorch implementation.211 212    The mean and standard-deviation are calculated per-dimension over213    the mini-batches and gamma and beta are learnable parameter vectors214    of size C (where C is the input size).215 216    During training, this layer keeps a running estimate of its computed mean217    and variance. The running sum is kept with a default momentum of 0.1.218 219    During evaluation, this running mean/variance is used for normalization.220 221    Because the BatchNorm is done over the `C` dimension, computing statistics222    on `(N, H, W)` slices, it's common terminology to call this Spatial BatchNorm223 224    Args:225        num_features: num_features from an expected input of226            size batch_size x num_features x height x width227        eps: a value added to the denominator for numerical stability.228            Default: 1e-5229        momentum: the value used for the running_mean and running_var230            computation. Default: 0.1231        affine: a boolean value that when set to ``True``, gives the layer learnable232            affine parameters. Default: ``True``233 234    Shape:235        - Input: :math:`(N, C, H, W)`236        - Output: :math:`(N, C, H, W)` (same shape as input)237 238    Examples:239        >>> # With Learnable Parameters240        >>> m = SynchronizedBatchNorm2d(100)241        >>> # Without Learnable Parameters242        >>> m = SynchronizedBatchNorm2d(100, affine=False)243        >>> input = torch.autograd.Variable(torch.randn(20, 100, 35, 45))244        >>> output = m(input)245    """246 247    def _check_input_dim(self, input):248        if input.dim() != 4:249            raise ValueError('expected 4D input (got {}D input)'250                             .format(input.dim()))251        super(SynchronizedBatchNorm2d, self)._check_input_dim(input)252 253 254class SynchronizedBatchNorm3d(_SynchronizedBatchNorm):255    r"""Applies Batch Normalization over a 5d input that is seen as a mini-batch256    of 4d inputs257 258    .. math::259 260        y = \frac{x - mean[x]}{ \sqrt{Var[x] + \epsilon}} * gamma + beta261 262    This module differs from the built-in PyTorch BatchNorm3d as the mean and263    standard-deviation are reduced across all devices during training.264 265    For example, when one uses `nn.DataParallel` to wrap the network during266    training, PyTorch's implementation normalize the tensor on each device using267    the statistics only on that device, which accelerated the computation and268    is also easy to implement, but the statistics might be inaccurate.269    Instead, in this synchronized version, the statistics will be computed270    over all training samples distributed on multiple devices.271    272    Note that, for one-GPU or CPU-only case, this module behaves exactly same273    as the built-in PyTorch implementation.274 275    The mean and standard-deviation are calculated per-dimension over276    the mini-batches and gamma and beta are learnable parameter vectors277    of size C (where C is the input size).278 279    During training, this layer keeps a running estimate of its computed mean280    and variance. The running sum is kept with a default momentum of 0.1.281 282    During evaluation, this running mean/variance is used for normalization.283 284    Because the BatchNorm is done over the `C` dimension, computing statistics285    on `(N, D, H, W)` slices, it's common terminology to call this Volumetric BatchNorm286    or Spatio-temporal BatchNorm287 288    Args:289        num_features: num_features from an expected input of290            size batch_size x num_features x depth x height x width291        eps: a value added to the denominator for numerical stability.292            Default: 1e-5293        momentum: the value used for the running_mean and running_var294            computation. Default: 0.1295        affine: a boolean value that when set to ``True``, gives the layer learnable296            affine parameters. Default: ``True``297 298    Shape:299        - Input: :math:`(N, C, D, H, W)`300        - Output: :math:`(N, C, D, H, W)` (same shape as input)301 302    Examples:303        >>> # With Learnable Parameters304        >>> m = SynchronizedBatchNorm3d(100)305        >>> # Without Learnable Parameters306        >>> m = SynchronizedBatchNorm3d(100, affine=False)307        >>> input = torch.autograd.Variable(torch.randn(20, 100, 35, 45, 10))308        >>> output = m(input)309    """310 311    def _check_input_dim(self, input):312        if input.dim() != 5:313            raise ValueError('expected 5D input (got {}D input)'314                             .format(input.dim()))315        super(SynchronizedBatchNorm3d, self)._check_input_dim(input)316