Dev 2.3. (#242)
* Fix inputs of duration discriminator * Add LSTM * Update models.py * Update tensorboard scalar * Noise injection for minimizing modality gap * Update infer.py * support bf16 run * del unused_para flag * support bf16 config * add grad clip * fix(logger and grad):add dur grad,fix grad clip * Update webui_preprocess.py * Fix English G2P * fix(bert_gen):add pass * Pass SDP to DD * Update webui_preprocess.py * Update config.json * Update webui.py * Update chinese_bert.py * Upload webui for deploy * Update webui.py * torch.save as pt not npy * Update config.json * add freeze emo vq * Update webui_preprocess.py * Fix tone_sandhi.py * Comment up grad clip * Fix in-place addition * Add SLM discriminator * Add DDP for WD * Feat: Style text: make emotions and style similar to the style text by mixing bert (#240) (#241) * fix:(oldVersion210) Load on demand Emotion model * feat: update fastapi.py. 添加更多错误日志信息 * Switch pyopenjtalk to pyopenjtalk-prebuilt * fix: update fastapi.py. 2.2 reference适配 * Update resample.py * 修复Onnx导出的BUG (#237) * Add files via upload * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Add files via upload * Add files via upload * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Delete attentions_onnx.py * Delete models_onnx.py * Add files via upload * Add files via upload * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci * Update __init__.py * Update __init__.py * Update __init__.py * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- * Fix onnx * Format export * Feat: style-text and bert mixing (JA only) * Ensure the same tensor shape * Update * update gradio version * Fix * Style text for chinese and english (ver 2.2) * Style text for chinese and english (ver 2.1) * Style text in FastAPI * Translate style text desc in chinese --------- Co-authored-by: litagin02 <139731664+litagin02@users.noreply.github.com> Co-authored-by: Sora <654163754@qq.com> Co-authored-by: Sihan Wang <wangsihan1995@gmail.com> Co-authored-by: Ναρουσέ·μ·γιουμεμί·Χινακάννα <40709280+NaruseMioShirakana@users.noreply.github.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> * Remove CLAP * Revert "Remove CLAP" This reverts commit 62fd59bc837c580239840a2bc84b15e0663730fc. Revert * Remove CLAP * bf16 audo grad cilp * Update webui and infer utils * Update webui.py * Update webui.py * Update webui-preprocess.py * Update webui_preprocess.py --------- Co-authored-by: Sihan Wang <wangsihan1995@gmail.com> Co-authored-by: OedoSoldier <31711261+OedoSoldier@users.noreply.github.com> Co-authored-by: litagin02 <139731664+litagin02@users.noreply.github.com> Co-authored-by: Sora <654163754@qq.com> Co-authored-by: Ναρουσέ·μ·γιουμεμί·Χινακάννα <40709280+NaruseMioShirakana@users.noreply.github.com> Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com>
This commit is contained in:
@@ -18,13 +18,15 @@ def cleaned_text_to_sequence(cleaned_text, tones, language):
|
||||
return phones, tones, lang_ids
|
||||
|
||||
|
||||
def get_bert(norm_text, word2ph, language, device):
|
||||
def get_bert(norm_text, word2ph, language, device, style_text=None, style_weight=0.7):
|
||||
from .chinese_bert import get_bert_feature as zh_bert
|
||||
from .english_bert_mock import get_bert_feature as en_bert
|
||||
from .japanese_bert import get_bert_feature as jp_bert
|
||||
|
||||
lang_bert_func_map = {"ZH": zh_bert, "EN": en_bert, "JP": jp_bert}
|
||||
bert = lang_bert_func_map[language](norm_text, word2ph, device)
|
||||
bert = lang_bert_func_map[language](
|
||||
norm_text, word2ph, device, style_text, style_weight
|
||||
)
|
||||
return bert
|
||||
|
||||
|
||||
|
||||
@@ -12,7 +12,13 @@ tokenizer = AutoTokenizer.from_pretrained(LOCAL_PATH)
|
||||
models = dict()
|
||||
|
||||
|
||||
def get_bert_feature(text, word2ph, device=config.bert_gen_config.device):
|
||||
def get_bert_feature(
|
||||
text,
|
||||
word2ph,
|
||||
device=config.bert_gen_config.device,
|
||||
style_text=None,
|
||||
style_weight=0.7,
|
||||
):
|
||||
if (
|
||||
sys.platform == "darwin"
|
||||
and torch.backends.mps.is_available()
|
||||
@@ -29,12 +35,24 @@ def get_bert_feature(text, word2ph, device=config.bert_gen_config.device):
|
||||
inputs[i] = inputs[i].to(device)
|
||||
res = models[device](**inputs, output_hidden_states=True)
|
||||
res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||
|
||||
if style_text:
|
||||
style_inputs = tokenizer(style_text, return_tensors="pt")
|
||||
for i in style_inputs:
|
||||
style_inputs[i] = style_inputs[i].to(device)
|
||||
style_res = models[device](**style_inputs, output_hidden_states=True)
|
||||
style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||
style_res_mean = style_res.mean(0)
|
||||
assert len(word2ph) == len(text) + 2
|
||||
word2phone = word2ph
|
||||
phone_level_feature = []
|
||||
for i in range(len(word2phone)):
|
||||
repeat_feature = res[i].repeat(word2phone[i], 1)
|
||||
if style_text:
|
||||
repeat_feature = (
|
||||
res[i].repeat(word2phone[i], 1) * (1 - style_weight)
|
||||
+ style_res_mean.repeat(word2phone[i], 1) * style_weight
|
||||
)
|
||||
else:
|
||||
repeat_feature = res[i].repeat(word2phone[i], 1)
|
||||
phone_level_feature.append(repeat_feature)
|
||||
|
||||
phone_level_feature = torch.cat(phone_level_feature, dim=0)
|
||||
|
||||
100
text/english.py
100
text/english.py
@@ -5,6 +5,7 @@ from g2p_en import G2p
|
||||
from transformers import DebertaV2Tokenizer
|
||||
|
||||
from text import symbols
|
||||
from text.symbols import punctuation
|
||||
|
||||
current_file_path = os.path.dirname(__file__)
|
||||
CMU_DICT_PATH = os.path.join(current_file_path, "cmudict.rep")
|
||||
@@ -217,6 +218,8 @@ def refine_ph(phn):
|
||||
if re.search(r"\d$", phn):
|
||||
tone = int(phn[-1]) + 1
|
||||
phn = phn[:-1]
|
||||
else:
|
||||
tone = 3
|
||||
return phn.lower(), tone
|
||||
|
||||
|
||||
@@ -389,45 +392,84 @@ def sep_text(text):
|
||||
return words
|
||||
|
||||
|
||||
def text_to_words(text):
|
||||
tokens = tokenizer.tokenize(text)
|
||||
words = []
|
||||
for idx, t in enumerate(tokens):
|
||||
if t.startswith("▁"):
|
||||
words.append([t[1:]])
|
||||
else:
|
||||
if t in punctuation:
|
||||
if idx == len(tokens) - 1:
|
||||
words.append([f"{t}"])
|
||||
else:
|
||||
if (
|
||||
not tokens[idx + 1].startswith("▁")
|
||||
and tokens[idx + 1] not in punctuation
|
||||
):
|
||||
if idx == 0:
|
||||
words.append([])
|
||||
words[-1].append(f"{t}")
|
||||
else:
|
||||
words.append([f"{t}"])
|
||||
else:
|
||||
if idx == 0:
|
||||
words.append([])
|
||||
words[-1].append(f"{t}")
|
||||
return words
|
||||
|
||||
|
||||
def g2p(text):
|
||||
phones = []
|
||||
tones = []
|
||||
# word2ph = []
|
||||
words = sep_text(text)
|
||||
tokens = [tokenizer.tokenize(i) for i in words]
|
||||
phone_len = []
|
||||
# words = sep_text(text)
|
||||
# tokens = [tokenizer.tokenize(i) for i in words]
|
||||
words = text_to_words(text)
|
||||
|
||||
for word in words:
|
||||
if word.upper() in eng_dict:
|
||||
phns, tns = refine_syllables(eng_dict[word.upper()])
|
||||
phones.append([post_replace_ph(i) for i in phns])
|
||||
tones.append(tns)
|
||||
# word2ph.append(len(phns))
|
||||
else:
|
||||
phone_list = list(filter(lambda p: p != " ", _g2p(word)))
|
||||
phns = []
|
||||
tns = []
|
||||
for ph in phone_list:
|
||||
if ph in arpa:
|
||||
ph, tn = refine_ph(ph)
|
||||
phns.append(ph)
|
||||
tns.append(tn)
|
||||
else:
|
||||
phns.append(ph)
|
||||
tns.append(0)
|
||||
phones.append([post_replace_ph(i) for i in phns])
|
||||
tones.append(tns)
|
||||
# word2ph.append(len(phns))
|
||||
# phones = [post_replace_ph(i) for i in phones]
|
||||
temp_phones, temp_tones = [], []
|
||||
if len(word) > 1:
|
||||
if "'" in word:
|
||||
word = ["".join(word)]
|
||||
for w in word:
|
||||
if w in punctuation:
|
||||
temp_phones.append(w)
|
||||
temp_tones.append(0)
|
||||
continue
|
||||
if w.upper() in eng_dict:
|
||||
phns, tns = refine_syllables(eng_dict[w.upper()])
|
||||
temp_phones += [post_replace_ph(i) for i in phns]
|
||||
temp_tones += tns
|
||||
# w2ph.append(len(phns))
|
||||
else:
|
||||
phone_list = list(filter(lambda p: p != " ", _g2p(w)))
|
||||
phns = []
|
||||
tns = []
|
||||
for ph in phone_list:
|
||||
if ph in arpa:
|
||||
ph, tn = refine_ph(ph)
|
||||
phns.append(ph)
|
||||
tns.append(tn)
|
||||
else:
|
||||
phns.append(ph)
|
||||
tns.append(0)
|
||||
temp_phones += [post_replace_ph(i) for i in phns]
|
||||
temp_tones += tns
|
||||
phones += temp_phones
|
||||
tones += temp_tones
|
||||
phone_len.append(len(temp_phones))
|
||||
# phones = [post_replace_ph(i) for i in phones]
|
||||
|
||||
word2ph = []
|
||||
for token, phoneme in zip(tokens, phones):
|
||||
phone_len = len(phoneme)
|
||||
for token, pl in zip(words, phone_len):
|
||||
word_len = len(token)
|
||||
|
||||
aaa = distribute_phone(phone_len, word_len)
|
||||
aaa = distribute_phone(pl, word_len)
|
||||
word2ph += aaa
|
||||
|
||||
phones = ["_"] + [j for i in phones for j in i] + ["_"]
|
||||
tones = [0] + [j for i in tones for j in i] + [0]
|
||||
phones = ["_"] + phones + ["_"]
|
||||
tones = [0] + tones + [0]
|
||||
word2ph = [1] + word2ph + [1]
|
||||
assert len(phones) == len(tones), text
|
||||
assert len(phones) == sum(word2ph), text
|
||||
|
||||
@@ -13,7 +13,13 @@ tokenizer = DebertaV2Tokenizer.from_pretrained(LOCAL_PATH)
|
||||
models = dict()
|
||||
|
||||
|
||||
def get_bert_feature(text, word2ph, device=config.bert_gen_config.device):
|
||||
def get_bert_feature(
|
||||
text,
|
||||
word2ph,
|
||||
device=config.bert_gen_config.device,
|
||||
style_text=None,
|
||||
style_weight=0.7,
|
||||
):
|
||||
if (
|
||||
sys.platform == "darwin"
|
||||
and torch.backends.mps.is_available()
|
||||
@@ -30,11 +36,24 @@ def get_bert_feature(text, word2ph, device=config.bert_gen_config.device):
|
||||
inputs[i] = inputs[i].to(device)
|
||||
res = models[device](**inputs, output_hidden_states=True)
|
||||
res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||
if style_text:
|
||||
style_inputs = tokenizer(style_text, return_tensors="pt")
|
||||
for i in style_inputs:
|
||||
style_inputs[i] = style_inputs[i].to(device)
|
||||
style_res = models[device](**style_inputs, output_hidden_states=True)
|
||||
style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||
style_res_mean = style_res.mean(0)
|
||||
assert len(word2ph) == res.shape[0], (text, res.shape[0], len(word2ph))
|
||||
word2phone = word2ph
|
||||
phone_level_feature = []
|
||||
for i in range(len(word2phone)):
|
||||
repeat_feature = res[i].repeat(word2phone[i], 1)
|
||||
if style_text:
|
||||
repeat_feature = (
|
||||
res[i].repeat(word2phone[i], 1) * (1 - style_weight)
|
||||
+ style_res_mean.repeat(word2phone[i], 1) * style_weight
|
||||
)
|
||||
else:
|
||||
repeat_feature = res[i].repeat(word2phone[i], 1)
|
||||
phone_level_feature.append(repeat_feature)
|
||||
|
||||
phone_level_feature = torch.cat(phone_level_feature, dim=0)
|
||||
|
||||
@@ -13,8 +13,16 @@ tokenizer = AutoTokenizer.from_pretrained(LOCAL_PATH)
|
||||
models = dict()
|
||||
|
||||
|
||||
def get_bert_feature(text, word2ph, device=config.bert_gen_config.device):
|
||||
def get_bert_feature(
|
||||
text,
|
||||
word2ph,
|
||||
device=config.bert_gen_config.device,
|
||||
style_text=None,
|
||||
style_weight=0.7,
|
||||
):
|
||||
text = "".join(text2sep_kata(text)[0])
|
||||
if style_text:
|
||||
style_text = "".join(text2sep_kata(style_text)[0])
|
||||
if (
|
||||
sys.platform == "darwin"
|
||||
and torch.backends.mps.is_available()
|
||||
@@ -31,12 +39,25 @@ def get_bert_feature(text, word2ph, device=config.bert_gen_config.device):
|
||||
inputs[i] = inputs[i].to(device)
|
||||
res = models[device](**inputs, output_hidden_states=True)
|
||||
res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||
if style_text:
|
||||
style_inputs = tokenizer(style_text, return_tensors="pt")
|
||||
for i in style_inputs:
|
||||
style_inputs[i] = style_inputs[i].to(device)
|
||||
style_res = models[device](**style_inputs, output_hidden_states=True)
|
||||
style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu()
|
||||
style_res_mean = style_res.mean(0)
|
||||
|
||||
assert len(word2ph) == len(text) + 2
|
||||
word2phone = word2ph
|
||||
phone_level_feature = []
|
||||
for i in range(len(word2phone)):
|
||||
repeat_feature = res[i].repeat(word2phone[i], 1)
|
||||
if style_text:
|
||||
repeat_feature = (
|
||||
res[i].repeat(word2phone[i], 1) * (1 - style_weight)
|
||||
+ style_res_mean.repeat(word2phone[i], 1) * style_weight
|
||||
)
|
||||
else:
|
||||
repeat_feature = res[i].repeat(word2phone[i], 1)
|
||||
phone_level_feature.append(repeat_feature)
|
||||
|
||||
phone_level_feature = torch.cat(phone_level_feature, dim=0)
|
||||
|
||||
@@ -634,9 +634,11 @@ class ToneSandhi:
|
||||
# input seg: [('听', 'v'), ('一', 'm'), ('听', 'v')]
|
||||
# output seg: [['听一听', 'v']]
|
||||
def _merge_yi(self, seg: List[Tuple[str, str]]) -> List[Tuple[str, str]]:
|
||||
new_seg = []
|
||||
new_seg = [] * len(seg)
|
||||
# function 1
|
||||
for i, (word, pos) in enumerate(seg):
|
||||
i = 0
|
||||
while i < len(seg):
|
||||
word, pos = seg[i]
|
||||
if (
|
||||
i - 1 >= 0
|
||||
and word == "一"
|
||||
@@ -645,6 +647,7 @@ class ToneSandhi:
|
||||
and seg[i - 1][1] == "v"
|
||||
):
|
||||
new_seg[i - 1][0] = new_seg[i - 1][0] + "一" + new_seg[i - 1][0]
|
||||
i += 2
|
||||
else:
|
||||
if (
|
||||
i - 2 >= 0
|
||||
@@ -655,7 +658,8 @@ class ToneSandhi:
|
||||
continue
|
||||
else:
|
||||
new_seg.append([word, pos])
|
||||
seg = new_seg
|
||||
i += 1
|
||||
seg = [i for i in new_seg if len(i) > 0]
|
||||
new_seg = []
|
||||
# function 2
|
||||
for i, (word, pos) in enumerate(seg):
|
||||
|
||||
Reference in New Issue
Block a user