Refactor: run "hatch run style:fmt"
This commit is contained in:
@@ -75,18 +75,18 @@ cov = [
|
|||||||
[tool.hatch.envs.style]
|
[tool.hatch.envs.style]
|
||||||
detached = true
|
detached = true
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"black",
|
"black",
|
||||||
"isort",
|
"isort",
|
||||||
]
|
]
|
||||||
[tool.hatch.envs.style.scripts]
|
[tool.hatch.envs.style.scripts]
|
||||||
check = [
|
check = [
|
||||||
"black --check --diff .",
|
"black --check --diff .",
|
||||||
"isort --check-only --diff --profile black --gitignore --lai 2 .",
|
"isort --check-only --diff --profile black --gitignore --lai 2 .",
|
||||||
]
|
]
|
||||||
fmt = [
|
fmt = [
|
||||||
"black .",
|
"black .",
|
||||||
"isort --profile black --gitignore --lai 2 .",
|
"isort --profile black --gitignore --lai 2 .",
|
||||||
"check",
|
"check",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[tool.hatch.envs.test.matrix]]
|
[[tool.hatch.envs.test.matrix]]
|
||||||
|
|||||||
@@ -80,10 +80,14 @@ def load_model(
|
|||||||
if language == Languages.EN:
|
if language == Languages.EN:
|
||||||
model = cast(
|
model = cast(
|
||||||
DebertaV2Model,
|
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:
|
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
|
__loaded_models[language] = model
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Loaded the {language} BERT model from {pretrained_model_name_or_path}"
|
f"Loaded the {language} BERT model from {pretrained_model_name_or_path}"
|
||||||
@@ -135,9 +139,17 @@ def load_tokenizer(
|
|||||||
# BERT トークナイザーをロードし、辞書に格納して返す
|
# BERT トークナイザーをロードし、辞書に格納して返す
|
||||||
## 英語のみ DebertaV2Tokenizer でロードする必要がある
|
## 英語のみ DebertaV2Tokenizer でロードする必要がある
|
||||||
if language == Languages.EN:
|
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:
|
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
|
__loaded_tokenizers[language] = tokenizer
|
||||||
logger.info(
|
logger.info(
|
||||||
f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}"
|
f"Loaded the {language} BERT tokenizer from {pretrained_model_name_or_path}"
|
||||||
|
|||||||
@@ -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:
|
with open(Path(__file__).parent / "opencpop-strict.txt", "r", encoding="utf-8") as f:
|
||||||
__PINYIN_TO_SYMBOL_MAP = {
|
__PINYIN_TO_SYMBOL_MAP = {
|
||||||
line.split("\t")[0]: line.strip().split("\t")[1]
|
line.split("\t")[0]: line.strip().split("\t")[1] for line in f.readlines()
|
||||||
for line in f.readlines()
|
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user