From 08f7eaff95aca7b19190c622d6af699ef32d180e Mon Sep 17 00:00:00 2001
From: zhifu gao <zhifu.gzf@alibaba-inc.com>
Date: 星期六, 13 一月 2024 13:52:36 +0800
Subject: [PATCH] Dev cmz (#1244)
---
funasr/datasets/preprocessor.py | 96 ++++++++++++++++++++++++++++-------------------
1 files changed, 57 insertions(+), 39 deletions(-)
diff --git a/funasr/datasets/preprocessor.py b/funasr/datasets/preprocessor.py
index b303418..966cc94 100644
--- a/funasr/datasets/preprocessor.py
+++ b/funasr/datasets/preprocessor.py
@@ -705,55 +705,73 @@
return line
@classmethod
- def split_words_jieba(cls, text: str):
- input_list = text.split()
- token_list_all = []
- langauge_list = []
- token_list_tmp = []
- language_flag = None
- for token in input_list:
- if cls.isEnglish(token) and language_flag == 'Chinese':
+ def split_words(cls, text: str , seg_jieba: bool):
+ if seg_jieba == True:
+ input_list = text.split()
+ token_list_all = []
+ langauge_list = []
+ token_list_tmp = []
+ language_flag = None
+ for token in input_list:
+ if cls.isEnglish(token) and language_flag == 'Chinese':
+ token_list_all.append(token_list_tmp)
+ langauge_list.append('Chinese')
+ token_list_tmp = []
+ elif not cls.isEnglish(token) and language_flag == 'English':
+ token_list_all.append(token_list_tmp)
+ langauge_list.append('English')
+ token_list_tmp = []
+
+ token_list_tmp.append(token)
+
+ if cls.isEnglish(token):
+ language_flag = 'English'
+ else:
+ language_flag = 'Chinese'
+
+ if token_list_tmp:
token_list_all.append(token_list_tmp)
- langauge_list.append('Chinese')
- token_list_tmp = []
- elif not cls.isEnglish(token) and language_flag == 'English':
- token_list_all.append(token_list_tmp)
- langauge_list.append('English')
- token_list_tmp = []
+ langauge_list.append(language_flag)
- token_list_tmp.append(token)
+ result_list = []
+ for token_list_tmp, language_flag in zip(token_list_all, langauge_list):
+ if language_flag == 'English':
+ result_list.extend(token_list_tmp)
+ else:
+ seg_list = jieba.cut(cls.join_chinese_and_english(token_list_tmp), HMM=False)
+ result_list.extend(seg_list)
- if cls.isEnglish(token):
- language_flag = 'English'
- else:
- language_flag = 'Chinese'
+ return result_list
- if token_list_tmp:
- token_list_all.append(token_list_tmp)
- langauge_list.append(language_flag)
+ else:
+ words = []
+ segs = text.split()
+ for seg in segs:
+ # There is no space in seg.
+ current_word = ""
+ for c in seg:
+ if len(c.encode()) == 1:
+ # This is an ASCII char.
+ current_word += c
+ else:
+ # This is a Chinese char.
+ if len(current_word) > 0:
+ words.append(current_word)
+ current_word = ""
+ words.append(c)
+ if len(current_word) > 0:
+ words.append(current_word)
+ return words
- result_list = []
- for token_list_tmp, language_flag in zip(token_list_all, langauge_list):
- if language_flag == 'English':
- result_list.extend(token_list_tmp)
- else:
- seg_list = jieba.cut(cls.join_chinese_and_english(token_list_tmp), HMM=False)
- result_list.extend(seg_list)
-
- return result_list
def __call__(
self, uid: str, data: Dict[str, Union[list, str, np.ndarray]]
) -> Dict[str, Union[list, np.ndarray]]:
# Split words.
- if isinstance(data[self.text_name], str):
- if self.seg_jieba:
- # jieba.load_userdict(seg_dict_file)
- split_text = self.split_words_jieba(data[self.text_name])
- else:
- split_text = self.split_words(data[self.text_name])
- else:
- split_text = data[self.text_name]
+ data_in = data[self.text_name]
+ if isinstance(data[self.text_name], list):
+ data_in = " ".join(data[self.text_name])
+ split_text = self.split_words(data_in, self.seg_jieba)
data[self.text_name] = " ".join(split_text)
data = self._speech_process(data)
data = self._text_process(data)
--
Gitblit v1.9.1