diff --git a/default_style.py b/default_style.py index 763e291..67b6fc3 100644 --- a/default_style.py +++ b/default_style.py @@ -1,6 +1,6 @@ import os -from style_bert_vits2.logging import logger from style_bert_vits2.constants import DEFAULT_STYLE +from style_bert_vits2.logging import logger import numpy as np import json diff --git a/gen_yaml.py b/gen_yaml.py index 91301ac..76df206 100644 --- a/gen_yaml.py +++ b/gen_yaml.py @@ -1,7 +1,8 @@ +import argparse import os import shutil import yaml -import argparse + parser = argparse.ArgumentParser( description="config.ymlの生成。あらかじめ前準備をしたデータをバッチファイルなどで連続で学習する時にtrain_ms.pyより前に使用する。" diff --git a/style_bert_vits2/models/monotonic_alignment.py b/style_bert_vits2/models/monotonic_alignment.py index 0f393c1..b499ad0 100644 --- a/style_bert_vits2/models/monotonic_alignment.py +++ b/style_bert_vits2/models/monotonic_alignment.py @@ -28,7 +28,7 @@ def maximum_path(neg_cent: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: t_t_max = mask.sum(1)[:, 0].data.cpu().numpy().astype(int32) t_s_max = mask.sum(2)[:, 0].data.cpu().numpy().astype(int32) - maximum_path_jit(path, neg_cent, t_t_max, t_s_max) + __maximum_path_jit(path, neg_cent, t_t_max, t_s_max) return torch.from_numpy(path).to(device=device, dtype=dtype) @@ -43,7 +43,7 @@ def maximum_path(neg_cent: torch.Tensor, mask: torch.Tensor) -> torch.Tensor: nopython = True, nogil = True, ) # type: ignore -def maximum_path_jit(paths: Any, values: Any, t_ys: Any, t_xs: Any) -> None: +def __maximum_path_jit(paths: Any, values: Any, t_ys: Any, t_xs: Any) -> None: """ 与えられたパス、値、およびターゲットの y と x 座標を使用して JIT で最大パスを計算する diff --git a/style_bert_vits2/text_processing/__init__.py b/style_bert_vits2/text_processing/__init__.py index 5e77df6..4719525 100644 --- a/style_bert_vits2/text_processing/__init__.py +++ b/style_bert_vits2/text_processing/__init__.py @@ -8,7 +8,7 @@ from style_bert_vits2.text_processing.symbols import ( ) -_symbol_to_id = {s: i for i, s in enumerate(SYMBOLS)} +__symbol_to_id = {s: i for i, s in enumerate(SYMBOLS)} def extract_bert_feature( @@ -97,7 +97,7 @@ def cleaned_text_to_sequence(cleaned_phones: list[str], tones: list[int], langua tuple[list[int], list[int], list[int]]: List of integers corresponding to the symbols in the text """ - phones = [_symbol_to_id[symbol] for symbol in cleaned_phones] + phones = [__symbol_to_id[symbol] for symbol in cleaned_phones] tone_start = LANGUAGE_TONE_START_MAP[language] tones = [i + tone_start for i in tones] lang_id = LANGUAGE_ID_MAP[language] diff --git a/style_bert_vits2/text_processing/bert_models.py b/style_bert_vits2/text_processing/bert_models.py index 9132082..e8ef4b4 100644 --- a/style_bert_vits2/text_processing/bert_models.py +++ b/style_bert_vits2/text_processing/bert_models.py @@ -5,11 +5,13 @@ Style-Bert-VITS2 の学習・推論に必要な各言語ごとの BERT モデル 場合によっては多重にロードされて非効率なほか、BERT モデルのロード元のパスがハードコードされているためライブラリ化ができない。 そこで、ライブラリの利用前に、音声合成に利用する言語の BERT モデルだけを「明示的に」ロードできるようにした。 -一度 load_tokenizer() で当該言語の BERT モデルがロードされていれば、ライブラリ内部のどこからでもロード済みのモデル/トークナイザーを取得できる。 +一度 load_model/tokenizer() で当該言語の BERT モデルがロードされていれば、ライブラリ内部のどこからでもロード済みのモデル/トークナイザーを取得できる。 """ +import gc from typing import cast +import torch from transformers import ( AutoModelForMaskedLM, AutoTokenizer, @@ -25,10 +27,10 @@ from style_bert_vits2.logging import logger # 各言語ごとのロード済みの BERT モデルを格納する辞書 -loaded_models: dict[Languages, PreTrainedModel | DebertaV2Model] = {} +__loaded_models: dict[Languages, PreTrainedModel | DebertaV2Model] = {} # 各言語ごとのロード済みの BERT トークナイザーを格納する辞書 -loaded_tokenizers: dict[Languages, PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer] = {} +__loaded_tokenizers: dict[Languages, PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer] = {} def load_model( @@ -56,8 +58,8 @@ def load_model( """ # すでにロード済みの場合はそのまま返す - if language in loaded_models: - return loaded_models[language] + if language in __loaded_models: + return __loaded_models[language] # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 if pretrained_model_name_or_path is None: @@ -71,7 +73,7 @@ def load_model( model = cast(DebertaV2Model, DebertaV2Model.from_pretrained(pretrained_model_name_or_path)) else: model = AutoModelForMaskedLM.from_pretrained(pretrained_model_name_or_path) - loaded_models[language] = model + __loaded_models[language] = model logger.info(f"Loaded the {language} BERT model from {pretrained_model_name_or_path}") return model @@ -102,8 +104,8 @@ def load_tokenizer( """ # すでにロード済みの場合はそのまま返す - if language in loaded_tokenizers: - return loaded_tokenizers[language] + if language in __loaded_tokenizers: + return __loaded_tokenizers[language] # pretrained_model_name_or_path が指定されていない場合はデフォルトのパスを利用 if pretrained_model_name_or_path is None: @@ -117,7 +119,59 @@ def load_tokenizer( tokenizer = DebertaV2Tokenizer.from_pretrained(pretrained_model_name_or_path) else: tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path) - loaded_tokenizers[language] = tokenizer + __loaded_tokenizers[language] = tokenizer logger.info(f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}") return tokenizer + + +def unload_model(language: Languages) -> None: + """ + 指定された言語の BERT モデルをアンロードする + + Args: + language (Languages): アンロードする BERT モデルの言語 + """ + + if language in __loaded_models: + del __loaded_models[language] + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + logger.info(f"Unloaded the {language} BERT model") + + +def unload_tokenizer(language: Languages) -> None: + """ + 指定された言語の BERT トークナイザーをアンロードする + + Args: + language (Languages): アンロードする BERT トークナイザーの言語 + """ + + if language in __loaded_tokenizers: + del __loaded_tokenizers[language] + gc.collect() + if torch.cuda.is_available(): + torch.cuda.empty_cache() + logger.info(f"Unloaded the {language} BERT tokenizer") + + +def unload_all_models() -> None: + """ + すべての BERT モデルをアンロードする + """ + + for language in list(__loaded_models.keys()): + unload_model(language) + logger.info("Unloaded all BERT models") + + +def unload_all_tokenizers() -> None: + """ + すべての BERT トークナイザーをアンロードする + """ + + for language in list(__loaded_tokenizers.keys()): + unload_tokenizer(language) + logger.info("Unloaded all BERT tokenizers") diff --git a/style_bert_vits2/text_processing/chinese/bert_feature.py b/style_bert_vits2/text_processing/chinese/bert_feature.py index 25024cb..3178565 100644 --- a/style_bert_vits2/text_processing/chinese/bert_feature.py +++ b/style_bert_vits2/text_processing/chinese/bert_feature.py @@ -7,7 +7,7 @@ from style_bert_vits2.constants import Languages from style_bert_vits2.text_processing import bert_models -models: dict[torch.device | str, PreTrainedModel] = {} +__models: dict[torch.device | str, PreTrainedModel] = {} def extract_bert_feature( @@ -41,8 +41,8 @@ def extract_bert_feature( device = "cuda" if device == "cuda" and not torch.cuda.is_available(): device = "cpu" - if device not in models.keys(): - models[device] = bert_models.load_model(Languages.ZH).to(device) # type: ignore + if device not in __models.keys(): + __models[device] = bert_models.load_model(Languages.ZH).to(device) # type: ignore style_res_mean = None with torch.no_grad(): @@ -50,13 +50,13 @@ def extract_bert_feature( inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) # type: ignore - res = models[device](**inputs, output_hidden_states=True) + res = __models[device](**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: style_inputs[i] = style_inputs[i].to(device) # type: ignore - style_res = models[device](**style_inputs, output_hidden_states=True) + style_res = __models[device](**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0) diff --git a/style_bert_vits2/text_processing/english/bert_feature.py b/style_bert_vits2/text_processing/english/bert_feature.py index ec556c2..b29d531 100644 --- a/style_bert_vits2/text_processing/english/bert_feature.py +++ b/style_bert_vits2/text_processing/english/bert_feature.py @@ -7,7 +7,7 @@ from style_bert_vits2.constants import Languages from style_bert_vits2.text_processing import bert_models -models: dict[torch.device | str, PreTrainedModel] = {} +__models: dict[torch.device | str, PreTrainedModel] = {} def extract_bert_feature( @@ -41,8 +41,8 @@ def extract_bert_feature( device = "cuda" if device == "cuda" and not torch.cuda.is_available(): device = "cpu" - if device not in models.keys(): - models[device] = bert_models.load_model(Languages.EN).to(device) # type: ignore + if device not in __models.keys(): + __models[device] = bert_models.load_model(Languages.EN).to(device) # type: ignore style_res_mean = None with torch.no_grad(): @@ -50,13 +50,13 @@ def extract_bert_feature( inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) # type: ignore - res = models[device](**inputs, output_hidden_states=True) + res = __models[device](**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: style_inputs[i] = style_inputs[i].to(device) # type: ignore - style_res = models[device](**style_inputs, output_hidden_states=True) + style_res = __models[device](**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0) diff --git a/style_bert_vits2/text_processing/japanese/bert_feature.py b/style_bert_vits2/text_processing/japanese/bert_feature.py index 3ff9d7b..d1809fe 100644 --- a/style_bert_vits2/text_processing/japanese/bert_feature.py +++ b/style_bert_vits2/text_processing/japanese/bert_feature.py @@ -8,7 +8,7 @@ from style_bert_vits2.text_processing import bert_models from style_bert_vits2.text_processing.japanese.g2p import text_to_sep_kata -models: dict[torch.device | str, PreTrainedModel] = {} +__models: dict[torch.device | str, PreTrainedModel] = {} def extract_bert_feature( @@ -48,8 +48,8 @@ def extract_bert_feature( device = "cuda" if device == "cuda" and not torch.cuda.is_available(): device = "cpu" - if device not in models.keys(): - models[device] = bert_models.load_model(Languages.JP).to(device) # type: ignore + if device not in __models.keys(): + __models[device] = bert_models.load_model(Languages.JP).to(device) # type: ignore style_res_mean = None with torch.no_grad(): @@ -57,13 +57,13 @@ def extract_bert_feature( inputs = tokenizer(text, return_tensors="pt") for i in inputs: inputs[i] = inputs[i].to(device) # type: ignore - res = models[device](**inputs, output_hidden_states=True) + res = __models[device](**inputs, output_hidden_states=True) res = torch.cat(res["hidden_states"][-3:-2], -1)[0].cpu() if assist_text: style_inputs = tokenizer(assist_text, return_tensors="pt") for i in style_inputs: style_inputs[i] = style_inputs[i].to(device) # type: ignore - style_res = models[device](**style_inputs, output_hidden_states=True) + style_res = __models[device](**style_inputs, output_hidden_states=True) style_res = torch.cat(style_res["hidden_states"][-3:-2], -1)[0].cpu() style_res_mean = style_res.mean(0)