From 6763e150b7d38fa95d5d036d4f2c419af10c7d4b Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stardust=C2=B7=E5=87=8F?= Date: Tue, 5 Sep 2023 10:23:56 +0800 Subject: [PATCH] support ja --- webui.py | 24 +++++++++++++++++------- 1 file changed, 17 insertions(+), 7 deletions(-) diff --git a/webui.py b/webui.py index a51374d..edbf03f 100644 --- a/webui.py +++ b/webui.py @@ -40,27 +40,37 @@ def get_text(text, language_str, hps): word2ph[0] += 1 bert = get_bert(norm_text, word2ph, language_str) del word2ph + assert bert.shape[-1] == len(phone), phone - assert bert.shape[-1] == len(phone) - + if language_str=='ZH': + bert = bert + ja_bert = torch.zeros(768, len(phone)) + elif language_str=="JA": + ja_bert = bert + bert = torch.zeros(1024, len(phone)) + else: + bert = torch.zeros(1024, len(phone)) + ja_bert = torch.zeros(768, len(phone)) + assert bert.shape[-1] == len(phone), ( + bert.shape, len(phone), sum(word2ph), p1, p2, t1, t2, pold, pold2, word2ph, text, w2pho) phone = torch.LongTensor(phone) tone = torch.LongTensor(tone) language = torch.LongTensor(language) - - return bert, phone, tone, language - + return bert, ja_bert, phone, tone, language + def infer(text, sdp_ratio, noise_scale, noise_scale_w, length_scale, sid, language): global net_g - bert, phones, tones, lang_ids = get_text(text, language, hps) + bert, ja_bert, phones, tones, lang_ids = get_text(text, language, hps) with torch.no_grad(): x_tst=phones.to(device).unsqueeze(0) tones=tones.to(device).unsqueeze(0) lang_ids=lang_ids.to(device).unsqueeze(0) bert = bert.to(device).unsqueeze(0) + ja_bert = ja_bert.to(device).unsqueeze(0) x_tst_lengths = torch.LongTensor([phones.size(0)]).to(device) del phones speakers = torch.LongTensor([hps.data.spk2id[sid]]).to(device) - audio = net_g.infer(x_tst, x_tst_lengths, speakers, tones, lang_ids, bert, sdp_ratio=sdp_ratio + audio = net_g.infer(x_tst, x_tst_lengths, speakers, tones, lang_ids, bert, ja_bert, sdp_ratio=sdp_ratio , noise_scale=noise_scale, noise_scale_w=noise_scale_w, length_scale=length_scale)[0][0,0].data.cpu().float().numpy() del x_tst, tones, lang_ids, bert, x_tst_lengths, speakers return audio