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)