Add tokenizer longhaul cache
This commit is contained in:
@@ -259,6 +259,8 @@ set_tests_properties(test-thread-safety PROPERTIES FIXTURES_REQUIRED test-downlo
|
||||
|
||||
llama_build_and_test(test-arg-parser.cpp)
|
||||
llama_build_and_test(test-longhaul.cpp)
|
||||
llama_build_and_test(test-tokenizer-longhaul.cpp
|
||||
ARGS ${PROJECT_SOURCE_DIR}/models/ggml-vocab-qwen35.gguf)
|
||||
llama_build_and_test(test-token-cache.cpp)
|
||||
target_include_directories(test-token-cache PRIVATE ${PROJECT_SOURCE_DIR}/src)
|
||||
|
||||
|
||||
@@ -160,6 +160,21 @@ static void test(void) {
|
||||
assert(params.load_mode == LLAMA_LOAD_MODE_LONGHAUL);
|
||||
assert(params.longhaul_cache_bytes == 2ULL * 1024 * 1024 * 1024);
|
||||
|
||||
params.tokenizer_longhaul = false;
|
||||
params.tokenizer_longhaul_cache_bytes = 0;
|
||||
argv = {"binary_name", "--tokenizer-longhaul"};
|
||||
assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON));
|
||||
|
||||
argv = {"binary_name", "--tokenizer-longhaul", "--tokenizer-longhaul-cache", "64"};
|
||||
assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON));
|
||||
assert(params.tokenizer_longhaul);
|
||||
assert(params.tokenizer_longhaul_cache_bytes == 64ULL * 1024 * 1024);
|
||||
|
||||
params.tokenizer_longhaul = false;
|
||||
params.tokenizer_longhaul_cache_bytes = 0;
|
||||
argv = {"binary_name", "--tokenizer-longhaul-cache", "64"};
|
||||
assert(false == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON));
|
||||
|
||||
argv = {"binary_name", "--token-cache-size", "256", "--token-cache-dir", "/tmp/llama-token-cache-test.sqlite3"};
|
||||
assert(true == common_params_parse(argv.size(), list_str_to_char(argv).data(), params, LLAMA_EXAMPLE_COMMON));
|
||||
assert(params.token_cache_size_mib == 256);
|
||||
|
||||
@@ -0,0 +1,144 @@
|
||||
#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;
|
||||
}
|
||||
Reference in New Issue
Block a user