MCPcopy Create free account
hub / github.com/kwuking/TimeMixer / Normalize

Class Normalize

layers/StandardNorm.py:4–67  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

2import torch.nn as nn
3
4class 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)

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected