diff --git a/src/transcribe-tokenizer.cpp b/src/transcribe-tokenizer.cpp index b46c5647c..fe90cd4bf 100644 --- a/src/transcribe-tokenizer.cpp +++ b/src/transcribe-tokenizer.cpp @@ -507,6 +507,8 @@ transcribe_status Tokenizer::encode(const std::string & text, std::vector pretokenize_granite(const std::string & text) { return out; } +// --------------------------------------------------------------------------- +// Tekken (Mistral / Voxtral) pretokenizer. +// --------------------------------------------------------------------------- + +namespace { + +constexpr size_t k_no_match = static_cast(-1); + +// Pretoken end offsets (in codepoints) for the Tekken regex. The letter +// alternatives can fail after a partial match, so each is a small +// backtracking matcher. +std::vector tekken_split_offsets(const std::vector & cpts, size_t begin, size_t end) { + std::vector out; + + auto get_cpt = [&](size_t p) -> uint32_t { + return (begin <= p && p < end) ? cpts[p] : OOR; + }; + auto get_flags = [&](size_t p) -> CptFlags { + return (begin <= p && p < end) ? flags_from_cpt(cpts[p]) : CptFlags{}; + }; + + // [\p{Lu}\p{Lt}\p{Lm}\p{Lo}\p{M}], with ASCII-only case (see + // pretokenize_tekken in the header). + auto is_upper_class = [&](size_t p) { + const uint32_t c = get_cpt(p); + const CptFlags f = get_flags(p); + return (f.is_letter() && !(c >= 'a' && c <= 'z')) || f.is_accent_mark(); + }; + // [\p{Ll}\p{Lm}\p{Lo}\p{M}] + auto is_lower_class = [&](size_t p) { + const uint32_t c = get_cpt(p); + const CptFlags f = get_flags(p); + return (f.is_letter() && !(c >= 'A' && c <= 'Z')) || f.is_accent_mark(); + }; + // [^\r\n\p{L}\p{N}] (any in-range codepoint, marks included) + auto is_prefix = [&](size_t p) { + const uint32_t c = get_cpt(p); + const CptFlags f = get_flags(p); + return c != OOR && c != '\r' && c != '\n' && !f.is_letter() && !f.is_number(); + }; + // [^\s\p{L}\p{N}] + auto is_symbol = [&](size_t p) { + const CptFlags f = get_flags(p); + return get_cpt(p) != OOR && !f.is_whitespace() && !f.is_letter() && !f.is_number(); + }; + + // UPPER* LOWER+ starting at `s`. If no LOWER follows the upper run, + // backtrack to the last codepoint in the run that is in both classes. + auto match_upper_star_lower_plus = [&](size_t s) -> size_t { + size_t i = s; + while (is_upper_class(i)) { + i++; + } + if (is_lower_class(i)) { + while (is_lower_class(i)) { + i++; + } + return i; + } + for (size_t k = i; k > s; --k) { + if (is_lower_class(k - 1)) { + return k; + } + } + return k_no_match; + }; + // UPPER+ LOWER* starting at `s`. + auto match_upper_plus_lower_star = [&](size_t s) -> size_t { + if (!is_upper_class(s)) { + return k_no_match; + } + size_t i = s; + while (is_upper_class(i)) { + i++; + } + while (is_lower_class(i)) { + i++; + } + return i; + }; + + size_t prev = begin; + auto push_upto = [&](size_t p) { + assert(prev <= p && p <= end); + if (p > prev) { + out.push_back(p); + prev = p; + } + }; + + size_t pos = begin; + while (pos < end) { + const uint32_t cpt = get_cpt(pos); + const CptFlags flags = get_flags(pos); + + // Alternatives 1 and 2, each with an optional (greedy) prefix + // codepoint: tried with the prefix first, then without. + { + const bool has_prefix = is_prefix(pos); + size_t e = has_prefix ? match_upper_star_lower_plus(pos + 1) : k_no_match; + if (e == k_no_match) { + e = match_upper_star_lower_plus(pos); + } + if (e == k_no_match && has_prefix) { + e = match_upper_plus_lower_star(pos + 1); + } + if (e == k_no_match) { + e = match_upper_plus_lower_star(pos); + } + if (e != k_no_match) { + pos = e; + push_upto(pos); + continue; + } + } + + // \p{N} + if (flags.is_number()) { + pos++; + push_upto(pos); + continue; + } + + // ?[^\s\p{L}\p{N}]+ [\r\n/]* + { + const size_t s = (cpt == ' ' && is_symbol(pos + 1)) ? pos + 1 : pos; + if (is_symbol(s)) { + pos = s; + while (is_symbol(pos)) { + pos++; + } + for (uint32_t cn = get_cpt(pos); cn == '\r' || cn == '\n' || cn == '/'; cn = get_cpt(pos)) { + pos++; + } + push_upto(pos); + continue; + } + } + + // Whitespace handling: \s*[\r\n]+ | \s+(?!\S) | \s+ + size_t ws_count = 0; + size_t last_after_nl = 0; + while (get_flags(pos + ws_count).is_whitespace()) { + const uint32_t cw = get_cpt(pos + ws_count); + if (cw == '\r' || cw == '\n') { + last_after_nl = pos + ws_count + 1; + } + ws_count++; + } + if (last_after_nl > 0) { + pos = last_after_nl; + push_upto(pos); + continue; + } + if (ws_count > 1 && get_cpt(pos + ws_count) != OOR) { + pos += ws_count - 1; + push_upto(pos); + continue; + } + if (ws_count > 0) { + pos += ws_count; + push_upto(pos); + continue; + } + + // Fallback: emit one codepoint as its own pretoken. + pos++; + push_upto(pos); + } + + return out; +} + +} // namespace + +std::vector pretokenize_tekken(const std::string & text) { + std::vector out; + if (text.empty()) { + return out; + } + const auto cpts = cpts_from_utf8(text); + const auto ends = tekken_split_offsets(cpts, 0, cpts.size()); + + out.reserve(ends.size()); + size_t prev = 0; + for (size_t e : ends) { + std::string encoded; + encoded.reserve((e - prev) * 2); + for (size_t i = prev; i < e; ++i) { + const std::string u = cpt_to_utf8(cpts[i]); + for (char c : u) { + encoded += byte_to_unicode(static_cast(c)); + } + } + out.emplace_back(std::move(encoded)); + prev = e; + } + return out; +} + // --------------------------------------------------------------------------- // GPT-2 pretokenizer. // --------------------------------------------------------------------------- diff --git a/src/transcribe-unicode.h b/src/transcribe-unicode.h index 82fad06f2..8a9009432 100644 --- a/src/transcribe-unicode.h +++ b/src/transcribe-unicode.h @@ -50,6 +50,8 @@ struct CptFlags { bool is_whitespace() const { return (bits & WHITESPACE) != 0; } + bool is_accent_mark() const { return (bits & ACCENT_MARK) != 0; } + // True if any category bit in MASK_CATEGORIES (the low byte) is // set. Mirrors unicode_cpt_flags::as_uint() & MASK_CATEGORIES != 0 // from llama.cpp. Used by the pretokenizer to distinguish "known @@ -172,4 +174,21 @@ std::vector pretokenize_gpt2_raw_bytes(const std::string & text); // pretokenizer we'd emit 5380 where the reference produces (30, 198). std::vector pretokenize_granite(const std::string & text); +// Mistral Tekken pretokenizer (Voxtral). The tekken.json pattern is: +// +// [^\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+ +// +// Case is ASCII-only: A-Z is upper, a-z is lower, other letters are +// both. So "iPhone" -> "i", "Phone", but a non-ASCII case change inside +// a word does not split (rare; e.g. "нужноМАНА" differs from +// mistral-common). Combining marks join letter runs, so Devanagari / +// Thai words stay whole. +std::vector pretokenize_tekken(const std::string & text); + } // namespace transcribe::unicode diff --git a/tests/CMakeLists.txt b/tests/CMakeLists.txt index db3d64be3..10b5e134b 100644 --- a/tests/CMakeLists.txt +++ b/tests/CMakeLists.txt @@ -186,6 +186,22 @@ transcribe_apply_warnings(transcribe_batch_mask_unit) add_test(NAME transcribe_batch_mask_unit COMMAND transcribe_batch_mask_unit) +# ----------------------------------------------------------------------------- +# Tekken (Mistral / Voxtral) pretokenizer split (pure host, no model) +# ----------------------------------------------------------------------------- + +add_executable(transcribe_tekken_pretok_unit + tekken_pretok_unit.cpp) + +target_link_libraries(transcribe_tekken_pretok_unit PRIVATE transcribe) + +target_include_directories(transcribe_tekken_pretok_unit PRIVATE + ${CMAKE_SOURCE_DIR}/src) + +transcribe_apply_warnings(transcribe_tekken_pretok_unit) + +add_test(NAME transcribe_tekken_pretok_unit COMMAND transcribe_tekken_pretok_unit) + # ----------------------------------------------------------------------------- # Chunked-prefill causal mask geometry (pure host, no model) @@ -1161,6 +1177,22 @@ if(TRANSCRIBE_BUILD_REAL_MODEL_TESTS) COMMAND transcribe_whisper_tokenize_parity) set_tests_properties(transcribe_whisper_tokenize_parity PROPERTIES SKIP_RETURN_CODE 77) + + # Voxtral: token ids match mistral-common's Tekkenizer (Tekken + # pretokenizer + converter-rebuilt merges). Gated by + # TRANSCRIBE_VOXTRAL_GGUF. + add_executable(transcribe_voxtral_tokenize_parity + voxtral_tokenize_parity.cpp) + + target_link_libraries(transcribe_voxtral_tokenize_parity + PRIVATE transcribe) + + transcribe_apply_warnings(transcribe_voxtral_tokenize_parity) + + add_test(NAME transcribe_voxtral_tokenize_parity + COMMAND transcribe_voxtral_tokenize_parity) + set_tests_properties(transcribe_voxtral_tokenize_parity PROPERTIES + SKIP_RETURN_CODE 77) endif() # ----------------------------------------------------------------------------- diff --git a/tests/tekken_pretok_unit.cpp b/tests/tekken_pretok_unit.cpp new file mode 100644 index 000000000..b76a439ae --- /dev/null +++ b/tests/tekken_pretok_unit.cpp @@ -0,0 +1,84 @@ +// tekken_pretok_unit.cpp - pure-host test of the Tekken (Mistral / +// Voxtral) pretokenizer split, no model required. +// +// Expected splits are regex.findall() with the tekken.json pattern, +// except the ASCII-only case cases marked below. + +#include "transcribe-unicode.h" + +#include +#include +#include +#include + +namespace { + +int g_failures = 0; + +std::string byte_encode(const std::string & raw) { + std::string out; + for (char c : raw) { + out += transcribe::unicode::byte_to_unicode(static_cast(c)); + } + return out; +} + +struct Case { + const char * text; + std::vector expected; // raw UTF-8 pretokens +}; + +void check_case(const Case & c) { + const std::vector got = transcribe::unicode::pretokenize_tekken(c.text); + bool ok = got.size() == c.expected.size(); + for (size_t i = 0; ok && i < got.size(); ++i) { + ok = got[i] == byte_encode(c.expected[i]); + } + if (!ok) { + std::fprintf(stderr, "FAIL pretokenize_tekken(\"%s\"): got %zu pieces, expected %zu\n", c.text, got.size(), + c.expected.size()); + ++g_failures; + } +} + +} // namespace + +int main() { + const std::vector cases = { + // Case split: lower->Upper boundary splits, UPPER+lower stays one word. + { u8"iPhone McDonald HTTPServer", { u8"i", u8"Phone", u8" Mc", u8"Donald", u8" HTTPServer" } }, + { u8"ABC aBC AbC", { u8"ABC", u8" a", u8"BC", u8" Ab", u8"C" } }, + // No contraction alternative: the apostrophe prefixes the letters. + { u8"don't I'M", { u8"don", u8"'t", u8" I", u8"'M" } }, + // Symbol runs swallow a trailing [\r\n/]*. + { u8"!\n/x", { u8"!\n/", u8"x" } }, + { u8"a /\r\n/ b", { u8"a", u8" /\r\n/", u8" b" } }, + // ASCII-only case: non-ASCII letters sit in both classes, so the + // regex split before U+00C4 does not happen. + { u8"\u00d6l\u00c4nderung \u01c5emal", { u8"\u00d6l\u00c4nderung", u8" \u01c5emal" } }, + // Lm (U+02B0) sits in both classes. + { u8"a\u02b0B", { u8"a\u02b0", u8"B" } }, + // Combining marks join letter runs, or stand alone after a prefix. + { u8"e\u0301cole \u0301 !\u0301a", { u8"e\u0301cole", u8" \u0301", u8" !\u0301", u8"a" } }, + // Devanagari: vowel signs / virama are \p{M}, so words stay whole. + { u8"\u0928\u092e\u0938\u094d\u0924\u0947 \u0926\u0941\u0928\u093f\u092f\u093e", + { u8"\u0928\u092e\u0938\u094d\u0924\u0947", u8" \u0926\u0941\u0928\u093f\u092f\u093e" } }, + // Single-codepoint \p{N}, including non-ASCII digits / fractions. + { u8"12 \u0663\u00bd", { u8"1", u8"2", u8" ", u8"\u0663", u8"\u00bd" } }, + // Whitespace alternatives. + { u8" x \n\n y ", { u8" ", u8" x", u8" \n\n", u8" ", u8" y", u8" " } }, + // Supplementary-plane letters (binary-search path of the flags table). + { u8"\U00010400\U00010428 \U0001D400", { u8"\U00010400\U00010428", u8" \U0001D400" } }, + { u8"", {} }, + }; + for (const Case & c : cases) { + check_case(c); + } + + if (g_failures > 0) { + std::fprintf(stderr, "tekken_pretok_unit: %d failures\n", g_failures); + return EXIT_FAILURE; + } + std::fprintf(stdout, "tekken_pretok_unit: ok\n"); + return EXIT_SUCCESS; +} diff --git a/tests/voxtral_tokenize_parity.cpp b/tests/voxtral_tokenize_parity.cpp new file mode 100644 index 000000000..d7d150529 --- /dev/null +++ b/tests/voxtral_tokenize_parity.cpp @@ -0,0 +1,106 @@ +// voxtral_tokenize_parity.cpp - real-model gated test that +// transcribe_tokenize() on a Voxtral GGUF produces the same token ids as +// mistral-common's Tekkenizer. +// +// Expected ids: Tekkenizer.from_file(tekken.json).encode(text, bos=False, +// eos=False) for Voxtral-Mini-3B-2507. +// +// Gated by TRANSCRIBE_VOXTRAL_GGUF. Exits 77 (cmake SKIP_RETURN_CODE) +// when unset. + +#include "transcribe.h" + +#include + +#include +#include +#include +#include + +namespace { + +int g_failures = 0; + +bool file_exists(const std::string & path) { + struct stat st{}; + return ::stat(path.c_str(), &st) == 0; +} + +struct Case { + const char * text; + std::vector expected; +}; + +void check_case(const struct transcribe_model * model, const Case & c) { + int32_t tokens[64]; + const int n = transcribe_tokenize(model, c.text, tokens, 64); + if (n < 0 || static_cast(n) != c.expected.size()) { + std::fprintf(stderr, "FAIL tokenize(%s): n=%d, expected size %zu\n", c.text, n, c.expected.size()); + ++g_failures; + return; + } + for (int i = 0; i < n; ++i) { + if (tokens[i] != c.expected[i]) { + std::fprintf(stderr, "FAIL tokenize(%s): tokens[%d]=%d, expected %d\n", c.text, i, tokens[i], + c.expected[i]); + ++g_failures; + return; + } + } +} + +} // namespace + +int main() { + const char * env = std::getenv("TRANSCRIBE_VOXTRAL_GGUF"); + if (env == nullptr || env[0] == '\0') { + std::fprintf(stderr, + "voxtral_tokenize_parity: TRANSCRIBE_VOXTRAL_GGUF " + "not set; skipping.\n"); + return 77; + } + const std::string model_path = env; + if (!file_exists(model_path)) { + std::fprintf(stderr, "voxtral_tokenize_parity: model not found: %s\n", model_path.c_str()); + return 77; + } + + transcribe_model_load_params mp; + transcribe_model_load_params_init(&mp); + mp.backend = TRANSCRIBE_BACKEND_CPU; + struct transcribe_model * model = nullptr; + const transcribe_status st = transcribe_model_load_file(model_path.c_str(), &mp, &model); + if (st != TRANSCRIBE_OK || model == nullptr) { + std::fprintf(stderr, "FAIL load: %s\n", transcribe_status_string(st)); + return EXIT_FAILURE; + } + + const std::vector cases = { + // The language hint build_transcription_prompt encodes. + { "lang:de", { 9909, 1058, 1558 } }, + { "Transcribe this audio.", { 6881, 13089, 1593, 16023, 1046 } }, + { "HTTPServer iPhone McDonald", { 30499, 11473, 1623, 16742, 5303, 31609 } }, + { u8"\u0928\u092e\u0938\u094d\u0924\u0947 \u0926\u0941\u0928\u093f\u092f\u093e", + { 2485, 3525, 22475, 1803, 115801 } }, + { u8"\u0e2a\u0e27\u0e31\u0e2a\u0e14\u0e35\u0e04\u0e23\u0e31\u0e1a", + { 8335, 61695, 8335, 64836, 38686, 21941 } }, + { u8"\u00dcber Publikum na\u00efve \u00d6l\u00c4nderung", { 79904, 99618, 98355, 86231, 86169, 8793, 1851 } }, + { u8"e\u0301cole", { 1101, 1204, 1129, 17242 } }, + { "http://example.com/a/b\n", { 3809, 2345, 16609, 2354, 22139, 15836, 1010 } }, + { "3rd 1234", { 1051, 7989, 1032, 1049, 1050, 1051, 1052 } }, + { " multi spaces ", { 1032, 8549, 1032, 18971, 1256 } }, + { "", {} }, + }; + for (const Case & c : cases) { + check_case(model, c); + } + + transcribe_model_free(model); + + if (g_failures > 0) { + std::fprintf(stderr, "voxtral_tokenize_parity: %d failures\n", g_failures); + return EXIT_FAILURE; + } + std::fprintf(stdout, "voxtral_tokenize_parity: ok\n"); + return EXIT_SUCCESS; +}