diff --git a/bert/bert_models.json b/bert/bert_models.json index b1a2024..9b288b8 100644 --- a/bert/bert_models.json +++ b/bert/bert_models.json @@ -21,6 +21,6 @@ }, "deberta-v3-large-onnx": { "repo_id": "tsukumijima/deberta-v3-large-onnx", - "files": ["model.onnx"] + "files": ["spm.model", "model.onnx"] } } diff --git a/convert_bert_onnx.py b/convert_bert_onnx.py index f863f0a..e4f43a5 100644 --- a/convert_bert_onnx.py +++ b/convert_bert_onnx.py @@ -228,22 +228,24 @@ if __name__ == "__main__": # トークナイザーを Fast Tokenizer 用形式に変換して保存 if language == Languages.EN: - slow_tokenizer = DebertaV2Tokenizer.from_pretrained( + tokenizer = DebertaV2Tokenizer.from_pretrained( 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: - slow_tokenizer = AutoTokenizer.from_pretrained( + tokenizer = AutoTokenizer.from_pretrained( pretrained_model_name_or_path, 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: - slow_tokenizer = AutoTokenizer.from_pretrained( + tokenizer = AutoTokenizer.from_pretrained( pretrained_model_name_or_path, 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(f"[bold green]Tokenizer JSON saved to {tokenizer_json_path}[/bold green]") print(Rule(characters="=", style=Style(color="blue"))) diff --git a/style_bert_vits2/nlp/onnx_bert_models.py b/style_bert_vits2/nlp/onnx_bert_models.py index d4b095c..0a521f3 100644 --- a/style_bert_vits2/nlp/onnx_bert_models.py +++ b/style_bert_vits2/nlp/onnx_bert_models.py @@ -87,6 +87,15 @@ def load_model( 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 にファイルパスが指定された場合: # 既にダウンロード済みという前提のもと、モデルへのローカルパスを model_path に格納する else: @@ -138,7 +147,7 @@ def load_tokenizer( revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None) Returns: - Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2Tokenizer]: ロード済みの BERT トークナイザー + Union[PreTrainedTokenizer, PreTrainedTokenizerFast, DebertaV2TokenizerFast]: ロード済みの BERT トークナイザー """ # すでにロード済みの場合はそのまま返す