Forward pass for training. Args: input: Input audio/text data. input_lengths: Lengths of input. **kwargs: Additional keyword arguments.
(
self,
input: torch.Tensor,
input_lengths,
**kwargs,
)
| 147 | return self.n_mels * self.lfr_m |
| 148 | |
| 149 | def forward( |
| 150 | self, |
| 151 | input: torch.Tensor, |
| 152 | input_lengths, |
| 153 | **kwargs, |
| 154 | ) -> Tuple[torch.Tensor, torch.Tensor]: |
| 155 | """Forward pass for training. |
| 156 | |
| 157 | Args: |
| 158 | input: Input audio/text data. |
| 159 | input_lengths: Lengths of input. |
| 160 | **kwargs: Additional keyword arguments. |
| 161 | """ |
| 162 | batch_size = input.size(0) |
| 163 | feats = [] |
| 164 | feats_lens = [] |
| 165 | for i in range(batch_size): |
| 166 | waveform_length = input_lengths[i] |
| 167 | waveform = input[i][:waveform_length] |
| 168 | if self.upsacle_samples: |
| 169 | waveform = waveform * (1 << 15) |
| 170 | waveform = waveform.unsqueeze(0) |
| 171 | mat = kaldi.fbank( |
| 172 | waveform, |
| 173 | num_mel_bins=self.n_mels, |
| 174 | frame_length=min(self.frame_length,waveform_length/self.fs*1000), |
| 175 | frame_shift=self.frame_shift, |
| 176 | dither=self.dither, |
| 177 | energy_floor=0.0, |
| 178 | window_type=self.window, |
| 179 | sample_frequency=self.fs, |
| 180 | snip_edges=self.snip_edges, |
| 181 | ) |
| 182 | |
| 183 | if self.lfr_m != 1 or self.lfr_n != 1: |
| 184 | mat = apply_lfr(mat, self.lfr_m, self.lfr_n) |
| 185 | if self.cmvn is not None: |
| 186 | mat = apply_cmvn(mat, self.cmvn) |
| 187 | feat_length = mat.size(0) |
| 188 | feats.append(mat) |
| 189 | feats_lens.append(feat_length) |
| 190 | |
| 191 | feats_lens = torch.as_tensor(feats_lens) |
| 192 | if batch_size == 1: |
| 193 | feats_pad = feats[0][None, :, :] |
| 194 | else: |
| 195 | feats_pad = pad_sequence(feats, batch_first=True, padding_value=0.0) |
| 196 | return feats_pad, feats_lens |
| 197 | |
| 198 | def forward_fbank( |
| 199 | self, input: torch.Tensor, input_lengths: torch.Tensor |
nothing calls this directly
no test coverage detected