Files
longhaul.cpp/tests/test-tokenizer-longhaul.cpp
T

145 lines
5.2 KiB
C++

#include "llama.h"
#include "../src/llama-token-cache.h"
#include "../src/llama-tokenizer-longhaul.h"
#include <chrono>
#include <cstring>
#include <filesystem>
#include <string>
#include <thread>
#include <vector>
#undef NDEBUG
#include <cassert>
static std::vector<llama_token> tokenize(const llama_vocab * vocab, const std::string & text) {
int32_t count = llama_tokenize(vocab, text.data(), int32_t(text.size()), nullptr, 0, false, true);
assert(count < 0);
std::vector<llama_token> result(size_t(-count));
count = llama_tokenize(vocab, text.data(), int32_t(text.size()), result.data(), result.size(), false, true);
assert(count == int32_t(result.size()));
return result;
}
int main(int argc, char ** argv) {
namespace fs = std::filesystem;
const fs::path root = fs::temp_directory_path() /
("llama-tokenizer-longhaul-test-" + std::to_string(
std::chrono::high_resolution_clock::now().time_since_epoch().count()));
const std::string db = (root / "token-cache.sqlite3").string();
const llama_token_cache_params params = {
db.c_str(),
128 * 1024 * 1024,
};
assert(llama_token_cache_configure(params));
assert(llama_token_cache_clear());
std::vector<llama_tokenizer_longhaul_reverse> reverse;
std::vector<llama_tokenizer_longhaul_merge> merges;
for (int i = 0; i < 200; ++i) {
reverse.push_back({"token-" + std::to_string(i), i});
if (i > 0) {
merges.push_back({
"token-" + std::to_string(i - 1),
"token-" + std::to_string(i),
i - 1,
});
}
}
{
llama_tokenizer_longhaul_index index("source:test", 512);
std::string error;
assert(index.prepare(reverse, merges, "tokenizer-id", error));
assert(error.empty());
assert(index.n_merges() == merges.size());
for (int i = 0; i < 200; ++i) {
assert(index.text_to_token("token-" + std::to_string(i)) == i);
}
assert(index.text_to_token("missing") == LLAMA_TOKEN_NULL);
for (int i = 1; i < 200; ++i) {
assert(index.find_bpe_rank(
"token-" + std::to_string(i - 1),
"token-" + std::to_string(i)) == i - 1);
}
assert(index.find_bpe_rank("missing", "pair") == -1);
assert(index.cache_bytes_used() <= index.cache_capacity());
const auto loaded_merges = index.get_bpe_merges();
assert(loaded_merges.size() == merges.size());
assert(loaded_merges.front() == "token-0 token-1");
assert(loaded_merges.back() == "token-198 token-199");
}
// A warm open must use the completed artifact even without rebuild inputs.
{
llama_tokenizer_longhaul_index index("source:test", 128);
std::string error;
assert(index.prepare({}, {}, "tokenizer-id", error));
assert(index.text_to_token("token-42") == 42);
std::vector<std::thread> workers;
for (int thread = 0; thread < 4; ++thread) {
workers.emplace_back([&index, thread] {
for (int i = thread; i < 200; i += 4) {
assert(index.text_to_token("token-" + std::to_string(i)) == i);
}
});
}
for (auto & worker : workers) worker.join();
assert(index.cache_bytes_used() <= index.cache_capacity());
}
if (argc > 1) {
llama_backend_init();
llama_model_params normal_params = llama_model_default_params();
normal_params.vocab_only = true;
llama_model * normal = llama_model_load_from_file(argv[1], normal_params);
assert(normal != nullptr);
llama_model_kv_override overrides[2] = {};
std::strcpy(overrides[0].key, "general.architecture");
overrides[0].tag = LLAMA_KV_OVERRIDE_TYPE_STR;
std::strcpy(overrides[0].val_str, "qwen35moe");
llama_model_params longhaul_params = llama_model_default_params();
longhaul_params.vocab_only = true;
longhaul_params.kv_overrides = overrides;
longhaul_params.tokenizer_longhaul = true;
longhaul_params.tokenizer_longhaul_cache_bytes = 1024;
llama_model * longhaul = llama_model_load_from_file(argv[1], longhaul_params);
assert(longhaul != nullptr);
const llama_vocab * normal_vocab = llama_model_get_vocab(normal);
const llama_vocab * longhaul_vocab = llama_model_get_vocab(longhaul);
for (const std::string & text : {
"Hello, world!",
" tokenizer longhaul",
"Qwen3.5: 12345\n",
"<|im_end|>",
"Unicode: Καλημέρα 世界 🚚",
}) {
const auto expected = tokenize(normal_vocab, text);
const auto actual = tokenize(longhaul_vocab, text);
assert(actual == expected);
for (const auto token : actual) {
assert(std::strcmp(
llama_vocab_get_text(normal_vocab, token),
llama_vocab_get_text(longhaul_vocab, token)) == 0);
}
}
llama_model_free(longhaul);
llama_model_free(normal);
llama_backend_free();
}
assert(llama_token_cache_clear());
std::error_code ec;
fs::remove_all(root, ec);
return 0;
}