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

Function main

tools/extract_embedding.py:23–60  ·  view source on GitHub ↗
(args)

Source from the content-addressed store, hash-verified

21
22
23def main(args):
24 utt2wav, utt2spk = {}, {}
25 with open('{}/wav.scp'.format(args.dir)) as f:
26 for l in f:
27 l = l.replace('\n', '').split()
28 utt2wav[l[0]] = l[1]
29 with open('{}/utt2spk'.format(args.dir)) as f:
30 for l in f:
31 l = l.replace('\n', '').split()
32 utt2spk[l[0]] = l[1]
33
34 option = onnxruntime.SessionOptions()
35 option.graph_optimization_level = onnxruntime.GraphOptimizationLevel.ORT_ENABLE_ALL
36 option.intra_op_num_threads = 1
37 providers = ["CPUExecutionProvider"]
38 ort_session = onnxruntime.InferenceSession(args.onnx_path, sess_options=option, providers=providers)
39
40 utt2embedding, spk2embedding = {}, {}
41 for utt in tqdm(utt2wav.keys()):
42 audio, sample_rate = torchaudio.load(utt2wav[utt])
43 if sample_rate != 16000:
44 audio = torchaudio.transforms.Resample(orig_freq=sample_rate, new_freq=16000)(audio)
45 feat = kaldi.fbank(audio,
46 num_mel_bins=80,
47 dither=0,
48 sample_frequency=16000)
49 feat = feat - feat.mean(dim=0, keepdim=True)
50 embedding = ort_session.run(None, {ort_session.get_inputs()[0].name: feat.unsqueeze(dim=0).cpu().numpy()})[0].flatten().tolist()
51 utt2embedding[utt] = embedding
52 spk = utt2spk[utt]
53 if spk not in spk2embedding:
54 spk2embedding[spk] = []
55 spk2embedding[spk].append(embedding)
56 for k, v in spk2embedding.items():
57 spk2embedding[k] = torch.tensor(v).mean(dim=0).tolist()
58
59 torch.save(utt2embedding, '{}/utt2embedding.pt'.format(args.dir))
60 torch.save(spk2embedding, '{}/spk2embedding.pt'.format(args.dir))
61
62
63if __name__ == "__main__":

Callers 1

Calls 1

loadMethod · 0.80

Tested by

no test coverage detected