| 2 | import torch.nn as nn |
| 3 | |
| 4 | class Normalize(nn.Module): |
| 5 | def __init__(self, num_features: int, eps=1e-5, affine=False, subtract_last=False, non_norm=False): |
| 6 | """ |
| 7 | :param num_features: the number of features or channels |
| 8 | :param eps: a value added for numerical stability |
| 9 | :param affine: if True, RevIN has learnable affine parameters |
| 10 | """ |
| 11 | super(Normalize, self).__init__() |
| 12 | self.num_features = num_features |
| 13 | self.eps = eps |
| 14 | self.affine = affine |
| 15 | self.subtract_last = subtract_last |
| 16 | self.non_norm = non_norm |
| 17 | if self.affine: |
| 18 | self._init_params() |
| 19 | |
| 20 | def forward(self, x, mode: str): |
| 21 | if mode == 'norm': |
| 22 | self._get_statistics(x) |
| 23 | x = self._normalize(x) |
| 24 | elif mode == 'denorm': |
| 25 | x = self._denormalize(x) |
| 26 | else: |
| 27 | raise NotImplementedError |
| 28 | return x |
| 29 | |
| 30 | def _init_params(self): |
| 31 | # initialize RevIN params: (C,) |
| 32 | self.affine_weight = nn.Parameter(torch.ones(self.num_features)) |
| 33 | self.affine_bias = nn.Parameter(torch.zeros(self.num_features)) |
| 34 | |
| 35 | def _get_statistics(self, x): |
| 36 | dim2reduce = tuple(range(1, x.ndim - 1)) |
| 37 | if self.subtract_last: |
| 38 | self.last = x[:, -1, :].unsqueeze(1) |
| 39 | else: |
| 40 | self.mean = torch.mean(x, dim=dim2reduce, keepdim=True).detach() |
| 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: |
| 58 | return x |
| 59 | if self.affine: |
| 60 | x = x - self.affine_bias |
| 61 | x = x / (self.affine_weight + self.eps * self.eps) |