MCPcopy Create free account
hub / github.com/FunAudioLLM/FunMusic / main

Function main

tools/extract_speech_token.py:25–51  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

23
24
25def 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
54if __name__ == "__main__":

Callers 1

Calls 1

loadMethod · 0.80

Tested by

no test coverage detected