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

Method _wavelet_decompose

layers/DWT_Decomposition.py:85–100  ·  view source on GitHub ↗
(self, x)

Source from the content-addressed store, hash-verified

83 self.affine_bias = nn.Parameter(torch.zeros((self.level + 1, self.channel)))
84
85 def _wavelet_decompose(self, x):
86 # input: x shape: batch, channel, seq
87 yl, yh = self.dwt(x)
88
89 if self.affine:
90 yl = yl.transpose(1, 2) # batch, seq, channel
91 yl = yl * self.affine_weight[0]
92 yl = yl + self.affine_bias[0]
93 yl = yl.transpose(1, 2) # batch, channel, seq
94 for i in range(self.level):
95 yh_ = yh[i].transpose(1, 2) # batch, seq, channel
96 yh_ = yh_ * self.affine_weight[i + 1]
97 yh_ = yh_ + self.affine_bias[i + 1]
98 yh[i] = yh_.transpose(1, 2) # batch, channel, seq
99
100 return yl, yh
101
102 def _wavelet_reverse_decompose(self, yl, yh):
103 if self.affine:

Callers 1

transformMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected