Team Ai
Datasetpublic

Brunobkr/llama.cpp_AlgMor24_github

ΩFFFΣLLIa • llama.cpp • AlgMor24 ██████╗ ███████╗███████╗███████╗██╗ ██╗ ██╗ █████╗ ██╔═══██╗██╔════╝██╔════╝██╔════╝██║ ██║ ██║██╔══██╗ ██║ ██║█████╗ █████╗ █████╗ ██║ ██║ ██║███████║ ██║ ██║██╔══╝ ██╔══╝ ██╔══╝ ██║ ██║ ██║██╔══██║ ╚██████╔╝██║ ██║ ███████╗███████╗███████╗██║██║ ██║ ╚═════╝ ╚═╝ ╚═╝ ╚══════╝╚══════╝╚══════╝╚═╝╚═╝ ╚═╝ High-Performance LLM / VLM Inference & Autonomous Agentic Ecosystem… See the full description on the dataset page: https://huggingface.co/datasets/Brunobkr/llama.cpp_AlgMor24_github.

sourceHugging Faceupdated 2mo agoView on Hugging Face
0likes3kdownloads
test-save-load-state.cpp441 linesDownload Raw Back to tests
1#include "arg.h"2#include "common.h"3#include "log.h"4#include "llama-cpp.h"5 6#include <clocale>7#include <random>8#include <vector>9 10struct llama_batch_ptr {11    llama_batch batch;12 13    llama_batch_ptr(int32_t n_tokens, int32_t embd, int32_t n_seq_max)14        : batch{llama_batch_init(n_tokens, embd, n_seq_max)} {}15 16    ~llama_batch_ptr() { llama_batch_free(batch); }17 18    llama_batch_ptr(const llama_batch_ptr &) = delete;19    llama_batch_ptr & operator=(const llama_batch_ptr &) = delete;20    llama_batch_ptr(llama_batch_ptr &&) = default;21    llama_batch_ptr & operator=(llama_batch_ptr &&) = default;22 23    llama_batch & get() { return batch; }24    const llama_batch & get() const { return batch; }25};26 27static llama_tokens generate_tokens(llama_context * ctx, llama_sampler * smpl, int & n_past, int32_t n_predict, llama_seq_id seq_id) {28    llama_tokens result;29    llama_batch_ptr batch(1, 0, 1);30 31    for (int i = 0; i < n_predict; i++) {32        auto next_token = llama_sampler_sample(smpl, ctx, -1);33 34        LOG("%d ", next_token);35        result.push_back(next_token);36 37        common_batch_clear(batch.get());38        common_batch_add(batch.get(), next_token, n_past, {seq_id}, true);39 40        if (llama_decode(ctx, batch.get())) {41            LOG_ERR("\n%s: failed to evaluate\n", __func__);42            return {};43        }44        n_past++;45    }46 47    return result;48}49 50// Test 1: baseline51// - decode all but the last token52// - save state to disk53// - decode the last token54// - generate n_predict tokens55static llama_tokens test_baseline(struct llama_model * model, const struct common_params & params, const llama_tokens & tokens) {56    auto ctx = llama_context_ptr{llama_init_from_model(model, common_context_params_to_llama(params))};57 58    auto sparams = llama_sampler_chain_default_params();59    auto smpl = llama_sampler_ptr{llama_sampler_chain_init(sparams)};60    llama_sampler_chain_add(smpl.get(), llama_sampler_init_dist(params.sampling.seed));61 62    auto n_past = 0;63    if (!common_prompt_batch_decode(ctx.get(), tokens, (int)tokens.size(), n_past, params.n_batch, params.out_file, true)) {64        LOG_ERR("%s: failed to decode prompt\n", __func__);65        return {};66    }67 68    LOG("\n=== Test 1: baseline ===\n");69 70    auto result = generate_tokens(ctx.get(), smpl.get(), n_past, params.n_predict, 0);71    if (result.empty()) {72        return {};73    }74 75    LOG("\n");76 77    return result;78}79 80 81// Test 2: sequence removal isolation82// - decode the same prefix into two sequences83// - remove sequence 084// - verify that sequence 1 remains unchanged85static bool test_seq_rm_isolated(86        struct llama_model         * model,87        const struct common_params & params,88        const llama_tokens         & tokens) {89    auto params_ctx = common_context_params_to_llama(params);90    params_ctx.n_ctx      = 256;91    params_ctx.n_seq_max  = 2;92    params_ctx.kv_unified = true;93 94    auto ctx = llama_context_ptr{llama_init_from_model(model, params_ctx)};95    if (!ctx) {96        LOG_ERR("%s: failed to create context\n", __func__);97        return false;98    }99 100    LOG("\n=== Test 2: sequence removal isolation ===\n");101 102    const size_t n_tokens = tokens.size() < 128 ? tokens.size() : 128;103    for (llama_seq_id seq_id = 0; seq_id < 2; ++seq_id) {104        llama_batch_ptr batch(n_tokens, 0, 1);105        for (size_t i = 0; i < n_tokens; ++i) {106            common_batch_add(batch.get(), tokens[i], i, { seq_id }, false);107        }108 109        if (llama_decode(ctx.get(), batch.get())) {110            LOG_ERR("%s: failed to decode prompt for sequence %d\n", __func__, seq_id);111            return false;112        }113    }114 115    const auto get_seq_state = [&](llama_seq_id seq_id, std::vector<uint8_t> & state) {116        const size_t state_size = llama_state_seq_get_size(ctx.get(), seq_id);117        if (state_size == 0) {118            LOG_ERR("%s: sequence state is empty\n", __func__);119            return false;120        }121 122        state.resize(state_size);123        const size_t ncopy = llama_state_seq_get_data(ctx.get(), state.data(), state.size(), seq_id);124        if (ncopy != state.size()) {125            LOG_ERR("%s: sequence state length %zu does not match expected length %zu\n",126                    __func__, ncopy, state.size());127            return false;128        }129 130        return true;131    };132 133    std::vector<uint8_t> state_before;134    if (!get_seq_state(1, state_before)) {135        return false;136    }137 138    if (!llama_memory_seq_rm(llama_get_memory(ctx.get()), 0, -1, -1)) {139        LOG_ERR("%s: failed to remove sequence 0\n", __func__);140        return false;141    }142 143    std::vector<uint8_t> state_after;144    if (!get_seq_state(1, state_after)) {145        return false;146    }147 148    if (state_before != state_after) {149        LOG_ERR("%s: removing sequence 0 changed sequence 1\n", __func__);150        return false;151    }152 153    LOG("PASS\n");154    return true;155}156 157 158// Test 3: state load159// - create a new context160// - load state from file161// - replay the last prompt token162// - generate n_predict tokens and compare against expected result163static bool test_state_load(struct llama_model * model, const struct common_params & params, const llama_tokens & tokens, const llama_tokens & expected_result) {164    auto ctx = llama_context_ptr{llama_init_from_model(model, common_context_params_to_llama(params))};165 166    auto sparams = llama_sampler_chain_default_params();167    auto smpl = llama_sampler_ptr{llama_sampler_chain_init(sparams)};168    llama_sampler_chain_add(smpl.get(), llama_sampler_init_dist(params.sampling.seed));169 170    LOG("\n=== Test 3: state load ===\n");171 172    // Load state from file173    llama_tokens unused_sts(tokens.size());174    size_t n_token_count_out = 0;175 176    if (!llama_state_load_file(ctx.get(), params.out_file.data(), unused_sts.data(), unused_sts.size(), &n_token_count_out)) {177        LOG_ERR("\n%s: failed to load state\n", __func__);178        return false;179    }180 181    LOG_TRC("%s: loaded state with %zu tokens\n", __func__, n_token_count_out);182 183    // Replay last token184    int n_past = (int) n_token_count_out - 1;185    if (!common_replay_last_token(ctx.get(), tokens.back(), n_past)) {186        return false;187    }188    n_past++;189 190    // Generate tokens191    auto result = generate_tokens(ctx.get(), smpl.get(), n_past, params.n_predict, 0);192    if (result.empty()) {193        return false;194    }195 196    if (result != expected_result) {197        LOG_ERR("\n%s: error: generation differs from expected\n", __func__);198        return false;199    }200 201    LOG("\nPASS\n");202    return true;203}204 205 206// Test 4: seq copy (host)207// - create a multi-seq context208// - load state from file209// - replay the last prompt token210// - migrate KV cache from seq 0 to seq 1 via the CPU path211// - generate n_predict tokens on seq 1 and compare against expected result212static bool test_seq_cp_host(struct llama_model * model, const struct common_params & params, const llama_tokens & tokens, const llama_tokens & expected_result) {213    auto params_ctx = common_context_params_to_llama(params);214    params_ctx.n_seq_max = 2;215    auto ctx = llama_context_ptr{llama_init_from_model(model, params_ctx)};216 217    auto sparams = llama_sampler_chain_default_params();218    auto smpl = llama_sampler_ptr{llama_sampler_chain_init(sparams)};219    llama_sampler_chain_add(smpl.get(), llama_sampler_init_dist(params.sampling.seed));220 221    LOG("\n=== Test 4: seq copy (host) ===\n");222 223    // Load state from file224    llama_tokens unused_sts(tokens.size());225    size_t n_token_count_out = 0;226 227    if (!llama_state_load_file(ctx.get(), params.out_file.data(), unused_sts.data(), unused_sts.size(), &n_token_count_out)) {228        LOG_ERR("\n%s: failed to load state\n", __func__);229        return false;230    }231 232    LOG_TRC("%s: loaded state with %zu tokens\n", __func__, n_token_count_out);233 234    // Replay last token235    int n_past = (int) n_token_count_out - 1;236    if (!common_replay_last_token(ctx.get(), tokens.back(), n_past)) {237        return false;238    }239    n_past++;240 241    // Migrate KV cache from seq 0 to seq 1 (CPU path)242    {243        std::vector<uint8_t> seq_store(llama_state_seq_get_size(ctx.get(), 0));244        const size_t ncopy = llama_state_seq_get_data(ctx.get(), seq_store.data(), seq_store.size(), 0);245        if (ncopy != seq_store.size()) {246            LOG_ERR("\n%s: seq copy data length %zd does not match expected length %zd\n", __func__, ncopy, seq_store.size());247            return false;248        }249        LOG_TRC("%s: seq 0 copied, %zd bytes\n", __func__, ncopy);250 251        llama_memory_clear(llama_get_memory(ctx.get()), true);252        LOG_TRC("%s: kv cache cleared\n", __func__);253 254        const size_t nset = llama_state_seq_set_data(ctx.get(), seq_store.data(), seq_store.size(), 1);255        if (nset != seq_store.size()) {256            LOG_ERR("\n%s: seq set data length %zd does not match expected length %zd\n", __func__, nset, seq_store.size());257            return false;258        }259        LOG_TRC("%s: seq 1 restored, %zd bytes\n", __func__, nset);260    }261 262    // Generate tokens on seq 1263    auto result = generate_tokens(ctx.get(), smpl.get(), n_past, params.n_predict, 1);264    if (result.empty()) {265        return false;266    }267 268    if (result != expected_result) {269        LOG_ERR("\n%s: error: generation differs from expected\n", __func__);270        return false;271    }272 273    LOG("\nPASS\n");274    return true;275}276 277 278// Test 5: seq copy (device)279// - create a multi-seq context280// - load state from file281// - replay the last prompt token282// - migrate KV cache from seq 0 to seq 1 via the on-device path283// - generate n_predict tokens on seq 1 and compare against expected result284static bool test_seq_cp_device(struct llama_model * model, const struct common_params & params, const llama_tokens & tokens, const llama_tokens & expected_result) {285    auto params_ctx = common_context_params_to_llama(params);286    params_ctx.n_seq_max = 2;287    auto ctx = llama_context_ptr{llama_init_from_model(model, params_ctx)};288 289    auto sparams = llama_sampler_chain_default_params();290    auto smpl = llama_sampler_ptr{llama_sampler_chain_init(sparams)};291    llama_sampler_chain_add(smpl.get(), llama_sampler_init_dist(params.sampling.seed));292 293    LOG("\n=== Test 5: seq copy (device) ===\n");294 295    // Load state from file296    llama_tokens unused_sts(tokens.size());297    size_t n_token_count_out = 0;298 299    if (!llama_state_load_file(ctx.get(), params.out_file.data(), unused_sts.data(), unused_sts.size(), &n_token_count_out)) {300        LOG_ERR("\n%s: failed to load state\n", __func__);301        return false;302    }303 304    LOG_TRC("%s: loaded state with %zu tokens\n", __func__, n_token_count_out);305 306    // Replay last token307    int n_past = (int) n_token_count_out - 1;308    if (!common_replay_last_token(ctx.get(), tokens.back(), n_past)) {309        return false;310    }311    n_past++;312 313    // Migrate KV cache from seq 0 to seq 1 (on-device path)314    {315        std::vector<uint8_t> seq_store(llama_state_seq_get_size_ext(ctx.get(), 0, LLAMA_STATE_SEQ_FLAGS_ON_DEVICE));316        const size_t ncopy = llama_state_seq_get_data_ext(ctx.get(), seq_store.data(), seq_store.size(), 0, LLAMA_STATE_SEQ_FLAGS_ON_DEVICE);317        if (ncopy != seq_store.size()) {318            LOG_ERR("\n%s: seq copy data length %zd does not match expected length %zd\n", __func__, ncopy, seq_store.size());319            return false;320        }321        LOG_TRC("%s: seq 0 copied, %zd bytes\n", __func__, ncopy);322 323        llama_memory_clear(llama_get_memory(ctx.get()), true);324        LOG_TRC("%s: kv cache cleared\n", __func__);325 326        const size_t nset = llama_state_seq_set_data_ext(ctx.get(), seq_store.data(), seq_store.size(), 1, LLAMA_STATE_SEQ_FLAGS_ON_DEVICE);327        if (nset != seq_store.size()) {328            LOG_ERR("\n%s: seq set data length %zd does not match expected length %zd\n", __func__, nset, seq_store.size());329            return false;330        }331        LOG_TRC("%s: seq 1 restored, %zd bytes\n", __func__, nset);332    }333 334    // Generate tokens on seq 1335    auto result = generate_tokens(ctx.get(), smpl.get(), n_past, params.n_predict, 1);336    if (result.empty()) {337        return false;338    }339 340    if (result != expected_result) {341        LOG_ERR("\n%s: error: generation differs from expected\n", __func__);342        return false;343    }344 345    LOG("\nPASS\n");346    return true;347}348 349 350int main(int argc, char ** argv) {351    std::setlocale(LC_NUMERIC, "C");352 353    common_params params;354    params.prompt = "";355    params.n_batch = 100;356    params.out_file = "dump_state.bin";357    params.sampling.seed = 1234;358 359    common_init();360 361    if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_COMMON)) {362        return 1;363    }364 365    if (params.n_parallel == 1) {366        LOG_TRC("%s: n_parallel == 1, enabling unified kv cache\n", __func__);367        params.kv_unified = true;368    }369 370    if (params.n_predict < 0) {371        params.n_predict = 16;372    }373 374    ggml_backend_load_all();375 376    auto llama_init = common_init_from_params(params, true);377    auto * model = llama_init->model();378 379    if (model == nullptr) {380        LOG_ERR("%s: failed to init\n", __func__);381        return 1;382    }383 384    GGML_ASSERT(llama_init->context() == nullptr);385 386    // Tokenize prompt or generate random tokens387    llama_tokens tokens;388    if (params.prompt.empty()) {389        const int n_prompt = params.n_batch;390 391        // this path is useful for model files that do not have a tokenizer392        LOG_INF("%s: no prompt provided, generating %d (n_batch) random tokens\n", __func__, n_prompt);393 394        const auto * vocab = llama_model_get_vocab(model);395        const auto n_vocab = llama_vocab_n_tokens(vocab);396 397        std::mt19937 rng(params.sampling.seed);398        std::uniform_int_distribution<llama_token> dist(0, n_vocab - 1);399        for (int i = 0; i < n_prompt; i++) {400            tokens.push_back(dist(rng));401        }402    } else {403        LOG_INF("%s: tokenizing prompt '%s'\n", __func__, params.prompt.c_str());404 405        auto ctx = llama_context_ptr{llama_init_from_model(model, common_context_params_to_llama(params))};406        tokens = common_tokenize(ctx.get(), params.prompt, true);407    }408 409    LOG_INF("%s: the input prompt is %d tokens\n", __func__, (int)tokens.size());410 411    // Test 1: baseline (saves state to disk)412    auto result_baseline = test_baseline(model, params, tokens);413    if (result_baseline.empty()) {414        return 1;415    }416 417    // Test 2: sequence removal isolation418    if (!test_seq_rm_isolated(model, params, tokens)) {419        return 1;420    }421 422    // Test 3: state load423    if (!test_state_load(model, params, tokens, result_baseline)) {424        return 1;425    }426 427    // Test 4: seq copy (host)428    if (!test_seq_cp_host(model, params, tokens, result_baseline)) {429        return 1;430    }431 432    // Test 5: seq copy (device)433    if (!test_seq_cp_device(model, params, tokens, result_baseline)) {434        return 1;435    }436 437    LOG("\nAll tests passed.\n");438 439    return 0;440}441 
Brunobkr/llama.cpp_AlgMor24_github · Team Ai