| 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: |