gulabpatel/First-Order-Motion
0
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 