| 6 | |
| 7 | class Prenet(nn.Module): |
| 8 | def __init__(self, in_dim=80, out_dim=256, kernel=5, n_layers=3, strides=None): |
| 9 | super(Prenet, self).__init__() |
| 10 | padding = kernel // 2 |
| 11 | self.layers = [] |
| 12 | self.strides = strides if strides is not None else [1] * n_layers |
| 13 | for l in range(n_layers): |
| 14 | self.layers.append(nn.Sequential( |
| 15 | nn.Conv1d(in_dim, out_dim, kernel_size=kernel, padding=padding, stride=self.strides[l]), |
| 16 | nn.ReLU(), |
| 17 | nn.BatchNorm1d(out_dim) |
| 18 | )) |
| 19 | in_dim = out_dim |
| 20 | self.layers = nn.ModuleList(self.layers) |
| 21 | self.out_proj = nn.Linear(out_dim, out_dim) |
| 22 | |
| 23 | def forward(self, x): |
| 24 | """ |