Refactor and add merge
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, style_text=None, style_weight=0.7):
|
||||
def get_bert(
|
||||
norm_text, word2ph, language, device, assist_text=None, assist_text_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, style_text, style_weight
|
||||
norm_text, word2ph, device, assist_text, assist_text_weight
|
||||
)
|
||||
return bert
|
||||
|
||||
@@ -16,8 +16,8 @@ def get_bert_feature(
|
||||
text,
|
||||
word2ph,
|
||||
device=config.bert_gen_config.device,
|
||||
style_text=None,
|
||||
style_weight=0.7,
|
||||
assist_text=None,
|
||||
assist_text_weight=0.7,
|
||||
):
|
||||
if (
|
||||
sys.platform == "darwin"
|
||||
@@ -37,8 +37,8 @@ def get_bert_feature(
|
||||
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")
|
||||
if assist_text:
|
||||
style_inputs = tokenizer(assist_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)
|
||||
@@ -48,10 +48,10 @@ def get_bert_feature(
|
||||
word2phone = word2ph
|
||||
phone_level_feature = []
|
||||
for i in range(len(word2phone)):
|
||||
if style_text:
|
||||
if assist_text:
|
||||
repeat_feature = (
|
||||
res[i].repeat(word2phone[i], 1) * (1 - style_weight)
|
||||
+ style_res_mean.repeat(word2phone[i], 1) * style_weight
|
||||
res[i].repeat(word2phone[i], 1) * (1 - assist_text_weight)
|
||||
+ style_res_mean.repeat(word2phone[i], 1) * assist_text_weight
|
||||
)
|
||||
else:
|
||||
repeat_feature = res[i].repeat(word2phone[i], 1)
|
||||
|
||||
@@ -17,8 +17,8 @@ def get_bert_feature(
|
||||
text,
|
||||
word2ph,
|
||||
device=config.bert_gen_config.device,
|
||||
style_text=None,
|
||||
style_weight=0.7,
|
||||
assist_text=None,
|
||||
assist_text_weight=0.7,
|
||||
):
|
||||
if (
|
||||
sys.platform == "darwin"
|
||||
@@ -38,8 +38,8 @@ def get_bert_feature(
|
||||
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")
|
||||
if assist_text:
|
||||
style_inputs = tokenizer(assist_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)
|
||||
@@ -49,10 +49,10 @@ def get_bert_feature(
|
||||
word2phone = word2ph
|
||||
phone_level_feature = []
|
||||
for i in range(len(word2phone)):
|
||||
if style_text:
|
||||
if assist_text:
|
||||
repeat_feature = (
|
||||
res[i].repeat(word2phone[i], 1) * (1 - style_weight)
|
||||
+ style_res_mean.repeat(word2phone[i], 1) * style_weight
|
||||
res[i].repeat(word2phone[i], 1) * (1 - assist_text_weight)
|
||||
+ style_res_mean.repeat(word2phone[i], 1) * assist_text_weight
|
||||
)
|
||||
else:
|
||||
repeat_feature = res[i].repeat(word2phone[i], 1)
|
||||
|
||||
@@ -17,12 +17,12 @@ def get_bert_feature(
|
||||
text,
|
||||
word2ph,
|
||||
device=config.bert_gen_config.device,
|
||||
style_text=None,
|
||||
style_weight=0.7,
|
||||
assist_text=None,
|
||||
assist_text_weight=0.7,
|
||||
):
|
||||
text = "".join(text2sep_kata(text)[0])
|
||||
if style_text:
|
||||
style_text = "".join(text2sep_kata(style_text)[0])
|
||||
if assist_text:
|
||||
assist_text = "".join(text2sep_kata(assist_text)[0])
|
||||
if (
|
||||
sys.platform == "darwin"
|
||||
and torch.backends.mps.is_available()
|
||||
@@ -41,8 +41,8 @@ def get_bert_feature(
|
||||
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")
|
||||
if assist_text:
|
||||
style_inputs = tokenizer(assist_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)
|
||||
@@ -53,10 +53,10 @@ def get_bert_feature(
|
||||
word2phone = word2ph
|
||||
phone_level_feature = []
|
||||
for i in range(len(word2phone)):
|
||||
if style_text:
|
||||
if assist_text:
|
||||
repeat_feature = (
|
||||
res[i].repeat(word2phone[i], 1) * (1 - style_weight)
|
||||
+ style_res_mean.repeat(word2phone[i], 1) * style_weight
|
||||
res[i].repeat(word2phone[i], 1) * (1 - assist_text_weight)
|
||||
+ style_res_mean.repeat(word2phone[i], 1) * assist_text_weight
|
||||
)
|
||||
else:
|
||||
repeat_feature = res[i].repeat(word2phone[i], 1)
|
||||
|
||||
Reference in New Issue
Block a user