txya900619 commited on
Commit
64f639a
·
1 Parent(s): 239a818

feat: init hakka f5 tts

Browse files
.gitattributes CHANGED
@@ -33,3 +33,4 @@ saved_model/**/* filter=lfs diff=lfs merge=lfs -text
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
 
 
33
  *.zip filter=lfs diff=lfs merge=lfs -text
34
  *.zst filter=lfs diff=lfs merge=lfs -text
35
  *tfevents* filter=lfs diff=lfs merge=lfs -text
36
+ *.wav filter=lfs diff=lfs merge=lfs -text
app.py ADDED
@@ -0,0 +1,384 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import re
2
+ import tempfile
3
+ from importlib.resources import files
4
+
5
+ import gradio as gr
6
+ import soundfile as sf
7
+ import torch
8
+ import torchcodec
9
+ from cached_path import cached_path
10
+ from omegaconf import OmegaConf
11
+ from torchcodec.decoders import AudioDecoder
12
+
13
+ from pinyin.pinyin import get_pinyin, pinyin_configs
14
+
15
+ try:
16
+ import spaces
17
+
18
+ USING_SPACES = True
19
+ except ImportError:
20
+ USING_SPACES = False
21
+
22
+ from f5_tts.infer.utils_infer import (
23
+ device,
24
+ hop_length,
25
+ infer_process,
26
+ load_checkpoint,
27
+ load_vocoder,
28
+ mel_spec_type,
29
+ n_fft,
30
+ n_mel_channels,
31
+ ode_method,
32
+ preprocess_ref_audio_text,
33
+ remove_silence_for_generated_wav,
34
+ save_spectrogram,
35
+ target_sample_rate,
36
+ win_length,
37
+ )
38
+ from f5_tts.model import CFM, DiT
39
+ from f5_tts.model.utils import get_tokenizer
40
+
41
+
42
+ def gpu_decorator(func):
43
+ if USING_SPACES:
44
+ return spaces.GPU(func)
45
+ else:
46
+ return func
47
+
48
+
49
+ vocoder = load_vocoder()
50
+
51
+
52
+ def load_model(
53
+ model_cls,
54
+ model_cfg,
55
+ ckpt_path,
56
+ mel_spec_type=mel_spec_type,
57
+ vocab_file="",
58
+ ode_method=ode_method,
59
+ use_ema=True,
60
+ device=device,
61
+ fp16=False,
62
+ ):
63
+ if vocab_file == "":
64
+ vocab_file = str(files("f5_tts").joinpath("infer/examples/vocab.txt"))
65
+ tokenizer = "custom"
66
+
67
+ print("\nvocab : ", vocab_file)
68
+ print("token : ", tokenizer)
69
+ print("model : ", ckpt_path, "\n")
70
+
71
+ vocab_char_map, vocab_size = get_tokenizer(vocab_file, tokenizer)
72
+ model = CFM(
73
+ transformer=model_cls(
74
+ **model_cfg, text_num_embeds=vocab_size, mel_dim=n_mel_channels
75
+ ),
76
+ mel_spec_kwargs=dict(
77
+ n_fft=n_fft,
78
+ hop_length=hop_length,
79
+ win_length=win_length,
80
+ n_mel_channels=n_mel_channels,
81
+ target_sample_rate=target_sample_rate,
82
+ mel_spec_type=mel_spec_type,
83
+ ),
84
+ odeint_kwargs=dict(
85
+ method=ode_method,
86
+ ),
87
+ vocab_char_map=vocab_char_map,
88
+ ).to(device)
89
+
90
+ dtype = torch.float32 if mel_spec_type == "bigvgan" or not fp16 else None
91
+ model = load_checkpoint(model, ckpt_path, device, dtype=dtype, use_ema=use_ema)
92
+
93
+ return model
94
+
95
+
96
+ def load_f5tts(ckpt_path, vocab_path, fp16=False):
97
+ ckpt_path = str(cached_path(ckpt_path))
98
+ F5TTS_model_cfg = dict(
99
+ dim=1024, depth=22, heads=16, ff_mult=2, text_dim=512, conv_layers=4
100
+ )
101
+ vocab_path = str(cached_path(vocab_path))
102
+ return load_model(
103
+ DiT,
104
+ F5TTS_model_cfg,
105
+ ckpt_path,
106
+ vocab_file=vocab_path,
107
+ use_ema=False,
108
+ fp16=fp16,
109
+ )
110
+
111
+
112
+ OmegaConf.register_new_resolver("load_f5tts", load_f5tts)
113
+
114
+ models_config = OmegaConf.to_object(OmegaConf.load("configs/models.yaml"))
115
+
116
+ DEFAULT_MODEL_ID = list(models_config.keys())[0]
117
+ DEFAULT_DIALECT = list(pinyin_configs["lexicon"].keys())[0]
118
+
119
+
120
+ @gpu_decorator
121
+ def infer(
122
+ ref_audio_orig,
123
+ ref_text,
124
+ gen_text,
125
+ model,
126
+ remove_silence=False,
127
+ cross_fade_duration=0.15,
128
+ nfe_step=32,
129
+ speed=1,
130
+ show_info=gr.Info,
131
+ ):
132
+ if not ref_audio_orig:
133
+ gr.Warning("Please provide reference audio.")
134
+ return gr.update(), gr.update(), ref_text
135
+
136
+ if not gen_text.strip():
137
+ gr.Warning("Please enter text to generate.")
138
+ return gr.update(), gr.update(), ref_text
139
+
140
+ ref_audio, ref_text = preprocess_ref_audio_text(
141
+ ref_audio_orig, ref_text, show_info=show_info
142
+ )
143
+
144
+ ref_duration = AudioDecoder(ref_audio).metadata.duration_seconds_from_header
145
+ ref_text_len = len(re.split(r"[, ]", ref_text)) + ref_text.count(",")
146
+ gen_text_len = len(re.split(r"[, ]", gen_text)) + gen_text.count(",")
147
+ target_duration = ref_duration + ref_duration * gen_text_len / ref_text_len / speed
148
+ print(target_duration)
149
+
150
+ final_wave, final_sample_rate, combined_spectrogram = infer_process(
151
+ ref_audio,
152
+ ref_text,
153
+ gen_text,
154
+ model,
155
+ vocoder,
156
+ cross_fade_duration=cross_fade_duration,
157
+ nfe_step=nfe_step,
158
+ speed=speed,
159
+ show_info=show_info,
160
+ progress=gr.Progress(),
161
+ fix_duration=target_duration,
162
+ )
163
+
164
+ # Remove silence
165
+ if remove_silence:
166
+ with tempfile.NamedTemporaryFile(delete=False, suffix=".wav") as f:
167
+ sf.write(f.name, final_wave, final_sample_rate)
168
+ remove_silence_for_generated_wav(f.name)
169
+ final_wave = torchcodec.decoders.AudioDecoder(f.name).get_all_samples().data
170
+ final_wave = final_wave.squeeze().cpu().numpy()
171
+
172
+ # Save the spectrogram
173
+ with tempfile.NamedTemporaryFile(suffix=".png", delete=False) as tmp_spectrogram:
174
+ spectrogram_path = tmp_spectrogram.name
175
+ save_spectrogram(combined_spectrogram, spectrogram_path)
176
+
177
+ return (final_sample_rate, final_wave), spectrogram_path
178
+
179
+
180
+ def get_title():
181
+ with open("DEMO.md", encoding="utf-8") as tong:
182
+ return tong.readline().strip("# ")
183
+
184
+
185
+ with gr.Blocks(
186
+ title="臺灣客語語音生成系統",
187
+ ) as demo:
188
+ gr.Markdown(
189
+ """
190
+ # 臺灣客語語音合成系統
191
+ ### Taiwanese Hakka Text-to-Speech System
192
+ ### 研發團隊
193
+ - **[李鴻欣 Hung-Shin Lee](mailto:hungshinlee@gmail.com)**
194
+ - **[陳力瑋 Li-Wei Chen](mailto:wayne900619@gmail.com)**
195
+ ### 合作單位
196
+ - **[國立聯合大學智慧客家實驗室](https://www.gohakka.org)**
197
+ """
198
+ )
199
+
200
+ with gr.Row():
201
+ with gr.Column():
202
+ model_drop_down = gr.Dropdown(
203
+ models_config.keys(),
204
+ value=DEFAULT_MODEL_ID,
205
+ label="模型",
206
+ )
207
+
208
+ ref_dialect_radio = gr.Radio(
209
+ choices=pinyin_configs["lexicon"].keys(),
210
+ value=DEFAULT_DIALECT,
211
+ label="ref 腔調",
212
+ )
213
+
214
+ ref_text_input = gr.Textbox(
215
+ value="",
216
+ label="Reference Text",
217
+ )
218
+
219
+ ref_audio_input = gr.Audio(
220
+ type="filepath",
221
+ waveform_options=gr.WaveformOptions(
222
+ sample_rate=24000,
223
+ ),
224
+ label="Reference Audio",
225
+ )
226
+
227
+ gen_dialect_radio = gr.Radio(
228
+ choices=pinyin_configs["lexicon"].keys(),
229
+ value=DEFAULT_DIALECT,
230
+ label="gen 腔調",
231
+ )
232
+
233
+ gen_text_input = gr.Textbox(
234
+ label="Text to Generate",
235
+ value="",
236
+ )
237
+
238
+ generate_btn = gr.Button("Synthesize", variant="primary")
239
+
240
+ with gr.Accordion("Advanced Settings", open=False):
241
+ remove_silence = gr.Checkbox(
242
+ label="Remove Silences",
243
+ info="The model tends to produce silences, especially on longer audio. We can manually remove silences if needed. Note that this is an experimental feature and may produce strange results. This will also increase generation time.",
244
+ value=False,
245
+ )
246
+ speed_slider = gr.Slider(
247
+ label="Speed",
248
+ minimum=0.3,
249
+ maximum=2.0,
250
+ value=1.0,
251
+ step=0.1,
252
+ info="語速(越小越慢)",
253
+ )
254
+ nfe_slider = gr.Slider(
255
+ label="NFE Steps",
256
+ minimum=4,
257
+ maximum=64,
258
+ value=32,
259
+ step=2,
260
+ info="Set the number of denoising steps.",
261
+ )
262
+ cross_fade_duration_slider = gr.Slider(
263
+ label="Cross-Fade Duration (s)",
264
+ minimum=0.0,
265
+ maximum=1.0,
266
+ value=0.15,
267
+ step=0.01,
268
+ info="Set the duration of the cross-fade between audio clips.",
269
+ )
270
+ with gr.Column():
271
+ audio_output = gr.Audio(label="Synthesized Audio")
272
+ spectrogram_output = gr.Image(label="Spectrogram")
273
+
274
+ @gpu_decorator
275
+ def basic_tts(
276
+ model_drop_down: str,
277
+ ref_dialect_radio: str,
278
+ ref_audio_input: str,
279
+ ref_text_input: str,
280
+ gen_dialect_radio: str,
281
+ gen_text_input: str,
282
+ remove_silence: bool,
283
+ cross_fade_duration_slider: float,
284
+ nfe_slider: int,
285
+ speed_slider: float,
286
+ ):
287
+ ref_text_input = ref_text_input.strip()
288
+ if len(ref_text_input) == 0:
289
+ raise gr.Error("請勿輸入空字串。")
290
+
291
+ gen_text_input = gen_text_input.strip()
292
+ if len(gen_text_input) == 0:
293
+ raise gr.Error("請勿輸入空字串。")
294
+
295
+ words, pinyin, missing_words = get_pinyin(ref_text_input, ref_dialect_radio)
296
+ if len(missing_words) > 0:
297
+ raise gr.Error(
298
+ f"參考句子中的[{','.join(missing_words)}]目前無法轉成 pinyin。請嘗試其他句子。"
299
+ )
300
+ ref_text_input = pinyin
301
+
302
+ words, pinyin, missing_words = get_pinyin(gen_text_input, gen_dialect_radio)
303
+ if len(missing_words) > 0:
304
+ raise gr.Error(
305
+ f"生成句子中的[{','.join(missing_words)}]目前無法轉成 pinyin。請嘗試其他句子。"
306
+ )
307
+ gen_text_input = pinyin
308
+
309
+ audio_out, spectrogram_path = infer(
310
+ ref_audio_input,
311
+ ref_text_input,
312
+ gen_text_input,
313
+ models_config[model_drop_down],
314
+ remove_silence,
315
+ cross_fade_duration=cross_fade_duration_slider,
316
+ nfe_step=nfe_slider,
317
+ speed=speed_slider,
318
+ )
319
+ return audio_out, spectrogram_path
320
+
321
+ generate_btn.click(
322
+ basic_tts,
323
+ inputs=[
324
+ model_drop_down,
325
+ ref_dialect_radio,
326
+ ref_audio_input,
327
+ ref_text_input,
328
+ gen_dialect_radio,
329
+ gen_text_input,
330
+ remove_silence,
331
+ cross_fade_duration_slider,
332
+ nfe_slider,
333
+ speed_slider,
334
+ ],
335
+ outputs=[audio_output, spectrogram_output],
336
+ )
337
+
338
+ gr.Examples(
339
+ [
340
+ [
341
+ "./ref_wav/0000001_0.15-0.93.wav",
342
+ "恁早",
343
+ "sixian",
344
+ "食飯愛正經食,正毋會食到半出半入",
345
+ "sixian",
346
+ ],
347
+ [
348
+ "./ref_wav/0000002_0.15-2.73.wav",
349
+ "你今晡日著到恁派頭",
350
+ "sixian",
351
+ "食飯愛正經食,正毋會食到半出半入",
352
+ "sixian",
353
+ ],
354
+ [
355
+ "./ref_wav/0000002_0.15-2.73.wav",
356
+ "你今晡日著到恁派頭",
357
+ "sixian",
358
+ "歸條路吊等長長个花燈,祈求風調雨順,歸屋下人个心願,親像花燈下燒暖个光華",
359
+ "sixian",
360
+ ],
361
+ ],
362
+ label="範例",
363
+ inputs=[
364
+ ref_audio_input,
365
+ ref_text_input,
366
+ ref_dialect_radio,
367
+ gen_text_input,
368
+ gen_dialect_radio,
369
+ ],
370
+ )
371
+
372
+
373
+ demo.launch(
374
+ css="@import url(https://tauhu.tw/tauhu-oo.css);",
375
+ theme=gr.themes.Default(
376
+ font=(
377
+ "tauhu-oo",
378
+ gr.themes.GoogleFont("Source Sans Pro"),
379
+ "ui-sans-serif",
380
+ "system-ui",
381
+ "sans-serif",
382
+ )
383
+ ),
384
+ )
configs/models.yaml ADDED
@@ -0,0 +1 @@
 
 
1
+ step-488070: ${load_f5tts:hf://formospeech/f5-tts-hita-finetune-v1/model_481032.safetensors,hf://formospeech/f5-tts-hita-finetune-v1/vocab.txt}
configs/pinyin.yaml ADDED
@@ -0,0 +1,11 @@
 
 
 
 
 
 
 
 
 
 
 
 
1
+ delimiter_list: ${gh_download:FormoSpeech/FormoG2P, hakka/normalize/delimiters.json}
2
+ replace_dict: ${gh_download:FormoSpeech/FormoG2P, hakka/normalize/replaced_words_htia.json}
3
+ v2f_dict: ${gh_download:FormoSpeech/FormoG2P, [hakka/normalize/v2f_goyu.json, hakka/normalize/v2f_htia.json]}
4
+ preserved_list: ${gh_download:FormoSpeech/FormoG2P, hakka/normalize/preserved_words_htia.json}
5
+ lexicon:
6
+ sixian: ${gh_download:FormoSpeech/FormoG2P, hakka/sixian.json}
7
+ hailu: ${gh_download:FormoSpeech/FormoG2P, hakka/hailu.json}
8
+ dapu: ${gh_download:FormoSpeech/FormoG2P, hakka/dapu.json}
9
+ nansixian: ${gh_download:FormoSpeech/FormoG2P, hakka/nansixian.json}
10
+ raoping: ${gh_download:FormoSpeech/FormoG2P, hakka/raoping.json}
11
+ zhaoan: ${gh_download:FormoSpeech/FormoG2P, hakka/zhaoan.json}
pinyin/__init__.py ADDED
@@ -0,0 +1,27 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import time
2
+
3
+ import requests
4
+ from omegaconf import OmegaConf
5
+
6
+
7
+ def gh_download(repo, path):
8
+ paths = [path] if isinstance(path, str) else path
9
+ result = None
10
+ for path in paths:
11
+ url = f"https://raw.githubusercontent.com/{repo}/refs/heads/main/{path}"
12
+ response = requests.get(url)
13
+ if response.status_code != 200:
14
+ print(f"Status code: {response.status_code}")
15
+ raise Exception(f"Failed to download {path} from {repo}")
16
+
17
+ if result is None:
18
+ result = response.json()
19
+ elif isinstance(result, list):
20
+ result.extend(response.json())
21
+ elif isinstance(result, dict):
22
+ result.update(response.json())
23
+ time.sleep(0.5)
24
+ return result
25
+
26
+
27
+ OmegaConf.register_new_resolver("gh_download", gh_download)
pinyin/convert_digits.py ADDED
@@ -0,0 +1,180 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2024 Hung-Shin Lee (hungshinlee@gmail.com)
2
+ # Apache 2.0
3
+
4
+ import itertools
5
+ import re
6
+
7
+ c_basic = "零一二三四五六七八九"
8
+ d2c = {str(d): c for d, c in enumerate(c_basic)}
9
+ d2c["."] = "點"
10
+
11
+
12
+ def num4year(matched):
13
+ def _num4year(num):
14
+ return "{}".format("".join([c_basic[int(i)] for i in num]))
15
+
16
+ matched_str = matched.group(0)
17
+ for m in matched.groups():
18
+ matched_str = matched_str.replace(m, _num4year(m))
19
+ return matched_str
20
+
21
+
22
+ def num2chines_simple(matched):
23
+ return "{}".format("".join([d2c[i] for i in matched]))
24
+
25
+
26
+ def num4percent(matched):
27
+ matched = matched.group(1)
28
+ return "百分之{}".format(num2chinese(matched[:-1]))
29
+
30
+
31
+ def num4cellphone(matched):
32
+ matched = matched.group(1)
33
+ matched = matched.replace(" ", "").replace("-", "")
34
+ return "".join([c_basic[int(i)] for i in matched])
35
+
36
+
37
+ def num4er(matched): # 2 to 二
38
+ matched = matched.group(1)
39
+ return matched.replace("2", "二")
40
+
41
+
42
+ def num4liang(matched): # 2 to 兩
43
+ matched = matched.group(1)
44
+ return matched.replace("2", "兩")
45
+
46
+
47
+ def num4general(matched):
48
+ num = matched.group(1)
49
+ if re.match(r"[A-Za-z-─]", num[0]):
50
+ if len(num[1:]) < 3:
51
+ # MP3 or F-16
52
+ return "{}{}".format(num[0], num2chinese(num[1:]))
53
+ else:
54
+ # AM104
55
+ return "{}{}".format(num[0], num2chines_simple(num[1:]))
56
+
57
+ else:
58
+ if re.match(r"[0-9]", num[0]):
59
+ return "{}".format(num2chinese(num))
60
+ else:
61
+ return "{}{}".format(num[0], num2chinese(num[1:]))
62
+
63
+
64
+ def parse_num(text: str) -> str:
65
+ # year
66
+ text = re.sub(r"([0-9]{4})[到至]([0-9]{4})年", num4year, text)
67
+ text = re.sub(r"([0-9]{4})年", num4year, text)
68
+
69
+ # percentage
70
+ text = re.sub(r"([0-9]+\.?[0-9]?%)", num4percent, text)
71
+
72
+ # cellphone
73
+ text = re.sub(r"([0-9]{4}\s?-\s?[0-9]{6})", num4cellphone, text)
74
+
75
+ # single 2 to 二
76
+ text = re.sub(r"([^\d]2[診樓月號])", num4er, text)
77
+ text = re.sub(r"([初]2[^\d])", num4er, text)
78
+
79
+ # single 2 to 兩
80
+ text = re.sub(r"([^\d]2[^\d])", num4liang, text)
81
+
82
+ # general number
83
+ text = re.sub(r"([^0-9]?[0-9]+\.?[0-9]?)", num4general, text)
84
+
85
+ return text
86
+
87
+
88
+ def num2chinese(num, big=False, simp=False, o=False, twoalt=True) -> str:
89
+ """
90
+ Converts numbers to Chinese representations.
91
+ https://gist.github.com/gumblex/0d65cad2ba607fd14de7
92
+ `big` : use financial characters.
93
+ `simp` : use simplified characters instead of traditional characters.
94
+ `o` : use 〇 for zero.
95
+ `twoalt`: use 两/兩 for two when appropriate.
96
+ Note that `o` and `twoalt` is ignored when `big` is used,
97
+ and `twoalt` is ignored when `o` is used for formal representations.
98
+ """
99
+ # check num first
100
+ nd = str(num)
101
+ if abs(float(nd)) >= 1e48:
102
+ raise ValueError("number out of range")
103
+ elif "e" in nd:
104
+ raise ValueError("scientific notation is not supported")
105
+ c_symbol = "正负点" if simp else "正負點"
106
+ if o: # formal
107
+ twoalt = False
108
+ if big:
109
+ c_basic = "零壹贰叁肆伍陆柒捌玖" if simp else "零壹貳參肆伍陸柒捌玖"
110
+ c_unit1 = "拾佰仟"
111
+ c_twoalt = "贰" if simp else "貳"
112
+ else:
113
+ c_basic = "〇一二三四五六七八九" if o else "零一二三四五六七八九"
114
+ c_unit1 = "十百千"
115
+ if twoalt:
116
+ c_twoalt = "两" if simp else "兩"
117
+ else:
118
+ c_twoalt = "二"
119
+ c_unit2 = "万亿兆京垓秭穰沟涧正载" if simp else "萬億兆京垓秭穰溝澗正載"
120
+
121
+ def revuniq(l):
122
+ return "".join(k for k, g in itertools.groupby(reversed(l)))
123
+
124
+ nd = str(num)
125
+ result = []
126
+ if nd[0] == "+":
127
+ result.append(c_symbol[0])
128
+ elif nd[0] == "-":
129
+ result.append(c_symbol[1])
130
+ if "." in nd:
131
+ integer, remainder = nd.lstrip("+-").split(".")
132
+ else:
133
+ integer, remainder = nd.lstrip("+-"), None
134
+ if int(integer):
135
+ splitted = [integer[max(i - 4, 0) : i] for i in range(len(integer), 0, -4)]
136
+ intresult = []
137
+ for nu, unit in enumerate(splitted):
138
+ # special cases
139
+ if int(unit) == 0: # 0000
140
+ intresult.append(c_basic[0])
141
+ continue
142
+ elif nu > 0 and int(unit) == 2: # 0002
143
+ intresult.append(c_twoalt + c_unit2[nu - 1])
144
+ continue
145
+ ulist = []
146
+ unit = unit.zfill(4)
147
+ for nc, ch in enumerate(reversed(unit)):
148
+ if ch == "0":
149
+ if ulist: # ???0
150
+ ulist.append(c_basic[0])
151
+ elif nc == 0:
152
+ ulist.append(c_basic[int(ch)])
153
+ elif nc == 1 and ch == "1" and all([i == "0" for i in unit[: nc + 1]]):
154
+ # special case for tens
155
+ # edit the 'elif' if you don't like
156
+ # 十四, 三千零十四, 三千三百一十四
157
+ ulist.append(c_unit1[0])
158
+ elif nc > 1 and ch == "2":
159
+ ulist.append(c_twoalt + c_unit1[nc - 1])
160
+ else:
161
+ ulist.append(c_basic[int(ch)] + c_unit1[nc - 1])
162
+ # print(ulist)
163
+ ustr = revuniq(ulist)
164
+ if nu == 0:
165
+ intresult.append(ustr)
166
+ else:
167
+ intresult.append(ustr + c_unit2[nu - 1])
168
+ result.append(revuniq(intresult).strip(c_basic[0]))
169
+ else:
170
+ result.append(c_basic[0])
171
+ if remainder:
172
+ result.append(c_symbol[2])
173
+ result.append("".join(c_basic[int(ch)] for ch in remainder))
174
+ return "".join(result)
175
+
176
+
177
+ if __name__ == "__main__":
178
+ text = "若手機仔幾多號?吾手機仔係0964-498042。"
179
+
180
+ print(f"{text} -> {parse_num(text)}")
pinyin/pinyin.py ADDED
@@ -0,0 +1,90 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ import os
2
+ import re
3
+ from pathlib import Path
4
+
5
+ import jieba
6
+ from omegaconf import OmegaConf
7
+
8
+ from pinyin.convert_digits import parse_num
9
+ from pinyin.proc_text import (
10
+ apply_v2f,
11
+ normalize_text,
12
+ prep_regex,
13
+ run_jieba,
14
+ update_jieba_dict,
15
+ )
16
+
17
+ pinyin_configs = OmegaConf.to_object(OmegaConf.load("configs/pinyin.yaml"))
18
+ for key in pinyin_configs["preserved_list"]:
19
+ pinyin_configs["v2f_dict"].pop(key, None)
20
+ delimiter_regex, replace_regex, v2f_regex = prep_regex(
21
+ pinyin_configs["delimiter_list"],
22
+ pinyin_configs["replace_dict"],
23
+ pinyin_configs["v2f_dict"],
24
+ )
25
+
26
+
27
+ def get_pinyin(raw_text: str, dialect: str) -> tuple[str, str, str, list[str]]:
28
+ pinyin_split = re.split(r"([a-z]+\d+)", raw_text)
29
+
30
+ final_words = []
31
+ final_pinyin = []
32
+ final_missing_words = []
33
+ for hanzi_or_pinyin in pinyin_split:
34
+ if len(hanzi_or_pinyin.strip()) == 0:
35
+ continue
36
+
37
+ if re.search(r"[a-z]+\d+", hanzi_or_pinyin):
38
+ final_words.append(hanzi_or_pinyin)
39
+ final_pinyin.append(hanzi_or_pinyin)
40
+ else:
41
+ words, pinyin, missing_words = parse_hanzi_to_pinyin(
42
+ hanzi_or_pinyin, dialect
43
+ )
44
+ final_words.extend(words)
45
+ final_pinyin.extend(pinyin)
46
+ final_missing_words.extend(missing_words)
47
+
48
+ if len(final_pinyin) == 0 or len(final_missing_words) > 0:
49
+ return final_words, final_pinyin, final_missing_words
50
+
51
+ final_words = " ".join(final_words).replace(" , ", ",")
52
+ final_pinyin = " ".join(final_pinyin).replace(" , ", ",")
53
+
54
+ return final_words, final_pinyin, final_missing_words
55
+
56
+
57
+ def parse_hanzi_to_pinyin(
58
+ hanzi: str, dialect: str
59
+ ) -> tuple[list[str], list[str], list[str], list[str]]:
60
+ lexicon = pinyin_configs["lexicon"][dialect]
61
+ update_jieba_dict(
62
+ list(lexicon.keys()), Path(os.path.dirname(jieba.__file__)) / "dict.txt"
63
+ )
64
+
65
+ text = normalize_text(hanzi, pinyin_configs["replace_dict"], replace_regex)
66
+ text = parse_num(text)
67
+ text_parts = [s.strip() for s in re.split(delimiter_regex, text) if s.strip()]
68
+ text = ",".join(text_parts)
69
+ word_list = run_jieba(text)
70
+ word_list = apply_v2f(word_list, pinyin_configs["v2f_dict"], v2f_regex)
71
+ word_list = run_jieba("".join(word_list))
72
+
73
+ final_words = []
74
+ final_pinyin = []
75
+ missing_words = []
76
+ for word in word_list:
77
+ if not bool(word.strip()):
78
+ continue
79
+ if word == ",":
80
+ final_words.append(",")
81
+ final_pinyin.append(",")
82
+ elif word not in lexicon:
83
+ final_words.append(word)
84
+ missing_words.append(word)
85
+ else:
86
+ final_words.append(f"{word}")
87
+ # NOTE 只有 lexicon[word] 中的第一個才被考慮
88
+ final_pinyin.append(lexicon[word]["pinyin"][0])
89
+
90
+ return final_words, final_pinyin, missing_words
pinyin/proc_text.py ADDED
@@ -0,0 +1,84 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ # Copyright 2024 Hung-Shin Lee (hungshinlee@gmail.com)
2
+ # Apache 2.0
3
+
4
+ import re
5
+ from pathlib import Path
6
+ from unicodedata import normalize
7
+
8
+ import jieba
9
+ import opencc
10
+
11
+ jieba.setLogLevel(20)
12
+ jieba.re_han_default = re.compile(r"([\u2e80-\U000e01efa-zA-Z0-9+#&\._%\-']+)", re.U)
13
+
14
+ s2tw_converter = opencc.OpenCC("s2tw.json")
15
+
16
+
17
+ def update_jieba_dict(
18
+ lexicon: list,
19
+ jieba_dict_path: Path,
20
+ high_freq_words: list = [],
21
+ high_freq_words_weight: int = 10,
22
+ ) -> list:
23
+ lexicon = sorted(set(lexicon))
24
+
25
+ jieba_dict_path.unlink(missing_ok=True)
26
+ Path("/tmp/jieba.cache").unlink(missing_ok=True)
27
+
28
+ with jieba_dict_path.open("w", encoding="utf-8") as file:
29
+ for word in lexicon:
30
+ if word in high_freq_words:
31
+ file.write(f"{word} {len(word) * high_freq_words_weight}\n")
32
+ else:
33
+ file.write(f"{word} {len(word)}\n")
34
+
35
+ jieba.dt.initialized = False
36
+
37
+ return lexicon
38
+
39
+
40
+ def run_jieba(line: str) -> list:
41
+ # NOTE JIEBA 處理多行文本的結果會失去原本的行結構
42
+
43
+ seg_list = list(jieba.cut(line, cut_all=False, HMM=False))
44
+
45
+ return seg_list
46
+
47
+
48
+ def normalize_text(text: str, replace_dict: dict, replace_regex: str) -> str:
49
+ def replace_match(match):
50
+ return replace_dict[match.group(0)]
51
+
52
+ text = re.sub(r"\x08", "", text)
53
+ text = re.sub(r"\ufeff", "", text)
54
+ text = re.sub(r"\u0010", "", text)
55
+ text = normalize("NFKC", text)
56
+ text = re.sub(replace_regex, replace_match, text)
57
+ text = " ".join(text.split()).upper()
58
+
59
+ return text
60
+
61
+
62
+ def apply_v2f(word_list: list, v2f_dict: dict, v2f_regex: str) -> list:
63
+ result = []
64
+ for word in word_list:
65
+ result.append(re.sub(v2f_regex, lambda x: v2f_dict[x.group(0)], word))
66
+
67
+ return result
68
+
69
+
70
+ def prep_regex(
71
+ delimiter_list: list, replace_dict: dict = {}, v2f_dict: dict = {}
72
+ ) -> tuple[str, str, str]:
73
+ delimiter_regex = "|".join(map(re.escape, delimiter_list))
74
+
75
+ replace_regex = ""
76
+ if len(replace_dict):
77
+ sorted_keys = sorted(replace_dict.keys(), key=len, reverse=True)
78
+ replace_regex = "|".join(map(re.escape, sorted_keys))
79
+
80
+ v2f_regex = ""
81
+ if len(v2f_dict):
82
+ v2f_regex = "|".join(map(re.escape, v2f_dict.keys()))
83
+
84
+ return delimiter_regex, replace_regex, v2f_regex
ref_wav/0000001_0.15-0.93.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:78ef5c480c70cb01e6065a59bf896b4598a638a075d5a1d287342c018b7a7129
3
+ size 37484
ref_wav/0000002_0.15-2.73.wav ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ version https://git-lfs.github.com/spec/v1
2
+ oid sha256:560d43f7e4dc1cc7a994f98ca9870a8c5dad7be1a89fe2d26b534519370a5bd2
3
+ size 123884
requirements.txt ADDED
@@ -0,0 +1,3 @@
 
 
 
 
1
+ omegaconf
2
+ opencc
3
+ f5-tts