From 33d153505a3c25fdd1f70f513b2877b18c6476ab Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Stardust=C2=B7=E5=87=8F?= Date: Wed, 8 Nov 2023 17:09:15 +0800 Subject: [PATCH] sync emo branch (#159) MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit * Update README.md * 更新 bert_models.json * fix * Update data_utils.py * Update infer.py * performance improve * Feat: support auto split in webui (#158) * Feat: support auto split in webui * [pre-commit.ci] auto fixes from pre-commit.com hooks for more information, see https://pre-commit.ci --------- Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> --------- Co-authored-by: Sora Co-authored-by: pre-commit-ci[bot] <66853113+pre-commit-ci[bot]@users.noreply.github.com> --- README.md | 4 +++- bert/bert_models.json | 2 +- data_utils.py | 24 +++++------------------- infer.py | 10 +++++----- server_fastapi.py | 2 ++ tools/log.py | 4 +++- train_ms.py | 4 ++-- 7 files changed, 21 insertions(+), 29 deletions(-) diff --git a/README.md b/README.md index 0e9edbf..a76aa78 100644 --- a/README.md +++ b/README.md @@ -1,7 +1,9 @@ # Bert-VITS2 VITS2 Backbone with bert - +# 紧急通知 +我们在2.0版本中发现了重大bug,该bug导致日文和英文bert被置0后训练,即失去bert效果。 +我们将重炼2.0版本底模,已经开炉的建议关炉静候。 ## 请注意,本项目核心思路来源于[anyvoiceai/MassTTS](https://github.com/anyvoiceai/MassTTS) 一个非常好的tts项目 ## MassTTS的演示demo为[ai版峰哥锐评峰哥本人,并找回了在金三角失落的腰子](https://www.bilibili.com/video/BV1w24y1c7z9) diff --git a/bert/bert_models.json b/bert/bert_models.json index 61d4e6a..78721c3 100644 --- a/bert/bert_models.json +++ b/bert/bert_models.json @@ -1,7 +1,7 @@ { "deberta-v2-large-japanese": { "repo_id": "ku-nlp/deberta-v2-large-japanese", - "files": ["spm.model", "pytorch_model.bin"] + "files": ["pytorch_model.bin"] }, "chinese-roberta-wwm-ext-large": { "repo_id": "hfl/chinese-roberta-wwm-ext-large", diff --git a/data_utils.py b/data_utils.py index d6526c4..1e297b2 100644 --- a/data_utils.py +++ b/data_utils.py @@ -145,38 +145,24 @@ class TextAudioSpeakerLoader(torch.utils.data.Dataset): word2ph[0] += 1 bert_path = wav_path.replace(".wav", ".bert.pt") try: - bert = torch.load(bert_path) - assert bert.shape[-1] == len(phone) + bert_ori = torch.load(bert_path) + assert bert_ori.shape[-1] == len(phone) except Exception as e: logger.warn("Bert load Failed") logger.warn(e) if language_str == "ZH": - bert = bert + bert = bert_ori ja_bert = torch.zeros(1024, len(phone)) en_bert = torch.zeros(1024, len(phone)) elif language_str == "JP": bert = torch.zeros(1024, len(phone)) - ja_bert = bert + ja_bert = bert_ori en_bert = torch.zeros(1024, len(phone)) elif language_str == "EN": bert = torch.zeros(1024, len(phone)) ja_bert = torch.zeros(1024, len(phone)) - en_bert = bert - assert bert.shape[-1] == len(phone), ( - bert.shape, - len(phone), - sum(word2ph), - p1, - p2, - t1, - t2, - pold, - pold2, - word2ph, - text, - w2pho, - ) + en_bert = bert_ori phone = torch.LongTensor(phone) tone = torch.LongTensor(tone) language = torch.LongTensor(language) diff --git a/infer.py b/infer.py index 8497998..441e52c 100644 --- a/infer.py +++ b/infer.py @@ -85,22 +85,22 @@ def get_text(text, language_str, hps, device): for i in range(len(word2ph)): word2ph[i] = word2ph[i] * 2 word2ph[0] += 1 - bert = get_bert(norm_text, word2ph, language_str, device) + bert_ori = get_bert(norm_text, word2ph, language_str, device) del word2ph - assert bert.shape[-1] == len(phone), phone + assert bert_ori.shape[-1] == len(phone), phone if language_str == "ZH": - bert = bert + bert = bert_ori ja_bert = torch.zeros(1024, len(phone)) en_bert = torch.zeros(1024, len(phone)) elif language_str == "JP": bert = torch.zeros(1024, len(phone)) - ja_bert = bert + ja_bert = bert_ori en_bert = torch.zeros(1024, len(phone)) elif language_str == "EN": bert = torch.zeros(1024, len(phone)) ja_bert = torch.zeros(1024, len(phone)) - en_bert = bert + en_bert = bert_ori else: raise ValueError("language_str should be ZH, JP or EN") diff --git a/server_fastapi.py b/server_fastapi.py index 1f22860..00debb3 100644 --- a/server_fastapi.py +++ b/server_fastapi.py @@ -5,6 +5,7 @@ import logging import gc import random +import gradio import numpy as np import utils from fastapi import FastAPI, Query, Request @@ -245,6 +246,7 @@ if __name__ == "__main__": ) audios.append(np.zeros((int)(44100 * 0.3))) audio = np.concatenate(audios) + audio = gradio.processing_utils.convert_to_16_bit_wav(audio) wavContent = BytesIO() wavfile.write( wavContent, loaded_models.models[model_id].hps.data.sampling_rate, audio diff --git a/tools/log.py b/tools/log.py index 626b9e9..47e51ae 100644 --- a/tools/log.py +++ b/tools/log.py @@ -9,6 +9,8 @@ import sys logger.remove() # 自定义格式并添加到标准输出 -log_format = "{time:MM-DD HH:mm:ss} [{level}] | {message}" +log_format = ( + "{time:MM-DD HH:mm:ss} [{level}] | {file}:{line} | {message}" +) logger.add(sys.stdout, format=log_format) diff --git a/train_ms.py b/train_ms.py index 3c25e32..37933d1 100644 --- a/train_ms.py +++ b/train_ms.py @@ -206,8 +206,8 @@ def run(): ) else: optim_dur_disc = None - net_g = DDP(net_g, device_ids=[rank], find_unused_parameters=True) - net_d = DDP(net_d, device_ids=[rank], find_unused_parameters=True) + net_g = DDP(net_g, device_ids=[rank]) + net_d = DDP(net_d, device_ids=[rank]) dur_resume_lr = None if net_dur_disc is not None: net_dur_disc = DDP(net_dur_disc, device_ids=[rank], find_unused_parameters=True)