Fix: ONNX inference of BERT models fails in some Windows environments

This commit is contained in:
tsukumi
2024-09-23 11:33:17 +09:00
parent f9607ad891
commit 5ca6e184d7
3 changed files with 17 additions and 17 deletions

View File

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

View File

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

View File

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