Ling-3.0-tiny-RKNN / src /tokenizer.cpp
Sariel00's picture
Publish Ling-3.0-tiny RKNN engine and model
3fd1a35 verified
Raw History Blame Contribute Delete
10.6 kB
#include "ling3/tokenizer.h"
#include <algorithm>
#include <array>
#include <cstring>
#include <limits>
#include <stdexcept>
#include <unordered_map>
#include <unordered_set>
#include <utility>
#if LING3_WITH_ICU
#include <unicode/normalizer2.h>
#include <unicode/regex.h>
#include <unicode/stringpiece.h>
#include <unicode/unistr.h>
#endif
namespace ling3 {
namespace {
constexpr std::array<char, 8> 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<std::uint64_t>(left) << 32U) | right;
}
class Reader {
public:
explicit Reader(std::span<const std::byte> data) : data_(data) {}
template <typename T>
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<const char *>(data_.data() + offset_);
offset_ += count;
return {begin, count};
}
std::size_t remaining() const noexcept { return data_.size() - offset_; }
private:
std::span<const std::byte> 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<std::uint32_t, 256> byte_ids {};
std::vector<std::string_view> pieces;
std::unordered_map<std::uint64_t, Merge> merges;
std::vector<AddedToken> added;
std::unordered_set<std::uint32_t> special_ids;
#if LING3_WITH_ICU
const icu::Normalizer2 * normalizer = nullptr;
std::unique_ptr<icu::RegexPattern> split_pattern;
#endif
explicit Impl(std::span<const std::byte> 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<std::uint32_t>();
const auto merge_count = reader.Scalar<std::uint32_t>();
const auto added_count = reader.Scalar<std::uint32_t>();
const auto flags = reader.Scalar<std::uint32_t>();
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<std::uint32_t>();
pieces.reserve(vocab_count);
for (std::uint32_t id = 0; id < vocab_count; ++id) {
pieces.push_back(reader.Bytes(reader.Scalar<std::uint32_t>()));
}
merges.reserve(static_cast<std::size_t>(merge_count) * 2);
for (std::uint32_t rank = 0; rank < merge_count; ++rank) {
const auto left = reader.Scalar<std::uint32_t>();
const auto right = reader.Scalar<std::uint32_t>();
const auto result = reader.Scalar<std::uint32_t>();
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<std::uint32_t>();
const auto bytes = reader.Scalar<std::uint32_t>();
const auto token_flags = reader.Scalar<std::uint32_t>();
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<std::uint32_t> & output) const {
if (text.empty()) return;
std::vector<std::uint32_t> 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<std::uint32_t>::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<std::ptrdiff_t>(best_index + 1));
}
output.insert(output.end(), symbols.begin(), symbols.end());
}
void EncodeOrdinary(std::string_view text, std::vector<std::uint32_t> & 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<std::int32_t>(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<icu::RegexMatcher> 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<std::uint32_t> Encode(std::string_view text) const {
std::vector<std::uint32_t> 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<const std::byte> data)
: impl_(std::make_unique<Impl>(data)) {}
Tokenizer::~Tokenizer() = default;
Tokenizer::Tokenizer(Tokenizer &&) noexcept = default;
Tokenizer & Tokenizer::operator=(Tokenizer &&) noexcept = default;
std::vector<std::uint32_t> Tokenizer::Encode(std::string_view text) const {
return impl_->Encode(text);
}
std::string Tokenizer::Decode(
std::span<const std::uint32_t> 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