| 54 | |
| 55 | |
| 56 | class WaveGenerator(nn.Module): |
| 57 | def __init__( |
| 58 | self, |
| 59 | input_channel, |
| 60 | channels, |
| 61 | rates, |
| 62 | kernel_sizes, |
| 63 | d_out: int = 1, |
| 64 | ): |
| 65 | super().__init__() |
| 66 | |
| 67 | # Add first conv layer |
| 68 | layers = [WNConv1d(input_channel, channels, kernel_size=7, padding=3)] |
| 69 | |
| 70 | # Add upsampling + MRF blocks |
| 71 | for i, (kernel_size, stride) in enumerate(zip(kernel_sizes, rates)): |
| 72 | input_dim = channels // 2**i |
| 73 | output_dim = channels // 2 ** (i + 1) |
| 74 | layers += [DecoderBlock(input_dim, output_dim, kernel_size, stride)] |
| 75 | |
| 76 | # Add final conv layer |
| 77 | layers += [ |
| 78 | Snake1d(output_dim), |
| 79 | WNConv1d(output_dim, d_out, kernel_size=7, padding=3), |
| 80 | nn.Tanh(), |
| 81 | ] |
| 82 | |
| 83 | self.model = nn.Sequential(*layers) |
| 84 | |
| 85 | self.apply(init_weights) |
| 86 | |
| 87 | def forward(self, x): |
| 88 | return self.model(x) |
no outgoing calls
no test coverage detected