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
0likes3.1kdownloads
test-recurrent-state-rollback.cpp225 linesDownload Raw Back to tests
1#include "arg.h"2#include "common.h"3#include "llama.h"4 5#include <algorithm>6#include <clocale>7#include <cmath>8#include <cstdio>9#include <vector>10 11static llama_context * make_ctx(const common_params & params, llama_model * model) {12    auto cparams = common_context_params_to_llama(params);13    cparams.n_seq_max = 1;14    cparams.n_rs_seq  = 8;15    cparams.n_batch   = std::max(cparams.n_batch,  (uint32_t) (cparams.n_rs_seq + 1));16    cparams.n_ubatch  = std::max(cparams.n_ubatch, (uint32_t) (cparams.n_rs_seq + 1));17    return llama_init_from_model(model, cparams);18}19 20static bool decode_tokens(llama_context * ctx, const std::vector<llama_token> & tokens, uint32_t count) {21    llama_batch batch = llama_batch_init(count, 0, 1);22    for (uint32_t pos = 0; pos < count; ++pos) {23        common_batch_add(batch, tokens[pos], pos, { 0 }, pos + 1 == count);24    }25    const bool ok = llama_decode(ctx, batch) == 0;26    llama_batch_free(batch);27    return ok;28}29 30static bool decode_one(llama_context * ctx, llama_token tok, llama_pos pos) {31    llama_batch batch = llama_batch_init(1, 0, 1);32    common_batch_add(batch, tok, pos, { 0 }, true);33    const bool ok = llama_decode(ctx, batch) == 0;34    llama_batch_free(batch);35    return ok;36}37 38int main(int argc, char ** argv) {39    std::setlocale(LC_NUMERIC, "C");40 41    common_params params;42    params.sampling.seed = 1234;43    params.n_predict = 1;44 45    common_init();46 47    if (!common_params_parse(argc, argv, params, LLAMA_EXAMPLE_COMMON)) {48        return 1;49    }50 51    ggml_backend_load_all();52 53    common_init_result_ptr llama_init = common_init_from_params(params);54    llama_model * model = llama_init->model();55    if (model == nullptr) {56        fprintf(stderr, "%s : failed to init model\n", __func__);57        return 1;58    }59 60    if (!llama_model_is_recurrent(model) && !llama_model_is_hybrid(model)) {61        fprintf(stderr, "%s : skipping for non-recurrent model\n", __func__);62        return 0;63    }64 65    const llama_vocab * vocab   = llama_model_get_vocab(model);66    const int           n_vocab = llama_vocab_n_tokens(vocab);67 68    llama_context * ctx_src = make_ctx(params, model);69    llama_context * ctx_dst = make_ctx(params, model);70    if (ctx_src == nullptr || ctx_dst == nullptr) {71        fprintf(stderr, "%s : failed to init contexts\n", __func__);72        return 1;73    }74 75    if (llama_n_rs_seq(ctx_src) == 0) {76        fprintf(stderr, "%s : skipping because n_rs_seq is disabled\n", __func__);77        llama_free(ctx_src);78        llama_free(ctx_dst);79        return 0;80    }81 82    std::vector<llama_token> tokens;83    if (llama_vocab_type(vocab) == LLAMA_VOCAB_TYPE_NONE) {84        tokens = { 1, 2, 3, 4, 5, 6, 7, 8, 9 };85    } else {86        tokens = common_tokenize(ctx_src, "The quick brown fox jumps over the lazy dog", true);87    }88    const uint32_t n_rs_seq = llama_n_rs_seq(ctx_src);89    constexpr uint32_t n_rollback = 3;90    if (n_rs_seq < n_rollback) {91        fprintf(stderr, "%s : skipping because n_rs_seq is too small\n", __func__);92        llama_free(ctx_src);93        llama_free(ctx_dst);94        return 0;95    }96    if (tokens.empty()) {97        fprintf(stderr, "%s : not enough prompt tokens\n", __func__);98        return 1;99    }100    tokens.resize(n_rs_seq + 1, tokens.back());101 102    const uint32_t  n_tokens     = tokens.size();103    const llama_pos rollback_pos = (llama_pos) n_tokens - n_rollback;104 105    // Decode the full prompt on the source, then roll back three positions.106    // Replaying them crosses DSV4's ratio-4 compressor boundary.107    // Rollback leaves the recurrent memory in a snapshot state (rs_idx != 0).108    if (!decode_tokens(ctx_src, tokens, n_tokens)) {109        fprintf(stderr, "%s : failed to decode prompt\n", __func__);110        return 1;111    }112    if (!llama_memory_seq_rm(llama_get_memory(ctx_src), 0, rollback_pos, -1)) {113        fprintf(stderr, "%s : rollback failed\n", __func__);114        return 1;115    }116 117    // Save the rolled-back state and restore it into a fresh context.118    common_prompt_checkpoint ckpt;119    ckpt.update_tgt(ctx_src, 0, 0);120    ckpt.load_tgt(ctx_dst, 0, 0);121 122    constexpr float eps = 1e-5f;123    std::vector<std::vector<float>> logits_src_replay(n_rollback);124    const auto replay_and_compare = [&](const char * mode) {125        for (uint32_t i = 0; i < n_rollback; ++i) {126            const llama_pos pos = rollback_pos + i;127            if (!decode_one(ctx_src, tokens[pos], pos) ||128                !decode_one(ctx_dst, tokens[pos], pos)) {129                fprintf(stderr, "%s : %s replay failed at position %d\n", __func__, mode, pos);130                return false;131            }132 133            const float * logits_src = llama_get_logits_ith(ctx_src, 0);134            const float * logits_dst = llama_get_logits_ith(ctx_dst, 0);135            if (logits_src == nullptr || logits_dst == nullptr) {136                fprintf(stderr, "%s : missing %s logits at position %d\n", __func__, mode, pos);137                return false;138            }139 140            logits_src_replay[i].assign(logits_src, logits_src + n_vocab);141            for (int token = 0; token < n_vocab; ++token) {142                if (std::fabs(logits_src[token] - logits_dst[token]) > eps) {143                    fprintf(stderr, "%s : %s logits mismatch at position %d, token %d (%g != %g)\n",144                            __func__, mode, pos, token, (double) logits_src[token], (double) logits_dst[token]);145                    return false;146                }147            }148        }149        return true;150    };151    if (!replay_and_compare("full")) {152        return 1;153    }154 155    if (!llama_memory_seq_rm(llama_get_memory(ctx_src), 0, rollback_pos, -1) ||156        !llama_memory_seq_rm(llama_get_memory(ctx_dst), 0, rollback_pos, -1)) {157        fprintf(stderr, "%s : partial rollback failed\n", __func__);158        return 1;159    }160 161    constexpr llama_state_seq_flags partial_flags = LLAMA_STATE_SEQ_FLAGS_PARTIAL_ONLY;162    common_prompt_checkpoint ckpt_partial;163    ckpt_partial.update_tgt(ctx_src, 0, partial_flags);164    ckpt_partial.load_tgt(ctx_dst, 0, partial_flags);165 166    if (!replay_and_compare("partial")) {167        return 1;168    }169 170    // Repeat the load into a context that already has its own rollback state:171    // groups 1..n_rs_seq hold a different prompt's history, and rs_idx[0] is172    // non-zero at load time. The restore must wipe that state and still match.173    llama_context * ctx_dirty = make_ctx(params, model);174    if (ctx_dirty == nullptr) {175        fprintf(stderr, "%s : failed to init dirty ctx\n", __func__);176        return 1;177    }178 179    std::vector<llama_token> noise = tokens;180    for (auto & t : noise) {181        t = (t + 1) % n_vocab;182        if (t < 0) {183            t = 0;184        }185    }186    if (!decode_tokens(ctx_dirty, noise, n_tokens)) {187        fprintf(stderr, "%s : dirty prompt decode failed\n", __func__);188        return 1;189    }190    if (!llama_memory_seq_rm(llama_get_memory(ctx_dirty), 0, rollback_pos, -1)) {191        fprintf(stderr, "%s : dirty rollback failed\n", __func__);192        return 1;193    }194 195    ckpt.load_tgt(ctx_dirty, 0, 0);196 197    for (uint32_t i = 0; i < n_rollback; ++i) {198        const llama_pos pos = rollback_pos + i;199        if (!decode_one(ctx_dirty, tokens[pos], pos)) {200            fprintf(stderr, "%s : dirty replay failed at position %d\n", __func__, pos);201            return 1;202        }203 204        const float * logits_dirty = llama_get_logits_ith(ctx_dirty, 0);205        if (logits_dirty == nullptr) {206            fprintf(stderr, "%s : missing dirty logits at position %d\n", __func__, pos);207            return 1;208        }209 210        for (int token = 0; token < n_vocab; ++token) {211            if (std::fabs(logits_src_replay[i][token] - logits_dirty[token]) > eps) {212                fprintf(stderr, "%s : dirty-ctx logits mismatch at position %d, token %d (%g != %g)\n",213                        __func__, pos, token, (double) logits_src_replay[i][token], (double) logits_dirty[token]);214                return 1;215            }216        }217    }218 219    fprintf(stderr, "%s : recurrent rollback checkpoint restored successfully\n", __func__);220    llama_free(ctx_src);221    llama_free(ctx_dst);222    llama_free(ctx_dirty);223    return 0;224}225 
Brunobkr/llama.cpp_AlgMor24_github · Team Ai