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 typing import TYPE_CHECKING, Any, Optional, Sequence, Union
|
||||
from typing import TYPE_CHECKING, Any, cast, Optional, Sequence, Union
|
||||
|
||||
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": inputs["input_ids"],
|
||||
"attention_mask": inputs["attention_mask"],
|
||||
"input_ids": cast(NDArray[Any], inputs["input_ids"]).astype(np.int64),
|
||||
"attention_mask": cast(NDArray[Any], inputs["attention_mask"]).astype(np.int64),
|
||||
},
|
||||
)[0]
|
||||
|
||||
@@ -119,8 +119,8 @@ def extract_bert_feature_onnx(
|
||||
style_res = session.run(
|
||||
[output_name],
|
||||
{
|
||||
"input_ids": style_inputs["input_ids"],
|
||||
"attention_mask": style_inputs["attention_mask"],
|
||||
"input_ids": cast(NDArray[Any], style_inputs["input_ids"]).astype(np.int64),
|
||||
"attention_mask": cast(NDArray[Any], style_inputs["attention_mask"]).astype(np.int64),
|
||||
},
|
||||
)[0]
|
||||
style_res_mean = np.mean(style_res, axis=0)
|
||||
|
||||
Reference in New Issue
Block a user