From c5acc04e2df3316c284c3ab75575498934314560 Mon Sep 17 00:00:00 2001
From: 游雁 <zhifu.gzf@alibaba-inc.com>
Date: 星期四, 30 三月 2023 16:35:03 +0800
Subject: [PATCH] Merge branch 'dev_cmz2' of github.com:alibaba-damo-academy/FunASR into dev_cmz2 add
---
funasr/runtime/python/onnxruntime/funasr_onnx/punc_bin.py | 8 +++-----
1 files changed, 3 insertions(+), 5 deletions(-)
diff --git a/funasr/runtime/python/onnxruntime/funasr_onnx/punc_bin.py b/funasr/runtime/python/onnxruntime/funasr_onnx/punc_bin.py
index 034475c..3f649bc 100644
--- a/funasr/runtime/python/onnxruntime/funasr_onnx/punc_bin.py
+++ b/funasr/runtime/python/onnxruntime/funasr_onnx/punc_bin.py
@@ -76,9 +76,8 @@
try:
outputs = self.infer(data['text'], data['text_lengths'])
y = outputs[0]
- _, indices = y.view(-1, y.shape[-1]).topk(1, dim=1)
- punctuations = indices
- assert punctuations.size()[0] == len(mini_sentence)
+ punctuations = np.argmax(y,axis=-1)[0]
+ assert punctuations.size == len(mini_sentence)
except ONNXRuntimeError:
logging.warning("error")
@@ -102,8 +101,7 @@
mini_sentence = mini_sentence[0:sentenceEnd + 1]
punctuations = punctuations[0:sentenceEnd + 1]
- punctuations_np = punctuations.cpu().numpy()
- new_mini_sentence_punc += [int(x) for x in punctuations_np]
+ new_mini_sentence_punc += [int(x) for x in punctuations]
words_with_punc = []
for i in range(len(mini_sentence)):
if i > 0:
--
Gitblit v1.9.1