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

Class ASRFeatExtractor

fireredasr/data/asr_feat.py:10–39  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

8
9
10class 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

Callers 1

from_pretrainedMethod · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected