| 100 | return yl, yh |
| 101 | |
| 102 | def _wavelet_reverse_decompose(self, yl, yh): |
| 103 | if self.affine: |
| 104 | yl = yl.transpose(1, 2) # batch, seq, channel |
| 105 | yl = yl - self.affine_bias[0] |
| 106 | yl = yl / (self.affine_weight[0] + self.eps) |
| 107 | yl = yl.transpose(1, 2) # batch, channel, seq |
| 108 | for i in range(self.level): |
| 109 | yh_ = yh[i].transpose(1, 2) # batch, seq, channel |
| 110 | yh_ = yh_ - self.affine_bias[i + 1] |
| 111 | yh_ = yh_ / (self.affine_weight[i + 1] + self.eps) |
| 112 | yh[i] = yh_.transpose(1, 2) # batch, channel, seq |
| 113 | |
| 114 | x = self.idwt((yl, yh)) |
| 115 | return x # shape: batch, channel, seq |
| 116 | |
| 117 | |
| 118 | ############################################################################################### |