From 07d246b98b2f476eb7f6a0aa4856313d010cbb62 Mon Sep 17 00:00:00 2001 From: tsukumi Date: Tue, 12 Mar 2024 18:27:08 +0000 Subject: [PATCH] Refactor: run "hatch run style:fmt" --- pyproject.toml | 14 +++++++------- style_bert_vits2/nlp/bert_models.py | 20 ++++++++++++++++---- style_bert_vits2/nlp/chinese/g2p.py | 3 +-- 3 files changed, 24 insertions(+), 13 deletions(-) diff --git a/pyproject.toml b/pyproject.toml index e8c218d..a5f53ec 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -75,18 +75,18 @@ cov = [ [tool.hatch.envs.style] detached = true dependencies = [ - "black", - "isort", + "black", + "isort", ] [tool.hatch.envs.style.scripts] check = [ - "black --check --diff .", - "isort --check-only --diff --profile black --gitignore --lai 2 .", + "black --check --diff .", + "isort --check-only --diff --profile black --gitignore --lai 2 .", ] fmt = [ - "black .", - "isort --profile black --gitignore --lai 2 .", - "check", + "black .", + "isort --profile black --gitignore --lai 2 .", + "check", ] [[tool.hatch.envs.test.matrix]] diff --git a/style_bert_vits2/nlp/bert_models.py b/style_bert_vits2/nlp/bert_models.py index eb84eb7..1166846 100644 --- a/style_bert_vits2/nlp/bert_models.py +++ b/style_bert_vits2/nlp/bert_models.py @@ -80,10 +80,14 @@ def load_model( if language == Languages.EN: model = cast( DebertaV2Model, - DebertaV2Model.from_pretrained(pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision), + DebertaV2Model.from_pretrained( + pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision + ), ) else: - model = AutoModelForMaskedLM.from_pretrained(pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision) + model = AutoModelForMaskedLM.from_pretrained( + pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision + ) __loaded_models[language] = model logger.info( f"Loaded the {language} BERT model from {pretrained_model_name_or_path}" @@ -135,9 +139,17 @@ def load_tokenizer( # BERT トークナイザーをロードし、辞書に格納して返す ## 英語のみ DebertaV2Tokenizer でロードする必要がある if language == Languages.EN: - tokenizer = DebertaV2Tokenizer.from_pretrained(pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision) + tokenizer = DebertaV2Tokenizer.from_pretrained( + pretrained_model_name_or_path, + cache_dir=cache_dir, + revision=revision, + ) else: - tokenizer = AutoTokenizer.from_pretrained(pretrained_model_name_or_path, cache_dir=cache_dir, revision=revision) + tokenizer = AutoTokenizer.from_pretrained( + pretrained_model_name_or_path, + cache_dir=cache_dir, + revision=revision, + ) __loaded_tokenizers[language] = tokenizer logger.info( f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}" diff --git a/style_bert_vits2/nlp/chinese/g2p.py b/style_bert_vits2/nlp/chinese/g2p.py index 4ce2b26..f38e09f 100644 --- a/style_bert_vits2/nlp/chinese/g2p.py +++ b/style_bert_vits2/nlp/chinese/g2p.py @@ -10,8 +10,7 @@ from style_bert_vits2.nlp.symbols import PUNCTUATIONS with open(Path(__file__).parent / "opencpop-strict.txt", "r", encoding="utf-8") as f: __PINYIN_TO_SYMBOL_MAP = { - line.split("\t")[0]: line.strip().split("\t")[1] - for line in f.readlines() + line.split("\t")[0]: line.strip().split("\t")[1] for line in f.readlines() }