| | |
| | | |
| | | try{ |
| | | m_session = std::make_unique<Ort::Session>(env_, punc_model.c_str(), session_options); |
| | | LOG(INFO) << "Successfully load model from " << punc_model; |
| | | } |
| | | catch (std::exception const &e) { |
| | | LOG(ERROR) << "Error when load punc onnx model: " << e.what(); |
| | | exit(0); |
| | | exit(-1); |
| | | } |
| | | // read inputnames outputnames |
| | | string strName; |
| | |
| | | { |
| | | } |
| | | |
| | | string CTTransformer::AddPunc(const char* sz_input) |
| | | string CTTransformer::AddPunc(const char* sz_input, std::string language) |
| | | { |
| | | string strResult; |
| | | vector<string> strOut; |
| | |
| | | vector<string> WordWithPunc; |
| | | for (int i = 0; i < InputStr.size(); i++) |
| | | { |
| | | if (i > 0 && !(InputStr[i][0] & 0x80) && (i + 1) <InputStr.size() && !(InputStr[i+1][0] & 0x80))// �м��Ӣ�ģ� |
| | | // if (i > 0 && !(InputStr[i][0] & 0x80) && (i + 1) <InputStr.size() && !(InputStr[i+1][0] & 0x80))// �м��Ӣ�ģ� |
| | | if (i > 0 && !(InputStr[i-1][0] & 0x80) && !(InputStr[i][0] & 0x80)) |
| | | { |
| | | InputStr[i] = InputStr[i]+ " "; |
| | | InputStr[i] = " " + InputStr[i]; |
| | | } |
| | | WordWithPunc.push_back(InputStr[i]); |
| | | |
| | |
| | | NewPuncOut.assign(NewPunctuation.begin(), NewPunctuation.end() - 1); |
| | | NewPuncOut.push_back(PERIOD_INDEX); |
| | | } |
| | | else if (NewString[NewString.size() - 1] == m_tokenizer.Id2Punc(PERIOD_INDEX) && NewString[NewString.size() - 1] == m_tokenizer.Id2Punc(QUESTION_INDEX)) |
| | | else if (NewString[NewString.size() - 1] != m_tokenizer.Id2Punc(PERIOD_INDEX) && NewString[NewString.size() - 1] != m_tokenizer.Id2Punc(QUESTION_INDEX)) |
| | | { |
| | | NewSentenceOut = NewString; |
| | | NewSentenceOut.push_back(m_tokenizer.Id2Punc(PERIOD_INDEX)); |
| | |
| | | } |
| | | } |
| | | } |
| | | for (auto& item : NewSentenceOut) |
| | | |
| | | for (auto& item : NewSentenceOut){ |
| | | strResult += item; |
| | | } |
| | | |
| | | if(language == "en-bpe"){ |
| | | std::vector<std::string> chineseSymbols; |
| | | chineseSymbols.push_back(","); |
| | | chineseSymbols.push_back("。"); |
| | | chineseSymbols.push_back("、"); |
| | | chineseSymbols.push_back("?"); |
| | | |
| | | std::string englishSymbols = ",.,?"; |
| | | for (size_t i = 0; i < chineseSymbols.size(); i++) { |
| | | size_t pos = 0; |
| | | while ((pos = strResult.find(chineseSymbols[i], pos)) != std::string::npos) { |
| | | strResult.replace(pos, 3, 1, englishSymbols[i]); |
| | | pos++; |
| | | } |
| | | } |
| | | } |
| | | |
| | | return strResult; |
| | | } |
| | | |
| | |
| | | catch (std::exception const &e) |
| | | { |
| | | LOG(ERROR) << "Error when run punc onnx forword: " << (e.what()); |
| | | exit(0); |
| | | } |
| | | return punction; |
| | | } |
| | | |
| | | } // namespace funasr |
| | | } // namespace funasr |