Fix: ONNX inference of BERT models fails in some Windows environments
This commit is contained in:
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
@@ -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)
|
||||||
|
|||||||
Reference in New Issue
Block a user