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)