MCPcopy Create free account
hub / github.com/alpha2phi/python-apps / InstanceNormalization

Class InstanceNormalization

cartoon-camera/backend/network/Transformer.py:153–177  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

151
152
153class InstanceNormalization(nn.Module):
154 def __init__(self, dim, eps=1e-9):
155 super(InstanceNormalization, self).__init__()
156 self.scale = nn.Parameter(torch.FloatTensor(dim))
157 self.shift = nn.Parameter(torch.FloatTensor(dim))
158 self.eps = eps
159 self._reset_parameters()
160
161 def _reset_parameters(self):
162 self.scale.data.uniform_()
163 self.shift.data.zero_()
164
165 def __call__(self, x):
166 n = x.size(2) * x.size(3)
167 t = x.view(x.size(0), x.size(1), n)
168 mean = torch.mean(t, 2).unsqueeze(2).unsqueeze(3).expand_as(x)
169 # Calculate the biased var. torch.var returns unbiased var
170 var = torch.var(t, 2).unsqueeze(2).unsqueeze(3).expand_as(x) * ((n - 1) / float(n))
171 scale_broadcast = self.scale.unsqueeze(1).unsqueeze(1).unsqueeze(0)
172 scale_broadcast = scale_broadcast.expand_as(x)
173 shift_broadcast = self.shift.unsqueeze(1).unsqueeze(1).unsqueeze(0)
174 shift_broadcast = shift_broadcast.expand_as(x)
175 out = (x - mean) / torch.sqrt(var + self.eps)
176 out = out * scale_broadcast + shift_broadcast
177 return out

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected