(
self,
encoder_text_hidden_states: Optional[torch.Tensor] = None,
text_attention_mask: Optional[torch.LongTensor] = None,
speaker_embeds: Optional[torch.FloatTensor] = None,
lyric_token_idx: Optional[torch.LongTensor] = None,
lyric_mask: Optional[torch.LongTensor] = None,
)
| 374 | return prompt_prenet_out |
| 375 | |
| 376 | def encode( |
| 377 | self, |
| 378 | encoder_text_hidden_states: Optional[torch.Tensor] = None, |
| 379 | text_attention_mask: Optional[torch.LongTensor] = None, |
| 380 | speaker_embeds: Optional[torch.FloatTensor] = None, |
| 381 | lyric_token_idx: Optional[torch.LongTensor] = None, |
| 382 | lyric_mask: Optional[torch.LongTensor] = None, |
| 383 | ): |
| 384 | |
| 385 | bs = encoder_text_hidden_states.shape[0] |
| 386 | device = encoder_text_hidden_states.device |
| 387 | |
| 388 | # speaker embedding |
| 389 | encoder_spk_hidden_states = self.speaker_embedder(speaker_embeds).unsqueeze(1) |
| 390 | speaker_mask = torch.ones(bs, 1, device=device) |
| 391 | |
| 392 | # genre embedding |
| 393 | encoder_text_hidden_states = self.genre_embedder(encoder_text_hidden_states) |
| 394 | |
| 395 | # lyric |
| 396 | encoder_lyric_hidden_states = self.forward_lyric_encoder( |
| 397 | lyric_token_idx=lyric_token_idx, |
| 398 | lyric_mask=lyric_mask, |
| 399 | ) |
| 400 | |
| 401 | encoder_hidden_states = torch.cat( |
| 402 | [ |
| 403 | encoder_spk_hidden_states, |
| 404 | encoder_text_hidden_states, |
| 405 | encoder_lyric_hidden_states, |
| 406 | ], |
| 407 | dim=1, |
| 408 | ) |
| 409 | encoder_hidden_mask = torch.cat( |
| 410 | [speaker_mask, text_attention_mask, lyric_mask], dim=1 |
| 411 | ) |
| 412 | return encoder_hidden_states, encoder_hidden_mask |
| 413 | |
| 414 | def decode( |
| 415 | self, |
no test coverage detected