mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-27 20:02:15 +00:00
parent
1f066d4dc4
commit
95ac7a2f35
4 changed files with 164 additions and 306 deletions
|
|
@ -10,9 +10,6 @@
|
|||
#include "core/common/utf8_util.h"
|
||||
#include "re2/re2.h"
|
||||
|
||||
#include <codecvt>
|
||||
#include <locale>
|
||||
|
||||
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<int64_t>& input_dims) const;
|
||||
|
||||
Status SeparatorTokenize(OpKernelContext* context, size_t N, size_t C,
|
||||
const std::vector<int64_t>& input_dims) const;
|
||||
Status SeparatorExpressionTokenizer(OpKernelContext* context, size_t N, size_t C,
|
||||
const std::vector<int64_t>& input_dims) const;
|
||||
|
||||
Status ExpressionTokenize(OpKernelContext* ctx,
|
||||
size_t N, size_t C,
|
||||
const std::vector<int64_t>& input_dims) const;
|
||||
Status TokenExpression(OpKernelContext* ctx,
|
||||
size_t N, size_t C,
|
||||
const std::vector<int64_t>& input_dims) const;
|
||||
|
||||
bool mark_;
|
||||
bool mark_{false};
|
||||
std::string pad_value_;
|
||||
int64_t mincharnum_;
|
||||
bool char_tokenezation_;
|
||||
struct SearchData;
|
||||
std::unique_ptr<SearchData> search_data_;
|
||||
int64_t mincharnum_{0};
|
||||
bool char_tokenezation_{false};
|
||||
std::vector<std::unique_ptr<re2::RE2>> separators_;
|
||||
std::unique_ptr<re2::RE2> 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 CharT, class Value>
|
||||
class TernarySearchTree {
|
||||
private:
|
||||
struct Node {
|
||||
std::unique_ptr<Node> left_;
|
||||
std::unique_ptr<Node> mid_;
|
||||
std::unique_ptr<Node> 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<Node> 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<Node> 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<wchar_t, SearchValue> 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<SearchData> sd(std::make_unique<SearchData>());
|
||||
std::wstring_convert<std::codecvt_utf8<wchar_t>> 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<re2::RE2> 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<re2::RE2> 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<int64_t>& 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<int64_t>& input_dims) const {
|
||||
using namespace re2;
|
||||
std::vector<std::vector<StringPiece>> 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<std::codecvt_utf8<wchar_t>> 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<std::vector<std::wstring>> tokenized_strings;
|
||||
tokenized_strings.reserve(N * C);
|
||||
auto X = ctx->Input<Tensor>(0);
|
||||
auto const input_data = X->template Data<std::string>();
|
||||
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<const unsigned char*>(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<Match> 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<StringPiece> 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<StringPiece> 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<const unsigned char*>(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<const unsigned char*>(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<int64_t>& input_dims) const {
|
||||
Status Tokenizer::TokenExpression(OpKernelContext* ctx,
|
||||
size_t N, size_t C,
|
||||
const std::vector<int64_t>& 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<const unsigned char*>(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<const unsigned char*>(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;
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -1,7 +1,6 @@
|
|||
// Copyright (c) Microsoft Corporation. All rights reserved.
|
||||
// Licensed under the MIT License.
|
||||
|
||||
#include <codecvt>
|
||||
#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<std::string>& seps,
|
||||
void InitTestAttr(OpTester& test, bool mark, const std::vector<std::string>& 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<std::string> 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<int64_t> dims{2};
|
||||
std::vector<std::string> input{u8"Абсу中文", u8"Коñó"};
|
||||
|
|
@ -357,12 +355,10 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersCompleteMatchEmpt
|
|||
// Test entire separators match so we get nothing
|
||||
// in the output
|
||||
{
|
||||
std::vector<std::string> 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<int64_t> dims{2};
|
||||
std::vector<std::string> input{u8"Абсу中文", u8"Коñó"};
|
||||
|
|
@ -382,12 +378,10 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersCompleteMatchEmpt
|
|||
TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersStartMatchC) {
|
||||
// Match the start
|
||||
{
|
||||
std::vector<std::string> 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<int64_t> dims{2};
|
||||
std::vector<std::string> input{u8"Абсу中文", u8"Коñó"};
|
||||
|
|
@ -413,12 +407,10 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersStartMatchC) {
|
|||
TEST(ContribOpTest, TokenizerWithSeparators_MixCharsWithMarkersEndMatchC) {
|
||||
// Match the end
|
||||
{
|
||||
std::vector<std::string> 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<int64_t> dims{2};
|
||||
std::vector<std::string> 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<std::string> 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<int64_t> dims{2};
|
||||
std::vector<std::string> 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<std::string> 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<int64_t> dims{2};
|
||||
std::vector<std::string> 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<std::string> 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<int64_t> dims{2, 2};
|
||||
std::vector<std::string> 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<std::string> separators = {
|
||||
u8"су",
|
||||
u8"Абсу"};
|
||||
std::vector<std::string> separators = {u8"су", u8"Абсу"};
|
||||
|
||||
OpTester test("Tokenizer", opset_ver, domain);
|
||||
InitTestAttr(test, false, separators, 1);
|
||||
|
|
@ -691,7 +675,7 @@ TEST(ContribOpTest, TokenizerWithSeparators_MixCharCommonPrefixC) {
|
|||
|
||||
test.AddOutput<std::string>("Y", output_dims, output);
|
||||
test.Run(OpTester::ExpectResult::kExpectSuccess);
|
||||
}
|
||||
} // namespace test
|
||||
|
||||
TEST(ContribOpTest, TokenizerExpression_RegEx) {
|
||||
OpTester test("Tokenizer", opset_ver, domain);
|
||||
|
|
|
|||
Loading…
Reference in a new issue