fix: use tokenizer-specific pre-tokenization rules (#1975)

This commit is contained in:
leejet 2026-09-15 02:37:52 +08:00 committed by GitHub
parent 07a85c74cb
commit 59c23bce0d
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
9 changed files with 92 additions and 1021 deletions

View File

@ -2,6 +2,7 @@
#include <algorithm> #include <algorithm>
#include <sstream> #include <sstream>
#include <stdexcept>
#include "core/util.h" #include "core/util.h"
#include "tokenize_util.h" #include "tokenize_util.h"
@ -31,8 +32,37 @@ std::vector<std::pair<int, std::u32string>> BPETokenizer::bytes_to_unicode() {
return byte_unicode_pairs; return byte_unicode_pairs;
} }
std::vector<std::string> BPETokenizer::token_split(const std::string& text) const { BPETokenizer::BPETokenizer(const std::string& pattern) {
return ::token_split(text); if (!pattern.empty()) {
split_regex_ = std::make_unique<sd::Regex>();
std::string error;
if (!split_regex_->compile(pattern, &error)) {
throw std::runtime_error("invalid tokenizer regex: " + error);
}
}
}
bool BPETokenizer::token_split(const std::string& text, std::vector<std::string>& tokens, std::string* error) const {
tokens.clear();
if (error) {
error->clear();
}
if (!split_regex_) {
if (!text.empty()) {
tokens.push_back(text);
}
return true;
}
std::vector<sd::Regex::Match> matches;
if (!split_regex_->find_matches(text, matches, error)) {
return false;
}
for (const auto& match : matches) {
if (match.first != match.second) {
tokens.push_back(text.substr(match.first, match.second - match.first));
}
}
return true;
} }
std::vector<std::u32string> BPETokenizer::split_utf32(const std::string& text, char32_t delimiter) { std::vector<std::u32string> BPETokenizer::split_utf32(const std::string& text, char32_t delimiter) {
@ -154,7 +184,10 @@ bool BPETokenizer::encode(const std::string& text, std::vector<int>& result, on_
token_strs.push_back(splited_text); token_strs.push_back(splited_text);
continue; continue;
} }
auto tokens = token_split(splited_text); std::vector<std::string> tokens;
if (!token_split(splited_text, tokens, error)) {
return false;
}
for (auto& token : tokens) { for (auto& token : tokens) {
if (on_new_token_cb != nullptr) { if (on_new_token_cb != nullptr) {
bool skip = on_new_token_cb(token, bpe_tokens); bool skip = on_new_token_cb(token, bpe_tokens);

View File

@ -5,15 +5,19 @@
#include <cstdint> #include <cstdint>
#include <functional> #include <functional>
#include <map> #include <map>
#include <regex> #include <memory>
#include <set> #include <set>
#include <string> #include <string>
#include <utility> #include <utility>
#include <vector> #include <vector>
#include "core/regex.h"
#include "tokenizer.h" #include "tokenizer.h"
class BPETokenizer : public Tokenizer { class BPETokenizer : public Tokenizer {
private:
std::unique_ptr<sd::Regex> split_regex_;
protected: protected:
std::map<int, std::u32string> byte_encoder; std::map<int, std::u32string> byte_encoder;
std::map<std::u32string, int> byte_decoder; std::map<std::u32string, int> byte_decoder;
@ -28,12 +32,12 @@ protected:
protected: protected:
static std::vector<std::pair<int, std::u32string>> bytes_to_unicode(); static std::vector<std::pair<int, std::u32string>> bytes_to_unicode();
static std::vector<std::u32string> split_utf32(const std::string& text, char32_t delimiter = U'\n'); static std::vector<std::u32string> split_utf32(const std::string& text, char32_t delimiter = U'\n');
virtual std::vector<std::string> token_split(const std::string& text) const; bool token_split(const std::string& text, std::vector<std::string>& tokens, std::string* error = nullptr) const;
std::vector<std::u32string> bpe(const std::u32string& token) const; std::vector<std::u32string> bpe(const std::u32string& token) const;
std::string decode_token(int token_id) const override; std::string decode_token(int token_id) const override;
public: public:
BPETokenizer() = default; explicit BPETokenizer(const std::string& pattern);
virtual ~BPETokenizer() = default; virtual ~BPETokenizer() = default;
bool encode(const std::string& text, std::vector<int>& tokens, on_new_token_cb_t on_new_token_cb = nullptr, std::string* error = nullptr) override; bool encode(const std::string& text, std::vector<int>& tokens, on_new_token_cb_t on_new_token_cb = nullptr, std::string* error = nullptr) override;

View File

@ -8,10 +8,10 @@
#include "core/util.h" #include "core/util.h"
#include "ggml.h" #include "ggml.h"
#include "tokenize_util.h"
#include "vocab/vocab.h" #include "vocab/vocab.h"
CLIPTokenizer::CLIPTokenizer(int pad_token_id, const std::string& merges_utf8_str) { CLIPTokenizer::CLIPTokenizer(int pad_token_id, const std::string& merges_utf8_str)
: BPETokenizer(R"((?i:'s|'t|'re|'ve|'m|'ll|'d)|\p{L}+|\p{N}|[^\s\p{L}\p{N}]+)") {
UNK_TOKEN = "<|endoftext|>"; UNK_TOKEN = "<|endoftext|>";
BOS_TOKEN = "<|startoftext|>"; BOS_TOKEN = "<|startoftext|>";
EOS_TOKEN = "<|endoftext|>"; EOS_TOKEN = "<|endoftext|>";
@ -101,17 +101,3 @@ std::string CLIPTokenizer::normalize(const std::string& text) const {
std::transform(normalized_text.begin(), normalized_text.end(), normalized_text.begin(), [](unsigned char c) { return static_cast<char>(std::tolower(c)); }); std::transform(normalized_text.begin(), normalized_text.end(), normalized_text.begin(), [](unsigned char c) { return static_cast<char>(std::tolower(c)); });
return normalized_text; return normalized_text;
} }
std::vector<std::string> CLIPTokenizer::token_split(const std::string& text) const {
std::regex clip_pat(R"('s|'t|'re|'ve|'m|'ll|'d|[[:alpha:]]+|[[:digit:]]|[^[:space:][:alpha:][:digit:]]+)",
std::regex::icase);
std::sregex_iterator iter(text.begin(), text.end(), clip_pat);
std::sregex_iterator end;
std::vector<std::string> result;
for (; iter != end; ++iter) {
result.emplace_back(iter->str());
}
return result;
}

View File

@ -11,7 +11,6 @@ class CLIPTokenizer : public BPETokenizer {
protected: protected:
void load_from_merges(const std::string& merges_utf8_str); void load_from_merges(const std::string& merges_utf8_str);
std::string normalize(const std::string& text) const override; std::string normalize(const std::string& text) const override;
std::vector<std::string> token_split(const std::string& text) const override;
public: public:
explicit CLIPTokenizer(int pad_token_id = 49407, const std::string& merges_utf8_str = ""); explicit CLIPTokenizer(int pad_token_id = 49407, const std::string& merges_utf8_str = "");

View File

@ -46,7 +46,9 @@ void GemmaTokenizer::load_from_merges(const std::string& merges_utf8_str, const
bpe_len = rank; bpe_len = rank;
} }
GemmaTokenizer::GemmaTokenizer(const std::string& merges_utf8_str, const std::string& vocab_utf8_str) { GemmaTokenizer::GemmaTokenizer(const std::string& merges_utf8_str, const std::string& vocab_utf8_str)
: BPETokenizer("") {
// Gemma replaces spaces with metaspace before its literal-space Split, so no regex boundaries apply.
byte_level_bpe = false; byte_level_bpe = false;
byte_fallback = true; byte_fallback = true;
add_bos_token = true; add_bos_token = true;

View File

@ -42,7 +42,8 @@ void MistralTokenizer::load_from_merges(const std::string& merges_utf8_str, cons
bpe_len = rank; bpe_len = rank;
} }
MistralTokenizer::MistralTokenizer(const std::string& merges_utf8_str, const std::string& vocab_utf8_str) { MistralTokenizer::MistralTokenizer(const std::string& merges_utf8_str, const std::string& vocab_utf8_str)
: BPETokenizer(R"([^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}]*[\p{Ll}\p{Lm}\p{Lo}\p{M}]+|[^\r\n\p{L}\p{N}]?[\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}]+[\p{Ll}\p{Lm}\p{Lo}\p{M}]*|\p{N}| ?[^\s\p{L}\p{N}]+[\r\n/]*|\s*[\r\n]+|\s+(?!\S)|\s+)") {
add_bos_token = true; add_bos_token = true;
UNK_TOKEN = "<unk>"; UNK_TOKEN = "<unk>";

View File

@ -87,7 +87,8 @@ Qwen2Tokenizer::Qwen2Tokenizer(const std::string& merges_utf8_str)
} }
Qwen2Tokenizer::Qwen2Tokenizer(const std::string& merges_utf8_str, Qwen2Tokenizer::Qwen2Tokenizer(const std::string& merges_utf8_str,
const std::vector<std::string>& special_tokens_override) { const std::vector<std::string>& special_tokens_override)
: BPETokenizer(R"((?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+)") {
UNK_TOKEN = "<|endoftext|>"; UNK_TOKEN = "<|endoftext|>";
EOS_TOKEN = "<|endoftext|>"; EOS_TOKEN = "<|endoftext|>";
PAD_TOKEN = "<|endoftext|>"; PAD_TOKEN = "<|endoftext|>";

File diff suppressed because it is too large Load Diff

View File

@ -4,7 +4,6 @@
#include <string> #include <string>
#include <vector> #include <vector>
std::vector<std::string> token_split(const std::string& text);
std::vector<std::string> split_with_special_tokens(const std::string& text, const std::vector<std::string>& special_tokens); std::vector<std::string> split_with_special_tokens(const std::string& text, const std::vector<std::string>& special_tokens);
#endif // __SD_TOKENIZERS_BPE_TOKENIZE_UTIL_H__ #endif // __SD_TOKENIZERS_BPE_TOKENIZE_UTIL_H__