MCPcopy Create free account
hub / github.com/FireRedTeam/FireRedASR / Conv2dSubsampling

Class Conv2dSubsampling

fireredasr/models/module/conformer_encoder.py:81–117  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

79
80
81class 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
120class RelPositionalEncoding(torch.nn.Module):

Callers 1

__init__Method · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected