Fix: spm.model is missing

This commit is contained in:
tsukumi
2024-12-19 12:05:33 +09:00
parent 2833fb5eeb
commit ebd249bca8
3 changed files with 19 additions and 8 deletions

View File

@@ -21,6 +21,6 @@
}, },
"deberta-v3-large-onnx": { "deberta-v3-large-onnx": {
"repo_id": "tsukumijima/deberta-v3-large-onnx", "repo_id": "tsukumijima/deberta-v3-large-onnx",
"files": ["model.onnx"] "files": ["spm.model", "model.onnx"]
} }
} }

View File

@@ -228,22 +228,24 @@ if __name__ == "__main__":
# トークナイザーを Fast Tokenizer 用形式に変換して保存 # トークナイザーを Fast Tokenizer 用形式に変換して保存
if language == Languages.EN: if language == Languages.EN:
slow_tokenizer = DebertaV2Tokenizer.from_pretrained( tokenizer = DebertaV2Tokenizer.from_pretrained(
pretrained_model_name_or_path, pretrained_model_name_or_path,
) )
convert_slow_tokenizer(slow_tokenizer).save(str(tokenizer_json_path)) convert_slow_tokenizer(tokenizer).save(str(tokenizer_json_path))
elif language == Languages.JP: elif language == Languages.JP:
slow_tokenizer = AutoTokenizer.from_pretrained( tokenizer = AutoTokenizer.from_pretrained(
pretrained_model_name_or_path, pretrained_model_name_or_path,
use_fast=False, # 明示的に Slow Tokenizer を使う use_fast=False, # 明示的に Slow Tokenizer を使う
) )
BertConverter(slow_tokenizer).converted().save(str(tokenizer_json_path)) BertConverter(tokenizer).converted().save(str(tokenizer_json_path))
elif language == Languages.ZH: elif language == Languages.ZH:
slow_tokenizer = AutoTokenizer.from_pretrained( tokenizer = AutoTokenizer.from_pretrained(
pretrained_model_name_or_path, pretrained_model_name_or_path,
use_fast=False, # 明示的に Slow Tokenizer を使う use_fast=False, # 明示的に Slow Tokenizer を使う
) )
convert_slow_tokenizer(slow_tokenizer).save(str(tokenizer_json_path)) convert_slow_tokenizer(tokenizer).save(str(tokenizer_json_path))
else:
assert False, "Invalid language"
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

@@ -87,6 +87,15 @@ def load_model(
revision=revision, revision=revision,
) )
) )
# 英語用 BERT のみ、spm.model もダウンロードする
# Fast 版の BERT トークナイザーでは不要なはずだが、念のため
if language == Languages.EN:
hf_hub_download(
repo_id=pretrained_model_name_or_path,
filename="spm.model",
cache_dir=cache_dir,
revision=revision,
)
# pretrained_model_name_or_path にファイルパスが指定された場合: # pretrained_model_name_or_path にファイルパスが指定された場合:
# 既にダウンロード済みという前提のもと、モデルへのローカルパスを model_path に格納する # 既にダウンロード済みという前提のもと、モデルへのローカルパスを model_path に格納する
else: else:
@@ -138,7 +147,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 トークナイザー
""" """
# すでにロード済みの場合はそのまま返す # すでにロード済みの場合はそのまま返す