Improve: Support ONNX inference, add ONNX conversion script

This commit is contained in:
tsukumi
2024-09-17 16:55:03 +09:00
parent f42b9f0f89
commit 5e2c83c6c7
21 changed files with 44483 additions and 376 deletions

View File

@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Optional, Sequence
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union
from numpy.typing import NDArray
@@ -60,8 +60,7 @@ def extract_bert_feature_onnx(
text: str,
word2ph: list[int],
language: Languages,
onnx_providers: list[str],
onnx_provider_options: Optional[Sequence[dict[str, Any]]],
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
assist_text: Optional[str] = None,
assist_text_weight: float = 0.7,
) -> NDArray[Any]:
@@ -73,7 +72,6 @@ def extract_bert_feature_onnx(
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
language (Languages): テキストの言語
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
onnx_provider_options (Optional[dict[str, Any]]): ONNX 推論で利用する ExecutionProvider のオプション
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
@@ -90,7 +88,6 @@ def extract_bert_feature_onnx(
text,
word2ph,
onnx_providers,
onnx_provider_options,
assist_text,
assist_text_weight,
)

View File

@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, Optional, Sequence
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union
import numpy as np
from numpy.typing import NDArray
@@ -86,8 +86,7 @@ def extract_bert_feature(
def extract_bert_feature_onnx(
text: str,
word2ph: list[int],
onnx_providers: list[str],
onnx_provider_options: Optional[Sequence[dict[str, Any]]],
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]],
assist_text: Optional[str] = None,
assist_text_weight: float = 0.7,
) -> NDArray[Any]:
@@ -98,7 +97,6 @@ def extract_bert_feature_onnx(
text (str): 日本語のテキスト
word2ph (list[int]): 元のテキストの各文字に音素が何個割り当てられるかを表すリスト
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
onnx_provider_options (Optional[dict[str, Any]]): ONNX 推論で利用する ExecutionProvider のオプション
assist_text (Optional[str], optional): 補助テキスト (デフォルト: None)
assist_text_weight (float, optional): 補助テキストの重み (デフォルト: 0.7)
@@ -118,7 +116,6 @@ def extract_bert_feature_onnx(
session = onnx_bert_models.load_model(
language=Languages.JP,
onnx_providers=onnx_providers,
onnx_provider_options=onnx_provider_options,
)
output_name = session.get_outputs()[0].name
res = session.run(

View File

@@ -39,11 +39,10 @@ __loaded_tokenizers: dict[
def load_model(
language: Languages,
pretrained_model_name_or_path: Optional[str] = None,
onnx_providers: list[str] = ["CPUExecutionProvider"],
onnx_provider_options: Optional[Sequence[dict[str, Any]]] = None,
onnx_providers: Sequence[Union[str, tuple[str, dict[str, Any]]]] = ["CPUExecutionProvider"],
cache_dir: Optional[str] = None,
revision: str = "main",
) -> onnxruntime.InferenceSession:
) -> onnxruntime.InferenceSession: # fmt: skip
"""
指定された言語の ONNX 版 BERT モデルをロードし、ロード済みの ONNX 版 BERT モデルを返す。
一度ロードされていれば、ロード済みの ONNX 版 BERT モデルを即座に返す。
@@ -59,7 +58,6 @@ def load_model(
language (Languages): ロードする学習済みモデルの対象言語
pretrained_model_name_or_path (Optional[str]): ロードする学習済みモデルの名前またはパス。指定しない場合はデフォルトのパスが利用される (デフォルト: None)
onnx_providers (list[str]): ONNX 推論で利用する ExecutionProvider (CPUExecutionProvider, CUDAExecutionProvider など)
onnx_provider_options (Optional[dict[str, Any]]): ONNX 推論で利用する ExecutionProvider のオプション
cache_dir (Optional[str]): モデルのキャッシュディレクトリ。指定しない場合はデフォルトのキャッシュディレクトリが利用される (デフォルト: None)
revision (str): モデルの Hugging Face 上の Git リビジョン。指定しない場合は最新の main ブランチの内容が利用される (デフォルト: None)
@@ -98,7 +96,6 @@ def load_model(
__loaded_models[language] = onnxruntime.InferenceSession(
model_path,
providers=onnx_providers,
provider_options=onnx_provider_options,
)
logger.info(
f"Loaded the {language} ONNX BERT model from {pretrained_model_name_or_path}"