MelStyleEncoder
| 683 | |
| 684 | |
| 685 | class MelStyleEncoder(nn.Module): |
| 686 | """MelStyleEncoder""" |
| 687 | |
| 688 | def __init__( |
| 689 | self, |
| 690 | n_mel_channels=80, |
| 691 | style_hidden=128, |
| 692 | style_vector_dim=256, |
| 693 | style_kernel_size=5, |
| 694 | style_head=2, |
| 695 | dropout=0.1, |
| 696 | ): |
| 697 | super(MelStyleEncoder, self).__init__() |
| 698 | self.in_dim = n_mel_channels |
| 699 | self.hidden_dim = style_hidden |
| 700 | self.out_dim = style_vector_dim |
| 701 | self.kernel_size = style_kernel_size |
| 702 | self.n_head = style_head |
| 703 | self.dropout = dropout |
| 704 | |
| 705 | self.spectral = nn.Sequential( |
| 706 | LinearNorm(self.in_dim, self.hidden_dim), |
| 707 | Mish(), |
| 708 | nn.Dropout(self.dropout), |
| 709 | LinearNorm(self.hidden_dim, self.hidden_dim), |
| 710 | Mish(), |
| 711 | nn.Dropout(self.dropout), |
| 712 | ) |
| 713 | |
| 714 | self.temporal = nn.Sequential( |
| 715 | Conv1dGLU(self.hidden_dim, self.hidden_dim, self.kernel_size, self.dropout), |
| 716 | Conv1dGLU(self.hidden_dim, self.hidden_dim, self.kernel_size, self.dropout), |
| 717 | ) |
| 718 | |
| 719 | self.slf_attn = MultiHeadAttention( |
| 720 | self.n_head, |
| 721 | self.hidden_dim, |
| 722 | self.hidden_dim // self.n_head, |
| 723 | self.hidden_dim // self.n_head, |
| 724 | self.dropout, |
| 725 | ) |
| 726 | |
| 727 | self.fc = LinearNorm(self.hidden_dim, self.out_dim) |
| 728 | |
| 729 | def temporal_avg_pool(self, x, mask=None): |
| 730 | if mask is None: |
| 731 | out = torch.mean(x, dim=1) |
| 732 | else: |
| 733 | len_ = (~mask).sum(dim=1).unsqueeze(1) |
| 734 | x = x.masked_fill(mask.unsqueeze(-1), 0) |
| 735 | x = x.sum(dim=1) |
| 736 | out = torch.div(x, len_) |
| 737 | return out |
| 738 | |
| 739 | def forward(self, x, mask=None): |
| 740 | x = x.transpose(1, 2) |
| 741 | if mask is not None: |
| 742 | mask = (mask.int() == 0).squeeze(1) |