From 95ac7a2f3597b560638667bbdd6030d929bba7a4 Mon Sep 17 00:00:00 2001 From: Dmitri Smirnov Date: Wed, 24 Apr 2019 10:23:45 -0700 Subject: [PATCH] Implement separators as regex (#857) * Implement separators as regex --- onnxruntime/contrib_ops/cpu/tokenizer.cc | 382 ++++++------------ onnxruntime/core/common/utf8_util.h | 17 + .../core/graph/contrib_ops/contrib_defs.cc | 13 +- .../test/contrib_ops/tokenizer_test.cc | 58 +-- 4 files changed, 164 insertions(+), 306 deletions(-) diff --git a/onnxruntime/contrib_ops/cpu/tokenizer.cc b/onnxruntime/contrib_ops/cpu/tokenizer.cc index c65e08344e..56e5c9d3d8 100644 --- a/onnxruntime/contrib_ops/cpu/tokenizer.cc +++ b/onnxruntime/contrib_ops/cpu/tokenizer.cc @@ -10,9 +10,6 @@ #include "core/common/utf8_util.h" #include "re2/re2.h" -#include -#include - namespace onnxruntime { namespace contrib { @@ -29,19 +26,18 @@ class Tokenizer final : public OpKernel { Status CharTokenize(OpKernelContext* context, size_t N, size_t C, const std::vector& input_dims) const; - Status SeparatorTokenize(OpKernelContext* context, size_t N, size_t C, - const std::vector& input_dims) const; + Status SeparatorExpressionTokenizer(OpKernelContext* context, size_t N, size_t C, + const std::vector& input_dims) const; - Status ExpressionTokenize(OpKernelContext* ctx, - size_t N, size_t C, - const std::vector& input_dims) const; + Status TokenExpression(OpKernelContext* ctx, + size_t N, size_t C, + const std::vector& input_dims) const; - bool mark_; + bool mark_{false}; std::string pad_value_; - int64_t mincharnum_; - bool char_tokenezation_; - struct SearchData; - std::unique_ptr search_data_; + int64_t mincharnum_{0}; + bool char_tokenezation_{false}; + std::vector> separators_; std::unique_ptr regex_; }; @@ -56,177 +52,12 @@ ONNX_CPU_OPERATOR_TYPED_MS_KERNEL( contrib::Tokenizer); namespace tokenizer_details { - const char start_text = 0x2; const char end_text = 0x3; - -const std::string conv_error("Conversion Error"); -const std::wstring wconv_error(L"Conversion Error"); - -// Use a Trie like structure for searching multiple strings -// at once but convert it to a ternary tree for saving space. -// We insert separators in the same order they are specified. -// Template parameter is a CharT which can be a char/wchar_t -// or anything else that supports operator ><,== as long as -// this is not a variable length sequence. We convert utf8 to utf16 -// before inserting. -// Value is a supplementary information useful for search hit -// and is present in the nodes that terminate the whole search pattern -template -class TernarySearchTree { - private: - struct Node { - std::unique_ptr left_; - std::unique_ptr mid_; - std::unique_ptr right_; - CharT c_; // character - Value value_; - bool has_val_; - - explicit Node(CharT c) : c_(c), value_(), has_val_(false) { - } - ~Node() = default; - }; - - struct GetState { - const CharT* const str_; - const size_t len_; - size_t depth_; - const Value* result_; - }; - - public: - TernarySearchTree() = default; - ~TernarySearchTree() = default; - - /** - * Returns a ptr to an associated value and nullptr on search miss. - * Must use default constructed state. - */ - const Value* get(const CharT* str, size_t len) const { - if (len == 0) { - return nullptr; - } - GetState get_state{str, len, 0, nullptr}; - get(root_.get(), get_state); - return get_state.result_; - } - - /** - * Returns true if successful and false on empty strings - * and duplicates. - */ - bool put(const CharT* str, size_t len, const Value& v) { - if (len < 1) { - assert(false); - return false; - } - Node* new_root = put(root_.get(), str, len, v, 0); - if (new_root != nullptr) { - root_.release(); - root_.reset(new_root); - return true; - } - return false; - } - - private: - void update_state(const Node* node, GetState& state) const { - if (node->has_val_) { - if (state.result_ == nullptr) { - state.result_ = &node->value_; - } else if (node->value_ < *state.result_) { - state.result_ = &node->value_; - } - } - } - void get(const Node* node, GetState& state) const { - if (node == nullptr) { - return; - } - assert(state.depth_ < state.len_); - CharT c = state.str_[state.depth_]; - if (c < node->c_) { - get(node->left_.get(), state); - return; - } else if (c > node->c_) { - get(node->right_.get(), state); - return; - } else if (state.depth_ < (state.len_ - 1)) { - // Check if we have a match at this node - update_state(node, state); - if (node->mid_ != nullptr) { - ++state.depth_; - get(node->mid_.get(), state); - } - return; - } - update_state(node, state); - } - - Node* put(Node* node, const CharT* str, size_t len, const Value& v, size_t depth) { - CharT c = str[depth]; - - std::unique_ptr new_node; - if (node == nullptr) { - new_node.reset(new Node(c)); - } - - Node* new_link = nullptr; - Node* n = (node != nullptr) ? node : new_node.get(); - if (c < n->c_) { - new_link = put(n->left_.get(), str, len, v, depth); - if (new_link != nullptr) { - n->left_.release(); - n->left_.reset(new_link); - } - } else if (c > n->c_) { - new_link = put(n->right_.get(), str, len, v, depth); - if (new_link != nullptr) { - n->right_.release(); - n->right_.reset(new_link); - } - } else if (depth < (len - 1)) { - new_link = put(n->mid_.get(), str, len, v, depth + 1); - if (new_link != nullptr) { - n->mid_.release(); - n->mid_.reset(new_link); - } - } else { - if (!n->has_val_) { - n->value_ = v; - n->has_val_ = true; - new_link = n; - } - } - if (new_link != nullptr) { - new_node.release(); - return n; - } - return nullptr; - } - std::unique_ptr root_; -}; - -// We store the length of the original pattern within the -// Ternary Tree. This allows us to cut out the length of the matching -// separator from the original string. -struct SearchValue { - size_t w_len; - int priority_; - bool operator<(const SearchValue& o) const { - return priority_ < o.priority_; - } -}; - } // namespace tokenizer_details using namespace tokenizer_details; -struct Tokenizer::SearchData { - TernarySearchTree tst_; -}; - Tokenizer::Tokenizer(const OpKernelInfo& info) : OpKernel(info) { int64_t mark = 0; auto status = info.GetAttr("mark", &mark); @@ -248,38 +79,37 @@ Tokenizer::Tokenizer(const OpKernelInfo& info) : OpKernel(info) { status = info.GetAttr("tokenexp", &tokenexp); ORT_ENFORCE(status.IsOK(), "Either one of the separators OR tokenexp attributes required but none is set"); ORT_ENFORCE(!tokenexp.empty(), "Expecting a non-empty tokenexp"); + char_tokenezation_ = (tokenexp == "."); } else { - ORT_ENFORCE(!separators.empty(), "Expect at least one item within separators"); + ORT_ENFORCE(!separators.empty(), "separators must not be empty"); + if (separators.size() == 1 && separators[0].empty()) { + char_tokenezation_ = true; + } } - char_tokenezation_ = (separators.size() == 1 && - separators[0].empty()); - ORT_ENFORCE(!char_tokenezation_ || mincharnum_ < 2, "mincharnum is too big for char level tokenezation"); // Check if we have separators or tokenexp if (!char_tokenezation_) { if (!separators.empty()) { - std::unique_ptr sd(std::make_unique()); - std::wstring_convert> converter(conv_error, wconv_error); - int priority = 0; // earlier search patterns get priority + re2::RE2::Options options; + options.set_longest_match(true); for (const auto& sep : separators) { - ORT_ENFORCE(!sep.empty(), "No empty separators allowed"); - std::wstring wsep = converter.from_bytes(sep); - ORT_ENFORCE(wsep != wconv_error, "Separator strings contains invalid utf8 chars"); - bool result = sd->tst_.put(wsep.c_str(), wsep.length(), {wsep.length(), priority}); - ORT_ENFORCE(result, "duplicate separator detected"); - ++priority; + std::unique_ptr regex(new re2::RE2(sep, options)); + if (!regex->ok()) { + ORT_THROW("Can not digest separators: ", sep, " ", regex->error()); + } + separators_.push_back(std::move(regex)); } - search_data_.swap(sd); } else { // Use tokenexp + assert(!tokenexp.empty()); re2::RE2::Options options; options.set_longest_match(true); std::unique_ptr regex(new re2::RE2(tokenexp, options)); if (!regex->ok()) { - ORT_THROW("Can not digest regex: ", regex->error()); + ORT_THROW("Can not digest tokenexp: ", regex->error()); } regex_.swap(regex); } @@ -362,80 +192,88 @@ Status Tokenizer::CharTokenize(OpKernelContext* ctx, size_t N, size_t C, return Status::OK(); } -Status Tokenizer::SeparatorTokenize(OpKernelContext* ctx, - size_t N, size_t C, - const std::vector& input_dims) const { - struct Match { - int priority_; - size_t offset_; - size_t size_; - // create a conflict for overlapping matches - // thus if they overlap neither is less than the other - // and they are considered equal - bool operator<(const Match& o) const { - return (offset_ + size_) <= o.offset_; - } - }; +Status Tokenizer::SeparatorExpressionTokenizer(OpKernelContext* ctx, + size_t N, size_t C, + const std::vector& input_dims) const { + using namespace re2; + std::vector> rows; + rows.reserve(N * C); + + // We do not constraint the search to match + // on the beginning or end of the string + const RE2::Anchor anchor = RE2::UNANCHORED; - std::wstring_convert> converter(conv_error, wconv_error); // Scan all strings and attempt to find separators in them // collect all the output tokens here size_t max_tokens = 0; - std::vector> tokenized_strings; - tokenized_strings.reserve(N * C); auto X = ctx->Input(0); auto const input_data = X->template Data(); auto curr_input = input_data; auto const last = input_data + N * C; while (curr_input != last) { const auto& s = *curr_input; - std::wstring wstr = converter.from_bytes(s); - if (wstr == wconv_error) { + size_t utf8_chars = 0; // length in utf8 chars + if (!utf8_validate(reinterpret_cast(s.data()), s.size(), + utf8_chars)) { return Status(common::ONNXRUNTIME, common::INVALID_ARGUMENT, - "Invalid utf8 chars in the input: " + s); + "Input string contains invalid utf8 chars: " + s); } - std::set matches; - const wchar_t* ws = wstr.c_str(); - size_t len_remaining = wstr.length(); - size_t offset = 0; - while (len_remaining > 0) { - const auto* val = search_data_->tst_.get(ws, len_remaining); - if (val != nullptr) { - auto p = matches.insert({val->priority_, offset, val->w_len}); - while (!p.second && val->priority_ < p.first->priority_) { - // if overlapping matches of the same pattern(priority), then - // the earlier match naturally wins - matches.erase(p.first); - p = matches.insert({val->priority_, offset, val->w_len}); - } - } - ++ws; - ++offset; - --len_remaining; - } + std::vector row{s}; - // Tokenize - tokenized_strings.emplace_back(); - auto& row_tokens = tokenized_strings.back(); - row_tokens.reserve(matches.size() + 1); - ws = wstr.c_str(); - offset = 0; - for (const auto& m : matches) { - assert(m.offset_ >= offset); - size_t sz = (m.offset_ - offset); - if (sz > 0 && sz >= size_t(mincharnum_)) { - row_tokens.emplace_back(ws, sz); - } - offset = m.offset_ + m.size_; - ws = wstr.c_str() + offset; - } - assert(offset <= wstr.length()); - if (offset < wstr.length()) { - row_tokens.emplace_back(ws, wstr.length() - offset); - } + for (const auto& sep : separators_) { + std::vector tokens; + for (const auto& text : row) { + const auto end_pos = text.length(); + size_t start_pos = 0; + StringPiece submatch; - max_tokens = std::max(max_tokens, row_tokens.size()); + bool match = true; + do { + match = sep->Match(text, start_pos, end_pos, anchor, &submatch, 1); + if (match) { + // Record pos/len + assert(submatch.data() != nullptr); + size_t match_pos = submatch.data() - text.data(); + assert(match_pos >= start_pos); + auto token_len = match_pos - start_pos; + utf8_chars = 0; + bool valid = utf8_len(reinterpret_cast(text.data() + start_pos), + token_len, utf8_chars); + if (!valid) { + return Status(common::ONNXRUNTIME, common::INVALID_ARGUMENT, + "Match contains invalid utf8 chars: " + submatch.as_string()); + } + if (utf8_chars >= size_t(mincharnum_)) { + tokens.emplace_back(text.data() + start_pos, token_len); + } + // Update starting position + // Guard against empty string match + auto match_len = submatch.length(); + if (match_len > 0) { + start_pos = match_pos + match_len; + } else { + size_t bytes = 0; + utf8_bytes(*submatch.data(), bytes); + start_pos = match_pos + bytes; + } + } else { + // record trailing token + auto trailing_len = end_pos - start_pos; + utf8_chars = 0; + utf8_len(reinterpret_cast(text.data() + start_pos), + trailing_len, utf8_chars); + if (utf8_chars >= size_t(mincharnum_)) { + tokens.emplace_back(text.data() + start_pos, trailing_len); + } + } + } while (match); + } // row + // Replace the row with the results of this tokenezation + row.swap(tokens); + } // separators_ + max_tokens = std::max(max_tokens, row.size()); + rows.push_back(std::move(row)); ++curr_input; } @@ -463,7 +301,8 @@ Status Tokenizer::SeparatorTokenize(OpKernelContext* ctx, const size_t max_output_index = N * C * max_tokens; #endif size_t output_index = 0; - for (auto& row : tokenized_strings) { + curr_input = input_data; + for (auto& row : rows) { #ifdef _DEBUG size_t c_idx = output_index; #endif @@ -472,8 +311,8 @@ Status Tokenizer::SeparatorTokenize(OpKernelContext* ctx, ++output_index; } // Output tokens for this row - for (auto& token : row) { - *(output_data + output_index) = converter.to_bytes(token); + for (const auto& token : row) { + (output_data + output_index)->assign(token.data(), token.size()); ++output_index; } if (mark_) { @@ -489,13 +328,14 @@ Status Tokenizer::SeparatorTokenize(OpKernelContext* ctx, assert(output_index <= max_output_index); assert((output_index - c_idx) <= max_tokens); #endif + ++curr_input; } return Status::OK(); } -Status Tokenizer::ExpressionTokenize(OpKernelContext* ctx, - size_t N, size_t C, - const std::vector& input_dims) const { +Status Tokenizer::TokenExpression(OpKernelContext* ctx, + size_t N, size_t C, + const std::vector& input_dims) const { using namespace re2; // Represents a token that will be output after // first is the index, second is the size; @@ -514,6 +354,14 @@ Status Tokenizer::ExpressionTokenize(OpKernelContext* ctx, while (curr_input != last) { const auto& s = *curr_input; + + size_t utf8_chars = 0; + if (!utf8_validate(reinterpret_cast(s.data()), s.size(), + utf8_chars)) { + return Status(common::ONNXRUNTIME, common::INVALID_ARGUMENT, + "Input string contains invalid utf8 chars: " + s); + } + tokens.emplace_back(); auto& row = tokens.back(); @@ -533,11 +381,18 @@ Status Tokenizer::ExpressionTokenize(OpKernelContext* ctx, // Guard against empty match and make // sure we make progress either way auto token_len = submatch.length(); - if (token_len > 0) { + utf8_chars = 0; + if (!utf8_len(reinterpret_cast(submatch.data()), token_len, utf8_chars)) { + return Status(common::ONNXRUNTIME, common::INVALID_ARGUMENT, + "Match contains invalid utf8 chars: " + submatch.as_string()); + } + if (utf8_chars >= size_t(mincharnum_)) { row.push_back(submatch); start_pos = match_pos + token_len; } else { - start_pos = match_pos + 1; + size_t bytes = 0; + utf8_bytes(*submatch.data(), bytes); + start_pos = match_pos + bytes; } } } while (match); @@ -645,10 +500,11 @@ Status Tokenizer::Compute(OpKernelContext* ctx) const { if (char_tokenezation_) { s = CharTokenize(ctx, N, C, input_dims); } else { - if (regex_ != nullptr) { - s = ExpressionTokenize(ctx, N, C, input_dims); + if (!separators_.empty()) { + s = SeparatorExpressionTokenizer(ctx, N, C, input_dims); } else { - s = SeparatorTokenize(ctx, N, C, input_dims); + assert(regex_ != nullptr); + s = TokenExpression(ctx, N, C, input_dims); } } return s; diff --git a/onnxruntime/core/common/utf8_util.h b/onnxruntime/core/common/utf8_util.h index bbc10bb094..218309f719 100644 --- a/onnxruntime/core/common/utf8_util.h +++ b/onnxruntime/core/common/utf8_util.h @@ -31,6 +31,23 @@ inline bool utf8_bytes(unsigned char ch, size_t& len) { return false; } +// Computes length of the utf8 string in characters +inline bool utf8_len(const unsigned char* s, size_t bytes, size_t& len) { + size_t result = 0; + while (bytes > 0) { + size_t char_bytes = 0; + bool valid = utf8_bytes(*s, char_bytes); + if (!valid || bytes < char_bytes) { + return false; + } + bytes -= char_bytes; + s += char_bytes; + ++result; + } + len = result; + return true; +} + inline bool utf8_validate(const unsigned char* s, size_t len, size_t& utf8_chars) { size_t utf8_len = 0; size_t idx = 0; diff --git a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc index ea0f6c00e5..dd5835d949 100644 --- a/onnxruntime/core/graph/contrib_ops/contrib_defs.cc +++ b/onnxruntime/core/graph/contrib_ops/contrib_defs.cc @@ -736,7 +736,8 @@ activation and leaky_relu_alpha.)DOC") If the maximum number of tokens found per input string is D, the output shape would be [N, C, D] when input shape is [N, C]. Similarly, if input shape is [C] then the output should be [C, D]. Tokenizer has two different operation modes. The first mode is selected when "tokenexp" is not set and "separators" is set. If "tokenexp" is set and "separators" is not set, - the second mode will be used. The first mode breaks each input string into tokens by removing separators. + the second mode will be used. The first mode breaks each input string into tokens by matching and removing separators. + "separators" is a list of strings which are regular expressions. "tokenexp" is a single regular expression. Let's assume "separators" is [" "] and consider an example. If input is @@ -751,6 +752,9 @@ activation and leaky_relu_alpha.)DOC") whose shape is [2, 5] because you can find at most 5 tokens per input string. Note that the input at most can have two axes, so 3-D and higher dimension are not supported. + If "separators" contains a single empty string, the Tokenizer will enter into character tokenezation mode. This means all strings + will be broken part into individual characters. + For each input string, the second mode searches matches of "tokenexp" and each match will be a token in Y. The matching of "tokenexp" is conducted greedily (i.e., a match should be as long as possible). This operator searches for the first match starting from the beginning of the considered string, @@ -798,14 +802,11 @@ of [N, 0] then [N, 0]. OPTIONAL) .Attr( "separators", - "an optional list of strings (type: AttributeProto::STRINGS), each single string in this attribute is a separator." + "an optional list of strings attribute that contains a list of separators - regular expressions to match separators" " Two consecutive segments in X connected by a separator would be divided into two tokens." " For example, if the input is \"Hello World!\" and this attribute contains only one space character," " the corresponding output would be [\"Hello\", \"World!\"]. To achieve character-level tokenization," - " one should set the separators to [\"\"], which contains only one empty string." - " If 'separators' is a L-element array, there will be L rounds of tokenization using one stop word." - " More specifically, in the first round, the first element in 'separators' is used to tokenize each string in the input." - " Then, the second element in 'separators' will be used to tokenize the resulted strings produced at the first round.", + " one should set the 'separators' to [\"\"], which contains an empty string.", AttributeProto::STRINGS, OPTIONAL) .Attr( diff --git a/onnxruntime/test/contrib_ops/tokenizer_test.cc b/onnxruntime/test/contrib_ops/tokenizer_test.cc index a5741785e7..cdd382a952 100644 --- a/onnxruntime/test/contrib_ops/tokenizer_test.cc +++ b/onnxruntime/test/contrib_ops/tokenizer_test.cc @@ -1,7 +1,6 @@ // Copyright (c) Microsoft Corporation. All rights reserved. // Licensed under the MIT License. -#include #include "gtest/gtest.h" #include "test/providers/provider_test_utils.h" @@ -20,12 +19,13 @@ const int opset_ver = 1; using namespace tokenizer_test; -void InitTestAttr(OpTester& test, bool mark, const std::vector& seps, +void InitTestAttr(OpTester& test, bool mark, const std::vector& sepexp, int64_t mincharnum, const std::string& tokenexp = std::string()) { test.AddAttribute("mark", int64_t{mark}); - if (!seps.empty()) { - test.AddAttribute("separators", seps); + if (!sepexp.empty()) { + test.AddAttribute("separators", sepexp); } + if (!tokenexp.empty()) { test.AddAttribute("tokenexp", tokenexp); } @@ -325,12 +325,10 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersC) { // [C] dimensions // Output [C][D] { - std::vector separators = { - u8"у", - u8"ñ"}; + std::string sepexp = u8"(у|ñ)"; OpTester test("Tokenizer", opset_ver, domain); - InitTestAttr(test, true, separators, 1); + InitTestAttr(test, true, {sepexp}, 1); std::vector dims{2}; std::vector input{u8"Абсу中文", u8"Коñó"}; @@ -357,12 +355,10 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersCompleteMatchEmpt // Test entire separators match so we get nothing // in the output { - std::vector separators = { - u8"Абсу中文", - u8"Коñó"}; + std::string sepexp = u8"(Абсу中文)|(Коñó)"; OpTester test("Tokenizer", opset_ver, domain); - InitTestAttr(test, true, separators, 1); + InitTestAttr(test, true, {sepexp}, 1); std::vector dims{2}; std::vector input{u8"Абсу中文", u8"Коñó"}; @@ -382,12 +378,10 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersCompleteMatchEmpt TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersStartMatchC) { // Match the start { - std::vector separators = { - u8"А", - u8"К"}; + std::string sepexp = u8"(А)|(К)"; OpTester test("Tokenizer", opset_ver, domain); - InitTestAttr(test, true, separators, 1); + InitTestAttr(test, true, {sepexp}, 1); std::vector dims{2}; std::vector input{u8"Абсу中文", u8"Коñó"}; @@ -413,12 +407,10 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersStartMatchC) { TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersEndMatchC) { // Match the end { - std::vector separators = { - u8"文", - u8"ó"}; + std::string sepexp = u8"(文)|(ó)"; OpTester test("Tokenizer", opset_ver, domain); - InitTestAttr(test, true, separators, 1); + InitTestAttr(test, true, {sepexp}, 1); std::vector dims{2}; std::vector input{u8"Абсу中文", u8"Коñó"}; @@ -444,12 +436,10 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersEndMatchC) { TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersEndMatchAtLeast4CharsC) { // Match the end, require at least 4 chars { - std::vector separators = { - u8"文", - u8"ó"}; + std::string sepexp = u8"(文)|(ó)"; OpTester test("Tokenizer", opset_ver, domain); - InitTestAttr(test, true, separators, 4); + InitTestAttr(test, true, {sepexp}, 4); std::vector dims{2}; std::vector input{u8"Абсу中文", u8"Коñó"}; @@ -476,12 +466,10 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersEndMatchAtLeast4C TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersEmptyInputEmptyOutputC) { // Empty input for [C] should produce [C][0] { - std::vector separators = { - u8"文", - u8"ó"}; + std::string sepexp = u8"(文)|(ó)"; OpTester test("Tokenizer", opset_ver, domain); - InitTestAttr(test, true, separators, 4); + InitTestAttr(test, true, {sepexp}, 4); std::vector dims{2}; std::vector input{u8"", u8""}; @@ -495,17 +483,15 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersEmptyInputEmptyOu test.Run(OpTester::ExpectResult::kExpectSuccess); } -} +} // namespace test TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersEmptyInputEmptyOutputNC) { // Empty input for [N][C] should produce [N][C][0] { - std::vector separators = { - u8"文", - u8"ó"}; + std::string sepexp = u8"(文)|(ó)"; OpTester test("Tokenizer", opset_ver, domain); - InitTestAttr(test, true, separators, 4); + InitTestAttr(test, true, {sepexp}, 4); std::vector dims{2, 2}; std::vector input{u8"", u8"文", u8"ó", u8""}; @@ -528,9 +514,7 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharsNoMarkersSeparatorsOverlapSh { // In this case the first pattern must match first // and there would be no match for the second - std::vector separators = { - u8"су", - u8"Абсу"}; + std::vector separators = {u8"су", u8"Абсу"}; OpTester test("Tokenizer", opset_ver, domain); InitTestAttr(test, false, separators, 1); @@ -691,7 +675,7 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharCommonPrefixC) { test.AddOutput("Y", output_dims, output); test.Run(OpTester::ExpectResult::kExpectSuccess); -} +} // namespace test TEST(ContribOpTest, TokenizerExpression_RegEx) { OpTester test("Tokenizer", opset_ver, domain);