| 180 | |
| 181 | class TextEncoder(nn.Module): |
| 182 | def __init__( |
| 183 | self, |
| 184 | out_channels, |
| 185 | hidden_channels, |
| 186 | filter_channels, |
| 187 | n_heads, |
| 188 | n_layers, |
| 189 | kernel_size, |
| 190 | p_dropout, |
| 191 | latent_channels=192, |
| 192 | version = "v2", |
| 193 | ): |
| 194 | super().__init__() |
| 195 | self.out_channels = out_channels |
| 196 | self.hidden_channels = hidden_channels |
| 197 | self.filter_channels = filter_channels |
| 198 | self.n_heads = n_heads |
| 199 | self.n_layers = n_layers |
| 200 | self.kernel_size = kernel_size |
| 201 | self.p_dropout = p_dropout |
| 202 | self.latent_channels = latent_channels |
| 203 | self.version = version |
| 204 | |
| 205 | self.ssl_proj = nn.Conv1d(768, hidden_channels, 1) |
| 206 | |
| 207 | self.encoder_ssl = attentions.Encoder( |
| 208 | hidden_channels, |
| 209 | filter_channels, |
| 210 | n_heads, |
| 211 | n_layers // 2, |
| 212 | kernel_size, |
| 213 | p_dropout, |
| 214 | ) |
| 215 | |
| 216 | self.encoder_text = attentions.Encoder( |
| 217 | hidden_channels, filter_channels, n_heads, n_layers, kernel_size, p_dropout |
| 218 | ) |
| 219 | |
| 220 | if self.version == "v1": |
| 221 | symbols = symbols_v1.symbols |
| 222 | else: |
| 223 | symbols = symbols_v2.symbols |
| 224 | self.text_embedding = nn.Embedding(len(symbols), hidden_channels) |
| 225 | |
| 226 | self.mrte = MRTE() |
| 227 | |
| 228 | self.encoder2 = attentions.Encoder( |
| 229 | hidden_channels, |
| 230 | filter_channels, |
| 231 | n_heads, |
| 232 | n_layers // 2, |
| 233 | kernel_size, |
| 234 | p_dropout, |
| 235 | ) |
| 236 | |
| 237 | self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1) |
| 238 | |
| 239 | def forward(self, y, y_lengths, text, text_lengths, ge, speed=1,test=None): |