(self, x)
| 41 | self.stdev = torch.sqrt(torch.var(x, dim=dim2reduce, keepdim=True, unbiased=False) + self.eps).detach() |
| 42 | |
| 43 | def _normalize(self, x): |
| 44 | if self.non_norm: |
| 45 | return x |
| 46 | if self.subtract_last: |
| 47 | x = x - self.last |
| 48 | else: |
| 49 | x = x - self.mean |
| 50 | x = x / self.stdev |
| 51 | if self.affine: |
| 52 | x = x * self.affine_weight |
| 53 | x = x + self.affine_bias |
| 54 | return x |
| 55 | |
| 56 | def _denormalize(self, x): |
| 57 | if self.non_norm: |