| 79 | |
| 80 | |
| 81 | class Conv2dSubsampling(nn.Module): |
| 82 | def __init__(self, idim, d_model, out_channels=32): |
| 83 | super().__init__() |
| 84 | self.conv = nn.ModuleList([ |
| 85 | nn.Conv2d(1, out_channels, 3, 2), |
| 86 | nn.ReLU(), |
| 87 | nn.Conv2d(out_channels, out_channels, 3, 2), |
| 88 | nn.ReLU(), |
| 89 | ]) |
| 90 | subsample_idim = ((idim - 1) // 2 - 1) // 2 |
| 91 | self.out = nn.Linear(out_channels * subsample_idim, d_model) |
| 92 | |
| 93 | self.subsampling = 4 |
| 94 | left_context = right_context = 3 # both exclude currect frame |
| 95 | self.context = left_context + 1 + right_context # 7 |
| 96 | |
| 97 | def forward(self, x, x_mask, input_lengths): |
| 98 | x = x.unsqueeze(1) |
| 99 | for layer in self.conv: |
| 100 | x = layer(x) |
| 101 | N, C, T, D = x.size() |
| 102 | x = self.out(x.transpose(1, 2).contiguous().view(N, T, C * D)) |
| 103 | |
| 104 | # Arithmetically calculate the output lengths, which is more robust for TRT. |
| 105 | # The conv layers in Conv2dSubsampling have kernel_size=3, stride=2 |
| 106 | # which corresponds to a length calculation of (L-3)//2 + 1 |
| 107 | output_lengths = (input_lengths - 3) // 2 + 1 |
| 108 | output_lengths = (output_lengths - 3) // 2 + 1 |
| 109 | |
| 110 | # Re-create the mask from the correctly calculated lengths. |
| 111 | max_len = x.size(1) |
| 112 | device = x.device |
| 113 | indices = torch.arange(max_len, device=device).expand(N, -1) |
| 114 | mask = indices < output_lengths.unsqueeze(1) |
| 115 | mask = mask.unsqueeze(1) # (N, 1, T_out) |
| 116 | |
| 117 | return x, output_lengths, mask |
| 118 | |
| 119 | |
| 120 | class RelPositionalEncoding(torch.nn.Module): |