145 lines
5.2 KiB
C++
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;
|
|
}
|