#include "ling3/tokenizer.h" #include #include #include #include #include #include #include #include #if LING3_WITH_ICU #include #include #include #include #endif namespace ling3 { namespace { constexpr std::array kMagic = {'L', '3', 'T', 'O', 'K', '2', '\0', '\0'}; std::uint64_t PairKey(std::uint32_t left, std::uint32_t right) noexcept { return (static_cast(left) << 32U) | right; } class Reader { public: explicit Reader(std::span data) : data_(data) {} template T Scalar() { if (offset_ > data_.size() || sizeof(T) > data_.size() - offset_) { throw std::runtime_error("tokenizer asset is truncated"); } T value; std::memcpy(&value, data_.data() + offset_, sizeof(value)); offset_ += sizeof(value); return value; } std::string_view Bytes(std::size_t count) { if (offset_ > data_.size() || count > data_.size() - offset_) { throw std::runtime_error("tokenizer asset is truncated"); } const auto * begin = reinterpret_cast(data_.data() + offset_); offset_ += count; return {begin, count}; } std::size_t remaining() const noexcept { return data_.size() - offset_; } private: std::span data_; std::size_t offset_ = 0; }; struct Merge { std::uint32_t rank = 0; std::uint32_t result = 0; }; struct AddedToken { std::uint32_t id = 0; std::string_view content; bool special = false; }; } // namespace struct Tokenizer::Impl { std::array byte_ids {}; std::vector pieces; std::unordered_map merges; std::vector added; std::unordered_set special_ids; #if LING3_WITH_ICU const icu::Normalizer2 * normalizer = nullptr; std::unique_ptr split_pattern; #endif explicit Impl(std::span data) { Reader reader(data); const auto magic = reader.Bytes(kMagic.size()); if (!std::equal(kMagic.begin(), kMagic.end(), magic.begin(), magic.end())) { throw std::runtime_error("tokenizer asset has invalid magic"); } const auto vocab_count = reader.Scalar(); const auto merge_count = reader.Scalar(); const auto added_count = reader.Scalar(); const auto flags = reader.Scalar(); if (vocab_count != 157184 || merge_count == 0 || added_count == 0 || flags != 0) { throw std::runtime_error("tokenizer asset header is incompatible with Ling-3.0-tiny"); } for (auto & token : byte_ids) token = reader.Scalar(); pieces.reserve(vocab_count); for (std::uint32_t id = 0; id < vocab_count; ++id) { pieces.push_back(reader.Bytes(reader.Scalar())); } merges.reserve(static_cast(merge_count) * 2); for (std::uint32_t rank = 0; rank < merge_count; ++rank) { const auto left = reader.Scalar(); const auto right = reader.Scalar(); const auto result = reader.Scalar(); if (left >= vocab_count || right >= vocab_count || result >= vocab_count || !merges.emplace(PairKey(left, right), Merge {rank, result}).second) { throw std::runtime_error("tokenizer asset contains an invalid BPE merge"); } } added.reserve(added_count); for (std::uint32_t index = 0; index < added_count; ++index) { const auto id = reader.Scalar(); const auto bytes = reader.Scalar(); const auto token_flags = reader.Scalar(); const auto content = reader.Bytes(bytes); if (id >= vocab_count || content.empty() || (token_flags & ~1U) != 0 || pieces[id] != content) { throw std::runtime_error("tokenizer asset contains an invalid AddedToken"); } added.push_back({id, content, (token_flags & 1U) != 0}); if ((token_flags & 1U) != 0) special_ids.insert(id); } if (reader.remaining() != 0) { throw std::runtime_error("tokenizer asset has trailing bytes"); } std::stable_sort(added.begin(), added.end(), [](const AddedToken & left, const AddedToken & right) { return left.content.size() > right.content.size(); }); #if LING3_WITH_ICU UErrorCode status = U_ZERO_ERROR; normalizer = icu::Normalizer2::getNFCInstance(status); if (U_FAILURE(status) || normalizer == nullptr) { throw std::runtime_error("ICU failed to initialize the NFC normalizer"); } static constexpr auto pattern = R"REGEX('(?i:[sdmt]|ll|ve|re)|[^\r\n\p{L}\p{N}]?+\p{L}+|\p{N}| ?[^\s\p{L}\p{N}]++[\r\n]*|\s*[\r\n]|\s+(?!\S)|\s+)REGEX"; status = U_ZERO_ERROR; split_pattern.reset(icu::RegexPattern::compile( icu::UnicodeString::fromUTF8(pattern), 0, status)); if (U_FAILURE(status) || split_pattern == nullptr) { throw std::runtime_error("ICU failed to compile the official tokenizer regex"); } #endif } void EncodeBpe(std::string_view text, std::vector & output) const { if (text.empty()) return; std::vector symbols; symbols.reserve(text.size()); for (unsigned char byte : text) symbols.push_back(byte_ids[byte]); while (symbols.size() > 1) { std::uint32_t best_rank = std::numeric_limits::max(); std::uint32_t best_result = 0; std::size_t best_index = symbols.size(); for (std::size_t index = 0; index + 1 < symbols.size(); ++index) { const auto found = merges.find(PairKey(symbols[index], symbols[index + 1])); if (found != merges.end() && found->second.rank < best_rank) { best_rank = found->second.rank; best_result = found->second.result; best_index = index; } } if (best_index == symbols.size()) break; symbols[best_index] = best_result; symbols.erase(symbols.begin() + static_cast(best_index + 1)); } output.insert(output.end(), symbols.begin(), symbols.end()); } void EncodeOrdinary(std::string_view text, std::vector & output) const { if (text.empty()) return; #if LING3_WITH_ICU UErrorCode status = U_ZERO_ERROR; const auto source = icu::UnicodeString::fromUTF8( icu::StringPiece(text.data(), static_cast(text.size()))); icu::UnicodeString normalized; normalizer->normalize(source, normalized, status); if (U_FAILURE(status)) throw std::runtime_error("ICU NFC normalization failed"); std::unique_ptr matcher(split_pattern->matcher(normalized, status)); if (U_FAILURE(status) || matcher == nullptr) { throw std::runtime_error("ICU tokenizer matcher creation failed"); } std::int32_t consumed = 0; while (matcher->find(status)) { const std::int32_t begin = matcher->start(status); const std::int32_t end = matcher->end(status); if (U_FAILURE(status) || begin != consumed || end <= begin) { throw std::runtime_error("official tokenizer regex did not partition the input"); } std::string utf8; normalized.tempSubStringBetween(begin, end).toUTF8String(utf8); EncodeBpe(utf8, output); consumed = end; } if (U_FAILURE(status) || consumed != normalized.length()) { throw std::runtime_error("official tokenizer regex left unmatched input"); } #else EncodeBpe(text, output); #endif } std::vector Encode(std::string_view text) const { std::vector output; std::size_t cursor = 0; while (cursor < text.size()) { std::size_t match_offset = text.size(); const AddedToken * match = nullptr; for (const auto & token : added) { const auto found = text.find(token.content, cursor); if (found < match_offset) { match_offset = found; match = &token; } else if (found == match_offset && match != nullptr && token.content.size() > match->content.size()) { match = &token; } } if (match == nullptr) { EncodeOrdinary(text.substr(cursor), output); break; } EncodeOrdinary(text.substr(cursor, match_offset - cursor), output); output.push_back(match->id); cursor = match_offset + match->content.size(); } return output; } }; Tokenizer::Tokenizer(std::span data) : impl_(std::make_unique(data)) {} Tokenizer::~Tokenizer() = default; Tokenizer::Tokenizer(Tokenizer &&) noexcept = default; Tokenizer & Tokenizer::operator=(Tokenizer &&) noexcept = default; std::vector Tokenizer::Encode(std::string_view text) const { return impl_->Encode(text); } std::string Tokenizer::Decode( std::span tokens, bool skip_special) const { std::string output; for (std::uint32_t token : tokens) { if (token >= impl_->pieces.size()) throw std::out_of_range("token ID is out of range"); if (skip_special && impl_->special_ids.contains(token)) continue; output.append(impl_->pieces[token]); } return output; } std::string_view Tokenizer::Piece(std::uint32_t token) const { if (token >= impl_->pieces.size()) throw std::out_of_range("token ID is out of range"); return impl_->pieces[token]; } std::size_t Tokenizer::vocab_size() const noexcept { return impl_->pieces.size(); } bool Tokenizer::uses_official_unicode_rules() const noexcept { #if LING3_WITH_ICU return true; #else return false; #endif } } // namespace ling3