(args)
| 21 | |
| 22 | |
| 23 | def 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 | |
| 63 | if __name__ == "__main__": |
no test coverage detected