Refactor: run "hatch run style:fmt"

This commit is contained in:
tsukumi
2024-03-12 18:27:08 +00:00
parent e8a76e547b
commit 07d246b98b
3 changed files with 24 additions and 13 deletions

View File

@@ -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]]

View File

@@ -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}"

View File

@@ -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()
}