| 151 | |
| 152 | |
| 153 | class 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 |