(self, input_)
| 47 | init.zeros_(self.bias) |
| 48 | |
| 49 | def forward(self, input_): |
| 50 | batchsize, channels, height, width = input_.size() |
| 51 | numel = batchsize * height * width |
| 52 | input_ = input_.permute(1, 0, 2, 3).contiguous().view(channels, numel) |
| 53 | sum_ = input_.sum(1) |
| 54 | sum_of_square = input_.pow(2).sum(1) |
| 55 | mean = sum_ / numel |
| 56 | sumvar = sum_of_square - sum_ * mean |
| 57 | |
| 58 | self.running_mean = ( |
| 59 | 1 - self.momentum |
| 60 | ) * self.running_mean + self.momentum * mean.detach() |
| 61 | unbias_var = sumvar / (numel - 1) |
| 62 | self.running_var = ( |
| 63 | 1 - self.momentum |
| 64 | ) * self.running_var + self.momentum * unbias_var.detach() |
| 65 | |
| 66 | bias_var = sumvar / numel |
| 67 | inv_std = 1 / (bias_var + self.eps).pow(0.5) |
| 68 | output = (input_ - mean.unsqueeze(1)) * inv_std.unsqueeze( |
| 69 | 1 |
| 70 | ) * self.weight.unsqueeze(1) + self.bias.unsqueeze(1) |
| 71 | |
| 72 | return ( |
| 73 | output.view(channels, batchsize, height, width) |
| 74 | .permute(1, 0, 2, 3) |
| 75 | .contiguous() |
| 76 | ) |
nothing calls this directly
no outgoing calls
no test coverage detected