Fix: The conversion script for English BERT to Fast Tokenizer was incorrect, so tokenization was not performed correctly

This commit is contained in:
tsukumi
2024-12-19 11:10:33 +09:00
parent c1fce3fec7
commit 2833fb5eeb
6 changed files with 1024100 additions and 256046 deletions

View File

@@ -6,9 +6,9 @@
"normalizer": { "normalizer": {
"type": "BertNormalizer", "type": "BertNormalizer",
"clean_text": true, "clean_text": true,
"handle_chinese_chars": false, "handle_chinese_chars": true,
"strip_accents": false, "strip_accents": null,
"lowercase": false "lowercase": true
}, },
"pre_tokenizer": { "pre_tokenizer": {
"type": "BertPreTokenizer" "type": "BertPreTokenizer"

View File

@@ -6,9 +6,9 @@
"normalizer": { "normalizer": {
"type": "BertNormalizer", "type": "BertNormalizer",
"clean_text": true, "clean_text": true,
"handle_chinese_chars": false, "handle_chinese_chars": true,
"strip_accents": false, "strip_accents": null,
"lowercase": false "lowercase": true
}, },
"pre_tokenizer": { "pre_tokenizer": {
"type": "BertPreTokenizer" "type": "BertPreTokenizer"

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

View File

@@ -15,8 +15,8 @@ from rich import print
from rich.rule import Rule from rich.rule import Rule
from rich.style import Style from rich.style import Style
from torch import nn from torch import nn
from transformers import PreTrainedTokenizerBase from transformers import AutoTokenizer, DebertaV2Tokenizer, PreTrainedTokenizerBase
from transformers.convert_slow_tokenizer import BertConverter from transformers.convert_slow_tokenizer import BertConverter, convert_slow_tokenizer
from style_bert_vits2.constants import DEFAULT_BERT_MODEL_PATHS, Languages from style_bert_vits2.constants import DEFAULT_BERT_MODEL_PATHS, Languages
from style_bert_vits2.nlp import bert_models from style_bert_vits2.nlp import bert_models
@@ -227,9 +227,23 @@ if __name__ == "__main__":
print(Rule(characters="=", style=Style(color="blue"))) print(Rule(characters="=", style=Style(color="blue")))
# トークナイザーを Fast Tokenizer 用形式に変換して保存 # トークナイザーを Fast Tokenizer 用形式に変換して保存
tokenizer = bert_models.load_tokenizer(language) if language == Languages.EN:
converter = BertConverter(tokenizer) slow_tokenizer = DebertaV2Tokenizer.from_pretrained(
converter.converted().save(str(tokenizer_json_path)) pretrained_model_name_or_path,
)
convert_slow_tokenizer(slow_tokenizer).save(str(tokenizer_json_path))
elif language == Languages.JP:
slow_tokenizer = AutoTokenizer.from_pretrained(
pretrained_model_name_or_path,
use_fast=False, # 明示的に Slow Tokenizer を使う
)
BertConverter(slow_tokenizer).converted().save(str(tokenizer_json_path))
elif language == Languages.ZH:
slow_tokenizer = AutoTokenizer.from_pretrained(
pretrained_model_name_or_path,
use_fast=False, # 明示的に Slow Tokenizer を使う
)
convert_slow_tokenizer(slow_tokenizer).save(str(tokenizer_json_path))
print(Rule(characters="=", style=Style(color="blue"))) print(Rule(characters="=", style=Style(color="blue")))
print(f"[bold green]Tokenizer JSON saved to {tokenizer_json_path}[/bold green]") print(f"[bold green]Tokenizer JSON saved to {tokenizer_json_path}[/bold green]")
print(Rule(characters="=", style=Style(color="blue"))) print(Rule(characters="=", style=Style(color="blue")))

View File

@@ -142,7 +142,7 @@ def load_tokenizer(
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None) revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
Returns: Returns:
Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]: ロード済みの BERT トークナイザー Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2TokenizerFast]: ロード済みの BERT トークナイザー
""" """
# すでにロード済みの場合はそのまま返す # すでにロード済みの場合はそのまま返す
@@ -251,8 +251,6 @@ def unload_tokenizer(language: Languages) -> None:
language (Languages): アンロードする BERT トークナイザーの言語 language (Languages): アンロードする BERT トークナイザーの言語
""" """
import torch
if language in __loaded_tokenizers: if language in __loaded_tokenizers:
del __loaded_tokenizers[language] del __loaded_tokenizers[language]
gc.collect() gc.collect()