Views
No views yet
1
2from typing import Iterable, List, Tuple
3
4import jieba
5import onnxruntime as ort
6import soundfile as sf
7import torch
8
9
10class Lexicon:
11 def __init__(self, lexion_filename: str, tokens_filename: str):
12 tokens = dict()
13 with open(tokens_filename, encoding="utf-8") as f:
14 for line in f:
15 s, i = line.split()
16 tokens[s] = int(i)
17
18 lexicon = dict()
19 with open(lexion_filename, encoding="utf-8") as f:
20 for line in f:
21 splits = line.split()
22 word_or_phrase = splits[0]
23 phone_tone_list = splits[1:]
24 assert len(phone_tone_list) & 1 == 0, len(phone_tone_list)
25 phones = phone_tone_list[: len(phone_tone_list) // 2]
26 phones = [tokens[p] for p in phones]
27
28 tones = phone_tone_list[len(phone_tone_list) // 2 :]
29 tones = [int(t) for t in tones]
30
31 lexicon[word_or_phrase] = (phones, tones)
32
33
34 lexicon["呣"] = lexicon["母"]
35 lexicon["嗯"] = lexicon["恩"]
36 self.lexicon = lexicon
37
38 punctuation = ["!", "?", "…", ",", ".", "'", "-"]
39 for p in punctuation:
40 i = tokens[p]
41 tone = 0
42 self.lexicon[p] = ([i], [tone])
43 self.lexicon[" "] = ([tokens["_"]], [0])
44
45 def _convert(self, text: str) -> Tuple[List[int], List[int]]:
46 phones = []
47 tones = []
48
49 if text == ",":
50 text = ","
51 elif text == "。":
52 text = "."
53 elif text == "!":
54 text = "!"
55 elif text == "?":
56 text = "?"
57
58 if text not in self.lexicon:
59 print("t", text)
60 if len(text) > 1:
61 for w in text:
62 print("w", w)
63 p, t = self.convert(w)
64 if p:
65 phones += p
66 tones += t
67 return phones, tones
68
69 phones, tones = self.lexicon[text]
70 return phones, tones
71
72 def convert(self, text_list: Iterable[str]) -> Tuple[List[int], List[int]]:
73 phones = []
74 tones = []
75 for text in text_list:
76 print(text)
77 p, t = self._convert(text)
78 phones += p
79 tones += t
80 return phones, tones
81
82
83class OnnxModel:
84 def __init__(self, filename):
85 session_opts = ort.SessionOptions()
86 session_opts.inter_op_num_threads = 1
87 session_opts.intra_op_num_threads = 4
88
89 self.session_opts = session_opts
90 self.model = ort.InferenceSession(
91 filename,
92 sess_options=self.session_opts,
93 providers=["CPUExecutionProvider"],
94 )
95 meta = self.model.get_modelmeta().custom_metadata_map
96 self.bert_dim = int(meta["bert_dim"])
97 self.ja_bert_dim = int(meta["ja_bert_dim"])
98 self.add_blank = int(meta["add_blank"])
99 self.sample_rate = int(meta["sample_rate"])
100 self.speaker_id = int(meta["speaker_id"])
101 self.lang_id = int(meta["lang_id"])
102 self.sample_rate = int(meta["sample_rate"])
103
104 def __call__(self, x, tones):
105 """
106 Args:
107 x: 1-D int64 torch tensor
108 tones: 1-D int64 torch tensor
109 """
110 x = x.unsqueeze(0)
111 tones = tones.unsqueeze(0)
112
113 print(x.shape, tones.shape)
114 sid = torch.tensor([self.speaker_id], dtype=torch.int64)
115 noise_scale = torch.tensor([0.6], dtype=torch.float32)
116 length_scale = torch.tensor([1.0], dtype=torch.float32)
117 noise_scale_w = torch.tensor([0.8], dtype=torch.float32)
118
119 x_lengths = torch.tensor([x.shape[-1]], dtype=torch.int64)
120
121 y = self.model.run(
122 ["y"],
123 {
124 "x": x.numpy(),
125 "x_lengths": x_lengths.numpy(),
126 "tones": tones.numpy(),
127 "sid": sid.numpy(),
128 "noise_scale": noise_scale.numpy(),
129 "noise_scale_w": noise_scale_w.numpy(),
130 "length_scale": length_scale.numpy(),
131 },
132 )[0][0][0]
133 return y
134
135
136def main():
137 lexicon = Lexicon(lexion_filename="./lexicon.txt", tokens_filename="./tokens.txt")
138
139 text = "这是一个使用 next generation kaldi 的 text to speech 中英文例子. Thank you! 你觉得如何呢? are you ok? Fantastic! How about you?"
140 text = text.lower() # this step is crutial for split words correctly
141 s = jieba.cut(text, HMM=True)
142
143 phones, tones = lexicon.convert(s)
144
145 model = OnnxModel("./model.onnx")
146
147 if model.add_blank:
148 new_phones = [0] * (2 * len(phones) + 1)
149 new_tones = [0] * (2 * len(tones) + 1)
150
151 new_phones[1::2] = phones
152 new_tones[1::2] = tones
153
154 phones = new_phones
155 tones = new_tones
156
157 phones = torch.tensor(phones, dtype=torch.int64)
158 tones = torch.tensor(tones, dtype=torch.int64)
159
160 print(phones.shape, tones.shape)
161
162 y = model(x=phones, tones=tones)
163 sf.write("./test.wav", y, model.sample_rate)
164
165
166if __name__ == "__main__":
167 main()
168