mirror of
https://github.com/storytold/storyteller-ml.git
synced 2026-10-09 00:09:55 +00:00
Added support for V2 models
This commit is contained in:
+14
-15
@@ -1,18 +1,10 @@
|
||||
|
||||
# Base CUDA image
|
||||
|
||||
FROM cnstark/pytorch:2.0.1-py3.9.17-cuda11.8.0-ubuntu20.04
|
||||
|
||||
# Set environment variables
|
||||
ENV DEBIAN_FRONTEND=noninteractive
|
||||
ENV TZ=Etc/UTC
|
||||
|
||||
# Create directories for Python installation and model code
|
||||
RUN mkdir -p /python_install /model_code/GPT-SoVITS/GPT_SoVITS/pretrained_models
|
||||
|
||||
# Set the working directory for Python installation
|
||||
WORKDIR /python_install
|
||||
|
||||
# Install necessary packages and git-lfs
|
||||
RUN apt-get update && \
|
||||
apt-get install -y --no-install-recommends \
|
||||
@@ -35,6 +27,12 @@ RUN apt-get update && \
|
||||
add-apt-repository ppa:deadsnakes/ppa && \
|
||||
rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# Create directories for Python installation and model code
|
||||
RUN mkdir -p /python_install /model_code/GPT-SoVITS/GPT_SoVITS/pretrained_models
|
||||
|
||||
# Set the working directory for Python installation
|
||||
WORKDIR /python_install
|
||||
|
||||
# Copy the requirements.txt and install Python dependencies
|
||||
COPY requirements.txt /python_install/
|
||||
RUN python3 -m venv python && \
|
||||
@@ -44,20 +42,21 @@ RUN python3 -m venv python && \
|
||||
|
||||
# Clone Hugging Face repository and move models
|
||||
RUN git clone https://huggingface.co/lj1995/GPT-SoVITS /python_install/GPT-SoVITS && \
|
||||
mkdir -p /model_code/GPT-SoVITS/GPT_SoVITS/pretrained_models && \
|
||||
mv /python_install/GPT-SoVITS/* /model_code/GPT-SoVITS/GPT_SoVITS/pretrained_models/
|
||||
|
||||
# Move Docker, GPT_SoVITS, requirements.txt, and tools to /model_code/GPT-SoVITS
|
||||
WORKDIR /model_code/GPT-SoVITS
|
||||
RUN mv /python_install/requirements.txt /model_code/GPT-SoVITS/
|
||||
|
||||
# Set the working directory for model code
|
||||
WORKDIR /model_code/GPT-SoVITS
|
||||
# Download, unzip, rename, and place G2PW models (Essential for Chinese TTS)
|
||||
RUN apt-get update && apt-get install -y unzip && \
|
||||
curl -L -o G2PWModel.zip https://paddlespeech.bj.bcebos.com/Parakeet/released_models/g2p/G2PWModel_1.1.zip && \
|
||||
unzip G2PWModel.zip -d /model_code/GPT-SoVITS/text && \
|
||||
mv /model_code/GPT-SoVITS/text/G2PWModel_1.1 /model_code/GPT-SoVITS/text/G2PWModel && \
|
||||
rm G2PWModel.zip
|
||||
|
||||
# Copy the remaining contents to the model_code directory
|
||||
COPY . /model_code/GPT-SoVITS
|
||||
|
||||
# Change working directory to the project folder
|
||||
WORKDIR /model_code/GPT-SoVITS
|
||||
|
||||
# Run an interactive shell by default and activate the virtual environment
|
||||
CMD ["bash", "-c", ". /python_install/python/bin/activate && exec bash"]
|
||||
CMD ["/bin/bash", "-c", ". /python_install/python/bin/activate && exec /bin/bash"]
|
||||
@@ -50,8 +50,27 @@ logging.getLogger("asyncio").setLevel(logging.ERROR)
|
||||
logging.getLogger("charset_normalizer").setLevel(logging.ERROR)
|
||||
logging.getLogger("torchaudio._extension").setLevel(logging.ERROR)
|
||||
|
||||
# Define language dictionary
|
||||
dict_language = {
|
||||
# Define version and pretrained models
|
||||
version = os.environ.get("version", "v2")
|
||||
pretrained_sovits_name = [
|
||||
"GPT_SoVITS/pretrained_models/gsv-v2final-pretrained/s2G2333k.pth",
|
||||
"GPT_SoVITS/pretrained_models/s2G488k.pth"
|
||||
]
|
||||
pretrained_gpt_name = [
|
||||
"GPT_SoVITS/pretrained_models/gsv-v2final-pretrained/s1bert25hz-5kh-longer-epoch=12-step=369668.ckpt",
|
||||
"GPT_SoVITS/pretrained_models/s1bert25hz-2kh-longer-epoch=68e-step=50232.ckpt"
|
||||
]
|
||||
|
||||
_ = [[], []]
|
||||
for i in range(2):
|
||||
if os.path.exists(pretrained_gpt_name[i]):
|
||||
_[0].append(pretrained_gpt_name[i])
|
||||
if os.path.exists(pretrained_sovits_name[i]):
|
||||
_[-1].append(pretrained_sovits_name[i])
|
||||
pretrained_gpt_name, pretrained_sovits_name = _
|
||||
|
||||
# Define language dictionaries for v1 and v2
|
||||
dict_language_v1 = {
|
||||
"chinese": "all_zh",
|
||||
"english": "en",
|
||||
"japanese": "all_ja",
|
||||
@@ -59,14 +78,23 @@ dict_language = {
|
||||
"japanese+english": "ja",
|
||||
"automatic": "auto",
|
||||
}
|
||||
dict_language_v2 = {
|
||||
"chinese": "all_zh",
|
||||
"english": "en",
|
||||
"japanese": "all_ja",
|
||||
"chinese+english": "zh",
|
||||
"japanese+english": "ja",
|
||||
"cantonese": "all_yue",
|
||||
"korean": "all_ko",
|
||||
"cantonese+english": "yue",
|
||||
"korean+english": "ko",
|
||||
"automatic": "auto",
|
||||
"automatic(cantonese)": "auto_yue",
|
||||
}
|
||||
|
||||
# Define punctuation set
|
||||
punctuation = set(['!', '?', '…', ',', '.', '-'," "])
|
||||
|
||||
# Define default paths for pretrained models
|
||||
pretrained_gpt_path = "GPT_SoVITS/pretrained_models/s1bert25hz-2kh-longer-epoch=68e-step=50232.ckpt"
|
||||
pretrained_sovits_path = "GPT_SoVITS/pretrained_models/s2G488k.pth"
|
||||
|
||||
# Paths for cnhubert and bert models
|
||||
cnhubert_base_path = os.environ.get("cnhubert_base_path", "GPT_SoVITS/pretrained_models/chinese-hubert-base")
|
||||
bert_path = os.environ.get("bert_path", "GPT_SoVITS/pretrained_models/chinese-roberta-wwm-ext-large")
|
||||
@@ -148,11 +176,16 @@ class DictToAttrRecursive(dict):
|
||||
|
||||
# Function to change SoVITS weights
|
||||
def change_sovits_weights(sovits_path):
|
||||
global vq_model, hps
|
||||
global vq_model, hps, version, dict_language
|
||||
dict_s2 = torch.load(sovits_path, map_location="cpu")
|
||||
hps = dict_s2["config"]
|
||||
hps = DictToAttrRecursive(hps)
|
||||
hps.model.semantic_frame_rate = "25hz"
|
||||
if dict_s2['weight']['enc_p.text_embedding.weight'].shape[0] == 322:
|
||||
hps.model.version = "v1"
|
||||
else:
|
||||
hps.model.version = "v2"
|
||||
version = hps.model.version
|
||||
vq_model = SynthesizerTrn(
|
||||
hps.data.filter_length // 2 + 1,
|
||||
hps.train.segment_size // hps.data.hop_length,
|
||||
@@ -169,6 +202,7 @@ def change_sovits_weights(sovits_path):
|
||||
print(vq_model.load_state_dict(dict_s2["weight"], strict=False))
|
||||
with open("./sweight.txt", "w", encoding="utf-8") as f:
|
||||
f.write(sovits_path)
|
||||
dict_language = dict_language_v1 if version == 'v1' else dict_language_v2
|
||||
|
||||
# Function to change GPT weights
|
||||
def change_gpt_weights(gpt_path):
|
||||
@@ -204,9 +238,9 @@ def get_spepc(hps, filename):
|
||||
return spec
|
||||
|
||||
# Function to clean text
|
||||
def clean_text_inf(text, language):
|
||||
phones, word2ph, norm_text = clean_text(text, language)
|
||||
phones = cleaned_text_to_sequence(phones)
|
||||
def clean_text_inf(text, language, version):
|
||||
phones, word2ph, norm_text = clean_text(text, language, version)
|
||||
phones = cleaned_text_to_sequence(phones, version)
|
||||
return phones, word2ph, norm_text
|
||||
|
||||
dtype = torch.float16 if is_half else torch.float32
|
||||
@@ -234,8 +268,8 @@ def get_first(text):
|
||||
return text
|
||||
|
||||
# Function to get phones and BERT embeddings
|
||||
def get_phones_and_bert(text, language):
|
||||
if language in {"en", "all_zh", "all_ja"}:
|
||||
def get_phones_and_bert(text, language, version):
|
||||
if language in {"en", "all_zh", "all_ja", "all_ko", "all_yue"}:
|
||||
language = language.replace("all_", "")
|
||||
if language == "en":
|
||||
LangSegment.setfilters(["en"])
|
||||
@@ -247,30 +281,35 @@ def get_phones_and_bert(text, language):
|
||||
if language == "zh":
|
||||
if re.search(r'[A-Za-z]', formattext):
|
||||
formattext = re.sub(r'[a-z]', lambda x: x.group(0).upper(), formattext)
|
||||
formattext = chinese.text_normalize(formattext)
|
||||
return get_phones_and_bert(formattext, "zh")
|
||||
formattext = chinese.mix_text_normalize(formattext)
|
||||
return get_phones_and_bert(formattext, "zh", version)
|
||||
else:
|
||||
phones, word2ph, norm_text = clean_text_inf(formattext, language)
|
||||
|
||||
bert = get_bert_feature(norm_text, word2ph).to(device)
|
||||
phones, word2ph, norm_text = clean_text_inf(formattext, language, version)
|
||||
bert = get_bert_feature(norm_text, word2ph).to(device)
|
||||
elif language == "yue" and re.search(r'[A-Za-z]', formattext):
|
||||
formattext = re.sub(r'[a-z]', lambda x: x.group(0).upper(), formattext)
|
||||
formattext = chinese.mix_text_normalize(formattext)
|
||||
return get_phones_and_bert(formattext, "yue", version)
|
||||
else:
|
||||
phones, word2ph, norm_text = clean_text_inf(formattext, language)
|
||||
phones, word2ph, norm_text = clean_text_inf(formattext, language, version)
|
||||
bert = torch.zeros(
|
||||
(1024, len(phones)),
|
||||
dtype=torch.float16 if is_half else torch.float32,
|
||||
).to(device)
|
||||
elif language in {"zh", "ja", "auto"}:
|
||||
elif language in {"zh", "ja", "ko", "yue", "auto", "auto_yue"}:
|
||||
textlist = []
|
||||
langlist = []
|
||||
LangSegment.setfilters(["zh", "ja", "en", "ko"])
|
||||
if language == "auto":
|
||||
for tmp in LangSegment.getTexts(text):
|
||||
if tmp["lang"] == "ko":
|
||||
langlist.append("zh")
|
||||
textlist.append(tmp["text"])
|
||||
else:
|
||||
langlist.append(tmp["lang"])
|
||||
textlist.append(tmp["text"])
|
||||
langlist.append(tmp["lang"])
|
||||
textlist.append(tmp["text"])
|
||||
elif language == "auto_yue":
|
||||
for tmp in LangSegment.getTexts(text):
|
||||
if tmp["lang"] == "zh":
|
||||
tmp["lang"] = "yue"
|
||||
langlist.append(tmp["lang"])
|
||||
textlist.append(tmp["text"])
|
||||
else:
|
||||
for tmp in LangSegment.getTexts(text):
|
||||
if tmp["lang"] == "en":
|
||||
@@ -285,7 +324,7 @@ def get_phones_and_bert(text, language):
|
||||
norm_text_list = []
|
||||
for i in range(len(textlist)):
|
||||
lang = langlist[i]
|
||||
phones, word2ph, norm_text = clean_text_inf(textlist[i], lang)
|
||||
phones, word2ph, norm_text = clean_text_inf(textlist[i], lang, version)
|
||||
bert = get_bert_inf(phones, word2ph, norm_text, lang)
|
||||
phones_list.append(phones)
|
||||
norm_text_list.append(norm_text)
|
||||
@@ -387,7 +426,7 @@ def get_tts_wav(ref_wav_path, prompt_text, prompt_language, text, text_language,
|
||||
audio_opt = []
|
||||
|
||||
if not ref_free:
|
||||
phones1, bert1, norm_text1 = get_phones_and_bert(prompt_text, prompt_language)
|
||||
phones1, bert1, norm_text1 = get_phones_and_bert(prompt_text, prompt_language, version)
|
||||
|
||||
t3_start = ttime()
|
||||
|
||||
@@ -397,7 +436,7 @@ def get_tts_wav(ref_wav_path, prompt_text, prompt_language, text, text_language,
|
||||
if text[-1] not in splits:
|
||||
text += "。" if text_language != "en" else "."
|
||||
print("Actual target text (per sentence):", text)
|
||||
phones2, bert2, norm_text2 = get_phones_and_bert(text, text_language)
|
||||
phones2, bert2, norm_text2 = get_phones_and_bert(text, text_language, version)
|
||||
print("Processed text (per sentence):", norm_text2)
|
||||
|
||||
if not ref_free:
|
||||
@@ -558,8 +597,8 @@ def gptsovits_inference(gpt_model_path, sovits_model_path, ref_wav_path, prompt_
|
||||
if not os.path.exists(ref_wav_path):
|
||||
print("You must input a reference audio path!")
|
||||
return None
|
||||
gpt_path = gpt_model_path if gpt_model_path else pretrained_gpt_path
|
||||
sovits_path = sovits_model_path if sovits_model_path else pretrained_sovits_path
|
||||
gpt_path = gpt_model_path if gpt_model_path else pretrained_gpt_name[0]
|
||||
sovits_path = sovits_model_path if sovits_model_path else pretrained_sovits_name[0]
|
||||
|
||||
if loaded_gpt_path != gpt_path:
|
||||
change_gpt_weights(gpt_path)
|
||||
@@ -598,8 +637,8 @@ def main():
|
||||
parser.add_argument("--sovits_model", type=str, help="Path to SoVITS model checkpoint")
|
||||
parser.add_argument("--ref_audio", type=str, help="Path to reference wav file")
|
||||
parser.add_argument("--ref_text", type=str, help="Path to reference text file (optional)")
|
||||
parser.add_argument("--target_language", type=str, choices=dict_language.keys(), help="Language of the target text")
|
||||
parser.add_argument("--ref_language", type=str, choices=dict_language.keys(), help="Language of the reference text")
|
||||
parser.add_argument("--target_language", type=str, choices=dict_language_v1.keys() | dict_language_v2.keys(), help="Language of the target text")
|
||||
parser.add_argument("--ref_language", type=str, choices=dict_language_v1.keys() | dict_language_v2.keys(), help="Language of the reference text")
|
||||
parser.add_argument("--target_text", type=str, help="Path to the target text file")
|
||||
parser.add_argument("--how_to_cut", type=str, default="No slice", help="How to cut the text for synthesis")
|
||||
parser.add_argument("--top_k", type=int, default=20, help="Top K sampling")
|
||||
@@ -610,8 +649,8 @@ def main():
|
||||
|
||||
args = parser.parse_args()
|
||||
|
||||
gpt_model_path = args.gpt_model if args.gpt_model else (pretrained_gpt_path if args.use_pretrained_gpt else None)
|
||||
sovits_model_path = args.sovits_model if args.sovits_model else (pretrained_sovits_path if args.use_pretrained_sovits else None)
|
||||
gpt_model_path = args.gpt_model if args.gpt_model else (pretrained_gpt_name[0] if args.use_pretrained_gpt else None)
|
||||
sovits_model_path = args.sovits_model if args.sovits_model else (pretrained_sovits_name[0] if args.use_pretrained_sovits else None)
|
||||
|
||||
if not gpt_model_path or not sovits_model_path:
|
||||
parser.error("--gpt_model and --sovits_model are required when not using pretrained models")
|
||||
|
||||
Reference in New Issue
Block a user