| 8 | |
| 9 | |
| 10 | class ASRFeatExtractor: |
| 11 | def __init__(self, kaldi_cmvn_file): |
| 12 | self.cmvn = CMVN(kaldi_cmvn_file) if kaldi_cmvn_file != "" else None |
| 13 | self.fbank = KaldifeatFbank(num_mel_bins=80, frame_length=25, |
| 14 | frame_shift=10, dither=0.0) |
| 15 | |
| 16 | def __call__(self, wav_paths): |
| 17 | feats = [] |
| 18 | durs = [] |
| 19 | for wav_path in wav_paths: |
| 20 | sample_rate, wav_np = kaldiio.load_mat(wav_path) |
| 21 | dur = wav_np.shape[0] / sample_rate |
| 22 | fbank = self.fbank((sample_rate, wav_np)) |
| 23 | if self.cmvn is not None: |
| 24 | fbank = self.cmvn(fbank) |
| 25 | fbank = torch.from_numpy(fbank).float() |
| 26 | feats.append(fbank) |
| 27 | durs.append(dur) |
| 28 | lengths = torch.tensor([feat.size(0) for feat in feats]).long() |
| 29 | feats_pad = self.pad_feat(feats, 0.0) |
| 30 | return feats_pad, lengths, durs |
| 31 | |
| 32 | def pad_feat(self, xs, pad_value): |
| 33 | # type: (List[Tensor], int) -> Tensor |
| 34 | n_batch = len(xs) |
| 35 | max_len = max([xs[i].size(0) for i in range(n_batch)]) |
| 36 | pad = torch.ones(n_batch, max_len, *xs[0].size()[1:]).to(xs[0].device).to(xs[0].dtype).fill_(pad_value) |
| 37 | for i in range(n_batch): |
| 38 | pad[i, :xs[i].size(0)] = xs[i] |
| 39 | return pad |
| 40 | |
| 41 | |
| 42 | |