MCPcopy Create free account
hub / github.com/FreeformRobotics/OTS / _SynchronizedBatchNorm

Class _SynchronizedBatchNorm

lib/nn/modules/batchnorm.py:38–139  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

36
37
38class _SynchronizedBatchNorm(_BatchNorm):
39 def __init__(self, num_features, eps=1e-5, momentum=0.001, 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 = False
45 self._parallel_id = None
46 self._slave_pipe = None
47
48 # customed batch norm statistics
49 self._moving_average_fraction = 1. - momentum
50 self.register_buffer('_tmp_running_mean', torch.zeros(self.num_features))
51 self.register_buffer('_tmp_running_var', torch.ones(self.num_features))
52 self.register_buffer('_running_iter', torch.ones(1))
53 self._tmp_running_mean = self.running_mean.clone() * self._running_iter
54 self._tmp_running_var = self.running_var.clone() * self._running_iter
55
56 def forward(self, input):
57 # If it is not parallel computation or is in evaluation mode, use PyTorch's implementation.
58 if not (self._is_parallel and self.training):
59 return F.batch_norm(
60 input, self.running_mean, self.running_var, self.weight, self.bias,
61 self.training, self.momentum, self.eps)
62
63 # Resize the input to (B, C, -1).
64 input_shape = input.size()
65 input = input.view(input.size(0), self.num_features, -1)
66
67 # Compute the sum and square-sum.
68 sum_size = input.size(0) * input.size(2)
69 input_sum = _sum_ft(input)
70 input_ssum = _sum_ft(input ** 2)
71
72 # Reduce-and-broadcast the statistics.
73 if self._parallel_id == 0:
74 mean, inv_std = self._sync_master.run_master(_ChildMessage(input_sum, input_ssum, sum_size))
75 else:
76 mean, inv_std = self._slave_pipe.run_slave(_ChildMessage(input_sum, input_ssum, sum_size))
77
78 # Compute the output.
79 if self.affine:
80 # MJY:: Fuse the multiplication for speed.
81 output = (input - _unsqueeze_ft(mean)) * _unsqueeze_ft(inv_std * self.weight) + _unsqueeze_ft(self.bias)
82 else:
83 output = (input - _unsqueeze_ft(mean)) * _unsqueeze_ft(inv_std)
84
85 # Reshape it.
86 return output.view(input_shape)
87
88 def __data_parallel_replicate__(self, ctx, copy_id):
89 self._is_parallel = True
90 self._parallel_id = copy_id
91
92 # parallel_id == 0 means master device.
93 if self._parallel_id == 0:
94 ctx.sync_master = self._sync_master
95 else:

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected