Fix: maintain compatibility with Python 3.9
This commit is contained in:
@@ -43,8 +43,8 @@ from style_bert_vits2.constants import (
|
|||||||
)
|
)
|
||||||
from style_bert_vits2.logging import logger
|
from style_bert_vits2.logging import logger
|
||||||
from style_bert_vits2.nlp import bert_models
|
from style_bert_vits2.nlp import bert_models
|
||||||
from style_bert_vits2.nlp.japanese import normalize_text
|
|
||||||
from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone
|
from style_bert_vits2.nlp.japanese.g2p_utils import g2kata_tone, kata_tone2phone_tone
|
||||||
|
from style_bert_vits2.nlp.japanese.normalizer import normalize_text
|
||||||
from style_bert_vits2.nlp.japanese.user_dict import (
|
from style_bert_vits2.nlp.japanese.user_dict import (
|
||||||
apply_word,
|
apply_word,
|
||||||
delete_word,
|
delete_word,
|
||||||
|
|||||||
@@ -5,7 +5,7 @@
|
|||||||
|
|
||||||
import torch
|
import torch
|
||||||
from torch.nn import functional as F
|
from torch.nn import functional as F
|
||||||
from typing import Any, Optional
|
from typing import Any, Optional, Union
|
||||||
|
|
||||||
|
|
||||||
def init_weights(m: torch.nn.Module, mean: float = 0.0, std: float = 0.01) -> None:
|
def init_weights(m: torch.nn.Module, mean: float = 0.0, std: float = 0.01) -> None:
|
||||||
@@ -180,12 +180,12 @@ def generate_path(duration: torch.Tensor, mask: torch.Tensor) -> torch.Tensor:
|
|||||||
return path
|
return path
|
||||||
|
|
||||||
|
|
||||||
def clip_grad_value_(parameters: torch.Tensor | list[torch.Tensor], clip_value: Optional[float], norm_type: float = 2.0) -> float:
|
def clip_grad_value_(parameters: Union[torch.Tensor, list[torch.Tensor]], clip_value: Optional[float], norm_type: float = 2.0) -> float:
|
||||||
"""
|
"""
|
||||||
勾配の値をクリップする
|
勾配の値をクリップする
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
parameters (torch.Tensor | list[torch.Tensor]): クリップするパラメータ
|
parameters (Union[torch.Tensor, list[torch.Tensor]]): クリップするパラメータ
|
||||||
clip_value (Optional[float]): クリップする値。None の場合はクリップしない
|
clip_value (Optional[float]): クリップする値。None の場合はクリップしない
|
||||||
norm_type (float): ノルムの種類
|
norm_type (float): ノルムの種類
|
||||||
|
|
||||||
|
|||||||
@@ -9,7 +9,7 @@ Style-Bert-VITS2 の学習・推論に必要な各言語ごとの BERT モデル
|
|||||||
"""
|
"""
|
||||||
|
|
||||||
import gc
|
import gc
|
||||||
from typing import cast, Optional
|
from typing import cast, Optional, Union
|
||||||
|
|
||||||
import torch
|
import torch
|
||||||
from transformers import (
|
from transformers import (
|
||||||
@@ -27,16 +27,16 @@ from style_bert_vits2.logging import logger
|
|||||||
|
|
||||||
|
|
||||||
# 各言語ごとのロード済みの BERT モデルを格納する辞書
|
# 各言語ごとのロード済みの BERT モデルを格納する辞書
|
||||||
__loaded_models: dict[Languages, PreTrainedModel | DebertaV2Model] = {}
|
__loaded_models: dict[Languages, Union[PreTrainedModel, DebertaV2Model]] = {}
|
||||||
|
|
||||||
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
# 各言語ごとのロード済みの BERT トークナイザーを格納する辞書
|
||||||
__loaded_tokenizers: dict[Languages, PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer] = {}
|
__loaded_tokenizers: dict[Languages, Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]] = {}
|
||||||
|
|
||||||
|
|
||||||
def load_model(
|
def load_model(
|
||||||
language: Languages,
|
language: Languages,
|
||||||
pretrained_model_name_or_path: Optional[str] = None,
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
) -> PreTrainedModel | DebertaV2Model:
|
) -> Union[PreTrainedModel, DebertaV2Model]:
|
||||||
"""
|
"""
|
||||||
指定された言語の BERT モデルをロードし、ロード済みの BERT モデルを返す。
|
指定された言語の BERT モデルをロードし、ロード済みの BERT モデルを返す。
|
||||||
一度ロードされていれば、ロード済みの BERT モデルを即座に返す。
|
一度ロードされていれば、ロード済みの BERT モデルを即座に返す。
|
||||||
@@ -54,7 +54,7 @@ def load_model(
|
|||||||
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
PreTrainedModel | DebertaV2Model: ロード済みの BERT モデル
|
Union[PreTrainedModel, DebertaV2Model]: ロード済みの BERT モデル
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# すでにロード済みの場合はそのまま返す
|
# すでにロード済みの場合はそのまま返す
|
||||||
@@ -82,7 +82,7 @@ def load_model(
|
|||||||
def load_tokenizer(
|
def load_tokenizer(
|
||||||
language: Languages,
|
language: Languages,
|
||||||
pretrained_model_name_or_path: Optional[str] = None,
|
pretrained_model_name_or_path: Optional[str] = None,
|
||||||
) -> PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer:
|
) -> Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]:
|
||||||
"""
|
"""
|
||||||
指定された言語の BERT モデルをロードし、ロード済みの BERT トークナイザーを返す。
|
指定された言語の BERT モデルをロードし、ロード済みの BERT トークナイザーを返す。
|
||||||
一度ロードされていれば、ロード済みの BERT トークナイザーを即座に返す。
|
一度ロードされていれば、ロード済みの BERT トークナイザーを即座に返す。
|
||||||
@@ -100,7 +100,7 @@ def load_tokenizer(
|
|||||||
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
PreTrainedTokenizer | PreTrainedTokenizerFast | DebertaV2Tokenizer: ロード済みの BERT トークナイザー
|
Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]: ロード済みの BERT トークナイザー
|
||||||
"""
|
"""
|
||||||
|
|
||||||
# すでにロード済みの場合はそのまま返す
|
# すでにロード済みの場合はそのまま返す
|
||||||
|
|||||||
@@ -8,13 +8,13 @@ from style_bert_vits2.constants import Languages
|
|||||||
from style_bert_vits2.nlp import bert_models
|
from style_bert_vits2.nlp import bert_models
|
||||||
|
|
||||||
|
|
||||||
__models: dict[torch.device | str, PreTrainedModel] = {}
|
__models: dict[str, PreTrainedModel] = {}
|
||||||
|
|
||||||
|
|
||||||
def extract_bert_feature(
|
def extract_bert_feature(
|
||||||
text: str,
|
text: str,
|
||||||
word2ph: list[int],
|
word2ph: list[int],
|
||||||
device: torch.device | str,
|
device: str,
|
||||||
assist_text: Optional[str] = None,
|
assist_text: Optional[str] = None,
|
||||||
assist_text_weight: float = 0.7,
|
assist_text_weight: float = 0.7,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
@@ -24,7 +24,7 @@ def extract_bert_feature(
|
|||||||
Args:
|
Args:
|
||||||
text (str): 中国語のテキスト
|
text (str): 中国語のテキスト
|
||||||
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
device (torch.device | str): 推論に利用するデバイス
|
device (str): 推論に利用するデバイス
|
||||||
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
|
|||||||
@@ -1,2 +0,0 @@
|
|||||||
from style_bert_vits2.nlp.japanese.g2p import g2p # noqa: F401
|
|
||||||
from style_bert_vits2.nlp.japanese.normalizer import normalize_text # noqa: F401
|
|
||||||
|
|||||||
@@ -9,13 +9,13 @@ from style_bert_vits2.nlp import bert_models
|
|||||||
from style_bert_vits2.nlp.japanese.g2p import text_to_sep_kata
|
from style_bert_vits2.nlp.japanese.g2p import text_to_sep_kata
|
||||||
|
|
||||||
|
|
||||||
__models: dict[torch.device | str, PreTrainedModel] = {}
|
__models: dict[str, PreTrainedModel] = {}
|
||||||
|
|
||||||
|
|
||||||
def extract_bert_feature(
|
def extract_bert_feature(
|
||||||
text: str,
|
text: str,
|
||||||
word2ph: list[int],
|
word2ph: list[int],
|
||||||
device: torch.device | str,
|
device: str,
|
||||||
assist_text: Optional[str] = None,
|
assist_text: Optional[str] = None,
|
||||||
assist_text_weight: float = 0.7,
|
assist_text_weight: float = 0.7,
|
||||||
) -> torch.Tensor:
|
) -> torch.Tensor:
|
||||||
@@ -25,7 +25,7 @@ def extract_bert_feature(
|
|||||||
Args:
|
Args:
|
||||||
text (str): 日本語のテキスト
|
text (str): 日本語のテキスト
|
||||||
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
|
||||||
device (torch.device | str): 推論に利用するデバイス
|
device (str): 推論に利用するデバイス
|
||||||
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
|
||||||
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user