(args)
| 23 | |
| 24 | |
| 25 | def main(args): |
| 26 | utt2wav = {} |
| 27 | with open('{}/wav.scp'.format(args.dir)) as f: |
| 28 | for l in f: |
| 29 | l = l.replace('\n', '').split() |
| 30 | utt2wav[l[0]] = l[1] |
| 31 | |
| 32 | option = onnxruntime.SessionOptions() |
| 33 | option.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_ENABLE_ALL |
| 34 | option.intra_op_num_threads = 1 |
| 35 | providers = ["CUDAExecutionProvider"] |
| 36 | ort_session = onnxruntime.InferenceSession(args.onnx_path, sess_options=option, providers=providers) |
| 37 | |
| 38 | utt2speech_token = {} |
| 39 | for utt in tqdm(utt2wav.keys()): |
| 40 | audio, sample_rate = torchaudio.load(utt2wav[utt]) |
| 41 | if sample_rate != 16000: |
| 42 | audio = torchaudio.transforms.Resample(orig_freq=sample_rate, new_freq=16000)(audio) |
| 43 | if audio.shape[1] / 16000 > 30: |
| 44 | logging.warning('do not support extract speech token for audio longer than 30s') |
| 45 | speech_token = [] |
| 46 | else: |
| 47 | feat = whisper.log_mel_spectrogram(audio, n_mels=128) |
| 48 | speech_token = ort_session.run(None, {ort_session.get_inputs()[0].name: feat.detach().cpu().numpy(), |
| 49 | ort_session.get_inputs()[1].name: np.array([feat.shape[2]], dtype=np.int32)})[0].flatten().tolist() |
| 50 | utt2speech_token[utt] = speech_token |
| 51 | torch.save(utt2speech_token, '{}/utt2speech_token.pt'.format(args.dir)) |
| 52 | |
| 53 | |
| 54 | if __name__ == "__main__": |
no test coverage detected