Initilize duration predictor module. Args: idim (int): Input dimension. n_layers (int, optional): Number of convolutional layers. n_chans (int, optional): Number of channels of convolutional layers. kernel_size (int, optional): Kernel size of c
(self, idim, n_layers=2, n_chans=384, kernel_size=3, dropout_rate=0.1, offset=1.0, padding='SAME')
| 68 | """ |
| 69 | |
| 70 | def __init__(self, idim, n_layers=2, n_chans=384, kernel_size=3, dropout_rate=0.1, offset=1.0, padding='SAME'): |
| 71 | """Initilize duration predictor module. |
| 72 | Args: |
| 73 | idim (int): Input dimension. |
| 74 | n_layers (int, optional): Number of convolutional layers. |
| 75 | n_chans (int, optional): Number of channels of convolutional layers. |
| 76 | kernel_size (int, optional): Kernel size of convolutional layers. |
| 77 | dropout_rate (float, optional): Dropout rate. |
| 78 | offset (float, optional): Offset value to avoid nan in log domain. |
| 79 | """ |
| 80 | super(DurationPredictor, self).__init__() |
| 81 | self.offset = offset |
| 82 | self.conv = torch.nn.ModuleList() |
| 83 | self.kernel_size = kernel_size |
| 84 | self.padding = padding |
| 85 | for idx in range(n_layers): |
| 86 | in_chans = idim if idx == 0 else n_chans |
| 87 | self.conv += [torch.nn.Sequential( |
| 88 | torch.nn.ConstantPad1d(((kernel_size - 1) // 2, (kernel_size - 1) // 2) |
| 89 | if padding == 'SAME' |
| 90 | else (kernel_size - 1, 0), 0), |
| 91 | torch.nn.Conv1d(in_chans, n_chans, kernel_size, stride=1, padding=0), |
| 92 | torch.nn.ReLU(), |
| 93 | LayerNorm(n_chans, dim=1), |
| 94 | torch.nn.Dropout(dropout_rate) |
| 95 | )] |
| 96 | if hparams['dur_loss'] in ['mse', 'huber']: |
| 97 | odims = 1 |
| 98 | elif hparams['dur_loss'] == 'mog': |
| 99 | odims = 15 |
| 100 | elif hparams['dur_loss'] == 'crf': |
| 101 | odims = 32 |
| 102 | from torchcrf import CRF |
| 103 | self.crf = CRF(odims, batch_first=True) |
| 104 | self.linear = torch.nn.Linear(n_chans, odims) |
| 105 | |
| 106 | def _forward(self, xs, x_masks=None, is_inference=False): |
| 107 | xs = xs.transpose(1, -1) # (B, idim, Tmax) |