From 5ca6e184d757cec3a24b5141fc410f6bbd129d9a Mon Sep 17 00:00:00 2001 From: tsukumi Date: Mon, 23 Sep 2024 11:33:17 +0900 Subject: [PATCH] Fix: ONNX inference of BERT models fails in some Windows environments --- style_bert_vits2/nlp/chinese/bert_feature.py | 14 +++++++------- style_bert_vits2/nlp/english/bert_feature.py | 10 +++++----- style_bert_vits2/nlp/japanese/bert_feature.py | 10 +++++----- 3 files changed, 17 insertions(+), 17 deletions(-) diff --git a/style_bert_vits2/nlp/chinese/bert_feature.py b/style_bert_vits2/nlp/chinese/bert_feature.py index 40a705a..64cdf72 100644 --- a/style_bert_vits2/nlp/chinese/bert_feature.py +++ b/style_bert_vits2/nlp/chinese/bert_feature.py @@ -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,9 +108,9 @@ def extract_bert_feature_onnx( res = session.run( [output_name], { - "input_ids": inputs["input_ids"], - "token_type_ids": inputs["token_type_ids"], - "attention_mask": inputs["attention_mask"], + "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), }, )[0] @@ -120,9 +120,9 @@ def extract_bert_feature_onnx( style_res = session.run( [output_name], { - "input_ids": style_inputs["input_ids"], - "token_type_ids": style_inputs["token_type_ids"], - "attention_mask": style_inputs["attention_mask"], + "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), }, )[0] style_res_mean = np.mean(style_res, axis=0) diff --git a/style_bert_vits2/nlp/english/bert_feature.py b/style_bert_vits2/nlp/english/bert_feature.py index 6ff143d..ae55239 100644 --- a/style_bert_vits2/nlp/english/bert_feature.py +++ b/style_bert_vits2/nlp/english/bert_feature.py @@ -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) diff --git a/style_bert_vits2/nlp/japanese/bert_feature.py b/style_bert_vits2/nlp/japanese/bert_feature.py index 9a91249..8f91906 100644 --- a/style_bert_vits2/nlp/japanese/bert_feature.py +++ b/style_bert_vits2/nlp/japanese/bert_feature.py @@ -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 @@ -121,8 +121,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] @@ -132,8 +132,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)