Improve: Make ONNX inference tests more rigorous

This commit is contained in:
tsukumi
2024-09-23 12:27:29 +09:00
parent 2cdb5844cb
commit c792bf9a0d
4 changed files with 42 additions and 25 deletions

View File

@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, cast, Optional, Sequence, Union
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union, cast
import numpy as np
from numpy.typing import NDArray
@@ -108,9 +108,9 @@ def extract_bert_feature_onnx(
res = session.run(
[output_name],
{
"input_ids": cast(NDArray[Any], inputs["input_ids"]).astype(np.int64),
"token_type_ids": cast(NDArray[Any], inputs["token_type_ids"]).astype(np.int64),
"attention_mask": cast(NDArray[Any], inputs["attention_mask"]).astype(np.int64),
"input_ids": inputs["input_ids"].astype(np.int64), # type: ignore
"token_type_ids": inputs["token_type_ids"].astype(np.int64), # type: ignore
"attention_mask": inputs["attention_mask"].astype(np.int64), # type: ignore
},
)[0]
@@ -120,9 +120,9 @@ def extract_bert_feature_onnx(
style_res = session.run(
[output_name],
{
"input_ids": cast(NDArray[Any], style_inputs["input_ids"]).astype(np.int64),
"token_type_ids": cast(NDArray[Any], style_inputs["token_type_ids"]).astype(np.int64),
"attention_mask": cast(NDArray[Any], style_inputs["attention_mask"]).astype(np.int64),
"input_ids": style_inputs["input_ids"].astype(np.int64), # type: ignore
"token_type_ids": style_inputs["token_type_ids"].astype(np.int64), # type: ignore
"attention_mask": style_inputs["attention_mask"].astype(np.int64), # type: ignore
},
)[0]
style_res_mean = np.mean(style_res, axis=0)

View File

@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, cast, Optional, Sequence, Union
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union, cast
import numpy as np
from numpy.typing import NDArray
@@ -108,8 +108,8 @@ def extract_bert_feature_onnx(
res = session.run(
[output_name],
{
"input_ids": cast(NDArray[Any], inputs["input_ids"]).astype(np.int64),
"attention_mask": cast(NDArray[Any], inputs["attention_mask"]).astype(np.int64),
"input_ids": inputs["input_ids"].astype(np.int64), # type: ignore
"attention_mask": inputs["attention_mask"].astype(np.int64), # type: ignore
},
)[0]
@@ -119,8 +119,8 @@ def extract_bert_feature_onnx(
style_res = session.run(
[output_name],
{
"input_ids": cast(NDArray[Any], style_inputs["input_ids"]).astype(np.int64),
"attention_mask": cast(NDArray[Any], style_inputs["attention_mask"]).astype(np.int64),
"input_ids": style_inputs["input_ids"].astype(np.int64), # type: ignore
"attention_mask": style_inputs["attention_mask"].astype(np.int64), # type: ignore
},
)[0]
style_res_mean = np.mean(style_res, axis=0)

View File

@@ -1,6 +1,6 @@
from __future__ import annotations
from typing import TYPE_CHECKING, Any, cast, Optional, Sequence, Union
from typing import TYPE_CHECKING, Any, Optional, Sequence, Union, cast
import numpy as np
from numpy.typing import NDArray
@@ -121,8 +121,8 @@ def extract_bert_feature_onnx(
res = session.run(
[output_name],
{
"input_ids": cast(NDArray[Any], inputs["input_ids"]).astype(np.int64),
"attention_mask": cast(NDArray[Any], inputs["attention_mask"]).astype(np.int64),
"input_ids": inputs["input_ids"].astype(np.int64), # type: ignore
"attention_mask": inputs["attention_mask"].astype(np.int64), # type: ignore
},
)[0]
@@ -132,8 +132,8 @@ def extract_bert_feature_onnx(
style_res = session.run(
[output_name],
{
"input_ids": cast(NDArray[Any], style_inputs["input_ids"]).astype(np.int64),
"attention_mask": cast(NDArray[Any], style_inputs["attention_mask"]).astype(np.int64),
"input_ids": style_inputs["input_ids"].astype(np.int64), # type: ignore
"attention_mask": style_inputs["attention_mask"].astype(np.int64), # type: ignore
},
)[0]
style_res_mean = np.mean(style_res, axis=0)